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"