diff --git a/CHANGELOG.md b/CHANGELOG.md index 63e2c4b7..2fb9bf12 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -861,6 +861,10 @@ and this project uses [Semantic Versioning](https://semver.org/spec/v2.0.0.html) ### Added +- An opt-in, time-window-only SQLite observation store for measured model-group + routing, configured with `--routing-observation-window-seconds` alongside + `--state-db`; it shares completed attempt evidence across gateway processes + without inventing decay or cross-model equivalence (ADR 0042). - Verbose/debug logging (ADR 0005): a new stdlib-only `debug_logging.py` module, a `--log-level {DEBUG,INFO,WARNING,ERROR,CRITICAL}` CLI flag with a `--verbose`/`--debug` shorthand (default unchanged: `WARNING`), and new @@ -986,6 +990,18 @@ and this project uses [Semantic Versioning](https://semver.org/spec/v2.0.0.html) ### Fixed +- A blank-string `seed` or `top_logprobs` on `/v1/responses` no longer forces + the provider-only (non-streamed) execution path for `orchestrator/auto` / + `orchestrator/free`. Both omit-equivalence checks now pop the field instead + of leaving the raw blank string behind, matching the existing Chat + Completions convention and keeping the streamed-orchestration decision + (`_responses_virtual_requires_provider_path`) honest about what was + actually supplied. +- The opt-in durable routing-observation store now prunes rows by the shared + database's largest registered routing-observation window, not by whichever + writer happens to have the shortest local window. Mixed-window gateway + processes therefore keep physically bounded storage without erasing active + evidence required by a longer-window peer. - `TaskOrchestrator._invoke`'s route/Conduct primary chat call now classifies a `ProviderUpstreamError` (5xx, 429, network) directly from its own already-computed `retryable` flag (`tool_fallback.classify_provider_transport_failure`) diff --git a/README.md b/README.md index 6eedf575..4506e721 100644 --- a/README.md +++ b/README.md @@ -81,6 +81,7 @@ HTTP serving is hardened for local lab use: - `/healthz` is a minimal unauthenticated process probe; use the administrator-authenticated `/readyz` endpoint for secret-free orchestration, sync-routing, and optional batch dependency status. Liveness stays available during optional dependency degradation. - Full orchestration traces are not returned by default. Set `include_orchestration_trace: true` per standard Chat request or start with `--expose-trace-by-default` when the caller is trusted. Requests that also use `tools` or `response_format` still fail with `unsupported_trace_disclosure`; remove the trace flag for structured or tool requests. - State is in-memory by default. Pass `--state-db PATH` (or `CONTEXTUAL_ORCHESTRATOR_STATE_DB`) to persist workflow runs, evaluation runs, audit, and analytics to a stdlib sqlite file so they survive a restart; without it, behavior is unchanged. +- Model-group routing observations remain process-local unless an operator explicitly passes `--routing-observation-window-seconds SECONDS` with `--state-db PATH`. That opt-in SQLite store shares only completed success/failure and measured latency/token observations inside the selected time window; it applies no decay or cross-model quality weighting. Replay remains per-router by its configured window, while transactional pruning uses the shared database's largest registered routing-observation window so a short-window process does not delete evidence still required by a longer-window peer. Non-stream success-observation persistence errors fail closed; provider-failure recording preserves the original failure and logs the durable-evidence outage. If a stream has already emitted provider bytes, a write failure is logged and cannot change the completed response. - Response caching is off by default. Pass `--cache-ttl SECONDS` to serve identical requests (same messages + mode) from an in-memory TTL+LRU cache and skip the provider calls; `0` disables it. - `ModelClient.batch_chat(agent, {custom_id: messages})` runs many requests through the provider's Batch API (async, 24h completion window, typically ~50% cheaper) — suited to evaluation/benchmark workloads, not latency-sensitive chat. The mock path answers synchronously. diff --git a/contextual_orchestrator/__init__.py b/contextual_orchestrator/__init__.py index 9640f6a0..cbab6cc4 100644 --- a/contextual_orchestrator/__init__.py +++ b/contextual_orchestrator/__init__.py @@ -88,6 +88,11 @@ ResponseCacheProvider, build_response_cache_key, ) +from .routing_observation_store import ( + RoutingObservation, + RoutingObservationStore, + SqliteRoutingObservationStore, +) from .tool_fallback import ( MAX_TOOL_RETRY_ATTEMPTS, ToolExecutionError, @@ -149,6 +154,9 @@ "ResponseCacheProvider", "RedisResponseCacheProvider", "build_response_cache_key", + "RoutingObservation", + "RoutingObservationStore", + "SqliteRoutingObservationStore", # routing / batch "RoutingPolicy", "RoutingHints", diff --git a/contextual_orchestrator/__main__.py b/contextual_orchestrator/__main__.py index 9c6150bb..a6f91751 100644 --- a/contextual_orchestrator/__main__.py +++ b/contextual_orchestrator/__main__.py @@ -945,6 +945,12 @@ def main(argv: list[str] | None = None) -> None: parser.add_argument("--agents", default="examples/agents.mock.json", help="Agent config JSON.") parser.add_argument("--state-db", default=os.environ.get("CONTEXTUAL_ORCHESTRATOR_STATE_DB", "") or None, help="Optional sqlite path to persist runs/audit/analytics across restarts (default: in-memory).") + parser.add_argument( + "--routing-observation-window-seconds", + type=_positive_int, + default=None, + help="Persist measured routing observations in state-db for this explicit time window; decay is not applied.", + ) parser.add_argument("--mode", choices=["auto", "route", "conduct"], default="auto") parser.add_argument("--serve", action="store_true", help="Run the chat completions HTTP server.") parser.add_argument( @@ -1056,6 +1062,8 @@ def main(argv: list[str] | None = None) -> None: ) _add_log_level_arguments(parser) args = parser.parse_args(arguments) + if args.routing_observation_window_seconds is not None and not args.state_db: + parser.error("--routing-observation-window-seconds requires --state-db") client = ModelClient( ca_bundle=args.provider_ca_bundle, @@ -1069,6 +1077,7 @@ def main(argv: list[str] | None = None) -> None: load_agents(args.agents), client=client, state_db=args.state_db, + routing_observation_window_seconds=args.routing_observation_window_seconds, agents_db=args.agents_db, budget_max_output_tokens=args.budget_max_output_tokens, budget_max_cost_usd=args.budget_max_cost_usd, diff --git a/contextual_orchestrator/model_group.py b/contextual_orchestrator/model_group.py index 590a2e2f..0a0066fe 100644 --- a/contextual_orchestrator/model_group.py +++ b/contextual_orchestrator/model_group.py @@ -35,6 +35,7 @@ from collections.abc import Callable from .conventions import require_object_name +from .routing_observation_store import RoutingObservationStore #: Smoothing gain of the latency EWMA (Jacobson 1988 uses alpha = 1/8). EWMA_LATENCY_GAIN = 0.125 @@ -58,6 +59,10 @@ _GROUP_NAME_NORMALIZE_RE = re.compile(r"[-\s]+") +class RoutingObservationPersistenceError(RuntimeError): + """A configured durable routing observation could not be recorded.""" + + def canonical_group_name(raw_name: str) -> str: """Normalize a group alias to its canonical snake_case name. @@ -75,7 +80,12 @@ def canonical_group_name(raw_name: str) -> str: class ModelGroupRouter: - """Thread-safe, in-memory measured-performance ledger for group members.""" + """Thread-safe measured-performance ledger for model-group members. + + The default ledger is process-local. Supplying a + :class:`RoutingObservationStore` makes the same measured observations + visible to other gateway processes inside an explicit time window. + """ def __init__( self, @@ -83,7 +93,10 @@ def __init__( ewma_gain: float = EWMA_LATENCY_GAIN, min_latency_seconds: float = MIN_ROUTING_LATENCY_SECONDS, prior_resolver: Callable[[str], tuple[float, float]] | None = None, - clock: Callable[[], float] = time.monotonic, + observation_context_resolver: Callable[[str], str] | None = None, + observation_store: RoutingObservationStore | None = None, + ledger_name: str = "transport", + clock: Callable[[], float] | None = None, ) -> None: if not 0 < ewma_gain <= 1: raise ValueError("ewma_gain must be within (0, 1]") @@ -92,8 +105,30 @@ def __init__( self._ewma_gain = float(ewma_gain) self._min_latency_seconds = float(min_latency_seconds) self._prior_resolver = prior_resolver + self._observation_context_resolver = observation_context_resolver + if observation_store is not None and any( + not callable(getattr(observation_store, name, None)) + for name in ("append", "load", "delete_members") + ): + raise TypeError("observation_store must implement the routing observation contract") + if type(ledger_name) is not str or not ledger_name.strip(): + raise ValueError("ledger_name must be a non-empty string") + store_clock = getattr(observation_store, "now", None) + if clock is None and callable(store_clock): + clock = store_clock + if clock is None: + clock = time.time + if not callable(clock): + raise TypeError("clock must be callable") + self._observation_store = observation_store + self._ledger_name = ledger_name.strip() self._clock = clock self._lock = threading.Lock() + # Store I/O stays serialized with in-memory updates under the same lock + # ordering so refresh=False reads never observe a partially refreshed row set. + self._observation_io_lock = threading.Lock() + self._member_contexts: dict[str, str] = {} + self._retired_member_contexts: dict[str, str] = {} # member_id -> {"alpha", "beta", "ewma", "ewma_tps"}; ewma/ewma_tps are # None until the first observation of each kind arrives. self._members: dict[str, dict[str, float | None]] = {} @@ -104,6 +139,8 @@ def __init__( def register_member(self, member_id: str) -> None: """Ensure a member exists in the ledger (idempotent, keeps history).""" with self._lock: + self._retired_member_contexts.pop(member_id, None) + self._member_contexts[member_id] = self._resolve_context_key(member_id) self._members.setdefault(member_id, self._blank_state(member_id)) def _blank_state(self, member_id: str) -> dict[str, float | None]: @@ -121,24 +158,67 @@ def _blank_state(self, member_id: str) -> dict[str, float | None]: "ewma_tps": None, } + @staticmethod + def _float_value(value: float | None, default: float = 0.0) -> float: + """Read an optional numeric state value without changing valid zeroes.""" + return default if value is None else float(value) + def forget_members(self, keep_member_ids: set[str]) -> None: """Drop ledger rows for members that left every group.""" + if self._observation_store is None: + with self._lock: + for member_id in list(self._members): + if member_id not in keep_member_ids: + self._retired_member_contexts[member_id] = self._member_contexts.get( + member_id, member_id + ) + del self._members[member_id] + self._member_contexts.pop(member_id, None) + self._minute_observations.pop(member_id, None) + self._max_observed_rpm.pop(member_id, None) + self._max_observed_tpm.pop(member_id, None) + return with self._lock: - for member_id in list(self._members): - if member_id not in keep_member_ids: - del self._members[member_id] - self._minute_observations.pop(member_id, None) - self._max_observed_rpm.pop(member_id, None) - self._max_observed_tpm.pop(member_id, None) + with self._observation_io_lock: + removed = set(self._members) - keep_member_ids + self._observation_store.delete_members(self._ledger_name, removed) + for member_id in removed: + if member_id not in keep_member_ids: + self._retired_member_contexts[member_id] = self._member_contexts.get( + member_id, member_id + ) + self._members.pop(member_id, None) + self._member_contexts.pop(member_id, None) + self._minute_observations.pop(member_id, None) + self._max_observed_rpm.pop(member_id, None) + self._max_observed_tpm.pop(member_id, None) def reset_members(self, member_ids: set[str]) -> None: """Discard measurements whose group context changed.""" + if self._observation_store is None: + with self._lock: + for member_id in member_ids: + self._retired_member_contexts[member_id] = self._member_contexts.get( + member_id, member_id + ) + self._members.pop(member_id, None) + self._member_contexts.pop(member_id, None) + self._minute_observations.pop(member_id, None) + self._max_observed_rpm.pop(member_id, None) + self._max_observed_tpm.pop(member_id, None) + return with self._lock: - for member_id in member_ids: - self._members.pop(member_id, None) - self._minute_observations.pop(member_id, None) - self._max_observed_rpm.pop(member_id, None) - self._max_observed_tpm.pop(member_id, None) + with self._observation_io_lock: + self._observation_store.delete_members(self._ledger_name, member_ids) + for member_id in member_ids: + self._retired_member_contexts[member_id] = self._member_contexts.get( + member_id, member_id + ) + self._members.pop(member_id, None) + self._member_contexts.pop(member_id, None) + self._minute_observations.pop(member_id, None) + self._max_observed_rpm.pop(member_id, None) + self._max_observed_tpm.pop(member_id, None) def update_prior( self, @@ -167,8 +247,10 @@ def update_prior( raise ValueError(f"{name} must be finite and non-negative") with self._lock: state = self._ensure_locked(member_id) - old_alpha = float(state.get("prior_alpha") or 0.0) - old_beta = float(state.get("prior_beta") or 0.0) + if state is None: + return + old_alpha = self._float_value(state.get("prior_alpha")) + old_beta = self._float_value(state.get("prior_beta")) delta_alpha = float(prior_alpha) - old_alpha delta_beta = float(prior_beta) - old_beta state["prior_alpha"] = float(prior_alpha) @@ -182,6 +264,9 @@ def observe_success( latency_seconds: float | None, output_tokens: int | None = None, total_tokens: int | None = None, + *, + observation_context_key: str | None = None, + observed_at: float | None = None, ) -> None: """Record one successful attempt, with wall-clock latency when honestly known. @@ -197,8 +282,10 @@ def observe_success( duration that would not honestly describe this one attempt. """ clamped: float | None + latency: float | None if latency_seconds is None: clamped = None + latency = None else: if isinstance(latency_seconds, bool) or not isinstance(latency_seconds, (int, float)): raise TypeError("latency_seconds must be a real number") @@ -228,92 +315,313 @@ def observe_success( raise ValueError("output_tokens must be representable as a finite float") from None if not math.isfinite(throughput_sample): raise ValueError("output_tokens must be representable as a finite float") + when = self._resolve_observed_at(observed_at) with self._lock: - state = self._ensure_locked(member_id) - state["alpha"] = float(state["alpha"]) + 1.0 - if clamped is not None: - ewma = state["ewma"] - state["ewma"] = ( - clamped - if ewma is None - else (1.0 - self._ewma_gain) * float(ewma) + self._ewma_gain * clamped + with self._observation_io_lock: + self._persist_observation( + member_id, + observation_context_key=observation_context_key, + observed_at=when, + success=True, + latency_seconds=latency, + output_tokens=output_tokens, ) - if throughput_sample is not None: - tps = state["ewma_tps"] - state["ewma_tps"] = ( - throughput_sample - if tps is None - else (1.0 - self._ewma_gain) * float(tps) - + self._ewma_gain * throughput_sample + state = self._ensure_locked( + member_id, + observation_context_key=observation_context_key, ) - observations = self._minute_observations.setdefault(member_id, deque()) - now = self._clock() - observations.append((now, total_tokens)) - cutoff = now - RATE_OBSERVATION_WINDOW_SECONDS - while observations and observations[0][0] <= cutoff: - observations.popleft() - self._max_observed_rpm[member_id] = max( - self._max_observed_rpm.get(member_id, 0), len(observations) - ) - self._max_observed_tpm[member_id] = max( - self._max_observed_tpm.get(member_id, 0), - sum(tokens or 0 for _, tokens in observations), - ) - - def observe_failure(self, member_id: str) -> None: + if state is not None: + self._apply_success_locked(state, clamped, throughput_sample) + observations = self._minute_observations.setdefault(member_id, deque()) + observations.append((when, total_tokens)) + cutoff = when - RATE_OBSERVATION_WINDOW_SECONDS + while observations and observations[0][0] <= cutoff: + observations.popleft() + self._max_observed_rpm[member_id] = max( + self._max_observed_rpm.get(member_id, 0), len(observations) + ) + self._max_observed_tpm[member_id] = max( + self._max_observed_tpm.get(member_id, 0), + sum(tokens or 0 for _, tokens in observations), + ) + + def observe_failure( + self, + member_id: str, + *, + observation_context_key: str | None = None, + observed_at: float | None = None, + ) -> None: """Record one failed attempt (stability evidence only; no latency).""" + when = self._resolve_observed_at(observed_at) with self._lock: - state = self._ensure_locked(member_id) - state["beta"] = float(state["beta"]) + 1.0 + with self._observation_io_lock: + self._persist_observation( + member_id, + observation_context_key=observation_context_key, + observed_at=when, + success=False, + ) + state = self._ensure_locked( + member_id, + observation_context_key=observation_context_key, + ) + if state is not None: + self._apply_failure_locked(state) + + def _resolve_observed_at(self, observed_at: float | None) -> float: + """Resolve one observation timestamp before any router lock can delay it.""" + if observed_at is None: + when = float(self._clock()) + else: + when = float(observed_at) + if not math.isfinite(when): + raise ValueError("observed_at must be finite") + return when + + def _persist_observation( + self, + member_id: str, + *, + observation_context_key: str | None = None, + observed_at: float | None = None, + success: bool, + latency_seconds: float | None = None, + output_tokens: int | None = None, + ) -> None: + """Persist one observation while keeping storage failures identifiable.""" + if self._observation_store is None: + return + when = self._resolve_observed_at(observed_at) + try: + self._observation_store.append( + self._ledger_name, + member_id, + context_key=( + observation_context_key + if observation_context_key is not None + else self._context_key_locked( + member_id, + observation_context_key=observation_context_key, + ) + ), + observed_at=when, + success=success, + latency_seconds=latency_seconds, + output_tokens=output_tokens, + ) + except Exception as exc: # noqa: BLE001 - preserve the durable boundary type + raise RoutingObservationPersistenceError( + "durable routing observation could not be recorded" + ) from exc + + def refresh(self) -> None: + """Reload current-window observations from the shared store.""" + if self._observation_store is None: + return + with self._lock: + with self._observation_io_lock: + observations = self._observation_store.load( + self._ledger_name, + active_contexts=dict(self._member_contexts), + ) + # ponytail: replay the bounded window for cross-process + # correctness; add a sequence cursor only after measured fleet + # load requires it. + member_ids = tuple(self._members) + prior_by_member = { + member_id: ( + self._float_value( + state.get("prior_alpha"), BETA_PRIOR_SUCCESS_COUNT + ), + self._float_value( + state.get("prior_beta"), BETA_PRIOR_FAILURE_COUNT + ), + ) + for member_id, state in self._members.items() + } + rebuilt = { + member_id: self._blank_state(member_id) for member_id in member_ids + } + for member_id, (prior_alpha, prior_beta) in prior_by_member.items(): + state = rebuilt[member_id] + state["alpha"] = prior_alpha + state["beta"] = prior_beta + state["prior_alpha"] = prior_alpha + state["prior_beta"] = prior_beta + for observation in observations: + if observation.member_id not in rebuilt: + continue + state = rebuilt[observation.member_id] + if observation.success: + if observation.latency_seconds is None: + continue + latency = max(float(observation.latency_seconds), self._min_latency_seconds) + throughput = ( + None + if observation.output_tokens is None + else float(observation.output_tokens) / latency + ) + self._apply_success_locked(state, latency, throughput) + else: + self._apply_failure_locked(state) + self._members = rebuilt def member_score(self, member_id: str) -> float: """Return the expected successful responses per second for a member.""" + self.refresh() with self._lock: return self._score_locked(member_id) - def member_observation_count(self, member_id: str) -> int: + def member_observation_count(self, member_id: str, *, refresh: bool = True) -> int: """Total completed attempts recorded for one member (success + failure).""" + if refresh: + self.refresh() with self._lock: state = self._members.get(member_id) if state is None: return 0 - alpha = float(state["alpha"]) - float(state.get("prior_alpha", BETA_PRIOR_SUCCESS_COUNT)) - beta = float(state["beta"]) - float(state.get("prior_beta", BETA_PRIOR_FAILURE_COUNT)) + alpha = self._float_value(state["alpha"]) - self._float_value( + state.get("prior_alpha"), BETA_PRIOR_SUCCESS_COUNT + ) + beta = self._float_value(state["beta"]) - self._float_value( + state.get("prior_beta"), BETA_PRIOR_FAILURE_COUNT + ) return int(max(alpha, 0.0)) + int(max(beta, 0.0)) - def ranked_member_ids(self, member_ids: list[str] | tuple[str, ...]) -> list[str]: + def ranked_member_ids( + self, + member_ids: list[str] | tuple[str, ...], + *, + refresh: bool = True, + ) -> list[str]: """Order member ids best-first by measured score, preserving input ties.""" - scored = {member_id: self.member_score(member_id) for member_id in member_ids} - return sorted(member_ids, key=lambda member_id: -scored[member_id]) + if refresh: + self.refresh() + with self._lock: + scored = {member_id: self._score_locked(member_id) for member_id in member_ids} + return sorted(member_ids, key=lambda member_id: -scored[member_id]) - def member_report(self, member_id: str) -> dict[str, float | int | None]: + def member_report( + self, member_id: str, *, refresh: bool = True + ) -> dict[str, float | int | None]: """One member's measured evidence row for admin/analytics surfaces.""" + if refresh: + self.refresh() with self._lock: return self._report_locked(member_id) - def snapshot(self) -> dict[str, dict[str, float | int | None]]: + def snapshot( + self, *, refresh: bool = True + ) -> dict[str, dict[str, float | int | None]]: """Copy of every member's report keyed by member id.""" + if refresh: + self.refresh() with self._lock: return {member_id: self._report_locked(member_id) for member_id in self._members} # --- internal helpers (callers must hold ``self._lock``) --------------- - def _ensure_locked(self, member_id: str) -> dict[str, float | None]: + def _ensure_locked( + self, + member_id: str, + *, + observation_context_key: str | None = None, + ) -> dict[str, float | None] | None: + current_context = self._member_contexts.get(member_id) + if current_context is None: + if member_id in self._retired_member_contexts: + return None + self._member_contexts[member_id] = ( + observation_context_key + if observation_context_key is not None + else self._resolve_context_key(member_id) + ) + elif observation_context_key is not None and observation_context_key != current_context: + return None return self._members.setdefault(member_id, self._blank_state(member_id)) + def _resolve_context_key(self, member_id: str) -> str: + context_key = ( + self._observation_context_resolver(member_id) + if self._observation_context_resolver is not None + else member_id + ) + if type(context_key) is not str or not context_key: + raise ValueError("observation_context_resolver must return a non-empty string") + return context_key + + def _context_key_locked( + self, + member_id: str, + *, + observation_context_key: str | None = None, + ) -> str: + if member_id not in self._member_contexts: + if member_id in self._retired_member_contexts: + return ( + observation_context_key + if observation_context_key is not None + else self._resolve_context_key(member_id) + ) + self._member_contexts[member_id] = ( + observation_context_key + if observation_context_key is not None + else self._resolve_context_key(member_id) + ) + return self._member_contexts[member_id] + + def _apply_success_locked( + self, + state: dict[str, float | None], + latency: float | None, + throughput_sample: float | None, + ) -> None: + """Apply one already-validated success while the router lock is held. + + ``latency`` is ``None`` when this success carries no honest + single-attempt wall-clock timing (for example, one shared Batch API + call covering several answers) -- stability evidence (``alpha``) is + still recorded, but the latency EWMA is left untouched rather than + updated from a duration that would not honestly describe this one + attempt. + """ + state["alpha"] = self._float_value(state["alpha"]) + 1.0 + if latency is not None: + ewma = state["ewma"] + state["ewma"] = ( + latency + if ewma is None + else (1.0 - self._ewma_gain) * self._float_value(ewma) + + self._ewma_gain * latency + ) + if throughput_sample is not None: + tps = state["ewma_tps"] + state["ewma_tps"] = ( + throughput_sample + if tps is None + else (1.0 - self._ewma_gain) * self._float_value(tps) + + self._ewma_gain * throughput_sample + ) + + @staticmethod + def _apply_failure_locked(state: dict[str, float | None]) -> None: + """Apply one already-validated failure while the router lock is held.""" + state["beta"] = ModelGroupRouter._float_value(state["beta"]) + 1.0 + def _score_locked(self, member_id: str) -> float: state = self._members.get(member_id) if state is None: return UNOBSERVED_MEMBER_SCORE - alpha = float(state["alpha"]) - beta = float(state["beta"]) + alpha = self._float_value(state["alpha"]) + beta = self._float_value(state["beta"]) stability = alpha / (alpha + beta) ewma = state["ewma"] if ewma is None: # Unobserved members share one neutral reference latency so their # scores are identical and static ordering survives untouched. return stability / 1.0 - return stability / max(float(ewma), self._min_latency_seconds) + return stability / max(self._float_value(ewma), self._min_latency_seconds) def _report_locked(self, member_id: str) -> dict[str, float | int | None]: state = self._members.get(member_id) @@ -329,20 +637,26 @@ def _report_locked(self, member_id: str) -> dict[str, float | int | None]: "failure_count": 0, "score": UNOBSERVED_MEMBER_SCORE, } - alpha = float(state["alpha"]) - beta = float(state["beta"]) + alpha = self._float_value(state["alpha"]) + beta = self._float_value(state["beta"]) ewma = state["ewma"] ewma_tps = state["ewma_tps"] return { "success_posterior_mean": round(alpha / (alpha + beta), 6), - "ewma_latency_seconds": None if ewma is None else round(float(ewma), 6), + "ewma_latency_seconds": None if ewma is None else round(self._float_value(ewma), 6), "ewma_tokens_per_second": ( - None if ewma_tps is None else round(float(ewma_tps), 6) + None if ewma_tps is None else round(self._float_value(ewma_tps), 6) + ), + "success_count": int( + alpha + - self._float_value(state.get("prior_alpha"), BETA_PRIOR_SUCCESS_COUNT) + ), + "failure_count": int( + beta + - self._float_value(state.get("prior_beta"), BETA_PRIOR_FAILURE_COUNT) ), "max_observed_rpm": self._max_observed_rpm.get(member_id, 0), "max_observed_tpm": self._max_observed_tpm.get(member_id, 0), "rate_observation_window_seconds": int(RATE_OBSERVATION_WINDOW_SECONDS), - "success_count": int(alpha - float(state.get("prior_alpha", BETA_PRIOR_SUCCESS_COUNT))), - "failure_count": int(beta - float(state.get("prior_beta", BETA_PRIOR_FAILURE_COUNT))), "score": round(self._score_locked(member_id), 9), } diff --git a/contextual_orchestrator/orchestrator.py b/contextual_orchestrator/orchestrator.py index 9d168b3b..0892fec6 100644 --- a/contextual_orchestrator/orchestrator.py +++ b/contextual_orchestrator/orchestrator.py @@ -44,7 +44,11 @@ from .conventions import require_object_name from .credentials import NotConfigured, get_credential from .release_authorization import evaluate_release_authorization -from .model_group import ModelGroupRouter, canonical_group_name +from .model_group import ( + ModelGroupRouter, + RoutingObservationPersistenceError, + canonical_group_name, +) from .openrouter_uptime import OpenRouterUptimeCollector from .benchmark_priors import resolve_quality_prior from .endpoint_race import EndpointAttempt, EndpointEquivalenceContract, race_first_valid @@ -81,6 +85,7 @@ ) from .response_cache import ResponseCacheProvider, build_response_cache_key from .psychometric_routing import PsychometricRoutingEvidence +from .routing_observation_store import SqliteRoutingObservationStore from .reasoning_effort_profile import ( ReasoningEffortProfile, apply_request_profile, @@ -186,6 +191,8 @@ def _request_endpoint_partition() -> str: "ValueError", }) +_LOGGER = logging.getLogger(__name__) + def _safe_provider_probe_error_type(exc: Exception) -> str: """Keep provider diagnostics package-owned instead of echoing exception classes.""" @@ -345,6 +352,23 @@ class FastMLSIRMJudgeComponents: format_error: type[Exception] +@dataclass(frozen=True) +class _InvocationResult: + """Provider invocation payload plus the durable-observation context used.""" + + output: str + served_id: str + served_model: str + usage: dict[str, Any] | None + observation_context_key: str + + def __iter__(self): + yield self.output + yield self.served_id + yield self.served_model + yield self.usage + + def _resolve_fast_mlsirm_components() -> FastMLSIRMJudgeComponents | None: """Resolve the fast-mlsirm adapter symbols without importing unconditionally.""" try: @@ -3619,7 +3643,9 @@ class _StateStore: def __init__(self, path: str) -> None: self._lock = threading.Lock() - self._conn = sqlite3.connect(path, check_same_thread=False) + # Keep the persistent state connection's busy wait aligned with the + # short-lived routing-observation connections sharing this file. + self._conn = sqlite3.connect(path, check_same_thread=False, timeout=30.0) try: self._conn.execute("BEGIN IMMEDIATE") self._migrate_legacy_table() @@ -3854,6 +3880,7 @@ def __init__( cache_provider: ResponseCacheProvider | None = None, role_effort_catalog: dict[str, ReasoningEffortProfile] | None = None, pii_key_name: str = DEFAULT_PII_KEY_NAME, + routing_observation_window_seconds: int | None = None, allow_empty_agents: bool = False, token_counter: Any = None, ) -> None: @@ -3867,21 +3894,49 @@ def __init__( self.agents = [agent for agent in self.candidates if not agent.disabled] if not self.agents and not allow_empty_agents: # pragma: no cover raise ValueError("at least one enabled agent is required") + if routing_observation_window_seconds is not None and ( + isinstance(routing_observation_window_seconds, bool) + or type(routing_observation_window_seconds) is not int + or routing_observation_window_seconds < 1 + ): + raise ValueError("routing_observation_window_seconds must be a positive integer") + if routing_observation_window_seconds is not None and state_db is None: + raise ValueError("routing_observation_window_seconds requires state_db") + self._routing_observation_store = ( + SqliteRoutingObservationStore( + state_db, + routing_observation_window_seconds, + start_heartbeat=False, + ) + if state_db is not None and routing_observation_window_seconds is not None + else None + ) # Measured speed/stability routing inside model groups (global: every - # selection path below funnels through _ranked_agents). Ledger state is - # process-local by design: it reflects this instance's observed traffic - # and resets on restart, never carrying stale evidence across pools. - self._group_router = ModelGroupRouter() + # selection path below funnels through _ranked_agents). The optional + # store shares only the explicitly configured time window across processes. + self._group_router = ModelGroupRouter( + observation_context_resolver=self._routing_observation_context_for_member, + observation_store=self._routing_observation_store, + ledger_name="transport", + ) # Quality ledger: identical estimator family as the transport ledger but # fed by real-time fast-mlsirm judge verdicts on final answers, so # measured accuracy -- not transport success -- steers future routing. - self._quality_router = ModelGroupRouter(prior_resolver=resolve_quality_prior) + self._quality_router = ModelGroupRouter( + prior_resolver=resolve_quality_prior, + observation_context_resolver=self._routing_observation_context_for_member, + observation_store=self._routing_observation_store, + ledger_name="quality", + ) self._psychometric_router = PsychometricRoutingEvidence( max_contexts=self.EVIDENCE_CACHE_MAX_ENTRIES ) for grouped in self.candidates: self._group_router.register_member(grouped.id) self._quality_router.register_member(grouped.id) + if self._routing_observation_store is not None: + self._group_router.refresh() + self._quality_router.refresh() self._openrouter_collector = OpenRouterUptimeCollector( self.candidates, @@ -3971,6 +4026,8 @@ def __init__( self._pii_encryptors: dict[str, PiiFieldEncryptor] = {} if self._store is not None: self._reload_state() + if self._routing_observation_store is not None: + self._routing_observation_store.start_heartbeat() def close(self) -> None: """Release optional durable resources owned by this orchestrator.""" @@ -3980,6 +4037,8 @@ def close(self) -> None: self._pool_store.close() if self._store is not None: self._store.close() + if self._routing_observation_store is not None: + self._routing_observation_store.close() @contextmanager def request_policy(self, zdr_only: bool = False): @@ -4287,16 +4346,22 @@ def proxy_completion( ) measured = bool(agent.group_name or requested_model == self.FREE_MODEL) started_at = time.perf_counter() + observation_context_key = self._routing_observation_context_for_agent(agent) try: result = self.client.proxy_send(agent, endpoint, upstream) except Exception as exc: request_too_large = _is_request_too_large_error(exc) if measured and not request_too_large: - self._group_router.observe_failure(agent.id) + self._record_group_failure( + agent.id, + observation_context_key=observation_context_key, + ) raise if measured: self._group_router.observe_success( - agent.id, time.perf_counter() - started_at + agent.id, + time.perf_counter() - started_at, + observation_context_key=observation_context_key, ) return result @@ -4352,6 +4417,7 @@ def proxy_completion( every_failure_was_request_too_large = True for candidate in candidates: started_at = time.perf_counter() + observation_context_key = self._routing_observation_context_for_agent(candidate) candidate_payload = dict(upstream) candidate_payload["model"] = candidate.model if isinstance(file_replicas, dict): @@ -4390,12 +4456,17 @@ def proxy_completion( if not (request_too_large or capability_mismatch): self._record_failure(candidate.id) if candidate.group_name and not (request_too_large or capability_mismatch): - self._group_router.observe_failure(candidate.id) + self._record_group_failure( + candidate.id, + observation_context_key=observation_context_key, + ) continue self._record_success(candidate.id) if candidate.group_name: self._group_router.observe_success( - candidate.id, time.perf_counter() - started_at + candidate.id, + time.perf_counter() - started_at, + observation_context_key=observation_context_key, ) return result if last_failure is not None and every_failure_was_request_too_large: @@ -4629,7 +4700,7 @@ def _orchestrated_provider_completion( active_profile = effort_profile or self._role_effort_profile("synthesizer") virtual_model = requested_model in { None, - "contextual-orchestrator", + self.GATEWAY_DEFAULT_MODEL, self.AUTO_MODEL, self.FREE_MODEL, } @@ -4827,8 +4898,9 @@ def send_synthesis( and not _is_request_too_large_error(exc) and not isinstance(exc, EffortProfileError) ): - self._group_router.observe_failure(final_agent.id) + self._record_group_failure_for_agent(final_agent) raise + synthesis_context_key = self._routing_observation_context_for_agent(final_agent) synthesis_output = provider_output(final_agent, raw) synthesis_step = { "id": len(workflow["trace"]), @@ -4887,8 +4959,9 @@ def send_synthesis( if not _is_request_too_large_error(exc): self._record_failure(final_agent.id) if final_agent.group_name and not _is_request_too_large_error(exc): - self._group_router.observe_failure(final_agent.id) + self._record_group_failure_for_agent(final_agent) raise + repair_context_key = self._routing_observation_context_for_agent(final_agent) repaired_output = provider_output(final_agent, repaired) repair_error = _structured_output_error(repaired_output, response_format) if repair_error is None: @@ -4912,7 +4985,7 @@ def send_synthesis( failed_agent = final_agent self._record_failure(failed_agent.id) if failed_agent.group_name: - self._group_router.observe_failure(failed_agent.id) + self._record_group_failure_for_agent(failed_agent) if not virtual_model: raise ProviderResponseError( "structured synthesis and repair violated response_format" @@ -4936,8 +5009,14 @@ def send_synthesis( synthesis_started = time.perf_counter() self._record_success(final_agent.id) if final_agent.group_name: - self._group_router.observe_success( - final_agent.id, time.perf_counter() - synthesis_started + self._record_group_success_for_agent( + final_agent, + time.perf_counter() - synthesis_started, + observation_context_key=( + repair_context_key + if repair_step is not None + else synthesis_context_key + ), ) if response_request: raw.setdefault("output_text", synthesis_output) @@ -5205,6 +5284,7 @@ def stream_route( agent = self._requested_agent(model_name) or self._select_agent( text, "worker", free_only=model_name == self.FREE_MODEL ) + observation_context_key = self._routing_observation_context_for_agent(agent) parts: list[str] = [] effort_profile = self._role_effort_profile("worker") stream_kwargs: dict[str, Any] = {} @@ -5220,27 +5300,52 @@ def stream_route( yield delta except Exception: if agent.group_name or model_name == self.FREE_MODEL: - self._group_router.observe_failure(agent.id) + self._record_group_failure( + agent.id, + observation_context_key=observation_context_key, + ) raise usage = self.client.take_usage() if hasattr(self.client, "take_usage") else None if usage_callback is not None: usage_callback(usage) if agent.group_name or model_name == self.FREE_MODEL: - self._group_router.observe_success(agent.id, time.perf_counter() - started_at) + try: + self._group_router.observe_success( + agent.id, + time.perf_counter() - started_at, + observation_context_key=observation_context_key, + ) + except RoutingObservationPersistenceError: + # ponytail: preserve an already-emitted stream; surface the + # degraded durable-evidence state through the structured log. + _LOGGER.error( + "durable routing observation failed after streamed response completion" + ) answer = "".join(parts) # Real-time judging after the stream: already-sent bytes cannot be # recalled, so the verdict never changes this response -- it feeds the # quality ledger so measured accuracy steers future member ordering, # and it is persisted for audit. latency_seconds = time.perf_counter() - started_at - verification = self._realtime_route_judge( - text=text, - answer=answer, - served_id=agent.id, - latency_seconds=latency_seconds, - usage=usage, - free_only=model_name == self.FREE_MODEL, - ) + try: + verification = self._realtime_route_judge( + text=text, + answer=answer, + served_id=agent.id, + latency_seconds=latency_seconds, + usage=usage, + free_only=model_name == self.FREE_MODEL, + observation_context_key=observation_context_key, + ) + except RoutingObservationPersistenceError: + _LOGGER.error( + "durable routing observation failed after streamed response completion" + ) + verification = { + "accepted": False, + "reason": "routing observation persistence failed after stream completion", + "verifier_output": "", + } trace_step = { "id": 0, "role": "worker", @@ -5522,6 +5627,7 @@ def batch_route(self, prompts: list[str]) -> list[dict[str, Any]]: answers: dict[int, dict[str, Any]] = {} prepared_rows: dict[int, dict[str, Any]] = {} run_ids: dict[int, str] = {} + observation_context_keys: dict[int, str] = {} for agent_id, requests in requests_by_agent.items(): # A prior group's spend is already reflected in the budget meter # by its own pending persist below, so a later group must not @@ -5551,9 +5657,11 @@ def batch_route(self, prompts: list[str]) -> list[dict[str, Any]]: result=answers[pending_index], row=prepared_rows[pending_index], run_id=run_ids[pending_index], + observation_context_key=observation_context_keys[pending_index], ) raise BudgetExceededError("spend budget exceeded", detail=budget) agent = agents_by_id[agent_id] + observation_context_key = self._routing_observation_context_for_agent(agent) effort_profile = self._role_effort_profile("worker") batch_started_at = time.perf_counter() batch = ( @@ -5612,6 +5720,7 @@ def batch_route(self, prompts: list[str]) -> list[dict[str, Any]]: ) prepared_rows[index] = row run_ids[index] = run_id + observation_context_keys[index] = observation_context_key raise for custom_id, result in results.items(): # _validate_batch_results already pinned every result key to the @@ -5638,6 +5747,7 @@ def batch_route(self, prompts: list[str]) -> list[dict[str, Any]]: ) prepared_rows[index] = row run_ids[index] = run_id + observation_context_keys[index] = observation_context_key records: list[dict[str, Any]] = [] for index, (prompt, agent) in enumerate(selected): @@ -5666,6 +5776,7 @@ def batch_route(self, prompts: list[str]) -> list[dict[str, Any]]: result=answers[index], row=prepared_rows[index], run_id=run_ids[index], + observation_context_key=observation_context_keys[index], ) ) return records @@ -5723,6 +5834,7 @@ def _finalize_batch_row( result: dict[str, Any], row: dict[str, Any], run_id: str, + observation_context_key: str, ) -> dict[str, Any]: """Judge one already-persisted pending batch row and record it as complete. @@ -5754,6 +5866,7 @@ def _finalize_batch_row( latency_seconds=None, usage=result.get("usage"), free_only=False, + observation_context_key=observation_context_key, ) row["realtime_judge"] = { "accepted": verification["accepted"], @@ -6055,7 +6168,10 @@ def get_model_group(self, group_name: str) -> dict[str, Any]: for capability in sorted(MODEL_CAPABILITIES) if any(capability in agent.tags for agent in members) }, - "members": [self._agent_to_admin_payload(self._agent(agent_id)) for agent_id in ranked_ids], + "members": [ + self._agent_to_admin_payload(self._agent(agent_id), refresh=False) + for agent_id in ranked_ids + ], } def set_model_group(self, group_name: str, member_agent_ids: list[str]) -> dict[str, Any]: @@ -6290,13 +6406,14 @@ def route_once( break tried_ids.add(candidate.id) start = time.perf_counter() - attempt_answer, attempt_served_id, _attempt_served_model, attempt_usage = self._invoke( + invocation = self._invoke( candidate, messages, text=text, role="worker", allowed_agent_ids=allowed_agent_ids, ) + attempt_answer, attempt_served_id, _attempt_served_model, attempt_usage = invocation latency_seconds = time.perf_counter() - start row = { "id": attempt_index, @@ -6322,6 +6439,7 @@ def route_once( latency_seconds=latency_seconds, usage=attempt_usage, free_only=free_only, + observation_context_key=getattr(invocation, "observation_context_key", None), prompt_context=prompt_context, ) row["realtime_judge"] = { @@ -6362,6 +6480,7 @@ def _realtime_route_judge( latency_seconds: float | None, usage: dict[str, Any] | None, free_only: bool, + observation_context_key: str | None = None, prompt_context: str | None = None, ) -> dict[str, Any]: """Judge one direct-route answer now and feed the quality ledger. @@ -6377,13 +6496,24 @@ def _realtime_route_judge( """ output_tokens = self._usage_completion_tokens(usage) - def _record(accepted: bool, irt_row: tuple[int, ...] = ()) -> None: + def _record( + accepted: bool, + irt_row: tuple[int, ...] = (), + *, + observation_context_key: str | None = None, + ) -> None: if accepted: self._quality_router.observe_success( - served_id, latency_seconds, output_tokens=output_tokens + served_id, + latency_seconds, + output_tokens=output_tokens, + observation_context_key=observation_context_key, ) else: - self._quality_router.observe_failure(served_id) + self._record_quality_failure( + served_id, + observation_context_key=observation_context_key, + ) if prompt_context is not None: self._observe_contextual_quality( prompt_context, @@ -6413,7 +6543,11 @@ def _record(accepted: bool, irt_row: tuple[int, ...] = ()) -> None: and all(type(value) is int and value in (0, 1) for value in raw_irt_row) else () ) - _record(accepted, irt_row) + _record( + accepted, + irt_row, + observation_context_key=observation_context_key, + ) return base @staticmethod @@ -6970,11 +7104,16 @@ def _measured_member_order(self, member_ids: list[str]) -> list[str]: throughput/stability ledger decides; with no evidence at all the caller's input order survives untouched. No synthetic scores. """ + self._quality_router.refresh() judged_quality = any( - self._quality_router.member_observation_count(member_id) > 0 + self._quality_router.member_observation_count(member_id, refresh=False) > 0 for member_id in member_ids ) - router = self._quality_router if judged_quality else self._group_router + if judged_quality: + router = self._quality_router + else: + self._group_router.refresh() + router = self._group_router if _LOGGER.isEnabledFor(logging.DEBUG): for member_id in member_ids: _LOGGER.debug( @@ -6983,7 +7122,7 @@ def _measured_member_order(self, member_ids: list[str]) -> list[str]: judged_quality, router.member_score(member_id), ) - return router.ranked_member_ids(member_ids) + return router.ranked_member_ids(member_ids, refresh=False) def _psychometric_order( self, candidates: list[ModelAgent], prompt_context: str | None @@ -7483,15 +7622,21 @@ def _record_race_attempt( error: BaseException | None, *, capability: str, + observation_context_key: str | None = None, ) -> None: """Share race completion evidence with normal stability/circuit ledgers.""" self._record_endpoint_attempt(endpoint_id, value, error, capability=capability) if error is not None and not _is_request_too_large_error(error): - self._group_router.observe_failure(endpoint_id) + self._record_group_failure( + endpoint_id, + observation_context_key=observation_context_key, + ) self._record_failure(endpoint_id) def _race_attempt_collector( - self, capability: str + self, + capability: str, + observation_contexts: dict[str, str] | None = None, ) -> tuple[ Callable[[str, Any | None, BaseException | None], None], Callable[[str | None], None], @@ -7512,7 +7657,15 @@ def completed( error: BaseException | None, ) -> None: self._record_race_attempt( - endpoint_id, value, error, capability=capability + endpoint_id, + value, + error, + capability=capability, + observation_context_key=( + None + if observation_contexts is None + else observation_contexts.get(endpoint_id) + ), ) if error is not None or value is None: return @@ -7560,6 +7713,10 @@ def proxy_capability( raise ValueError( "immediate_race endpoint count exceeds the supported concurrency capacity" ) + observation_contexts = { + agent.id: self._routing_observation_context_for_agent(agent) + for agent in race_members + } def call(agent: ModelAgent) -> dict[str, Any] | tuple[bytes, str]: payload = { key: value for key, value in body.items() @@ -7577,7 +7734,10 @@ def call(agent: ModelAgent) -> dict[str, Any] | tuple[bytes, str]: ) contract = EndpointEquivalenceContract(**race_members[0].endpoint_equivalence) # type: ignore[arg-type] - attempt_completed, finalize_attempts = self._race_attempt_collector(capability) + attempt_completed, finalize_attempts = self._race_attempt_collector( + capability, + observation_contexts, + ) try: outcome = race_first_valid( [ @@ -7611,11 +7771,16 @@ def call(agent: ModelAgent) -> dict[str, Any] | tuple[bytes, str]: if outcome is not None: self._record_endpoint_race(outcome, capability=capability) self._group_router.observe_success( - outcome.winner_endpoint_id, outcome.completion_ms / 1000 + outcome.winner_endpoint_id, + outcome.completion_ms / 1000, + observation_context_key=observation_contexts.get( + outcome.winner_endpoint_id + ), ) return outcome.value last_error: Exception | None = None for agent in candidates: + observation_context_key = self._routing_observation_context_for_agent(agent) payload = { key: value for key, value in body.items() @@ -7642,15 +7807,24 @@ def call(agent: ModelAgent) -> dict[str, Any] | tuple[bytes, str]: every_failure_was_request_too_large and request_too_large ) if not request_too_large: - self._group_router.observe_failure(agent.id) + self._record_group_failure( + agent.id, + observation_context_key=observation_context_key, + ) continue if selection_sink is not None: selected_result = selection_sink(agent, result) self._group_router.observe_success( - agent.id, time.perf_counter() - started_at + agent.id, + time.perf_counter() - started_at, + observation_context_key=observation_context_key, ) return selected_result - self._group_router.observe_success(agent.id, time.perf_counter() - started_at) + self._group_router.observe_success( + agent.id, + time.perf_counter() - started_at, + observation_context_key=observation_context_key, + ) return result if saw_failure and every_failure_was_request_too_large: raise ProviderRequestTooLargeError( @@ -7670,7 +7844,7 @@ def _invoke( allowed_agent_ids: set[str] | None = None, eligibility_role: str | None = None, excluded_agent_ids: set[str] | None = None, - ) -> tuple[str, str, str, dict[str, Any] | None]: + ) -> _InvocationResult: """Call an agent with bounded, safety-aware tool retry and failover. ``ModelClient`` handles provider transport retries. This layer classifies @@ -7715,6 +7889,10 @@ def _invoke( ) effort_profile = self._role_effort_profile(role) request_settings = self.client.request_settings_snapshot() + observation_contexts = { + agent.id: self._routing_observation_context_for_agent(agent) + for agent in race_members + } def call(agent: ModelAgent) -> tuple[str, str, str, dict[str, Any] | None]: with self.client.request_settings(**request_settings): @@ -7727,7 +7905,10 @@ def call(agent: ModelAgent) -> tuple[str, str, str, dict[str, Any] | None]: return output, agent.id, agent.model, usage contract = EndpointEquivalenceContract(**race_members[0].endpoint_equivalence) # type: ignore[arg-type] - attempt_completed, finalize_attempts = self._race_attempt_collector("text") + attempt_completed, finalize_attempts = self._race_attempt_collector( + "text", + observation_contexts, + ) try: outcome = race_first_valid( [ @@ -7761,8 +7942,19 @@ def call(agent: ModelAgent) -> tuple[str, str, str, dict[str, Any] | None]: outcome.winner_endpoint_id, outcome.completion_ms / 1000, output_tokens=output_tokens, + observation_context_key=observation_contexts.get( + outcome.winner_endpoint_id + ), + ) + return _InvocationResult( + output=outcome.value[0], + served_id=outcome.value[1], + served_model=outcome.value[2], + usage=outcome.value[3], + observation_context_key=observation_contexts.get( + outcome.winner_endpoint_id, outcome.winner_endpoint_id + ), ) - return outcome.value retry_limit = min(self.tool_retry_attempts, MAX_TOOL_RETRY_ATTEMPTS) bounded_provider_response_failures = 0 last_provider_response_error: ProviderResponseError | None = None @@ -7773,6 +7965,7 @@ def call(agent: ModelAgent) -> tuple[str, str, str, dict[str, Any] | None]: last_upstream_error: ProviderUpstreamError | None = None for agent in candidates: retry_attempt = 0 + observation_context_key = self._routing_observation_context_for_agent(agent) while True: try: attempt_start = time.perf_counter() @@ -7787,7 +7980,10 @@ def call(agent: ModelAgent) -> tuple[str, str, str, dict[str, Any] | None]: break every_failure_was_request_too_large = False if agent.group_name or allowed_agent_ids is not None: - self._group_router.observe_failure(agent.id) + self._record_group_failure( + agent.id, + observation_context_key=observation_context_key, + ) if isinstance(exc, ToolFallbackStoppedError): # Deliberately terminal, even inside a free/auto virtual # pool with untried candidates remaining: every path that @@ -7878,10 +8074,17 @@ def call(agent: ModelAgent) -> tuple[str, str, str, dict[str, Any] | None]: agent.id, time.perf_counter() - attempt_start, output_tokens=output_tokens, + observation_context_key=observation_context_key, total_tokens=total_tokens, ) self._record_success(agent.id) - return output, agent.id, agent.model, usage + return _InvocationResult( + output=output, + served_id=agent.id, + served_model=agent.model, + usage=usage, + observation_context_key=observation_context_key, + ) if ( last_provider_response_error is not None and bounded_provider_response_failures == len(candidates) @@ -8014,18 +8217,123 @@ def _record_failure(self, agent_id: str) -> None: self.circuit_reset_seconds, ) + def _record_group_failure( + self, + agent_id: str, + *, + observation_context_key: str | None = None, + observed_at: float | None = None, + ) -> None: + """Record provider failure without replacing an already active error.""" + try: + self._group_router.observe_failure( + agent_id, + observation_context_key=observation_context_key, + observed_at=observed_at, + ) + except RoutingObservationPersistenceError: + _LOGGER.error( + "durable routing observation failed while recording provider failure" + ) + + def _record_group_failure_for_agent( + self, + agent: ModelAgent, + *, + observed_at: float | None = None, + ) -> None: + """Record provider failure using the current agent shape as the context key.""" + self._record_group_failure( + agent.id, + observation_context_key=self._routing_observation_context_for_agent(agent), + observed_at=observed_at, + ) + + def _record_quality_failure( + self, + agent_id: str, + *, + observation_context_key: str | None = None, + observed_at: float | None = None, + ) -> None: + """Record judge failure without aborting a usable-answer failover.""" + try: + self._quality_router.observe_failure( + agent_id, + observation_context_key=observation_context_key, + observed_at=observed_at, + ) + except RoutingObservationPersistenceError: + _LOGGER.error( + "durable quality observation failed while recording provider failure" + ) + def _record_success(self, agent_id: str) -> None: with self._circuit_lock: cleared = self._circuit.pop(agent_id, None) if cleared is not None and _LOGGER.isEnabledFor(logging.DEBUG): _LOGGER.debug("circuit_cleared agent_id=%s", agent_id) + def _record_group_success_for_agent( + self, + agent: ModelAgent, + latency_seconds: float, + *, + observation_context_key: str | None = None, + observed_at: float | None = None, + ) -> None: + """Record provider success using the current agent shape as the context key.""" + self._group_router.observe_success( + agent.id, + latency_seconds, + observation_context_key=( + observation_context_key + if observation_context_key is not None + else self._routing_observation_context_for_agent(agent) + ), + observed_at=observed_at, + ) + def _agent(self, agent_id: str) -> ModelAgent: for agent in self.candidates: if agent.id == agent_id and _agent_matches_request_endpoint(agent): return agent raise KeyError(agent_id) # pragma: no cover + def _routing_observation_context_for_agent(self, agent: ModelAgent) -> str: + """Stable context key for durable routing evidence tied to one agent shape.""" + payload = { + "auth_scheme": agent.auth_scheme, + "base_url": agent.base_url, + "credential_name": agent.credential_name, + "group_name": canonical_group_name(agent.group_name) if agent.group_name else "", + "id": agent.id, + "local_credential_key": agent.local_credential_key, + "model": agent.model, + "provider_name": agent.provider_name, + } + return hashlib.sha256( + json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8") + ).hexdigest() + + def _routing_observation_context_for_member(self, agent_id: str) -> str: + try: + agent = self._agent(agent_id) + except KeyError: + payload = { + "auth_scheme": "", + "base_url": "", + "group_name": "", + "id": agent_id, + "model": "", + "provider_name": "", + "removed_from_pool": True, + } + return hashlib.sha256( + json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8") + ).hexdigest() + return self._routing_observation_context_for_agent(agent) + def _agent_in_pool(self, agent_pool_id: str, worker_agent_id: str) -> ModelAgent: """Resolve an agent only through the pool boundary it can belong to. @@ -8358,7 +8666,9 @@ def _infer_provider_name(self, base_url: str) -> str: return base_url.split("//", 1)[-1].split("/", 1)[0] return base_url # pragma: no cover - def _agent_to_admin_payload(self, agent: ModelAgent) -> dict[str, Any]: + def _agent_to_admin_payload( + self, agent: ModelAgent, *, refresh: bool = True + ) -> dict[str, Any]: return { "id": agent.id, "model": agent.model, @@ -8372,16 +8682,25 @@ def _agent_to_admin_payload(self, agent: ModelAgent) -> dict[str, Any]: "context_window": agent.context_window, "stream_usage_supported": agent.stream_usage_supported, "group_name": agent.group_name, - "group_routing": self._group_router.member_report(agent.id) if agent.group_name else None, + "group_routing": ( + self._group_router.member_report(agent.id, refresh=refresh) + if agent.group_name + else None + ), } - def list_agents(self, page_number: int = 1, page_size: int = 10) -> list[dict[str, Any]]: + def list_agents( + self, page_number: int = 1, page_size: int = 10, *, refresh: bool = True + ) -> list[dict[str, Any]]: """Return a paginated admin-safe view of configured agents.""" if page_number < 1 or page_size < 1: # pragma: no cover raise ValueError("page_number/page_size must be >= 1") start = (page_number - 1) * page_size end = start + page_size - return [self._agent_to_admin_payload(agent) for agent in self.candidates[start:end]] + return [ + self._agent_to_admin_payload(agent, refresh=refresh) + for agent in self.candidates[start:end] + ] def list_openai_models(self) -> dict[str, Any]: """Return an OpenAI-compatible ``/v1/models`` list from the agent pool. @@ -15116,15 +15435,26 @@ def admin_state( ) -> dict[str, Any]: """Build the admin console state payload from agents, policy, and audit data.""" agent_page_size = max(1, len(self.candidates)) + self._group_router.refresh() + self._quality_router.refresh() return { - "agents": self.list_agents(page_size=agent_page_size), + "agents": self.list_agents(page_size=agent_page_size, refresh=False), "policy": { **self.policy.as_dict(), "roles": list(self.ROLE_TAGS), }, "routing_evidence": { - "transport": self._group_router.snapshot(), - "quality": self._quality_router.snapshot(), + "transport": self._group_router.snapshot(refresh=False), + "quality": self._quality_router.snapshot(refresh=False), + }, + "routing_observation_policy": { + "enabled": self._routing_observation_store is not None, + "window_seconds": ( + None + if self._routing_observation_store is None + else self._routing_observation_store.window_seconds + ), + "retention_policy": "time_window_only" if self._routing_observation_store is not None else None, }, "recent_workflow_runs": [ self._shorten_run(run) diff --git a/contextual_orchestrator/provider_errors.py b/contextual_orchestrator/provider_errors.py index 3b862792..05753248 100644 --- a/contextual_orchestrator/provider_errors.py +++ b/contextual_orchestrator/provider_errors.py @@ -116,6 +116,22 @@ def provider_error_body(exc: urllib.error.HTTPError) -> bytes: return body +def _sanitize_provider_message_text(raw: object) -> str | None: + """Collapse one provider diagnostic to a bounded caller-safe sentence.""" + collapsed = "".join( + char + if char == "\t" or not (ord(char) < 0x20 or 0x7F <= ord(char) <= 0x9F) + else " " + for char in str(raw) + ).strip() + collapsed = collapsed[:MAX_SAFE_MESSAGE_CHARS] + if _SAFE_SCHEMA_DIAGNOSTIC.search(collapsed): + return _SAFE_SCHEMA_ERROR_SUMMARY + if not collapsed or _SENSITIVE_PROVIDER_MESSAGE.search(collapsed): + return None + return collapsed + + def safe_provider_message(exc: BaseException) -> str | None: """Extract one bounded, control-free diagnostic sentence from a failure. @@ -149,18 +165,7 @@ def safe_provider_message(exc: BaseException) -> str | None: else: first_arg = next((part for part in exc.args[:1] if isinstance(part, str)), "") raw = first_arg or type(exc).__name__ - collapsed = "".join( - char - if char == "\t" or not (ord(char) < 0x20 or 0x7F <= ord(char) <= 0x9F) - else " " - for char in str(raw) - ).strip() - collapsed = collapsed[:MAX_SAFE_MESSAGE_CHARS] - if _SAFE_SCHEMA_DIAGNOSTIC.search(collapsed): - return _SAFE_SCHEMA_ERROR_SUMMARY - if not collapsed or _SENSITIVE_PROVIDER_MESSAGE.search(collapsed): - return None - return collapsed + return _sanitize_provider_message_text(raw) class ProviderUpstreamError(RuntimeError): @@ -190,7 +195,10 @@ def __init__( self.provider_status = provider_status self.retryable = retryable self.transport = transport - super().__init__(message) + super().__init__( + _sanitize_provider_message_text(message) + or "provider diagnostic was redacted for safety" + ) @property def detail(self) -> dict[str, Any]: diff --git a/contextual_orchestrator/routing_observation_store.py b/contextual_orchestrator/routing_observation_store.py new file mode 100644 index 00000000..14c9159e --- /dev/null +++ b/contextual_orchestrator/routing_observation_store.py @@ -0,0 +1,462 @@ +"""Time-windowed durable observations for measured model-group routing. + +The store keeps one immutable row per completed provider attempt. A positive +operator-selected window limits which rows participate in a router's current +state; no decay, cross-model weighting, or inferred provider equivalence is +introduced here. +""" + +from __future__ import annotations + +from dataclasses import dataclass +import math +import os +import sqlite3 +import threading +import time +import uuid +from collections.abc import Callable, Iterable, Mapping +from typing import Protocol + + +@dataclass(frozen=True) +class RoutingObservation: + """One completed routing attempt restored from durable storage.""" + + member_id: str + success: bool + latency_seconds: float | None + output_tokens: int | None + + +class RoutingObservationStore(Protocol): + """Minimal store contract consumed by :class:`ModelGroupRouter`.""" + + def now(self) -> float: + """Return the store clock for default observation timestamps.""" + + def append( + self, + ledger_name: str, + member_id: str, + *, + context_key: str, + observed_at: float, + success: bool, + latency_seconds: float | None = None, + output_tokens: int | None = None, + ) -> None: + """Append one completed attempt to the current observation window.""" + + def load( + self, + ledger_name: str, + active_contexts: Mapping[str, str] | None = None, + ) -> list[RoutingObservation]: + """Return current-window observations in completion order.""" + + def delete_members(self, ledger_name: str, member_ids: Iterable[str]) -> None: + """Delete observations whose group membership is no longer valid.""" + + def close(self) -> None: + """Release store resources owned by the router.""" + + +class SqliteRoutingObservationStore: + """Share measured routing observations through a SQLite database. + + The store opens a short-lived connection for each operation so separate + gateway processes can use the same database. Rows older than the explicit + ``window_seconds`` are ignored and removed on writes. A retention window + is intentionally used instead of an invented decay coefficient. + """ + + _TABLE_NAME = "routing_observations" + _REGISTRATION_LEASE_WINDOW_MULTIPLIER = 2 + _HEARTBEAT_INTERVAL_MAX_SECONDS = 30.0 + _CREATE_TABLE_SQL = ( + "CREATE TABLE IF NOT EXISTS routing_observations (" + "observation_id INTEGER PRIMARY KEY AUTOINCREMENT, " + "ledger_name TEXT NOT NULL, member_id TEXT NOT NULL, context_key TEXT NOT NULL, " + "observed_at REAL NOT NULL, success INTEGER NOT NULL CHECK(success IN (0, 1)), " + "latency_seconds REAL, output_tokens INTEGER)" + ) + _CREATE_INDEX_SQL = ( + "CREATE INDEX IF NOT EXISTS routing_observations_ledger_time " + "ON routing_observations(ledger_name, member_id, context_key, observed_at, observation_id)" + ) + _CREATE_RETENTION_INDEX_SQL = ( + "CREATE INDEX IF NOT EXISTS routing_observations_observed_at " + "ON routing_observations(observed_at)" + ) + _METADATA_TABLE_NAME = "routing_observation_metadata" + _CREATE_METADATA_TABLE_SQL = ( + "CREATE TABLE IF NOT EXISTS routing_observation_metadata (" + "metadata_key TEXT PRIMARY KEY, metadata_value INTEGER NOT NULL)" + ) + _MAX_RETENTION_WINDOW_KEY = "max_retention_window_seconds" + _CREATE_REGISTRATIONS_TABLE_SQL = ( + "CREATE TABLE IF NOT EXISTS routing_observation_registrations (" + "registration_id TEXT PRIMARY KEY, window_seconds INTEGER NOT NULL, " + "lease_expires_at REAL NOT NULL)" + ) + + def __init__( + self, + path: str | os.PathLike[str], + window_seconds: int, + *, + clock: Callable[[], float] = time.time, + start_heartbeat: bool = True, + ) -> None: + if not isinstance(path, (str, os.PathLike)) or not str(path): + raise TypeError("path must be a non-empty filesystem path") + path_text = os.fspath(path) + if path_text == ":memory:" or ( + path_text.startswith("file:") + and (path_text.startswith("file::memory:") or "mode=memory" in path_text) + ): + raise ValueError("path must be a durable SQLite filesystem path, not an in-memory database") + if isinstance(window_seconds, bool) or type(window_seconds) is not int or window_seconds < 1: + raise ValueError("window_seconds must be a positive integer") + if not callable(clock): + raise TypeError("clock must be callable") + self._path = path_text + self._window_seconds = window_seconds + self._clock = clock + self._lock = threading.Lock() + self._registration_id = uuid.uuid4().hex + self._closed = False + self._heartbeat_stop = threading.Event() + with self._lock: + connection = self._connect() + try: + connection.execute("BEGIN IMMEDIATE") + connection.execute(self._CREATE_TABLE_SQL) + connection.execute(self._CREATE_METADATA_TABLE_SQL) + connection.execute(self._CREATE_REGISTRATIONS_TABLE_SQL) + self._ensure_schema(connection) + connection.execute(self._CREATE_INDEX_SQL) + connection.execute(self._CREATE_RETENTION_INDEX_SQL) + self._register_retention_window(connection) + if start_heartbeat: + self._refresh_registration(connection) + connection.commit() + finally: + connection.close() + self._heartbeat_thread: threading.Thread | None = None + if start_heartbeat: + self.start_heartbeat() + + def start_heartbeat(self) -> None: + """Start lease renewal after the owning service finishes initialization.""" + if self._heartbeat_thread is not None: + return + self._heartbeat_thread = threading.Thread( + target=self._heartbeat_registration, + name=f"routing-observation-heartbeat-{self._registration_id[:8]}", + daemon=True, + ) + self._heartbeat_thread.start() + + @property + def window_seconds(self) -> int: + """Return the operator-selected observation retention window.""" + return self._window_seconds + + def _connect(self) -> sqlite3.Connection: + """Open one cross-process-safe SQLite connection.""" + return sqlite3.connect(self._path, timeout=30.0) + + def _now(self) -> float: + """Return a finite wall-clock value used for the retention boundary.""" + value = float(self._clock()) + if not math.isfinite(value): + raise ValueError("clock must return a finite number") + return value + + def now(self) -> float: + """Return the caller-visible clock used for default observations.""" + return self._now() + + @staticmethod + def _ensure_schema(connection: sqlite3.Connection) -> None: + """Backfill additive columns required by the durable observation contract.""" + columns = { + str(row[1]) + for row in connection.execute("PRAGMA table_info(routing_observations)").fetchall() + } + if "context_key" not in columns: + connection.execute( + "ALTER TABLE routing_observations " + "ADD COLUMN context_key TEXT NOT NULL DEFAULT ''" + ) + registration_columns = { + str(row[1]) + for row in connection.execute( + "PRAGMA table_info(routing_observation_registrations)" + ).fetchall() + } + if "lease_expires_at" not in registration_columns: + connection.execute( + "ALTER TABLE routing_observation_registrations " + "ADD COLUMN lease_expires_at REAL NOT NULL DEFAULT 0" + ) + + def _register_retention_window(self, connection: sqlite3.Connection) -> None: + """Retain the historical maximum-window value for schema audit only.""" + connection.execute( + "INSERT INTO routing_observation_metadata(metadata_key, metadata_value) " + "VALUES(?, ?) " + "ON CONFLICT(metadata_key) DO UPDATE SET " + "metadata_value = MAX(routing_observation_metadata.metadata_value, excluded.metadata_value)", + (self._MAX_RETENTION_WINDOW_KEY, self._window_seconds), + ) + + def _refresh_registration( + self, connection: sqlite3.Connection, *, now: float | None = None + ) -> float: + """Renew this active store's bounded lease and discard crashed peers.""" + current = self._now() if now is None else now + connection.execute( + "DELETE FROM routing_observation_registrations WHERE lease_expires_at <= ?", + (current,), + ) + connection.execute( + "INSERT INTO routing_observation_registrations " + "(registration_id, window_seconds, lease_expires_at) VALUES (?, ?, ?) " + "ON CONFLICT(registration_id) DO UPDATE SET " + "window_seconds = excluded.window_seconds, " + "lease_expires_at = excluded.lease_expires_at", + ( + self._registration_id, + self._window_seconds, + current + + float( + self._window_seconds + * self._REGISTRATION_LEASE_WINDOW_MULTIPLIER + ), + ), + ) + return current + + def _heartbeat_registration(self) -> None: + """Renew this live store's lease independently of routing traffic.""" + interval = min( + float(self._window_seconds), self._HEARTBEAT_INTERVAL_MAX_SECONDS + ) + while not self._heartbeat_stop.is_set(): + with self._lock: + if self._closed: + return + connection = None + try: + connection = self._connect() + connection.execute("BEGIN IMMEDIATE") + self._refresh_registration(connection) + connection.commit() + except Exception: + if connection is not None: + connection.rollback() + finally: + if connection is not None: + connection.close() + if self._heartbeat_stop.wait(interval): + return + + def _retention_cutoff(self, connection: sqlite3.Connection) -> float: + """Return the physical prune boundary from the shared database-wide window.""" + current = self._refresh_registration(connection) + row = connection.execute( + "SELECT MAX(window_seconds) FROM routing_observation_registrations", + ).fetchone() + max_window_seconds = ( + self._window_seconds + if row is None or row[0] is None + else max(int(row[0]), 1) + ) + return current - float(max_window_seconds) + + @staticmethod + def _validate_ledger_name(ledger_name: str) -> None: + """Validate the fixed logical ledger identifier.""" + if type(ledger_name) is not str or not ledger_name.strip(): + raise ValueError("ledger_name must be a non-empty string") + + @staticmethod + def _validate_member_id(member_id: str) -> None: + """Validate an opaque member identifier before persistence.""" + if type(member_id) is not str or not member_id: + raise ValueError("member_id must be a non-empty string") + + @staticmethod + def _validate_context_key(context_key: str) -> None: + """Validate the member-context key used to reject stale rows.""" + if type(context_key) is not str or not context_key: + raise ValueError("context_key must be a non-empty string") + + def append( + self, + ledger_name: str, + member_id: str, + *, + context_key: str, + observed_at: float, + success: bool, + latency_seconds: float | None = None, + output_tokens: int | None = None, + ) -> None: + """Append one validated attempt and prune rows outside the shared retention window.""" + self._validate_ledger_name(ledger_name) + self._validate_member_id(member_id) + self._validate_context_key(context_key) + if type(success) is not bool: + raise TypeError("success must be a boolean") + latency: float | None = None + if success: + if latency_seconds is not None: + if isinstance(latency_seconds, bool) or not isinstance(latency_seconds, (int, float)): + raise TypeError("latency_seconds must be numeric when provided") + latency = float(latency_seconds) + if not math.isfinite(latency) or latency < 0: + raise ValueError("latency_seconds must be finite and nonnegative") + if output_tokens is not None and ( + isinstance(output_tokens, bool) + or type(output_tokens) is not int + or output_tokens <= 0 + ): + raise ValueError("output_tokens must be a positive integer when provided") + elif latency_seconds is not None or output_tokens is not None: + raise ValueError("failed observations cannot contain success-only measurements") + when = float(observed_at) + if not math.isfinite(when): + raise ValueError("observed_at must be finite") + connection = self._connect() + with self._lock: + try: + connection.execute("BEGIN IMMEDIATE") + connection.execute( + "INSERT INTO routing_observations " + "(ledger_name, member_id, context_key, observed_at, success, latency_seconds, output_tokens) " + "VALUES (?, ?, ?, ?, ?, ?, ?)", + ( + ledger_name.strip(), + member_id, + context_key, + when, + int(success), + latency, + None if not success else output_tokens, + ), + ) + connection.execute( + "DELETE FROM routing_observations WHERE observed_at < ?", + (self._retention_cutoff(connection),), + ) + connection.commit() + except Exception: + connection.rollback() + raise + finally: + connection.close() + + def load( + self, + ledger_name: str, + active_contexts: Mapping[str, str] | None = None, + ) -> list[RoutingObservation]: + """Return only observations still inside this router's configured replay window.""" + self._validate_ledger_name(ledger_name) + cutoff = self._now() - self._window_seconds + active = dict(active_contexts or {}) + for member_id, context_key in active.items(): + self._validate_member_id(member_id) + self._validate_context_key(context_key) + connection = self._connect() + with self._lock: + try: + connection.execute("BEGIN IMMEDIATE") + if active: + placeholders = ", ".join(["(?, ?)"] * len(active)) + params: list[object] = [ledger_name.strip(), cutoff] + for member_id, context_key in active.items(): + params.extend((member_id, context_key)) + rows = connection.execute( + "SELECT member_id, success, latency_seconds, output_tokens " + "FROM routing_observations " + "WHERE ledger_name = ? AND observed_at >= ? " + f"AND (member_id, context_key) IN ({placeholders}) " + "ORDER BY observed_at, observation_id", + tuple(params), + ).fetchall() + else: + rows = connection.execute( + "SELECT member_id, success, latency_seconds, output_tokens " + "FROM routing_observations " + "WHERE ledger_name = ? AND observed_at >= ? " + "ORDER BY observed_at, observation_id", + (ledger_name.strip(), cutoff), + ).fetchall() + connection.commit() + except Exception: + connection.rollback() + raise + finally: + connection.close() + return [ + RoutingObservation( + member_id=row[0], + success=bool(row[1]), + latency_seconds=row[2], + output_tokens=row[3], + ) + for row in rows + ] + + def delete_members(self, ledger_name: str, member_ids: Iterable[str]) -> None: + """Remove stale group-context observations without touching other ledgers.""" + self._validate_ledger_name(ledger_name) + members = tuple(dict.fromkeys(member_ids)) + for member_id in members: + self._validate_member_id(member_id) + if not members: + return + connection = self._connect() + with self._lock: + try: + connection.execute("BEGIN IMMEDIATE") + self._refresh_registration(connection) + for member_id in members: + connection.execute( + "DELETE FROM routing_observations WHERE ledger_name = ? AND member_id = ?", + (ledger_name.strip(), member_id), + ) + connection.commit() + except Exception: + connection.rollback() + raise + finally: + connection.close() + + def close(self) -> None: + """Release this store's retention-window registration.""" + self._heartbeat_stop.set() + if self._heartbeat_thread is not None: + self._heartbeat_thread.join() + with self._lock: + if self._closed: + return + connection = self._connect() + try: + connection.execute( + "DELETE FROM routing_observation_registrations " + "WHERE registration_id = ?", + (self._registration_id,), + ) + connection.commit() + self._closed = True + finally: + connection.close() + + +__all__ = ["RoutingObservation", "RoutingObservationStore", "SqliteRoutingObservationStore"] diff --git a/contextual_orchestrator/server.py b/contextual_orchestrator/server.py index eb6a7751..579e485f 100644 --- a/contextual_orchestrator/server.py +++ b/contextual_orchestrator/server.py @@ -18,7 +18,6 @@ import tempfile import threading import time -import traceback import urllib.error import urllib.parse from typing import Any, Callable, Mapping @@ -1282,6 +1281,11 @@ def _validate_responses_logprobs(body: dict[str, Any]) -> None: if "top_logprobs" in body: tlp = body.get("top_logprobs") if tlp is None or (isinstance(tlp, str) and not tlp.strip()): + # Pop rather than leave a raw blank/None value in body: later + # provider-only routing checks (_responses_virtual_requires_provider_path) + # read body.get("top_logprobs") directly and must see an omitted + # field, not a truthy empty string. + body.pop("top_logprobs", None) return if body.get("logprobs") is not True: raise RequestError( @@ -1346,6 +1350,11 @@ def _validate_responses_seed(body: dict[str, Any]) -> int | None: message="seed must be an integer", ) if seed is None: + # Pop rather than leave a raw blank/None value in body: later + # provider-only routing checks (_responses_virtual_requires_provider_path) + # read body.get("seed") directly and must see an omitted field, not a + # truthy empty string. + body.pop("seed", None) return None body["seed"] = seed if seed < -(2**63) or seed > (2**63 - 1): @@ -6285,9 +6294,8 @@ def do_GET(self) -> None: # noqa: N802 _provider_upstream_message(exc), exc.detail, ) - except Exception: - traceback.print_exc() - self._send_error(500, "internal_error", "internal server error") + except Exception as exc: + self._send_internal_error(exc) def do_PATCH(self) -> None: # noqa: N802 """Apply an authenticated agent-pool worker update.""" @@ -6329,9 +6337,8 @@ def do_PATCH(self) -> None: # noqa: N802 _provider_upstream_message(exc), exc.detail, ) - except Exception: - traceback.print_exc() - self._send_error(500, "internal_error", "internal server error") + except Exception as exc: + self._send_internal_error(exc) def do_DELETE(self) -> None: # noqa: N802 """Delete an authenticated agent-pool worker resource.""" @@ -6422,9 +6429,8 @@ def do_DELETE(self) -> None: # noqa: N802 self._send_error(400, "invalid_request", str(exc)) except KeyError as exc: self._send_error(404, "agent_not_found", str(exc)) - except Exception: - traceback.print_exc() - self._send_error(500, "internal_error", "internal server error") + except Exception as exc: + self._send_internal_error(exc) def do_POST(self) -> None: # noqa: N802 """Dispatch authenticated completion, agent, and simulation writes.""" @@ -7296,6 +7302,9 @@ def register_video_job(agent: ModelAgent, provider_result: dict[str, Any]) -> di if remaining_timeout <= 0: break attempt_started_at = time.perf_counter() + observation_context_key = ( + orchestrator._routing_observation_context_for_agent(embedding_agent) + ) try: document = self._run(lambda agent=embedding_agent: coordinator.complete_embeddings_batch( inputs, @@ -7309,18 +7318,25 @@ def register_video_job(agent: ModelAgent, provider_result: dict[str, Any]) -> di )) except Exception as exc: # noqa: BLE001 - measured member failover last_embedding_error = exc - orchestrator._group_router.observe_failure(embedding_agent.id) + orchestrator._record_group_failure( + embedding_agent.id, + observation_context_key=observation_context_key, + ) continue if document.get("status") == "completed": orchestrator._group_router.observe_success( embedding_agent.id, time.perf_counter() - attempt_started_at, + observation_context_key=observation_context_key, ) break last_embedding_error = RuntimeError( f"embedding member ended with {document.get('status', 'unknown')}" ) - orchestrator._group_router.observe_failure(embedding_agent.id) + orchestrator._record_group_failure( + embedding_agent.id, + observation_context_key=observation_context_key, + ) document = None if document is None: raise RequestError( @@ -7400,6 +7416,9 @@ def register_video_job(agent: ModelAgent, provider_result: dict[str, Any]) -> di last_embedding_error: Exception | None = None for embedding_agent in embedding_agents: attempt_started_at = time.perf_counter() + observation_context_key = ( + orchestrator._routing_observation_context_for_agent(embedding_agent) + ) try: document = self._run(lambda agent=embedding_agent: coordinator.complete_embeddings_batch( inputs, @@ -7412,12 +7431,16 @@ def register_video_job(agent: ModelAgent, provider_result: dict[str, Any]) -> di )) except Exception as exc: # noqa: BLE001 - measured member failover last_embedding_error = exc - orchestrator._group_router.observe_failure(embedding_agent.id) + orchestrator._record_group_failure( + embedding_agent.id, + observation_context_key=observation_context_key, + ) continue if document.get("status") == "completed": orchestrator._group_router.observe_success( embedding_agent.id, time.perf_counter() - attempt_started_at, + observation_context_key=observation_context_key, ) break if document is None: @@ -7998,9 +8021,8 @@ def register_video_job(agent: ModelAgent, provider_result: dict[str, Any]) -> di _provider_upstream_message(exc), exc.detail, ) - except Exception: - traceback.print_exc() - self._send_error(500, "internal_error", "internal server error") + except Exception as exc: + self._send_internal_error(exc) finally: if endpoint_policy is not None: endpoint_policy.__exit__(None, None, None) @@ -8196,6 +8218,23 @@ def _send_error( _LOGGER.warning("request_failed status=%s code=%s", status, code) self._send(_error_payload(code, message, {"request_id": uuid.uuid4().hex, **(detail or {})}), status) + def _send_internal_error(self, exc: Exception) -> None: + """Log one internal failure with a request ID and return a generic 500.""" + request_id = uuid.uuid4().hex + _LOGGER.error( + "internal_request_error request_id=%s error_type=%s", + request_id, + type(exc).__name__, + ) + self._send( + _error_payload( + "internal_error", + "internal server error", + {"request_id": request_id}, + ), + 500, + ) + def _write_response(self, writer: Callable[[], None]) -> bool: """Run a response-writing callback, swallowing a dead-peer disconnect. diff --git a/docs/library_research.md b/docs/library_research.md index 3cee4438..d9fd0f15 100644 --- a/docs/library_research.md +++ b/docs/library_research.md @@ -80,6 +80,22 @@ Buyer next action: call `default_role_effort_catalog()` / `run_equal_budget_abla and keep route/conduct defaults unchanged until `production_default_change_allowed` returns true. +## Time-windowed routing-observation persistence (2026-08-30) + +| Area | Researched | Decision | Skipped | +|---|---|---|---| +| Shared replay store | Existing stdlib `sqlite3`; repository `state_db`; ADR 0042 routing-observation contract | Keep the opt-in routing-observation store on stdlib SQLite and persist a database-wide maximum retention window in metadata so transactional pruning stays physically bounded without letting a short-window process delete a longer-window peer's evidence. | WAL tuning, a second queue/service, active-lease coordination, or a speculative PostgreSQL migration for this bounded slice. | +| Replay/prune policy | Existing router `window_seconds` replay boundary; SQLite transaction semantics | Reuse per-router `window_seconds` for logical replay, but prune rows by the shared maximum registered window during writes. This preserves completion-order replay and cross-process safety while keeping storage bounded. | Unbounded history, calibrated decay, row-count heuristics, and inferred provider equivalence. | + +### Active-retention correction (2026-09-04) + +Physical pruning is governed by unexpired rows in +`routing_observation_registrations`, not the historical +`max_retention_window_seconds` metadata row. The metadata key is retained only +for schema compatibility and audit history; it does not authorize retention or +prevent pruning after a process lease expires. A store registers only when its +owner finishes initialization and starts the heartbeat. + ## Discovery output ceilings and context windows (2026-08-31) Issue #927 needs real per-model output-ceiling and context-window metadata from diff --git a/docs/model-group-product-technical-spec.md b/docs/model-group-product-technical-spec.md index ca7d962e..179a12ff 100644 --- a/docs/model-group-product-technical-spec.md +++ b/docs/model-group-product-technical-spec.md @@ -110,6 +110,15 @@ classDiagram +agent_id: text +group_name: text } + class RoutingObservation { + +observation_id: integer + +ledger_name: text + +member_id: text + +observed_at: real + +success: boolean + +latency_seconds: real? + +output_tokens: integer? + } class VideoJobRecord { +gateway_job_id: text +provider_job_id: text @@ -125,6 +134,7 @@ classDiagram } ModelGroup "1" --> "0..*" ModelGroupMember AgentPool "1" --> "0..1" ModelGroupMember + ModelGroupMember "1" --> "0..*" RoutingObservation VideoJobRecord "1" --> "0..1" VideoJobUsage ``` @@ -159,11 +169,18 @@ transactionally into these relations. - Touch targets, responsive layout, typography/color tokens, navigation, form validation, performance, and accessibility remain governed by the existing Admin design system and regression contracts. -- Multi-replica routing observation calibration, cancellable conducted-answer - streaming, and opt-in spend-capped live canaries remain explicitly tracked in - `docs/product-technical-gap-baseline.md`. Video job durability still depends - on the configured shared job registry; standalone mode makes no durability - claim. +- Operators may opt into multi-process observation sharing with + `--routing-observation-window-seconds` alongside `--state-db`. The normalized + `routing_observations` table replays only the selected time window and uses + retention-only semantics. Physical pruning is bounded by the shared + database's largest registered routing-observation window so shorter-window + peers cannot erase longer-window evidence; calibrated decay, cross-model + quality weights, and production horizontal-scaling claims remain outside this + slice. +- Video job durability still depends on the configured shared job registry; + standalone mode makes no durability claim. +- Cancellable conducted-answer streaming and opt-in spend-capped live + canaries remain explicitly tracked in `docs/product-technical-gap-baseline.md`. ## Evidence and limits diff --git a/docs/planning/adrs/0032-model-group-cost-aware-discovery.md b/docs/planning/adrs/0032-model-group-cost-aware-discovery.md index 201d141f..7f219faa 100644 --- a/docs/planning/adrs/0032-model-group-cost-aware-discovery.md +++ b/docs/planning/adrs/0032-model-group-cost-aware-discovery.md @@ -98,7 +98,7 @@ sequenceDiagram ## Data, web, and operational boundaries -Group membership survives restart in normalized `model_group` and `model_group_member` relations; the agent JSON payload no longer duplicates `group_name`. Startup migrates legacy payload membership transactionally without dropping agent configuration. Authenticated REST resources provide `GET/POST /api/v1/model_groups` and `GET/PATCH/DELETE /api/v1/model_groups/{group_name}`; deleting a group retains its provider agents. The existing worker-agent create/PATCH API also accepts `group_name`. Admin provides a keyboard/native-form editor backed by those same resources and shows capability coverage, posterior success probability, and EWMA latency instead of fabricated capacity/success figures. The observation ledger intentionally resets on process restart: persisting it without a measurement horizon would let stale provider incidents dominate current routing. Add a normalized, time-windowed observation table when multi-instance aggregation and an explicit retention/decay policy are specified. +Group membership survives restart in normalized `model_group` and `model_group_member` relations; the agent JSON payload no longer duplicates `group_name`. Startup migrates legacy payload membership transactionally without dropping agent configuration. Authenticated REST resources provide `GET/POST /api/v1/model_groups` and `GET/PATCH/DELETE /api/v1/model_groups/{group_name}`; deleting a group retains its provider agents. The existing worker-agent create/PATCH API also accepts `group_name`. Admin provides a keyboard/native-form editor backed by those same resources and shows capability coverage, posterior success probability, and EWMA latency instead of fabricated capacity/success figures. The observation ledger remains process-local by default. ADR 0039 adds an explicitly opt-in, time-window-only SQLite implementation for sharing completed observations; calibrated decay, cross-model weighting, and production horizontal scaling remain outside this slice. ```mermaid erDiagram @@ -123,7 +123,7 @@ erDiagram - Capability tests cover all eight requested model surfaces, group-scoped measured selection, binary speech preservation, and OpenRouter modality metadata without paid inference. An opt-in live test may use a currently free model, but the deterministic contract suite never assumes that a transient free model will remain listed. - OpenCode Zen `/zen/v1/models` availability is joined to Models.dev cost/modality metadata; if either catalog lacks matching structured cost evidence, retain `unknown` rather than infer a price. - Gap: response quality is not yet in this intra-model score. Distinct-model composition must use calibrated evaluation evidence (for example fast-mlsirm), not a hand-authored weight. -- Gap: multi-replica telemetry needs a time-windowed durable store and concurrency-safe aggregation before production horizontal scaling. +- Gap: the opt-in store provides retention-only multi-process observation sharing; calibrated decay, bounded replay at fleet scale, and production horizontal scaling still need measured design evidence. - Gap: final answer deltas for conducted workflows begin after synthesis; true token streaming across dependent workflow steps would require a cancellable asynchronous execution graph. diff --git a/docs/planning/adrs/0040-streamed-responses-usage-boundary.md b/docs/planning/adrs/0040-streamed-responses-usage-boundary.md index b98621e9..213ae30b 100644 --- a/docs/planning/adrs/0040-streamed-responses-usage-boundary.md +++ b/docs/planning/adrs/0040-streamed-responses-usage-boundary.md @@ -22,7 +22,7 @@ success_criteria: source: "tests/test_orchestrated_responses_stream.py" --- -# Record streamed Responses usage at the workflow boundary +# ADR 0040: Record streamed Responses usage at the workflow boundary ## Context diff --git a/docs/planning/adrs/0042-time-windowed-routing-observations.md b/docs/planning/adrs/0042-time-windowed-routing-observations.md new file mode 100644 index 00000000..b9086ff0 --- /dev/null +++ b/docs/planning/adrs/0042-time-windowed-routing-observations.md @@ -0,0 +1,91 @@ +--- +id: "0042" +title: "Opt-in time-windowed routing observations" +status: proposed +proposed_date: "2026-08-29" +deciders: + - "repository maintainer" +affected_components: + - "contextual_orchestrator/routing_observation_store.py" + - "contextual_orchestrator/model_group.py" + - "contextual_orchestrator/orchestrator.py" +related: + - path: "docs/planning/adrs/0032-model-group-cost-aware-discovery.md" + relation: extends +success_criteria: + - metric: "restart continuity" + target: "an explicitly configured state database restores current-window observations" + source: "tests/test_routing_observation_store.py" + - metric: "retention boundary" + target: "observations outside the configured wall-clock window do not affect routing" + source: "tests/test_routing_observation_store.py" +--- + +# Opt-in time-windowed routing observations + +## Context + +Measured model-group routing currently keeps its Beta-Bernoulli success ledger +and Jacobson-style latency EWMA in one gateway process. That is useful local +evidence, but it resets on restart and cannot be shared by replicas. Persisting +unbounded history would allow an old provider incident to dominate a current +route, while decay or cross-model weights would introduce uncalibrated policy. + +## Decision + +Add a normalized `routing_observations` SQLite table behind the explicit +`--routing-observation-window-seconds SECONDS` option. The option requires the +existing `--state-db PATH`; without both settings, routing behavior stays +process-local and unchanged. Each transport or quality ledger writes one row +per completed attempt with its ledger name, opaque member id, wall-clock time, +success flag, measured latency when successful, and provider-reported output +tokens when available. + +Each router refresh replays only rows at or after `now - window_seconds`, in +completion order, using the existing estimator and priors. Writes prune rows +transactionally, but the physical prune boundary is the shared database's +largest registered observation window rather than the current writer's local +window. Separate short-lived SQLite connections and a transactional write +boundary permit multiple gateway processes to use the same database without a +short-window process deleting evidence still required by a longer-window peer; +a storage error propagates so configured durable evidence is never silently +reported as local-only evidence. Non-stream success-observation requests +therefore fail closed. Provider-failure recording preserves the active provider +failure and logs the durable-evidence outage so failover can continue. A stream +may already have emitted provider bytes before its post-completion observation +write; that write failure is logged as degraded durable evidence and cannot +change the already-emitted response. Removed group members delete their +persisted rows. The Admin state +reports whether the policy is enabled, its window, and the literal +`time_window_only` retention policy. + +This decision deliberately does not add calibrated decay, a fleet-wide sequence +cursor, a row-count policy, cross-model quality weighting, inferred provider +equivalence, or a production horizontal-scaling claim. Full-window replay is a +small correctness-first implementation; add incremental replay or a different +shared store only after measured fleet load justifies it. Provider-reported token +counts remain optional and are never inferred. + +## Consequences + +- Operators can preserve current-window transport and quality evidence across + restarts and cooperating gateway processes. +- The default remains process-local, so existing deployments do not create a + new database or change routing behavior. +- The selected window is a retention boundary, not a statistical calibration; + production defaults remain unchanged until fleet-scale evidence and a decay + policy are separately accepted. +- SQLite is appropriate for the current opt-in bounded slice. A high-throughput + deployment may need a shared service or incremental replay after measurement, + but that is intentionally not built speculatively here. + +## References + +Jacobson, V. (1988). Congestion avoidance and control. *ACM SIGCOMM Computer +Communication Review, 18*(4), 314–329. https://doi.org/10.1145/52325.52356 + +Ong, I., Almahairi, A., Wu, V., Chiang, W.-L., Wu, T., Gonzalez, J. E., Kadous, +M. W., & Stoica, I. (2024). *RouteLLM: Learning to route LLMs with preference +data* [Preprint]. arXiv. https://arxiv.org/abs/2406.18665 + +SQLite. (2026). *Transaction*. https://www.sqlite.org/lang_transaction.html diff --git a/docs/product-technical-gap-baseline.md b/docs/product-technical-gap-baseline.md index 8198808f..60f43c3a 100644 --- a/docs/product-technical-gap-baseline.md +++ b/docs/product-technical-gap-baseline.md @@ -925,6 +925,113 @@ A short status comment was left on each of `#868`, `#857`, `#906`, `#911`, and `#912` recording the pin-bump-did-not-fix-it finding so the next pass (human or agent) does not re-diagnose the same sidecar failure from scratch. +## 2026-08-30 PR #911 review-finding verification pass + +Re-verified the CodeRabbit/Devin findings on this branch against current code +before merging protected `main`. `model_group.py`'s lock ordering +(`observe_success`/`observe_failure`/`refresh` now consistently take +`self._lock` outside `self._observation_io_lock`) and the +`_measured_member_order` double-refresh nitpick were already fixed on this +head. The exception-masking concern (`observe_failure` swallowing the +original error) does not reproduce: every call site in `orchestrator.py` and +`server.py` goes through `_record_group_failure`/`_record_quality_failure`, +which catch `RoutingObservationPersistenceError` internally, log it, and let +the caller's bare `raise` preserve the original exception. The +"`zdr_only` leaking into provider payloads" finding does not reproduce +either: `_ORCHESTRATION_ONLY_KEYS` (which includes `zdr_only`) is applied +consistently at every `upstream = {...}` provider-payload construction site +in `orchestrator.py`, and `batch_routing.py` documents `zdr_only` as a +selection-policy-only field never forwarded upstream. The ADR 0042 +(renumbered from its original 0039 due to a collision with an +independently-landed, unrelated ADR of the same number -- see +`docs/planning/adrs/0042-time-windowed-routing-observations.md`) `status: +proposed` field is correct as-is for an unmerged PR. The previously flagged +"Partially closed — multi-instance telemetry" doc edit is no longer present +in the current diff against protected `main` (superseded by an earlier +merge). Fixed: `tests/test_measured_routing_evidence.py`'s +`fail_first_quality_append` fake now checks the `ledger_name` argument +(`"quality"`) before raising, instead of unconditionally failing whichever +ledger's `append` happens to be called first. Also re-verified on this exact +PR head: empty-string `seed`/`top_logprobs` omission is now covered by +`tests/test_orchestrated_responses_stream.py::test_streamed_orchestrated_responses_allows_blank_seed_and_top_logprobs`, +and `text.format.type="text"` no longer forces the provider-only Responses +path because `_responses_virtual_requires_provider_path` restricts that branch +to actual provider-preserving controls. Newly fixed on this head: manually +raised `ProviderUpstreamError` instances can no longer echo raw provider/client +diagnostics into HTTP or SSE customer responses, so oversized text, private +URLs, and credential-bearing strings are reduced to a bounded redacted +sentence before `_provider_upstream_message()` adds caller guidance. +Regression coverage now includes one direct HTTP failure path and both chat and +Responses SSE failure surfaces. `record_stream_usage` SSE-completion exception +handling remains independently verified on the branch and was not reopened +here. + +## 2026-08-30 durable routing-observation context and ordering slice + +PR #911 exact head `1e14b2fed0e9043ce3fd639c57aede4c40f77bf7` left durable +routing observations keyed only by `ledger_name` and `member_id`, with write +time assigned when SQLite acquired the transaction lock and old rows pruned by +whichever writer had the shortest configured window. That produced four linked +correctness gaps on the live branch: reused agent IDs could inherit unrelated +history after restart, in-flight old-group attempts could repopulate a newly +assigned member after `reset_members`, concurrent writers could replay rows in +lock-acquisition order instead of completion order, and one short-window +gateway could delete evidence still inside another gateway's longer window. + +The bounded fix keeps the opt-in SQLite slice but changes the durable contract: +each row now carries a caller-supplied completion timestamp plus an active +member-context key derived from the serving agent shape, router refresh loads +only rows whose `(member_id, context_key)` matches the current member set, and +shared writers prune transactionally by the database's largest registered +routing-observation window instead of their local window. This preserves +restart continuity for the same agent shape while failing closed on stale or +reassigned evidence and prevents a short-window peer from deleting evidence a +longer-window peer still needs. + +Exact branch evidence on August 30, 2026 used `uv run pytest -q` from the clean +`commercial-loop-20260830-pr911-ordering-refresh` worktree: +`tests/test_routing_observation_store.py` passed (`24 passed in 1.74s`), then +`tests/test_routing_observation_store.py tests/test_measured_routing_evidence.py tests/test_model_group.py` +passed (`86 passed in 2.91s`). Added regressions cover physical pruning of +rows older than the shared maximum retention window while preserving mixed-window +coexistence. Hosted protected checks, independent review, conflict resolution, +and normal merge remain required before this becomes protected-main evidence. + +The September 4 review found that the historical maximum-window metadata no +longer controls physical pruning after active leases were introduced. Current +authority is the set of unexpired registration leases; the metadata row remains +schema-compatibility and audit history only. Initialization now registers a +lease only when the owning orchestrator reaches heartbeat startup, preventing a +failed constructor from extending retention. Batch success without honest +per-item latency remains durable, incomplete embedding attempts keep failover +when observation persistence is unavailable, and batch quality observations +retain the execution-time member context across reassignment. These repairs are +open-PR evidence until current-head review, full tests, and protected merge pass. + +## 2026-08-29 bounded routing-observation durability slice + +The multi-instance routing gap is now partially implemented behind an explicit +operator choice. `--routing-observation-window-seconds SECONDS` requires +`--state-db PATH` and enables the normalized `routing_observations` SQLite table. +Each completed transport or quality attempt is stored in its logical ledger; +router reads replay only observations inside the selected wall-clock window, +and writes prune rows outside the database-wide maximum registered +routing-observation window. Separate gateway processes can therefore share +current-window success/failure, measured latency, and provider-reported +completion-token evidence through the existing state database without letting a +shorter-window peer erase a longer-window peer's active evidence. + +This is deliberately a retention-only slice: the default remains process-local; +there is no calibrated decay, cross-model quality weighting, inferred provider +equivalence, or claim of production horizontal-scaling readiness. Non-stream +success-observation persistence errors fail closed; provider-failure recording +preserves the original failure and logs the durable-evidence outage. After a +stream has emitted provider bytes, a write failure is logged and cannot change +the completed response. +Focused store/router/orchestrator/CLI coverage is maintained in +`tests/test_routing_observation_store.py`; protected `main` and hosted Checks +remain the release boundary for this implementation. + ## 2026-08-30 hourly loop: #868 test-mock fix, #857 narrow hardening, #906 stale-base merge Fresh status check confirmed #868/#911/#912 were still `BLOCKED` purely on the diff --git a/docs/product_planning.md b/docs/product_planning.md index fab5d421..c979b52c 100644 --- a/docs/product_planning.md +++ b/docs/product_planning.md @@ -43,7 +43,7 @@ Enterprise teams want the benefit of collective model intelligence without makin |---|---| | Admin dashboard | Fugu says the user should see one model; enterprise operators still need internal evidence. | | Agent pool | Fugu report makes pool configurability and provider exclusion first-class. | -| Model groups | Operators assert provider-neutral equivalence, then inspect and edit measured member routing without relying on transient example model names. | +| Model groups | Operators assert provider-neutral equivalence, then inspect and edit measured member routing without relying on transient example model names; an explicit SQLite time window can share current observations across gateway processes. | | Orchestration policy | Fugu and Fugu Ultra define the latency-quality frontier that operators must tune. | | Workflow run trace | TRINITY and Conductor both make role/step behavior central to coordination quality. | | Access list inspector | Conductor access lists are the concrete mechanism for context visibility and auditability. | @@ -81,6 +81,7 @@ fast-mlsirm contracts are available. - The management console prioritizes traceability and policy control over decorative SaaS chrome. - Every new product surface maps to one of: compatible API adoption, model/group pool management, cost and policy control, trace audit, access-list evidence, evaluation replay, or i18n. - Local runtime analytics are clearly labeled as process-local evidence and not production telemetry. +- Shared routing observations are opt-in, time-window-only evidence; they do not imply calibrated decay or production horizontal-scaling readiness. - Streamed Responses cost rollups are per-workflow and per-trace-step; missing provider usage is explicitly unavailable, never estimated from answer text. - Sales readiness is reported as local enterprise-pilot evidence with pass/warn/fail remediation, not as a production compliance certificate. - Commercial readiness is reported as high-value buyer due-diligence evidence with a KRW 2,000,000,000 target value caveat, not as a sale guarantee. diff --git a/tests/test_batch_optimizer.py b/tests/test_batch_optimizer.py index 350384a6..5cc52a62 100644 --- a/tests/test_batch_optimizer.py +++ b/tests/test_batch_optimizer.py @@ -87,6 +87,39 @@ def test_batch_route_persists_runs_with_usage() -> None: assert row["output_tokens"] == 18 # 3 x 6 reported completion tokens +def test_batch_route_judge_uses_execution_time_routing_context(monkeypatch) -> None: + original = ModelAgent( + "general_agent", + "model-x", + tags=("reasoning", "writing"), + credential_key="ORIGINAL_ACCOUNT", + ) + replacement = replace(original, credential_key="REPLACEMENT_ACCOUNT") + client = _CountingClient() + orchestrator = TaskOrchestrator([original], client=client) + expected_context = orchestrator._routing_observation_context_for_agent(original) + batch_chat = client.batch_chat + + def reassign_after_batch(*args, **kwargs): + result = batch_chat(*args, **kwargs) + orchestrator.agents = [replacement] + orchestrator.candidates = [replacement] + return result + + observed_contexts: list[str | None] = [] + + def judge(**kwargs): + observed_contexts.append(kwargs["observation_context_key"]) + return {"accepted": True, "reason": "test", "verifier_output": kwargs["answer"]} + + monkeypatch.setattr(client, "batch_chat", reassign_after_batch) + monkeypatch.setattr(orchestrator, "_realtime_route_judge", judge) + + orchestrator.batch_route(["task one"]) + + assert observed_contexts == [expected_context] + + def test_optimizer_use_batch_routes_via_batch_and_matches_serial() -> None: batch_client = _CountingClient() serial_client = _CountingClient() diff --git a/tests/test_embeddings_model_pool_http_honesty.py b/tests/test_embeddings_model_pool_http_honesty.py index 80e12fd0..4bcfbb01 100644 --- a/tests/test_embeddings_model_pool_http_honesty.py +++ b/tests/test_embeddings_model_pool_http_honesty.py @@ -11,7 +11,13 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1])) -from contextual_orchestrator import CostRoutingCoordinator, ModelAgent, TaskOrchestrator +from contextual_orchestrator import ( + CostRoutingCoordinator, + InMemoryConfigStore, + ModelAgent, + TaskOrchestrator, +) +from contextual_orchestrator.model_group import RoutingObservationPersistenceError from contextual_orchestrator.server import SecurityConfig, build_server _TEST_AUTH_TOKEN = "embeddings_model_pool_http_honesty_token" @@ -270,6 +276,118 @@ def test_http_batch_embeddings_auto_selects_enabled_embedding_agent() -> None: thread.join(timeout=5) +def test_embedding_attempts_keep_their_original_routing_context() -> None: + for path, input_key, value in ( + ("/v1/embeddings", "input", "alpha"), + ("/v1/batch/embeddings", "inputs", ["alpha"]), + ): + original = ModelAgent( + "embedding_agent", + "mock-planner", + base_url="mock://original", + tags=("embedding",), + ) + orchestrator = TaskOrchestrator( + [original, ModelAgent("survivor_agent", "mock-survivor", tags=("embedding",))] + ) + expected_context = orchestrator._routing_observation_context_for_agent(original) + counter = type("ExactSyntheticCounter", (), {"count_text": lambda self, text, model="": len(text)})() + coordinator = CostRoutingCoordinator( + orchestrator, InMemoryConfigStore(), embedding_token_counter=counter + ) + complete = coordinator.complete_embeddings_batch + + def complete_after_reassignment(*args, **kwargs): + document = complete(*args, **kwargs) + orchestrator.remove_agent("default", original.id) + orchestrator.add_agent( + "default", + { + "id": original.id, + "model": original.model, + "base_url": "mock://replacement", + "tags": ["embedding"], + }, + ) + return document + + coordinator.complete_embeddings_batch = complete_after_reassignment # type: ignore[method-assign] + observed_contexts: list[str | None] = [] + observe_success = orchestrator._group_router.observe_success + + def capture_success(*args, **kwargs): + observed_contexts.append(kwargs.get("observation_context_key")) + return observe_success(*args, **kwargs) + + orchestrator._group_router.observe_success = capture_success # type: ignore[method-assign] + server = build_server( + orchestrator, + port=0, + security=SecurityConfig(auth_token=_TEST_AUTH_TOKEN), + coordinator=coordinator, + ) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + status, body = _post( + server.server_address[1], + path, + {"model": original.model, input_key: value}, + ) + assert status == 200, body + assert observed_contexts == [expected_context] + finally: + server.shutdown() + thread.join(timeout=5) + orchestrator.close() + + +def test_incomplete_embedding_attempt_survives_observation_store_failure() -> None: + first = ModelAgent("embedding_first", "mock-planner", tags=("embedding",)) + second = ModelAgent("embedding_second", "mock-planner", tags=("embedding",)) + orchestrator = TaskOrchestrator([first, second]) + counter = type( + "ExactSyntheticCounter", + (), + {"count_text": lambda self, text, model="": len(text)}, + )() + coordinator = CostRoutingCoordinator( + orchestrator, + InMemoryConfigStore(), + embedding_token_counter=counter, + ) + complete = coordinator.complete_embeddings_batch + attempted: list[str | None] = [] + + def incomplete_then_complete(*args, **kwargs): + attempted.append(kwargs.get("agent_id")) + if kwargs.get("agent_id") == first.id: + return {"status": "failed"} + return complete(*args, **kwargs) + + def fail_observation(*args, **kwargs): + raise RoutingObservationPersistenceError("simulated store outage") + + coordinator.complete_embeddings_batch = incomplete_then_complete # type: ignore[method-assign] + orchestrator._group_router.observe_failure = fail_observation # type: ignore[method-assign] + server = build_server( + orchestrator, + port=0, + security=SecurityConfig(auth_token=_TEST_AUTH_TOKEN), + coordinator=coordinator, + ) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + status, body = _post(server.server_address[1], "/v1/embeddings", {"input": "alpha"}) + assert status == 200, body + assert attempted == [first.id, second.id] + finally: + server.shutdown() + thread.join(timeout=5) + orchestrator.close() + + if __name__ == "__main__": test_http_embeddings_rejects_model_outside_agent_pool() test_http_embeddings_accepts_model_in_agent_pool() diff --git a/tests/test_measured_routing_evidence.py b/tests/test_measured_routing_evidence.py index 5a0c35a3..1b0ff973 100644 --- a/tests/test_measured_routing_evidence.py +++ b/tests/test_measured_routing_evidence.py @@ -340,6 +340,61 @@ def judge(text, fallback, *, free_only=False): assert backup_quality["success_count"] == 1 +def test_route_once_keeps_judge_reject_failover_when_quality_write_fails( + monkeypatch: pytest.MonkeyPatch, tmp_path, caplog +) -> None: + agents = [ + ModelAgent("primary_worker", "mock", tags=("reasoning",), priority=5), + ModelAgent("backup_worker", "mock", tags=("reasoning",), priority=1), + ] + orchestrator = TaskOrchestrator( + agents, + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + + def fake_invoke(agent, messages, **kwargs): + if agent.id == "primary_worker": + return "weak answer", agent.id, agent.model, {"completion_tokens": 10} + return "strong answer", agent.id, agent.model, None + + monkeypatch.setattr(orchestrator, "_invoke", fake_invoke) + monkeypatch.setattr( + orchestrator, + "_model_judge_verification", + lambda _text, fallback, *, free_only=False: { + "accepted": "strong" in fallback["verifier_output"], + "reason": "verdict", + "verifier_output": fallback["verifier_output"], + "judge": "model", + }, + ) + original_append = orchestrator._routing_observation_store.append + quality_append_calls = 0 + + def fail_first_quality_append(*args, **kwargs): + nonlocal quality_append_calls + ledger_name = args[0] if args else kwargs.get("ledger_name") + if ledger_name == "quality": + quality_append_calls += 1 + if quality_append_calls == 1: + raise OSError("simulated quality storage outage") + return original_append(*args, **kwargs) + + try: + monkeypatch.setattr( + orchestrator._routing_observation_store, + "append", + fail_first_quality_append, + ) + result = orchestrator.route_once([{"role": "user", "content": "do work"}]) + assert result["answer"] == "strong answer" + assert result["trace"][-1]["agent_id"] == "backup_worker" + assert "durable quality observation failed" in caplog.text + finally: + orchestrator.close() + + def test_policy_realtime_judge_must_be_boolean() -> None: from contextual_orchestrator.orchestrator import OrchestrationPolicy diff --git a/tests/test_model_group.py b/tests/test_model_group.py index 4cd8c36d..da4cdb15 100644 --- a/tests/test_model_group.py +++ b/tests/test_model_group.py @@ -121,6 +121,43 @@ def test_group_reassignment_discards_old_group_measurements() -> None: assert orchestrator._group_router.member_report(member.id)["success_count"] == 0 +def test_missing_member_context_falls_back_to_a_stable_hash() -> None: + orchestrator = TaskOrchestrator([_agent("member_one", "vendor/model")]) + + first = orchestrator._routing_observation_context_for_member("missing_member") + second = orchestrator._routing_observation_context_for_member("missing_member") + + assert first == second + assert isinstance(first, str) + assert len(first) == 64 + + +def test_record_group_success_for_agent_uses_the_supplied_agent_shape() -> None: + first = _agent("member_one", "vendor/one") + second = ModelAgent( + "member_two", + "vendor/two", + "https://alternate.example/v1", + provider_name="alternate_provider", + group_name="shared_reasoning_model", + ) + orchestrator = TaskOrchestrator([first, second]) + observed: list[tuple[str, str]] = [] + original = orchestrator._group_router.observe_success + + def capture(member_id: str, latency_seconds: float, **kwargs) -> None: + observed.append((member_id, kwargs["observation_context_key"])) + original(member_id, latency_seconds, **kwargs) + + orchestrator._group_router.observe_success = capture # type: ignore[method-assign] + + orchestrator._record_group_success_for_agent(second, 0.1) + + assert observed == [ + (second.id, orchestrator._routing_observation_context_for_agent(second)) + ] + + def test_load_agents_accepts_sidecar_list_catalog(tmp_path) -> None: catalog = tmp_path / "catalog.json" catalog.write_text( diff --git a/tests/test_openrouter_uptime.py b/tests/test_openrouter_uptime.py index 1bdf6abe..053dd506 100644 --- a/tests/test_openrouter_uptime.py +++ b/tests/test_openrouter_uptime.py @@ -135,6 +135,20 @@ def test_update_prior_rejects_invalid_components() -> None: router.update_prior("member_b", float("nan"), 0.0) +def test_update_prior_ignores_retired_member() -> None: + """A late telemetry poll cannot kill the collector after pool removal.""" + collector, group_router, quality_router, _ = _collectors(100.0) + agent = _agents()[0] + group_router.forget_members({"other_provider_member"}) + quality_router.forget_members({"other_provider_member"}) + + collector._poll_agent(agent) + + assert collector.window_evidence(agent.id) == (1.0, 0.0) + assert agent.id not in group_router.snapshot(refresh=False) + assert agent.id not in quality_router.snapshot(refresh=False) + + if __name__ == "__main__": test_start_without_openrouter_agents_is_inert() test_poll_folds_one_window_of_measured_mass() diff --git a/tests/test_orchestrated_responses_stream.py b/tests/test_orchestrated_responses_stream.py index 5df6f5db..3c8bd55c 100644 --- a/tests/test_orchestrated_responses_stream.py +++ b/tests/test_orchestrated_responses_stream.py @@ -107,6 +107,55 @@ def test_virtual_models_stream_openai_reasoning_summaries(model: str) -> None: )["trace"]} == {"free_worker"} +def test_streamed_orchestrated_responses_allows_blank_seed_and_top_logprobs() -> None: + """Blank-string ``seed`` / ``top_logprobs`` are omit-equivalent and must not + force the provider-only (non-streamed) Responses path. + + Regression: ``_validate_responses_seed`` and ``_validate_responses_logprobs`` + left the raw blank string in ``body`` on their omit branch, so + ``_responses_virtual_requires_provider_path``'s ``body.get("seed") is not + None`` / ``body.get("top_logprobs") is not None`` checks saw a truthy + ``""`` and wrongly rejected virtual streaming with ``invalid_stream`` even + though both fields were semantically omitted. + """ + token = "responses_stream_blank_controls_token" + orchestrator = TaskOrchestrator( + [ModelAgent("blank_control_worker", "mock-model", tags=("reasoning",), priority=100)] + ) + server = build_server(orchestrator, port=0, security=SecurityConfig(auth_token=token)) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + request = urllib.request.Request( + f"http://127.0.0.1:{server.server_address[1]}/v1/responses", + data=json.dumps( + { + "model": "orchestrator/auto", + "input": "Blank seed and top_logprobs must not force provider-only.", + "stream": True, + "seed": "", + "top_logprobs": "", + } + ).encode(), + headers={"content-type": "application/json", "authorization": f"Bearer {token}"}, + method="POST", + ) + with urllib.request.urlopen(request, timeout=10) as response: + assert response.headers["content-type"].startswith("text/event-stream") + stream = response.read().decode() + finally: + server.shutdown() + thread.join(timeout=5) + + types = [ + json.loads(line[6:])["type"] + for line in stream.splitlines() + if line.startswith("data: {") + ] + assert types[0] == "response.created" + assert types[-1] == "response.completed" + + def test_streamed_responses_records_unavailable_usage_without_estimating_answer() -> None: token = "responses_stream_usage_token" orchestrator = TaskOrchestrator([ diff --git a/tests/test_passthrough_provider_failover.py b/tests/test_passthrough_provider_failover.py index 87751f2a..4b94ff47 100644 --- a/tests/test_passthrough_provider_failover.py +++ b/tests/test_passthrough_provider_failover.py @@ -624,6 +624,87 @@ def proxy_send_once( ] +def test_structured_repair_records_success_with_final_provider_context(monkeypatch) -> None: + """A repair-time 413 must not attach the fallback success to the first account.""" + class RepairFailoverClient(SequencedProxyClient): + def proxy_send_once( + self, agent: ModelAgent, endpoint: str, payload: dict[str, Any] + ) -> dict[str, Any]: + del endpoint + self.calls.append((agent.id, deepcopy(payload))) + if len(self.calls) == 1: + return { + "model": agent.model, + "choices": [{"message": {"content": "not json"}}], + } + if agent.id == "primary_agent": + raise _http_error(413) + return { + "model": agent.model, + "choices": [{"message": {"content": "{}"}}], + } + + proxy_send = proxy_send_once + + client = RepairFailoverClient({}) + orchestrator = TaskOrchestrator( + [ + ModelAgent( + "primary_agent", + "primary-model", + priority=10, + provider_name="primary", + credential_key="PRIMARY_ACCOUNT", + group_name="shared_model", + tags=("response_format",), + ), + ModelAgent( + "fallback_agent", + "fallback-model", + priority=1, + provider_name="fallback", + credential_key="FALLBACK_ACCOUNT", + group_name="shared_model", + tags=("response_format",), + ), + ], + client=client, + ) + orchestrator.conduct = lambda *args, **kwargs: { # type: ignore[method-assign] + "mode": "conduct", + "answer": "evidence", + "trace": [], + "verification": {"accepted": True, "reason": "test", "verifier_output": ""}, + } + observations: list[tuple[str, str | None]] = [] + monkeypatch.setattr( + orchestrator._group_router, + "observe_success", + lambda member_id, latency_seconds, **kwargs: observations.append( + (member_id, kwargs.get("observation_context_key")) + ), + ) + + result = orchestrator.proxy_completion( + { + "model": TaskOrchestrator.AUTO_MODEL, + "messages": [{"role": "user", "content": "repair this"}], + "response_format": { + "type": "json_schema", + "json_schema": {"schema": {"type": "object"}}, + }, + }, + single_agent=False, + ) + + fallback = orchestrator._agent("fallback_agent") + assert result["model"] == "fallback-model" + assert observations[-1] == ( + "fallback_agent", + orchestrator._routing_observation_context_for_agent(fallback), + ) + + def test_all_virtual_candidates_rejecting_size_preserves_request_too_large() -> None: """Exhausted 413 routing remains a client-visible size error, not HTTP 500.""" client = SequencedProxyClient( diff --git a/tests/test_provider_error_taxonomy.py b/tests/test_provider_error_taxonomy.py index 17f5ebd9..d912d94f 100644 --- a/tests/test_provider_error_taxonomy.py +++ b/tests/test_provider_error_taxonomy.py @@ -9,6 +9,7 @@ import io import json +import re import socket import ssl import sys @@ -20,6 +21,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from contextual_orchestrator import server as server_module # noqa: E402 from contextual_orchestrator import ModelAgent, TaskOrchestrator # noqa: E402 from contextual_orchestrator.orchestrator import ModelClient, is_transient_error # noqa: E402 from contextual_orchestrator.provider_errors import ( # noqa: E402 @@ -495,10 +497,85 @@ def test_chat_completions_returns_openai_compatible_rate_limit_error() -> None: assert error["detail"]["retryable"] is True +def test_chat_completions_redacts_sensitive_manual_provider_message() -> None: + """HTTP error payloads must not echo secrets or private endpoints from manual failures.""" + + class SensitiveUpstream(ModelClient): + def chat(self, agent: ModelAgent, messages: list, temperature: float = 0.2) -> str: # type: ignore[override] + del messages, temperature + raise ProviderUpstreamError( + agent_id=agent.id, + model=agent.model, + error_code="api_error", + message="Authorization: Bearer secret-token at http://10.0.0.9/internal " + ("x" * 400), + client_status=502, + retryable=True, + ) + + orchestrator = TaskOrchestrator( + [ModelAgent("worker_agent", "gpt-x", tags=("reasoning",))], + client=SensitiveUpstream(), + ) + orchestrator._triage_fn = lambda text: False + token = "taxonomy_redaction_token" + server = build_server(orchestrator, port=0, security=SecurityConfig(auth_token=token)) + threading.Thread(target=server.serve_forever, daemon=True).start() + try: + status, body = _post( + f"http://127.0.0.1:{server.server_address[1]}/v1/chat/completions", + {"model": "gpt-x", "messages": [{"role": "user", "content": "hello"}]}, + token, + ) + finally: + server.shutdown() + + assert status == 502 + message = body["error"]["message"] + assert "secret-token" not in message + assert "http://10.0.0.9/internal" not in message + assert "provider diagnostic was redacted for safety" in message + assert len(message) < 300 + + +def test_chat_completions_internal_errors_log_only_request_id_and_type() -> None: + """Unhandled server exceptions log structured metadata without a traceback dump.""" + + class ExplodingClient(ModelClient): + def chat(self, agent: ModelAgent, messages: list, temperature: float = 0.2) -> str: # type: ignore[override] + del agent, messages, temperature + raise RuntimeError("sensitive stack details should stay server-side") + + orchestrator = TaskOrchestrator( + [ModelAgent("worker_agent", "gpt-x", tags=("reasoning",))], + client=ExplodingClient(), + ) + orchestrator._triage_fn = lambda text: False + token = "taxonomy_internal_error_token" + server = build_server(orchestrator, port=0, security=SecurityConfig(auth_token=token)) + threading.Thread(target=server.serve_forever, daemon=True).start() + try: + with patch.object(server_module._LOGGER, "error") as log_error: + status, body = _post( + f"http://127.0.0.1:{server.server_address[1]}/v1/chat/completions", + {"model": "gpt-x", "messages": [{"role": "user", "content": "hello"}]}, + token, + ) + finally: + server.shutdown() + + assert status == 500 + request_id = body["error"]["detail"]["request_id"] + assert re.fullmatch(r"[0-9a-f]{32}", request_id) + assert log_error.call_count == 1 + assert log_error.call_args.args == ( + "internal_request_error request_id=%s error_type=%s", + request_id, + "RuntimeError", + ) + + def test_guidance_table_covers_every_documented_code() -> None: """Every classified code has caller guidance; unknown codes get a default.""" - from contextual_orchestrator import server as server_module - codes = {surface[1] for surface in PROVIDER_STATUS_SURFACES.values()} codes.update({"tls_verification_failed", "tls_failure", "provider_connection_error"}) for code in codes: diff --git a/tests/test_routing_observation_store.py b/tests/test_routing_observation_store.py new file mode 100644 index 00000000..7e2d36c8 --- /dev/null +++ b/tests/test_routing_observation_store.py @@ -0,0 +1,1051 @@ +"""Contracts for the time-windowed durable routing-observation boundary.""" + +from __future__ import annotations + +import sqlite3 +import threading +import time + +import pytest + +from contextual_orchestrator.__main__ import main +from contextual_orchestrator.model_group import ModelGroupRouter +from contextual_orchestrator.orchestrator import ModelAgent, TaskOrchestrator +from contextual_orchestrator.routing_observation_store import ( + SqliteRoutingObservationStore, +) + + +class _Clock: + def __init__(self, value: float = 100.0) -> None: + self.value = value + + def __call__(self) -> float: + return self.value + + +def test_store_shares_current_window_and_keeps_ledgers_separate(tmp_path) -> None: + clock = _Clock() + path = tmp_path / "routing.sqlite" + first = SqliteRoutingObservationStore(path, 10, clock=clock) + second = SqliteRoutingObservationStore(path, 10, clock=clock) + + first.append( + "transport", + "member_a", + context_key="member_a:v1", + observed_at=clock(), + success=True, + latency_seconds=0.2, + output_tokens=20, + ) + first.append( + "transport", + "member_b", + context_key="member_b:v1", + observed_at=clock(), + success=False, + ) + first.append( + "quality", + "member_a", + context_key="member_a:v1", + observed_at=clock(), + success=True, + latency_seconds=0.5, + ) + + assert second.window_seconds == 10 + assert [(row.member_id, row.success) for row in second.load("transport")] == [ + ("member_a", True), + ("member_b", False), + ] + assert [row.member_id for row in second.load("quality")] == ["member_a"] + + clock.value = 111 + assert second.load("transport") == [] + first.append( + "transport", + "member_c", + context_key="member_c:v1", + observed_at=clock(), + success=False, + ) + with sqlite3.connect(path) as connection: + assert connection.execute("SELECT count(*) FROM routing_observations").fetchone()[0] == 1 + first.close() + second.close() + + +def test_store_persists_success_without_invented_latency(tmp_path) -> None: + store = SqliteRoutingObservationStore(tmp_path / "routing.sqlite", 60) + try: + store.append( + "quality", + "member_a", + context_key="member_a:v1", + observed_at=store.now(), + success=True, + latency_seconds=None, + output_tokens=7, + ) + rows = store.load("quality") + finally: + store.close() + + assert [(row.success, row.latency_seconds, row.output_tokens) for row in rows] == [ + (True, None, 7) + ] + + +def test_store_creates_retention_index_for_prune_path(tmp_path) -> None: + store = SqliteRoutingObservationStore(tmp_path / "routing.sqlite", 60) + try: + with sqlite3.connect(tmp_path / "routing.sqlite") as connection: + index_names = { + row[1] for row in connection.execute("PRAGMA index_list(routing_observations)") + } + finally: + store.close() + + assert "routing_observations_observed_at" in index_names + + +def test_store_migrates_prelease_registration_schema(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + with sqlite3.connect(path) as connection: + connection.execute( + "CREATE TABLE routing_observation_registrations (" + "registration_id TEXT PRIMARY KEY, window_seconds INTEGER NOT NULL)" + ) + connection.execute( + "INSERT INTO routing_observation_registrations VALUES (?, ?)", + ("crashed_legacy_process", 3600), + ) + + store = SqliteRoutingObservationStore(path, 60, clock=_Clock(100.0)) + try: + with sqlite3.connect(path) as connection: + columns = { + row[1] + for row in connection.execute( + "PRAGMA table_info(routing_observation_registrations)" + ) + } + registrations = connection.execute( + "SELECT registration_id FROM routing_observation_registrations" + ).fetchall() + finally: + store.close() + + assert "lease_expires_at" in columns + assert registrations == [(store._registration_id,)] + + +def test_concurrent_stores_atomically_migrate_prelease_registration_schema(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + with sqlite3.connect(path) as connection: + connection.execute( + "CREATE TABLE routing_observation_registrations (" + "registration_id TEXT PRIMARY KEY, window_seconds INTEGER NOT NULL)" + ) + + ready = threading.Barrier(2) + stores: list[SqliteRoutingObservationStore] = [] + errors: list[Exception] = [] + + def build() -> None: + try: + ready.wait(timeout=1.0) + stores.append(SqliteRoutingObservationStore(path, 60)) + except Exception as exc: # pragma: no cover - test assertion path + errors.append(exc) + + threads = [threading.Thread(target=build) for _ in range(2)] + for thread in threads: + thread.start() + for thread in threads: + thread.join(timeout=2.0) + + try: + assert not errors + assert all(not thread.is_alive() for thread in threads) + with sqlite3.connect(path) as connection: + columns = { + row[1] + for row in connection.execute( + "PRAGMA table_info(routing_observation_registrations)" + ) + } + assert "lease_expires_at" in columns + finally: + for store in stores: + store.close() + + +def test_store_deletes_only_requested_members(tmp_path) -> None: + clock = _Clock(100.0) + store = SqliteRoutingObservationStore(tmp_path / "routing.sqlite", 60, clock=clock) + store.append("transport", "member_a", context_key="member_a:v1", observed_at=91.0, success=False) + store.append("transport", "member_b", context_key="member_b:v1", observed_at=92.0, success=False) + store.append("quality", "member_a", context_key="member_a:v1", observed_at=93.0, success=False) + + store.delete_members("transport", ["member_a", "member_a"]) + + assert [row.member_id for row in store.load("transport")] == ["member_b"] + assert [row.member_id for row in store.load("quality")] == ["member_a"] + store.delete_members("transport", []) + + +@pytest.mark.parametrize( + ("path", "window", "clock", "error"), + [ + ("", 1, None, TypeError), + ("routing.sqlite", 0, None, ValueError), + ("routing.sqlite", True, None, ValueError), + ("routing.sqlite", 1, object(), TypeError), + ], +) +def test_store_constructor_validates_configuration(path, window, clock, error, tmp_path) -> None: + with pytest.raises(error): + SqliteRoutingObservationStore( + tmp_path / path if path else path, + window, + **({"clock": clock} if clock is not None else {}), + ) + + +@pytest.mark.parametrize("path", [":memory:", "file::memory:?cache=shared", "file:routing?mode=memory&cache=shared"]) +def test_store_rejects_in_memory_databases(path) -> None: + with pytest.raises(ValueError, match="durable SQLite filesystem path"): + SqliteRoutingObservationStore(path, 60) + + +def test_store_validates_attempt_shape(tmp_path) -> None: + store = SqliteRoutingObservationStore(tmp_path / "routing.sqlite", 60) + with pytest.raises(ValueError): + store.append("", "member", context_key="member:v1", observed_at=1.0, success=False) + with pytest.raises(ValueError): + store.append("transport", "", context_key="member:v1", observed_at=1.0, success=False) + with pytest.raises(ValueError): + store.append("transport", "member", context_key="", observed_at=1.0, success=False) + with pytest.raises(TypeError): + store.append("transport", "member", context_key="member:v1", observed_at=1.0, success=1) # type: ignore[arg-type] + with pytest.raises(TypeError): + store.append("transport", "member", context_key="member:v1", observed_at=1.0, success=True, latency_seconds=True) + with pytest.raises(ValueError): + store.append("transport", "member", context_key="member:v1", observed_at=1.0, success=True, latency_seconds=-1) + with pytest.raises(ValueError): + store.append("transport", "member", context_key="member:v1", observed_at=1.0, success=False, latency_seconds=0.1) + with pytest.raises(ValueError): + store.append("transport", "member", context_key="member:v1", observed_at=1.0, success=True, latency_seconds=0.1, output_tokens=0) + with pytest.raises(ValueError): + store.append("transport", "member", context_key="member:v1", observed_at=1.0, success=False, output_tokens=1) + with pytest.raises(ValueError): + store.append("transport", "member", context_key="member:v1", observed_at=float("nan"), success=False) + store.close() + + +def test_router_restores_and_refreshes_shared_observations(tmp_path) -> None: + clock = _Clock() + store = SqliteRoutingObservationStore(tmp_path / "routing.sqlite", 10, clock=clock) + first = ModelGroupRouter(observation_store=store, ledger_name="transport") + first.register_member("member_a") + first.register_member("member_b") + first.observe_success("member_a", 0.2, output_tokens=20) + first.observe_failure("member_b") + + restored = ModelGroupRouter(observation_store=store, ledger_name="transport") + restored.register_member("member_a") + restored.register_member("member_b") + restored.refresh() + + assert restored.member_observation_count("member_a") == 1 + assert restored.member_observation_count("member_b") == 1 + assert restored.member_report("member_a")["ewma_latency_seconds"] == 0.2 + assert restored.ranked_member_ids(["member_b", "member_a"]) == ["member_a", "member_b"] + + clock.value = 111 + restored.refresh() + assert restored.member_observation_count("member_a") == 0 + assert restored.member_observation_count("member_b") == 0 + + +def test_router_preserves_updated_priors_during_refresh(tmp_path) -> None: + store = SqliteRoutingObservationStore(tmp_path / "routing.sqlite", 60) + router = ModelGroupRouter( + prior_resolver=lambda _member_id: (1.0, 1.0), + observation_store=store, + ledger_name="quality", + ) + router.register_member("member_a") + router.update_prior("member_a", 3.0, 4.0) + router.refresh() + + report = router.member_report("member_a") + assert report["success_posterior_mean"] == 0.428571 + assert report["success_count"] == 0 + assert report["failure_count"] == 0 + + +def test_router_serializes_store_operations_with_memory_updates() -> None: + class _LockCheckingStore: + router: ModelGroupRouter | None = None + + def _assert_lock_order(self) -> None: + assert self.router is not None + assert self.router._lock.locked() + assert self.router._observation_io_lock.locked() + + def append(self, *args, **kwargs) -> None: + self._assert_lock_order() + + def load(self, ledger_name: str, active_contexts=None) -> list: + self._assert_lock_order() + return [] + + def delete_members(self, ledger_name: str, member_ids) -> None: + self._assert_lock_order() + + store = _LockCheckingStore() + router = ModelGroupRouter(observation_store=store) + store.router = router + router.register_member("member_a") + router.observe_success("member_a", 0.2) + router.observe_failure("member_a") + router.refresh() + router.reset_members({"member_a"}) + router.forget_members(set()) + + +def test_measured_member_order_refreshes_each_ledger_once(tmp_path, monkeypatch) -> None: + agents = [ + ModelAgent("member_a", "mock-a", group_name="shared_model"), + ModelAgent("member_b", "mock-b", group_name="shared_model"), + ModelAgent("member_c", "mock-c", group_name="shared_model"), + ] + orchestrator = TaskOrchestrator( + agents, + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + try: + quality_refreshes = 0 + transport_refreshes = 0 + quality_refresh = orchestrator._quality_router.refresh + transport_refresh = orchestrator._group_router.refresh + + def count_quality_refresh() -> None: + nonlocal quality_refreshes + quality_refreshes += 1 + quality_refresh() + + def count_transport_refresh() -> None: + nonlocal transport_refreshes + transport_refreshes += 1 + transport_refresh() + + monkeypatch.setattr(orchestrator._quality_router, "refresh", count_quality_refresh) + monkeypatch.setattr(orchestrator._group_router, "refresh", count_transport_refresh) + + assert orchestrator._measured_member_order([agent.id for agent in agents]) == [ + agent.id for agent in agents + ] + assert quality_refreshes == 1 + assert transport_refreshes == 1 + finally: + orchestrator.close() + + +def test_admin_state_refreshes_each_routing_ledger_once(tmp_path, monkeypatch) -> None: + agents = [ + ModelAgent("member_a", "mock-a", group_name="shared_model"), + ModelAgent("member_b", "mock-b", group_name="shared_model"), + ] + orchestrator = TaskOrchestrator( + agents, + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + refreshes = {"transport": 0, "quality": 0} + original_transport_refresh = orchestrator._group_router.refresh + original_quality_refresh = orchestrator._quality_router.refresh + + def count_transport_refresh() -> None: + refreshes["transport"] += 1 + original_transport_refresh() + + def count_quality_refresh() -> None: + refreshes["quality"] += 1 + original_quality_refresh() + + try: + monkeypatch.setattr(orchestrator._group_router, "refresh", count_transport_refresh) + monkeypatch.setattr(orchestrator._quality_router, "refresh", count_quality_refresh) + orchestrator.admin_state() + assert refreshes == {"transport": 1, "quality": 1} + finally: + orchestrator.close() + + +def test_group_listing_refreshes_each_routing_ledger_once(tmp_path, monkeypatch) -> None: + agents = [ + ModelAgent("member_a", "mock-a", group_name="shared_model"), + ModelAgent("member_b", "mock-b", group_name="shared_model"), + ] + orchestrator = TaskOrchestrator( + agents, + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + refreshes = {"transport": 0, "quality": 0} + original_transport = orchestrator._group_router.refresh + original_quality = orchestrator._quality_router.refresh + + def refresh_transport() -> None: + refreshes["transport"] += 1 + original_transport() + + def refresh_quality() -> None: + refreshes["quality"] += 1 + original_quality() + + try: + monkeypatch.setattr(orchestrator._group_router, "refresh", refresh_transport) + monkeypatch.setattr(orchestrator._quality_router, "refresh", refresh_quality) + orchestrator.list_model_groups() + assert refreshes == {"transport": 1, "quality": 1} + finally: + orchestrator.close() + + +def test_routing_context_changes_with_credential_identity() -> None: + orchestrator = TaskOrchestrator([ModelAgent("route_member", "mock-model")]) + first = ModelAgent( + "route_member", "mock-model", base_url="local://worker", credential_key="ACCOUNT_ONE", local_credential_key="LOCAL_ONE" + ) + second = ModelAgent( + "route_member", "mock-model", base_url="local://worker", credential_key="ACCOUNT_TWO", local_credential_key="LOCAL_ONE" + ) + third = ModelAgent( + "route_member", "mock-model", base_url="local://worker", credential_key="ACCOUNT_ONE", local_credential_key="LOCAL_TWO" + ) + + assert orchestrator._routing_observation_context_for_agent(first) != ( + orchestrator._routing_observation_context_for_agent(second) + ) + assert orchestrator._routing_observation_context_for_agent(first) != ( + orchestrator._routing_observation_context_for_agent(third) + ) + + +def test_stream_judgement_keeps_served_agent_context_after_reassignment(monkeypatch) -> None: + served = ModelAgent( + "route_member", "mock-model", group_name="shared_model", credential_key="ACCOUNT_ONE" + ) + replacement = ModelAgent( + "route_member", "mock-model", group_name="shared_model", credential_key="ACCOUNT_TWO" + ) + orchestrator = TaskOrchestrator([served]) + expected_context = orchestrator._routing_observation_context_for_agent(served) + observed_contexts: list[str | None] = [] + + def stream_chat(*args, **kwargs): + del args, kwargs + orchestrator.agents = [replacement] + yield "answer" + + def judge(**kwargs): + observed_contexts.append(kwargs["observation_context_key"]) + return {"accepted": True, "reason": "test", "verifier_output": "answer"} + + monkeypatch.setattr(orchestrator.client, "stream_chat", stream_chat) + monkeypatch.setattr(orchestrator, "_realtime_route_judge", judge) + + assert list(orchestrator.stream_route([{"role": "user", "content": "hello"}])) == [ + "answer" + ] + assert observed_contexts == [expected_context] + + +def test_stream_preserves_emitted_answer_when_observation_write_fails( + tmp_path, monkeypatch, caplog +) -> None: + orchestrator = TaskOrchestrator( + [ModelAgent("member_a", "mock-model", group_name="shared_model")], + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + + def fail_append(*args, **kwargs): + raise OSError("simulated storage outage") + + try: + monkeypatch.setattr(orchestrator._routing_observation_store, "append", fail_append) + assert list( + orchestrator.stream_route( + [{"role": "user", "content": "stream this"}], + model_name="contextual-orchestrator", + ) + ) + assert "durable routing observation failed" in caplog.text + finally: + orchestrator.close() + + +def test_provider_failure_observation_outage_does_not_mask_active_failure( + tmp_path, monkeypatch, caplog +) -> None: + orchestrator = TaskOrchestrator( + [ModelAgent("member_a", "mock-model", group_name="shared_model")], + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + + def fail_append(*args, **kwargs): + raise OSError("simulated storage outage") + + provider_error = RuntimeError("provider failure") + try: + monkeypatch.setattr(orchestrator._routing_observation_store, "append", fail_append) + try: + raise provider_error + except RuntimeError as error: + orchestrator._record_group_failure("member_a") + assert error is provider_error + assert "durable routing observation failed" in caplog.text + finally: + orchestrator.close() + + +def test_removed_agent_observation_uses_captured_context_without_resolving_pool_member( + tmp_path, monkeypatch +) -> None: + removed = ModelAgent("member_a", "mock-model", group_name="shared_model") + survivor = ModelAgent("member_b", "mock-model", group_name="shared_model") + orchestrator = TaskOrchestrator( + [removed, survivor], + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + context_key = orchestrator._routing_observation_context_for_agent(removed) + orchestrator.remove_agent("default", removed.id) + + def fail_if_resolved(member_id: str) -> str: + raise AssertionError(f"unexpected resolver call for {member_id}") + + try: + monkeypatch.setattr(orchestrator._group_router, "_resolve_context_key", fail_if_resolved) + orchestrator._group_router.observe_failure( + removed.id, + observation_context_key=context_key, + ) + with sqlite3.connect(tmp_path / "state.sqlite") as connection: + rows = connection.execute( + "SELECT member_id, context_key FROM routing_observations" + ).fetchall() + assert rows == [(removed.id, context_key)] + finally: + orchestrator.close() + + +def test_removed_agent_observation_without_captured_context_uses_stable_fallback( + tmp_path, +) -> None: + removed = ModelAgent("member_a", "mock-model", group_name="shared_model") + survivor = ModelAgent("member_b", "mock-model", group_name="shared_model") + orchestrator = TaskOrchestrator( + [removed, survivor], + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + try: + orchestrator.remove_agent("default", removed.id) + expected_context = orchestrator._routing_observation_context_for_member(removed.id) + orchestrator._record_group_failure(removed.id) + with sqlite3.connect(tmp_path / "state.sqlite") as connection: + rows = connection.execute( + "SELECT member_id, context_key FROM routing_observations" + ).fetchall() + assert rows == [(removed.id, expected_context)] + assert orchestrator._routing_observation_context_for_member(removed.id) == expected_context + finally: + orchestrator.close() + + +def test_removed_agent_observation_does_not_recreate_runtime_member_state(tmp_path) -> None: + removed = ModelAgent("member_a", "mock-model", group_name="shared_model") + survivor = ModelAgent("member_b", "mock-model", group_name="shared_model") + orchestrator = TaskOrchestrator( + [removed, survivor], + state_db=str(tmp_path / "state.sqlite"), + routing_observation_window_seconds=60, + ) + context_key = orchestrator._routing_observation_context_for_agent(removed) + try: + orchestrator.remove_agent("default", removed.id) + assert removed.id not in orchestrator._group_router.snapshot(refresh=False) + orchestrator._group_router.observe_success( + removed.id, + 0.2, + observation_context_key=context_key, + ) + assert removed.id not in orchestrator._group_router.snapshot(refresh=False) + with sqlite3.connect(tmp_path / "state.sqlite") as connection: + rows = connection.execute( + "SELECT member_id, context_key, success FROM routing_observations" + ).fetchall() + assert rows == [(removed.id, context_key, 1)] + finally: + orchestrator.close() + + +def test_router_rejects_incomplete_store_contract() -> None: + with pytest.raises(TypeError): + ModelGroupRouter(observation_store=object()) # type: ignore[arg-type] + + +def test_router_uses_store_public_clock_for_default_observed_at() -> None: + captured: list[float] = [] + + class _Store: + def now(self) -> float: + return 123.5 + + def append(self, ledger_name, member_id, **kwargs) -> None: + del ledger_name, member_id + captured.append(kwargs["observed_at"]) + + def load(self, ledger_name: str, active_contexts=None) -> list: + del ledger_name, active_contexts + return [] + + def delete_members(self, ledger_name: str, member_ids) -> None: + del ledger_name, member_ids + + router = ModelGroupRouter(observation_store=_Store()) + + router.observe_failure("member_a") + + assert captured == [123.5] + + +def test_router_deletes_persisted_context_when_members_leave(tmp_path) -> None: + store = SqliteRoutingObservationStore(tmp_path / "routing.sqlite", 60) + router = ModelGroupRouter(observation_store=store) + router.register_member("member_a") + router.observe_failure("member_a") + router.reset_members({"member_a"}) + router.register_member("member_a") + assert router.member_observation_count("member_a") == 0 + router.observe_failure("member_a") + router.forget_members(set()) + assert store.load("transport") == [] + + +def test_store_load_ignores_stale_member_context_after_restart(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + store = SqliteRoutingObservationStore(path, 60, clock=_Clock(100.0)) + store.append( + "transport", + "member_a", + context_key="member_a:v1", + observed_at=90.0, + success=True, + latency_seconds=0.2, + ) + + rows = store.load("transport", {"member_a": "member_a:v2"}) + + assert rows == [] + + +def test_store_load_orders_by_observed_completion_time_not_insert_order(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + store = SqliteRoutingObservationStore(path, 60, clock=_Clock(100.0)) + store.append( + "transport", + "member_a", + context_key="member_a:v1", + observed_at=95.0, + success=False, + ) + store.append( + "transport", + "member_a", + context_key="member_a:v1", + observed_at=90.0, + success=True, + latency_seconds=0.1, + ) + + rows = store.load("transport", {"member_a": "member_a:v1"}) + + assert [(row.success, row.latency_seconds) for row in rows] == [ + (True, 0.1), + (False, None), + ] + + +def test_router_captures_default_observed_at_before_lock_acquisition(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + store = SqliteRoutingObservationStore(path, 60, clock=_Clock(100.0)) + router = ModelGroupRouter(observation_store=store) + first_timestamp_captured = threading.Event() + allow_first_attempt = threading.Event() + observed_times = iter((90.0, 95.0)) + + def resolve(observed_at: float | None) -> float: + if observed_at is not None: + return float(observed_at) + when = next(observed_times) + if when == 90.0: + first_timestamp_captured.set() + assert allow_first_attempt.wait(timeout=1.0) + return when + + router._resolve_observed_at = resolve # type: ignore[method-assign] + + first = threading.Thread( + target=router.observe_success, + args=("member_a", 0.1), + name="member-a-observation", + ) + second = threading.Thread( + target=router.observe_success, + args=("member_b", 0.1), + name="member-b-observation", + ) + + first.start() + assert first_timestamp_captured.wait(timeout=1.0) + second.start() + second.join(timeout=1.0) + assert not second.is_alive() + allow_first_attempt.set() + first.join(timeout=1.0) + assert not first.is_alive() + + assert [row.member_id for row in store.load("transport")] == [ + "member_a", + "member_b", + ] + + +def test_shorter_writer_window_does_not_delete_longer_window_evidence(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + long_window = SqliteRoutingObservationStore(path, 60, clock=_Clock(100.0)) + short_window = SqliteRoutingObservationStore(path, 10, clock=_Clock(100.0)) + long_window.append( + "transport", + "member_a", + context_key="member_a:v1", + observed_at=50.0, + success=False, + ) + + short_window.append( + "transport", + "member_b", + context_key="member_b:v1", + observed_at=100.0, + success=False, + ) + + rows = long_window.load( + "transport", + {"member_a": "member_a:v1", "member_b": "member_b:v1"}, + ) + + assert [row.member_id for row in rows] == ["member_a", "member_b"] + + +def test_closed_long_window_no_longer_controls_physical_retention(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + clock = _Clock(100.0) + long_window = SqliteRoutingObservationStore(path, 60, clock=clock) + short_window = SqliteRoutingObservationStore(path, 10, clock=clock) + long_window.append( + "transport", "member_a", context_key="member_a:v1", observed_at=50.0, success=False + ) + long_window.close() + + short_window.append( + "transport", "member_b", context_key="member_b:v1", observed_at=100.0, success=False + ) + + with sqlite3.connect(path) as connection: + rows = connection.execute( + "SELECT member_id FROM routing_observations ORDER BY observation_id" + ).fetchall() + assert rows == [("member_b",)] + + +def test_crashed_long_window_registration_expires_and_no_longer_controls_retention( + tmp_path, +) -> None: + path = tmp_path / "routing.sqlite" + clock = _Clock(100.0) + crashed_long_window = SqliteRoutingObservationStore(path, 60, clock=clock) + short_window = SqliteRoutingObservationStore(path, 10, clock=clock) + crashed_long_window.append( + "transport", + "member_a", + context_key="member_a:v1", + observed_at=50.0, + success=False, + ) + + # Simulate an ungraceful exit: do not call close() on the long-window + # instance. Once its bounded lease expires, a live writer removes both the + # stale registration and evidence outside every active window. + clock.value = 221.0 + short_window.append( + "transport", + "member_b", + context_key="member_b:v1", + observed_at=clock(), + success=False, + ) + + with sqlite3.connect(path) as connection: + registrations = connection.execute( + "SELECT registration_id FROM routing_observation_registrations" + ).fetchall() + observations = connection.execute( + "SELECT member_id FROM routing_observations ORDER BY observation_id" + ).fetchall() + + assert registrations == [(short_window._registration_id,)] + assert observations == [("member_b",)] + + +def test_idle_live_long_window_heartbeat_preserves_replay_evidence( + tmp_path, monkeypatch +) -> None: + path = tmp_path / "routing.sqlite" + clock = _Clock(100.0) + monkeypatch.setattr( + SqliteRoutingObservationStore, "_HEARTBEAT_INTERVAL_MAX_SECONDS", 0.01 + ) + long_window = SqliteRoutingObservationStore(path, 60, clock=clock) + short_window = SqliteRoutingObservationStore(path, 10, clock=clock) + try: + long_window.append( + "transport", + "member_a", + context_key="member_a:v1", + observed_at=170.0, + success=False, + ) + + # The long-window gateway receives no more routing traffic, but its + # independent heartbeat must keep its retention registration live. + clock.value = 221.0 + deadline = time.monotonic() + 1.0 + while time.monotonic() < deadline: + with sqlite3.connect(path) as connection: + lease = connection.execute( + "SELECT lease_expires_at FROM routing_observation_registrations " + "WHERE registration_id = ?", + (long_window._registration_id,), + ).fetchone() + if lease is not None and lease[0] > clock.value: + break + time.sleep(0.01) + else: # pragma: no cover - assertion diagnostic + raise AssertionError("idle store heartbeat did not renew its lease") + + short_window.append( + "transport", + "member_b", + context_key="member_b:v1", + observed_at=clock(), + success=False, + ) + + assert [ + row.member_id + for row in long_window.load( + "transport", + {"member_a": "member_a:v1", "member_b": "member_b:v1"}, + ) + ] == ["member_a", "member_b"] + finally: + long_window.close() + short_window.close() + + +def test_heartbeat_retries_after_connection_failure(tmp_path, monkeypatch) -> None: + path = tmp_path / "routing.sqlite" + clock = _Clock(100.0) + monkeypatch.setattr( + SqliteRoutingObservationStore, "_HEARTBEAT_INTERVAL_MAX_SECONDS", 0.01 + ) + store = SqliteRoutingObservationStore(path, 60, clock=clock, start_heartbeat=False) + connect = store._connect + attempts = 0 + + def flaky_connect(): + nonlocal attempts + attempts += 1 + if attempts == 1: + raise sqlite3.OperationalError("temporary connection failure") + return connect() + + store._connect = flaky_connect + clock.value = 200.0 + store.start_heartbeat() + try: + deadline = time.monotonic() + 1.0 + lease = None + while lease is None and time.monotonic() < deadline: + time.sleep(0.01) + with sqlite3.connect(path) as connection: + lease = connection.execute( + "SELECT lease_expires_at FROM routing_observation_registrations " + "WHERE registration_id = ?", + (store._registration_id,), + ).fetchone() + assert attempts >= 2 + assert lease is not None and lease[0] > clock.value + finally: + store.close() + + +def test_store_prunes_only_rows_older_than_database_retention_window(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + clock = _Clock(50.0) + short_window = SqliteRoutingObservationStore(path, 10, clock=clock) + long_window = SqliteRoutingObservationStore(path, 60, clock=clock) + short_window.append( + "transport", + "member_a", + context_key="member_a:v1", + observed_at=50.0, + success=False, + ) + long_window.append( + "transport", + "member_b", + context_key="member_b:v1", + observed_at=90.0, + success=False, + ) + + clock.value = 111.0 + short_window.append( + "transport", + "member_c", + context_key="member_c:v1", + observed_at=111.0, + success=False, + ) + + with sqlite3.connect(path) as connection: + rows = connection.execute( + "SELECT member_id FROM routing_observations ORDER BY observed_at, observation_id" + ).fetchall() + + assert rows == [("member_b",), ("member_c",)] + + +def test_concurrent_window_registration_keeps_longest_retention(tmp_path) -> None: + path = tmp_path / "routing.sqlite" + ready = threading.Barrier(2) + errors: list[Exception] = [] + + def build(window_seconds: int) -> None: + try: + ready.wait(timeout=1.0) + SqliteRoutingObservationStore(path, window_seconds, clock=_Clock(100.0)).close() + except Exception as exc: # pragma: no cover - test assertion path + errors.append(exc) + + short = threading.Thread(target=build, args=(10,), name="short-window") + long = threading.Thread(target=build, args=(60,), name="long-window") + short.start() + long.start() + short.join(timeout=1.0) + long.join(timeout=1.0) + + assert not errors + with sqlite3.connect(path) as connection: + row = connection.execute( + "SELECT metadata_value FROM routing_observation_metadata WHERE metadata_key = ?", + (SqliteRoutingObservationStore._MAX_RETENTION_WINDOW_KEY,), + ).fetchone() + + assert row == (60,) + + +def test_task_orchestrator_opt_in_store_survives_restart_and_reports_policy(tmp_path) -> None: + path = tmp_path / "state.sqlite" + agents = [ModelAgent("member_a", "mock-model", group_name="shared_model")] + first = TaskOrchestrator( + agents, + state_db=str(path), + routing_observation_window_seconds=60, + ) + try: + first._group_router.observe_success("member_a", 0.25) + assert first.admin_state()["routing_observation_policy"] == { + "enabled": True, + "window_seconds": 60, + "retention_policy": "time_window_only", + } + finally: + first.close() + + second = TaskOrchestrator( + agents, + state_db=str(path), + routing_observation_window_seconds=60, + ) + try: + assert second._group_router.member_observation_count("member_a") == 1 + finally: + second.close() + + +def test_task_orchestrator_requires_state_db_for_durable_observations() -> None: + with pytest.raises(ValueError, match="requires state_db"): + TaskOrchestrator( + [ModelAgent("member_a", "mock-model")], + routing_observation_window_seconds=60, + ) + + +def test_failed_orchestrator_initialization_does_not_start_retention_heartbeat( + tmp_path, monkeypatch +) -> None: + started: list[SqliteRoutingObservationStore] = [] + original = SqliteRoutingObservationStore.start_heartbeat + + def record_start(store): + started.append(store) + original(store) + + monkeypatch.setattr(SqliteRoutingObservationStore, "start_heartbeat", record_start) + + path = tmp_path / "state.sqlite" + with pytest.raises(ValueError, match="tool_retry_attempts"): + TaskOrchestrator( + [ModelAgent("member_a", "mock-model")], + state_db=str(path), + routing_observation_window_seconds=60, + tool_retry_attempts=-1, + ) + + assert started == [] + with sqlite3.connect(path) as connection: + assert connection.execute( + "SELECT count(*) FROM routing_observation_registrations" + ).fetchone() == (0,) + + +def test_cli_requires_state_db_for_durable_observations(capsys) -> None: + with pytest.raises(SystemExit) as exc_info: + main(["--routing-observation-window-seconds", "60"]) + assert exc_info.value.code == 2 + assert "requires --state-db" in capsys.readouterr().err diff --git a/tests/test_true_streaming.py b/tests/test_true_streaming.py index c79ad45d..8520410a 100644 --- a/tests/test_true_streaming.py +++ b/tests/test_true_streaming.py @@ -524,6 +524,44 @@ def stream_route(self, *_args, **_kwargs): assert errors[0]["error_detail"]["retryable"] is True +def test_chat_stream_redacts_sensitive_manual_provider_message() -> None: + """Manual upstream failures must be sanitized before reaching chat SSE callers.""" + server = build_server(TaskOrchestrator([ModelAgent("general_agent", "m-model")]), port=0) + handler = server.RequestHandlerClass.__new__(server.RequestHandlerClass) + frames: list[str] = [] + + class Orchestrator: + def stream_route(self, *_args, **_kwargs): + raise ProviderUpstreamError( + agent_id="worker_agent", + model="gpt-x", + error_code="api_error", + message="Authorization: Bearer secret-token at http://10.0.0.9/internal " + ("x" * 400), + client_status=502, + retryable=True, + ) + yield "unreachable" + + handler._begin_sse = lambda: True + handler._write_sse = lambda frame: frames.append(frame) is None or True + try: + handler._stream_route_completion( + Orchestrator(), SecurityConfig(auth_token="stream-token"), [], "gpt-x" + ) + finally: + server.server_close() + + error = next( + json.loads(frame.removeprefix("data: ")) + for frame in frames + if frame.startswith("data: {") and '"error_code"' in frame + ) + message = error["error_message"] + assert "secret-token" not in message + assert "http://10.0.0.9/internal" not in message + assert "provider diagnostic was redacted for safety" in message + + def test_responses_stream_preserves_classified_provider_error_payload() -> None: """A failed Responses SSE event keeps the actionable upstream taxonomy.""" server = build_server(TaskOrchestrator([ModelAgent("general_agent", "m-model")]), port=0) @@ -561,6 +599,43 @@ def conduct(self, *_args, **_kwargs): assert payload["response"]["error"]["detail"]["retryable"] is True +def test_responses_stream_redacts_sensitive_manual_provider_message() -> None: + """Manual upstream failures must be sanitized before reaching Responses SSE callers.""" + server = build_server(TaskOrchestrator([ModelAgent("general_agent", "m-model")]), port=0) + handler = server.RequestHandlerClass.__new__(server.RequestHandlerClass) + frames: list[str] = [] + + class Orchestrator: + def would_route(self, *_args, **_kwargs): + return False + + def conduct(self, *_args, **_kwargs): + raise ProviderUpstreamError( + agent_id="worker_agent", + model="gpt-x", + error_code="api_error", + message="Authorization: Bearer secret-token at http://10.0.0.9/internal " + ("x" * 400), + client_status=502, + retryable=True, + ) + + handler._begin_sse = lambda: True + handler._write_sse = lambda frame: frames.append(frame) is None or True + try: + assert handler._stream_orchestrated_response( + Orchestrator(), SecurityConfig(auth_token="stream-token"), [], "gpt-x" + ) is False + finally: + server.server_close() + + failed_frame = next(frame for frame in frames if "response.failed" in frame) + payload = json.loads(failed_frame.split("data: ", 1)[1]) + message = payload["response"]["error"]["message"] + assert "secret-token" not in message + assert "http://10.0.0.9/internal" not in message + assert "provider diagnostic was redacted for safety" in message + + if __name__ == "__main__": for name, fn in sorted(globals().items()): if name.startswith("test_") and callable(fn):