Feature/delta team/epic r04 #7
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
@@ -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?",
|
||||
|
||||
@@ -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 "<empty registry>"
|
||||
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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
+19
-7
@@ -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
|
||||
|
||||
|
||||
+24
-17
@@ -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
|
||||
|
||||
+26
-10
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
+69
-4
@@ -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))
|
||||
# .../<checkout>/tests/conftest.py -> .../<checkout>
|
||||
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()
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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("<think>secret plan</think>Visible answer")
|
||||
|
||||
assert cleaned == "Visible answer"
|
||||
@@ -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.
|
||||
"""
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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", {})
|
||||
@@ -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]
|
||||
@@ -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)
|
||||
+35
-28
@@ -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
|
||||
|
||||
+31
-24
@@ -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 ""
|
||||
|
||||
+26
-27
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user