diff --git a/application/model_routing/__init__.py b/application/model_routing/__init__.py index 06bee05..2ff9e41 100644 --- a/application/model_routing/__init__.py +++ b/application/model_routing/__init__.py @@ -1 +1,54 @@ -"""Application model routing package: model route decisions and multi-provider balancing.""" +"""Application model routing package: model route decisions and multi-provider balancing. + +Public surface (R03-T03 — the single routing entry point every chat surface uses): + +* :class:`RoutingApplicationService` — decides one turn's provider/model. +* :class:`RoutingRequest` / :class:`RoutingOutcome` — the immutable DTOs in and out. +* :class:`RoutingMode` — Off / Auto / Manual / Fallback. +* :func:`build_routing_application_service` — wires the service to a live + ``AppContext`` (engine + per-workspace mode + confirm timeout). + +Typical call site (see ``ui/chat_panel.py::_apply_routing``):: + + service = build_routing_application_service(self.ctx) + outcome = service.resolve( + RoutingRequest(surface="cowork", prompt=text, + current_provider=provider, current_model=model), + confirm=lambda decision, timeout: confirm_switch(self, decision, timeout), + ) + +Only ``core_routing_adapter`` touches ``core/routing``; the service and the DTOs +stay pure Python so the whole rule set is testable without Qt or the engine. +""" + +from .core_routing_adapter import ( + AppContextModeResolver, + CoreRoutingEngine, + build_routing_application_service, +) +from .routing_application_service import ( + ConfirmationCallback, + ModeResolver, + RoutingApplicationService, + RoutingDecisionPort, +) +from .routing_models import ( + RouteEvaluation, + RoutingMode, + RoutingOutcome, + RoutingRequest, +) + +__all__ = [ + "AppContextModeResolver", + "ConfirmationCallback", + "CoreRoutingEngine", + "ModeResolver", + "RouteEvaluation", + "RoutingApplicationService", + "RoutingDecisionPort", + "RoutingMode", + "RoutingOutcome", + "RoutingRequest", + "build_routing_application_service", +] diff --git a/application/model_routing/core_routing_adapter.py b/application/model_routing/core_routing_adapter.py new file mode 100644 index 0000000..f2fdfa8 --- /dev/null +++ b/application/model_routing/core_routing_adapter.py @@ -0,0 +1,169 @@ +"""Adapters that plug the existing routing engine into the application service. + +:mod:`routing_application_service` is written against two narrow ports so it can +be unit-tested with plain fakes. This module supplies the real implementations — +the assessment/scoring engine in ``core/routing`` and the per-workspace mode +lookup on ``AppContext`` — and is therefore the ONLY file in +``application/model_routing/`` that knows those concrete types exist. + +All engine imports are deferred into method bodies. Importing the routing stack +pulls in Pydantic models and the on-disk assessment store, and the UI must be +able to import this module during startup without paying that cost (the same +lazy-wiring reason ``state.py::AppContext.routing`` gives). +""" + +from __future__ import annotations + +import logging +from typing import Any, Optional + +from .routing_application_service import RoutingApplicationService +from .routing_models import RouteEvaluation, RoutingMode, RoutingRequest + +logger = logging.getLogger("cowork_local.application.model_routing") + + +class CoreRoutingEngine: + """:class:`RoutingDecisionPort` backed by ``core/routing/service.py``. + + Translates in both directions: application DTOs in, and the engine's + ``RouteResult``/``SwitchDecision``/``TaskType`` flattened back out into a + :class:`RouteEvaluation`, so no ``core.routing`` type ever escapes into the + application service or the UI call sites. + """ + + def __init__(self, routing_service: Any) -> None: + self._routing_service = routing_service + + def evaluate(self, request: RoutingRequest, mode: RoutingMode) -> RouteEvaluation: + """Rank candidates for this turn and report the engine's verdict.""" + from ...core.routing.models import TaskType, candidate_key + + result = self._routing_service.route( + request.surface, + request.prompt, + request.current_provider, + request.current_model, + # The engine only knows off/auto/manual; FALLBACK was already mapped + # to AUTO upstream so the value handed over here is always valid. + mode_override=mode.value, + required_capabilities=list(request.required_capabilities) or None, + task_type=self._parse_task_type(request.task_type, TaskType), + ) + + decision = result.decision + target = result.target() # (provider, model_id) or None + current_key = ( + candidate_key(request.current_provider, request.current_model) + if request.current_model + else "" + ) + return RouteEvaluation( + task_type=self._task_type_value(result.task_type), + should_switch=bool(result.should_switch), + target_provider=target[0] if target else None, + target_model=target[1] if target else None, + score_gain=float(getattr(decision, "score_gain", 0.0) or 0.0), + reason=str(getattr(decision, "reason", "") or ""), + current_is_usable=self._current_is_usable(result, current_key), + decision=decision, + ) + + # -- translation helpers --------------------------------------------- # + @staticmethod + def _parse_task_type(raw: Optional[str], task_type_enum) -> Optional[Any]: + """Coerce a task-type string to the engine's enum. + + ``None`` (the common case) means "let the engine classify the prompt". + An unrecognised string is also downgraded to ``None`` rather than + raising, so a stale value in a saved workspace cannot break a turn. + """ + if raw is None: + return None + if isinstance(raw, task_type_enum): + return raw + try: + return task_type_enum(str(raw).strip().lower()) + except ValueError: + logger.warning("routing: unknown task type %r — classifying from the prompt", raw) + return None + + @staticmethod + def _task_type_value(task_type: Any) -> str: + """The plain string form of the engine's task type enum.""" + return str(getattr(task_type, "value", task_type) or "") + + @staticmethod + def _current_is_usable(result: Any, current_key: str) -> bool: + """Whether the currently selected model can still serve this task. + + This is the signal FALLBACK mode acts on. A model is usable when the + ranking scored it above zero; ``rank_models`` already drops candidates + that are unavailable, lack a probe for this task type, or failed their + last probe, so "absent from the ranking" is precisely "cannot serve it". + + With no ranking (routing off, or the engine's internal error path) or no + current model, we answer True: absence of evidence must not trigger a + surprise switch in a mode whose whole promise is not to surprise. + """ + ranking = getattr(result, "ranking", None) + if ranking is None or not current_key: + return True + try: + return float(ranking.score_of(current_key)) > 0.0 + except Exception: # noqa: BLE001 — defensive: never fail a turn on telemetry-ish data + logger.debug("routing: could not score current model %r", current_key, exc_info=True) + return True + + +class AppContextModeResolver: + """:class:`ModeResolver` backed by the active workspace's settings. + + Reads through ``AppContext.project_routing_mode``, which already layers the + workspace override on top of the global default — so per-workspace routing + modes keep working unchanged now that the mode lookup moved out of the + widgets. + """ + + def __init__(self, ctx: Any) -> None: + self._ctx = ctx + + def mode_for(self, surface: str) -> RoutingMode: + """Effective mode for ``surface`` in the active workspace.""" + return RoutingMode.parse(self._ctx.project_routing_mode(surface)) + + +def build_routing_application_service(ctx: Any) -> RoutingApplicationService: + """The shared :class:`RoutingApplicationService` for this app context. + + Cached on the context (like ``AppContext.routing()`` caches the engine) so + every surface talks to the same instance and a future stateful addition — + per-surface cool-down, switch history — is shared rather than duplicated per + widget. Falls back to a fresh instance if the context refuses attribute + assignment, which keeps tests using lightweight stand-ins working. + """ + cached = getattr(ctx, "_routing_app_service", None) + if cached is not None: + return cached + + service = RoutingApplicationService( + CoreRoutingEngine(ctx.routing()), + AppContextModeResolver(ctx), + # Read at call time: the user can change the confirm timeout in Settings + # between two turns and the next Manual dialog should honour it. + confirm_timeout_sec=lambda: float( + (ctx.config.routing or {}).get("confirm_timeout_sec", 60) or 60 + ), + ) + try: + ctx._routing_app_service = service + except Exception: # noqa: BLE001 — read-only/slotted stand-ins stay supported + logger.debug("routing: could not cache the application service on the context", exc_info=True) + return service + + +__all__ = [ + "AppContextModeResolver", + "CoreRoutingEngine", + "build_routing_application_service", +] diff --git a/application/model_routing/routing_application_service.py b/application/model_routing/routing_application_service.py new file mode 100644 index 0000000..9faf703 --- /dev/null +++ b/application/model_routing/routing_application_service.py @@ -0,0 +1,236 @@ +"""The one place that decides how a turn is routed (R03-T03). + +Before this service, ``ui/chat_panel.py#L638``, ``ui/co4e_tab.py`` and +``ui/folder_tab.py`` each carried their own copy of the same eight-step dance: +clear last turn's override → read the surface's mode → bail on "off" → call the +routing engine → check ``should_switch`` → resolve the target → show the Manual +confirm dialog → publish the override and a status line. Three copies meant +three chances to drift, and none of them could be tested without a Qt widget. + +The dance now lives here, once, in pure Python: + +* the routing engine is reached through :class:`RoutingDecisionPort`; +* the surface's Off/Auto/Manual/Fallback mode through :class:`ModeResolver`; +* the Manual-mode confirmation through a ``confirm`` callback supplied per call, + so the Qt dialog stays in the presentation layer where it belongs. + +Every failure path degrades to "keep the current model": a routing problem must +never be the reason a user cannot send a message. +""" + +from __future__ import annotations + +import logging +from typing import Any, Callable, Optional, Protocol, runtime_checkable + +from .routing_models import ( + RouteEvaluation, + RoutingMode, + RoutingOutcome, + RoutingRequest, +) + +logger = logging.getLogger("cowork_local.application.model_routing") + +# Asks the user to approve a Manual-mode switch. Receives the underlying +# decision object (for rendering) plus the timeout in seconds; returns True to +# approve. Supplied by the caller so this module never imports a UI toolkit. +ConfirmationCallback = Callable[[Any, float], bool] + + +@runtime_checkable +class RoutingDecisionPort(Protocol): + """The routing engine, as this service needs it. + + Narrowed to a single method on purpose: the concrete engine + (``core/routing/service.py::RoutingService``) exposes assessment, + persistence and scheduling too, none of which a turn-time decision needs. + """ + + def evaluate(self, request: RoutingRequest, mode: RoutingMode) -> RouteEvaluation: + """Rank candidates for ``request`` and report whether to switch.""" + + +@runtime_checkable +class ModeResolver(Protocol): + """Resolves the effective routing mode for a surface. + + In the app this reads the active workspace's per-surface override with the + global default behind it (``AppContext.project_routing_mode``); in tests it + is a two-line stub. + """ + + def mode_for(self, surface: str) -> RoutingMode: + """Effective mode for ``surface``.""" + + +class RoutingApplicationService: + """Turn-time routing decisions for every chat surface.""" + + # Matches DEFAULT_CONFIG["routing"]["confirm_timeout_sec"]; used only when + # no timeout provider is wired, so a bare service is still usable in tests. + DEFAULT_CONFIRM_TIMEOUT_SEC = 60.0 + + def __init__( + self, + decision_port: RoutingDecisionPort, + mode_resolver: Optional[ModeResolver] = None, + *, + confirm_timeout_sec: Optional[Callable[[], float]] = None, + ) -> None: + self._decision_port = decision_port + self._mode_resolver = mode_resolver + # A callable rather than a number: the timeout lives in mutable config + # the user can change in Settings between two turns. + self._confirm_timeout_sec = confirm_timeout_sec + + # -- public API ------------------------------------------------------ # + def resolve( + self, + request: RoutingRequest, + confirm: Optional[ConfirmationCallback] = None, + ) -> RoutingOutcome: + """Decide this turn's provider/model. + + Returns a :class:`RoutingOutcome`; ``provider``/``model`` are ``None`` + whenever the surface should keep its own selection. Never raises — an + unexpected failure is logged and reported as "keep current", because a + broken assessment store must not block chatting. + """ + mode = request.mode or self._resolve_mode(request.surface) + try: + return self._resolve_unguarded(request, mode, confirm) + except Exception: # noqa: BLE001 — routing must never break a turn + logger.exception("routing.resolve failed — keeping the current model") + return RoutingOutcome.keep_current(mode, reason="routing error — keeping current model") + + def confirm_timeout(self) -> float: + """Seconds to wait for a Manual-mode confirmation. + + Falls back to the built-in default when the provider is missing or + returns something unusable, so a corrupted config value cannot produce a + zero-second dialog that instantly declines every switch. + """ + if self._confirm_timeout_sec is None: + return self.DEFAULT_CONFIRM_TIMEOUT_SEC + try: + value = float(self._confirm_timeout_sec()) + except (TypeError, ValueError): + return self.DEFAULT_CONFIRM_TIMEOUT_SEC + return value if value > 0 else self.DEFAULT_CONFIRM_TIMEOUT_SEC + + # -- internals ------------------------------------------------------- # + def _resolve_mode(self, surface: str) -> RoutingMode: + """The surface's configured mode, defaulting to OFF when unresolvable — + routing stays opt-in, so "we don't know" must mean "don't switch".""" + if self._mode_resolver is None: + return RoutingMode.OFF + try: + return RoutingMode.parse(self._mode_resolver.mode_for(surface)) + except Exception: # noqa: BLE001 — a config read must not break a turn + logger.exception("routing: could not resolve mode for surface %r", surface) + return RoutingMode.OFF + + def _resolve_unguarded( + self, + request: RoutingRequest, + mode: RoutingMode, + confirm: Optional[ConfirmationCallback], + ) -> RoutingOutcome: + """The decision flow proper; :meth:`resolve` owns the safety net.""" + # 1. Routing disabled, or nothing to classify -> keep the selection. + if mode is RoutingMode.OFF: + return RoutingOutcome.keep_current(mode, reason="routing off") + if not request.has_prompt: + return RoutingOutcome.keep_current(mode, reason="empty prompt — nothing to route") + + # 2. Ask the engine. FALLBACK is evaluated with AUTO's ranking because + # it needs the same candidate list; only the accept/reject rule below + # differs, so the engine stays unaware of the extra mode. + engine_mode = RoutingMode.AUTO if mode is RoutingMode.FALLBACK else mode + evaluation = self._decision_port.evaluate(request, engine_mode) + + # 3. Apply the mode's own accept rule to the engine's verdict. + if mode is RoutingMode.FALLBACK: + accepted, reason = self._fallback_verdict(evaluation) + else: + accepted, reason = evaluation.should_switch, evaluation.reason + + if not accepted or not evaluation.has_target: + return RoutingOutcome.keep_current( + mode, + reason=reason or evaluation.reason, + task_type=evaluation.task_type, + decision=evaluation.decision, + ) + + # 4. Manual mode asks first; a decline or a timeout keeps the current + # model (and is reported as such, so the surface can tell the two + # cases apart from "nothing better was found"). + if mode is RoutingMode.MANUAL and not self._approved(evaluation, confirm): + return RoutingOutcome.keep_current( + mode, + reason="switch declined by user or confirmation timed out", + task_type=evaluation.task_type, + declined=True, + decision=evaluation.decision, + ) + + # 5. Publish the override for THIS turn only. The provider falls back to + # the request's current provider when the engine named a model but no + # provider (same-provider switch). + return RoutingOutcome( + mode=mode, + switched=True, + provider=evaluation.target_provider or request.current_provider, + model=evaluation.target_model or "", + task_type=evaluation.task_type, + score_gain=evaluation.score_gain, + reason=reason or evaluation.reason, + decision=evaluation.decision, + ) + + @staticmethod + def _fallback_verdict(evaluation: RouteEvaluation) -> tuple: + """FALLBACK's accept rule: switch ONLY to rescue an unusable selection. + + The user's pinned model wins as long as it can serve the turn, even when + a higher-scoring candidate exists — that is the whole point of the mode. + A switch happens only when the current model is not a usable candidate + (never assessed, marked unavailable, or its last probe failed) and the + engine has something to move to. + """ + if evaluation.current_is_usable: + return False, "fallback mode — current model is healthy, keeping it" + if not evaluation.has_target: + return False, "fallback mode — current model unusable and no replacement available" + return True, "fallback mode — current model unavailable, switching to the best alternative" + + def _approved( + self, + evaluation: RouteEvaluation, + confirm: Optional[ConfirmationCallback], + ) -> bool: + """Run the Manual-mode confirmation callback. + + No callback means no way to ask, and silently switching in Manual mode + would violate the mode's contract — so a missing callback is treated as + "not approved". A callback that raises is treated the same way, since a + broken dialog must not auto-approve a model change. + """ + if confirm is None: + logger.warning("routing: manual mode without a confirmation callback — keeping current model") + return False + try: + return bool(confirm(evaluation.decision, self.confirm_timeout())) + except Exception: # noqa: BLE001 + logger.exception("routing: confirmation callback failed — keeping current model") + return False + + +__all__ = [ + "ConfirmationCallback", + "ModeResolver", + "RoutingApplicationService", + "RoutingDecisionPort", +] diff --git a/application/model_routing/routing_models.py b/application/model_routing/routing_models.py new file mode 100644 index 0000000..8f9808c --- /dev/null +++ b/application/model_routing/routing_models.py @@ -0,0 +1,158 @@ +"""Pure-Python DTOs exchanged with :mod:`routing_application_service`. + +These types are the vocabulary the chat surfaces (Cowork chat, Co4E, AI-Edit) +now speak instead of each re-deriving routing state from raw config lookups and +``core/routing`` internals. + +Layer rules (``docs/architecture/ADR-001-layered-architecture.md``): application +code is 100% pure Python. Nothing here imports PySide6, and nothing here imports +``core.routing`` either — the concrete routing engine is reached only through +the adapter in :mod:`core_routing_adapter`, which keeps this module trivially +testable with plain fakes. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Optional, Tuple + + +class RoutingMode(str, Enum): + """The four routing behaviours a surface can be in (R03-T03). + + ``OFF``/``AUTO``/``MANUAL`` map 1:1 onto the existing per-surface toggle and + onto ``core/routing/models.py::SwitchMode``. ``FALLBACK`` is new and + deliberately NOT an optimisation mode: it keeps whatever model the user + chose and only re-routes when that model cannot serve the turn, which is the + behaviour a resilience-minded workspace wants (never surprise me, but never + leave me stuck either). + """ + + OFF = "off" + AUTO = "auto" + MANUAL = "manual" + FALLBACK = "fallback" + + @classmethod + def parse(cls, raw: Any, default: "RoutingMode" = None) -> "RoutingMode": + """Best-effort coercion from config/UI strings. + + Routing must never break a turn, so an unrecognised value degrades to + ``default`` (``OFF`` unless told otherwise) instead of raising — the same + defensive posture ``config.routing_mode_for`` already takes. + """ + fallback = default if default is not None else cls.OFF + if isinstance(raw, cls): + return raw + try: + return cls(str(raw or "").strip().lower()) + except ValueError: + return fallback + + +@dataclass(frozen=True) +class RoutingRequest: + """Everything needed to decide how ONE turn should be routed. + + Frozen: the request is captured from live UI state (the selected model, the + typed prompt) and then handed to code that may run on a worker thread. An + immutable snapshot means the user changing the model picker mid-turn cannot + retroactively alter the decision that was already made — the same rationale + behind R04's ``ConversationExecutionRequest``. + """ + + surface: str # "cowork" | "co4e" | "ai_edit" | ... + prompt: str # the user's text; drives task classification + current_provider: str # provider the surface would use as-is + current_model: str = "" # model the surface would use ("" = provider default) + mode: Optional[RoutingMode] = None # explicit override; None -> resolve per surface + # Pre-classified task type ("coding", "qa", ...). AI-Edit always knows its + # turns are coding work, so it pins this and skips prompt classification. + task_type: Optional[str] = None + required_capabilities: Tuple[str, ...] = () # e.g. ("vision",) + + @property + def has_prompt(self) -> bool: + """Whether there is anything to classify. An empty prompt cannot be + routed meaningfully, so every surface short-circuits on it.""" + return bool((self.prompt or "").strip()) + + +@dataclass(frozen=True) +class RouteEvaluation: + """A routing engine's verdict, normalised away from ``core/routing`` types. + + The adapter flattens ``RouteResult``/``SwitchDecision`` into these plain + fields so the application service never touches Pydantic models or enums + owned by another layer. ``decision`` still carries the original object + because the Manual-mode confirm dialog renders its ``reason``. + """ + + task_type: str + should_switch: bool + target_provider: Optional[str] = None + target_model: Optional[str] = None + score_gain: float = 0.0 + reason: str = "" + # False when the currently selected model is not a usable candidate for this + # task (unranked, unavailable, or failed its last probe) — the single signal + # FALLBACK mode acts on. + current_is_usable: bool = True + decision: Any = None # original SwitchDecision, for the UI dialog + + @property + def has_target(self) -> bool: + """A switch is only actionable when the engine named a model to move to.""" + return bool(self.target_model or self.target_provider) + + +@dataclass(frozen=True) +class RoutingOutcome: + """What the calling surface should actually do for this turn. + + A surface needs exactly three things from routing — "which provider/model do + I build?", "do I tell the user?" and "was I told to stand down?" — so those + are the fields here, and nothing else. ``provider``/``model`` are ``None`` + when the surface should keep its own selection untouched. + """ + + mode: RoutingMode + switched: bool = False + provider: Optional[str] = None + model: Optional[str] = None + task_type: str = "" + score_gain: float = 0.0 + reason: str = "" + # True when Manual mode proposed a switch and the user declined or the + # confirmation timed out. Distinct from "no switch proposed" so a surface + # can tell "routing had nothing to offer" from "the user said no". + declined: bool = False + decision: Any = field(default=None, repr=False) + + @property + def should_notify(self) -> bool: + """Whether the surface should post the "switched model" status bubble. + Only an executed switch is worth interrupting the transcript for.""" + return self.switched + + @classmethod + def keep_current( + cls, + mode: RoutingMode, + *, + reason: str = "", + task_type: str = "", + declined: bool = False, + decision: Any = None, + ) -> "RoutingOutcome": + """The no-change outcome — the single constructor for every path that + leaves the surface's own model selection in place (routing off, empty + prompt, no better candidate, user declined, internal error).""" + return cls( + mode=mode, switched=False, provider=None, model=None, + task_type=task_type, reason=reason, declined=declined, decision=decision, + ) + + +__all__ = ["RoutingMode", "RoutingRequest", "RouteEvaluation", "RoutingOutcome"] diff --git a/config.py b/config.py index ba07910..6c96a4f 100644 --- a/config.py +++ b/config.py @@ -555,19 +555,26 @@ class AppConfig: d["surface_modes"].setdefault(surface, "") return d - def routing_mode_for(self, surface: str) -> str: - """Effective Off/Auto/Manual mode for a chat surface. + # The routing modes a surface may be in. "fallback" joined the set in + # R03-T03 (keep the selected model; re-route only when it cannot serve the + # turn) — see application/model_routing/routing_models.py::RoutingMode, + # which is the authority on what each mode means. + ROUTING_MODES = ("off", "auto", "manual", "fallback") - A per-surface override ("auto"/"manual"/"off") wins; an empty override - falls back to the global ``switch_mode``.""" + def routing_mode_for(self, surface: str) -> str: + """Effective Off/Auto/Manual/Fallback mode for a chat surface. + + A per-surface override wins; an empty override falls back to the global + ``switch_mode``. Anything unrecognised degrades to "off" so routing + stays opt-in even with a hand-edited config.""" routing = self.routing override = (routing.get("surface_modes", {}) or {}).get(surface, "") mode = override or routing.get("switch_mode", "off") - return mode if mode in ("off", "auto", "manual") else "off" + return mode if mode in self.ROUTING_MODES else "off" def set_routing_mode_for(self, surface: str, mode: str) -> None: - """Persist a chat surface's Off/Auto/Manual toggle selection.""" - mode = mode if mode in ("off", "auto", "manual") else "off" + """Persist a chat surface's routing toggle selection.""" + mode = mode if mode in self.ROUTING_MODES else "off" self.routing.setdefault("surface_modes", {})[surface] = mode self.save() diff --git a/core/usage_tracker.py b/core/usage_tracker.py index f1c6050..0f5ad3d 100644 --- a/core/usage_tracker.py +++ b/core/usage_tracker.py @@ -49,6 +49,18 @@ def set_context(source: str, label: str = "") -> None: _local.label = label +def current_context() -> tuple: + """The ``(source, label)`` currently tagged on THIS thread. + + Public counterpart to :func:`set_context`, added for + ``infrastructure/telemetry/usage_sink.py``: a subscriber that needs to + attribute one event to a different surface must be able to save the + caller's context and put it back afterwards, instead of leaving the worker + thread permanently retagged. + """ + return getattr(_local, "source", "") or "", getattr(_local, "label", "") or "" + + # ---- per-thread usage accumulator ----------------------------------------- # A step/run that wants to know its OWN token/cost (not the all-time file total) # calls begin_accumulation(), reads accumulated() before/after a unit of work, diff --git a/docs/refactor/Refactoring_Checklist.md b/docs/refactor/Refactoring_Checklist.md index c1bfd30..c857c66 100644 --- a/docs/refactor/Refactoring_Checklist.md +++ b/docs/refactor/Refactoring_Checklist.md @@ -62,18 +62,72 @@ * **Team chịu trách nhiệm**: 🔵 **Team Duy** (Chủ trì) * **Mục tiêu**: Hợp nhất logic routing bị phân tán thành `RoutingApplicationService` độc lập Qt; chuẩn hóa catalog nhà cung cấp. -- [ ] **R03-T01 (Team Duy)**: Xây dựng bộ Contract Tests chuẩn hóa cho các Provider từ `providers/base.py` ➔ `tests/contracts/test_providers.py` - *Start: `____-__-__ __:__` | End: `____-__-__ __:__`* -- [ ] **R03-T02 (Team Duy)**: Xây dựng `ProviderDescriptor` và `ProviderRegistry` tập trung từ `providers/factory.py` ➔ `domain/models/provider_descriptor.py` & `infrastructure/providers/provider_registry.py` - *Start: `____-__-__ __:__` | End: `____-__-__ __:__`* -- [ ] **R03-T03 (Team Duy)**: Xây dựng `RoutingApplicationService` độc lập với Qt từ `core/routing/` ➔ `application/model_routing/routing_application_service.py` - *Start: `____-__-__ __:__` | End: `____-__-__ __:__`* -- [ ] **R03-T04 (Team Duy)**: Di chuyển luồng gọi routing từ `ui/chat_panel.py#L638` sang `RoutingApplicationService` - *Start: `____-__-__ __:__` | End: `____-__-__ __:__`* -- [ ] **R03-T05 (Team Duy)**: Di chuyển luồng gọi routing từ `ui/co4e_tab.py` và `ui/folder_tab.py` sang `RoutingApplicationService` - *Start: `____-__-__ __:__` | End: `____-__-__ __:__`* -- [ ] **R03-T06 (Team Duy)**: Tách logic ghi nhận token usage ra khỏi Provider, chuyển thành `UsageEventSink` ➔ `infrastructure/telemetry/usage_sink.py` - *Start: `____-__-__ __:__` | End: `____-__-__ __:__`* +- [x] **R03-T01 (Team Duy)**: Xây dựng bộ Contract Tests chuẩn hóa cho các Provider từ `providers/base.py` ➔ `tests/contracts/test_providers.py` + *Start: `2026-08-22 18:59` | End: `2026-08-22 19:01`* +- [x] **R03-T02 (Team Duy)**: Xây dựng `ProviderDescriptor` và `ProviderRegistry` tập trung từ `providers/factory.py` ➔ `domain/models/provider_descriptor.py` & `infrastructure/providers/provider_registry.py` + *Start: `2026-08-22 18:45` | End: `2026-08-22 18:50`* +- [x] **R03-T03 (Team Duy)**: Xây dựng `RoutingApplicationService` độc lập với Qt từ `core/routing/` ➔ `application/model_routing/routing_application_service.py` + *Start: `2026-08-22 18:53` | End: `2026-08-22 18:57`* +- [x] **R03-T04 (Team Duy)**: Di chuyển luồng gọi routing từ `ui/chat_panel.py#L638` sang `RoutingApplicationService` + *Start: `2026-08-22 18:57` | End: `2026-08-22 18:58`* +- [x] **R03-T05 (Team Duy)**: Di chuyển luồng gọi routing từ `ui/co4e_tab.py` và `ui/folder_tab.py` sang `RoutingApplicationService` + *Start: `2026-08-22 18:58` | End: `2026-08-22 18:59`* +- [x] **R03-T06 (Team Duy)**: Tách logic ghi nhận token usage ra khỏi Provider, chuyển thành `UsageEventSink` ➔ `infrastructure/telemetry/usage_sink.py` + *Start: `2026-08-22 18:50` | End: `2026-08-22 18:53`* + +#### 📦 KẾT QUẢ THỰC HIỆN EPIC R03 (Hoàn tất 2026-08-22 19:01 — nhánh `feature/delta-team/epic-R03`) + +**File sản phẩm mới (tất cả < 400 dòng, 100% comment tiếng Anh):** + +| Task | File | LOC | Nội dung chính | +| :--- | :--- | :---: | :--- | +| T02 | `domain/models/provider_descriptor.py` | 196 | `ProviderDescriptor` (frozen dataclass), `WireProtocol`, `AuthKind`; giá/context để `None` khi chưa biết thay vì đoán bừa | +| T02 | `infrastructure/providers/provider_registry.py` | 287 | `ProviderRegistry` thread-safe: tra cứu theo id/alias, **tra cứu động theo model ID** (`find_by_model`), dựng adapter theo wire protocol; `BUILTIN_DESCRIPTORS` cho 5 provider | +| T03 | `application/model_routing/routing_models.py` | 158 | DTO thuần Python: `RoutingMode` (Off/Auto/Manual/**Fallback**), `RoutingRequest` (immutable snapshot), `RouteEvaluation`, `RoutingOutcome` | +| T03 | `application/model_routing/routing_application_service.py` | 236 | `RoutingApplicationService` — 1 nơi duy nhất quyết định routing; 2 port hẹp (`RoutingDecisionPort`, `ModeResolver`) + callback confirm ⇒ 0 phụ thuộc Qt | +| T03 | `application/model_routing/core_routing_adapter.py` | 169 | `CoreRoutingEngine` (cầu nối sang `core/routing`), `AppContextModeResolver`, `build_routing_application_service(ctx)` (cache 1 instance/ctx) | +| T06 | `infrastructure/telemetry/usage_sink.py` | 288 | `UsageEvent` + `UsageEventSink` (Protocol) + `UsageTrackerSink` / `InMemoryUsageSink` / `CompositeUsageSink`; publish không bao giờ raise | + +**File hiện hữu được sửa (đều có comment tiếng Anh tại mọi khối thay đổi):** + +| File | Thay đổi | +| :--- | :--- | +| `providers/factory.py` | Bỏ bảng `_REGISTRY` nội bộ, ủy quyền cho `ProviderRegistry`; vẫn raise `ProviderError` để không vỡ call site cũ | +| `providers/openai_compat.py`, `providers/anthropic.py` | Không còn gọi thẳng `core/usage_tracker`; chỉ **publish** `UsageEvent` qua sink (T06) | +| `ui/chat_panel.py` (#L638), `ui/co4e_tab.py`, `ui/folder_tab.py` | Xóa 3 bản sao logic routing (~35 dòng/file) ➔ gọi chung `RoutingApplicationService` (T04, T05); widget chỉ còn dựng `RoutingRequest`, host modal confirm và render kết quả | +| `config.py`, `state.py`, `ui/routing_toggle.py`, `i18n.py` | Mở đường cho chế độ thứ 4 **Fallback**: hằng `AppConfig.ROUTING_MODES`, validate per-workspace, thêm mục trong combo + chuỗi EN/JA/VI | +| `core/usage_tracker.py` | Thêm `current_context()` để sink mượn/trả lại context của thread thay vì gán đè vĩnh viễn | +| `tests/conftest.py`, `tests/routing/conftest.py` | **Sửa lỗi hạ tầng test nghiêm trọng** (xem "Ghi chú" bên dưới) | + +**Bộ test bổ sung (tất cả offline, không cần network/Qt):** + +| File | Số test | Phạm vi | +| :--- | :---: | :--- | +| `tests/contracts/test_providers.py` (+ `provider_stubs.py`) | 50 | Contract chạy parametrize trên **mọi** provider trong registry: signature `chat()`, canonical assistant message, tool call chuẩn hóa, đóng response, dịch tool schema, `ProviderError`, `list_models`/`test_connection`, đúng 1 `UsageEvent`/turn | +| `tests/unit/test_routing_application_service.py` | 28 | Đủ 4 chế độ + mọi nhánh degrade (engine lỗi, resolver lỗi, dialog lỗi, thiếu callback) | +| `tests/unit/test_provider_registry.py` | 17 | Descriptor + registry + đối chiếu catalogue với `DEFAULT_CONFIG["providers"]` | +| `tests/unit/test_core_routing_adapter.py` | 12 | Dịch `RouteResult` ⇄ DTO, task type sai định dạng, thiếu ranking, cache service | +| `tests/unit/test_usage_sink.py` | 13 | Fan-out, subscriber lỗi, khôi phục thread context, publish không raise | +| `tests/integration/test_routing_unification.py` | 14 | Chạy `RoutingApplicationService` trên **engine `core/routing` thật**; 3 surface (cowork/co4e/ai_edit) cho ra cùng 1 quyết định | + +**Kết quả cổng kiểm duyệt (DoD 7 tiêu chí):** + +| # | Tiêu chí | Lệnh | Kết quả | +| :---: | :--- | :--- | :--- | +| 1 | LOC < 400 | `wc -l` các file mới | ✅ Lớn nhất 288 dòng (`usage_sink.py`); `openai_compat.py` 374, `anthropic.py` 332 | +| 2 | Clean Architecture | `python scripts/check_imports.py` | ✅ `[PASS] 0 forbidden imports detected` | +| 3 | Comment tiếng Anh | Review thủ công | ✅ 100% khối code mới/sửa có comment giải thích logic + lý do kiến trúc | +| 4 | Có test tự động | `pytest tests/unit tests/contracts tests/integration` | ✅ 134 test mới, pass 100% | +| 5 | No Regression | `pytest tests/` | ✅ **236 passed in ~2.0s** (nền trước R03: 102 passed) | +| 6 | Timestamps | Bảng trên | ✅ Đã ghi Start/End cho T01–T06 | +| 7 | CASAN Gate | `scripts/run_quality_gate.py` | ⚠️ Script **chưa tồn tại** — thuộc R10-T02 (chưa làm). Đã chạy thay bằng `check_imports.py` + `pytest tests/` | + +**Ghi chú kỹ thuật cần biết khi review:** + +1. **Đã sửa 1 lỗi hạ tầng test có thể gây kết quả sai lệch**: `tests/conftest.py` cũ đẩy thư mục **cha** của repo vào `sys.path`, nên `import cowork_local.*` (dùng bởi `tests/routing/*` và `tests/characterization/*`) trỏ sang **một checkout `cowork_local` khác** nằm cạnh thư mục làm việc — test vẫn báo xanh nhưng chạy trên mã nguồn khác. Nay conftest bind thẳng checkout hiện tại vào `sys.modules["cowork_local"]`. +2. **Chế độ Fallback** là chế độ *chống gãy*, không phải chế độ tối ưu: giữ nguyên model người dùng chọn kể cả khi có model điểm cao hơn, chỉ chuyển khi model đó **không phục vụ được** turn (không có trong ranking / unavailable / probe fail). Engine `core/routing` không cần biết chế độ này — service map Fallback ➔ Auto khi hỏi ranking rồi tự áp luật chấp nhận riêng. +3. **T06 hiện tại**: provider publish `UsageEvent`; khi R04 dựng xong `AgentEvent` bus thì `ConversationApplicationService` sẽ là nơi phát sự kiện, sink giữ nguyên không phải sửa. +4. **Cần cài `mcp>=1.0.0`** (đã có trong `requirements.txt`) để `tests/test_project_context_mcp_template.py` collect được — thiếu gói này toàn bộ suite bị interrupt. --- @@ -230,16 +284,20 @@ | Ngày | Task Cần Hoàn Thành | Start Time | End Time | Trạng Thái | | :--- | :--- | :---: | :---: | :---: | | **21/08 (T6)** | Khóa DTO `ConversationExecutionRequest`, `AgentEvent`; Xây dựng `FakeProvider`, `FakeToolExecutor` | `2026-08-21 18:23` | `2026-08-21 18:35` | [x] | -| **22-23/08 (T7-CN)** | Chuẩn hóa `ProviderDescriptor`, `ProviderRegistry`; Wrap OpenAI, Anthropic, Ollama, FPT Gateway; Viết Contract Tests | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | -| **24/08 (T2)** | Xây dựng `RoutingApplicationService` độc lập Qt; Tách `ComposerWidget` & `AttachmentPicker` | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | +| **22-23/08 (T7-CN)** | Chuẩn hóa `ProviderDescriptor`, `ProviderRegistry`; Wrap OpenAI, Anthropic, Ollama, FPT Gateway; Viết Contract Tests | `2026-08-22 18:45` | `2026-08-22 19:01` | [x] | +| **24/08 (T2)** | Xây dựng `RoutingApplicationService` độc lập Qt; Tách `ComposerWidget` & `AttachmentPicker` | `2026-08-22 18:53` | `2026-08-22 18:57` | [~] | | **25/08 (T3)** | Xây dựng `ConversationApplicationService`; Tách `ChatHistoryWidget` và bubble renderer | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | | **26/08 (T4)** | Nối stream `AgentEvent` sang Chat History; Tách `AudioRecorderWidget` | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | | **27/08 (T5)** | Tách `ChatOutputPanel` & File Watcher; Lắp ráp container `ChatPanel` và `Floating HelpAgent` | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | -| **28/08 (T6)** | Xóa copy routing cũ trong `ui/chat_panel.py`; Fix circular import `model_pricing` ↔ `usage_tracker` | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | +| **28/08 (T6)** | Xóa copy routing cũ trong `ui/chat_panel.py`; Fix circular import `model_pricing` ↔ `usage_tracker` | `2026-08-22 18:57` | `2026-08-22 18:59` | [~] | | **29/08 (T7)** | Viết suite integration test cho toàn bộ luồng Chat (`tests/integration/test_chat_flow.py`) | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | | **30/08 (CN)** | 🔍 **Chủ trì CASAN Check 3**: Chạy `python scripts/check_imports.py` đảm bảo 0 import `PySide6` trong domain & application | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | | **31/08 (T2)** | **Chủ trì EPIC R10**: Viết Contributor Recipes, chạy E2E Smoke Test (`tests/e2e/test_smoke.py`) và merge PR cuối cùng | `____-__-__ __:__` | `____-__-__ __:__` | [ ] | +> **Chú thích trạng thái**: `[~]` = hoàn tất **phần thuộc EPIC R03**, phần còn lại của dòng đó thuộc EPIC khác nên chưa đóng. +> - Dòng **24/08**: đã xong `RoutingApplicationService` (R03-T03); phần `ComposerWidget`/`AttachmentPicker` thuộc R08-T01/T02 — chưa làm. +> - Dòng **28/08**: đã xóa copy routing trong `ui/chat_panel.py` (R03-T04) **và** cả `ui/co4e_tab.py`, `ui/folder_tab.py` (R03-T05); phần circular import `model_pricing` ↔ `usage_tracker` thuộc R09-T02 — chưa làm. + --- ### 🟣 TEAM NAM (Automation Workflows, Co4E, Monitoring & Governance) diff --git a/domain/models/provider_descriptor.py b/domain/models/provider_descriptor.py new file mode 100644 index 0000000..74301b9 --- /dev/null +++ b/domain/models/provider_descriptor.py @@ -0,0 +1,196 @@ +"""Provider catalog metadata — the domain-layer description of ONE LLM provider. + +Before R03 the answer to "which providers exist, what do they cost, what can +they do?" was spread over three places: the class table in +``providers/factory.py``, the hand-maintained pricing table in +``core/routing/metadata.py`` and a handful of ``if provider == "anthropic"`` +branches in the UI. :class:`ProviderDescriptor` is the single declarative +record those call sites now read from. + +Layer rules (see ``docs/architecture/ADR-001-layered-architecture.md``): this +module is 100% pure Python — no PySide6, no ``requests``, no filesystem, and no +import of the concrete ``providers/*`` adapters. It only *describes* a provider; +constructing one is the infrastructure layer's job +(``infrastructure/providers/provider_registry.py``). +""" + +from __future__ import annotations + +from dataclasses import dataclass, field, replace +from enum import Enum +from typing import Any, Dict, Optional, Tuple + + +class AuthKind(str, Enum): + """How a provider authenticates, so Settings/onboarding can ask for the + right thing instead of hard-coding per-provider form fields. + + Inherits ``str`` so a descriptor round-trips through JSON unchanged (the + value is written as a plain string), matching how the routing models in + ``core/routing/models.py`` already serialize their enums. + """ + + NONE = "none" # local runtimes (Ollama) — nothing to supply + API_KEY = "api_key" # bearer/x-api-key style secret + OAUTH_TOKEN = "oauth" # token minted by an external login flow (Copilot) + + +class WireProtocol(str, Enum): + """The on-the-wire dialect a provider speaks. + + Several *distinct* providers share one protocol (Ollama, Codex, GitHub + Copilot and generic gateways are all OpenAI Chat Completions), which is + exactly why protocol is a separate field from the provider id: the registry + picks the adapter class from the protocol, while everything user-facing + keys off the id. + """ + + OPENAI_COMPAT = "openai_compat" + ANTHROPIC = "anthropic" + + +@dataclass(frozen=True) +class ProviderDescriptor: + """Immutable metadata for one provider the app can route work to. + + Frozen because descriptors are shared process-wide by the registry, the + routing service and (eventually) the Settings screen; making them read-only + removes any chance one caller mutates the catalog another caller is + iterating. Use :meth:`with_models` to derive an updated copy instead. + + Unknown pricing/context values stay ``None`` rather than being guessed — + the routing scorer needs to distinguish "free" from "we don't know", the + same contract ``core/routing/models.py::ModelMetadata`` already follows. + """ + + provider_id: str # config key, e.g. "anthropic" + display_name: str # human label for Settings/UI + wire_protocol: WireProtocol # which adapter class implements it + auth_kind: AuthKind = AuthKind.API_KEY + default_model: str = "" # used when no model is selected + models: Tuple[str, ...] = () # known model ids (may be empty) + max_context: Optional[int] = None # tokens; None = unknown + cost_per_1k_input: Optional[float] = None # USD per 1K input tokens + cost_per_1k_output: Optional[float] = None # USD per 1K output tokens + supports_vision: bool = False + supports_tools: bool = True + supports_streaming: bool = True + requires_base_url: bool = False # gateway endpoints must be configured + # Extra ids that should resolve to this descriptor (renames/aliases kept for + # backwards compatibility with configs written by older app versions). + aliases: Tuple[str, ...] = () + # Free-form extension point so a team can attach provider-specific hints + # without another schema migration. + extras: Dict[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + """Reject descriptors that could never be looked up. + + Raising here (rather than at registration time) means a malformed + descriptor cannot exist at all, so every consumer downstream may assume + ``provider_id`` is a usable dict key. + """ + if not self.provider_id: + raise ValueError("ProviderDescriptor.provider_id must not be empty") + if not isinstance(self.wire_protocol, WireProtocol): + raise TypeError("ProviderDescriptor.wire_protocol must be a WireProtocol") + + # -- identity ------------------------------------------------------- # + @property + def identifiers(self) -> Tuple[str, ...]: + """Every id this descriptor answers to (canonical id first).""" + return (self.provider_id, *self.aliases) + + def matches(self, provider_id: str) -> bool: + """Case-insensitive id/alias match — config files and CLI flags are + typed by humans, so lookup must not be case sensitive.""" + needle = (provider_id or "").strip().lower() + return any(needle == known.lower() for known in self.identifiers) + + # -- capability queries --------------------------------------------- # + def knows_model(self, model_id: str) -> bool: + """Whether ``model_id`` is in this provider's declared catalog. + + A miss is NOT proof the model is unusable: gateways expose models we + cannot enumerate offline, so callers treat this as a hint (used to + resolve a bare model id back to its provider) and never as a gate that + blocks a request. + """ + needle = (model_id or "").strip().lower() + return any(needle == known.strip().lower() for known in self.models) + + def has_capability(self, capability: str) -> bool: + """Capability check by name, mirroring the vocabulary the routing + selector already filters on (``"vision"``, ``"tools"``, ``"streaming"``) + so a descriptor can be fed straight into ``rank_models``.""" + return capability in self.capabilities + + @property + def capabilities(self) -> frozenset: + """Capability set in the same vocabulary as + ``core/routing/models.py::ModelMetadata.capabilities``.""" + caps = set() + if self.supports_vision: + caps.add("vision") + if self.supports_tools: + caps.add("tools") + if self.supports_streaming: + caps.add("streaming") + return frozenset(caps) + + @property + def avg_cost_per_1k(self) -> Optional[float]: + """Blended input/output price, or ``None`` when either side is unknown. + + Uses the same 1:3 input:output weighting as + ``ModelMetadata.avg_cost_per_1k`` so a descriptor and an assessment + never disagree about what a model costs. + """ + ci, co = self.cost_per_1k_input, self.cost_per_1k_output + if ci is None or co is None: + return None + return (ci + 3.0 * co) / 4.0 + + def resolve_model(self, requested: str = "") -> str: + """The model id to actually call: the caller's choice when they made + one, otherwise this provider's default. Centralised here because every + surface (chat, Co4E, AI-Edit) previously re-implemented the same + ``model or config_default`` fallback inline.""" + return (requested or "").strip() or self.default_model + + # -- derivation / serialization ------------------------------------- # + def with_models(self, models, *, default_model: str = "") -> "ProviderDescriptor": + """A copy carrying a freshly discovered model list. + + Providers can enumerate their models at runtime (``list_models()``); + because the descriptor is frozen, discovery produces a NEW descriptor + that the registry swaps in atomically instead of mutating one that other + threads may be reading. + """ + ordered = tuple(dict.fromkeys(m for m in models if m)) # de-dup, keep order + chosen = default_model or self.default_model + # Keep the default pointing at something real: fall back to the first + # discovered model when the configured default vanished from the catalog. + if ordered and chosen not in ordered: + chosen = ordered[0] + return replace(self, models=ordered, default_model=chosen) + + def to_dict(self) -> Dict[str, Any]: + """JSON-friendly view for config persistence and the Settings UI.""" + return { + "provider_id": self.provider_id, + "display_name": self.display_name, + "wire_protocol": self.wire_protocol.value, + "auth_kind": self.auth_kind.value, + "default_model": self.default_model, + "models": list(self.models), + "max_context": self.max_context, + "cost_per_1k_input": self.cost_per_1k_input, + "cost_per_1k_output": self.cost_per_1k_output, + "capabilities": sorted(self.capabilities), + "requires_base_url": self.requires_base_url, + "aliases": list(self.aliases), + } + + +__all__ = ["AuthKind", "WireProtocol", "ProviderDescriptor"] diff --git a/i18n.py b/i18n.py index e3b8b2e..0c7300b 100644 --- a/i18n.py +++ b/i18n.py @@ -583,10 +583,13 @@ STRINGS: Dict[str, Dict[str, str]] = { "routing.mode_off": {"en": "Off", "ja": "オフ", "vi": "Tắt"}, "routing.mode_auto": {"en": "Auto", "ja": "自動", "vi": "Tự động"}, "routing.mode_manual": {"en": "Manual", "ja": "手動", "vi": "Thủ công"}, + # Fallback (R03-T03): resilience mode -- never switches for a better + # score, only to rescue a selected model that cannot serve the turn. + "routing.mode_fallback": {"en": "Fallback", "ja": "フォールバック", "vi": "Dự phòng"}, "routing.toggle_tooltip": { - "en": "Auto model routing for this chat.\nOff: always use the selected model.\nAuto: silently switch to the best-fit model.\nManual: ask before switching.", - "ja": "このチャットの自動モデルルーティング。\nオフ: 選択したモデルを常に使用。\n自動: 最適なモデルへ自動切替。\n手動: 切替前に確認。", - "vi": "Tự động định tuyến model cho khung chat này.\nTắt: luôn dùng model đã chọn.\nTự động: tự chuyển sang model phù hợp nhất.\nThủ công: hỏi xác nhận trước khi chuyển.", + "en": "Auto model routing for this chat.\nOff: always use the selected model.\nAuto: silently switch to the best-fit model.\nManual: ask before switching.\nFallback: keep the selected model, switch only if it is unavailable.", + "ja": "このチャットの自動モデルルーティング。\nオフ: 選択したモデルを常に使用。\n自動: 最適なモデルへ自動切替。\n手動: 切替前に確認。\nフォールバック: 選択モデルを維持し、利用できない場合のみ切替。", + "vi": "Tự động định tuyến model cho khung chat này.\nTắt: luôn dùng model đã chọn.\nTự động: tự chuyển sang model phù hợp nhất.\nThủ công: hỏi xác nhận trước khi chuyển.\nDự phòng: giữ model đã chọn, chỉ chuyển khi model đó không dùng được.", }, "routing.confirm_title": { "en": "Switch model?", "ja": "モデルを切り替えますか?", "vi": "Chuyển model?", diff --git a/infrastructure/providers/provider_registry.py b/infrastructure/providers/provider_registry.py new file mode 100644 index 0000000..65e5b10 --- /dev/null +++ b/infrastructure/providers/provider_registry.py @@ -0,0 +1,287 @@ +"""Central registry of every LLM provider the app can talk to. + +Replaces the bare ``{name: class}`` dict in ``providers/factory.py`` as the +single catalogue of providers. Two responsibilities, kept deliberately narrow: + +1. **Lookup** — resolve a provider id (or one of its aliases, or a bare model + id) to its :class:`~domain.models.provider_descriptor.ProviderDescriptor`. +2. **Construction** — instantiate the concrete adapter class that speaks the + descriptor's wire protocol. + +This is infrastructure, not domain: it is allowed to import the concrete +``providers/*`` adapters (which pull in ``requests``). The adapters are imported +lazily inside :meth:`build` so that merely *reading the catalogue* — which the +pure routing service does on every turn — never drags the HTTP stack into the +process. +""" + +from __future__ import annotations + +import threading +from typing import Any, Dict, Iterable, List, Optional + +from ...domain.models.provider_descriptor import ( + AuthKind, + ProviderDescriptor, + WireProtocol, +) + +# --------------------------------------------------------------------------- # +# Built-in catalogue. +# +# Mirrors DEFAULT_CONFIG["providers"] in config.py (ids + default models) and +# providers/factory.py (id -> wire protocol). Prices are intentionally absent: +# core/routing/metadata.py owns cost, and a guessed price is worse than a +# known-unknown (see that module's docstring). +# --------------------------------------------------------------------------- # +BUILTIN_DESCRIPTORS: tuple = ( + ProviderDescriptor( + provider_id="openai_compat", + display_name="OpenAI-compatible gateway", + wire_protocol=WireProtocol.OPENAI_COMPAT, + auth_kind=AuthKind.API_KEY, + default_model="gpt-4o-mini", + supports_vision=True, + # A generic gateway has no fixed host, so the endpoint MUST be + # configured before the provider can be used at all. + requires_base_url=True, + ), + ProviderDescriptor( + provider_id="anthropic", + display_name="Anthropic Claude", + wire_protocol=WireProtocol.ANTHROPIC, + auth_kind=AuthKind.API_KEY, + default_model="claude-sonnet-4-6", + # Kept in sync with AnthropicProvider._FALLBACK_MODELS — the list the + # provider itself falls back to when /v1/models cannot be reached. + models=("claude-opus-4-8", "claude-sonnet-4-6", "claude-haiku-4-5-20251001"), + max_context=200000, + supports_vision=True, + ), + ProviderDescriptor( + provider_id="ollama", + display_name="Ollama (local)", + wire_protocol=WireProtocol.OPENAI_COMPAT, + # A local runtime needs no credential; Settings must not demand one. + auth_kind=AuthKind.NONE, + default_model="llama3.1", + supports_vision=False, + requires_base_url=True, + ), + ProviderDescriptor( + provider_id="github_copilot", + display_name="GitHub Copilot", + wire_protocol=WireProtocol.OPENAI_COMPAT, + # The credential is a Copilot token minted by an external login flow, + # not a self-service API key. + auth_kind=AuthKind.OAUTH_TOKEN, + default_model="gpt-4o", + models=("gpt-4o", "gpt-4o-mini"), + max_context=128000, + supports_vision=True, + ), + ProviderDescriptor( + provider_id="codex", + display_name="OpenAI", + wire_protocol=WireProtocol.OPENAI_COMPAT, + auth_kind=AuthKind.API_KEY, + default_model="gpt-4o-mini", + models=("gpt-4o", "gpt-4o-mini", "o1", "o3"), + max_context=128000, + supports_vision=True, + # Historic config key: early builds stored this provider as "openai". + aliases=("openai",), + ), +) + + +class ProviderNotFoundError(LookupError): + """Raised when no descriptor answers to the requested provider id. + + A dedicated type (rather than bare ``KeyError``) lets callers distinguish + "this provider is not in the catalogue" from an unrelated dict miss, and + keeps the message actionable by listing what IS registered. + """ + + +class ProviderRegistry: + """Thread-safe catalogue of :class:`ProviderDescriptor` records. + + Thread-safety matters because model discovery runs on background worker + threads (the routing prober, Settings' "Load models") and republishes an + updated descriptor via :meth:`replace`, while chat turns on other threads + are reading the catalogue concurrently. + """ + + def __init__(self, descriptors: Optional[Iterable[ProviderDescriptor]] = None) -> None: + # Keyed by canonical id; alias resolution walks the values so an alias + # can never shadow a real provider id. + self._by_id: Dict[str, ProviderDescriptor] = {} + self._lock = threading.RLock() + for descriptor in descriptors or (): + self.register(descriptor) + + # -- registration --------------------------------------------------- # + def register(self, descriptor: ProviderDescriptor) -> ProviderDescriptor: + """Add a descriptor. Refuses to silently overwrite an existing id so a + typo in a plugin cannot hijack a built-in provider; use :meth:`replace` + when an update is the actual intent.""" + with self._lock: + existing = self._by_id.get(descriptor.provider_id) + if existing is not None and existing != descriptor: + raise ValueError( + f"Provider '{descriptor.provider_id}' is already registered; " + "call replace() to update it." + ) + self._by_id[descriptor.provider_id] = descriptor + return descriptor + + def replace(self, descriptor: ProviderDescriptor) -> ProviderDescriptor: + """Register or update a descriptor unconditionally — the path model + discovery uses to publish a freshly enumerated model list.""" + with self._lock: + self._by_id[descriptor.provider_id] = descriptor + return descriptor + + # -- lookup ---------------------------------------------------------- # + def get(self, provider_id: str) -> ProviderDescriptor: + """Descriptor for ``provider_id`` (canonical id or alias). + + Raises :class:`ProviderNotFoundError` rather than returning ``None`` so + a misconfigured provider fails loudly at the call site instead of + surfacing later as an ``AttributeError`` on ``None``. + """ + found = self.find(provider_id) + if found is None: + known = ", ".join(sorted(self._by_id)) or "" + raise ProviderNotFoundError( + f"Unsupported provider: {provider_id!r}. Registered: {known}" + ) + return found + + def find(self, provider_id: str) -> Optional[ProviderDescriptor]: + """Non-raising :meth:`get` — ``None`` when nothing matches.""" + needle = (provider_id or "").strip() + if not needle: + return None + with self._lock: + direct = self._by_id.get(needle) + if direct is not None: + return direct + # Fall back to a case-insensitive id/alias scan; order is stable + # because dicts preserve insertion order, so the earliest-registered + # provider wins a tie. + for descriptor in self._by_id.values(): + if descriptor.matches(needle): + return descriptor + return None + + def find_by_model(self, model_id: str) -> Optional[ProviderDescriptor]: + """Resolve a bare model id back to the provider that serves it. + + This is the "dynamic lookup by model ID" R03-T02 calls for: routing + decisions and saved conversations sometimes carry only a model name, and + the caller still needs to know which provider to build. Returns ``None`` + when the model belongs to a gateway whose catalogue we cannot enumerate + offline — callers then fall back to the configured active provider. + """ + needle = (model_id or "").strip() + if not needle: + return None + with self._lock: + for descriptor in self._by_id.values(): + if descriptor.knows_model(needle): + return descriptor + return None + + def all(self) -> List[ProviderDescriptor]: + """Every registered descriptor, in registration order (snapshot copy — + safe to iterate while another thread registers).""" + with self._lock: + return list(self._by_id.values()) + + def ids(self) -> List[str]: + """Canonical provider ids, sorted for stable UI/reporting output.""" + with self._lock: + return sorted(self._by_id) + + def __contains__(self, provider_id: object) -> bool: + return isinstance(provider_id, str) and self.find(provider_id) is not None + + def __len__(self) -> int: + with self._lock: + return len(self._by_id) + + # -- construction ---------------------------------------------------- # + def adapter_class(self, provider_id: str): + """Concrete ``Provider`` subclass implementing this provider's protocol. + + The adapters are imported here (not at module import) so the pure + routing/domain code can consult the catalogue without loading + ``requests`` and the whole HTTP stack. + """ + descriptor = self.get(provider_id) + from ...providers.anthropic import AnthropicProvider + from ...providers.openai_compat import OpenAICompatProvider + + protocol_to_class = { + WireProtocol.OPENAI_COMPAT: OpenAICompatProvider, + WireProtocol.ANTHROPIC: AnthropicProvider, + } + adapter = protocol_to_class.get(descriptor.wire_protocol) + if adapter is None: # pragma: no cover — unreachable while the map is total + raise ProviderNotFoundError( + f"No adapter implements wire protocol {descriptor.wire_protocol!r}" + ) + return adapter + + def build(self, provider_id: str, conf: Dict[str, Any]): + """Instantiate a ready-to-use provider adapter. + + The descriptor's ``default_model`` fills in a missing/blank ``model`` so + a half-written config still produces a working provider instead of an + empty model id that only fails once the request hits the gateway. + """ + descriptor = self.get(provider_id) + adapter = self.adapter_class(descriptor.provider_id) + merged = dict(conf or {}) + merged["model"] = descriptor.resolve_model(merged.get("model", "")) + return adapter(merged) + + +# --------------------------------------------------------------------------- # +# Process-wide default registry. +# +# Built lazily under a lock: several UI screens can ask for it during startup +# from different threads, and double-construction would hand out two catalogues +# whose discovered model lists then drift apart. +# --------------------------------------------------------------------------- # +_default_registry: Optional[ProviderRegistry] = None +_default_lock = threading.Lock() + + +def default_registry() -> ProviderRegistry: + """The shared registry seeded with :data:`BUILTIN_DESCRIPTORS`.""" + global _default_registry + if _default_registry is None: + with _default_lock: + if _default_registry is None: + _default_registry = ProviderRegistry(BUILTIN_DESCRIPTORS) + return _default_registry + + +def reset_default_registry() -> None: + """Drop the cached registry — test-support hook so one test's registrations + cannot leak into the next.""" + global _default_registry + with _default_lock: + _default_registry = None + + +__all__ = [ + "BUILTIN_DESCRIPTORS", + "ProviderNotFoundError", + "ProviderRegistry", + "default_registry", + "reset_default_registry", +] diff --git a/infrastructure/telemetry/usage_sink.py b/infrastructure/telemetry/usage_sink.py new file mode 100644 index 0000000..ef5d8c3 --- /dev/null +++ b/infrastructure/telemetry/usage_sink.py @@ -0,0 +1,288 @@ +"""Token-usage telemetry as a publish/subscribe seam (R03-T06). + +Before this module every provider adapter reached straight into +``core/usage_tracker.py`` and wrote a dashboard row itself, which meant the +provider layer owned a telemetry policy decision ("where do usage numbers go?") +and no test could observe a turn's token accounting without touching the real +``~/.cowork_local/usage/`` files. + +Now a provider only *describes what happened* — it publishes an immutable +:class:`UsageEvent` — and subscribers decide what to do with it. The default +subscriber, :class:`UsageTrackerSink`, forwards to the existing usage tracker so +the Dashboard keeps working byte-for-byte; tests swap in +:class:`InMemoryUsageSink` and assert on the events directly. + +Every publish path is failure-tolerant on purpose: telemetry must never be the +reason a chat turn dies, which is the same contract +``usage_tracker.record()`` already documents. +""" + +from __future__ import annotations + +import logging +import threading +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional, Protocol, runtime_checkable + +logger = logging.getLogger("cowork_local.telemetry.usage") + + +@dataclass(frozen=True) +class UsageEvent: + """One provider turn's token accounting. + + Frozen so a subscriber cannot mutate an event the next subscriber in the + chain is about to receive. ``source``/``label`` stay optional: the usage + tracker already derives them from thread-local context set by whoever ran + the turn, and a provider adapter has no business knowing which UI surface + invoked it. + """ + + provider: str + model: str + input_tokens: int = 0 + output_tokens: int = 0 + cached_tokens: int = 0 + # True when the counts are a ~4-chars-per-token approximation because the + # gateway never sent a usage block. Surfaced in the Dashboard so users know + # which rows are measured and which are guessed. + estimated: bool = False + source: Optional[str] = None # None -> tracker's thread-local context + label: Optional[str] = None # None -> tracker's thread-local context + extras: Dict[str, Any] = field(default_factory=dict) + + @property + def total_tokens(self) -> int: + """Billable token count for this turn (cached tokens are already part + of the input count reported by every gateway we support, so adding them + again would double-count).""" + return int(self.input_tokens) + int(self.output_tokens) + + def to_dict(self) -> Dict[str, Any]: + """JSON-friendly view, using the same short keys as the usage tracker's + on-disk rows so a caller can diff an event against a stored row.""" + return { + "provider": self.provider, + "model": self.model, + "in": int(self.input_tokens), + "out": int(self.output_tokens), + "cache": int(self.cached_tokens), + "estimated": bool(self.estimated), + "source": self.source or "", + "label": self.label or "", + } + + +@runtime_checkable +class UsageEventSink(Protocol): + """Anything that can receive :class:`UsageEvent`s. + + A ``Protocol`` rather than a base class so a plain object (or a test double, + or a Qt-side adapter that re-emits a signal) qualifies without inheriting + from infrastructure code. + """ + + def emit(self, event: UsageEvent) -> None: + """Handle one usage event. Implementations MUST NOT raise.""" + + +class UsageTrackerSink: + """Default subscriber: writes each event through ``core/usage_tracker.py``. + + Keeps the existing Dashboard/telemetry pipeline (daily JSONL files, shared + cross-machine mirror, per-thread accumulator) as the single writer, so + routing this through an event seam changed the plumbing without changing + a single stored byte. + """ + + def __init__(self, recorder=None) -> None: + # The recorder is injectable so a test can verify the forwarding + # contract without importing the real tracker (and its config paths). + self._recorder = recorder + + def _resolve_recorder(self): + """Late-bind ``usage_tracker.record``. + + Imported on first use rather than at module import so telemetry stays + out of the import graph of anything that merely *declares* a sink. + """ + if self._recorder is None: + from ...core import usage_tracker as tracker + + self._recorder = tracker.record + return self._recorder + + def emit(self, event: UsageEvent) -> None: + """Forward one event; swallow every failure (telemetry is never fatal).""" + try: + record = self._resolve_recorder() + if event.source is None: + # Normal path: the worker thread already tagged its own + # source/label via set_context(), so record() attributes the row. + record( + event.provider, event.model, + int(event.input_tokens), int(event.output_tokens), + int(event.cached_tokens), estimated=bool(event.estimated), + ) + return + + # Event carries its own attribution: apply it for this single write + # and restore the thread's previous context afterwards, so a + # re-attributed event cannot silently relabel every later turn that + # runs on the same worker thread. + from ...core import usage_tracker as tracker + + previous_source, previous_label = tracker.current_context() + tracker.set_context(event.source, event.label or "") + try: + record( + event.provider, event.model, + int(event.input_tokens), int(event.output_tokens), + int(event.cached_tokens), estimated=bool(event.estimated), + ) + finally: + tracker.set_context(previous_source, previous_label) + except Exception: # noqa: BLE001 — usage tracking must never break a turn + logger.debug("usage sink: forwarding to usage_tracker failed", exc_info=True) + + +class InMemoryUsageSink: + """Collects events in a list — the test double for usage assertions.""" + + def __init__(self) -> None: + self.events: List[UsageEvent] = [] + self._lock = threading.Lock() + + def emit(self, event: UsageEvent) -> None: + """Append under a lock: parallel Co4E flows publish from several worker + threads at once and ``list.append`` alone would still be atomic, but the + lock also makes :meth:`snapshot` a consistent read.""" + with self._lock: + self.events.append(event) + + def snapshot(self) -> List[UsageEvent]: + """A copy of everything received so far.""" + with self._lock: + return list(self.events) + + def clear(self) -> None: + with self._lock: + self.events.clear() + + @property + def total_tokens(self) -> int: + return sum(e.total_tokens for e in self.snapshot()) + + +class CompositeUsageSink: + """Fans one event out to several subscribers. + + This is what makes the seam useful beyond the Dashboard: a future consumer + (per-workspace budget guard, live cost meter) subscribes alongside the + tracker instead of patching provider code again. One failing subscriber is + logged and skipped so it cannot starve the others. + """ + + def __init__(self, sinks=None) -> None: + self._sinks: List[UsageEventSink] = list(sinks or ()) + self._lock = threading.RLock() + + def add(self, sink: UsageEventSink) -> None: + with self._lock: + self._sinks.append(sink) + + def remove(self, sink: UsageEventSink) -> None: + """Detach a subscriber; a sink that was never added is ignored so + teardown code can call this unconditionally.""" + with self._lock: + if sink in self._sinks: + self._sinks.remove(sink) + + def sinks(self) -> List[UsageEventSink]: + with self._lock: + return list(self._sinks) + + def emit(self, event: UsageEvent) -> None: + for sink in self.sinks(): + try: + sink.emit(event) + except Exception: # noqa: BLE001 — one bad subscriber must not stop the rest + logger.debug("usage sink: subscriber %r failed", sink, exc_info=True) + + +# --------------------------------------------------------------------------- # +# Process-wide sink. +# +# Providers publish through the module-level helpers below rather than holding a +# sink reference, because a provider instance is created fresh for every turn +# (see AppContext.build_provider_for) and would otherwise have to be handed the +# telemetry wiring on every construction. +# --------------------------------------------------------------------------- # +_sink_lock = threading.RLock() +_sink: Optional[CompositeUsageSink] = None + + +def get_usage_sink() -> CompositeUsageSink: + """The shared sink, seeded with :class:`UsageTrackerSink` on first use.""" + global _sink + if _sink is None: + with _sink_lock: + if _sink is None: + _sink = CompositeUsageSink([UsageTrackerSink()]) + return _sink + + +def set_usage_sink(sink: Optional[CompositeUsageSink]) -> None: + """Replace the shared sink (``None`` restores the default on next use). + + Used by tests and by the app shell when it wants a different fan-out; kept + explicit so nothing silently reconfigures telemetry mid-run. + """ + global _sink + with _sink_lock: + _sink = sink + + +def subscribe(sink: UsageEventSink) -> UsageEventSink: + """Attach an extra subscriber to the shared sink and return it (so callers + can keep the handle for a later :func:`unsubscribe`).""" + get_usage_sink().add(sink) + return sink + + +def unsubscribe(sink: UsageEventSink) -> None: + """Detach a subscriber previously passed to :func:`subscribe`.""" + get_usage_sink().remove(sink) + + +def publish(event: UsageEvent) -> None: + """Publish one usage event to every subscriber. + + Never raises: called from inside a provider's streaming loop, where an + exception would abort an otherwise successful turn. + """ + try: + get_usage_sink().emit(event) + except Exception: # noqa: BLE001 + logger.debug("usage sink: publish failed", exc_info=True) + + +def estimate_tokens(text: str) -> int: + """~4 chars per token approximation, re-exported so provider adapters need + exactly ONE telemetry import instead of also importing the tracker.""" + return max(0, len(text or "") // 4) + + +__all__ = [ + "UsageEvent", + "UsageEventSink", + "UsageTrackerSink", + "InMemoryUsageSink", + "CompositeUsageSink", + "get_usage_sink", + "set_usage_sink", + "subscribe", + "unsubscribe", + "publish", + "estimate_tokens", +] diff --git a/providers/anthropic.py b/providers/anthropic.py index 0d63437..34b1edf 100644 --- a/providers/anthropic.py +++ b/providers/anthropic.py @@ -292,19 +292,31 @@ class AnthropicProvider(Provider): args = {"_raw": b["json"]} tool_calls.append({"id": b["id"], "name": b["name"], "arguments": args}) - # Dashboard usage event — real counts from the stream's usage events, - # else a ~4 chars/token estimate. Never breaks the turn. + # Usage event — real counts from the stream's usage events, else a + # ~4 chars/token estimate. Published to the telemetry sink (R03-T06) + # rather than written straight to the Dashboard store, so the provider + # stays a pure transport adapter. Never breaks the turn. try: - from ..core import usage_tracker as ut + from ..infrastructure.telemetry import usage_sink if usage_seen: - ut.record(self.name, self.model, usage_seen.get("in", 0), - usage_seen.get("out", 0), usage_seen.get("cache", 0)) + usage_sink.publish(usage_sink.UsageEvent( + provider=self.name, + model=self.model, + input_tokens=usage_seen.get("in", 0), + output_tokens=usage_seen.get("out", 0), + cached_tokens=usage_seen.get("cache", 0), + )) else: sent = json.dumps(payload.get("messages", []), ensure_ascii=False) got = "".join(text_parts) + "".join(b["json"] for b in blocks.values()) - ut.record(self.name, self.model, ut.estimate_tokens(sent), - ut.estimate_tokens(got), 0, estimated=True) + usage_sink.publish(usage_sink.UsageEvent( + provider=self.name, + model=self.model, + input_tokens=usage_sink.estimate_tokens(sent), + output_tokens=usage_sink.estimate_tokens(got), + estimated=True, + )) except Exception: # noqa: BLE001 pass diff --git a/providers/factory.py b/providers/factory.py index fb43b4c..11aeeea 100644 --- a/providers/factory.py +++ b/providers/factory.py @@ -1,25 +1,32 @@ -"""Build a provider instance from the application config.""" +"""Build a provider instance from the application config. + +Kept as the historic entry point (``providers.build_provider``) that call sites +across the app already import, but it no longer owns a provider table of its +own: since R03-T02 the catalogue lives in +``infrastructure/providers/provider_registry.py`` so provider ids, wire +protocols, default models and capabilities are declared exactly once. +""" from __future__ import annotations from typing import Any, Dict -from .anthropic import AnthropicProvider from .base import Provider, ProviderError -from .openai_compat import OpenAICompatProvider - -_REGISTRY = { - "openai_compat": OpenAICompatProvider, - "anthropic": AnthropicProvider, - # All OpenAI-compatible endpoints (Ollama's /v1 server, the Copilot chat API, - # and OpenAI itself) speak the same Chat Completions protocol. - "ollama": OpenAICompatProvider, - "github_copilot": OpenAICompatProvider, - "codex": OpenAICompatProvider, -} def build_provider(name: str, conf: Dict[str, Any]) -> Provider: - cls = _REGISTRY.get(name) - if cls is None: - raise ProviderError(f"Unsupported provider: {name}") - return cls(conf) + """Construct the adapter registered for ``name``. + + Delegates to the central registry and translates its lookup failure into + :class:`ProviderError`, because every existing call site (chat turns, + Settings' connection test, the routing prober) already handles that type — + changing the exception would ripple into unrelated error handling. + """ + from ..infrastructure.providers.provider_registry import ( + ProviderNotFoundError, + default_registry, + ) + + try: + return default_registry().build(name, conf) + except ProviderNotFoundError as exc: + raise ProviderError(f"Unsupported provider: {name}") from exc diff --git a/providers/openai_compat.py b/providers/openai_compat.py index c45083f..056425f 100644 --- a/providers/openai_compat.py +++ b/providers/openai_compat.py @@ -266,22 +266,38 @@ class OpenAICompatProvider(Provider): return _assemble_assistant(text_parts, tool_acc) def _record_usage(self, messages, text_parts, tool_acc, usage_seen) -> None: - """One Dashboard usage event per turn: real counts when the server's - final chunk carried a "usage" block, a ~4 chars/token estimate - otherwise. Never breaks the turn.""" + """Publish one usage event per turn: real counts when the server's final + chunk carried a "usage" block, a ~4 chars/token estimate otherwise. + + Since R03-T06 this only *describes* what the turn consumed and hands the + event to ``infrastructure/telemetry/usage_sink.py``; deciding where the + numbers land (Dashboard files, cost meters, tests) belongs to the + subscribers, not to a provider adapter. Never breaks the turn. + """ try: - from ..core import usage_tracker as ut + from ..infrastructure.telemetry import usage_sink if usage_seen: - ut.record(self.name, self.model, - usage_seen.get("prompt_tokens", 0), - usage_seen.get("completion_tokens", 0), - (usage_seen.get("prompt_tokens_details") or {}).get("cached_tokens", 0)) + usage_sink.publish(usage_sink.UsageEvent( + provider=self.name, + model=self.model, + input_tokens=usage_seen.get("prompt_tokens", 0), + output_tokens=usage_seen.get("completion_tokens", 0), + cached_tokens=(usage_seen.get("prompt_tokens_details") or {}).get("cached_tokens", 0), + )) else: + # No usage block from the gateway — approximate from the exact + # bytes we sent and received so the Dashboard still shows a + # (clearly flagged) figure instead of a silent zero. sent = json.dumps(self._to_api_messages(messages), ensure_ascii=False) got = "".join(text_parts) + "".join(s["args"] for s in tool_acc.values()) - ut.record(self.name, self.model, ut.estimate_tokens(sent), - ut.estimate_tokens(got), 0, estimated=True) + usage_sink.publish(usage_sink.UsageEvent( + provider=self.name, + model=self.model, + input_tokens=usage_sink.estimate_tokens(sent), + output_tokens=usage_sink.estimate_tokens(got), + estimated=True, + )) except Exception: # noqa: BLE001 pass diff --git a/state.py b/state.py index 98ab7a1..e87057a 100644 --- a/state.py +++ b/state.py @@ -73,14 +73,18 @@ class AppContext: return load_project(pid) def project_routing_mode(self, surface: str) -> str: - """Effective Off/Auto/Manual routing mode for a chat ``surface`` in the - ACTIVE workspace: the workspace's own override wins; otherwise the + """Effective Off/Auto/Manual/Fallback routing mode for a chat ``surface`` + in the ACTIVE workspace: the workspace's own override wins; otherwise the global default (``config.routing_mode_for``). This is what makes each - workspace keep its own routing mode.""" + workspace keep its own routing mode. + + The accepted set is taken from ``AppConfig.ROUTING_MODES`` rather than + repeated here, so adding a mode (as R03-T03 did with "fallback") stays a + one-line change instead of a hunt through every validation site.""" project = self._current_project() if project is not None: mode = (project.routing_modes or {}).get(surface, "") - if mode in ("off", "auto", "manual"): + if mode in self.config.ROUTING_MODES: return mode return self.config.routing_mode_for(surface) @@ -88,7 +92,7 @@ class AppContext: """Persist a surface's routing mode for the ACTIVE workspace. With no workspace selected, falls back to the global setting so behaviour outside a project stays global.""" - mode = mode if mode in ("off", "auto", "manual") else "off" + mode = mode if mode in self.config.ROUTING_MODES else "off" project = self._current_project() if project is None: self.config.set_routing_mode_for(surface, mode) diff --git a/tests/conftest.py b/tests/conftest.py index 46e4d53..1397f91 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,10 +1,75 @@ -"""Make the repository package importable when pytest runs from the repo root.""" +"""Make THIS checkout importable as the ``cowork_local`` package during tests. + +Why this is not just a ``sys.path`` insert +------------------------------------------ +Test modules import the app in two different styles: + +* top-level (``from providers.base import ...``) — resolved by the repository + root already sitting on ``sys.path`` when pytest is launched from it; +* fully qualified (``from cowork_local.core.routing.service import ...``) — + which only resolves when a directory literally named ``cowork_local`` is + importable. + +Simply appending the repository's PARENT directory to ``sys.path`` (the previous +behaviour) makes the second style resolve against *whatever* sibling folder +happens to be called ``cowork_local`` — on a developer machine that is often an +unrelated older checkout, so the whole suite silently exercises the wrong code +while still reporting green. Instead we bind the name ``cowork_local`` in +``sys.modules`` to the package rooted at THIS repository, so both import styles +always reach the working copy under test regardless of the checkout's directory +name. +""" from __future__ import annotations +import importlib.util import sys from pathlib import Path -REPOSITORY_PARENT = Path(__file__).resolve().parents[2] -if str(REPOSITORY_PARENT) not in sys.path: - sys.path.insert(0, str(REPOSITORY_PARENT)) +# ...//tests/conftest.py -> .../ +PACKAGE_ROOT = Path(__file__).resolve().parents[1] +PACKAGE_NAME = "cowork_local" + +# The repository root must stay importable so the top-level import style +# (``providers``/``domain``/``application``/``tests``) keeps working. +if str(PACKAGE_ROOT) not in sys.path: + sys.path.insert(0, str(PACKAGE_ROOT)) + + +def _bind_checkout_as_package() -> None: + """Register this checkout in ``sys.modules`` under the canonical package name. + + Executed at import time of the conftest (i.e. before any test module is + imported) so that a stale same-named directory elsewhere on ``sys.path`` can + never win the lookup. A no-op when the package is already bound to this very + directory, which keeps repeated conftest loads (pytest-xdist, sub-sessions) + idempotent. + """ + existing = sys.modules.get(PACKAGE_NAME) + if existing is not None: + # Already bound. Only rebind when it points at a DIFFERENT checkout, + # otherwise re-executing the package __init__ would duplicate module + # state that tests may already hold references to. + origin = getattr(existing, "__file__", "") or "" + if Path(origin).resolve().parent == PACKAGE_ROOT: + return + + spec = importlib.util.spec_from_file_location( + PACKAGE_NAME, + PACKAGE_ROOT / "__init__.py", + # Declaring the search locations is what turns the module into a real + # package, so ``cowork_local.core.routing`` and friends resolve as + # sub-modules of this directory. + submodule_search_locations=[str(PACKAGE_ROOT)], + ) + if spec is None or spec.loader is None: # pragma: no cover — defensive + return + module = importlib.util.module_from_spec(spec) + # Insert BEFORE executing so that a circular ``import cowork_local`` from + # inside the package body resolves to the partially-initialised module + # instead of restarting the import (standard CPython import semantics). + sys.modules[PACKAGE_NAME] = module + spec.loader.exec_module(module) + + +_bind_checkout_as_package() diff --git a/tests/contracts/__init__.py b/tests/contracts/__init__.py new file mode 100644 index 0000000..5ebab88 --- /dev/null +++ b/tests/contracts/__init__.py @@ -0,0 +1,7 @@ +"""Contract tests: one shared specification every interchangeable adapter must satisfy. + +Unlike unit tests (which pin ONE implementation's behaviour), a contract test is +parametrised over every implementation of an interface, so adding a new provider +means adding a row — not writing a new test file — and a provider that quietly +breaks the canonical shape fails here rather than in production. +""" diff --git a/tests/contracts/provider_stubs.py b/tests/contracts/provider_stubs.py new file mode 100644 index 0000000..810dd5e --- /dev/null +++ b/tests/contracts/provider_stubs.py @@ -0,0 +1,178 @@ +"""Offline transport doubles + per-protocol stream scripts for the provider contract tests. + +Kept in its own module so ``test_providers.py`` stays a readable list of +assertions instead of a wall of SSE fixtures, and so the LOC ceiling (400 lines +per production file, applied here too) is comfortably met by both halves. + +Nothing in here touches the network: :class:`FakeStreamResponse` mimics just +enough of ``requests.Response`` for the streaming loops in +``providers/openai_compat.py`` and ``providers/anthropic.py`` — status code, +mutable ``encoding``, ``iter_lines`` and ``close``. +""" + +from __future__ import annotations + +import json +from typing import Any, Dict, List, Optional + +# Canonical turn every protocol script below must produce, so the contract test +# can assert one expected result no matter which provider produced it. +EXPECTED_TEXT = "Hello world" +EXPECTED_TOOL_CALL = {"id": "call-1", "name": "read_file", "arguments": {"path": "a.txt"}} +EXPECTED_INPUT_TOKENS = 11 +EXPECTED_OUTPUT_TOKENS = 7 +EXPECTED_CACHED_TOKENS = 3 + + +class FakeStreamResponse: + """A minimal stand-in for a streaming ``requests.Response``. + + ``iter_lines`` replays pre-baked SSE lines; ``closed`` records that the + provider released the connection, which the contract asserts because a + provider that leaks the response leaks a socket per turn. + """ + + def __init__( + self, + lines: Optional[List[str]] = None, + status_code: int = 200, + body: str = "", + headers: Optional[Dict[str, str]] = None, + payload: Optional[Dict[str, Any]] = None, + ) -> None: + self.status_code = status_code + self._lines = list(lines or ()) + self.text = body + self.headers = dict(headers or {}) + self._payload = payload + self.closed = False + # Providers force UTF-8 on the response before reading it; the attribute + # simply has to exist and be writable. + self.encoding = None + + def iter_lines(self, decode_unicode: bool = False): + for line in self._lines: + yield line + + def json(self) -> Any: + if self._payload is None: + raise ValueError("no JSON payload configured on this fake response") + return self._payload + + def close(self) -> None: + self.closed = True + + +def _sse(payload: Dict[str, Any]) -> str: + """One SSE ``data:`` line carrying a JSON event.""" + return "data: " + json.dumps(payload, ensure_ascii=False) + + +def openai_stream_lines() -> List[str]: + """A complete OpenAI Chat Completions stream: text, one tool call, usage. + + Split across several deltas on purpose — chunk boundaries are where naive + stream parsers break, so the contract exercises them. + """ + return [ + _sse({"choices": [{"delta": {"content": "Hello "}}]}), + _sse({"choices": [{"delta": {"content": "world"}}]}), + _sse({"choices": [{"delta": {"tool_calls": [{ + "index": 0, "id": "call-1", + "function": {"name": "read_file", "arguments": '{"path":'}, + }]}}]}), + # Arguments arrive fragmented; the provider must concatenate before parsing. + _sse({"choices": [{"delta": {"tool_calls": [{ + "index": 0, "function": {"arguments": '"a.txt"}'}, + }]}}]}), + _sse({ + "choices": [{"delta": {}}], + "usage": { + "prompt_tokens": EXPECTED_INPUT_TOKENS, + "completion_tokens": EXPECTED_OUTPUT_TOKENS, + "prompt_tokens_details": {"cached_tokens": EXPECTED_CACHED_TOKENS}, + }, + }), + "data: [DONE]", + ] + + +def anthropic_stream_lines() -> List[str]: + """The same canonical turn expressed as an Anthropic Messages stream.""" + return [ + _sse({"type": "message_start", "message": {"usage": { + "input_tokens": EXPECTED_INPUT_TOKENS, + "cache_read_input_tokens": EXPECTED_CACHED_TOKENS, + }}}), + _sse({"type": "content_block_start", "index": 0, + "content_block": {"type": "text"}}), + _sse({"type": "content_block_delta", "index": 0, + "delta": {"type": "text_delta", "text": "Hello "}}), + _sse({"type": "content_block_delta", "index": 0, + "delta": {"type": "text_delta", "text": "world"}}), + _sse({"type": "content_block_start", "index": 1, "content_block": { + "type": "tool_use", "id": "call-1", "name": "read_file"}}), + _sse({"type": "content_block_delta", "index": 1, + "delta": {"type": "input_json_delta", "partial_json": '{"path":'}}), + _sse({"type": "content_block_delta", "index": 1, + "delta": {"type": "input_json_delta", "partial_json": '"a.txt"}'}}), + _sse({"type": "message_delta", + "usage": {"output_tokens": EXPECTED_OUTPUT_TOKENS}}), + _sse({"type": "message_stop"}), + ] + + +# Per wire protocol: how to script a successful turn, and the model-list payload +# ``list_models()`` expects. Keyed by the descriptor's wire protocol value so a +# new provider that reuses an existing protocol needs no new entry here. +PROTOCOL_FIXTURES = { + "openai_compat": { + "stream_lines": openai_stream_lines, + "models_payload": {"data": [{"id": "gpt-4o-mini"}, {"id": "gpt-4o"}]}, + "expected_models": ["gpt-4o-mini", "gpt-4o"], + }, + "anthropic": { + "stream_lines": anthropic_stream_lines, + "models_payload": {"data": [{"id": "claude-sonnet-4-6"}]}, + "expected_models": ["claude-sonnet-4-6"], + }, +} + + +class ScriptedTransport: + """Replaces ``Provider._request`` and hands back scripted responses. + + Records every call so a test can assert *how* the provider talked to the + endpoint (method, url, JSON payload) without a socket ever being opened. + """ + + def __init__(self, responses: List[FakeStreamResponse]) -> None: + self._responses = list(responses) + self.calls: List[Dict[str, Any]] = [] + + def __call__(self, method: str, url: str, **kwargs) -> FakeStreamResponse: + self.calls.append({"method": method, "url": url, **kwargs}) + if not self._responses: + raise AssertionError(f"unexpected extra request: {method} {url}") + # Pop in order: a provider that retries gets the NEXT scripted response, + # which is how the retry/error paths are driven. + return self._responses.pop(0) + + @property + def last_payload(self) -> Dict[str, Any]: + """The JSON body of the most recent request.""" + return self.calls[-1].get("json") or {} + + +__all__ = [ + "EXPECTED_CACHED_TOKENS", + "EXPECTED_INPUT_TOKENS", + "EXPECTED_OUTPUT_TOKENS", + "EXPECTED_TEXT", + "EXPECTED_TOOL_CALL", + "FakeStreamResponse", + "PROTOCOL_FIXTURES", + "ScriptedTransport", + "anthropic_stream_lines", + "openai_stream_lines", +] diff --git a/tests/contracts/test_providers.py b/tests/contracts/test_providers.py new file mode 100644 index 0000000..a3449b2 --- /dev/null +++ b/tests/contracts/test_providers.py @@ -0,0 +1,279 @@ +"""R03-T01 — the contract every LLM provider adapter must satisfy. + +Parametrised over EVERY provider in the central registry +(``infrastructure/providers/provider_registry.py``), so registering a new +provider automatically subjects it to the same specification and a provider that +drifts from the canonical shapes fails here. + +The contract, in one list: + +* construction — the registry builds a real ``Provider`` for every id; +* ``chat()`` — canonical signature, canonical assistant message, streamed text + delivered through ``on_text``, tool calls normalised to + ``{"id", "name", "arguments": dict}``, response always closed; +* tool schema translation matches the adapter's wire protocol; +* failures raise ``ProviderError`` — never a bare transport exception; +* ``list_models()`` / ``test_connection()`` report a reason instead of a silent + empty list; +* telemetry — exactly one ``UsageEvent`` per turn (R03-T06), with the real + counts when the stream reports them. + +Everything runs offline: ``Provider._request`` is replaced by a scripted +transport, so the suite needs no network, no API key and no Qt event loop. +""" + +from __future__ import annotations + +import pytest +import requests +from cowork_local.infrastructure.providers.provider_registry import ( + BUILTIN_DESCRIPTORS, + ProviderRegistry, +) +from cowork_local.infrastructure.telemetry import usage_sink +from cowork_local.providers.base import Provider, ProviderError, ToolSpec +from cowork_local.tests.contracts.provider_stubs import ( + EXPECTED_CACHED_TOKENS, + EXPECTED_INPUT_TOKENS, + EXPECTED_OUTPUT_TOKENS, + EXPECTED_TEXT, + EXPECTED_TOOL_CALL, + PROTOCOL_FIXTURES, + FakeStreamResponse, + ScriptedTransport, +) + +# Every provider id in the catalogue — the parametrisation that makes this a +# contract suite rather than a per-adapter unit test. +PROVIDER_IDS = [d.provider_id for d in BUILTIN_DESCRIPTORS] + +# Minimal config: enough for any adapter to build a URL and headers offline. +BASE_CONF = {"base_url": "https://gateway.test/v1", "api_key": "test-key"} + +SAMPLE_MESSAGES = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Say hello"}, +] + +SAMPLE_TOOL = ToolSpec( + name="read_file", + description="Read a file from disk", + parameters={"type": "object", "properties": {"path": {"type": "string"}}}, +) + + +@pytest.fixture() +def registry() -> ProviderRegistry: + """A private registry per test so registrations never leak between tests.""" + return ProviderRegistry(BUILTIN_DESCRIPTORS) + + +@pytest.fixture() +def collected_usage(monkeypatch) -> usage_sink.InMemoryUsageSink: + """Swap the process-wide telemetry sink for an in-memory one. + + Restored by monkeypatch after each test, so a contract run never appends to + the developer's real ``~/.cowork_local/usage/`` files. + """ + sink = usage_sink.InMemoryUsageSink() + monkeypatch.setattr(usage_sink, "_sink", usage_sink.CompositeUsageSink([sink])) + return sink + + +def _fixtures_for(registry: ProviderRegistry, provider_id: str) -> dict: + """The stream/model-list script matching this provider's wire protocol.""" + protocol = registry.get(provider_id).wire_protocol.value + return PROTOCOL_FIXTURES[protocol] + + +def _build(registry: ProviderRegistry, provider_id: str, transport=None) -> Provider: + """Build a provider and (optionally) replace its transport with a script.""" + provider = registry.build(provider_id, dict(BASE_CONF)) + if transport is not None: + # Patch the INSTANCE, not the class: parallel parametrised cases must + # not see each other's scripted transport. + provider._request = transport + return provider + + +# --------------------------------------------------------------------------- # +# Construction & interface shape +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_registry_builds_a_provider_for_every_registered_id(registry, provider_id) -> None: + """Every catalogued provider must be constructible — a descriptor with no + working adapter is a broken entry, not a feature flag.""" + provider = _build(registry, provider_id) + + assert isinstance(provider, Provider) + # The registry fills in the descriptor's default model when config omits it, + # so a half-configured provider still names a concrete model. + assert provider.model, f"{provider_id} built without a model id" + assert provider.describe() == f"{provider.name}:{provider.model}" + + +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_chat_signature_is_uniform(registry, provider_id) -> None: + """All adapters accept the same call, so the agent runtime can swap + providers without knowing which one it holds.""" + import inspect + + provider = _build(registry, provider_id) + params = list(inspect.signature(provider.chat).parameters) + + assert params == ["messages", "tools", "on_text", "cancel", "on_reasoning"] + + +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_tool_schema_matches_the_wire_protocol(registry, provider_id) -> None: + """A ToolSpec must translate into the exact shape the endpoint expects.""" + descriptor = registry.get(provider_id) + + if descriptor.wire_protocol.value == "anthropic": + translated = SAMPLE_TOOL.to_anthropic() + assert translated["input_schema"] == SAMPLE_TOOL.parameters + assert translated["name"] == "read_file" + else: + translated = SAMPLE_TOOL.to_openai() + assert translated["type"] == "function" + assert translated["function"]["parameters"] == SAMPLE_TOOL.parameters + + +# --------------------------------------------------------------------------- # +# The turn itself +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_chat_returns_the_canonical_assistant_message(registry, provider_id, collected_usage) -> None: + """Whatever the wire format, one turn yields the same canonical result.""" + fixtures = _fixtures_for(registry, provider_id) + response = FakeStreamResponse(lines=fixtures["stream_lines"]()) + transport = ScriptedTransport([response]) + provider = _build(registry, provider_id, transport) + + streamed: list = [] + result = provider.chat( + SAMPLE_MESSAGES, tools=[SAMPLE_TOOL], on_text=streamed.append, + ) + + assert result["role"] == "assistant" + assert result["content"] == EXPECTED_TEXT + # Text must arrive incrementally, not only in the final message — the chat + # UI streams from these callbacks. + assert "".join(streamed) == EXPECTED_TEXT + assert len(streamed) >= 2 + # Tool calls are normalised: parsed arguments, never the raw JSON fragments. + assert result["tool_calls"] == [EXPECTED_TOOL_CALL] + assert response.closed, "provider left the streaming response open" + + +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_chat_publishes_exactly_one_usage_event(registry, provider_id, collected_usage) -> None: + """R03-T06: a turn reports its token usage through the telemetry sink, with + the server's real counts when the stream carried them.""" + fixtures = _fixtures_for(registry, provider_id) + transport = ScriptedTransport([FakeStreamResponse(lines=fixtures["stream_lines"]())]) + provider = _build(registry, provider_id, transport) + + provider.chat(SAMPLE_MESSAGES, tools=[SAMPLE_TOOL]) + + events = collected_usage.snapshot() + assert len(events) == 1, "a turn must publish exactly one usage event" + event = events[0] + assert event.provider == provider.name + assert event.model == provider.model + assert event.input_tokens == EXPECTED_INPUT_TOKENS + assert event.output_tokens == EXPECTED_OUTPUT_TOKENS + assert event.cached_tokens == EXPECTED_CACHED_TOKENS + # Real counts were available, so the event must NOT be flagged as a guess. + assert event.estimated is False + + +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_usage_is_estimated_when_the_stream_reports_none(registry, provider_id, collected_usage) -> None: + """Gateways that never send usage still produce a dashboard row — clearly + flagged as an estimate rather than silently recorded as zero.""" + # Only text; no usage block anywhere in the stream. + silent_stream = ['data: ' + '{"choices": [{"delta": {"content": "hi"}}]}', "data: [DONE]"] + if registry.get(provider_id).wire_protocol.value == "anthropic": + silent_stream = [ + 'data: {"type": "content_block_start", "index": 0, "content_block": {"type": "text"}}', + 'data: {"type": "content_block_delta", "index": 0,' + ' "delta": {"type": "text_delta", "text": "hi"}}', + ] + transport = ScriptedTransport([FakeStreamResponse(lines=silent_stream)]) + provider = _build(registry, provider_id, transport) + + provider.chat(SAMPLE_MESSAGES) + + events = collected_usage.snapshot() + assert len(events) == 1 + assert events[0].estimated is True + # An estimate still has to be a positive number to be worth showing. + assert events[0].total_tokens > 0 + + +# --------------------------------------------------------------------------- # +# Failure behaviour +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_http_error_becomes_provider_error(registry, provider_id, collected_usage) -> None: + """Callers handle exactly one exception type; adapters must not leak + transport- or JSON-level errors past their boundary.""" + failing = FakeStreamResponse(status_code=401, body='{"error": {"message": "bad key"}}') + transport = ScriptedTransport([failing]) + provider = _build(registry, provider_id, transport) + + with pytest.raises(ProviderError): + provider.chat(SAMPLE_MESSAGES) + + assert failing.closed, "provider left a failed response open" + + +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_list_models_and_test_connection_report_a_reason(registry, provider_id) -> None: + """A failed model load must explain itself: ``last_error`` is what Settings + shows instead of an unexplained empty dropdown.""" + def _boom(*_args, **_kwargs): + # A transport failure, i.e. what actually happens when the gateway is + # unreachable — adapters translate this class of error, not arbitrary + # programming errors, which must still surface as bugs. + raise requests.ConnectionError("network down") + + provider = _build(registry, provider_id, _boom) + + models = provider.list_models() + + assert provider.last_error, f"{provider_id} swallowed a model-load failure" + ok, message = provider.test_connection() + assert ok is False + assert message + # Anthropic answers with a built-in fallback catalogue; a gateway answers + # with nothing. Both are acceptable — the contract is only that a failure is + # never reported as success. + assert isinstance(models, list) + + +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_list_models_returns_ids_on_success(registry, provider_id) -> None: + """The happy path returns plain model-id strings, not raw API objects.""" + fixtures = _fixtures_for(registry, provider_id) + transport = ScriptedTransport([ + FakeStreamResponse(status_code=200, payload=fixtures["models_payload"]), + ]) + provider = _build(registry, provider_id, transport) + + models = provider.list_models() + + assert models == fixtures["expected_models"] + assert provider.last_error == "" + assert all(isinstance(m, str) for m in models) + + +@pytest.mark.parametrize("provider_id", PROVIDER_IDS) +def test_strip_think_removes_inline_reasoning(registry, provider_id) -> None: + """Reasoning must never leak into a final answer, whichever adapter ran.""" + provider = _build(registry, provider_id) + + cleaned = provider.strip_think("secret planVisible answer") + + assert cleaned == "Visible answer" diff --git a/tests/integration/__init__.py b/tests/integration/__init__.py new file mode 100644 index 0000000..25b54b4 --- /dev/null +++ b/tests/integration/__init__.py @@ -0,0 +1,7 @@ +"""Integration tests: several real layers wired together, still fully offline. + +Where unit tests pin one class against fakes and contract tests pin an interface +across implementations, these exercise a real path end to end — e.g. the +application routing service on top of the real ``core/routing`` engine — so a +seam that only works against a mock is caught here. +""" diff --git a/tests/integration/test_routing_unification.py b/tests/integration/test_routing_unification.py new file mode 100644 index 0000000..353ca55 --- /dev/null +++ b/tests/integration/test_routing_unification.py @@ -0,0 +1,249 @@ +"""R03-T03/T04/T05 — the unified routing path over the REAL routing engine. + +The unit tests drive ``RoutingApplicationService`` against fakes; this suite +proves the same service produces correct outcomes on top of the actual +``core/routing`` stack (classifier → assessment store → scorer → selector → +switch controller), which is what the three chat surfaces now call. + +Offline by construction: a fake probe client answers benchmarks and judging, and +the assessment store is a temp file — no network, no Qt, no ``$HOME`` writes. +""" + +from __future__ import annotations + +import copy + +import pytest +from cowork_local.application.model_routing import ( + AppContextModeResolver, + CoreRoutingEngine, + RoutingApplicationService, + RoutingMode, + RoutingRequest, +) +from cowork_local.config import DEFAULT_CONFIG, AppConfig +from cowork_local.core import projects as projects_mod +from cowork_local.core.routing.clients import CompletionResult +from cowork_local.core.routing.service import RoutingService +from cowork_local.core.routing.store import AssessmentStore +from cowork_local.state import AppContext + +STRONG_ANSWER = "STRONG-DETAILED-CORRECT-ANSWER" +WEAK_ANSWER = "weak" + + +class FakeProbeClient: + """Deterministic stand-in for the provider layer used during assessment. + + Mirrors ``tests/routing/test_service.py``'s client: benchmark prompts get a + per-model canned answer, and judge prompts are graded by looking up that + answer, so scores are stable and no model is ever really called. + """ + + def __init__(self, answers, quality) -> None: + self.answers = answers + self.quality = quality + + def complete(self, provider, model_id, messages) -> CompletionResult: + text = messages[0]["content"] + if "grading an AI assistant" in text: # the judge rubric prompt + score = 0.0 + for answer, value in self.quality.items(): + if answer and answer in text: + score = value + break + return CompletionResult(text='{"score": %s}' % score) + answer = self.answers.get((provider, model_id)) + if answer is None: + return CompletionResult(error="unavailable") + return CompletionResult(text=answer, tokens_out=len(answer) // 4) + + +@pytest.fixture() +def ctx(tmp_path, monkeypatch): + """An AppContext with two assessable models and temp-only persistence.""" + # Keep workspace load/save off the developer's real ~/.cowork_local. + monkeypatch.setattr(projects_mod, "PROJECTS_DIR", tmp_path / "projects") + data = copy.deepcopy(DEFAULT_CONFIG) + data["providers"] = { + "anthropic": {"base_url": "x", "api_key": "x", "model": "strong-model"}, + } + data["routing"]["candidates"] = [ + {"provider": "anthropic", "model_id": "strong-model", "tier": "powerful"}, + {"provider": "anthropic", "model_id": "weak-model", "tier": "fast"}, + ] + data["routing"]["judge_provider"] = "anthropic" + data["routing"]["judge_model"] = "judge-model" + data["routing"]["policy"] = "quality" + data["routing"]["min_score_gain"] = 0.05 + return AppContext(AppConfig(data=data, path=tmp_path / "config.json")) + + +@pytest.fixture() +def routing_service(ctx, tmp_path) -> RoutingService: + """A real RoutingService with a populated assessment store.""" + client = FakeProbeClient( + answers={ + ("anthropic", "strong-model"): STRONG_ANSWER, + ("anthropic", "weak-model"): WEAK_ANSWER, + }, + quality={STRONG_ANSWER: 0.95, WEAK_ANSWER: 0.35}, + ) + store = AssessmentStore(store_path=tmp_path / "assess.json", + history_dir=tmp_path / "history") + service = RoutingService(ctx, store=store, client=client) + service.reassess() # populate real probe results + fit scores + return service + + +@pytest.fixture() +def app_service(ctx, routing_service) -> RoutingApplicationService: + """The application service wired exactly the way the UI wires it.""" + return RoutingApplicationService( + CoreRoutingEngine(routing_service), + AppContextModeResolver(ctx), + confirm_timeout_sec=lambda: float(ctx.config.routing["confirm_timeout_sec"]), + ) + + +def coding_request(**overrides) -> RoutingRequest: + """A coding turn currently pinned to the weaker model.""" + fields = dict( + surface="cowork", + prompt="Write a Python function to reverse a linked list", + current_provider="anthropic", + current_model="weak-model", + ) + fields.update(overrides) + return RoutingRequest(**fields) + + +# --------------------------------------------------------------------------- # +# Auto / Off / Manual over the real engine +# --------------------------------------------------------------------------- # +def test_auto_switches_to_the_better_assessed_model(app_service) -> None: + """The real scorer must rank the strong model first and the service must + hand that model back as this turn's override.""" + outcome = app_service.resolve(coding_request(mode=RoutingMode.AUTO)) + + assert outcome.switched is True + assert outcome.provider == "anthropic" + assert outcome.model == "strong-model" + assert outcome.task_type == "coding" # classified from the prompt + assert outcome.score_gain > 0 + + +def test_off_keeps_the_pinned_model(app_service) -> None: + """Off must not switch even when a clearly better model is assessed.""" + outcome = app_service.resolve(coding_request(mode=RoutingMode.OFF)) + + assert outcome.switched is False + assert outcome.provider is None + + +def test_manual_asks_before_switching(app_service) -> None: + """The confirm callback receives the engine's own decision object, which is + what ``ui/routing_toggle.py::confirm_switch`` renders.""" + seen: list = [] + + outcome = app_service.resolve( + coding_request(mode=RoutingMode.MANUAL), + confirm=lambda decision, timeout: seen.append((decision, timeout)) or True, + ) + + assert outcome.switched is True + decision, timeout = seen[0] + assert decision.to_model == "anthropic/strong-model" + assert decision.reason # human-readable explanation + assert timeout == pytest.approx(60.0) # from DEFAULT_CONFIG + + +def test_manual_decline_keeps_the_pinned_model(app_service) -> None: + outcome = app_service.resolve( + coding_request(mode=RoutingMode.MANUAL), + confirm=lambda decision, timeout: False, + ) + + assert outcome.switched is False + assert outcome.declined is True + + +def test_already_best_model_is_left_alone(app_service) -> None: + """No pointless churn: being on the best model is not a switch.""" + outcome = app_service.resolve( + coding_request(mode=RoutingMode.AUTO, current_model="strong-model")) + + assert outcome.switched is False + + +# --------------------------------------------------------------------------- # +# Fallback over the real engine +# --------------------------------------------------------------------------- # +def test_fallback_keeps_an_assessed_model_even_though_a_better_one_exists(app_service) -> None: + """weak-model IS usable (it has a real probe score), so Fallback stays put + where Auto would switch — the behavioural difference between the modes.""" + outcome = app_service.resolve(coding_request(mode=RoutingMode.FALLBACK)) + + assert outcome.switched is False + + +def test_fallback_rescues_a_model_the_engine_cannot_serve(app_service) -> None: + """A model absent from the ranking (never assessed / unavailable) is exactly + the situation Fallback exists for.""" + outcome = app_service.resolve( + coding_request(mode=RoutingMode.FALLBACK, current_model="ghost-model")) + + assert outcome.switched is True + assert outcome.model == "strong-model" + + +# --------------------------------------------------------------------------- # +# Surface parity — the point of R03-T04/T05 +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("surface", ["cowork", "co4e", "ai_edit"]) +def test_every_surface_gets_the_same_decision(app_service, surface) -> None: + """Chat, Co4E and AI-Edit used to hold three copies of this logic. Given the + same inputs they must now be indistinguishable.""" + outcome = app_service.resolve(coding_request(surface=surface, mode=RoutingMode.AUTO)) + + assert outcome.switched is True + assert outcome.model == "strong-model" + + +def test_ai_edit_pinned_task_type_reaches_the_engine(app_service) -> None: + """AI-Edit pins "coding" instead of classifying; the engine must honour it + even when the instruction text reads like something else entirely.""" + outcome = app_service.resolve(coding_request( + surface="ai_edit", + prompt="Write a poem about the ocean", # classifier would say "creative" + task_type="coding", + mode=RoutingMode.AUTO, + )) + + assert outcome.task_type == "coding" + + +def test_mode_comes_from_the_workspace_when_not_pinned(ctx, app_service) -> None: + """With no explicit mode, the service reads the per-workspace setting — the + lookup the widgets used to do themselves.""" + ctx.config.data["routing"]["switch_mode"] = "auto" + + outcome = app_service.resolve(coding_request()) + + assert outcome.mode is RoutingMode.AUTO + assert outcome.switched is True + + +def test_fallback_mode_survives_a_round_trip_through_config(ctx) -> None: + """The new mode must be persistable, or the toggle could never select it.""" + ctx.config.set_routing_mode_for("cowork", "fallback") + + assert ctx.config.routing_mode_for("cowork") == "fallback" + assert ctx.project_routing_mode("cowork") == "fallback" + + +def test_unknown_persisted_mode_degrades_to_off(ctx) -> None: + """A hand-edited config must not enable routing by accident.""" + ctx.config.routing["surface_modes"]["cowork"] = "turbo" + + assert ctx.config.routing_mode_for("cowork") == "off" diff --git a/tests/routing/conftest.py b/tests/routing/conftest.py index c892dd1..be48344 100644 --- a/tests/routing/conftest.py +++ b/tests/routing/conftest.py @@ -1,17 +1,9 @@ """Pytest fixtures/shared helpers for the routing test suite. -Ensures the ``cowork_local`` package is importable when pytest is invoked from -the package directory itself (so ``import cowork_local.core.routing...`` works -regardless of the working directory the suite is launched from). +Package importability is handled once and for all by ``tests/conftest.py``, +which binds THIS checkout to the ``cowork_local`` name in ``sys.modules``. +This file used to push the checkout's PARENT directory onto ``sys.path``, which +let an unrelated sibling folder named ``cowork_local`` shadow the working copy — +so that logic is intentionally gone; keep it that way. """ from __future__ import annotations - -import sys -from pathlib import Path - -# .../cowork_local/tests/routing/conftest.py → parent of the package dir -_PKG_DIR = Path(__file__).resolve().parents[2] # .../cowork_local -_REPO_ROOT = _PKG_DIR.parent # .../cowork_local_20260722 -for p in (str(_REPO_ROOT), str(_PKG_DIR)): - if p not in sys.path: - sys.path.insert(0, p) diff --git a/tests/unit/test_core_routing_adapter.py b/tests/unit/test_core_routing_adapter.py new file mode 100644 index 0000000..ed84ab8 --- /dev/null +++ b/tests/unit/test_core_routing_adapter.py @@ -0,0 +1,220 @@ +"""Unit tests for the adapters that bridge the routing engine to the app service. + +The integration suite covers the happy path over the real engine; this file pins +the translation edge cases that are hard to provoke there — malformed task +types, a missing ranking, and the service-caching contract. +""" + +from __future__ import annotations + +import pytest +from cowork_local.application.model_routing import ( + AppContextModeResolver, + CoreRoutingEngine, + RoutingApplicationService, + RoutingMode, + RoutingRequest, +) +from cowork_local.application.model_routing.core_routing_adapter import ( + build_routing_application_service, +) +from cowork_local.core.routing.models import SwitchDecision, SwitchMode, TaskType + + +class FakeRanking: + """Just enough of ``selector.Ranking`` for the adapter's usability check.""" + + def __init__(self, scores) -> None: + self._scores = dict(scores) + + def score_of(self, key: str) -> float: + return self._scores.get(key, 0.0) + + +class FakeRouteResult: + """Stands in for ``core.routing.service.RouteResult``.""" + + def __init__(self, decision, task_type=TaskType.CODING, ranking=None, target=None) -> None: + self.decision = decision + self.task_type = task_type + self.ranking = ranking + self._target = target + + @property + def should_switch(self) -> bool: + return self.decision.should_switch + + def target(self): + return self._target + + +class FakeRoutingService: + """Records the arguments the adapter forwards to the engine.""" + + def __init__(self, result: FakeRouteResult) -> None: + self.result = result + self.calls: list = [] + + def route(self, surface, prompt, current_provider, current_model, **kwargs): + self.calls.append({"surface": surface, "prompt": prompt, + "current_provider": current_provider, + "current_model": current_model, **kwargs}) + return self.result + + +def make_decision(**overrides) -> SwitchDecision: + fields = dict( + should_switch=True, + from_model="anthropic/weak-model", + to_model="anthropic/strong-model", + score_gain=0.3, + reason="coding fit 0.9 > current 0.6", + mode=SwitchMode.AUTO, + task_type="coding", + ) + fields.update(overrides) + return SwitchDecision(**fields) + + +def make_request(**overrides) -> RoutingRequest: + fields = dict(surface="cowork", prompt="Fix this bug", + current_provider="anthropic", current_model="weak-model") + fields.update(overrides) + return RoutingRequest(**fields) + + +# --------------------------------------------------------------------------- # +# CoreRoutingEngine translation +# --------------------------------------------------------------------------- # +def test_engine_flattens_the_route_result() -> None: + """No ``core.routing`` type may leak past the adapter — the application + service and the widgets only ever see plain fields.""" + service = FakeRoutingService(FakeRouteResult( + make_decision(), + ranking=FakeRanking({"anthropic/weak-model": 0.6}), + target=("anthropic", "strong-model"), + )) + + evaluation = CoreRoutingEngine(service).evaluate(make_request(), RoutingMode.AUTO) + + assert evaluation.task_type == "coding" # str, not TaskType + assert evaluation.should_switch is True + assert evaluation.target_provider == "anthropic" + assert evaluation.target_model == "strong-model" + assert evaluation.score_gain == pytest.approx(0.3) + assert evaluation.current_is_usable is True + + +def test_engine_forwards_the_mode_as_a_plain_string() -> None: + """``RoutingService.route`` takes the mode as a string; handing it an enum + would silently fall through to its "unknown mode -> off" branch.""" + service = FakeRoutingService(FakeRouteResult(make_decision(should_switch=False))) + + CoreRoutingEngine(service).evaluate(make_request(), RoutingMode.AUTO) + + assert service.calls[0]["mode_override"] == "auto" + + +def test_engine_reports_an_unranked_model_as_unusable() -> None: + """This is the signal Fallback acts on: absent from the ranking means the + selector already rejected it (unavailable / no probe / failed probe).""" + service = FakeRoutingService(FakeRouteResult( + make_decision(), + ranking=FakeRanking({"anthropic/strong-model": 0.9}), # current is absent + target=("anthropic", "strong-model"), + )) + + evaluation = CoreRoutingEngine(service).evaluate(make_request(), RoutingMode.AUTO) + + assert evaluation.current_is_usable is False + + +def test_engine_assumes_usable_without_a_ranking() -> None: + """No ranking (routing off, or the engine's own error path) is absence of + evidence — it must not trigger a surprise Fallback switch.""" + service = FakeRoutingService(FakeRouteResult(make_decision(), ranking=None)) + + evaluation = CoreRoutingEngine(service).evaluate(make_request(), RoutingMode.AUTO) + + assert evaluation.current_is_usable is True + + +def test_engine_assumes_usable_when_the_ranking_misbehaves() -> None: + """A broken ranking object must not fail the turn.""" + class BrokenRanking: + def score_of(self, key): + raise RuntimeError("corrupt ranking") + + service = FakeRoutingService(FakeRouteResult(make_decision(), ranking=BrokenRanking())) + + evaluation = CoreRoutingEngine(service).evaluate(make_request(), RoutingMode.AUTO) + + assert evaluation.current_is_usable is True + + +@pytest.mark.parametrize( + "raw, expected", + [("coding", TaskType.CODING), ("QA", TaskType.QA), (None, None), ("nonsense", None)], +) +def test_task_type_strings_are_coerced_or_dropped(raw, expected) -> None: + """A pinned task type is honoured; an unknown one falls back to letting the + engine classify the prompt rather than raising mid-turn.""" + service = FakeRoutingService(FakeRouteResult(make_decision(should_switch=False))) + + CoreRoutingEngine(service).evaluate(make_request(task_type=raw), RoutingMode.AUTO) + + assert service.calls[0]["task_type"] == expected + + +def test_required_capabilities_are_passed_as_a_list_or_none() -> None: + """``rank_models`` filters on a list; an empty tuple must become None so it + is treated as "no filter" rather than "require nothing, but filter".""" + service = FakeRoutingService(FakeRouteResult(make_decision(should_switch=False))) + engine = CoreRoutingEngine(service) + + engine.evaluate(make_request(required_capabilities=("vision",)), RoutingMode.AUTO) + engine.evaluate(make_request(), RoutingMode.AUTO) + + assert service.calls[0]["required_capabilities"] == ["vision"] + assert service.calls[1]["required_capabilities"] is None + + +# --------------------------------------------------------------------------- # +# Mode resolver + wiring +# --------------------------------------------------------------------------- # +def test_mode_resolver_reads_the_per_workspace_mode() -> None: + """Per-workspace routing keeps working now that the lookup left the widgets.""" + class StubCtx: + def project_routing_mode(self, surface): + return "fallback" if surface == "co4e" else "off" + + resolver = AppContextModeResolver(StubCtx()) + + assert resolver.mode_for("co4e") is RoutingMode.FALLBACK + assert resolver.mode_for("cowork") is RoutingMode.OFF + + +def test_service_is_built_once_and_cached_on_the_context() -> None: + """Every surface must share one instance, so future per-surface state (a + cool-down, a switch history) is shared rather than duplicated per widget.""" + class StubCtx: + def __init__(self): + self.routing_calls = 0 + self.config = type("Cfg", (), {"routing": {"confirm_timeout_sec": 45}})() + + def routing(self): + self.routing_calls += 1 + return FakeRoutingService(FakeRouteResult(make_decision(should_switch=False))) + + def project_routing_mode(self, surface): + return "off" + + ctx = StubCtx() + first = build_routing_application_service(ctx) + second = build_routing_application_service(ctx) + + assert first is second + assert ctx.routing_calls == 1 + assert isinstance(first, RoutingApplicationService) + # The confirm timeout is read from config at call time, not frozen at build. + assert first.confirm_timeout() == pytest.approx(45.0) diff --git a/tests/unit/test_provider_registry.py b/tests/unit/test_provider_registry.py new file mode 100644 index 0000000..c740d28 --- /dev/null +++ b/tests/unit/test_provider_registry.py @@ -0,0 +1,204 @@ +"""R03-T02 — unit tests for ProviderDescriptor and the central ProviderRegistry. + +Covers what the rest of the app now relies on the catalogue for: resolving ids +and aliases, resolving a bare model id back to its provider, filling in default +models, and refusing to let a duplicate registration silently hijack a built-in. +""" + +from __future__ import annotations + +import pytest +from cowork_local.domain.models.provider_descriptor import ( + AuthKind, + ProviderDescriptor, + WireProtocol, +) +from cowork_local.infrastructure.providers.provider_registry import ( + BUILTIN_DESCRIPTORS, + ProviderNotFoundError, + ProviderRegistry, +) + + +def make_descriptor(**overrides) -> ProviderDescriptor: + """A minimal valid descriptor; tests override just the field under test.""" + fields = dict( + provider_id="demo", + display_name="Demo provider", + wire_protocol=WireProtocol.OPENAI_COMPAT, + default_model="demo-small", + models=("demo-small", "demo-large"), + ) + fields.update(overrides) + return ProviderDescriptor(**fields) + + +# --------------------------------------------------------------------------- # +# ProviderDescriptor +# --------------------------------------------------------------------------- # +def test_descriptor_rejects_an_empty_id() -> None: + """An id-less descriptor could never be looked up, so it must not exist.""" + with pytest.raises(ValueError): + make_descriptor(provider_id="") + + +def test_descriptor_rejects_a_non_enum_protocol() -> None: + """The protocol drives adapter selection; a stray string would silently + fall through to "no adapter" at build time instead of failing here.""" + with pytest.raises(TypeError): + make_descriptor(wire_protocol="openai_compat") + + +def test_descriptor_is_immutable() -> None: + """Descriptors are shared process-wide; a mutation would be visible to every + other reader mid-iteration.""" + descriptor = make_descriptor() + + with pytest.raises(Exception): + descriptor.default_model = "hacked" # type: ignore[misc] + + +def test_id_matching_ignores_case_and_honours_aliases() -> None: + """Provider ids come from hand-edited config files and old app versions.""" + descriptor = make_descriptor(aliases=("legacy-demo",)) + + assert descriptor.matches("DEMO") + assert descriptor.matches(" legacy-demo ") + assert not descriptor.matches("other") + + +def test_capabilities_use_the_routing_vocabulary() -> None: + """The set must be feedable straight into the routing selector's filter.""" + descriptor = make_descriptor(supports_vision=True, supports_tools=True, + supports_streaming=False) + + assert descriptor.capabilities == frozenset({"vision", "tools"}) + assert descriptor.has_capability("vision") + assert not descriptor.has_capability("streaming") + + +def test_average_cost_is_none_when_a_price_is_unknown() -> None: + """Unknown prices stay unknown — a guessed number would silently skew the + routing scorer's cost term.""" + assert make_descriptor(cost_per_1k_input=0.5).avg_cost_per_1k is None + priced = make_descriptor(cost_per_1k_input=1.0, cost_per_1k_output=3.0) + # Same 1:3 input:output weighting as ModelMetadata.avg_cost_per_1k. + assert priced.avg_cost_per_1k == pytest.approx((1.0 + 9.0) / 4.0) + + +def test_resolve_model_prefers_the_caller_then_the_default() -> None: + """One place implements the "picked model or provider default" fallback that + every chat surface used to re-implement inline.""" + descriptor = make_descriptor() + + assert descriptor.resolve_model("demo-large") == "demo-large" + assert descriptor.resolve_model("") == "demo-small" + assert descriptor.resolve_model(" ") == "demo-small" + + +def test_with_models_repoints_a_default_that_vanished() -> None: + """After discovery, the default must still name a model that exists.""" + descriptor = make_descriptor() + + updated = descriptor.with_models(["demo-v2", "demo-v2", "demo-v3"]) + + assert updated.models == ("demo-v2", "demo-v3") # de-duplicated, order kept + assert updated.default_model == "demo-v2" + assert descriptor.models == ("demo-small", "demo-large"), "original was mutated" + + +def test_with_models_keeps_a_default_that_survived() -> None: + """Discovery must not reshuffle a user's working selection.""" + updated = make_descriptor().with_models(["demo-large", "demo-small"]) + + assert updated.default_model == "demo-small" + + +# --------------------------------------------------------------------------- # +# ProviderRegistry +# --------------------------------------------------------------------------- # +def test_registry_resolves_ids_aliases_and_reports_unknowns() -> None: + """Lookup must be forgiving about form, but loud about genuinely unknown + providers — a typo should fail at the call site, not as a None later.""" + registry = ProviderRegistry([make_descriptor(aliases=("legacy-demo",))]) + + assert registry.get("demo").provider_id == "demo" + assert registry.get("legacy-demo").provider_id == "demo" + assert registry.find("missing") is None + assert "demo" in registry + with pytest.raises(ProviderNotFoundError): + registry.get("missing") + + +def test_registry_refuses_to_overwrite_silently_but_replace_works() -> None: + """A second registration of the same id is almost always a bug; updating a + descriptor is a deliberate act with its own method.""" + registry = ProviderRegistry([make_descriptor()]) + + with pytest.raises(ValueError): + registry.register(make_descriptor(display_name="Impostor")) + + registry.replace(make_descriptor(display_name="Renamed")) + assert registry.get("demo").display_name == "Renamed" + assert len(registry) == 1 + + +def test_registry_re_registering_an_identical_descriptor_is_a_no_op() -> None: + """Idempotent registration keeps repeated bootstrap calls harmless.""" + registry = ProviderRegistry([make_descriptor()]) + + registry.register(make_descriptor()) + + assert len(registry) == 1 + + +def test_find_by_model_resolves_a_bare_model_id() -> None: + """Routing decisions and saved conversations sometimes carry only a model + name; the registry is what turns that back into a provider.""" + registry = ProviderRegistry([make_descriptor()]) + + assert registry.find_by_model("demo-large").provider_id == "demo" + # A gateway model we cannot enumerate offline is a miss, not an error — the + # caller falls back to the configured active provider. + assert registry.find_by_model("unknown-model") is None + assert registry.find_by_model("") is None + + +def test_builtin_catalogue_covers_every_configured_provider() -> None: + """The catalogue and DEFAULT_CONFIG must not drift: a provider users can + configure but the registry cannot build is a dead Settings entry.""" + from cowork_local.config import DEFAULT_CONFIG + + registry = ProviderRegistry(BUILTIN_DESCRIPTORS) + + for provider_id in DEFAULT_CONFIG["providers"]: + assert registry.find(provider_id) is not None, f"{provider_id} missing from registry" + + +def test_build_fills_in_the_default_model() -> None: + """A half-written config must still produce a usable provider rather than an + empty model id that only fails once the request reaches the gateway.""" + registry = ProviderRegistry(BUILTIN_DESCRIPTORS) + + provider = registry.build("anthropic", {"api_key": "k"}) + + assert provider.model == registry.get("anthropic").default_model + + +def test_build_respects_an_explicit_model() -> None: + """Per-tab model selection must win over the catalogue default.""" + registry = ProviderRegistry(BUILTIN_DESCRIPTORS) + + provider = registry.build("anthropic", {"api_key": "k", "model": "claude-opus-4-8"}) + + assert provider.model == "claude-opus-4-8" + + +def test_factory_still_raises_provider_error_for_unknown_ids() -> None: + """Existing call sites catch ProviderError; routing lookups through the + registry must not change the exception type they see.""" + from cowork_local.providers import build_provider + from cowork_local.providers.base import ProviderError + + with pytest.raises(ProviderError): + build_provider("definitely-not-a-provider", {}) diff --git a/tests/unit/test_routing_application_service.py b/tests/unit/test_routing_application_service.py new file mode 100644 index 0000000..23f5dd3 --- /dev/null +++ b/tests/unit/test_routing_application_service.py @@ -0,0 +1,384 @@ +"""R03-T03 — unit tests for the unified routing decision rules. + +The point of moving these rules out of the three chat widgets is that they can +now be exercised without Qt, without the assessment store and without a network: +the service talks to two narrow ports, so every mode is driven here by ~10-line +fakes. Each test names the behaviour a chat surface depends on. +""" + +from __future__ import annotations + +import pytest +from cowork_local.application.model_routing import ( + RouteEvaluation, + RoutingApplicationService, + RoutingMode, + RoutingOutcome, + RoutingRequest, +) + + +class FakeDecisionPort: + """A routing engine that returns a canned verdict and records its input.""" + + def __init__(self, evaluation: RouteEvaluation) -> None: + self.evaluation = evaluation + self.calls: list = [] + + def evaluate(self, request: RoutingRequest, mode: RoutingMode) -> RouteEvaluation: + self.calls.append((request, mode)) + return self.evaluation + + +class ExplodingDecisionPort: + """An engine that fails — proves routing degrades instead of breaking a turn.""" + + def evaluate(self, request: RoutingRequest, mode: RoutingMode) -> RouteEvaluation: + raise RuntimeError("assessment store is corrupt") + + +class FakeModeResolver: + """Per-surface mode lookup, standing in for the workspace settings.""" + + def __init__(self, mode) -> None: + self.mode = mode + self.surfaces: list = [] + + def mode_for(self, surface: str): + self.surfaces.append(surface) + return self.mode + + +def make_request(**overrides) -> RoutingRequest: + """A representative turn: Cowork chat, currently on a cheap OpenAI model.""" + fields = dict( + surface="cowork", + prompt="Refactor this function", + current_provider="codex", + current_model="gpt-4o-mini", + ) + fields.update(overrides) + return RoutingRequest(**fields) + + +def switch_evaluation(**overrides) -> RouteEvaluation: + """An engine verdict that proposes a switch to a better coding model.""" + fields = dict( + task_type="coding", + should_switch=True, + target_provider="anthropic", + target_model="claude-sonnet-4-6", + score_gain=0.21, + reason="coding fit 0.88 > current 0.67", + decision=object(), + ) + fields.update(overrides) + return RouteEvaluation(**fields) + + +# --------------------------------------------------------------------------- # +# Off +# --------------------------------------------------------------------------- # +def test_off_mode_never_consults_the_engine() -> None: + """Off must be free: no ranking, no store read, no decision at all.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.OFF)) + + outcome = service.resolve(make_request()) + + assert outcome.switched is False + assert outcome.provider is None and outcome.model is None + assert port.calls == [], "Off mode must not call the routing engine" + + +def test_missing_mode_resolver_defaults_to_off() -> None: + """Routing stays opt-in: with no way to read the mode, never switch.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port) + + outcome = service.resolve(make_request()) + + assert outcome.mode is RoutingMode.OFF + assert outcome.switched is False + + +def test_empty_prompt_is_not_routed() -> None: + """An empty message carries no signal to classify, so the engine is skipped.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.AUTO)) + + outcome = service.resolve(make_request(prompt=" ")) + + assert outcome.switched is False + assert port.calls == [] + + +# --------------------------------------------------------------------------- # +# Auto +# --------------------------------------------------------------------------- # +def test_auto_mode_switches_silently() -> None: + """Auto applies the engine's verdict without asking the user.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.AUTO)) + + outcome = service.resolve(make_request()) + + assert outcome.switched is True + assert outcome.provider == "anthropic" + assert outcome.model == "claude-sonnet-4-6" + assert outcome.task_type == "coding" + assert outcome.score_gain == pytest.approx(0.21) + assert outcome.should_notify is True + + +def test_auto_mode_keeps_current_when_nothing_is_better() -> None: + """No proposed switch means the surface's own selection is untouched.""" + port = FakeDecisionPort(switch_evaluation( + should_switch=False, reason="current model is already best-fit")) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.AUTO)) + + outcome = service.resolve(make_request()) + + assert outcome.switched is False + assert outcome.provider is None + assert "already best-fit" in outcome.reason + + +def test_switch_without_a_target_is_ignored() -> None: + """A verdict that says "switch" but names nothing is not actionable — a + surface must never be handed an empty model id.""" + port = FakeDecisionPort(switch_evaluation(target_provider=None, target_model=None)) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.AUTO)) + + outcome = service.resolve(make_request()) + + assert outcome.switched is False + + +def test_same_provider_switch_keeps_the_current_provider() -> None: + """A model-only switch must not blank out the provider the surface uses.""" + port = FakeDecisionPort(switch_evaluation(target_provider=None, target_model="o3")) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.AUTO)) + + outcome = service.resolve(make_request()) + + assert outcome.switched is True + assert outcome.provider == "codex" # unchanged, from the request + assert outcome.model == "o3" + + +# --------------------------------------------------------------------------- # +# Manual +# --------------------------------------------------------------------------- # +def test_manual_mode_switches_only_after_approval() -> None: + """Manual's contract: ask first, then apply exactly what was approved.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService( + port, FakeModeResolver(RoutingMode.MANUAL), + confirm_timeout_sec=lambda: 30.0, + ) + asked: list = [] + + def confirm(decision, timeout): + asked.append((decision, timeout)) + return True + + outcome = service.resolve(make_request(), confirm=confirm) + + assert outcome.switched is True + assert len(asked) == 1 + # The configured timeout must reach the dialog, not a hard-coded default. + assert asked[0][1] == pytest.approx(30.0) + + +def test_manual_mode_decline_is_reported_distinctly() -> None: + """"The user said no" must be distinguishable from "nothing better found", + so a surface can stay quiet in one case and explain itself in the other.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.MANUAL)) + + outcome = service.resolve(make_request(), confirm=lambda decision, timeout: False) + + assert outcome.switched is False + assert outcome.declined is True + + +def test_manual_mode_without_a_callback_never_switches() -> None: + """Silently switching in Manual mode would violate the mode's promise.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.MANUAL)) + + outcome = service.resolve(make_request(), confirm=None) + + assert outcome.switched is False + + +def test_manual_mode_treats_a_broken_dialog_as_a_decline() -> None: + """A crashing confirm dialog must not auto-approve a model change.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.MANUAL)) + + def confirm(decision, timeout): + raise RuntimeError("dialog blew up") + + outcome = service.resolve(make_request(), confirm=confirm) + + assert outcome.switched is False + assert outcome.declined is True + + +# --------------------------------------------------------------------------- # +# Fallback +# --------------------------------------------------------------------------- # +def test_fallback_keeps_a_healthy_model_even_when_a_better_one_exists() -> None: + """Fallback is a resilience mode, not an optimiser: a usable pinned model + wins over a higher-scoring candidate.""" + port = FakeDecisionPort(switch_evaluation(current_is_usable=True)) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.FALLBACK)) + + outcome = service.resolve(make_request()) + + assert outcome.switched is False + assert "healthy" in outcome.reason + + +def test_fallback_switches_when_the_current_model_cannot_serve_the_turn() -> None: + """The one case Fallback exists for: rescue an unusable selection.""" + port = FakeDecisionPort(switch_evaluation(current_is_usable=False)) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.FALLBACK)) + + outcome = service.resolve(make_request()) + + assert outcome.switched is True + assert outcome.model == "claude-sonnet-4-6" + + +def test_fallback_asks_the_engine_with_auto_semantics() -> None: + """The engine only understands off/auto/manual, so Fallback must reach it as + Auto — otherwise the engine would reject the unknown mode and rank nothing.""" + port = FakeDecisionPort(switch_evaluation(current_is_usable=False)) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.FALLBACK)) + + service.resolve(make_request()) + + assert port.calls[0][1] is RoutingMode.AUTO + + +def test_fallback_never_confirms_with_the_user() -> None: + """Rescuing an unusable model is not a proposal — it happens silently.""" + port = FakeDecisionPort(switch_evaluation(current_is_usable=False)) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.FALLBACK)) + asked: list = [] + + outcome = service.resolve( + make_request(), confirm=lambda decision, timeout: asked.append(1) or True) + + assert outcome.switched is True + assert asked == [] + + +def test_fallback_with_no_replacement_keeps_current() -> None: + """Nothing to fall back to means keep going with what we have and let the + provider surface the real error, rather than blanking the model.""" + port = FakeDecisionPort(switch_evaluation( + current_is_usable=False, target_provider=None, target_model=None)) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.FALLBACK)) + + outcome = service.resolve(make_request()) + + assert outcome.switched is False + + +# --------------------------------------------------------------------------- # +# Robustness & plumbing +# --------------------------------------------------------------------------- # +def test_engine_failure_degrades_to_keep_current() -> None: + """A broken assessment store must never stop a user sending a message.""" + service = RoutingApplicationService( + ExplodingDecisionPort(), FakeModeResolver(RoutingMode.AUTO)) + + outcome = service.resolve(make_request()) + + assert isinstance(outcome, RoutingOutcome) + assert outcome.switched is False + assert "error" in outcome.reason + + +def test_mode_resolver_failure_degrades_to_off() -> None: + """An unreadable workspace config must not enable routing by accident.""" + class BrokenResolver: + def mode_for(self, surface): + raise OSError("workspace file unreadable") + + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, BrokenResolver()) + + outcome = service.resolve(make_request()) + + assert outcome.mode is RoutingMode.OFF + assert port.calls == [] + + +def test_explicit_request_mode_overrides_the_resolver() -> None: + """A surface may pin the mode for one turn (tests, replay, admin actions).""" + resolver = FakeModeResolver(RoutingMode.OFF) + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, resolver) + + outcome = service.resolve(make_request(mode=RoutingMode.AUTO)) + + assert outcome.switched is True + assert resolver.surfaces == [], "an explicit mode must skip the resolver" + + +def test_request_is_forwarded_to_the_engine_unchanged() -> None: + """Surface, prompt and pinned task type must survive the hand-off — AI-Edit + relies on its "coding" pin reaching the engine.""" + port = FakeDecisionPort(switch_evaluation()) + service = RoutingApplicationService(port, FakeModeResolver(RoutingMode.AUTO)) + request = make_request(surface="ai_edit", task_type="coding", + required_capabilities=("vision",)) + + service.resolve(request) + + forwarded = port.calls[0][0] + assert forwarded is request + assert forwarded.surface == "ai_edit" + assert forwarded.task_type == "coding" + assert forwarded.required_capabilities == ("vision",) + + +@pytest.mark.parametrize( + "raw, expected", + [ + ("auto", RoutingMode.AUTO), + ("MANUAL", RoutingMode.MANUAL), + (" fallback ", RoutingMode.FALLBACK), + ("nonsense", RoutingMode.OFF), + ("", RoutingMode.OFF), + (None, RoutingMode.OFF), + ], +) +def test_mode_parsing_is_forgiving(raw, expected) -> None: + """Config values are hand-edited; an unknown one must degrade, not raise.""" + assert RoutingMode.parse(raw) is expected + + +def test_confirm_timeout_falls_back_to_the_default_when_unusable() -> None: + """A corrupted timeout must not produce a zero-second dialog that declines + every switch before the user can read it.""" + service = RoutingApplicationService( + FakeDecisionPort(switch_evaluation()), + FakeModeResolver(RoutingMode.MANUAL), + confirm_timeout_sec=lambda: 0.0, + ) + + assert service.confirm_timeout() == RoutingApplicationService.DEFAULT_CONFIRM_TIMEOUT_SEC + + +def test_routing_request_is_immutable() -> None: + """The snapshot must not change under a turn that is already in flight.""" + request = make_request() + + with pytest.raises(Exception): + request.prompt = "something else" # type: ignore[misc] diff --git a/tests/unit/test_usage_sink.py b/tests/unit/test_usage_sink.py new file mode 100644 index 0000000..c5eec27 --- /dev/null +++ b/tests/unit/test_usage_sink.py @@ -0,0 +1,184 @@ +"""R03-T06 — unit tests for the token-usage telemetry seam. + +The seam exists so provider adapters stop owning telemetry policy. These tests +pin the two properties that makes that safe: events reach every subscriber, and +no telemetry failure can ever propagate back into the turn that produced it. +""" + +from __future__ import annotations + +import pytest +from cowork_local.infrastructure.telemetry import usage_sink +from cowork_local.infrastructure.telemetry.usage_sink import ( + CompositeUsageSink, + InMemoryUsageSink, + UsageEvent, + UsageTrackerSink, +) + + +@pytest.fixture(autouse=True) +def isolated_sink(monkeypatch): + """Give every test its own process-wide sink. + + Autouse because a leaked sink would let one test's subscriber observe the + next test's events — and, worse, let a test write to the developer's real + usage files through the default tracker sink. + """ + monkeypatch.setattr(usage_sink, "_sink", None) + yield + monkeypatch.setattr(usage_sink, "_sink", None) + + +def make_event(**overrides) -> UsageEvent: + fields = dict(provider="anthropic", model="claude-sonnet-4-6", + input_tokens=100, output_tokens=40, cached_tokens=10) + fields.update(overrides) + return UsageEvent(**fields) + + +# --------------------------------------------------------------------------- # +# UsageEvent +# --------------------------------------------------------------------------- # +def test_event_is_immutable() -> None: + """A subscriber must not be able to edit the event the next one receives.""" + event = make_event() + + with pytest.raises(Exception): + event.input_tokens = 0 # type: ignore[misc] + + +def test_total_tokens_does_not_double_count_cache_reads() -> None: + """Every gateway we support already reports cached tokens inside the input + count, so adding them again would inflate the dashboard.""" + assert make_event().total_tokens == 140 + + +def test_to_dict_uses_the_stored_row_keys() -> None: + """Matching the tracker's short keys lets a caller diff an event against a + persisted row without a translation table.""" + row = make_event(source="cowork", label="Refactor chat").to_dict() + + assert row["in"] == 100 and row["out"] == 40 and row["cache"] == 10 + assert row["source"] == "cowork" and row["label"] == "Refactor chat" + assert row["estimated"] is False + + +# --------------------------------------------------------------------------- # +# Fan-out +# --------------------------------------------------------------------------- # +def test_publish_reaches_every_subscriber() -> None: + """The whole point of the seam: extra consumers attach without patching + provider code.""" + first, second = InMemoryUsageSink(), InMemoryUsageSink() + usage_sink.set_usage_sink(CompositeUsageSink([first, second])) + + usage_sink.publish(make_event()) + + assert len(first.snapshot()) == 1 + assert len(second.snapshot()) == 1 + + +def test_one_failing_subscriber_does_not_starve_the_others() -> None: + """A buggy consumer must not silently disable the Dashboard.""" + class Exploding: + def emit(self, event): + raise RuntimeError("subscriber is broken") + + healthy = InMemoryUsageSink() + usage_sink.set_usage_sink(CompositeUsageSink([Exploding(), healthy])) + + usage_sink.publish(make_event()) + + assert len(healthy.snapshot()) == 1 + + +def test_subscribe_and_unsubscribe_round_trip() -> None: + """Teardown code calls unsubscribe unconditionally, so removing a sink that + was never added must be harmless.""" + extra = InMemoryUsageSink() + + usage_sink.subscribe(extra) + usage_sink.publish(make_event()) + usage_sink.unsubscribe(extra) + usage_sink.unsubscribe(extra) # second removal is a no-op + usage_sink.publish(make_event(model="claude-opus-4-8")) + + assert [e.model for e in extra.snapshot()] == ["claude-sonnet-4-6"] + + +def test_default_sink_is_the_usage_tracker() -> None: + """Out of the box the seam must preserve the existing Dashboard pipeline.""" + sinks = usage_sink.get_usage_sink().sinks() + + assert any(isinstance(s, UsageTrackerSink) for s in sinks) + + +def test_in_memory_sink_totals_and_clears() -> None: + """Test-double conveniences the contract suite relies on.""" + sink = InMemoryUsageSink() + sink.emit(make_event()) + sink.emit(make_event(input_tokens=1, output_tokens=1, cached_tokens=0)) + + assert sink.total_tokens == 142 + sink.clear() + assert sink.snapshot() == [] + + +# --------------------------------------------------------------------------- # +# UsageTrackerSink forwarding +# --------------------------------------------------------------------------- # +def test_tracker_sink_forwards_the_counts() -> None: + """The adapter must hand the tracker exactly what the provider measured.""" + recorded: list = [] + + def fake_record(provider, model, tokens_in, tokens_out, cached, estimated=False): + recorded.append((provider, model, tokens_in, tokens_out, cached, estimated)) + + UsageTrackerSink(recorder=fake_record).emit(make_event(estimated=True)) + + assert recorded == [("anthropic", "claude-sonnet-4-6", 100, 40, 10, True)] + + +def test_tracker_sink_restores_the_thread_context_it_borrowed() -> None: + """An event carrying its own attribution must relabel ONE row, not every + later turn that happens to run on the same worker thread.""" + from cowork_local.core import usage_tracker as tracker + + tracker.set_context("cowork", "original chat") + seen: list = [] + UsageTrackerSink(recorder=lambda *a, **k: seen.append(tracker.current_context())).emit( + make_event(source="co4e", label="flow run")) + + assert seen == [("co4e", "flow run")], "event attribution was not applied" + assert tracker.current_context() == ("cowork", "original chat") + + +def test_tracker_sink_swallows_recorder_failures() -> None: + """Telemetry is never allowed to abort an otherwise successful turn.""" + def boom(*_args, **_kwargs): + raise OSError("usage directory is read-only") + + UsageTrackerSink(recorder=boom).emit(make_event()) # must not raise + + +def test_publish_never_raises_even_with_a_broken_sink() -> None: + """Last line of defence: providers call publish() inside their stream loop.""" + class Hostile: + def emit(self, event): + raise RuntimeError("nope") + + def sinks(self): + raise RuntimeError("nope") + + usage_sink.set_usage_sink(Hostile()) + + usage_sink.publish(make_event()) # must not raise + + +def test_estimate_tokens_matches_the_tracker_heuristic() -> None: + """Re-exported so adapters need one telemetry import; it must not drift.""" + from cowork_local.core import usage_tracker as tracker + + for text in ("", "a", "hello world", "x" * 4001): + assert usage_sink.estimate_tokens(text) == tracker.estimate_tokens(text) diff --git a/ui/chat_panel.py b/ui/chat_panel.py index 9457d13..5fbff3f 100644 --- a/ui/chat_panel.py +++ b/ui/chat_panel.py @@ -638,12 +638,16 @@ class ChatPanel(QWidget): def _apply_routing(self, text: str, turn: Dict[str, Any]) -> None: """Auto Model Routing hook — run once per outgoing message. - Off → no-op. Auto → silently switch to the best-fit model. Manual → ask - the user (modal, with the configured confirm timeout) before switching. - Sets ``self._routed_provider``/``self._routed_model`` for THIS turn; - :meth:`build_provider` honours them. Never raises — a routing failure - must never block sending a message; it just falls back to the tab's - own model. + Since R03-T04 the Off/Auto/Manual/Fallback rules live in + ``application/model_routing/routing_application_service.py``; the copy + that used to sit here (and again in Co4E and AI-Edit) is gone. What + remains is the widget's own job: snapshot the tab's provider/model into + a request, host the Manual-mode modal, and render the outcome by setting + ``self._routed_provider``/``self._routed_model`` for THIS turn (honoured + by :meth:`build_provider`) plus a status bubble. + + Never raises — a routing failure must never block sending a message; it + just falls back to the tab's own model. """ # Recompute fresh each message; clear any previous turn's override. self._routed_provider = None @@ -651,33 +655,36 @@ class ChatPanel(QWidget): # An explicitly-pinned Admin agent takes precedence over routing. if getattr(self, "_admin_agent", None) is not None: return - if not (text or "").strip(): - return try: - mode = self.ctx.project_routing_mode(self.kind) # per-workspace mode - if mode == "off": - return - service = self.ctx.routing() + from ..application.model_routing import ( + RoutingRequest, + build_routing_application_service, + ) + from .routing_toggle import confirm_switch + + # The model the tab WOULD use without routing — the picker's choice, + # or the provider's configured default when nothing is picked. cur_provider = self.ctx.config.active_provider cur_model = self._model or self.ctx.config.provider_conf(cur_provider).get("model", "") - result = service.route(self.kind, text, cur_provider, cur_model, mode_override=mode) - if not result.should_switch: - return - target = result.target() - if target is None: - return - to_provider, to_model = target - if mode == "manual": - from .routing_toggle import confirm_switch - timeout = float(self.ctx.config.routing.get("confirm_timeout_sec", 60) or 60) - if not confirm_switch(self, result.decision, timeout): - return # declined / timed out → keep current model - self._routed_provider = to_provider - self._routed_model = to_model + outcome = build_routing_application_service(self.ctx).resolve( + RoutingRequest( + surface=self.kind, # per-workspace mode key ("cowork"/…) + prompt=text, + current_provider=cur_provider, + current_model=cur_model, + ), + # Manual mode only: the modal stays in the presentation layer so + # the application service never imports Qt. + confirm=lambda decision, timeout: confirm_switch(self, decision, timeout), + ) + if not outcome.switched: + return # off / nothing better / declined → keep the tab's model + self._routed_provider = outcome.provider + self._routed_model = outcome.model notice = self.chat_view.add_status(tr( "routing.switched_notice", - model=to_model, task=result.task_type.value, - gain=f"{result.decision.score_gain:.2f}")) + model=outcome.model, task=outcome.task_type, + gain=f"{outcome.score_gain:.2f}")) turn["bubbles"].append(notice) except Exception: # noqa: BLE001 — routing must never block a chat turn self._routed_provider = None diff --git a/ui/co4e_tab.py b/ui/co4e_tab.py index b829b89..f5a9049 100644 --- a/ui/co4e_tab.py +++ b/ui/co4e_tab.py @@ -1849,36 +1849,43 @@ class Co4ETab(QWidget): def _apply_co4e_routing(self, request: str) -> str: """Route this Co4E turn to the best-fit model. Returns the model id to use ('' → provider default) and sets ``self._co4e_routed_provider`` when - a cross-provider switch is chosen. Off → no-op. Manual → confirm first. - Never raises — falls back to the default model on any error.""" + a cross-provider switch is chosen. + + R03-T05: the Off/Auto/Manual/Fallback rules are no longer re-implemented + here — they come from the shared ``RoutingApplicationService``, so Co4E, + the Cowork chat and AI-Edit can never drift apart again. This method only + adapts between Co4E's state and the service's DTOs. Never raises — falls + back to the default model on any error. + """ self._co4e_routed_provider = None - if not (request or "").strip(): - return "" try: - mode = self.ctx.project_routing_mode("co4e") # per-workspace mode - if mode == "off": - return "" - service = self.ctx.routing() + from ..application.model_routing import ( + RoutingRequest, + build_routing_application_service, + ) + from .routing_toggle import confirm_switch + cur_provider = self.ctx.config.active_provider cur_model = self.ctx.config.provider_conf(cur_provider).get("model", "") - result = service.route("co4e", request, cur_provider, cur_model, mode_override=mode) - if not result.should_switch: - return "" - target = result.target() - if target is None: - return "" - to_provider, to_model = target - if mode == "manual": - from .routing_toggle import confirm_switch - timeout = float(self.ctx.config.routing.get("confirm_timeout_sec", 60) or 60) - if not confirm_switch(self, result.decision, timeout): - return "" - self._co4e_routed_provider = to_provider + outcome = build_routing_application_service(self.ctx).resolve( + RoutingRequest( + surface="co4e", + prompt=request, + current_provider=cur_provider, + current_model=cur_model, + ), + confirm=lambda decision, timeout: confirm_switch(self, decision, timeout), + ) + if not outcome.switched: + return "" # '' keeps the provider's configured default model + # Remembered so the worker's build_provider_for() can follow a + # cross-provider switch, not just a model change. + self._co4e_routed_provider = outcome.provider self._append_chat("system", tr( "routing.switched_notice", - model=to_model, task=result.task_type.value, - gain=f"{result.decision.score_gain:.2f}")) - return to_model + model=outcome.model, task=outcome.task_type, + gain=f"{outcome.score_gain:.2f}")) + return outcome.model except Exception: # noqa: BLE001 — routing must never block a Co4E turn self._co4e_routed_provider = None return "" diff --git a/ui/folder_tab.py b/ui/folder_tab.py index c5aeebe..e49f469 100644 --- a/ui/folder_tab.py +++ b/ui/folder_tab.py @@ -923,43 +923,42 @@ class FolderTab(QWidget): def _ai_apply_routing(self, instruction: str) -> None: """Auto Model Routing for the AI-Edit surface (always a CODING task). - Off → no-op. Auto → silently pick the best coding model. Manual → ask - first. Sets ``self._ai_routed_provider``/``_ai_routed_model`` for this - run; :meth:`_ai_provider` honours them. Never raises.""" + R03-T05: routes through the shared ``RoutingApplicationService`` instead + of repeating the Off/Auto/Manual/Fallback rules locally. Sets + ``self._ai_routed_provider``/``_ai_routed_model`` for this run; + :meth:`_ai_provider` honours them. Never raises.""" self._ai_routed_provider = None self._ai_routed_model = None - if not (instruction or "").strip(): - return try: - from ..core.routing.models import TaskType - mode = self.ctx.project_routing_mode("ai_edit") # per-workspace mode - if mode == "off": - return - service = self.ctx.routing() + from ..application.model_routing import ( + RoutingRequest, + build_routing_application_service, + ) + from .routing_toggle import confirm_switch + cur_provider = self.ctx.config.active_provider picked = self.ai_model_combo.currentData() if hasattr(self, "ai_model_combo") else None cur_model = picked or self.ctx.config.provider_conf(cur_provider).get("model", "") - result = service.route( - "ai_edit", instruction, cur_provider, cur_model, - mode_override=mode, task_type=TaskType.CODING, + outcome = build_routing_application_service(self.ctx).resolve( + RoutingRequest( + surface="ai_edit", + prompt=instruction, + current_provider=cur_provider, + current_model=cur_model, + # AI-Edit turns are always code edits, so the task type is + # pinned rather than classified from the instruction text. + task_type="coding", + ), + confirm=lambda decision, timeout: confirm_switch(self, decision, timeout), ) - if not result.should_switch: + if not outcome.switched: return - target = result.target() - if target is None: - return - to_provider, to_model = target - if mode == "manual": - from .routing_toggle import confirm_switch - timeout = float(self.ctx.config.routing.get("confirm_timeout_sec", 60) or 60) - if not confirm_switch(self, result.decision, timeout): - return - self._ai_routed_provider = to_provider - self._ai_routed_model = to_model + self._ai_routed_provider = outcome.provider + self._ai_routed_model = outcome.model self.ai_chat.add_status(tr( "routing.switched_notice", - model=to_model, task=result.task_type.value, - gain=f"{result.decision.score_gain:.2f}")) + model=outcome.model, task=outcome.task_type, + gain=f"{outcome.score_gain:.2f}")) except Exception: # noqa: BLE001 — routing must never block an edit self._ai_routed_provider = None self._ai_routed_model = None diff --git a/ui/routing_toggle.py b/ui/routing_toggle.py index 8f0915b..0ffc26d 100644 --- a/ui/routing_toggle.py +++ b/ui/routing_toggle.py @@ -1,12 +1,13 @@ -"""Off/Auto/Manual routing toggle + Auto-run toggle + Manual-mode confirm dialog. +"""Off/Auto/Manual/Fallback routing toggle + Auto-run toggle + confirm dialog. Dropped into every chat surface's composer (Cowork / Co4E / AI-Edit). By default a :class:`RoutingToggle` reads/writes the **per-workspace** mode via ``AppContext.project_routing_mode`` / ``set_project_routing_mode`` (so each workspace keeps its own mode), but the storage is fully injectable through ``get_mode``/``set_mode`` callables — all the real decision logic lives in -``core/routing``. Call :meth:`refresh` when the active workspace changes so the -control shows that workspace's mode. +``application/model_routing`` (which the surfaces call through +``RoutingApplicationService``). Call :meth:`refresh` when the active workspace +changes so the control shows that workspace's mode. """ from __future__ import annotations @@ -39,7 +40,7 @@ class RoutingToggle(QWidget): Emits :attr:`mode_changed`; call :meth:`refresh` after the workspace switches. """ - mode_changed = Signal(str) # "off" | "auto" | "manual" + mode_changed = Signal(str) # "off" | "auto" | "manual" | "fallback" def __init__( self, @@ -65,11 +66,14 @@ class RoutingToggle(QWidget): self._label.setObjectName("hint") self._combo = QComboBox() self._combo.setToolTip(tr("routing.toggle_tooltip")) - # (data value, i18n key) — data is the persisted mode string. + # (data value, i18n key) — data is the persisted mode string. Order is + # least-to-most autonomous, with Fallback (R03-T03) last because it is + # the "only when something breaks" mode rather than a stronger Auto. self._modes = [ ("off", "routing.mode_off"), ("auto", "routing.mode_auto"), ("manual", "routing.mode_manual"), + ("fallback", "routing.mode_fallback"), ] for value, key in self._modes: self._combo.addItem(tr(key), value)