Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 21 additions & 3 deletions agent/model_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -1933,8 +1933,19 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str
resp = client.get(f"{server_url}/v1/models/{model}")
if resp.status_code == 200:
data = resp.json()
# vLLM returns max_model_len
ctx = data.get("max_model_len") or data.get("context_length") or data.get("max_tokens")
# Read genuine context-WINDOW fields only. `max_tokens` is
# excluded on purpose: on an Anthropic-style /v1/models/{id}
# passthrough it is the max *output* tokens (e.g. 128000 for a
# 1M-context model), NOT the context window — reading it here
# collapses the window and triggers premature auto-compaction.
# `max_input_tokens` IS the input/context ceiling, so it's kept.
ctx = (
data.get("max_model_len")
or data.get("context_length")
or data.get("context_window")
or data.get("max_input_tokens")
or data.get("max_position_embeddings")
)
if ctx and isinstance(ctx, (int, float)):
return int(ctx)

Expand All @@ -1946,7 +1957,14 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str
models_list = data.get("data", [])
for m in models_list:
if _model_id_matches(m.get("id", ""), model):
ctx = m.get("max_model_len") or m.get("context_length") or m.get("max_tokens")
# Context-window fields only (see the max_tokens note above).
ctx = (
m.get("max_model_len")
or m.get("context_length")
or m.get("context_window")
or m.get("max_input_tokens")
or m.get("max_position_embeddings")
)
if ctx and isinstance(ctx, (int, float)):
return int(ctx)
except Exception:
Expand Down
55 changes: 55 additions & 0 deletions tests/agent/test_model_metadata_local_ctx.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,61 @@ def test_vllm_context_length_key(self):

assert result == 32768

def test_anthropic_passthrough_reads_max_input_tokens_not_max_tokens(self):
"""Regression: an Anthropic-style /v1/models/{id} passthrough returns
`max_input_tokens` (the context window) AND `max_tokens` (the max OUTPUT
tokens). The resolver must read max_input_tokens, never max_tokens —
reading max_tokens collapsed a 1M-context model to its 128K output cap
and triggered premature auto-compaction (~96K on a 1M model)."""
from agent.model_metadata import _query_local_context_length

# Shape mirrors the fleet's Claude proxy passthrough of GET
# /v1/models/claude-fable-5.
detail_resp = self._make_resp(200, {
"type": "model",
"id": "claude-fable-5",
"display_name": "Claude Fable 5",
"max_input_tokens": 1000000, # context window
"max_tokens": 128000, # max OUTPUT tokens — must be ignored here
})

client_mock = MagicMock()
client_mock.__enter__ = lambda s: client_mock
client_mock.__exit__ = MagicMock(return_value=False)
client_mock.post.return_value = self._make_resp(404, {})
client_mock.get.return_value = detail_resp

with patch("agent.model_metadata.detect_local_server_type", return_value="vllm"), \
patch("httpx.Client", return_value=client_mock):
result = _query_local_context_length("claude-fable-5", "http://localhost:18801/anthropic")

assert result == 1000000, (
f"must read max_input_tokens (context), not max_tokens (output); got {result}"
)

def test_models_list_ignores_max_tokens_output_cap(self):
"""Same guard on the /v1/models LIST branch: max_tokens (output) must
not be mistaken for the context window."""
from agent.model_metadata import _query_local_context_length

detail_miss = self._make_resp(404, {})
list_resp = self._make_resp(200, {"data": [
{"id": "claude-fable-5", "max_input_tokens": 1000000, "max_tokens": 128000},
]})

client_mock = MagicMock()
client_mock.__enter__ = lambda s: client_mock
client_mock.__exit__ = MagicMock(return_value=False)
client_mock.post.return_value = self._make_resp(404, {})
# first GET is /v1/models/{model} (miss), second is /v1/models (list)
client_mock.get.side_effect = [detail_miss, list_resp]

with patch("agent.model_metadata.detect_local_server_type", return_value="vllm"), \
patch("httpx.Client", return_value=client_mock):
result = _query_local_context_length("claude-fable-5", "http://localhost:18801/anthropic")

assert result == 1000000, f"list branch must ignore max_tokens output cap; got {result}"


class TestQueryLocalContextLengthModelsList:
"""_query_local_context_length: falls back to /v1/models list."""
Expand Down
Loading