diff --git a/src/ucode/agents/claude.py b/src/ucode/agents/claude.py index 5d824195e..d4c4f467a 100644 --- a/src/ucode/agents/claude.py +++ b/src/ucode/agents/claude.py @@ -83,6 +83,7 @@ external_provider_selected, ) from ucode.os_compatibility import subprocess_cross_os +from ucode.smart_routing import pricing from ucode.smart_routing import v2 as smart_routing_v2 from ucode.smart_routing.claude_hooks import ( FIRST_PROMPT_SOCKET_ENV, @@ -2108,6 +2109,17 @@ def launch( } settings_override = _merge_claude_settings(settings_override or {}, {"env": fallback_env}) os.environ.update(fallback_env) + if smart_routing_v2.savings_statusline_enabled() and workspace: + # Show "Smart routing off" in the status row when nothing is routed: a plain launch, or a + # routed one whose setup failed (`launch_claude` otherwise exits without reaching here). + settings_override = settings_override or {} + smart_routing_v2.install_savings_statusline( + settings_override, + CLAUDE_USER_SETTINGS_PATH, + price_cache=pricing.price_cache_path(APP_DIR, workspace), + routing_enabled=False, + baseline_session_start=False, + ) exec_or_spawn(_build_claude_argv(binary, launch_args, settings_override=settings_override)) diff --git a/src/ucode/databricks.py b/src/ucode/databricks.py index 72f9a30d4..13d9d116c 100644 --- a/src/ucode/databricks.py +++ b/src/ucode/databricks.py @@ -2092,6 +2092,38 @@ def fetch_external_model_prices(workspace: str, token: str) -> tuple[list[dict], return models, None +# Per-token rates for Databricks-hosted (`system.ai`) model services from the gateway's +# pay-per-token endpoint-rates API, which the smart-routing savings statusline prices usage with. +# The model names travel as repeated `databricks_hosted_model_services` query parameters. +_ENDPOINT_RATES_API_PATH = "/api/ai-gateway/v2/endpoint-rates:batchGet" +_SYSTEM_AI_MODEL_PREFIX = "system.ai." + + +def fetch_endpoint_rates( + workspace: str, token: str, model_services: list[str] +) -> tuple[list[dict], str | None]: + """Batch-fetch per-token rates for ``system.ai`` model services. + + Returns ``(rates, reason)`` with each rate the raw ``EndpointRate`` entry; models the caller + can't see are absent rather than errors. Names outside ``system.ai`` are dropped because the + API rejects them. ``reason`` is non-None on failure so callers omit the estimate, not fail. + """ + names = sorted({name for name in model_services if name.startswith(_SYSTEM_AI_MODEL_PREFIX)}) + if not names: + return [], "no system.ai model services to price" + query = urlencode([("databricks_hosted_model_services", name) for name in names]) + url = f"https://{workspace_hostname(workspace)}{_ENDPOINT_RATES_API_PATH}?{query}" + payload, reason = _http_get_json(url, token, timeout=30) + if reason is not None: + return [], reason + if not isinstance(payload, dict): + return [], "endpoint-rates returned an unexpected response shape" + rates = payload.get("model_service_rates") + if not isinstance(rates, list): + return [], None + return [rate for rate in rates if isinstance(rate, dict)], None + + # The `update_mask` paths a config PATCH sends. The server rejects paths outside its mutable set, # so this omits `spec_version` (an estore-internal format marker, still sent in the body; naming it # in the mask is the 400 this fixes) and the reserved/deprecated fields (`display_name`, `tracing`, diff --git a/src/ucode/mods.py b/src/ucode/mods.py index 4d2fdc3e8..cb045077e 100644 --- a/src/ucode/mods.py +++ b/src/ucode/mods.py @@ -63,8 +63,9 @@ def write_mod(plugin_dir: Path, mod: ClaudeMod) -> None: # ug's smart-routing UI mod: a register.ts entry that composes the concern -# modules it imports (today the status band). +# modules it imports (the status band, the subagent routing line, and the +# savings estimate the band appends). SMART_ROUTING_UI = ClaudeMod( source="register.ts", - extra=("smart-routing-status.ts", "subagent-routing.ts"), + extra=("smart-routing-status.ts", "subagent-routing.ts", "smart-routing-savings.ts"), ) diff --git a/src/ucode/smart_routing/claude_statusline.py b/src/ucode/smart_routing/claude_statusline.py new file mode 100644 index 000000000..badea7cf6 --- /dev/null +++ b/src/ucode/smart_routing/claude_statusline.py @@ -0,0 +1,769 @@ +"""Claude Code statusline for smart routing: a one-line row with the orchestrator plugin version, +whether routing is on, and an estimate of what routing saved this session. + +``v2`` installs this as the session's ``statusLine`` command, wrapping any statusline the user +already had. The savings estimate assumes every token the session used, main agent and subagents +alike, would otherwise have run on the baseline main-agent model:: + + saved = sum(tokens * price(baseline)) - sum(tokens * price(model that served them)) + +summed over every assistant response in the main and subagent transcripts, per token class. The +baseline is the main model the user chose; under first-prompt routing, where the router also picks +the main model, it is the model the session started on, before routing switched it, until the main +model changes again. The router switches it once, before the first answer, so a later change is the +user's and makes their choice the baseline again. A negative result (routing chose pricier models) +is shown as a cost increase. + +The baseline is fixed per response when it is first read, never recomputed: a subagent response is +priced against the main model in effect at its timestamp (recorded as the main transcript is +read), and a main-agent response is its own baseline, so switching models with ``/model`` never +reprices earlier work. + +Under Claude Code mods there is no statusline; ``v2`` hands the mod a pricer command instead +(``mod_pricer_argv``), and the mod runs it with ``--mod-usage`` on its own token sums. + +Claude Code debounces refreshes and cancels a run that is still going when the next one starts, +so this module stays stdlib-only (ug's CLI imports take about a second) and reads transcripts +incrementally, resuming from offsets kept in a per-session state file. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import re +import shlex +import sys +import tempfile +import time +from collections.abc import Callable +from dataclasses import dataclass, field +from datetime import UTC, datetime +from decimal import ROUND_HALF_UP, Decimal, InvalidOperation +from pathlib import Path +from typing import Any + +# Only ``ucode.constants`` may be imported here besides ``pricing``: anything heavier (``session_env`` +# pulls in ``config_io``, hence tomlkit and Rich) would get a refresh cancelled. +from ucode.constants import SMART_ROUTING_ENV_KEYS, TRUTHY_ENV_VALUES +from ucode.smart_routing.pricing import ( + ModelPrice, + TokenUsage, + model_key, + read_price_cache, + token_cost, +) + +MODULE = "ucode.smart_routing.claude_statusline" +STATE_DIRNAME = "claude-savings" +STATE_RETENTION_SECONDS = 7 * 24 * 60 * 60 +_STATE_VERSION = 2 +# Env var naming the hooks' session-controls file; duplicated from ``session_env.SESSION_ENV_VAR`` +# (a test keeps them equal) because importing that module is too heavy for a refresh. +SESSION_ENV_FILE_ENV_VAR = "UCODE_SESSION_ENV_FILE" +# The mods' pricer command (a JSON argv) and the token sums it prices (a file the mod writes next to +# the session env file); the mod reads both names, so change them together. +PRICER_ENV_VAR = "UCODE_SAVINGS_PRICER" +MOD_USAGE_FILENAME = "mod-usage.json" +_MOD_USAGE_VERSION = 1 +# statusLine options that shape how the row renders rather than what it runs. +_PRESERVED_STATUS_LINE_KEYS = ("padding", "refreshInterval", "hideVimModeIndicator") +_SAFE_SESSION_ID_RE = re.compile(r"[A-Za-z0-9][A-Za-z0-9._-]{0,127}") +_CENT = Decimal("0.01") +# The Claude Code plugin whose skill drives smart routing's subagent delegation; the row shows its +# installed version. ug doesn't install it, so the row reads whatever Claude Code recorded on disk. +_ORCHESTRATOR_PLUGIN_NAME = "model-orchestrator" +_SYNTHETIC_MODEL = "" + +# The baseline model key for a response, given the key of the model that served it, its timestamp +# (epoch seconds, None when the record has none), and the main agent's model timeline as read so +# far (``_FileTotals.models``). +BaselineFor = Callable[[str, float | None, list[list[Any]]], str] + + +def effective_status_line( + settings: dict, *, user_settings_path: Path, project_dir: Path +) -> dict | None: + """The statusLine Claude Code would run without ug's, or None. + + ``settings`` is the per-launch ``--settings`` document (caller settings merged with ug's), which + outranks project-local, project, and user settings. Managed settings outrank ug's statusline as + well, so there is nothing to wrap there: Claude Code shows the managed statusline instead. + """ + from ucode.config_io import read_json_safe + + sources = ( + settings, + read_json_safe(project_dir / ".claude" / "settings.local.json"), + read_json_safe(project_dir / ".claude" / "settings.json"), + read_json_safe(user_settings_path), + ) + for source in sources: + status_line = source.get("statusLine") + if not isinstance(status_line, dict): + continue + command = status_line.get("command") + if isinstance(command, str) and command.strip(): + return status_line + return None + + +def savings_status_line( + original: dict | None, + *, + python: str, + state_dir: Path, + price_cache: Path, + routing_enabled: bool, + baseline_session_start: bool, +) -> dict: + """The ``statusLine`` setting that prints ``original``'s row(s), then the smart-routing row.""" + argv = [*_module_argv(python), "--state-dir", str(state_dir)] + argv += ["--price-cache", str(price_cache)] + if routing_enabled: + argv.append("--routing-enabled") + if baseline_session_start: + argv.append("--baseline-session-start") + savings = shlex.join(argv) + preserved = { + key: original[key] for key in _PRESERVED_STATUS_LINE_KEYS if original and key in original + } + command = savings if original is None else _chain_commands(original["command"], savings) + return {"type": "command", "command": command, **preserved} + + +def mod_pricer_argv(*, python: str, price_cache: Path, baseline_session_start: bool) -> list[str]: + """The command the Claude Code mod runs (plus ``--mod-usage PATH``) to price its token sums.""" + argv = [*_module_argv(python), "--price-cache", str(price_cache)] + if baseline_session_start: + argv.append("--baseline-session-start") + return argv + + +def _module_argv(python: str) -> list[str]: + # -P keeps the launch directory off sys.path, so a project file can't shadow a stdlib module. + return [python, "-P", "-m", MODULE] + + +def _chain_commands(original: str, savings: str) -> str: + """Run ``original`` as Claude Code would have, then ``savings``, both on the same stdin. + + The original runs in a subshell of whichever shell Claude Code uses, so its shell-specific + syntax keeps working and an ``exit`` in it can't skip the savings row. Command substitution + drops its trailing newlines, so the savings row always starts on a line of its own. + """ + return "\n".join( + [ + "ug_statusline_input=$(cat)", + "ug_statusline_base=$(printf '%s' \"$ug_statusline_input\" | (", + original, + ") )", + '[ -n "$ug_statusline_base" ] && printf \'%s\\n\' "$ug_statusline_base"', + f"printf '%s' \"$ug_statusline_input\" | {savings}", + ] + ) + + +def prune_state(state_dir: Path, *, now: float | None = None) -> None: + """Delete session state untouched for a week, including temp files from cancelled runs.""" + current = now if now is not None else time.time() + try: + entries = list(state_dir.iterdir()) + except OSError: + return + for path in entries: + try: + if path.is_file() and current - path.stat().st_mtime > STATE_RETENTION_SECONDS: + path.unlink() + except OSError: + continue + + +def _epoch(raw: object) -> float | None: + """Epoch seconds of a transcript record's ISO-8601 ``timestamp``, or None.""" + if not isinstance(raw, str) or not raw: + return None + try: + moment = datetime.fromisoformat(raw) + except ValueError: + return None + return (moment if moment.tzinfo else moment.replace(tzinfo=UTC)).timestamp() + + +def _response_costs( + prices: dict[str, ModelPrice], served_key: str, baseline_key: str, tokens: TokenUsage +) -> tuple[Decimal, Decimal] | None: + """``(actual, baseline)`` dollars for one response, or None when either model can't be priced.""" + served_price = prices.get(served_key) + baseline_price = prices.get(baseline_key) + if served_price is None or baseline_price is None: + return None + actual = token_cost(served_price, tokens) + baseline = token_cost(baseline_price, tokens) + if actual is None or baseline is None: + return None + return actual, baseline + + +def _claimable( + actual: Decimal, baseline: Decimal, *, rerouted: bool +) -> tuple[Decimal, Decimal] | None: + """``(saved, baseline)``, or None until routing changed some model and there is a cost to compare.""" + if not rerouted or baseline <= 0: + return None + return baseline - actual, baseline + + +def _own_baseline(served_key: str, when: float | None, models: list[list[Any]]) -> str: + """A response is its own baseline, so it is never counted as rerouted.""" + return served_key + + +def _timeline_baseline(fallback: str) -> BaselineFor: + """A subagent response's baseline is the main model in effect when it ran. + + The timeline is the main agent's ``[epoch, model_key]`` changes, in transcript order. A response + before the first recorded change gets the earliest model, one without a timestamp the latest, + and ``fallback`` applies when the main transcript shows no model at all. + """ + + def baseline_for(served_key: str, when: float | None, models: list[list[Any]]) -> str: + if not models: + return fallback + if when is None: + return models[-1][1] + for epoch, key in reversed(models): + if epoch <= when: + return key + return models[0][1] + + return baseline_for + + +def _first_prompt_baseline(start_key: str, after: BaselineFor) -> BaselineFor: + """The session's start model while the main agent is still on its first model, then ``after``. + + The router switches the main model once, before its first answer, so any later change is the + user's: pricing past it against the start model would credit routing with the user's choice. + """ + + def baseline_for(served_key: str, when: float | None, models: list[list[Any]]) -> str: + if len(models) < 2 or (when is not None and when < models[1][0]): + return start_key + return after(served_key, when, models) + + return baseline_for + + +@dataclass +class _FileTotals: + """Running cost totals for one transcript file, resumable from ``offset``.""" + + offset: int = 0 + actual: Decimal = Decimal(0) + baseline: Decimal = Decimal(0) + # Responses served by a model other than the baseline, i.e. where routing changed the model. + rerouted: int = 0 + unpriced: list[str] = field(default_factory=list) + # The last response's contribution, replaced (not re-added) when its id repeats. + last_id: str | None = None + last_actual: Decimal = Decimal(0) + last_baseline: Decimal = Decimal(0) + last_rerouted: int = 0 + # Main transcript only: ``[epoch, model_key]`` each time the main agent's model changed, so a + # later ``/model`` switch doesn't reprice the work done before it. + models: list[list[Any]] = field(default_factory=list) + + def dump(self) -> dict[str, Any]: + return { + "offset": self.offset, + "actual": str(self.actual), + "baseline": str(self.baseline), + "rerouted": self.rerouted, + "unpriced": self.unpriced, + "last_id": self.last_id, + "last_actual": str(self.last_actual), + "last_baseline": str(self.last_baseline), + "last_rerouted": self.last_rerouted, + "models": self.models, + } + + @classmethod + def load(cls, raw: object) -> _FileTotals: + if not isinstance(raw, dict): + return cls() + try: + unpriced = raw.get("unpriced") + last_id = raw.get("last_id") + models = raw.get("models") + return cls( + offset=int(raw.get("offset", 0)), + actual=Decimal(str(raw.get("actual", 0))), + baseline=Decimal(str(raw.get("baseline", 0))), + rerouted=int(raw.get("rerouted", 0)), + unpriced=[str(model) for model in unpriced] if isinstance(unpriced, list) else [], + last_id=last_id if isinstance(last_id, str) else None, + last_actual=Decimal(str(raw.get("last_actual", 0))), + last_baseline=Decimal(str(raw.get("last_baseline", 0))), + last_rerouted=int(raw.get("last_rerouted", 0)), + models=[[float(epoch), str(key)] for epoch, key in models] + if isinstance(models, list) + else [], + ) + except (TypeError, ValueError, InvalidOperation): + return cls() + + def _note_main_model(self, key: str, when: float | None) -> None: + if self.models and self.models[-1][1] == key: + return + # A record with no timestamp sorts with the one before it. + epoch = when if when is not None else (self.models[-1][0] if self.models else 0.0) + self.models.append([epoch, key]) + + def add_response( + self, + record: object, + prices: dict[str, ModelPrice], + baseline_for: BaselineFor, + *, + timeline: list[list[Any]] | None = None, + ) -> None: + """Fold in one transcript record. + + ``timeline`` is the main agent's, for a subagent transcript; without one this is the main + transcript, which records its own in ``models`` as it is read. + """ + if not isinstance(record, dict) or record.get("type") != "assistant": + return + message = record.get("message") + usage = message.get("usage") if isinstance(message, dict) else None + if not isinstance(message, dict) or not isinstance(usage, dict): + return + tokens = TokenUsage.from_message_usage(usage) + if not any(tokens): + return + raw_model = message.get("model") + model = raw_model if isinstance(raw_model, str) else "" + served_key = model_key(model) + when = _epoch(record.get("timestamp")) + if timeline is None and model and model != _SYNTHETIC_MODEL: + self._note_main_model(served_key, when) + baseline_key = baseline_for(served_key, when, self.models if timeline is None else timeline) + costs = _response_costs(prices, served_key, baseline_key, tokens) + if costs is None: + if model not in self.unpriced: + self.unpriced.append(model) + actual = baseline = Decimal(0) + else: + actual, baseline = costs + rerouted = int(served_key != baseline_key) + + # Claude Code writes one record per content block, each repeating the response's usage, + # and one response's records are contiguous in its transcript. + identity = message.get("id") or record.get("uuid") + identity = identity if isinstance(identity, str) and identity else None + if identity is not None and identity == self.last_id: + self.actual -= self.last_actual + self.baseline -= self.last_baseline + self.rerouted -= self.last_rerouted + self.actual += actual + self.baseline += baseline + self.rerouted += rerouted + self.last_id = identity + self.last_actual, self.last_baseline, self.last_rerouted = actual, baseline, rerouted + + def consume( + self, + path: Path, + prices: dict[str, ModelPrice], + baseline_for: BaselineFor, + *, + timeline: list[list[Any]] | None = None, + ) -> _FileTotals: + """Fold in complete lines appended since ``offset``; returns the updated totals.""" + try: + size = path.stat().st_size + except OSError: + return self + totals = self if size >= self.offset else _FileTotals() # rewritten: start over + if size == totals.offset: + return totals + try: + with path.open("rb") as handle: + handle.seek(totals.offset) + data = handle.read(size - totals.offset) + except OSError: + return totals + end = data.rfind(b"\n") + if end < 0: + return totals # only a partial line so far; Claude Code is still writing it + totals.offset += end + 1 + for line in data[: end + 1].splitlines(): + try: + record = json.loads(line) + except ValueError: + continue + totals.add_response(record, prices, baseline_for, timeline=timeline) + return totals + + +def _model_ref(raw: object) -> dict[str, str] | None: + if not isinstance(raw, dict) or not isinstance(raw.get("id"), str) or not raw["id"]: + return None + display_name = raw.get("display_name") + return { + "id": raw["id"], + "display_name": display_name + if isinstance(display_name, str) and display_name + else raw["id"], + } + + +def _state_path(state_dir: Path, session_id: str) -> Path: + if _SAFE_SESSION_ID_RE.fullmatch(session_id): + name = session_id + else: + name = hashlib.sha256(session_id.encode("utf-8")).hexdigest() + return state_dir / f"{name}.json" + + +def _read_state(path: Path) -> dict[str, Any]: + try: + state = json.loads(path.read_text(encoding="utf-8")) + except (OSError, ValueError): + return {"version": _STATE_VERSION} + if not isinstance(state, dict): + return {"version": _STATE_VERSION} + if state.get("version") != _STATE_VERSION: + # Only start_model survives a format change: it can't be re-captured once routing has + # switched the session's model, and the rest is rebuilt from the transcripts. + start = _model_ref(state.get("start_model")) + return {"version": _STATE_VERSION, **({} if start is None else {"start_model": start})} + return state + + +def _write_state(path: Path, state: dict[str, Any]) -> None: + path.parent.mkdir(mode=0o700, parents=True, exist_ok=True) + fd, temporary = tempfile.mkstemp(dir=path.parent, prefix=f".{path.name}.", suffix=".tmp") + try: + with os.fdopen(fd, "w", encoding="utf-8") as handle: + json.dump(state, handle, separators=(",", ":")) + os.replace(temporary, path) + except BaseException: + Path(temporary).unlink(missing_ok=True) + raise + + +def _transcript_paths(transcript: Path) -> list[Path]: + """The main transcript, then its subagents' (``/subagents/agent-*.jsonl``).""" + subagents = transcript.with_suffix("") / "subagents" + return [transcript, *sorted(subagents.glob("agent-*.jsonl"))] + + +def _claude_config_dir() -> Path: + """Claude Code's config directory, honoring ``CLAUDE_CONFIG_DIR`` as Claude Code does.""" + configured = os.environ.get("CLAUDE_CONFIG_DIR", "").strip() + return Path(configured) if configured else Path.home() / ".claude" + + +def orchestrator_plugin_version(config_dir: Path | None = None) -> str | None: + """The installed ``model-orchestrator`` plugin version, or None when it can't be determined. + + ug doesn't install the plugin, so this reads whatever Claude Code recorded in + ``plugins/installed_plugins.json``; a missing file, unexpected shape, or absent plugin yields + None and the row simply omits the version rather than inventing one. + """ + base = config_dir if config_dir is not None else _claude_config_dir() + try: + payload = json.loads((base / "plugins" / "installed_plugins.json").read_text("utf-8")) + except (OSError, ValueError): + return None + plugins = payload.get("plugins") if isinstance(payload, dict) else None + if not isinstance(plugins, dict): + return None + for key, records in plugins.items(): + # Keys are "@"; match on the name so the marketplace can vary. + if not isinstance(key, str) or key.split("@", 1)[0] != _ORCHESTRATOR_PLUGIN_NAME: + continue + for record in records if isinstance(records, list) else []: + version = record.get("version") if isinstance(record, dict) else None + if isinstance(version, str) and version and version != "unknown": + return version + return None + + +def _savings_figures(saved: Decimal, baseline: Decimal) -> tuple[str, Decimal]: + """``saved``'s magnitude as dollars (to the cent) and as a whole percent of ``baseline``.""" + percent = Decimal(0) + if baseline > 0: + percent = (abs(saved) / baseline * 100).quantize(Decimal(1), rounding=ROUND_HALF_UP) + magnitude = abs(saved) + amount = ( + "<$0.01" if magnitude < _CENT else f"${magnitude.quantize(_CENT, rounding=ROUND_HALF_UP):,}" + ) + return amount, percent + + +def _savings_text(saved: Decimal, baseline: Decimal) -> str: + """The savings clause: a money-bag estimate when positive, an honest cost increase when not.""" + amount, percent = _savings_figures(saved, baseline) + if saved >= 0: + return f"๐Ÿ’ฐ Est. saved with smart routing: {amount} ({percent}%)" + return f"Smart routing cost {amount} more ({percent}%)" + + +def _mod_savings_text(saved: Decimal, baseline: Decimal) -> str: + """The mod's savings segment; the mod's band already says it is about smart routing.""" + amount, percent = _savings_figures(saved, baseline) + if saved >= 0: + return f"๐Ÿ’ฐ Est. saved {amount} ({percent}%)" + return f"cost {amount} more ({percent}%)" + + +def _compute_savings( + raw: str, *, state_dir: Path, price_cache: Path, baseline_session_start: bool +) -> tuple[Decimal, Decimal] | None: + """The ``(saved, baseline)`` dollars for one statusline payload, or None with nothing to claim. + + None until some response ran on a model other than its baseline, and whenever any response + can't be priced: an undercounted figure would be worse than none. The caller then shows "on". + """ + try: + payload = json.loads(raw) + except ValueError: + return None + if not isinstance(payload, dict): + return None + session_id = payload.get("session_id") + transcript = payload.get("transcript_path") + current = _model_ref(payload.get("model")) + if not isinstance(session_id, str) or not session_id: + return None + if not isinstance(transcript, str) or not transcript or current is None: + return None + + state_path = _state_path(state_dir, session_id) + state = _read_state(state_path) + start = _model_ref(state.get("start_model")) + # Recorded on the session's first refresh, before first-prompt routing can switch the model. + if start is None: + start = state["start_model"] = current + _write_state(state_path, state) + + cached = read_price_cache(price_cache) + if cached is None: + return None + prices, fingerprint = cached + # First-prompt routing prices work against the model the session started on. Otherwise the + # baseline is whatever main model each response ran beside, so it isn't part of the key. + start_key = model_key(start["id"]) if baseline_session_start else None + + key = [start_key, fingerprint] + files = state.get("files") + if state.get("key") != key or not isinstance(files, dict): + state["key"], files = key, {} + main_path, *subagent_paths = _transcript_paths(Path(transcript)) + actual_total = baseline_total = Decimal(0) + rerouted = 0 + unpriced = False + + def fold( + path: Path, baseline_for: BaselineFor, *, timeline: list[list[Any]] | None + ) -> _FileTotals: + nonlocal actual_total, baseline_total, rerouted, unpriced + totals = _FileTotals.load(files.get(str(path))).consume( + path, prices, baseline_for, timeline=timeline + ) + files[str(path)] = totals.dump() + actual_total += totals.actual + baseline_total += totals.baseline + rerouted += totals.rerouted + unpriced = unpriced or bool(totals.unpriced) + return totals + + main_baseline: BaselineFor = _own_baseline + subagent_baseline = _timeline_baseline(model_key(current["id"])) + if start_key is not None: + main_baseline = _first_prompt_baseline(start_key, main_baseline) + subagent_baseline = _first_prompt_baseline(start_key, subagent_baseline) + # The main transcript is read first so its model timeline covers the subagent responses. + main_totals = fold(main_path, main_baseline, timeline=None) + for path in subagent_paths: + fold(path, subagent_baseline, timeline=main_totals.models) + state["files"] = files + _write_state(state_path, state) + + if unpriced: + return None + return _claimable(actual_total, baseline_total, rerouted=rerouted > 0) + + +def _mod_usage_savings( + usage_path: Path, *, price_cache: Path, baseline_session_start: bool +) -> tuple[Decimal, Decimal] | None: + """The ``(saved, baseline)`` dollars for the mod's token sums, or None with nothing to claim. + + ``usage_path`` holds the mod's tokens pre-aggregated per (agent, baseline, served); a malformed + or baseline-less entry, like an unpriced one, voids the estimate rather than undercount it. + Under first-prompt routing, entries from before the user changed the main model use the + session's start model as their baseline. A file that can't be read as the mod's format at all is + an error, not an empty estimate. + """ + document = json.loads(usage_path.read_text(encoding="utf-8")) + if not isinstance(document, dict) or document.get("version") != _MOD_USAGE_VERSION: + raise ValueError(f"unsupported mod usage file: {usage_path}") + entries = document.get("entries") + if not isinstance(entries, list): + raise ValueError(f"unsupported mod usage file: {usage_path}") + fixed_baseline = None + if baseline_session_start: + start = document.get("start_model") + if not isinstance(start, str) or not start: + return None + fixed_baseline = start + cached = read_price_cache(price_cache) + if cached is None: + return None + prices = cached[0] + + actual_total = baseline_total = Decimal(0) + rerouted = False + for entry in entries: + if not isinstance(entry, dict): + return None + usage, served = entry.get("usage"), entry.get("served") + baseline = entry.get("baseline") + if fixed_baseline is not None and entry.get("before_user_switch") is True: + baseline = fixed_baseline + if not isinstance(usage, dict) or not isinstance(served, str) or not served: + return None + if not isinstance(baseline, str) or not baseline: + return None + # The mod's usage has no cache-write TTL split, and ug launches Claude with 1-hour caching. + tokens = TokenUsage.from_message_usage(usage, uncovered_writes_1h=True) + if not any(tokens): + continue + served_key, baseline_key = model_key(served), model_key(baseline) + costs = _response_costs(prices, served_key, baseline_key, tokens) + if costs is None: + return None + actual_total += costs[0] + baseline_total += costs[1] + rerouted = rerouted or served_key != baseline_key + return _claimable(actual_total, baseline_total, rerouted=rerouted) + + +def mod_usage_output(usage_path: Path, *, price_cache: Path, baseline_session_start: bool) -> str: + """The one JSON line the mod reads: its savings segment and plugin segment, each or null.""" + # Each segment fails alone, and the mod must get a parseable line whatever goes wrong. + savings = plugin = None + try: + computed = _mod_usage_savings( + usage_path, price_cache=price_cache, baseline_session_start=baseline_session_start + ) + savings = None if computed is None else _mod_savings_text(*computed) + except Exception: # noqa: BLE001 + pass + try: + version = orchestrator_plugin_version() + plugin = f"plugin v{version}" if version else None + except Exception: # noqa: BLE001 + pass + return json.dumps({"savings": savings, "plugin": plugin}, ensure_ascii=False) + + +def _routing_on(launch_enabled: bool) -> bool: + """Whether smart routing is on now, following the session's ``/smart-router`` toggle. + + The launch decides whether a session is routed at all: a plain launch never routes, and could + inherit an outer routed session's file. For a routed one, the session file (which the toggle + rewrites but the statusLine command can't be) overrides the launch-time flag when it names a + routing flag; an empty file, as after turning routing back on, leaves the launch's choice. + """ + if not launch_enabled: + return False + path = os.environ.get(SESSION_ENV_FILE_ENV_VAR, "").strip() + if not path: + return True + try: + values = json.loads(Path(path).read_text(encoding="utf-8")) + except (OSError, ValueError): + return True + if not isinstance(values, dict): + return True + flags = [values[name] for name in SMART_ROUTING_ENV_KEYS if name in values] + if not flags: + return True + return any( + isinstance(flag, str) and flag.strip().lower() in TRUTHY_ENV_VALUES for flag in flags + ) + + +def render( + raw: str, + *, + routing_enabled: bool, + state_dir: Path, + price_cache: Path, + baseline_session_start: bool, +) -> str: + """The one-line smart-routing status row. + + ``off`` when routing is disabled (at launch, or since by ``/smart-router``), ``on`` once it is + on but no estimate exists yet, otherwise the estimate. The orchestrator plugin version is + appended when it can be read. + """ + version = orchestrator_plugin_version() + plugin = f" ยท smart router plugin v{version}" if version else "" + if not _routing_on(routing_enabled): + return f"Smart routing off{plugin}" + computed = _compute_savings( + raw, + state_dir=state_dir, + price_cache=price_cache, + baseline_session_start=baseline_session_start, + ) + if computed is None: + return f"Smart routing on{plugin}" + return f"{_savings_text(*computed)}{plugin}" + + +def main(argv: list[str] | None = None) -> int: + parser = argparse.ArgumentParser( + prog=MODULE, + description="Print the smart-routing status row for a Claude Code session, or (with " + "--mod-usage) price the Claude Code mod's token sums as a JSON line.", + ) + source = parser.add_mutually_exclusive_group(required=True) + source.add_argument("--state-dir", type=Path) + source.add_argument("--mod-usage", type=Path) + parser.add_argument("--price-cache", type=Path, required=True) + parser.add_argument("--routing-enabled", action="store_true") + parser.add_argument("--baseline-session-start", action="store_true") + args = parser.parse_args(argv) + try: + if args.mod_usage is not None: + line = mod_usage_output( + args.mod_usage, + price_cache=args.price_cache, + baseline_session_start=args.baseline_session_start, + ) + else: + line = render( + sys.stdin.read(), + routing_enabled=args.routing_enabled, + state_dir=args.state_dir, + price_cache=args.price_cache, + baseline_session_start=args.baseline_session_start, + ) + # Write bytes so the middle-dot survives a non-UTF-8 stdout locale. + sys.stdout.buffer.write((line + "\n").encode("utf-8")) + except Exception: # noqa: BLE001 - a status error must never break the user's status row + return 0 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/src/ucode/smart_routing/pricing.py b/src/ucode/smart_routing/pricing.py new file mode 100644 index 000000000..7176f06d2 --- /dev/null +++ b/src/ucode/smart_routing/pricing.py @@ -0,0 +1,292 @@ +"""Per-token model prices and the cost of coding-agent token usage. + +The launcher fetches rates from the AI Gateway's endpoint-rates API (``prices_from_endpoint_rates``) +and caches them per workspace (``price_cache_path``) for the savings statusline to read; this module +defines that cache and the arithmetic. A token class with no rate makes a response unpriceable, and +the statusline hides the estimate rather than undercount it. The API returns per-million-token +dollar rates for input, output, and cache tokens (5-minute and 1-hour writes, and reads), which the +estimate needs because fixed multipliers don't hold across models (Opus 5.5 cache reads bill at +0.05x input, Opus 4.8's at 0.1x). A model whose org has no DBU-to-dollar conversion stays unpriced. + +Stdlib-only on purpose: the savings statusline imports this on every Claude Code refresh, and the +CLI's usual imports (Rich, Typer, the Databricks SDK, even ``urllib.request``) cost enough startup +for Claude Code to cancel the run. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import re +import tempfile +import time +from collections.abc import Iterable, Mapping +from dataclasses import dataclass, fields +from decimal import Decimal, InvalidOperation +from pathlib import Path +from typing import Any, NamedTuple + +_PRICE_CACHE_VERSION = 1 +_MILLION = Decimal(1_000_000) +_CONTEXT_SUFFIX_RE = re.compile(r"\[[^\]]*\]$") +_ANTHROPIC_AIGW_PREFIX_RE = re.compile(r"^anthropic-aigw-[0-9a-f]{8}-") +# The gateway reports some served models by their provider id, e.g. Bedrock's +# `anthropic.claude-haiku-4-5-20251001-v1:0` for a request to `system.ai.claude-haiku-4-5`. +_PROVIDER_PREFIX_RE = re.compile(r"^(?:(?:us|eu|apac|au|jp|global)\.)?anthropic\.") +_PROVIDER_SUFFIX_RE = re.compile(r"(?:-20\d{6})?(?:-v\d+(?::\d+)?)?$") +# Other served models are reported by a deployment id that drifts from the model service name: +# `glm-5.3-flash` for `system.ai.glm-5-3-flash`, `glm-5-3-colo-on-sp-v1` for `system.ai.glm-5-3`. +_DEPLOYMENT_SUFFIX_RE = re.compile(r"-colo-on-[a-z0-9]+$") + + +def model_key(model: str) -> str: + """Collapse the spellings of one model id that Claude Code and the gateway use. + + Mirrors ``routing.unwrap_anthropic_gateway_model`` + ``routing.normalize_model`` without + importing ``routing``, whose ``urllib.request`` import alone costs ~0.2s of statusline startup. + Also drops a ``[1m]`` context-window selector (a window, not a price), and the provider prefix, + date/version suffix and deployment spelling of a served id, so it keys the same as the + requested model. + """ + name = _CONTEXT_SUFFIX_RE.sub("", (model or "").strip().lower()) + name = _ANTHROPIC_AIGW_PREFIX_RE.sub("", name).rsplit("/", 1)[-1] + for prefix in ("databricks-", "system.ai."): + if name.startswith(prefix): + name = name[len(prefix) :] + break + name = _PROVIDER_PREFIX_RE.sub("", name).split("@", 1)[0] + name = _DEPLOYMENT_SUFFIX_RE.sub("", _PROVIDER_SUFFIX_RE.sub("", name)) + # Model service names spell versions with dashes. + return name.replace(".", "-") + + +@dataclass(frozen=True) +class ModelPrice: + """USD per million tokens by token class; ``None`` where the gateway reports no rate.""" + + input: Decimal | None = None + output: Decimal | None = None + cache_read: Decimal | None = None + cache_write_5m: Decimal | None = None + cache_write_1h: Decimal | None = None + # Rates for requests whose prompt exceeds the threshold (e.g. Sonnet 4.x above 200k tokens). + long_context_threshold: int | None = None + long_context: ModelPrice | None = None + + +_RATE_FIELDS = tuple( + field.name for field in fields(ModelPrice) if not field.name.startswith("long_context") +) + + +class TokenUsage(NamedTuple): + """One model response's billed tokens, split by the classes that are priced separately.""" + + input: int = 0 + cache_write_5m: int = 0 + cache_write_1h: int = 0 + cache_read: int = 0 + output: int = 0 + + @classmethod + def from_message_usage( + cls, usage: Mapping[str, Any], *, uncovered_writes_1h: bool = False + ) -> TokenUsage: + """Read an Anthropic Messages ``usage`` object as Claude Code records it in transcripts. + + Cache writes are split by TTL because a 1-hour write bills at twice the input rate and a + 5-minute write at 1.25x; ug enables 1-hour caching. Writes the breakdown doesn't cover are + billed as 5-minute writes, the API's default TTL. A caller whose usage never carries the + breakdown, but whose session ran with 1-hour caching, sets ``uncovered_writes_1h`` so those + writes aren't undercounted at the 5-minute rate. + """ + creation = usage.get("cache_creation") + creation = creation if isinstance(creation, Mapping) else {} + write_1h = _count(creation.get("ephemeral_1h_input_tokens")) + write_5m = _count(creation.get("ephemeral_5m_input_tokens")) + uncovered = max(_count(usage.get("cache_creation_input_tokens")) - write_1h - write_5m, 0) + if uncovered_writes_1h: + write_1h += uncovered + else: + write_5m += uncovered + return cls( + input=_count(usage.get("input_tokens")), + cache_write_5m=write_5m, + cache_write_1h=write_1h, + cache_read=_count(usage.get("cache_read_input_tokens")), + output=_count(usage.get("output_tokens")), + ) + + @property + def prompt_tokens(self) -> int: + """Every input-side token, which is what long-context price thresholds measure.""" + return self.input + self.cache_write_5m + self.cache_write_1h + self.cache_read + + +def _count(value: object) -> int: + return value if isinstance(value, int) and not isinstance(value, bool) and value > 0 else 0 + + +def token_cost(price: ModelPrice, tokens: TokenUsage) -> Decimal | None: + """USD cost of one response, or None when a token class it used has no rate. + + Returning None rather than pricing a missing rate at zero keeps callers from showing a figure + that silently undercounts. + """ + if ( + price.long_context is not None + and price.long_context_threshold is not None + and tokens.prompt_tokens > price.long_context_threshold + ): + price = price.long_context + total = Decimal(0) + for count, rate in ( + (tokens.input, price.input), + (tokens.cache_write_5m, price.cache_write_5m), + (tokens.cache_write_1h, price.cache_write_1h), + (tokens.cache_read, price.cache_read), + (tokens.output, price.output), + ): + if count <= 0: + continue + if rate is None: + return None + total += Decimal(count) * rate + return total / _MILLION + + +def _rate(raw: object) -> Decimal | None: + if isinstance(raw, bool) or not isinstance(raw, int | float | str) or not str(raw).strip(): + return None + try: + rate = Decimal(str(raw).strip()) + except InvalidOperation: + return None + return rate if rate.is_finite() and rate >= 0 else None + + +def _dump_price(price: ModelPrice) -> dict[str, Any]: + dumped: dict[str, Any] = { + name: str(value) for name in _RATE_FIELDS if (value := getattr(price, name)) is not None + } + if price.long_context is not None and price.long_context_threshold is not None: + dumped["long_context_threshold"] = price.long_context_threshold + dumped["long_context"] = _dump_price(price.long_context) + return dumped + + +def _load_price(raw: object) -> ModelPrice | None: + if not isinstance(raw, Mapping): + return None + rates = {name: _rate(raw.get(name)) for name in _RATE_FIELDS} + threshold = raw.get("long_context_threshold") + long_context = _load_price(raw.get("long_context")) + if isinstance(threshold, bool) or not isinstance(threshold, int) or long_context is None: + threshold, long_context = None, None + return ModelPrice(**rates, long_context_threshold=threshold, long_context=long_context) + + +# The endpoint-rates API groups each model's costs by unit; the statusline prices in dollars. +_DOLLARS_UNIT = "USD" +# `TokenType` enum names โ†’ the ``ModelPrice`` field each bills. CACHE_CREATION is the API's default +# 5-minute cache write; CACHE_CREATION_1H is the 1-hour write ug enables. +_TOKEN_TYPE_FIELDS = { + "TOKEN_TYPE_INPUT": "input", + "TOKEN_TYPE_OUTPUT": "output", + "TOKEN_TYPE_CACHE_CREATION": "cache_write_5m", + "TOKEN_TYPE_CACHE_CREATION_1H": "cache_write_1h", + "TOKEN_TYPE_CACHE_READ": "cache_read", +} + + +def prices_from_endpoint_rates(rates: Iterable[object]) -> dict[str, ModelPrice]: + """Map endpoint-rates ``EndpointRate`` entries to prices keyed by their model service. + + Reads the ``USD`` group of each entry's ``costs`` (a per-million-token ``cost`` per + ``token_type``). The API omits the dollar group when the org has no DBU-to-dollar conversion + configured; those models are left unpriced rather than shown in DBUs. + """ + prices: dict[str, ModelPrice] = {} + for rate in rates: + if not isinstance(rate, Mapping): + continue + service = rate.get("model_service") + costs = rate.get("costs") + if not isinstance(service, str) or not service or not isinstance(costs, list): + continue + dollars = next( + (c for c in costs if isinstance(c, Mapping) and c.get("unit") == _DOLLARS_UNIT), None + ) + token_costs = dollars.get("token_costs") if isinstance(dollars, Mapping) else None + if not isinstance(token_costs, list): + continue + rate_fields: dict[str, Decimal] = {} + for entry in token_costs: + if not isinstance(entry, Mapping): + continue + field = _TOKEN_TYPE_FIELDS.get(entry.get("token_type")) + value = _rate(entry.get("cost")) + if field is not None and value is not None: + rate_fields[field] = value + price = ModelPrice( + input=rate_fields.get("input"), + output=rate_fields.get("output"), + cache_read=rate_fields.get("cache_read"), + cache_write_5m=rate_fields.get("cache_write_5m"), + cache_write_1h=rate_fields.get("cache_write_1h"), + ) + if price.input is not None or price.output is not None: + prices[service] = price + return prices + + +def price_cache_path(app_dir: Path, workspace: str) -> Path: + """The workspace's price cache; dollar rates depend on its org's DBU conversion.""" + host = workspace.strip().lower().removeprefix("https://").rstrip("/") + digest = hashlib.sha256(host.encode("utf-8")).hexdigest()[:16] + return app_dir / f"model-prices-{digest}.json" + + +def write_price_cache( + path: Path, prices: Mapping[str, ModelPrice], *, now: float | None = None +) -> None: + """Atomically replace the price cache so a concurrent statusline never reads a partial file. + + ``prices`` is keyed by any spelling of the model id; entries are stored under ``model_key``. + """ + payload = { + "version": _PRICE_CACHE_VERSION, + "fetched_at": now if now is not None else time.time(), + "models": {model_key(model): _dump_price(price) for model, price in prices.items()}, + } + path.parent.mkdir(parents=True, exist_ok=True) + fd, temporary = tempfile.mkstemp(dir=path.parent, prefix=f".{path.name}.", suffix=".tmp") + try: + with os.fdopen(fd, "w", encoding="utf-8") as handle: + json.dump(payload, handle, separators=(",", ":")) + os.replace(temporary, path) + except BaseException: + Path(temporary).unlink(missing_ok=True) + raise + + +def read_price_cache(path: Path) -> tuple[dict[str, ModelPrice], str] | None: + """Load prices keyed by ``model_key``, plus a fingerprint that changes on every refresh.""" + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, ValueError): + return None + if not isinstance(payload, dict) or payload.get("version") != _PRICE_CACHE_VERSION: + return None + models = payload.get("models") + if not isinstance(models, dict): + return None + prices = { + key: price + for key, raw in models.items() + if isinstance(key, str) and (price := _load_price(raw)) is not None + } + if not prices: + return None + return prices, str(payload.get("fetched_at")) diff --git a/src/ucode/smart_routing/v2.py b/src/ucode/smart_routing/v2.py index 54f46f6e5..d96e66487 100644 --- a/src/ucode/smart_routing/v2.py +++ b/src/ucode/smart_routing/v2.py @@ -7,6 +7,7 @@ import socket import subprocess import sys +import threading import time import urllib.request from collections.abc import Callable, MutableMapping @@ -40,6 +41,7 @@ from ucode.databricks import ( AnthropicModelCatalog, build_auth_token_argv, + fetch_endpoint_rates, get_databricks_token, list_anthropic_model_catalog, list_anthropic_models, @@ -51,7 +53,13 @@ release_file_lock, ) from ucode.skills import SMART_ROUTER_SKILL, install_skill -from ucode.smart_routing import claude_routing, codex_interposer, routing +from ucode.smart_routing import ( + claude_routing, + claude_statusline, + codex_interposer, + pricing, + routing, +) from ucode.smart_routing.claude_hooks import ( FIRST_PROMPT_SOCKET_ENV, sync_first_prompt_hook, @@ -61,6 +69,7 @@ from ucode.smart_routing.session_env import SESSION_ENV_VAR, SESSION_PYTHON_ENV_VAR, start_session from ucode.ui import print_warning +ENABLE_SAVINGS_STATUSLINE_ENV_VAR = "ENABLE_SMART_ROUTING_SAVINGS" LEGACY_STATE_KEY = "smart_routing_enabled" CODEX_INTERPOSER_LOG = APP_DIR / "codex-v2-interposer.log" @@ -172,6 +181,100 @@ def first_prompt_routing_enabled(env: MutableMapping[str, str] | None = None) -> ) +def savings_statusline_enabled(env: MutableMapping[str, str] | None = None) -> bool: + """Whether a Claude launch shows the smart-routing status row (default on; ``0`` opts out).""" + source = os.environ if env is None else env + return source.get(ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "1") != "0" + + +def install_savings_statusline( + settings: dict, + user_settings_path: Path, + *, + price_cache: Path, + routing_enabled: bool, + baseline_session_start: bool, +) -> None: + """Point the per-launch ``statusLine`` at the smart-routing row, wrapping the user's own one. + + ``routing_enabled`` is False when the launch isn't routed (a plain launch, or routing setup + failed) so the row reads "off". The row reads per-token prices from ``price_cache``, since a + statusline refresh can't wait on the network; ``_start_savings_price_refresh`` fills it on + routed launches. + """ + state_dir = APP_DIR / claude_statusline.STATE_DIRNAME + claude_statusline.prune_state(state_dir) + original = claude_statusline.effective_status_line( + settings, user_settings_path=user_settings_path, project_dir=Path.cwd() + ) + settings["statusLine"] = claude_statusline.savings_status_line( + original, + python=sys.executable, + state_dir=state_dir, + price_cache=price_cache, + routing_enabled=routing_enabled, + baseline_session_start=baseline_session_start, + ) + + +# Claude settings env keys that name the main model a session may start on. +_MAIN_MODEL_ENV_KEYS = ( + "ANTHROPIC_MODEL", + "ANTHROPIC_DEFAULT_FABLE_MODEL", + "ANTHROPIC_DEFAULT_OPUS_MODEL", + "ANTHROPIC_DEFAULT_SONNET_MODEL", + "ANTHROPIC_DEFAULT_HAIKU_MODEL", +) + + +def _savings_model_services(models: list[str | None], settings: dict) -> list[str]: + """The ``system.ai`` model services a session's responses can be priced against. + + Covers the routable models plus the main model, which the baseline prices every token at and + which can come from a pinned ``--model`` or Claude's default-model env rather than the catalog. + """ + env = settings.get("env") + env = env if isinstance(env, dict) else {} + candidates = [*models, *(env.get(key) for key in _MAIN_MODEL_ENV_KEYS)] + services: set[str] = set() + for model in candidates: + if not isinstance(model, str) or not model: + continue + name = _canonical_claude_model_id(_unwrapped_claude_model_id(model.strip())) + name = name.removesuffix("[1m]") + if name.startswith("system.ai."): + services.add(name) + return sorted(services) + + +def _refresh_savings_prices( + workspace: str, token: str, model_services: list[str], price_cache: Path +) -> None: + """Fetch this launch's per-token rates and cache them for the savings statusline. + + On failure the previous cache stays; without one the row stays hidden. + """ + try: + rates, _reason = fetch_endpoint_rates(workspace, token, model_services) + prices = pricing.prices_from_endpoint_rates(rates) + if prices: + pricing.write_price_cache(price_cache, prices) + except Exception: # noqa: BLE001 - a background refresh must never surface in the agent's TUI + return + + +def _start_savings_price_refresh( + workspace: str, token: str, model_services: list[str], price_cache: Path +) -> None: + """Refresh prices off the launch path so a slow rates API never delays Claude's startup.""" + threading.Thread( + target=_refresh_savings_prices, + args=(workspace, token, model_services, price_cache), + name="ug-savings-prices", + daemon=True, + ).start() + + def enable_smart_routing( env: MutableMapping[str, str] | None = None, ) -> dict[str, str | None]: @@ -592,6 +695,32 @@ def launch_claude( ) if route_first_prompt: sync_first_prompt_hook(settings, hook_executable) + if savings_statusline_enabled(): + price_cache = pricing.price_cache_path(APP_DIR, workspace) + if mods_enabled: + # The mod draws the savings itself, so there is no statusline to install: it runs this + # command on its own token sums. + env[claude_statusline.PRICER_ENV_VAR] = json.dumps( + claude_statusline.mod_pricer_argv( + python=sys.executable, + price_cache=price_cache, + baseline_session_start=route_first_prompt, + ) + ) + else: + install_savings_statusline( + settings, + user_settings_path, + price_cache=price_cache, + routing_enabled=True, + baseline_session_start=route_first_prompt, + ) + _start_savings_price_refresh( + workspace, + token, + _savings_model_services([*model_ids, launch_model], settings), + price_cache, + ) model_setting = _ClaudeModelSettingGuard(user_settings_path) def route_prompt(prompt: str) -> claude_pty.FirstPromptRoute: diff --git a/tests/integration/README.md b/tests/integration/README.md index 55f5ac9aa..d61a83e55 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -257,7 +257,10 @@ configuration, each journey injects only a tracing-enabled CodingAgentConfig inp the suite's managed-config stub seam; the agents, inference, OTLP export, and table verification remain real. Each adds the prompt's UUID as a trace-safe `ug_integration_marker` attribute, resolves the destination table from the workspace tracing -configuration, waits 30 seconds, and queries that table through an existing SQL warehouse. +configuration, waits 30 seconds, and queries that table through an existing SQL warehouse +(a running one if any, else a serverless one). A query may wait up to 10 minutes for the +warehouse to start or queue it before its 180-second execution limit applies, and a timeout +reports which of the two it hit. The tests assert that a span with the marker arrived and identifies the requested model. Scoped discovery additionally requires Model Services diff --git a/tests/integration/utils/sql.py b/tests/integration/utils/sql.py index 03feb222e..40e879852 100644 --- a/tests/integration/utils/sql.py +++ b/tests/integration/utils/sql.py @@ -8,6 +8,10 @@ DEFAULT_TRACE_TABLE_PREFIX = "unity_gateway" TRACE_TABLE_SUFFIX = "_otel_spans" +# A statement stays PENDING while its warehouse starts or queues it; a stopped classic or pro +# warehouse can take minutes to start, which is warehouse latency, not a missing trace. +WAREHOUSE_WAIT_SECONDS = 600 +STATEMENT_RUN_SECONDS = 180 def resolve_trace_table(workspace: str, bearer: str) -> str: @@ -45,7 +49,7 @@ def _quote_identifier(identifier: str) -> str: def resolve_warehouse_id(workspace: str, bearer: str) -> str: - """Choose an existing warehouse, preferring one that is already running.""" + """Choose an existing warehouse: a running one, else a serverless one (quick to start).""" request = urllib.request.Request( f"{workspace.rstrip('/')}/api/2.0/sql/warehouses", headers={"Authorization": f"Bearer {bearer}"}, @@ -62,7 +66,15 @@ def resolve_warehouse_id(workspace: str, bearer: str) -> str: running = next( (warehouse for warehouse in warehouses if warehouse.get("state") == "RUNNING"), None ) - return str((running or warehouses[0])["id"]) + serverless = next( + ( + warehouse + for warehouse in warehouses + if warehouse.get("enable_serverless_compute") is True + ), + None, + ) + return str((running or serverless or warehouses[0])["id"]) def query_count( @@ -88,9 +100,23 @@ def query_count( with urllib.request.urlopen(request, timeout=60) as response: # noqa: S310 result = json.load(response) - deadline = time.monotonic() + 180 - while result.get("status", {}).get("state") in {"PENDING", "RUNNING"}: - assert time.monotonic() < deadline, "SQL statement did not finish within 180 seconds" + statement_id = result.get("statement_id") + submitted_at = time.monotonic() + running_since: float | None = None + while (state := result.get("status", {}).get("state")) in {"PENDING", "RUNNING"}: + now = time.monotonic() + if state == "RUNNING" and running_since is None: + running_since = now + if running_since is None: + assert now - submitted_at < WAREHOUSE_WAIT_SECONDS, ( + f"SQL statement {statement_id} stayed PENDING on warehouse {warehouse_id} " + f"for {WAREHOUSE_WAIT_SECONDS} seconds" + ) + else: + assert now - running_since < STATEMENT_RUN_SECONDS, ( + f"SQL statement {statement_id} on warehouse {warehouse_id} did not finish " + f"within {STATEMENT_RUN_SECONDS} seconds of running" + ) time.sleep(2) poll = urllib.request.Request( f"{workspace.rstrip('/')}/api/2.0/sql/statements/{result['statement_id']}", diff --git a/tests/test_agent_claude.py b/tests/test_agent_claude.py index 367b9daa3..3797b6eea 100644 --- a/tests/test_agent_claude.py +++ b/tests/test_agent_claude.py @@ -15,7 +15,7 @@ from ucode import constants, managed_files from ucode import databricks as db_mod from ucode.agents import LaunchOptions, claude -from ucode.smart_routing import claude_routing, v2 +from ucode.smart_routing import claude_routing, claude_statusline, v2 from ucode.state import MANAGED_OVERLAY_KEY WS = "https://example.databricks.com" @@ -2461,6 +2461,7 @@ def test_default_launch_keeps_existing_auth_path(self, monkeypatch): monkeypatch.delenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, raising=False) monkeypatch.delenv("ANTHROPIC_DEFAULT_MODEL", raising=False) monkeypatch.delenv("OAUTH_TOKEN", raising=False) + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "0") monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) @@ -2470,6 +2471,63 @@ def test_default_launch_keeps_existing_auth_path(self, monkeypatch): assert "ANTHROPIC_DEFAULT_MODEL" not in os.environ assert calls == [["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "--debug"]] + def test_savings_statusline_off_installed_on_vanilla_launch(self, monkeypatch): + # Default-on savings row + non-routed launch + a workspace => settings become inline JSON + # carrying the "off" status row (routing_enabled=False, so no --routing-enabled flag). + calls: list[list[str]] = [] + monkeypatch.delenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, raising=False) + monkeypatch.delenv(v2.ENABLE_SUBAGENT_ROUTING_ENV_VAR, raising=False) + monkeypatch.delenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, raising=False) + monkeypatch.delenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, raising=False) + monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") + monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) + + claude.launch({"workspace": WS, "profile": "test"}, ["--debug"], options=LaunchOptions()) + + assert calls[0][:2] == ["claude", "--settings"] + command = json.loads(calls[0][2])["statusLine"]["command"] + assert f"-m {claude_statusline.MODULE}" in command + assert "--routing-enabled" not in command + + @staticmethod + def _launch_after_routing_setup_failure(monkeypatch, savings: str) -> dict: + """Launch routed, with `launch_claude` failing to write its files, and return the settings.""" + calls: list[list[str]] = [] + monkeypatch.setenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, "1") + monkeypatch.setenv(v2.ENABLE_SUBAGENT_ROUTING_ENV_VAR, "1") + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, savings) + monkeypatch.setenv(claude.FIRST_PROMPT_SOCKET_ENV, "/tmp/first.sock") + monkeypatch.setenv("OAUTH_TOKEN", "stale") + monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") + monkeypatch.setattr( + v2, "launch_claude", Mock(side_effect=v2.ClaudeRoutingSetupError("disk full")) + ) + monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) + + claude.launch( + {"workspace": WS, "profile": "test"}, + ["--debug"], + options=LaunchOptions(launch_smart_routing=True), + ) + + return json.loads(calls[0][2]) + + def test_routing_setup_failure_launches_normally_with_the_off_row(self, monkeypatch): + # The fallback forces routing off for this launch, so the status row must say "off" + # rather than be missing (or be the routed launch's "on", which never got written). + settings = self._launch_after_routing_setup_failure(monkeypatch, "1") + + assert settings["env"][v2.ENABLE_SMART_ROUTING_ENV_VAR] == "0" + command = settings["statusLine"]["command"] + assert f"-m {claude_statusline.MODULE}" in command + assert "--routing-enabled" not in command + + def test_routing_setup_failure_respects_the_savings_opt_out(self, monkeypatch): + settings = self._launch_after_routing_setup_failure(monkeypatch, "0") + + assert settings["env"][v2.ENABLE_SMART_ROUTING_ENV_VAR] == "0" + assert "statusLine" not in settings + def test_windows_launch_preserves_prompt_as_literal_argv(self, monkeypatch, tmp_path): native_binary = tmp_path / "Claude Code" / "claude.exe" prompt = 'keep "quotes" & pipes | and %PATH% literal' @@ -2573,6 +2631,7 @@ def test_managed_picker_preserves_available_saved_model( monkeypatch.setattr(claude, "CLAUDE_SETTINGS_PATH", settings_path) monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "0") claude.launch( { @@ -2650,6 +2709,7 @@ def test_managed_picker_replaces_stale_saved_model_for_this_launch( def test_v2_noninteractive_launch_bypasses_first_prompt_routing(self, monkeypatch, tool_args): calls: list[list[str]] = [] monkeypatch.setenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, "1") + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "0") monkeypatch.setattr(v2, "launch_claude", Mock()) monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) @@ -2686,6 +2746,7 @@ def test_gateway_discovery_uses_direct_gateway(self, monkeypatch): calls: list[list[str]] = [] monkeypatch.delenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, raising=False) monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1") + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "0") monkeypatch.delenv("OAUTH_TOKEN", raising=False) monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) @@ -2700,6 +2761,7 @@ def test_gateway_discovery_enabled_under_provider(self, monkeypatch): calls: list[list[str]] = [] monkeypatch.delenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, raising=False) monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1") + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "0") monkeypatch.delenv("OAUTH_TOKEN", raising=False) monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) @@ -2718,6 +2780,29 @@ def test_gateway_discovery_enabled_under_provider(self, monkeypatch): assert calls == [["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "--debug"]] +class TestVanillaLaunchSavingsStatusline: + """A plain (non-routed) launch installs a "routing off" statusline row.""" + + def test_routing_off_command_omits_routing_enabled_flag(self, tmp_path, monkeypatch): + """install_savings_statusline with routing_enabled=False produces no --routing-enabled.""" + settings: dict = {} + user_settings_path = tmp_path / "settings.json" + user_settings_path.touch() + monkeypatch.setattr(v2, "APP_DIR", tmp_path) + + v2.install_savings_statusline( + settings, + user_settings_path, + price_cache=tmp_path / "prices.json", + routing_enabled=False, + baseline_session_start=False, + ) + + command = settings["statusLine"]["command"] + assert f"-m {claude_statusline.MODULE}" in command + assert "--routing-enabled" not in command + + class TestWriteToolConfigPrunesStaleModelEnv: """Stale ucode-managed model env keys (ANTHROPIC_MODEL, etc.) from earlier ucode versions must be removed on every launch โ€” otherwise they linger in diff --git a/tests/test_claude_mods_savings.py b/tests/test_claude_mods_savings.py new file mode 100644 index 000000000..08f85b58e --- /dev/null +++ b/tests/test_claude_mods_savings.py @@ -0,0 +1,116 @@ +"""Tests for the savings concern of ug's smart-routing Claude Code UI mod.""" + +from __future__ import annotations + +import json +from pathlib import Path + +from ucode import mods +from ucode.smart_routing.claude_statusline import ( + MOD_USAGE_FILENAME, + PRICER_ENV_VAR, + SESSION_ENV_FILE_ENV_VAR, +) + +MODS_DIR = Path(__file__).resolve().parents[1] / "typescript" / "claude-mods" +SAVINGS_TS = "smart-routing-savings.ts" + + +def _source(name: str) -> str: + return (MODS_DIR / name).read_text(encoding="utf-8") + + +def test_smart_routing_ui_ships_the_savings_concern(): + assert SAVINGS_TS in mods.SMART_ROUTING_UI.extra + + +def test_write_mod_copies_the_savings_concern(tmp_path): + plugin_dir = tmp_path / "routing-plugin" + + mods.write_mod(plugin_dir, mods.SMART_ROUTING_UI) + + hooks_dir = plugin_dir / "hooks" + assert json.loads((hooks_dir / "hooks.json").read_text()) == {"modules": ["./register.ts"]} + assert (hooks_dir / SAVINGS_TS).read_text() == _source(SAVINGS_TS) + + +def test_entry_composes_savings_and_the_band_draws_its_segments(): + entry = _source("register.ts") + assert "import { registerSavings } from './smart-routing-savings'" in entry + assert "registerSavings(on)" in entry + + status = _source("smart-routing-status.ts") + assert "import { savingsSegments } from './smart-routing-savings'" in status + assert "savingsSegments()" in status + + +def test_savings_concern_matches_the_python_contract(): + source = _source(SAVINGS_TS) + # $.env.get needs string literals, so the names appear verbatim. + assert f"'{PRICER_ENV_VAR}'" in source + assert f"'{MOD_USAGE_FILENAME}'" in source + assert "'--mod-usage'" in source + assert f"'{SESSION_ENV_FILE_ENV_VAR}'" in source + # One pricer line out: {"savings": ..., "plugin": ...}. + assert "out.savings" in source and "out.plugin" in source + # Document in: {"version": 1, "start_model": ..., "main_model": ..., "first_main_model": ..., + # "user_switched": ..., "entries": [{..., "before_user_switch": ...}]}; Python ignores the + # main-model fields, the mod reads them back to seed itself after a reload. + for field in ( + "version: 1,", + "start_model: startModel,", + "main_model: lastMainModel,", + "first_main_model: firstMainModel,", + "user_switched: userSwitched,", + "entries: Array.from(entries.values())", + "before_user_switch: beforeUserSwitch,", + ): + assert field in source, field + assert "saved.version !== 1" in source + for field in ("main_model", "start_model", "first_main_model", "user_switched"): + assert f"saved.{field}" in source, field + assert "item.before_user_switch === true" in source + + +def test_savings_concern_keeps_the_baseline_and_start_model_rules(): + source = _source(SAVINGS_TS) + # Subagents price against the main model's request id; main is its own baseline. + assert "lastMainModel = str(request.model) ?? lastMainModel" in source + assert "const base = agent === null ? served : baseline" in source + # Served is priced by the model service the request named, not the drifting id reported back. + assert "const served = str(request.model) ?? str(usage?.model)" in source + # The start model is only ever set once, so a later session.start cannot overwrite it. + assert "startModel ??=" in source + assert "startModel = " not in source.replace("startModel ??=", "") + # First-prompt routing switches the main model before its first request, so a main request + # on another model marks a user switch; each request is flagged before it runs. + assert "firstMainModel ??= lastMainModel" in source + assert "userSwitched ||= lastMainModel !== firstMainModel" in source + assert source.index("const beforeUserSwitch = !userSwitched") < source.index( + "const result = yield* next(e)" + ) + # Every turn.step counts (side queries are not turn steps): no turn-id filter. + assert "mainTurns" not in source + + +def test_savings_concern_hooks_the_events_it_needs(): + source = _source(SAVINGS_TS) + for event in ("session.start", "turn.step", "turn.complete"): + assert f"on('{event}'" in source + assert "on('turn.start'" not in source + + +def test_savings_concern_recovers_from_failures(): + source = _source(SAVINGS_TS) + # A failed pricer run drops the stale estimate; a failed write is retried next turn. + assert "savings: null, plugin: priced.plugin" in source + assert "dirty = true\n throw error" in source + + +def test_only_the_savings_concern_hooks_the_turn_events(): + # Registering an event twice without a matcher fails the whole hooks module, + # so the concerns composed by register.ts must not share these events. + for name in ("smart-routing-status.ts", "subagent-routing.ts", "register.ts"): + source = _source(name) + for event in ("session.start", "turn.start", "turn.step", "turn.complete"): + assert f"on('{event}'" not in source, (name, event) diff --git a/tests/test_claude_smart_routing_v2.py b/tests/test_claude_smart_routing_v2.py index 2ecd94498..cee11ef73 100644 --- a/tests/test_claude_smart_routing_v2.py +++ b/tests/test_claude_smart_routing_v2.py @@ -7,7 +7,9 @@ import sys import threading import time +from decimal import Decimal from pathlib import Path +from types import SimpleNamespace from unittest.mock import Mock import pytest @@ -15,7 +17,21 @@ from ucode import mods from ucode.agents import LaunchOptions, claude from ucode.databricks import AnthropicModelCatalog -from ucode.smart_routing import claude_hooks, claude_pty, routing, v2 +from ucode.smart_routing import claude_hooks, claude_pty, claude_statusline, pricing, routing, v2 +from ucode.smart_routing.pricing import ModelPrice + +_START_SAVINGS_PRICE_REFRESH = v2._start_savings_price_refresh + + +@pytest.fixture(autouse=True) +def no_savings_price_refresh(monkeypatch) -> Mock: + """Routed launches start a background price refresh; keep its thread and network call out. + + ``TestSavingsStatusline`` restores the real starter and runs it inline against a fake rates API. + """ + refresh = Mock() + monkeypatch.setattr(v2, "_start_savings_price_refresh", refresh) + return refresh def _plugin_agent_models(plugin_dir: Path) -> set[str]: @@ -214,6 +230,14 @@ def test_smart_routing_vars_accept_true_as_well_as_one(self): assert v2.first_prompt_routing_enabled({v2_var: "true"}) assert not v2.first_prompt_routing_enabled({v2_var: "true", sub_var: "True"}) + def test_savings_statusline_is_on_by_default_and_opts_out_with_zero(self, monkeypatch): + monkeypatch.delenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, raising=False) + assert v2.savings_statusline_enabled() + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "0") + assert not v2.savings_statusline_enabled() + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "1") + assert v2.savings_statusline_enabled() + class TestV2Launch: def test_strips_gateway_prefix_for_interposer(self): @@ -424,6 +448,8 @@ def test_subagent_only_launch_skips_first_prompt_routing(self, tmp_path, monkeyp user_settings = tmp_path / "settings.json" user_settings.write_text(json.dumps({"model": "opus"})) monkeypatch.delenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, raising=False) + # Opt out of the (default-on) savings row to check the user's statusline is left untouched. + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, "0") monkeypatch.setenv(v2.ENABLE_SUBAGENT_ROUTING_ENV_VAR, "1") monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "") monkeypatch.setenv("CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY", "") @@ -504,6 +530,227 @@ def send_signal(self, _signal): assert not captured["plugin_dir"].exists() # The model-setting guard is a first-prompt concern; user settings stay untouched. assert json.loads(user_settings.read_text()) == {"model": "opus"} + # With the savings row opted out, the user's own statusline is left alone. + assert "statusLine" not in settings + + +class TestSavingsStatusline: + """``ENABLE_SMART_ROUTING_SAVINGS`` puts the savings row in the per-launch settings.""" + + @staticmethod + def _launch( + monkeypatch, tmp_path, *, first_prompt: bool, mods: bool = False, savings: str = "1" + ) -> dict: + user_settings = tmp_path / "settings.json" + user_settings.write_text( + json.dumps({"statusLine": {"type": "command", "command": "my-line", "padding": 1}}) + ) + monkeypatch.chdir(tmp_path) + monkeypatch.setenv(v2.ENABLE_SAVINGS_STATUSLINE_ENV_VAR, savings) + if mods: + monkeypatch.setenv(v2.ENABLE_CLAUDE_CODE_MODS_ENV_VAR, "1") + # The mod's TypeScript is not under test here, only how the launch wires it. + monkeypatch.setattr(v2.mods, "write_mod", Mock()) + else: + monkeypatch.delenv(v2.ENABLE_CLAUDE_CODE_MODS_ENV_VAR, raising=False) + monkeypatch.setattr(v2, "_start_savings_price_refresh", _START_SAVINGS_PRICE_REFRESH) + if first_prompt: + monkeypatch.setenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, "1") + monkeypatch.delenv(v2.ENABLE_SUBAGENT_ROUTING_ENV_VAR, raising=False) + else: + monkeypatch.delenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, raising=False) + monkeypatch.setenv(v2.ENABLE_SUBAGENT_ROUTING_ENV_VAR, "1") + monkeypatch.setattr(v2, "APP_DIR", tmp_path) + monkeypatch.setattr(v2, "CLAUDE_PTY_LOG", tmp_path / "v2.log") + monkeypatch.setattr(v2, "_model_picker_catalog", lambda: None) + monkeypatch.setattr(v2, "get_databricks_token", lambda *_args, **_kwargs: "token") + monkeypatch.setattr(v2, "build_auth_token_argv", lambda *_args, **_kwargs: ["ug"]) + monkeypatch.setattr( + v2, + "list_anthropic_model_catalog", + lambda *_args: AnthropicModelCatalog( + model_ids=["system.ai.claude-opus-4-8"], model_id_to_display_name={} + ), + ) + captured: dict = {} + + def capture_settings(argv) -> None: + captured["settings"] = json.loads(Path(argv[argv.index("--settings") + 1]).read_text()) + + def fake_pty(argv, **_kwargs): + capture_settings(argv) + return 0 + + class FakeProcess: + def __init__(self, argv, **_kwargs): + capture_settings(argv) + + def wait(self): + return 0 + + def fake_rates(workspace, token, model_services): + captured["rates_request"] = (workspace, token, model_services) + return [ + { + "model_service": "system.ai.claude-opus-4-8", + "costs": [ + { + "unit": "USD", + "token_costs": [ + {"token_type": "TOKEN_TYPE_INPUT", "cost": 5}, + {"token_type": "TOKEN_TYPE_OUTPUT", "cost": 25}, + ], + } + ], + } + ], None + + class InlineThread: + """Runs the price refresh inline so the test can observe the cache it writes.""" + + def __init__(self, *, target, args, **_kwargs): + self._target, self._args = target, args + + def start(self): + self._target(*self._args) + + monkeypatch.setattr(claude_pty, "run_claude_pty", fake_pty) + monkeypatch.setattr(v2.subprocess, "Popen", FakeProcess) + monkeypatch.setattr(v2, "fetch_endpoint_rates", fake_rates) + monkeypatch.setattr(v2, "threading", SimpleNamespace(Thread=InlineThread)) + + with pytest.raises(SystemExit): + v2.launch_claude( + {"workspace": "https://example.com"}, + [], + binary="claude", + user_settings_path=user_settings, + launch_model="opus", + compose_settings=lambda _args: ({}, []), + launch_model_args=claude._launch_model_args, + model_name=claude._maybe_add_1m_suffix, + ) + return captured + + def test_subagent_only_wraps_the_user_statusline(self, monkeypatch, tmp_path): + captured = self._launch(monkeypatch, tmp_path, first_prompt=False) + status_line = captured["settings"]["statusLine"] + + assert status_line["type"] == "command" + assert status_line["padding"] == 1 + command = status_line["command"] + assert "\nmy-line\n" in command + assert f"-m {claude_statusline.MODULE}" in command + assert f"--state-dir {tmp_path / claude_statusline.STATE_DIRNAME}" in command + price_cache = pricing.price_cache_path(tmp_path, "https://example.com") + assert f"--price-cache {price_cache}" in command + assert "--routing-enabled" in command + # The baseline follows the main model the user chose. + assert "--baseline-session-start" not in command + # Only the mod runs a pricer; the statusline is self-contained. + assert claude_statusline.PRICER_ENV_VAR not in captured["settings"]["env"] + + def test_first_prompt_routing_uses_the_pre_routing_baseline(self, monkeypatch, tmp_path): + status_line = self._launch(monkeypatch, tmp_path, first_prompt=True)["settings"][ + "statusLine" + ] + + assert "--routing-enabled" in status_line["command"] + assert "--baseline-session-start" in status_line["command"] + + def test_caches_endpoint_rates_for_the_launch_models(self, monkeypatch, tmp_path): + captured = self._launch(monkeypatch, tmp_path, first_prompt=False) + + # The `opus` alias isn't a model service, so only the catalog id is priced. + assert captured["rates_request"] == ( + "https://example.com", + "token", + ["system.ai.claude-opus-4-8"], + ) + cached = pricing.read_price_cache(pricing.price_cache_path(tmp_path, "https://example.com")) + assert cached is not None + assert cached[0] == { + "claude-opus-4-8": ModelPrice(input=Decimal("5"), output=Decimal("25")) + } + + @pytest.mark.parametrize("first_prompt", [False, True]) + def test_mods_get_a_pricer_command_instead_of_a_statusline( + self, monkeypatch, tmp_path, first_prompt + ): + captured = self._launch(monkeypatch, tmp_path, first_prompt=first_prompt, mods=True) + + settings = captured["settings"] + # The mod draws the savings itself; the user's own statusline is left alone. + assert "statusLine" not in settings + price_cache = pricing.price_cache_path(tmp_path, "https://example.com") + assert json.loads(settings["env"][claude_statusline.PRICER_ENV_VAR]) == ( + claude_statusline.mod_pricer_argv( + python=sys.executable, + price_cache=price_cache, + baseline_session_start=first_prompt, + ) + ) + # The pricer reads the same cache, so the price refresh still runs. + assert captured["rates_request"] == ( + "https://example.com", + "token", + ["system.ai.claude-opus-4-8"], + ) + assert pricing.read_price_cache(price_cache) is not None + + def test_opting_out_of_savings_skips_the_mods_pricer_and_refresh(self, monkeypatch, tmp_path): + captured = self._launch(monkeypatch, tmp_path, first_prompt=False, mods=True, savings="0") + + assert claude_statusline.PRICER_ENV_VAR not in captured["settings"]["env"] + assert "statusLine" not in captured["settings"] + assert "rates_request" not in captured + + +class TestSavingsPriceRefresh: + def test_prices_routable_and_main_models_by_system_ai_name(self): + services = v2._savings_model_services( + [ + "system.ai.claude-opus-4-8[1m]", + "anthropic-aigw-73ea02b2-system.ai.glm-5-2", + "databricks-claude-sonnet-5", + "opus", + None, + ], + { + "env": { + "ANTHROPIC_DEFAULT_OPUS_MODEL": "system.ai.claude-opus-5-5", + "UNRELATED": "system.ai.not-a-main-model", + } + }, + ) + + assert services == [ + "system.ai.claude-opus-4-8", + "system.ai.claude-opus-5-5", + "system.ai.claude-sonnet-5", + "system.ai.glm-5-2", + ] + + @pytest.mark.parametrize( + "fetch", + [ + lambda *_args: ([], "HTTP 404 Not Found"), + lambda *_args: (_ for _ in ()).throw(RuntimeError("rates API down")), + ], + ) + def test_failed_refresh_keeps_the_previous_cache(self, monkeypatch, tmp_path, fetch): + cache = tmp_path / "model-prices.json" + pricing.write_price_cache( + cache, {"claude-opus-4-8": ModelPrice(input=Decimal("5"))}, now=1.0 + ) + monkeypatch.setattr(v2, "fetch_endpoint_rates", fetch) + + v2._refresh_savings_prices( + "https://example.com", "token", ["system.ai.claude-opus-4-8"], cache + ) + + cached = pricing.read_price_cache(cache) + assert cached is not None and cached[1] == "1.0" class TestV2ModelPickerDiscovery: diff --git a/tests/test_claude_statusline.py b/tests/test_claude_statusline.py new file mode 100644 index 000000000..c033a213f --- /dev/null +++ b/tests/test_claude_statusline.py @@ -0,0 +1,1154 @@ +"""Tests for the Claude Code smart-routing statusline row.""" + +from __future__ import annotations + +import hashlib +import json +import os +import shlex +import shutil +import subprocess +import sys +from datetime import datetime +from decimal import Decimal +from pathlib import Path + +import pytest + +from ucode.constants import ENABLE_SMART_ROUTING_ENV_VAR, ENABLE_SUBAGENT_ROUTING_ENV_VAR +from ucode.smart_routing import claude_statusline, pricing, session_env +from ucode.smart_routing.pricing import ModelPrice + +OPUS_ID = "system.ai.claude-opus-4-8" +SONNET_ID = "system.ai.claude-sonnet-5" +OPUS = {"id": OPUS_ID, "display_name": "Opus 4.8"} +SONNET = {"id": SONNET_ID, "display_name": "Sonnet 5"} +PRICES = { + OPUS_ID: ModelPrice( + input=Decimal("5"), + output=Decimal("25"), + cache_read=Decimal("0.5"), + cache_write_5m=Decimal("6.25"), + cache_write_1h=Decimal("10"), + ), + SONNET_ID: ModelPrice( + input=Decimal("2"), + output=Decimal("10"), + cache_read=Decimal("0.2"), + cache_write_5m=Decimal("2.5"), + cache_write_1h=Decimal("4"), + ), +} +# Opus main-agent response: 10*5 + 100k*0.5 + 1k*25 = $0.07505 (same either way). +MAIN_USAGE = {"input_tokens": 10, "cache_read_input_tokens": 100_000, "output_tokens": 1_000} +# Subagent response: $0.122 on Sonnet (1k*2 + 20k*4 + 4k*10), $0.305 at Opus rates. +SUBAGENT_USAGE = { + "input_tokens": 1_000, + "cache_creation_input_tokens": 20_000, + "cache_creation": {"ephemeral_1h_input_tokens": 20_000, "ephemeral_5m_input_tokens": 0}, + "output_tokens": 4_000, +} + + +# Transcript timestamps, oldest first. +T0, T1, T2, T3 = (f"2026-05-01T10:0{minute}:00.000Z" for minute in range(4)) +HAIKU_ID = "system.ai.claude-haiku-4-5" + + +def epoch(timestamp: str) -> float: + return datetime.fromisoformat(timestamp).timestamp() + + +def response( + message_id: str, model: str, usage: dict, block: str = "text", at: str | None = None +) -> dict: + record = { + "type": "assistant", + "uuid": f"{message_id}-{block}", + "message": {"id": message_id, "model": model, "usage": usage, "content": [{"type": block}]}, + } + if at is not None: + record["timestamp"] = at + return record + + +def append(path: Path, *records: dict, trailing_newline: bool = True) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + text = "\n".join(json.dumps(record) for record in records) + with path.open("a", encoding="utf-8") as handle: + handle.write(text + ("\n" if trailing_newline else "")) + + +def write_plugin_version( + config_dir: Path, + version: str, + *, + name: str = "model-orchestrator", + marketplace: str = "example-marketplace", +) -> None: + """Write a Claude Code ``installed_plugins.json`` recording ``name`` at ``version``.""" + plugins_dir = config_dir / "plugins" + plugins_dir.mkdir(parents=True, exist_ok=True) + (plugins_dir / "installed_plugins.json").write_text( + json.dumps( + { + "version": 2, + "plugins": {f"{name}@{marketplace}": [{"scope": "user", "version": version}]}, + } + ), + encoding="utf-8", + ) + + +@pytest.fixture(autouse=True) +def no_session_env_file(monkeypatch) -> None: + """Render reads the hooks' session file named by this env var; keep the developer's out.""" + monkeypatch.delenv(claude_statusline.SESSION_ENV_FILE_ENV_VAR, raising=False) + + +@pytest.fixture(autouse=True) +def claude_config_dir(tmp_path, monkeypatch) -> Path: + """Point the row's plugin-version lookup at an empty per-test config dir (no version by default). + + The row reads Claude Code's plugin state via ``CLAUDE_CONFIG_DIR``; isolating it keeps the + developer's real installed plugins out of the assertions. + """ + config_dir = tmp_path / "claude-config" + config_dir.mkdir(exist_ok=True) + monkeypatch.setenv("CLAUDE_CONFIG_DIR", str(config_dir)) + return config_dir + + +class Session: + """One Claude Code session's transcript layout plus ug's statusline state.""" + + def __init__(self, root: Path, session_id: str = "session-1") -> None: + self.session_id = session_id + self.transcript = root / "project" / "transcript.jsonl" + self.transcript.parent.mkdir(parents=True) + self.transcript.touch() + self.state_dir = root / "claude-savings" + self.price_cache = root / "model-prices.json" + + def subagent(self, agent_id: str) -> Path: + return self.transcript.with_suffix("") / "subagents" / f"agent-{agent_id}.jsonl" + + def payload(self, model: dict) -> str: + return json.dumps( + {"session_id": self.session_id, "transcript_path": str(self.transcript), "model": model} + ) + + def render( + self, + model: dict = OPUS, + *, + baseline_session_start: bool = False, + routing_enabled: bool = True, + ) -> str: + return claude_statusline.render( + self.payload(model), + routing_enabled=routing_enabled, + state_dir=self.state_dir, + price_cache=self.price_cache, + baseline_session_start=baseline_session_start, + ) + + +@pytest.fixture +def session(tmp_path) -> Session: + session = Session(tmp_path) + pricing.write_price_cache(session.price_cache, PRICES, now=1.0) + return session + + +class TestRender: + def test_prices_all_tokens_at_main_model_minus_actual_cost(self, session): + append( + session.transcript, + response("msg-main", OPUS_ID, MAIN_USAGE, "text"), + response("msg-main", OPUS_ID, MAIN_USAGE, "tool_use"), + ) + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + # Baseline $0.38005 (all tokens at Opus) - actual $0.19705 = $0.183 saved (48%). + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (48%)" + + def test_prices_a_served_bedrock_id_with_its_system_ai_rate(self, session): + haiku = ModelPrice( + input=Decimal("1"), + output=Decimal("5"), + cache_read=Decimal("0.1"), + cache_write_5m=Decimal("1.25"), + cache_write_1h=Decimal("2"), + ) + # Endpoint rates are keyed by system.ai name; Haiku responses carry its Bedrock id. + pricing.write_price_cache( + session.price_cache, {**PRICES, "system.ai.claude-haiku-4-5": haiku}, now=2.0 + ) + append( + session.subagent("a1"), + response("msg-sub", "anthropic.claude-haiku-4-5-20251001-v1:0", SUBAGENT_USAGE), + ) + + # $0.061 on Haiku (1k*1 + 20k*2 + 4k*5) vs $0.305 at Opus rates. + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: $0.24 (80%)" + + def test_counts_a_response_split_across_records_once(self, session): + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + once = session.render() + append(session.subagent("a2"), response("msg-other", SONNET_ID, SUBAGENT_USAGE, "text")) + append(session.subagent("a2"), response("msg-other", SONNET_ID, SUBAGENT_USAGE, "tool_use")) + + assert once == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)" + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: $0.37 (60%)" + + def test_reads_transcripts_incrementally(self, session): + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + first = session.render() + append( + session.subagent("a1"), + response("msg-late", SONNET_ID, SUBAGENT_USAGE), + trailing_newline=False, + ) + + assert first == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)" + # A line Claude Code is still writing is left for the next refresh. + assert session.render() == first + with session.subagent("a1").open("a", encoding="utf-8") as handle: + handle.write("\n") + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: $0.37 (60%)" + + def test_starts_over_when_a_transcript_is_rewritten(self, session): + append( + session.subagent("a1"), + *(response(f"m{i}", SONNET_ID, SUBAGENT_USAGE) for i in range(3)), + ) + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: $0.55 (60%)" + session.subagent("a1").write_text( + json.dumps(response("m0", SONNET_ID, SUBAGENT_USAGE)) + "\n" + ) + + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)" + + def test_shows_a_net_cost_increase_honestly(self, session): + append(session.transcript, response("msg-main", SONNET_ID, MAIN_USAGE)) + append(session.subagent("a1"), response("msg-sub", OPUS_ID, SUBAGENT_USAGE)) + + # Baseline $0.152 (all at Sonnet) vs actual $0.335: routing cost $0.183 more. + assert session.render(SONNET) == "Smart routing cost $0.18 more (120%)" + + def test_baseline_follows_the_main_model_the_user_chose(self, tmp_path, session): + other = Session(tmp_path / "other") + pricing.write_price_cache(other.price_cache, PRICES, now=1.0) + for each in (session, other): + append(each.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + assert session.render(OPUS) == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)" + # Nothing ran on a model other than a Sonnet baseline, so there is nothing to claim yet. + assert other.render(SONNET) == "Smart routing on" + + def test_first_prompt_routing_uses_the_pre_routing_model(self, session): + # The first refresh happens before the first prompt is routed: nothing to claim yet. + assert session.render(OPUS, baseline_session_start=True) == "Smart routing on" + append(session.transcript, response("msg-main", SONNET_ID, MAIN_USAGE)) + + # Main-agent tokens the router moved to Sonnet count toward savings too: + # $0.07505 at Opus vs 10*2 + 100k*0.2 + 1k*10 = $0.03002 at Sonnet. + assert ( + session.render(SONNET, baseline_session_start=True) + == "๐Ÿ’ฐ Est. saved with smart routing: $0.05 (60%)" + ) + + def test_reprices_when_the_price_cache_refreshes(self, session): + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + before = session.render() + pricing.write_price_cache( + session.price_cache, + { + **PRICES, + SONNET_ID: ModelPrice( + input=Decimal("5"), + output=Decimal("25"), + cache_read=Decimal("0.5"), + cache_write_5m=Decimal("6.25"), + cache_write_1h=Decimal("10"), + ), + }, + now=2.0, + ) + + assert before == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)" + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: <$0.01 (0%)" + + @pytest.mark.parametrize("model", ["system.ai.claude-mystery-1", ""]) + def test_shows_on_when_a_response_cannot_be_priced(self, session, model): + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + append(session.subagent("a2"), response("msg-x", model, SUBAGENT_USAGE)) + + # An unpriceable response would undercount the estimate, so the row falls back to "on". + assert session.render() == "Smart routing on" + + def test_shows_on_without_prices(self, session): + session.price_cache.unlink() + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + assert session.render() == "Smart routing on" + + def test_shows_on_when_the_baseline_model_is_unpriced(self, session): + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + model = {"id": "system.ai.claude-unknown", "display_name": "?"} + assert session.render(model) == "Smart routing on" + + def test_shows_the_orchestrator_plugin_version(self, session, claude_config_dir): + write_plugin_version(claude_config_dir, "0.4.4") + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + assert session.render() == ( + "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%) ยท smart router plugin v0.4.4" + ) + + def test_routing_disabled_shows_off(self, session): + # routing_enabled=False short-circuits before reading any transcript; empty is fine. + assert session.render(routing_enabled=False) == "Smart routing off" + + def test_routing_disabled_with_plugin_version(self, session, claude_config_dir): + write_plugin_version(claude_config_dir, "0.4.4") + assert ( + session.render(routing_enabled=False) + == "Smart routing off ยท smart router plugin v0.4.4" + ) + + def test_ignores_responses_without_usage(self, session): + append( + session.transcript, + response("msg-synthetic", "", {"input_tokens": 0, "output_tokens": 0}), + {"type": "user", "message": {"content": "hi"}}, + {"type": "assistant", "message": "not a dict"}, + ) + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)" + + @pytest.mark.parametrize( + "raw", + [ + "", + "not json", + "[]", + json.dumps({"session_id": "s", "model": OPUS}), + json.dumps({"session_id": "s", "transcript_path": "/x.jsonl", "model": {}}), + ], + ) + def test_malformed_payload_still_shows_the_row(self, session, raw): + # No usable payload means no savings, but the version/state row still renders. + assert ( + claude_statusline.render( + raw, + routing_enabled=True, + state_dir=session.state_dir, + price_cache=session.price_cache, + baseline_session_start=False, + ) + == "Smart routing on" + ) + + def test_unsafe_session_ids_do_not_escape_the_state_dir(self, tmp_path): + session = Session(tmp_path, session_id="../../escape") + pricing.write_price_cache(session.price_cache, PRICES) + + session.render() + + hashed = hashlib.sha256(b"../../escape").hexdigest() + assert [path.name for path in session.state_dir.iterdir()] == [f"{hashed}.json"] + + +class TestMainModelTimeline: + """The baseline is fixed per response, so `/model` never reprices work already done.""" + + def test_switching_the_main_model_does_not_reprice_earlier_work(self, session): + append(session.transcript, response("m1", OPUS_ID, MAIN_USAGE, at=T0)) + append(session.subagent("a1"), response("s1", SONNET_ID, SUBAGENT_USAGE, at=T1)) + # Baseline $0.38005 (main + subagent at Opus) - actual $0.19705. + assert session.render(OPUS) == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (48%)" + + # `/model sonnet`: the first subagent still counts against the Opus it ran beside, while + # work after the switch is measured against Sonnet (the old pricing would show "on"). + append(session.transcript, response("m2", SONNET_ID, MAIN_USAGE, at=T2)) + append(session.subagent("a1"), response("s2", SONNET_ID, SUBAGENT_USAGE, at=T3)) + + # Baseline $0.53207 (adds $0.03002 + $0.122 at Sonnet) - actual $0.34907. + assert session.render(SONNET) == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (34%)" + + def test_a_subagent_response_uses_the_main_model_in_effect_at_its_timestamp(self, session): + append( + session.transcript, + response("m1", OPUS_ID, MAIN_USAGE, at=T0), + response("m2", SONNET_ID, MAIN_USAGE, at=T2), + ) + append( + session.subagent("a1"), + response("s1", SONNET_ID, SUBAGENT_USAGE, at=T1), + response("s2", SONNET_ID, SUBAGENT_USAGE, at=T3), + ) + + # Only s1 ran beside Opus; s2 ran beside Sonnet, the main model by then. + assert session.render(SONNET) == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (34%)" + + def test_a_subagent_response_without_a_timestamp_uses_the_latest_main_model(self, session): + append( + session.transcript, + response("m1", OPUS_ID, MAIN_USAGE, at=T0), + response("m2", SONNET_ID, MAIN_USAGE, at=T2), + ) + append(session.subagent("a1"), response("s1", SONNET_ID, SUBAGENT_USAGE)) + + assert session.render(OPUS) == "Smart routing on" + + def test_a_subagent_response_before_the_first_main_record_uses_the_earliest_model( + self, session + ): + append(session.transcript, response("m1", OPUS_ID, MAIN_USAGE, at=T1)) + append(session.subagent("a1"), response("s1", SONNET_ID, SUBAGENT_USAGE, at=T0)) + + assert session.render(SONNET) == "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (48%)" + + def test_main_agent_responses_are_their_own_baseline(self, session): + # Choosing a cheaper main model with `/model` is the user's call, not routing's saving. + append(session.transcript, response("m1", SONNET_ID, MAIN_USAGE, at=T0)) + + assert session.render(OPUS) == "Smart routing on" + + def test_timeline_records_main_model_changes_only_and_persists(self, session): + append( + session.transcript, + response("m1", OPUS_ID, MAIN_USAGE, "text", at=T0), + response("m1", OPUS_ID, MAIN_USAGE, "tool_use", at=T0), + response("m2", OPUS_ID, MAIN_USAGE, at=T1), + ) + append(session.subagent("a1"), response("s1", HAIKU_ID, SUBAGENT_USAGE, at=T1)) + session.render() + append( + session.transcript, + response("m-synthetic", "", {"input_tokens": 5}, at=T1), + response("m3", SONNET_ID, MAIN_USAGE, at=T2), + response("m4", SONNET_ID, MAIN_USAGE, at=T3), + ) + session.render() + + state = claude_statusline._read_state( + claude_statusline._state_path(session.state_dir, session.session_id) + ) + assert state["files"][str(session.transcript)]["models"] == [ + [epoch(T0), "claude-opus-4-8"], + [epoch(T2), "claude-sonnet-5"], + ] + # Only the main transcript carries the timeline. + assert state["files"][str(session.subagent("a1"))]["models"] == [] + + def test_first_prompt_baseline_follows_a_later_user_switch(self, session): + assert session.render(OPUS, baseline_session_start=True) == "Smart routing on" + # The router kept Opus for the first answer, so the switch to Sonnet after it is the + # user's `/model`: neither m2 nor the subagent beside it is routing's saving (measuring + # both against the start model would claim $0.23). + append( + session.transcript, + response("m1", OPUS_ID, MAIN_USAGE, at=T0), + response("m2", SONNET_ID, MAIN_USAGE, at=T1), + ) + append(session.subagent("a1"), response("s1", SONNET_ID, SUBAGENT_USAGE, at=T2)) + + assert session.render(SONNET, baseline_session_start=True) == "Smart routing on" + + def test_first_prompt_routing_keeps_its_credit_after_a_user_switch(self, session): + assert session.render(OPUS, baseline_session_start=True) == "Smart routing on" + # The router picked Sonnet; its first answer and the subagent beside it count against Opus. + append(session.transcript, response("m1", SONNET_ID, MAIN_USAGE, at=T0)) + append(session.subagent("a1"), response("s1", SONNET_ID, SUBAGENT_USAGE, at=T1)) + # Baseline $0.38005 - actual $0.15202. + assert ( + session.render(SONNET, baseline_session_start=True) + == "๐Ÿ’ฐ Est. saved with smart routing: $0.23 (60%)" + ) + + # The user's `/model opus` then `/model sonnet` are their own baselines: the saving stays, + # over a baseline that grows by $0.10507 (against Opus it would be $0.27, 52%). + append( + session.transcript, + response("m2", OPUS_ID, MAIN_USAGE, at=T2), + response("m3", SONNET_ID, MAIN_USAGE, at=T3), + ) + assert ( + session.render(SONNET, baseline_session_start=True) + == "๐Ÿ’ฐ Est. saved with smart routing: $0.23 (47%)" + ) + + def test_discards_state_written_by_an_older_version(self, tmp_path): + path = tmp_path / "state.json" + path.write_text(json.dumps({"version": 1, "files": {"x": {}}})) + + assert claude_statusline._read_state(path) == {"version": claude_statusline._STATE_VERSION} + + def test_carries_the_start_model_over_from_an_older_version(self, tmp_path): + path = tmp_path / "state.json" + path.write_text( + json.dumps({"version": 1, "start_model": OPUS, "key": ["x", "y"], "files": {"x": {}}}) + ) + + assert claude_statusline._read_state(path) == { + "version": claude_statusline._STATE_VERSION, + "start_model": OPUS, + } + + @pytest.mark.parametrize( + "start_model", [None, "", "opus", {"display_name": "Opus"}, {"id": ""}] + ) + def test_drops_an_invalid_start_model_from_an_older_version(self, tmp_path, start_model): + path = tmp_path / "state.json" + path.write_text(json.dumps({"version": 1, "start_model": start_model})) + + assert claude_statusline._read_state(path) == {"version": claude_statusline._STATE_VERSION} + + def test_first_prompt_baseline_survives_a_state_version_change(self, session): + state_path = claude_statusline._state_path(session.state_dir, session.session_id) + state_path.parent.mkdir(parents=True) + state_path.write_text(json.dumps({"version": 1, "start_model": OPUS, "files": {}})) + append(session.transcript, response("msg-main", SONNET_ID, MAIN_USAGE)) + + # The session already runs on Sonnet; re-capturing it as the start model would hide the + # saving against Opus ($0.07505 vs $0.03002). + assert ( + session.render(SONNET, baseline_session_start=True) + == "๐Ÿ’ฐ Est. saved with smart routing: $0.05 (60%)" + ) + + +class TestSessionToggle: + """A `/smart-router` toggle rewrites the session file, which the statusLine command re-reads.""" + + SAVED = "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)" + OFF = {ENABLE_SMART_ROUTING_ENV_VAR: "0", ENABLE_SUBAGENT_ROUTING_ENV_VAR: "0"} + + @pytest.fixture + def toggle(self, session, tmp_path, monkeypatch) -> Path: + path = tmp_path / "session" / "env.json" + path.parent.mkdir() + monkeypatch.setenv(claude_statusline.SESSION_ENV_FILE_ENV_VAR, str(path)) + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + return path + + def test_env_var_name_matches_the_session_env_module(self): + assert claude_statusline.SESSION_ENV_FILE_ENV_VAR == session_env.SESSION_ENV_VAR + + @pytest.mark.parametrize( + ("contents", "shown"), + [ + ({}, "on"), + (OFF, "off"), + ({ENABLE_SMART_ROUTING_ENV_VAR: "0"}, "off"), + ({ENABLE_SUBAGENT_ROUTING_ENV_VAR: "1"}, "on"), + ({ENABLE_SUBAGENT_ROUTING_ENV_VAR: " TRUE "}, "on"), + ({ENABLE_SMART_ROUTING_ENV_VAR: "1", ENABLE_SUBAGENT_ROUTING_ENV_VAR: "0"}, "on"), + ({ENABLE_SUBAGENT_ROUTING_ENV_VAR: "yes"}, "off"), + ({"UNRELATED": "0"}, "on"), + ([], "on"), + ("nope", "on"), + ], + ) + def test_the_session_file_overrides_the_launch_flag(self, session, toggle, contents, shown): + toggle.write_text(json.dumps(contents)) + + expected = self.SAVED if shown == "on" else "Smart routing off" + assert session.render() == expected + + @pytest.mark.parametrize("damage", ["missing", "not json", ""]) + def test_an_unreadable_session_file_leaves_the_launch_flag_in_charge( + self, session, toggle, damage + ): + if damage != "missing": + toggle.write_text(damage) + + assert session.render() == self.SAVED + + def test_a_plain_launch_ignores_an_inherited_session_file(self, session, toggle): + toggle.write_text(json.dumps({ENABLE_SUBAGENT_ROUTING_ENV_VAR: "1"})) + + assert session.render(routing_enabled=False) == "Smart routing off" + + def test_toggling_takes_effect_on_the_next_refresh(self, session, toggle): + toggle.write_text(json.dumps({})) + assert session.render() == self.SAVED + + toggle.write_text(json.dumps(self.OFF)) + assert session.render() == "Smart routing off" + + toggle.write_text(json.dumps({})) + assert session.render() == self.SAVED + + +class TestModUsage: + """``--mod-usage`` prices the Claude Code mod's token sums into the mod's two segments.""" + + @staticmethod + def entry( + agent: str, + baseline: str | None, + served: str, + usage: dict, + *, + before_user_switch: bool = False, + ) -> dict: + found = { + "agent": agent, + "served": served, + "before_user_switch": before_user_switch, + "usage": usage, + } + return found if baseline is None else {**found, "baseline": baseline} + + @staticmethod + def write(tmp_path: Path, entries: list, *, start_model: str | None = OPUS_ID) -> Path: + path = tmp_path / "mod-usage.json" + path.write_text( + json.dumps({"version": 1, "start_model": start_model, "entries": entries}), + encoding="utf-8", + ) + return path + + @staticmethod + def price(session: Session, usage: Path, *, first_prompt: bool = False) -> dict: + line = claude_statusline.mod_usage_output( + usage, price_cache=session.price_cache, baseline_session_start=first_prompt + ) + assert "\n" not in line + return json.loads(line) + + def test_prices_tokens_against_each_entrys_baseline(self, session, tmp_path): + entries = [ + self.entry("main", OPUS_ID, OPUS_ID, MAIN_USAGE), + self.entry("a1", OPUS_ID, SONNET_ID, SUBAGENT_USAGE), + ] + + # Baseline $0.38005 - actual $0.19705. + assert self.price(session, self.write(tmp_path, entries)) == { + "savings": "๐Ÿ’ฐ Est. saved $0.18 (48%)", + "plugin": None, + } + + def test_an_entrys_baseline_is_the_main_model_when_its_request_ran(self, session, tmp_path): + entries = [ + self.entry("main", OPUS_ID, OPUS_ID, MAIN_USAGE), + self.entry("a1", OPUS_ID, SONNET_ID, SUBAGENT_USAGE), + self.entry("main", SONNET_ID, SONNET_ID, MAIN_USAGE), + self.entry("a1", SONNET_ID, SONNET_ID, SUBAGENT_USAGE), + ] + + assert self.price(session, self.write(tmp_path, entries))["savings"] == ( + "๐Ÿ’ฐ Est. saved $0.18 (34%)" + ) + + def test_reports_the_orchestrator_plugin_version(self, session, tmp_path, claude_config_dir): + write_plugin_version(claude_config_dir, "0.4.4") + + # No entries yet (the mod runs the pricer at session start): just the plugin segment. + assert self.price(session, self.write(tmp_path, [])) == { + "savings": None, + "plugin": "plugin v0.4.4", + } + + def test_shows_a_net_cost_increase_honestly(self, session, tmp_path): + entries = [ + self.entry("main", SONNET_ID, SONNET_ID, MAIN_USAGE), + self.entry("a1", SONNET_ID, OPUS_ID, SUBAGENT_USAGE), + ] + + # Baseline $0.152 vs actual $0.335. + assert self.price(session, self.write(tmp_path, entries))["savings"] == ( + "cost $0.18 more (120%)" + ) + + def test_nothing_to_claim_until_a_model_was_swapped(self, session, tmp_path): + entries = [self.entry("main", OPUS_ID, OPUS_ID, MAIN_USAGE)] + + assert self.price(session, self.write(tmp_path, entries))["savings"] is None + + def test_bills_entries_without_the_cache_write_split_as_one_hour_writes( + self, session, tmp_path + ): + usage = { + "input_tokens": 1_000, + "cache_creation_input_tokens": 20_000, + "output_tokens": 4_000, + } + + # ug launches Claude with 1-hour caching, so unsplit writes bill at the 1-hour rate, as + # SUBAGENT_USAGE does: $0.122 on Sonnet vs $0.305 at Opus rates. (5-minute: $0.092 vs $0.23.) + assert self.price( + session, self.write(tmp_path, [self.entry("a1", OPUS_ID, SONNET_ID, usage)]) + )["savings"] == ("๐Ÿ’ฐ Est. saved $0.18 (60%)") + + def test_honors_an_explicit_cache_write_split(self, session, tmp_path): + usage = { + "input_tokens": 1_000, + "cache_creation_input_tokens": 20_000, + "cache_creation": {"ephemeral_1h_input_tokens": 0, "ephemeral_5m_input_tokens": 20_000}, + "output_tokens": 4_000, + } + + # The split says 5-minute writes: $0.092 on Sonnet vs $0.23 at Opus rates. + assert self.price( + session, self.write(tmp_path, [self.entry("a1", OPUS_ID, SONNET_ID, usage)]) + )["savings"] == ("๐Ÿ’ฐ Est. saved $0.14 (60%)") + + def test_the_transcript_path_still_bills_unsplit_writes_as_five_minute_writes(self, session): + usage = { + "input_tokens": 1_000, + "cache_creation_input_tokens": 20_000, + "output_tokens": 4_000, + } + append(session.subagent("a1"), response("msg-sub", SONNET_ID, usage)) + + assert session.render() == "๐Ÿ’ฐ Est. saved with smart routing: $0.14 (60%)" + + def test_skips_entries_without_tokens(self, session, tmp_path): + entries = [ + self.entry("a1", OPUS_ID, SONNET_ID, SUBAGENT_USAGE), + self.entry("a2", "unpriced-baseline", "unpriced-served", {"input_tokens": 0}), + ] + + assert self.price(session, self.write(tmp_path, entries))["savings"] == ( + "๐Ÿ’ฐ Est. saved $0.18 (60%)" + ) + + def test_first_prompt_routing_prices_against_the_start_model(self, session, tmp_path): + # Under first-prompt routing the router picks the main model too, so an entry's own + # baseline (the routed model) would hide the saving; until the user switches the main + # model its presence or absence is ignored. + entries = [ + self.entry("main", SONNET_ID, SONNET_ID, MAIN_USAGE, before_user_switch=True), + self.entry("a1", None, SONNET_ID, SUBAGENT_USAGE, before_user_switch=True), + ] + usage = self.write(tmp_path, entries, start_model=OPUS_ID) + + # Baseline $0.38 at Opus vs actual $0.15202 on Sonnet (no fixed baseline: not rerouted). + assert self.price(session, usage, first_prompt=True)["savings"] == ( + "๐Ÿ’ฐ Est. saved $0.23 (60%)" + ) + assert self.price(session, usage)["savings"] is None + + def test_first_prompt_routing_prices_entries_after_a_user_switch_as_their_own( + self, session, tmp_path + ): + entries = [ + self.entry("main", SONNET_ID, SONNET_ID, MAIN_USAGE, before_user_switch=True), + self.entry("a1", SONNET_ID, SONNET_ID, SUBAGENT_USAGE, before_user_switch=True), + # The user's `/model opus`, then `/model sonnet`. + self.entry("main", OPUS_ID, OPUS_ID, MAIN_USAGE), + self.entry("main", SONNET_ID, SONNET_ID, MAIN_USAGE), + ] + usage = self.write(tmp_path, entries, start_model=OPUS_ID) + + # Baseline $0.48512 - actual $0.25709, as on the transcript path. + assert self.price(session, usage, first_prompt=True)["savings"] == ( + "๐Ÿ’ฐ Est. saved $0.23 (47%)" + ) + + @pytest.mark.parametrize("start_model", [None, ""]) + def test_first_prompt_routing_without_a_start_model_has_no_estimate( + self, session, tmp_path, start_model + ): + entries = [self.entry("a1", OPUS_ID, SONNET_ID, SUBAGENT_USAGE)] + usage = self.write(tmp_path, entries, start_model=start_model) + + assert self.price(session, usage, first_prompt=True)["savings"] is None + + @pytest.mark.parametrize( + "entries", + [ + # An unpriced served model, or an unpriced baseline, would undercount. + [("a1", OPUS_ID, "system.ai.claude-mystery-1")], + [("a1", "system.ai.claude-mystery-1", SONNET_ID)], + # Malformed entries. + [("a1", None, SONNET_ID)], + [("a1", "", SONNET_ID)], + [("a1", OPUS_ID, "")], + ], + ) + def test_an_unpriced_or_baseline_less_entry_voids_the_estimate( + self, session, tmp_path, entries + ): + good = self.entry("a0", OPUS_ID, SONNET_ID, SUBAGENT_USAGE) + bad = [self.entry(agent, base, served, SUBAGENT_USAGE) for agent, base, served in entries] + + assert self.price(session, self.write(tmp_path, [good, *bad]))["savings"] is None + + @pytest.mark.parametrize( + "bad", + ["nope", 7, {"baseline": OPUS_ID, "served": SONNET_ID}, {"baseline": OPUS_ID}], + ) + def test_a_malformed_entry_voids_the_estimate(self, session, tmp_path, bad): + good = self.entry("a0", OPUS_ID, SONNET_ID, SUBAGENT_USAGE) + + assert self.price(session, self.write(tmp_path, [good, bad]))["savings"] is None + + def test_without_prices_there_is_no_estimate_but_the_plugin_still_shows( + self, session, tmp_path, claude_config_dir + ): + write_plugin_version(claude_config_dir, "0.4.4") + session.price_cache.unlink() + entries = [self.entry("a1", OPUS_ID, SONNET_ID, SUBAGENT_USAGE)] + + assert self.price(session, self.write(tmp_path, entries)) == { + "savings": None, + "plugin": "plugin v0.4.4", + } + + @pytest.mark.parametrize( + "text", + [None, "not json", "[]", json.dumps({"version": 2, "entries": []}), '{"version": 1}'], + ) + def test_an_unusable_file_only_nulls_the_savings( + self, session, tmp_path, claude_config_dir, text + ): + write_plugin_version(claude_config_dir, "0.4.4") + usage = tmp_path / "mod-usage.json" + if text is not None: + usage.write_text(text, encoding="utf-8") + + assert self.price(session, usage) == {"savings": None, "plugin": "plugin v0.4.4"} + + def test_a_failing_plugin_lookup_only_nulls_the_plugin(self, session, tmp_path, monkeypatch): + def fail() -> str: + raise RuntimeError("plugin state unreadable") + + monkeypatch.setattr(claude_statusline, "orchestrator_plugin_version", fail) + entries = [self.entry("a1", OPUS_ID, SONNET_ID, SUBAGENT_USAGE)] + + assert self.price(session, self.write(tmp_path, entries)) == { + "savings": "๐Ÿ’ฐ Est. saved $0.18 (60%)", + "plugin": None, + } + + def run_cli(self, *args: str) -> subprocess.CompletedProcess: + return subprocess.run( + [sys.executable, "-P", "-m", claude_statusline.MODULE, *args], + capture_output=True, + timeout=30, + check=False, + ) + + def test_cli_prints_exactly_one_utf8_json_line(self, session, tmp_path): + entries = [self.entry("a1", OPUS_ID, SONNET_ID, SUBAGENT_USAGE)] + usage = self.write(tmp_path, entries) + + result = self.run_cli("--mod-usage", str(usage), "--price-cache", str(session.price_cache)) + + assert result.returncode == 0 + assert result.stdout.count(b"\n") == 1 + assert json.loads(result.stdout.decode("utf-8")) == { + "savings": "๐Ÿ’ฐ Est. saved $0.18 (60%)", + "plugin": None, + } + + def test_cli_never_fails_the_mod(self, session, tmp_path): + result = self.run_cli( + "--mod-usage", + str(tmp_path / "missing.json"), + "--price-cache", + str(session.price_cache), + "--baseline-session-start", + ) + + assert result.returncode == 0 + assert json.loads(result.stdout) == {"savings": None, "plugin": None} + + def test_state_dir_and_mod_usage_are_mutually_exclusive_and_one_is_required( + self, session, tmp_path + ): + cache = ["--price-cache", str(session.price_cache)] + + both = self.run_cli("--state-dir", str(tmp_path), "--mod-usage", str(tmp_path), *cache) + neither = self.run_cli(*cache) + + assert both.returncode == neither.returncode == 2 + + +class TestSavingsText: + def test_rounds_to_cents_and_whole_percent(self): + assert ( + claude_statusline._savings_text(Decimal("1234.565"), Decimal("2469.13")) + == "๐Ÿ’ฐ Est. saved with smart routing: $1,234.57 (50%)" + ) + + def test_negative_reads_as_a_cost_increase(self): + assert ( + claude_statusline._savings_text(Decimal("-0.04"), Decimal("0.40")) + == "Smart routing cost $0.04 more (10%)" + ) + + def test_sub_cent_amounts_are_not_shown_as_zero(self): + assert claude_statusline._savings_figures(Decimal("0.004"), Decimal("1")) == ( + "<$0.01", + Decimal(0), + ) + + def test_the_mod_segments_share_the_rounding(self): + assert ( + claude_statusline._mod_savings_text(Decimal("1234.565"), Decimal("2469.13")) + == "๐Ÿ’ฐ Est. saved $1,234.57 (50%)" + ) + assert ( + claude_statusline._mod_savings_text(Decimal("-0.04"), Decimal("0.40")) + == "cost $0.04 more (10%)" + ) + + +class TestOrchestratorPluginVersion: + def test_reads_the_installed_version(self, tmp_path): + write_plugin_version(tmp_path, "0.4.4") + + assert claude_statusline.orchestrator_plugin_version(tmp_path) == "0.4.4" + + def test_matches_regardless_of_marketplace_suffix(self, tmp_path): + write_plugin_version(tmp_path, "1.2.3", marketplace="some-other-marketplace") + + assert claude_statusline.orchestrator_plugin_version(tmp_path) == "1.2.3" + + def test_missing_file_is_none(self, tmp_path): + assert claude_statusline.orchestrator_plugin_version(tmp_path) is None + + def test_skips_unknown_version_and_other_plugins(self, tmp_path): + (tmp_path / "plugins").mkdir() + (tmp_path / "plugins" / "installed_plugins.json").write_text( + json.dumps( + { + "version": 2, + "plugins": { + "some-other-plugin@mp": [{"version": "9.9.9"}], + "model-orchestrator@mp": [{"version": "unknown"}], + }, + } + ), + encoding="utf-8", + ) + + assert claude_statusline.orchestrator_plugin_version(tmp_path) is None + + def test_unexpected_shape_is_none(self, tmp_path): + (tmp_path / "plugins").mkdir() + (tmp_path / "plugins" / "installed_plugins.json").write_text("[]", encoding="utf-8") + + assert claude_statusline.orchestrator_plugin_version(tmp_path) is None + + def test_honors_claude_config_dir_env(self, tmp_path, monkeypatch): + write_plugin_version(tmp_path, "0.5.0") + monkeypatch.setenv("CLAUDE_CONFIG_DIR", str(tmp_path)) + + assert claude_statusline.orchestrator_plugin_version() == "0.5.0" + + +class TestStatusLineSetting: + def test_wraps_the_highest_precedence_statusline(self, tmp_path): + user = tmp_path / "home" / "settings.json" + user.parent.mkdir() + user.write_text(json.dumps({"statusLine": {"type": "command", "command": "user"}})) + project = tmp_path / "project" + (project / ".claude").mkdir(parents=True) + (project / ".claude" / "settings.json").write_text( + json.dumps({"statusLine": {"type": "command", "command": "project"}}) + ) + (project / ".claude" / "settings.local.json").write_text( + json.dumps({"statusLine": {"type": "command", "command": "local"}}) + ) + + def effective(settings): + status_line = claude_statusline.effective_status_line( + settings, user_settings_path=user, project_dir=project + ) + return status_line and status_line["command"] + + assert effective({"statusLine": {"type": "command", "command": "caller"}}) == "caller" + assert effective({}) == "local" + (project / ".claude" / "settings.local.json").unlink() + assert effective({}) == "project" + (project / ".claude" / "settings.json").unlink() + assert effective({"statusLine": {"type": "command", "command": " "}}) == "user" + user.unlink() + assert effective({}) is None + + def test_runs_only_the_status_row_without_a_user_statusline(self, tmp_path): + setting = claude_statusline.savings_status_line( + None, + python="/opt/ug python/bin/python", + state_dir=tmp_path / "state", + price_cache=tmp_path / "prices.json", + routing_enabled=True, + baseline_session_start=True, + ) + + assert setting == { + "type": "command", + "command": shlex.join( + [ + "/opt/ug python/bin/python", + "-P", + "-m", + claude_statusline.MODULE, + "--state-dir", + str(tmp_path / "state"), + "--price-cache", + str(tmp_path / "prices.json"), + "--routing-enabled", + "--baseline-session-start", + ] + ), + } + + @pytest.mark.parametrize("first_prompt", [False, True]) + def test_mod_pricer_argv_runs_the_module_on_the_same_python(self, tmp_path, first_prompt): + argv = claude_statusline.mod_pricer_argv( + python="/opt/ug python/bin/python", + price_cache=tmp_path / "prices.json", + baseline_session_start=first_prompt, + ) + + assert argv == [ + "/opt/ug python/bin/python", + "-P", + "-m", + claude_statusline.MODULE, + "--price-cache", + str(tmp_path / "prices.json"), + *(["--baseline-session-start"] if first_prompt else []), + ] + + def test_mod_contract_names(self): + # typescript/claude-mods/smart-routing-savings.ts reads both names. + assert claude_statusline.PRICER_ENV_VAR == "UCODE_SAVINGS_PRICER" + assert claude_statusline.MOD_USAGE_FILENAME == "mod-usage.json" + + def test_routing_disabled_omits_routing_enabled_flag(self, tmp_path): + setting = claude_statusline.savings_status_line( + None, + python="/opt/ug python/bin/python", + state_dir=tmp_path / "state", + price_cache=tmp_path / "prices.json", + routing_enabled=False, + baseline_session_start=False, + ) + + assert "--routing-enabled" not in setting["command"] + assert f"-m {claude_statusline.MODULE}" in setting["command"] + + def test_keeps_the_user_statusline_render_options(self, tmp_path): + setting = claude_statusline.savings_status_line( + {"type": "command", "command": "echo hi", "padding": 2, "refreshInterval": 5}, + python=sys.executable, + state_dir=tmp_path, + price_cache=tmp_path / "prices.json", + routing_enabled=False, + baseline_session_start=False, + ) + + assert setting["padding"] == 2 + assert setting["refreshInterval"] == 5 + assert "echo hi" in setting["command"] + + +@pytest.mark.parametrize("shell", [shell for shell in ("sh", "bash", "zsh") if shutil.which(shell)]) +class TestWrappedCommand: + """Run the generated statusLine command the way Claude Code does, through a real shell.""" + + @staticmethod + def run(shell: str, session: Session, original: str) -> str: + setting = claude_statusline.savings_status_line( + {"type": "command", "command": original}, + python=sys.executable, + state_dir=session.state_dir, + price_cache=session.price_cache, + routing_enabled=True, + baseline_session_start=False, + ) + result = subprocess.run( + [shell, "-c", setting["command"]], + input=session.payload(OPUS), + capture_output=True, + text=True, + encoding="utf-8", + timeout=30, + check=True, + ) + return result.stdout + + def test_prints_the_user_row_then_the_status_row(self, shell, session): + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + # The user's command reads the same payload and prints without a trailing newline. + original = 'sed -n \'s/.*"display_name": *"\\([^"]*\\)".*/\\1/p\' | tr -d \'\\n\'' + output = self.run(shell, session, original) + + assert output == ("Opus 4.8\n๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)\n") + + def test_an_exit_in_the_user_command_does_not_skip_the_row(self, shell, session): + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + output = self.run(shell, session, "printf 'base\\n\\n'; exit 3") + + assert output == ("base\n๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)\n") + + def test_empty_user_output_leaves_only_the_status_row(self, shell, session): + append(session.subagent("a1"), response("msg-sub", SONNET_ID, SUBAGENT_USAGE)) + + assert self.run(shell, session, "true") == ( + "๐Ÿ’ฐ Est. saved with smart routing: $0.18 (60%)\n" + ) + + +class TestStartupCost: + def test_statusline_imports_stay_lightweight(self): + heavy = [ + "rich", + "typer", + "databricks", + "tomlkit", + "urllib.request", + "ucode.config_io", + "ucode.smart_routing.routing", + "ucode.smart_routing.session_env", + ] + result = subprocess.run( + [ + sys.executable, + "-P", + "-c", + f"import json, sys; import {claude_statusline.MODULE}; " + f"print(json.dumps([name for name in {heavy!r} if name in sys.modules]))", + ], + capture_output=True, + text=True, + timeout=30, + check=True, + ) + + assert json.loads(result.stdout) == [] + + +class TestPruneState: + def test_removes_only_stale_files(self, tmp_path): + state_dir = tmp_path / "state" + state_dir.mkdir() + stale = state_dir / "old.json" + fresh = state_dir / "new.json" + leftover = state_dir / ".new.json.abc.tmp" + for path in (stale, fresh, leftover): + path.write_text("{}") + now = fresh.stat().st_mtime + os.utime(stale, (now - claude_statusline.STATE_RETENTION_SECONDS - 1,) * 2) + os.utime(leftover, (now - claude_statusline.STATE_RETENTION_SECONDS - 1,) * 2) + + claude_statusline.prune_state(state_dir, now=now) + + assert sorted(path.name for path in state_dir.iterdir()) == ["new.json"] + + def test_missing_state_dir_is_fine(self, tmp_path): + claude_statusline.prune_state(tmp_path / "missing") diff --git a/tests/test_claude_windows_smart_routing.py b/tests/test_claude_windows_smart_routing.py index b6db0f462..bb5266f0f 100644 --- a/tests/test_claude_windows_smart_routing.py +++ b/tests/test_claude_windows_smart_routing.py @@ -54,6 +54,7 @@ def test_windows_subagent_routing_uses_native_binary_without_unix_imports(tmp_pa monkeypatch.delenv(v2.ENABLE_SUBAGENT_ROUTING_ENV_VAR, raising=False) monkeypatch.setattr(v2, "APP_DIR", tmp_path) monkeypatch.setattr(v2, "install_skill", lambda *_args, **_kwargs: None) + monkeypatch.setattr(v2, "_start_savings_price_refresh", lambda *_args, **_kwargs: None) monkeypatch.setattr(v2, "_model_picker_catalog", lambda: None) monkeypatch.setattr(v2, "get_databricks_token", lambda *_args, **_kwargs: "token") monkeypatch.setattr(v2, "build_auth_token_argv", lambda *_args, **_kwargs: ["ug"]) diff --git a/tests/test_databricks.py b/tests/test_databricks.py index 682c74abb..a58cbe294 100644 --- a/tests/test_databricks.py +++ b/tests/test_databricks.py @@ -4444,3 +4444,58 @@ def test_binary_mode_is_unchanged(self, monkeypatch): script = r"import sys; sys.stdout.buffer.write(b'\xff')" result = db_mod.run([sys.executable, "-c", script], capture_output=True) assert result.stdout == b"\xff" + + +class TestFetchEndpointRates: + def test_requests_system_ai_services_as_query_params(self, monkeypatch): + seen = {} + glm_rate = { + "model_service": "system.ai.glm-5-3", + "costs": [ + {"unit": "USD", "token_costs": [{"token_type": "TOKEN_TYPE_INPUT", "cost": 1.4}]} + ], + } + + def fake_get(url, token, **kwargs): + seen.update(url=url, token=token) + return {"model_service_rates": [glm_rate, "not-a-rate"]}, None + + monkeypatch.setattr(db_mod, "_http_get_json", fake_get) + + rates, reason = db_mod.fetch_endpoint_rates( + WS, "tok", ["system.ai.glm-5-3", "claude-opus-4-8", "system.ai.claude-sonnet-5"] + ) + + assert reason is None + assert rates == [glm_rate] + # Names outside system.ai are dropped (the API rejects them); the rest travel as repeated, + # sorted `databricks_hosted_model_services` query parameters. + assert seen["token"] == "tok" + assert seen["url"] == ( + f"{WS}/api/ai-gateway/v2/endpoint-rates:batchGet" + "?databricks_hosted_model_services=system.ai.claude-sonnet-5" + "&databricks_hosted_model_services=system.ai.glm-5-3" + ) + + def test_skips_the_request_without_system_ai_services(self, monkeypatch): + monkeypatch.setattr( + db_mod, "_http_get_json", lambda *args, **kwargs: pytest.fail("no request expected") + ) + + assert db_mod.fetch_endpoint_rates(WS, "tok", ["claude-opus-4-8"]) == ( + [], + "no system.ai model services to price", + ) + + @pytest.mark.parametrize( + ("response", "expected"), + [ + ((None, "HTTP 404 Not Found"), ([], "HTTP 404 Not Found")), + ((["unexpected"], None), ([], "endpoint-rates returned an unexpected response shape")), + (({}, None), ([], None)), + ], + ) + def test_reports_failures_without_raising(self, monkeypatch, response, expected): + monkeypatch.setattr(db_mod, "_http_get_json", lambda *args, **kwargs: response) + + assert db_mod.fetch_endpoint_rates(WS, "tok", ["system.ai.glm-5-3"]) == expected diff --git a/tests/test_smart_routing_pricing.py b/tests/test_smart_routing_pricing.py new file mode 100644 index 000000000..f3370c7aa --- /dev/null +++ b/tests/test_smart_routing_pricing.py @@ -0,0 +1,300 @@ +"""Tests for per-token model pricing used by the smart-routing savings statusline.""" + +from __future__ import annotations + +import json +from decimal import Decimal + +import pytest + +from ucode.smart_routing import pricing +from ucode.smart_routing.pricing import ModelPrice, TokenUsage + +# Catalog list prices (USD per million tokens) for the two Claude route arms. +OPUS = ModelPrice( + input=Decimal("5"), + output=Decimal("25"), + cache_read=Decimal("0.5"), + cache_write_5m=Decimal("6.25"), + cache_write_1h=Decimal("10"), +) +SONNET_4 = ModelPrice( + input=Decimal("3"), + output=Decimal("15"), + cache_read=Decimal("0.3"), + cache_write_5m=Decimal("3.75"), + long_context_threshold=200_000, + long_context=ModelPrice( + input=Decimal("6"), + output=Decimal("22.5"), + cache_read=Decimal("0.6"), + cache_write_5m=Decimal("7.5"), + ), +) + + +class TestModelKey: + @pytest.mark.parametrize( + "model", + [ + "claude-opus-4-8", + "system.ai.claude-opus-4-8", + "system.ai.claude-opus-4-8[1m]", + "databricks-claude-opus-4-8", + "anthropic-aigw-1a2b3c4d-claude-opus-4-8", + "SYSTEM.AI.Claude-Opus-4-8", + "anthropic.claude-opus-4-8", + "global.anthropic.claude-opus-4-8", + "us.anthropic.claude-opus-4-8-20260101-v1:0", + "claude-opus-4-8@default", + ], + ) + def test_collapses_gateway_spellings(self, model): + assert pricing.model_key(model) == "claude-opus-4-8" + + def test_matches_a_served_bedrock_id_to_the_requested_model(self): + # The gateway reports Haiku responses by their Bedrock id. + assert pricing.model_key("anthropic.claude-haiku-4-5-20251001-v1:0") == pricing.model_key( + "system.ai.claude-haiku-4-5" + ) + + def test_unwraps_non_claude_gateway_models(self): + assert pricing.model_key("anthropic-aigw-73ea02b2-system.ai.glm-5-2") == "glm-5-2" + + @pytest.mark.parametrize( + ("served", "requested"), + [ + ("glm-5.3-flash", "system.ai.glm-5-3-flash"), + ("glm-5-3-colo-on-sp-v1", "system.ai.glm-5-3"), + ("kimi-k3-colo-on-sp-v1", "system.ai.kimi-k3"), + ], + ) + def test_matches_a_served_deployment_id_to_the_requested_model(self, served, requested): + assert pricing.model_key(served) == pricing.model_key(requested) + + def test_keeps_distinct_models_distinct(self): + assert pricing.model_key("system.ai.claude-sonnet-5") != pricing.model_key( + "system.ai.claude-opus-4-8" + ) + + +class TestTokenUsage: + def test_splits_cache_writes_by_ttl(self): + tokens = TokenUsage.from_message_usage( + { + "input_tokens": 2, + "cache_creation_input_tokens": 6509, + "cache_read_input_tokens": 143654, + "output_tokens": 695, + "cache_creation": { + "ephemeral_1h_input_tokens": 6000, + "ephemeral_5m_input_tokens": 9, + }, + } + ) + + # The 500 writes the breakdown doesn't cover bill at the default 5-minute TTL. + assert tokens == TokenUsage( + input=2, cache_write_5m=509, cache_write_1h=6000, cache_read=143654, output=695 + ) + assert tokens.prompt_tokens == 2 + 6509 + 143654 + + def test_treats_missing_breakdown_as_five_minute_writes(self): + tokens = TokenUsage.from_message_usage({"cache_creation_input_tokens": 40}) + + assert tokens == TokenUsage(cache_write_5m=40) + + def test_can_bill_writes_the_breakdown_does_not_cover_as_one_hour_writes(self): + unsplit = {"cache_creation_input_tokens": 40} + partial = { + "cache_creation_input_tokens": 6509, + "cache_creation": {"ephemeral_1h_input_tokens": 6000, "ephemeral_5m_input_tokens": 9}, + } + + assert TokenUsage.from_message_usage(unsplit, uncovered_writes_1h=True) == TokenUsage( + cache_write_1h=40 + ) + # The breakdown's own 5-minute writes stay 5-minute; only the 500 uncovered move to 1-hour. + assert TokenUsage.from_message_usage(partial, uncovered_writes_1h=True) == TokenUsage( + cache_write_5m=9, cache_write_1h=6500 + ) + + @pytest.mark.parametrize("bad", [None, -5, True, "12", 1.5]) + def test_ignores_non_count_values(self, bad): + assert TokenUsage.from_message_usage({"input_tokens": bad, "output_tokens": 3}) == ( + TokenUsage(output=3) + ) + + +class TestTokenCost: + def test_prices_every_token_class_at_its_own_rate(self): + tokens = TokenUsage( + input=1_000, cache_write_5m=2_000, cache_write_1h=3_000, cache_read=100_000, output=500 + ) + + # 1k*5 + 2k*6.25 + 3k*10 + 100k*0.5 + 500*25 = 110,000 micro-dollars. + assert pricing.token_cost(OPUS, tokens) == Decimal("0.11") + + def test_uses_long_context_rates_above_the_threshold(self): + short = TokenUsage(input=100_000, output=1_000) + long = TokenUsage(input=150_000, cache_read=60_000, output=1_000) + + assert pricing.token_cost(SONNET_4, short) == Decimal("0.315") + assert pricing.token_cost(SONNET_4, long) == Decimal("0.9585") + + def test_returns_none_when_a_used_class_has_no_rate(self): + # SONNET_4 publishes no 1-hour write rate. + assert pricing.token_cost(SONNET_4, TokenUsage(cache_write_1h=10)) is None + + def test_missing_rate_for_an_unused_class_is_fine(self): + assert pricing.token_cost(SONNET_4, TokenUsage(input=100_000)) == Decimal("0.3") + + +class TestPriceCache: + def test_round_trips_prices_under_model_keys(self, tmp_path): + path = tmp_path / "cache" / "model-prices.json" + pricing.write_price_cache( + path, {"system.ai.claude-opus-4-8": OPUS, "claude-sonnet-4": SONNET_4}, now=100.0 + ) + + cached = pricing.read_price_cache(path) + + assert cached is not None + prices, fingerprint = cached + assert prices == {"claude-opus-4-8": OPUS, "claude-sonnet-4": SONNET_4} + assert fingerprint == "100.0" + assert [entry.name for entry in path.parent.iterdir()] == ["model-prices.json"] + + def test_fingerprint_changes_on_refresh(self, tmp_path): + path = tmp_path / "model-prices.json" + pricing.write_price_cache(path, {"claude-opus-4-8": OPUS}, now=100.0) + first = pricing.read_price_cache(path) + pricing.write_price_cache(path, {"claude-opus-4-8": OPUS}, now=200.0) + + assert first is not None and pricing.read_price_cache(path) is not None + assert first[1] != pricing.read_price_cache(path)[1] + + @pytest.mark.parametrize( + "content", + [ + "not json", + json.dumps({"version": 999, "models": {"claude-opus-4-8": {"input": "5"}}}), + json.dumps({"version": 1, "models": {}}), + json.dumps({"version": 1, "models": []}), + ], + ) + def test_unusable_cache_reads_as_absent(self, tmp_path, content): + path = tmp_path / "model-prices.json" + path.write_text(content) + + assert pricing.read_price_cache(path) is None + + def test_missing_cache_reads_as_absent(self, tmp_path): + assert pricing.read_price_cache(tmp_path / "missing.json") is None + + +class TestEndpointRates: + @staticmethod + def _rate(service: str, *, usd: dict | None = None, dbu: dict | None = None) -> dict: + """An ``EndpointRate`` entry: costs grouped by unit, each a token_type->cost (per 1M).""" + costs = [] + for unit, table in (("DBU", dbu), ("USD", usd)): + if table is not None: + costs.append( + { + "unit": unit, + "token_costs": [ + {"token_type": token_type, "cost": cost} + for token_type, cost in table.items() + ], + } + ) + return {"model_service": service, "costs": costs} + + def test_reads_dollar_rates_keyed_by_model_service(self): + rates = [ + self._rate( + "system.ai.claude-fable-5", + dbu={"TOKEN_TYPE_INPUT": 142.858}, # priced in USD; the DBU group is ignored + usd={"TOKEN_TYPE_INPUT": 10.00006, "TOKEN_TYPE_OUTPUT": 50.00002}, + ), + self._rate( + "system.ai.glm-5-3", + usd={"TOKEN_TYPE_INPUT": 1.4, "TOKEN_TYPE_OUTPUT": 4.3999998}, + ), + ] + + assert pricing.prices_from_endpoint_rates(rates) == { + "system.ai.claude-fable-5": ModelPrice( + input=Decimal("10.00006"), output=Decimal("50.00002") + ), + "system.ai.glm-5-3": ModelPrice(input=Decimal("1.4"), output=Decimal("4.3999998")), + } + + def test_reads_all_token_classes_including_cache(self): + rates = [ + self._rate( + "system.ai.claude-opus-4-8", + usd={ + "TOKEN_TYPE_INPUT": 5, + "TOKEN_TYPE_OUTPUT": 25, + "TOKEN_TYPE_CACHE_READ": 0.5, + "TOKEN_TYPE_CACHE_CREATION": 6.25, + "TOKEN_TYPE_CACHE_CREATION_1H": 10, + }, + ) + ] + + assert pricing.prices_from_endpoint_rates(rates) == { + "system.ai.claude-opus-4-8": ModelPrice( + input=Decimal("5"), + output=Decimal("25"), + cache_read=Decimal("0.5"), + cache_write_5m=Decimal("6.25"), + cache_write_1h=Decimal("10"), + ) + } + + def test_skips_models_without_a_dollar_group(self): + rates = [ + # No USD group: the org has no DBU-to-dollar conversion, so we can't price it. + self._rate("system.ai.glm-5-3", dbu={"TOKEN_TYPE_INPUT": 20}), + self._rate("system.ai.kimi-k3", usd={}), # dollar group present but empty + {"model_service": "system.ai.no-costs"}, # no costs field + { + "costs": [ + {"unit": "USD", "token_costs": [{"token_type": "TOKEN_TYPE_INPUT", "cost": 1}]} + ] + }, + "not-a-rate", + ] + + assert pricing.prices_from_endpoint_rates(rates) == {} + + def test_endpoint_prices_survive_the_cache(self, tmp_path): + path = tmp_path / "model-prices.json" + prices = pricing.prices_from_endpoint_rates( + [ + self._rate( + "system.ai.claude-haiku-4-5", + usd={"TOKEN_TYPE_INPUT": 1, "TOKEN_TYPE_OUTPUT": 5}, + ) + ] + ) + pricing.write_price_cache(path, prices) + + cached = pricing.read_price_cache(path) + + assert cached is not None + served = pricing.model_key("anthropic.claude-haiku-4-5-20251001-v1:0") + assert cached[0][served] == ModelPrice(input=Decimal("1"), output=Decimal("5")) + + +class TestPriceCachePath: + def test_is_stable_per_workspace_and_distinct_across_workspaces(self, tmp_path): + first = pricing.price_cache_path(tmp_path, "https://A.cloud.databricks.com/") + + assert first == pricing.price_cache_path(tmp_path, "https://a.cloud.databricks.com") + assert first != pricing.price_cache_path(tmp_path, "https://b.cloud.databricks.com") + assert first.parent == tmp_path + assert first.name.startswith("model-prices-") and first.suffix == ".json" diff --git a/typescript/claude-mods/register.ts b/typescript/claude-mods/register.ts index 63022c136..a0f15503a 100644 --- a/typescript/claude-mods/register.ts +++ b/typescript/claude-mods/register.ts @@ -1,14 +1,17 @@ import type { Register } from 'claude-code' +import { registerSavings } from './smart-routing-savings' import { registerStatusBand } from './smart-routing-status' import { registerSubagentRouting } from './subagent-routing' // Entry for ug's smart-routing UI mod. Claude Code loads one hooks module per // plugin, so this file imports and composes the concerns: the status band above -// the prompt and the compact per-subagent routing line. Written into the -// launch-scoped routing plugin under ENABLE_CLAUDE_CODE_MODS (see ucode.mods); -// visualization only, ug's Python hooks still route. +// the prompt, the compact per-subagent routing line, and the savings estimate the +// band appends. Written into the launch-scoped routing plugin under +// ENABLE_CLAUDE_CODE_MODS (see ucode.mods); visualization only, ug's Python hooks +// still route. export const register: Register = on => { registerStatusBand(on) registerSubagentRouting(on) + registerSavings(on) } diff --git a/typescript/claude-mods/smart-routing-savings.ts b/typescript/claude-mods/smart-routing-savings.ts new file mode 100644 index 000000000..f12c6bca9 --- /dev/null +++ b/typescript/claude-mods/smart-routing-savings.ts @@ -0,0 +1,290 @@ +import type { Register } from 'claude-code' + +type On = Parameters[0] + +// The savings concern of ug's smart-routing UI mod (composed by register.ts). +// It only collects: per turn request (main agent and subagents) it sums token +// usage by (agent, baseline, served, before_user_switch), where served is the model +// the request named (the model service billed; the id the API reports back is a +// gateway deployment's, which drifts) and baseline is what the request would have +// used without routing: the main model in effect for a subagent, served itself for +// main (routing never changes the main model mid-session). First-prompt routing +// switches the main model once, before the first request, so a later change is +// the user's; requests before it are flagged, and priced from start_model under +// that routing. +// Claude Code's side queries (title, compaction, helpers) are not turn steps, so +// they never count. A hooks module has no Node APIs and the price table lives in +// Python, so after each turn it writes the sums next to UCODE_SESSION_ENV_FILE, +// runs the pricer command ug put in UCODE_SAVINGS_PRICER (a JSON argv), and keeps +// the strings it prints for smart-routing-status.ts to draw. Everything is best +// effort and swallows errors. +// The env var and file names are kept in sync with ucode.smart_routing.claude_statusline +// (a test guards it). + +const USAGE_FILE = 'mod-usage.json' +const PRICER_TIMEOUT_MS = 20_000 +const TOKEN_KEYS = [ + 'input_tokens', + 'output_tokens', + 'cache_creation_input_tokens', + 'cache_read_input_tokens', +] as const +const CACHE_SPLIT_KEYS = ['ephemeral_5m_input_tokens', 'ephemeral_1h_input_tokens'] as const + +type Totals = Record<(typeof TOKEN_KEYS)[number], number> & { + cache_creation?: Partial> +} +type Entry = { + agent: string + baseline: string + served: string + before_user_switch: boolean + usage: Totals +} +type Priced = { savings: string | null; plugin: string | null } + +// Module state is shared with smart-routing-status.ts through savingsSegments(). +const entries = new Map() +// The model the session started on, before any first-prompt routing. Set once. +let startModel: string | null = null +// The model the main loop last named, so a subagent's baseline is a real model id +// ($.session.model() may return an alias). +let lastMainModel: string | null = null +// The model the main loop's first request named (first-prompt routing's pick). +// Set once. +let firstMainModel: string | null = null +// Whether a main request has since named another model, which only the user does. +let userSwitched = false +let seeding: Promise | undefined +let pricer: string[] | null | undefined +let priced: Priced = { savings: null, plugin: null } +let dirty = false +let running = false +let rerun = false + +const num = (value: unknown): number => + typeof value === 'number' && Number.isFinite(value) ? value : 0 + +const str = (value: unknown): string | null => (typeof value === 'string' && value ? value : null) + +// A hooks module has no Node APIs: env and files come through the mods API, and +// $.env.get must be called with a string literal. The pricer argv is only read +// once; absent or malformed means the feature is off. +async function readPricer($: any): Promise { + if (pricer === undefined) { + let found: string[] | null = null + try { + const parsed = JSON.parse(await $.env.get('UCODE_SAVINGS_PRICER')) + if (Array.isArray(parsed) && parsed.length > 0 && parsed.every(arg => typeof arg === 'string')) { + found = parsed + } + } catch {} + pricer = found + } + return pricer +} + +// Same directory derivation as subagent-routing.ts: the per-session dir ug's +// Python side writes next to UCODE_SESSION_ENV_FILE. +async function usagePath($: any): Promise { + let file: unknown + try { + file = await $.env.get('UCODE_SESSION_ENV_FILE') + } catch {} + if (typeof file !== 'string') return null + const cut = Math.max(file.lastIndexOf('/'), file.lastIndexOf('\\')) + return cut > 0 ? `${file.slice(0, cut)}/${USAGE_FILE}` : null +} + +async function mainModel($: any): Promise { + try { + return str(await $.session.model()) + } catch { + return null + } +} + +function accumulate( + agent: string, + baseline: string, + served: string, + beforeUserSwitch: boolean, + usage: any, +): void { + const key = JSON.stringify([agent, baseline, served, beforeUserSwitch]) + let entry = entries.get(key) + if (!entry) { + entry = { + agent, + baseline, + served, + before_user_switch: beforeUserSwitch, + usage: { input_tokens: 0, output_tokens: 0, cache_creation_input_tokens: 0, cache_read_input_tokens: 0 }, + } + entries.set(key, entry) + } + for (const k of TOKEN_KEYS) entry.usage[k] += num(usage[k]) + // The 5m/1h cache-write split is only summed where the API reported it. + const split = usage.cache_creation + if (split && typeof split === 'object') { + const sums = entry.usage.cache_creation ?? {} + for (const k of CACHE_SPLIT_KEYS) if (k in split) sums[k] = (sums[k] ?? 0) + num(split[k]) + entry.usage.cache_creation = sums + } +} + +// A mod reload resets module state while the session carries on, so the first +// hook of a fresh module instance picks up where the last one left off from the +// file it wrote. A missing or malformed file just means starting empty. +async function loadUsage($: any): Promise { + try { + const path = await usagePath($) + if (!path) return + const saved = JSON.parse(await $.fs.read(path)) + if (!saved || saved.version !== 1 || !Array.isArray(saved.entries)) return + for (const item of saved.entries) { + const agent = str(item?.agent) + const baseline = str(item?.baseline) + const served = str(item?.served) + if (agent && baseline && served && item.usage && typeof item.usage === 'object') { + accumulate(agent, baseline, served, item.before_user_switch === true, item.usage) + dirty = true + } + } + startModel ??= str(saved.start_model) + lastMainModel ??= str(saved.main_model) + firstMainModel ??= str(saved.first_main_model) + userSwitched ||= saved.user_switched === true + } catch {} +} + +// Whether the feature is on, with state seeded first (once per module instance, +// whichever hook runs first; concurrent hooks wait on the same load). +async function enable($: any): Promise { + if (!(await readPricer($))) return false + seeding ??= loadUsage($) + await seeding + return true +} + +// What the pricer printed on its last line. A failed run drops the estimate (it +// would be stale) but keeps the plugin version, which does not depend on usage. +async function runPricer($: any, argv: string[], path: string): Promise { + try { + const run = await $.process.run([...argv, '--mod-usage', path], { timeoutMs: PRICER_TIMEOUT_MS }) + if (run?.exitCode === 0) { + const out = JSON.parse(String(run.stdout ?? '').trim().split('\n').pop() ?? '') + if (out && typeof out === 'object') return { savings: str(out.savings), plugin: str(out.plugin) } + } + } catch {} + return { savings: null, plugin: priced.plugin } +} + +async function writeAndPrice($: any): Promise { + const argv = await readPricer($) + const path = await usagePath($) + if (!argv || !path) return + dirty = false + try { + await $.fs.write( + path, + JSON.stringify({ + version: 1, + start_model: startModel, + main_model: lastMainModel, + first_main_model: firstMainModel, + user_switched: userSwitched, + entries: Array.from(entries.values()), + }), + ) + } catch (error) { + // Retry at the next turn.complete rather than pricing a file that is not there. + dirty = true + throw error + } + const next = await runPricer($, argv, path) + if (next.savings !== priced.savings || next.plugin !== priced.plugin) { + priced = next + // Claude Code only redraws the band when its props change, not when module state does. + $.ui.invalidate('ui.render') + } +} + +// One pricer run at a time; a request that arrives mid-run just asks for another. +async function flush($: any): Promise { + if (running) { + rerun = true + return + } + running = true + try { + do { + rerun = false + await writeAndPrice($) + } while (rerun) + } catch { + } finally { + running = false + } +} + +// Strings for smart-routing-status.ts to append after its trailer, in order; +// empty until the pricer has something to show (or when the feature is off). +export const savingsSegments = (): string[] => + [priced.savings, priced.plugin].filter((segment): segment is string => segment !== null) + +export const registerSavings = (on: On): void => { + // This concern owns session.start for the composed module (a second + // unmatched hook on the same event fails the module load). It runs before the + // first prompt, so the start model is captured before any routing happens and + // the plugin version can show at once. session.start can fire again (/clear), + // by when first-prompt routing may have switched the model, so the start model + // is only ever set once. The flush is not awaited: it spawns a process and + // Claude Code waits on this hook before the first prompt. + on('session.start', async ($, e, next) => { + if (await enable($)) { + const model = await mainModel($) + startModel ??= model + void flush($) + } + return next(e) + }) + + // Fires for every request to a model, subagents included (e.agentId is set). + // yield* forwards the streamed response untouched; only the finished result's + // usage is read. + on('turn.step', async function* ($, e, next) { + const request = e as any + const agent = str(request.agentId) + const enabled = await enable($) + let baseline: string | null = null + if (enabled) { + if (agent === null) { + lastMainModel = str(request.model) ?? lastMainModel + firstMainModel ??= lastMainModel + userSwitched ||= lastMainModel !== firstMainModel + } else baseline = lastMainModel ?? (await mainModel($)) + } + // Read before the request runs: a concurrent main request may switch it. + const beforeUserSwitch = !userSwitched + const result = yield* next(e) + try { + const usage = (result as any)?.usage + const served = str(request.model) ?? str(usage?.model) + // Main is its own baseline: routing leaves the main model alone. + const base = agent === null ? served : baseline + if (enabled && usage && served && base) { + accumulate(agent ?? 'main', base, served, beforeUserSwitch, usage) + dirty = true + } + } catch {} + return result + }) + + // Price at the end of each turn, off the turn's critical path. + on('turn.complete', async ($, e, next) => { + const result = await next(e) + if ((await enable($)) && dirty) void flush($) + return result + }) +} diff --git a/typescript/claude-mods/smart-routing-status.ts b/typescript/claude-mods/smart-routing-status.ts index 857ed9041..dc3ddfefd 100644 --- a/typescript/claude-mods/smart-routing-status.ts +++ b/typescript/claude-mods/smart-routing-status.ts @@ -1,5 +1,7 @@ import type { Register } from 'claude-code' +import { savingsSegments } from './smart-routing-savings' + type On = Parameters[0] // The status band concern of ug's smart-routing UI mod (composed by register.ts). @@ -67,11 +69,13 @@ export const registerStatusBand = (on: On): void => { children: ['โ€” ' + (firstPrompt ? 'routing first prompt + subagents' : 'routing subagents')], }) : Text({ dimColor: true, children: ["โ€” Reenable with '/smart-router on'"] }) + // Python-priced savings/plugin strings; empty when off or nothing to show. + const extras = enabled ? savingsSegments().map(s => Text({ dimColor: true, children: ['ยท ' + s] })) : [] const line = Box({ flexDirection: 'row', columnGap: 1, - children: [Text({ bold: true, children: [BANNER] }), pill, trailer], + children: [Text({ bold: true, children: [BANNER] }), pill, trailer, ...extras], }) const children = theirs ? [line, theirs] : [line] return Box({ flexDirection: 'column', children })