From 837fbfd0eaf87283377d67cad4fb54cd347fe482 Mon Sep 17 00:00:00 2001 From: Alexandre Teixeira <111787685+alteixeira20@users.noreply.github.com> Date: Thu, 1 Oct 2026 21:59:53 +0100 Subject: [PATCH] feat(runtime): resolve compact-runtime context window at turn preparation The compact (clean v3) runtime had no effective context window: it learned a limit only reactively from a provider 400/413 and its terminal metrics carried no context_length. PR #41 addressed the reporting gap by probing provider metadata between the last model byte and [DONE], unauthenticated, and folded known-table and endpoint evidence into one "known" flag. Resolve the window once, before the first model request, instead: - src/agent_runtime/context_resolution.py adds a typed ContextResolution (effective value, evidence class, source, all observations, conflicts, provider_io, cached, secret-free probe errors). Evidence classes stay distinct: runtime_confirmed (llama.cpp /slots, /props, or a limit the provider stated this turn), provider_advertised (models catalog), operator_declared (client_runtime_context.model_context_window), known_table, unknown (0, never a default). - Selection is deterministic: runtime beats provider beats table; an operator declaration caps measured evidence and replaces weaker evidence. Disagreements are recorded as conflicts; a declaration below a measured value is a cap, above it a contradiction. - The provider probe forwards the turn's credentials only to the provider's own origin, runs URL resolution off the event loop, is bounded by one deadline, never raises, and caches remote results per credential fingerprint (shorter TTL for failures; local servers are re-probed). - stream_preview resolves at preparation (or accepts a supplied resolution), seeds the proactive trim budget from it when evidence is not unknown, and terminal metrics report only the stored resolution plus any limit the provider stated during the turn. Metrics perform no discovery. src/agent_loop.py and the regular runtime's legacy model_context probe are unchanged. A conftest guard keeps tests that drive the compact runtime with placeholder endpoints from performing real DNS/HTTP lookups. --- src/agent_runtime/context_resolution.py | 464 +++++++++++++++++++++ src/clean_agent_preview.py | 22 +- tests/conftest.py | 33 ++ tests/test_context_resolution.py | 513 ++++++++++++++++++++++++ 4 files changed, 1030 insertions(+), 2 deletions(-) create mode 100644 src/agent_runtime/context_resolution.py create mode 100644 tests/test_context_resolution.py diff --git a/src/agent_runtime/context_resolution.py b/src/agent_runtime/context_resolution.py new file mode 100644 index 000000000..e5064b518 --- /dev/null +++ b/src/agent_runtime/context_resolution.py @@ -0,0 +1,464 @@ +"""Effective model context window, resolved once per logical turn. + +A turn resolves the window it budgets against at preparation time, before any +model request, and keeps the value together with the evidence that chose it. +Terminal metrics report that stored resolution; they never start discovery. + +Evidence classes are kept apart instead of being folded into one "known" flag: + +* ``runtime_confirmed``: the serving process reported its active window + (llama.cpp ``/slots`` or ``/props``) or rejected a request of this turn with + an explicit limit. +* ``provider_advertised``: the provider's model catalog lists a window. +* ``operator_declared``: the client or operator declared a transport window. + It caps runtime or provider evidence and replaces weaker evidence. +* ``known_table``: the static ``KNOWN_CONTEXT_WINDOWS`` fallback. +* ``unknown``: nothing above is available. The value is 0, never a default. + +Any disagreement between sources is recorded as a conflict. An operator value +below a measured value is a cap, not a contradiction; an operator value above +it is a contradiction. + +Context sizing is not authority: nothing here grants or denies an operation. +""" +from __future__ import annotations + +import asyncio +from dataclasses import dataclass, replace +from enum import Enum +import hashlib +import json +import logging +import time +from typing import Any, Mapping, Optional +from urllib.parse import urlparse + +import httpx + +logger = logging.getLogger(__name__) + +# Upper bound for all provider metadata I/O of one turn preparation. A slow +# or unreachable metadata endpoint costs at most this much before the turn +# proceeds with whatever evidence it has. +PROBE_DEADLINE_SECONDS = 3.0 +# Remote provider metadata changes rarely; failures are retried sooner so a +# transient outage does not pin a turn to weaker evidence for long. Local +# servers are always re-probed because they can restart with another window. +PROBE_CACHE_TTL_SECONDS = 600.0 +PROBE_FAILURE_TTL_SECONDS = 60.0 + +# Headers that describe the chat request body rather than the caller. +_REQUEST_ONLY_HEADERS = frozenset({"content-type", "content-length", "accept", "accept-encoding"}) + + +class ContextEvidence(str, Enum): + RUNTIME_CONFIRMED = "runtime_confirmed" + PROVIDER_ADVERTISED = "provider_advertised" + OPERATOR_DECLARED = "operator_declared" + KNOWN_TABLE = "known_table" + UNKNOWN = "unknown" + + +@dataclass(frozen=True) +class ContextObservation: + evidence: ContextEvidence + value: int + source: str + + def to_dict(self) -> dict: + return {"evidence": self.evidence.value, "value": self.value, "source": self.source} + + +@dataclass(frozen=True) +class ContextConflict: + first: ContextObservation + second: ContextObservation + + def to_dict(self) -> dict: + return {"first": self.first.to_dict(), "second": self.second.to_dict()} + + +@dataclass(frozen=True) +class ContextResolution: + """The effective window of one turn and why it was chosen.""" + + effective: int + evidence: ContextEvidence + source: str + observations: tuple[ContextObservation, ...] = () + conflicts: tuple[ContextConflict, ...] = () + provider_io: bool = False + cached: bool = False + probe_errors: tuple[str, ...] = () + + @property + def mismatch(self) -> bool: + return bool(self.conflicts) + + @property + def budget_limit(self) -> int: + """Window the runtime may budget against; 0 means budget reactively.""" + return self.effective if self.evidence is not ContextEvidence.UNKNOWN else 0 + + def observe_runtime_limit(self, limit: Any, source: str = "provider_rejection") -> "ContextResolution": + """Fold a limit the provider stated during this turn. Performs no I/O.""" + try: + value = int(limit or 0) + except (TypeError, ValueError): + return self + if value <= 0: + return self + observation = ContextObservation(ContextEvidence.RUNTIME_CONFIRMED, value, source) + if observation in self.observations: + return self + combined = combine_observations((*self.observations, observation)) + return replace( + combined, + provider_io=self.provider_io, + cached=self.cached, + probe_errors=self.probe_errors, + ) + + def to_dict(self) -> dict: + return { + "effective": self.effective, + "evidence": self.evidence.value, + "source": self.source, + "mismatch": self.mismatch, + "conflicts": [conflict.to_dict() for conflict in self.conflicts], + "observations": [observation.to_dict() for observation in self.observations], + "provider_io": self.provider_io, + "cached": self.cached, + "probe_errors": list(self.probe_errors), + } + + +UNRESOLVED_CONTEXT = ContextResolution(0, ContextEvidence.UNKNOWN, "none") + +_MEASURED = (ContextEvidence.RUNTIME_CONFIRMED, ContextEvidence.PROVIDER_ADVERTISED) + + +def _conflicts(observations: tuple[ContextObservation, ...]) -> tuple[ContextConflict, ...]: + conflicts = [] + for index, first in enumerate(observations): + for second in observations[index + 1:]: + if first.value == second.value: + continue + classes = {first.evidence, second.evidence} + if ContextEvidence.OPERATOR_DECLARED in classes: + operator, other = ( + (first, second) if first.evidence is ContextEvidence.OPERATOR_DECLARED + else (second, first) + ) + # A declared window replaces the static table and may cap a + # measured window. Only a declaration above what the runtime + # or provider supports contradicts it. + if other.evidence not in _MEASURED or operator.value < other.value: + continue + conflicts.append(ContextConflict(first, second)) + return tuple(conflicts) + + +def _strongest(observations, evidence: ContextEvidence) -> Optional[ContextObservation]: + matching = [observation for observation in observations if observation.evidence is evidence] + return min(matching, key=lambda observation: observation.value) if matching else None + + +def combine_observations(observations) -> ContextResolution: + """Choose the effective window deterministically from observations. + + The smallest runtime-confirmed value wins, else the smallest provider + value. An operator declaration caps either, and replaces the known table + or an unknown window. The known table is used only when nothing stronger + exists. No observation yields an unknown window of 0. + """ + observations = tuple(observations) + measured = ( + _strongest(observations, ContextEvidence.RUNTIME_CONFIRMED) + or _strongest(observations, ContextEvidence.PROVIDER_ADVERTISED) + ) + operator = _strongest(observations, ContextEvidence.OPERATOR_DECLARED) + if measured and operator: + chosen = operator if operator.value < measured.value else measured + else: + chosen = measured or operator or _strongest(observations, ContextEvidence.KNOWN_TABLE) + conflicts = _conflicts(observations) + if chosen is None: + return replace(UNRESOLVED_CONTEXT, observations=observations, conflicts=conflicts) + return ContextResolution( + chosen.value, chosen.evidence, chosen.source, + observations=observations, conflicts=conflicts, + ) + + +# --------------------------------------------------------------------------- +# Provider metadata probe +# --------------------------------------------------------------------------- + +@dataclass(frozen=True) +class _ProbeResult: + observations: tuple[ContextObservation, ...] = () + errors: tuple[str, ...] = () + io: bool = False + + +_probe_cache: dict[tuple[str, str, str], tuple[float, _ProbeResult]] = {} + + +def clear_probe_cache() -> None: + _probe_cache.clear() + + +def _origin(url: str) -> tuple[str, str, Optional[int]]: + parsed = urlparse(url or "") + return (parsed.scheme.lower(), (parsed.hostname or "").lower(), parsed.port) + + +def _probe_headers(trusted_origins, target_url: str, headers: Optional[Mapping[str, Any]]) -> dict: + """Forward the turn's provider credentials only to the provider's origin.""" + if not headers or _origin(target_url) not in trusted_origins: + return {} + return { + str(name): str(value) for name, value in headers.items() + if value is not None and str(name).lower() not in _REQUEST_ONLY_HEADERS + } + + +def _auth_fingerprint(headers: Optional[Mapping[str, Any]]) -> str: + if not headers: + return "" + material = json.dumps( + sorted((str(k).lower(), str(v)) for k, v in headers.items() + if v is not None and str(k).lower() not in _REQUEST_ONLY_HEADERS), + separators=(",", ":"), + ) + return hashlib.sha256(material.encode("utf-8")).hexdigest()[:16] + + +def _serving_base(endpoint_url: str) -> str: + # Same derivation the regular runtime uses for llama.cpp server routes. + return endpoint_url.split("/v1")[0] if "/v1" in endpoint_url else endpoint_url.rsplit("/", 1)[0] + + +async def _get_json(client, url, headers, errors, label): + try: + response = await client.get(url, headers=headers) + except httpx.TimeoutException: + errors.append(f"{label}:timeout") + return None + except httpx.TransportError: + errors.append(f"{label}:transport_error") + return None + status = getattr(response, "status_code", 0) + if not (200 <= int(status or 0) < 300): + errors.append(f"{label}:http_{status}") + return None + try: + return response.json() + except Exception: + errors.append(f"{label}:invalid_payload") + return None + + +def _positive_int(value) -> int: + if isinstance(value, bool) or not isinstance(value, (int, float)) or value <= 0: + return 0 + return int(value) + + +async def _probe(endpoint_url, model, headers, is_local, observations, errors, timeout): + from src.copilot import is_copilot_base + from src.endpoint_resolver import build_models_url, resolve_url + from src.model_context import _model_ctx_from_entry + + trusted = {_origin(endpoint_url)} + async with httpx.AsyncClient(timeout=timeout) as client: + if is_local: + base = _serving_base(endpoint_url) + slots = await _get_json( + client, f"{base}/slots", _probe_headers(trusted, f"{base}/slots", headers), + errors, "slots", + ) + n_ctx = _positive_int(slots[0].get("n_ctx")) if ( + isinstance(slots, list) and slots and isinstance(slots[0], dict) + ) else 0 + if not n_ctx: + props = await _get_json( + client, f"{base}/props", _probe_headers(trusted, f"{base}/props", headers), + errors, "props", + ) + generation = props.get("default_generation_settings") if isinstance(props, dict) else None + n_ctx = _positive_int(generation.get("n_ctx")) if isinstance(generation, dict) else 0 + source = "llamacpp_props" + else: + source = "llamacpp_slots" + if n_ctx: + observations.append( + ContextObservation(ContextEvidence.RUNTIME_CONFIRMED, n_ctx, source) + ) + + # Copilot's catalog needs headers this layer does not own; an + # unauthenticated probe only fails. Its models are table-covered. + if is_copilot_base(endpoint_url): + errors.append("models:unsupported_endpoint") + return + # URL building may resolve the host (DNS, tailscale lookup); keep that + # off the event loop and inside the probe deadline. The resolved host + # is the same provider the chat request reaches. + models_url = await asyncio.to_thread(build_models_url, endpoint_url) + if not models_url: + errors.append("models:unsupported_endpoint") + return + trusted.add(_origin(await asyncio.to_thread(resolve_url, endpoint_url))) + payload = await _get_json( + client, models_url, _probe_headers(trusted, models_url, headers), + errors, "models", + ) + if payload is None: + return + entries = payload.get("data") if isinstance(payload, dict) else None + if not isinstance(entries, list): + errors.append("models:invalid_payload") + return + wanted = model.split("/")[-1] + for entry in entries: + if not isinstance(entry, dict): + continue + entry_id = str(entry.get("id") or "") + if entry_id == model or entry_id.split("/")[-1] == wanted: + value = _model_ctx_from_entry(entry) + if value: + observations.append(ContextObservation( + ContextEvidence.PROVIDER_ADVERTISED, int(value), "models_catalog", + )) + else: + errors.append("models:no_window_listed") + return + errors.append("models:model_not_listed") + + +async def probe_provider_context( + endpoint_url: str, + model: str, + *, + headers: Optional[Mapping[str, Any]] = None, + deadline_seconds: float = PROBE_DEADLINE_SECONDS, + is_local: Optional[bool] = None, +) -> _ProbeResult: + """Query provider metadata once, bounded by ``deadline_seconds``. + + Never raises: every failure is reported as a short, secret-free error code + so a turn can continue with other evidence. + """ + observations: list[ContextObservation] = [] + errors: list[str] = [] + if is_local is None: + is_local = await _is_local(endpoint_url) + timeout = max(0.1, float(deadline_seconds)) + try: + await asyncio.wait_for( + _probe(endpoint_url, model, headers, is_local, observations, errors, timeout), + timeout=timeout, + ) + except asyncio.TimeoutError: + errors.append("deadline_exceeded") + except Exception as exc: + logger.debug("Context window probe failed: %s", type(exc).__name__) + errors.append("probe_failed") + return _ProbeResult(tuple(observations), tuple(errors), io=True) + + +async def _is_local(endpoint_url: str) -> bool: + from src.model_context import is_local_endpoint + + try: + # Reads configured endpoints from the local database on the calling + # thread, as the regular runtime does. Moving it to worker threads + # gives SQLite sessions per-thread connections the app never uses. + return bool(is_local_endpoint(endpoint_url)) + except Exception: + return False + + +async def _cached_probe(endpoint_url, model, headers, deadline_seconds, clock): + is_local = await _is_local(endpoint_url) + key = (endpoint_url, model, _auth_fingerprint(headers)) + if not is_local: + cached = _probe_cache.get(key) + if cached and cached[0] > clock(): + return cached[1], True + result = await probe_provider_context( + endpoint_url, model, headers=headers, deadline_seconds=deadline_seconds, + is_local=is_local, + ) + if not is_local: + ttl = PROBE_CACHE_TTL_SECONDS if result.observations else PROBE_FAILURE_TTL_SECONDS + _probe_cache[key] = (clock() + ttl, result) + return result, False + + +def declared_context_window(client_runtime_context: Any) -> int: + """Operator/client declared transport window, or 0.""" + if not isinstance(client_runtime_context, Mapping): + return 0 + try: + value = int(client_runtime_context.get("model_context_window") or 0) + except (TypeError, ValueError): + return 0 + return value if value > 0 else 0 + + +async def resolve_effective_context( + endpoint_url: str, + model: str, + *, + headers: Optional[Mapping[str, Any]] = None, + client_runtime_context: Any = None, + deadline_seconds: float = PROBE_DEADLINE_SECONDS, + probe: bool = True, + clock=time.monotonic, +) -> ContextResolution: + """Resolve the effective context window for one turn preparation.""" + from src.model_context import _lookup_known + + observations: list[ContextObservation] = [] + errors: tuple[str, ...] = () + provider_io = cached = False + if probe and endpoint_url and model: + result, cached = await _cached_probe( + endpoint_url, model, headers, deadline_seconds, clock, + ) + observations.extend(result.observations) + errors = result.errors + provider_io = result.io and not cached + declared = declared_context_window(client_runtime_context) + if declared: + observations.append(ContextObservation( + ContextEvidence.OPERATOR_DECLARED, declared, "client_runtime_context", + )) + known = _lookup_known(model or "") + if known: + observations.append(ContextObservation(ContextEvidence.KNOWN_TABLE, int(known), "known_table")) + resolution = combine_observations(observations) + resolution = replace(resolution, provider_io=provider_io, cached=cached, probe_errors=errors) + if resolution.mismatch: + logger.info( + "Context window sources disagree for %s: %s", + model, [conflict.to_dict() for conflict in resolution.conflicts], + ) + return resolution + + +def context_metrics(resolution: Optional[ContextResolution], request_tokens: int) -> dict: + """Metrics fields derived only from a stored resolution. Performs no I/O.""" + resolution = resolution or UNRESOLVED_CONTEXT + length = resolution.budget_limit + percent = ( + min(round((request_tokens / length) * 100, 1), 100.0) + if length and request_tokens else 0 + ) + return { + "context_length": length, + "context_percent": percent, + "context_resolution": resolution.to_dict(), + } diff --git a/src/clean_agent_preview.py b/src/clean_agent_preview.py index 4c0302eb7..78dce14ff 100644 --- a/src/clean_agent_preview.py +++ b/src/clean_agent_preview.py @@ -4904,7 +4904,9 @@ async def preview_model_response(client, endpoint_url, headers, request, recover ) transport_attempts = 0 while True: - limit = recovery.get('context_limit') + # A provider-stated limit learned this turn overrides the window + # resolved at turn preparation. + limit = recovery.get('context_limit') or recovery.get('budget_limit') if limit: message_context = max(1, limit - estimate_tool_schema_tokens(request.get('tools')) @@ -4995,6 +4997,7 @@ async def stream_preview(*, endpoint_url, model, messages, headers, turn_contrac client_runtime_context=None, max_tokens=768, max_rounds=8, max_tool_calls=0, external_tool_schemas=None, temperature=0.0, + context_resolution=None, request_authority=MISSING_AUTHORITY, **ignored): from src.generation_sampling import validate_temperature temperature = validate_temperature(temperature) @@ -5249,7 +5252,16 @@ async def stream_preview(*, endpoint_url, model, messages, headers, turn_contrac static_fetch_failed_urls = set() entity_result_links = {} calendar_create_confirmation = '' - context_recovery = {} + # Resolve the effective window once, before any model request, so the + # turn budgets against it and metrics report exactly what it ran under. + # Terminal metrics must never start this discovery themselves. + if context_resolution is None: + from src.agent_runtime.context_resolution import resolve_effective_context + context_resolution = await resolve_effective_context( + endpoint_url, model, headers=headers, + client_runtime_context=client_runtime_context, + ) + context_recovery = {'budget_limit': context_resolution.budget_limit} successful_write = False editor_batch_pending = False editor_suggested_finds = [] @@ -8026,9 +8038,15 @@ async def stream_preview(*, endpoint_url, model, messages, headers, turn_contrac return elapsed = time.monotonic() - started ttft = first_token - started if first_token else None + from src.agent_runtime.context_resolution import context_metrics + # Report the stored resolution plus any limit the provider stated during + # this turn. This is pure bookkeeping; no metadata request happens here. + context_resolution = context_resolution.observe_runtime_limit( + context_recovery.get('context_limit')) yield event({'type': 'metrics', 'data': { 'email_task_scope': {**intent_accounting, 'failed': intent_scope_failed}, 'model': model, 'input_tokens': usage_in, 'output_tokens': usage_out, + **context_metrics(context_resolution, last_request_tokens), 'total_tokens': usage_in + usage_out, 'response_time': round(elapsed, 3), 'time_to_first_token': round(ttft, 3) if ttft is not None else None, 'tokens_per_second': round(usage_out / elapsed, 2) if elapsed > 0 else 0, diff --git a/tests/conftest.py b/tests/conftest.py index 0affe928a..c73185eaa 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -201,3 +201,36 @@ def _no_leaked_module_stubs(): "at teardown.", pytrace=False, ) + + +@pytest.fixture(autouse=True) +def _no_context_window_network_probe(request): + """Keep turn preparation from probing real provider metadata in tests. + + The compact runtime resolves its context window before the first model + request. Tests that drive it with placeholder endpoints must not perform + DNS or HTTP lookups; modules that exercise the probe opt in with a + module-level ``CONTEXT_PROBE_NETWORK = True`` and supply their own client. + """ + if getattr(request.module, "CONTEXT_PROBE_NETWORK", False): + yield + return + try: + from src.agent_runtime import context_resolution + except Exception: + yield + return + + async def _disabled_probe(endpoint_url, model, headers, is_local, observations, errors, timeout): + errors.append("probe_disabled_in_tests") + + # A private patcher keeps the shared ``monkeypatch`` fixture's teardown + # order unchanged for tests that check their own sys.modules hygiene. + patcher = pytest.MonkeyPatch() + patcher.setattr(context_resolution, "_probe", _disabled_probe) + context_resolution.clear_probe_cache() + try: + yield + finally: + patcher.undo() + context_resolution.clear_probe_cache() diff --git a/tests/test_context_resolution.py b/tests/test_context_resolution.py new file mode 100644 index 000000000..4290c8aca --- /dev/null +++ b/tests/test_context_resolution.py @@ -0,0 +1,513 @@ +"""Turn-preparation context window resolution and its compact-runtime use.""" +import asyncio +from types import SimpleNamespace +import json + +import httpx +import pytest + +import src.model_context as model_context +from src.agent_runtime import context_resolution as cr +from src.agent_runtime.context_resolution import ( + ContextEvidence, + ContextObservation, + UNRESOLVED_CONTEXT, + combine_observations, + context_metrics, + resolve_effective_context, +) +from src.clean_agent_preview import stream_preview +from src.tool_policy import ToolPolicy +from src.turn_contract import resolve_full_inventory_contract + +REMOTE = "http://provider.test/v1/chat/completions" +LOCAL = "http://127.0.0.1:8080/v1/chat/completions" +AUTH = {"Authorization": "Bearer secret-token", "Content-Type": "application/json"} + +# Opt out of the conftest guard; every test here installs a fake HTTP client. +CONTEXT_PROBE_NETWORK = True + + +def _obs(evidence, value, source="test"): + return ContextObservation(evidence, value, source) + + +class _Response: + def __init__(self, status=200, payload=None): + self.status_code = status + self._payload = payload + + def json(self): + if isinstance(self._payload, Exception): + raise self._payload + return self._payload + + +class _ModelStream: + def __init__(self, lines, status=200, text=""): + self.status_code = status + self.text = text + self._lines = lines + + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + async def aread(self): + return self.text.encode() + + def raise_for_status(self): + if self.status_code >= 400: + raise AssertionError(f"unexpected provider status {self.status_code}") + + async def aiter_lines(self): + for line in self._lines: + yield line + + +def _answer_lines(usage=None): + lines = ["data: " + json.dumps({"choices": [{"delta": {"content": "Done."}}]})] + if usage is not None: + lines.append("data: " + json.dumps({"choices": [{"delta": {}}], "usage": usage})) + lines.append("data: [DONE]") + return lines + + +class _FakeNetwork: + """One fake HTTP surface for both metadata GETs and model streams. + + The compact runtime and the probe share ``httpx.AsyncClient``; recording + every call in order proves which phase performed which I/O. + """ + + def __init__(self, routes=None, *, stream_responses=None, get_delay=0.0, get_error=None): + self.routes = routes or {} + self.stream_responses = list(stream_responses or []) + self.get_delay = get_delay + self.get_error = get_error + self.calls = [] + self.requests = [] + + def client_factory(self): + network = self + + class Client: + def __init__(self, **kwargs): + pass + + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + async def get(self, url, headers=None): + network.calls.append(("get", url, dict(headers or {}))) + if network.get_delay: + await asyncio.sleep(network.get_delay) + if network.get_error is not None: + raise network.get_error + for suffix, response in network.routes.items(): + if url.endswith(suffix): + return response + return _Response(404, None) + + def stream(self, method, url, headers=None, json=None): + network.calls.append(("stream", url, dict(headers or {}))) + network.requests.append(json) + if network.stream_responses: + return network.stream_responses.pop(0) + return _ModelStream(_answer_lines({"prompt_tokens": 2048, "completion_tokens": 3})) + + return Client + + +@pytest.fixture(autouse=True) +def _fresh_cache(monkeypatch): + cr.clear_probe_cache() + monkeypatch.setattr(model_context, "is_local_endpoint", lambda url: "127.0.0.1" in url) + # The real URL builder may resolve hosts; these tests stay off the network. + monkeypatch.setattr( + "src.endpoint_resolver.build_models_url", lambda base: base.split("/v1")[0] + "/v1/models", + ) + monkeypatch.setattr("src.endpoint_resolver.resolve_url", lambda url: url) + yield + cr.clear_probe_cache() + + +def _install(monkeypatch, network): + monkeypatch.setattr(cr.httpx, "AsyncClient", network.client_factory()) + + def legacy_get(url, *args, **kwargs): + # The legacy model_context probe is synchronous; record it so a + # terminal-metrics probe through that path is caught as well. + network.calls.append(("sync_get", url, dict(kwargs.get("headers") or {}))) + return _Response(404, None) + + monkeypatch.setattr(cr.httpx, "get", legacy_get) + return network + + +def _catalog(model, **fields): + return _Response(200, {"data": [{"id": model, **fields}]}) + + +# --------------------------------------------------------------------------- +# Pure combination rules +# --------------------------------------------------------------------------- + +def test_operator_declared_window_replaces_known_table_without_contradiction(): + resolution = combine_observations([ + _obs(ContextEvidence.KNOWN_TABLE, 128000), + _obs(ContextEvidence.OPERATOR_DECLARED, 200000), + ]) + assert resolution.effective == 200000 + assert resolution.evidence is ContextEvidence.OPERATOR_DECLARED + assert not resolution.mismatch + + +def test_operator_declared_window_caps_measured_window(): + resolution = combine_observations([ + _obs(ContextEvidence.PROVIDER_ADVERTISED, 131072), + _obs(ContextEvidence.OPERATOR_DECLARED, 32768), + ]) + assert (resolution.effective, resolution.evidence) == (32768, ContextEvidence.OPERATOR_DECLARED) + # A tighter declared transport limit is a cap, not a contradiction. + assert not resolution.mismatch + + +def test_operator_declaration_above_provider_is_contradiction_and_provider_wins(): + resolution = combine_observations([ + _obs(ContextEvidence.PROVIDER_ADVERTISED, 8192), + _obs(ContextEvidence.OPERATOR_DECLARED, 32768), + ]) + assert (resolution.effective, resolution.evidence) == (8192, ContextEvidence.PROVIDER_ADVERTISED) + assert resolution.mismatch + assert {c.evidence for c in (resolution.conflicts[0].first, resolution.conflicts[0].second)} == { + ContextEvidence.PROVIDER_ADVERTISED, ContextEvidence.OPERATOR_DECLARED, + } + + +def test_provider_advertised_beats_known_table_and_disagreement_is_visible(): + resolution = combine_observations([ + _obs(ContextEvidence.PROVIDER_ADVERTISED, 8192), + _obs(ContextEvidence.KNOWN_TABLE, 131072), + ]) + # The legacy probe takes max(api, table) for cloud endpoints. A static + # table is weaker evidence and must not override the provider silently. + assert (resolution.effective, resolution.evidence) == (8192, ContextEvidence.PROVIDER_ADVERTISED) + assert resolution.mismatch + + +def test_runtime_confirmed_beats_provider_advertised_and_records_mismatch(): + resolution = combine_observations([ + _obs(ContextEvidence.PROVIDER_ADVERTISED, 32768), + _obs(ContextEvidence.RUNTIME_CONFIRMED, 16384), + ]) + assert (resolution.effective, resolution.evidence) == (16384, ContextEvidence.RUNTIME_CONFIRMED) + assert resolution.mismatch + + +def test_no_evidence_is_unknown_zero_not_a_default(): + resolution = combine_observations([]) + assert resolution.effective == 0 + assert resolution.evidence is ContextEvidence.UNKNOWN + assert resolution.budget_limit == 0 + assert context_metrics(resolution, 500)["context_length"] == 0 + + +def test_runtime_limit_observation_is_pure_and_lowers_effective_window(): + base = combine_observations([_obs(ContextEvidence.PROVIDER_ADVERTISED, 8192)]) + updated = base.observe_runtime_limit(4096) + assert (updated.effective, updated.evidence, updated.source) == ( + 4096, ContextEvidence.RUNTIME_CONFIRMED, "provider_rejection", + ) + assert updated.mismatch + assert base.observe_runtime_limit(None) is base + assert base.observe_runtime_limit(0) is base + + +def test_context_metrics_reports_percent_against_stored_window(): + resolution = combine_observations([_obs(ContextEvidence.KNOWN_TABLE, 8000)]) + metrics = context_metrics(resolution, 2000) + assert metrics["context_length"] == 8000 + assert metrics["context_percent"] == 25.0 + assert metrics["context_resolution"]["evidence"] == "known_table" + assert metrics["context_resolution"]["mismatch"] is False + + +# --------------------------------------------------------------------------- +# Provider probe +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_provider_advertised_window_uses_turn_credentials(monkeypatch): + network = _install(monkeypatch, _FakeNetwork({"/models": _catalog("acme-model", max_model_len=8192)})) + resolution = await resolve_effective_context(REMOTE, "acme-model", headers=AUTH) + assert (resolution.effective, resolution.evidence, resolution.source) == ( + 8192, ContextEvidence.PROVIDER_ADVERTISED, "models_catalog", + ) + assert resolution.provider_io and not resolution.cached + [(kind, url, headers)] = network.calls + assert kind == "get" and url.endswith("/models") + assert headers == {"Authorization": "Bearer secret-token"} + + +@pytest.mark.asyncio +async def test_credentials_are_not_forwarded_to_another_origin(monkeypatch): + network = _install(monkeypatch, _FakeNetwork({"/models": _catalog("acme-model", max_model_len=8192)})) + monkeypatch.setattr( + "src.endpoint_resolver.build_models_url", lambda base: "http://catalog.elsewhere.test/v1/models", + ) + await resolve_effective_context(REMOTE, "acme-model", headers=AUTH) + assert network.calls[0][2] == {} + + +@pytest.mark.asyncio +async def test_credentials_follow_the_resolved_provider_host(monkeypatch): + network = _install(monkeypatch, _FakeNetwork({"/models": _catalog("acme-model", max_model_len=8192)})) + monkeypatch.setattr( + "src.endpoint_resolver.build_models_url", lambda base: "http://100.64.0.9/v1/models", + ) + monkeypatch.setattr( + "src.endpoint_resolver.resolve_url", lambda url: url.replace("provider.test", "100.64.0.9"), + ) + await resolve_effective_context(REMOTE, "acme-model", headers=AUTH) + assert network.calls[0][2] == {"Authorization": "Bearer secret-token"} + + +@pytest.mark.asyncio +async def test_slow_url_resolution_is_bounded_by_the_probe_deadline(monkeypatch): + import time as _time + + _install(monkeypatch, _FakeNetwork()) + + def slow_models_url(base): + _time.sleep(0.5) + return base.split("/v1")[0] + "/v1/models" + + monkeypatch.setattr("src.endpoint_resolver.build_models_url", slow_models_url) + loop = asyncio.get_running_loop() + started = loop.time() + resolution = await resolve_effective_context(REMOTE, "gpt-4o", deadline_seconds=0.05) + assert loop.time() - started < 0.4 + assert resolution.probe_errors == ("deadline_exceeded",) + assert resolution.evidence is ContextEvidence.KNOWN_TABLE + + +@pytest.mark.asyncio +async def test_known_table_fallback_when_provider_lists_no_window(monkeypatch): + _install(monkeypatch, _FakeNetwork({"/models": _catalog("gpt-4o-mini")})) + resolution = await resolve_effective_context(REMOTE, "gpt-4o-mini", headers=AUTH) + assert (resolution.effective, resolution.evidence) == (128000, ContextEvidence.KNOWN_TABLE) + assert "models:no_window_listed" in resolution.probe_errors + + +@pytest.mark.asyncio +async def test_unavailable_metadata_and_unknown_model_is_unknown(monkeypatch): + _install(monkeypatch, _FakeNetwork({"/models": _Response(503, None)})) + resolution = await resolve_effective_context(REMOTE, "mystery-model", headers=AUTH) + assert resolution.evidence is ContextEvidence.UNKNOWN + assert resolution.effective == 0 + assert resolution.probe_errors == ("models:http_503",) + + +@pytest.mark.asyncio +async def test_rejected_credentials_are_reported_without_leaking_them(monkeypatch): + _install(monkeypatch, _FakeNetwork({"/models": _Response(401, None)})) + resolution = await resolve_effective_context(REMOTE, "gpt-4o", headers=AUTH) + assert resolution.evidence is ContextEvidence.KNOWN_TABLE + assert resolution.probe_errors == ("models:http_401",) + assert "secret-token" not in json.dumps(resolution.to_dict()) + + +@pytest.mark.asyncio +async def test_provider_timeout_is_bounded_and_falls_back(monkeypatch): + _install(monkeypatch, _FakeNetwork(get_delay=5.0)) + loop = asyncio.get_running_loop() + started = loop.time() + resolution = await resolve_effective_context( + REMOTE, "gpt-4o", headers=AUTH, deadline_seconds=0.05, + ) + assert loop.time() - started < 1.0 + assert "deadline_exceeded" in resolution.probe_errors + assert (resolution.effective, resolution.evidence) == (128000, ContextEvidence.KNOWN_TABLE) + + +@pytest.mark.asyncio +async def test_provider_transport_failure_never_raises(monkeypatch): + _install(monkeypatch, _FakeNetwork(get_error=httpx.ConnectError("refused"))) + resolution = await resolve_effective_context( + REMOTE, "mystery-model", client_runtime_context={"model_context_window": 16384}, + ) + assert resolution.probe_errors == ("models:transport_error",) + assert (resolution.effective, resolution.evidence) == (16384, ContextEvidence.OPERATOR_DECLARED) + + +@pytest.mark.asyncio +async def test_local_runtime_confirmed_slots_and_provider_mismatch(monkeypatch): + network = _install(monkeypatch, _FakeNetwork({ + "/slots": _Response(200, [{"n_ctx": 16384}]), + "/models": _catalog("local-model", max_model_len=32768), + })) + resolution = await resolve_effective_context(LOCAL, "local-model") + assert (resolution.effective, resolution.evidence, resolution.source) == ( + 16384, ContextEvidence.RUNTIME_CONFIRMED, "llamacpp_slots", + ) + assert resolution.mismatch + assert [url.rsplit("/", 1)[-1] for _, url, _ in network.calls] == ["slots", "models"] + + +@pytest.mark.asyncio +async def test_local_runtime_confirmed_props_when_slots_disabled(monkeypatch): + _install(monkeypatch, _FakeNetwork({ + "/slots": _Response(501, None), + "/props": _Response(200, {"default_generation_settings": {"n_ctx": 8192}}), + "/models": _catalog("local-model"), + })) + resolution = await resolve_effective_context(LOCAL, "local-model") + assert (resolution.effective, resolution.source) == (8192, "llamacpp_props") + assert not resolution.mismatch + + +@pytest.mark.asyncio +async def test_remote_resolution_is_cached_per_credentials(monkeypatch): + network = _install(monkeypatch, _FakeNetwork({"/models": _catalog("acme-model", max_model_len=8192)})) + first = await resolve_effective_context(REMOTE, "acme-model", headers=AUTH) + second = await resolve_effective_context(REMOTE, "acme-model", headers=AUTH) + assert first.effective == second.effective == 8192 + assert second.cached and not second.provider_io + assert len(network.calls) == 1 + await resolve_effective_context(REMOTE, "acme-model", headers={"Authorization": "Bearer other"}) + assert len(network.calls) == 2 + + +@pytest.mark.asyncio +async def test_failed_remote_probe_expires_sooner(monkeypatch): + network = _install(monkeypatch, _FakeNetwork({"/models": _Response(503, None)})) + now = [1000.0] + await resolve_effective_context(REMOTE, "acme-model", clock=lambda: now[0]) + await resolve_effective_context(REMOTE, "acme-model", clock=lambda: now[0]) + assert len(network.calls) == 1 + now[0] += cr.PROBE_FAILURE_TTL_SECONDS + 1 + await resolve_effective_context(REMOTE, "acme-model", clock=lambda: now[0]) + assert len(network.calls) == 2 + + +@pytest.mark.asyncio +async def test_local_resolution_is_not_cached(monkeypatch): + network = _install(monkeypatch, _FakeNetwork({"/slots": _Response(200, [{"n_ctx": 4096}])})) + await resolve_effective_context(LOCAL, "local-model") + await resolve_effective_context(LOCAL, "local-model") + assert sum(1 for kind, url, _ in network.calls if url.endswith("/slots")) == 2 + + +# --------------------------------------------------------------------------- +# Compact runtime integration +# --------------------------------------------------------------------------- + +async def _run_preview(messages=None, **kwargs): + contract = resolve_full_inventory_contract(schemas=[], policy=ToolPolicy()) + raw = [chunk async for chunk in stream_preview( + endpoint_url=kwargs.pop("endpoint_url", REMOTE), model=kwargs.pop("model", "acme-model"), + messages=messages or [{"role": "user", "content": "hi"}], + headers=kwargs.pop("headers", AUTH), + turn_contract=contract, session_id="test", owner="test", + disabled_tools=set(), tool_policy=ToolPolicy(), **kwargs, + )] + assert raw[-1] == "data: [DONE]\n\n" + events = [json.loads(chunk[6:]) for chunk in raw if chunk.startswith("data: {")] + return next(event["data"] for event in events if event.get("type") == "metrics") + + +@pytest.mark.asyncio +async def test_compact_turn_resolves_once_before_model_and_metrics_do_no_discovery(monkeypatch): + network = _install(monkeypatch, _FakeNetwork({"/models": _catalog("acme-model", max_model_len=8192)})) + metrics = await _run_preview() + kinds = [kind for kind, _, _ in network.calls] + first_stream = kinds.index("stream") + # All discovery precedes the first model request; nothing after it, + # in particular nothing between the last model byte and [DONE]. + assert kinds[:first_stream] == ["get"] + assert kinds[first_stream:] == ["stream"] + assert metrics["context_length"] == 8192 + assert metrics["context_percent"] == 25.0 + assert metrics["context_resolution"]["evidence"] == "provider_advertised" + assert metrics["context_resolution"]["provider_io"] is True + + +@pytest.mark.asyncio +async def test_compact_turn_with_supplied_resolution_performs_no_metadata_io(monkeypatch): + network = _install(monkeypatch, _FakeNetwork()) + supplied = combine_observations([_obs(ContextEvidence.OPERATOR_DECLARED, 4096, "client_runtime_context")]) + metrics = await _run_preview(context_resolution=supplied) + assert [kind for kind, _, _ in network.calls] == ["stream"] + assert metrics["context_length"] == 4096 + assert metrics["context_resolution"] == supplied.to_dict() + + +@pytest.mark.asyncio +async def test_compact_metrics_report_unknown_window_as_zero(monkeypatch): + _install(monkeypatch, _FakeNetwork({"/models": _Response(404, None)})) + metrics = await _run_preview(model="mystery-model") + assert metrics["context_length"] == 0 + assert metrics["context_percent"] == 0 + assert metrics["context_resolution"]["evidence"] == "unknown" + assert metrics["context_resolution"]["probe_errors"] == ["models:http_404"] + + +@pytest.mark.asyncio +async def test_compact_runtime_budgets_against_resolved_window(monkeypatch): + turns = [] + for index in range(6): + turns.append({"role": "user", "content": f"question {index} " + "word " * 280}) + turns.append({"role": "assistant", "content": f"answer {index} " + "word " * 280}) + session = SimpleNamespace(history=turns) + + unbudgeted = _install(monkeypatch, _FakeNetwork()) + await _run_preview(history_session=session, context_resolution=UNRESOLVED_CONTEXT) + budgeted = _install(monkeypatch, _FakeNetwork()) + resolved = combine_observations([_obs(ContextEvidence.PROVIDER_ADVERTISED, 4096, "models_catalog")]) + await _run_preview(history_session=session, context_resolution=resolved) + + full = unbudgeted.requests[0]["messages"] + trimmed = budgeted.requests[0]["messages"] + # An unknown window keeps the previous reactive-only behavior. + assert len(full) > len(trimmed) + assert model_context.estimate_tokens(trimmed) < 4096 + assert trimmed[-1]["content"] == "hi" + + +@pytest.mark.asyncio +async def test_compact_metrics_fold_provider_stated_limit_without_new_discovery(monkeypatch): + rejection = ( + "This model's maximum context length is 4096 tokens. However, you requested " + "5000 tokens (4800 in the messages, 200 in the completion)." + ) + network = _install(monkeypatch, _FakeNetwork( + {"/models": _catalog("acme-model", max_model_len=8192)}, + stream_responses=[ + _ModelStream([], status=400, text=json.dumps({"error": {"message": rejection}})), + _ModelStream(_answer_lines({"prompt_tokens": 1024, "completion_tokens": 3})), + ], + )) + metrics = await _run_preview() + assert [kind for kind, _, _ in network.calls] == ["get", "stream", "stream"] + resolution = metrics["context_resolution"] + assert metrics["context_length"] == 4096 + assert (resolution["evidence"], resolution["source"]) == ("runtime_confirmed", "provider_rejection") + assert resolution["mismatch"] is True + + +@pytest.mark.asyncio +async def test_compact_turn_proceeds_when_metadata_probe_times_out(monkeypatch): + _install(monkeypatch, _FakeNetwork(get_error=httpx.ReadTimeout("slow catalog"))) + metrics = await _run_preview( + model="gpt-4o", client_runtime_context={"model_context_window": 65536}, + ) + resolution = metrics["context_resolution"] + assert resolution["probe_errors"] == ["models:timeout"] + assert metrics["context_length"] == 65536 + assert resolution["evidence"] == "operator_declared"