diff --git a/agent/model_metadata.py b/agent/model_metadata.py index feccf70cd2..940a5107ad 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -320,6 +320,16 @@ MINIMUM_CONTEXT_LENGTH = 64_000 # startup resolves the same model several times (banner, /model, compressor). Never persisted. _LOCAL_CTX_PROBE_TTL_SECONDS = 30.0 _LOCAL_CTX_PROBE_CACHE: Dict[tuple, tuple] = {} +# In-process (model, region) -> monotonic_ts memo of a FAILED Bedrock context probe. That probe pads +# prompts of 1.3M/2.2M tokens and attempts up to two converse calls (agent/bedrock_adapter.py +# _BEDROCK_PROBE_TIERS; the second tier is sent only when the first yields no parseable limit, i.e. on +# the failure path), and its failures are deliberately never persisted, so without a negative memo a +# model whose probe keeps failing (un-enabled model, opaque InternalServerException, unparseable length +# error) would re-send them on every resolution. Negative only, in memory only, and bounded — same +# reasoning as _ENDPOINT_PROBE_FAILURE_TTL_SECONDS: a failure is usually transient (expired SSO +# session, offline box), so it must expire rather than stick. +_BEDROCK_PROBE_FAILURE_TTL_SECONDS = 300.0 +_BEDROCK_PROBE_FAILURE_CACHE: Dict[tuple, float] = {} # Family-pattern fallbacks, used only when provider-aware sources all miss. # Lookups are longest-key-first substring matches, so dict order is cosmetic # and a specific key must be STRICTLY longer than its catch-all. @@ -1145,6 +1155,11 @@ def _invalidate_cached_context_length(model: str, base_url: str) -> None: bare, stripped = _strip_provider_prefix(model), (base_url or "").rstrip("/") _LOCAL_CTX_PROBE_CACHE.pop((bare, stripped), None) _LOCAL_CTX_PROBE_CACHE.pop(("ollama_show", bare, stripped), None) + # Same for a memoised Bedrock probe failure (keyed by region, which the caller does not know): + # the entry being dropped is the reason to ask the probe again, not to wait out its TTL. + for memo_key in list(_BEDROCK_PROBE_FAILURE_CACHE): # snapshot: another thread may be memoising + if memo_key[0] in (model, bare): + _BEDROCK_PROBE_FAILURE_CACHE.pop(memo_key, None) # Every key shape get_cached_context_length consults. stale_keys = {key, f"{model}@{base_url}", f"{key}/"} if not any(k in cache for k in stale_keys): @@ -1813,12 +1828,20 @@ def _validate_cached_context_length(model: str, base_url: str, cached: int, is_b return cached +def _bedrock_probe_failed_recently(model: str, region: str) -> bool: + """True while a failed Bedrock context probe for *model* in *region* is still memoised + (see _BEDROCK_PROBE_FAILURE_CACHE): answer from the static table without re-probing.""" + failed_at = _BEDROCK_PROBE_FAILURE_CACHE.get((model, region)) + return failed_at is not None and (time.monotonic() - failed_at) < _BEDROCK_PROBE_FAILURE_TTL_SECONDS + + def _resolve_bedrock_context_length(model: str, base_url: str) -> Optional[int]: """Step 1b: Bedrock static table + one cached live probe (Bedrock exposes no context window via - metadata APIs); None when boto3 is absent. Cached per model under base_url, else a synthetic - bedrock:// key so display/offline paths share it.""" + metadata APIs); None when boto3 is absent. Only a PROBED window is cached (the table answers a + call, never the cache), per model under base_url, else a synthetic bedrock:// key so + display/offline paths share it.""" try: - from agent.bedrock_adapter import get_bedrock_context_length, resolve_bedrock_region + from agent.bedrock_adapter import get_bedrock_context_length, probe_bedrock_context_length, resolve_bedrock_region except ImportError: return None # boto3 not installed — fall through to generic resolution cache_key_url = base_url or "bedrock://" @@ -1831,11 +1854,17 @@ def _resolve_bedrock_context_length(model: str, base_url: str) -> Optional[int]: if not region: with contextlib.suppress(Exception): region = resolve_bedrock_region() - ctx = get_bedrock_context_length(model, region=region, probe=bool(region)) - # Only persist probe-derived values (region present); a pure table fallback must not poison the cache. - if ctx and region: - save_context_length(model, cache_key_url, ctx) - return ctx + if region and not _bedrock_probe_failed_recently(model, region): + probed = probe_bedrock_context_length(model, region) + if probed: + # The probe is the only authoritative source, so it is the only thing worth persisting: + # a table fallback written here would be served forever (this branch runs before it), + # and the probe would never be consulted for the model again. + save_context_length(model, cache_key_url, probed) + _BEDROCK_PROBE_FAILURE_CACHE.pop((model, region), None) # success ends the failure window + return probed + _BEDROCK_PROBE_FAILURE_CACHE[(model, region)] = time.monotonic() + return get_bedrock_context_length(model, probe=False) # static table / default: answers this call only def _resolve_custom_endpoint_context_length(model: str, base_url: str, api_key: str, provider: str) -> int: diff --git a/tests/agent/test_model_metadata.py b/tests/agent/test_model_metadata.py index 868976e8e9..a4dee9044f 100644 --- a/tests/agent/test_model_metadata.py +++ b/tests/agent/test_model_metadata.py @@ -1292,6 +1292,73 @@ class TestBedrockContextResolution: assert mock_fetch.called +# ========================================================================= +# Bedrock context cache persistence — only a probe result may be persisted +# ========================================================================= + +class TestBedrockContextCachePersistence: + """``_resolve_bedrock_context_length`` persisted whatever + ``get_bedrock_context_length`` returned whenever a region was resolvable — + and ``resolve_bedrock_region()`` always resolves one (it ends in + ``or "us-east-1"``). A probe that returned None (expired SSO session, + offline, opaque server error) therefore froze the static table value — the + 128K default for a model with no table row — under ``model@`` or + ``model@bedrock://`` (written as ``model@bedrock:``), and the probe, which + the resolver treats as the only authoritative source, never ran for that + model again. + + Invariants: only a probe-derived window may be persisted, and a failed + probe is memoised in memory for ``_BEDROCK_PROBE_FAILURE_TTL_SECONDS`` so + it is not re-sent on every resolution, yet runs again once the memo lapses. + """ + + @pytest.fixture(autouse=True) + def _clear_bedrock_probe_memo(self): + """The failure memo is module-level state; it must not leak between tests.""" + from agent import model_metadata as mm + mm._BEDROCK_PROBE_FAILURE_CACHE.clear() + yield + mm._BEDROCK_PROBE_FAILURE_CACHE.clear() + + @patch("agent.bedrock_adapter.resolve_bedrock_region", return_value="us-east-1") + @patch("agent.bedrock_adapter.probe_bedrock_context_length", return_value=None) + def test_failed_probe_does_not_persist_static_fallback(self, mock_probe, mock_region, tmp_path): + """A failed probe answers from the table (the 128K default here) but writes nothing + to disk. On main this persists ``amazon.future-model-v1:0@bedrock:: 128000``.""" + from agent.bedrock_adapter import BEDROCK_DEFAULT_CONTEXT_LENGTH + model = "amazon.future-model-v1:0" # no BEDROCK_CONTEXT_LENGTHS row + cache_file = tmp_path / "context_length_cache.yaml" + with patch("agent.model_metadata._get_context_cache_path", return_value=cache_file): + assert get_model_context_length(model, provider="bedrock") == BEDROCK_DEFAULT_CONTEXT_LENGTH + assert get_cached_context_length(model, "bedrock://") is None + assert not cache_file.exists() + mock_probe.assert_called_once_with(model, "us-east-1") + + @patch("agent.bedrock_adapter.resolve_bedrock_region", return_value="us-east-1") + @patch("agent.bedrock_adapter.probe_bedrock_context_length", side_effect=[None, 1_000_000]) + def test_failed_probe_is_memoised_until_the_ttl_lapses(self, mock_probe, mock_region, tmp_path): + """Not persisting must not turn the probe into a per-resolution cost: inside the TTL a + second resolution answers from the table without re-probing and still writes nothing; + once the memo lapses the probe runs again and its window is what gets persisted. On + main the first call persists 128K and every later call serves it.""" + from agent import model_metadata as mm + from agent.bedrock_adapter import BEDROCK_DEFAULT_CONTEXT_LENGTH + model = "amazon.future-model-v1:0" + cache_file = tmp_path / "context_length_cache.yaml" + with patch("agent.model_metadata._get_context_cache_path", return_value=cache_file): + assert get_model_context_length(model, provider="bedrock") == BEDROCK_DEFAULT_CONTEXT_LENGTH + assert get_model_context_length(model, provider="bedrock") == BEDROCK_DEFAULT_CONTEXT_LENGTH + assert mock_probe.call_count == 1 + assert not cache_file.exists() # memoised in memory only + # Age the entry past the failure TTL (as tests/agent/test_probe_cache_followups.py does). + mm._BEDROCK_PROBE_FAILURE_CACHE[(model, "us-east-1")] = ( + time.monotonic() - mm._BEDROCK_PROBE_FAILURE_TTL_SECONDS - 1 + ) + assert get_model_context_length(model, provider="bedrock") == 1_000_000 + assert get_cached_context_length(model, "bedrock://") == 1_000_000 + assert mock_probe.call_count == 2 + + # ========================================================================= # _strip_provider_prefix — Ollama model:tag vs provider:model # =========================================================================