fix(gateway): preserve native compaction capability on resume

This commit is contained in:
Stephen Chin
2026-08-28 09:36:09 -07:00
committed by Teknium
parent 48a4201f40
commit c9b9b5e6c7
10 changed files with 115 additions and 3 deletions
+5
View File
@@ -613,6 +613,7 @@ def init_agent(
checkpoint_max_file_size_mb: int = 10,
pass_session_id: bool = False,
requested_provider: str = None,
capabilities: Optional[Dict[str, bool]] = None,
):
"""
Initialize the AI Agent.
@@ -712,6 +713,10 @@ def init_agent(
if isinstance(requested_provider, str) and requested_provider.strip()
else agent.provider
)
agent.capabilities = {
key: value for key, value in (capabilities or {}).items()
if isinstance(key, str) and isinstance(value, bool)
}
agent._credential_pool = credential_pool
agent.acp_command = acp_command or command
agent.acp_args = list(acp_args or args or [])
+4 -1
View File
@@ -203,7 +203,10 @@ def native_compaction_context_management(
return None
if not is_native_compaction_model(getattr(agent, "model", None)):
return None
if not is_direct_openai_route(
trusted_proxy = bool(
getattr(agent, "capabilities", {}).get("openai_native_compaction", False)
)
if not trusted_proxy and not is_direct_openai_route(
getattr(agent, "base_url", None), is_codex_backend=is_codex_backend
):
return None
+19 -1
View File
@@ -3116,6 +3116,8 @@ def _resolve_runtime_agent_kwargs_for_provider(provider: str) -> dict:
"args": list(runtime.get("args") or []),
"credential_pool": runtime.get("credential_pool"),
"request_overrides": dict(runtime.get("request_overrides") or {}),
"capabilities": dict(runtime.get("capabilities") or {}),
"max_tokens": runtime.get("max_output_tokens"),
}
@@ -8343,12 +8345,14 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
override_model = override.get("model", model)
override_runtime = {
"provider": override.get("provider"),
"requested_provider": override.get("requested_provider"),
"api_key": override.get("api_key"),
"base_url": override.get("base_url"),
"api_mode": override.get("api_mode"),
"max_tokens": override.get("max_tokens"),
"credential_pool": override.get("credential_pool"),
"request_overrides": override.get("request_overrides"),
"capabilities": dict(override.get("capabilities") or {}),
}
if override_runtime.get("api_key"):
if override_runtime.get("credential_pool") is None:
@@ -8503,6 +8507,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
"args": list(runtime_kwargs.get("args") or []),
"credential_pool": runtime_kwargs.get("credential_pool"),
"max_tokens": runtime_kwargs.get("max_tokens"),
"capabilities": dict(runtime_kwargs.get("capabilities") or {}),
}
base_request_overrides = dict(runtime_kwargs.get("request_overrides") or {})
route = {
@@ -27428,6 +27433,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
runtime.get("provider", ""),
runtime.get("requested_provider", ""),
runtime.get("api_mode", ""),
sorted((runtime.get("capabilities") or {}).items()),
sorted(enabled_toolsets) if enabled_toolsets else [],
# reasoning_config excluded — it's set per-message on the
# cached agent and doesn't affect system prompt or tools.
@@ -27496,6 +27502,9 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
override["request_overrides"] = dict(
runtime.get("request_overrides") or {}
)
override["requested_provider"] = runtime.get("requested_provider")
override["capabilities"] = dict(runtime.get("capabilities") or {})
override["max_tokens"] = runtime.get("max_tokens")
if not override.get("base_url"):
override["base_url"] = runtime.get("base_url")
except Exception:
@@ -27526,7 +27535,16 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
if not override:
return model, runtime_kwargs
model = override.get("model", model)
for key in ("provider", "api_key", "base_url", "api_mode", "credential_pool"):
for key in (
"provider",
"requested_provider",
"api_key",
"base_url",
"api_mode",
"credential_pool",
"capabilities",
"max_tokens",
):
val = override.get(key)
if val is not None:
runtime_kwargs[key] = val
+11 -1
View File
@@ -1480,7 +1480,7 @@ def _normalize_custom_provider_entry(
"models_discovered",
"context_length", "rate_limit_delay",
"request_timeout_seconds", "stale_timeout_seconds",
"discover_models", "extra_body", "extra_headers",
"discover_models", "extra_body", "extra_headers", "capabilities",
"ssl_ca_cert", "ssl_verify",
}
for camel, snake in _CAMEL_ALIASES.items():
@@ -1605,6 +1605,16 @@ def _normalize_custom_provider_entry(
if models_discovered:
normalized["models_discovered"] = True
capabilities = entry.get("capabilities")
if isinstance(capabilities, dict):
normalized_capabilities = {
key: value
for key, value in capabilities.items()
if isinstance(key, str) and isinstance(value, bool)
}
if normalized_capabilities:
normalized["capabilities"] = normalized_capabilities
context_length = entry.get("context_length")
if isinstance(context_length, int) and context_length > 0:
normalized["context_length"] = context_length
+32
View File
@@ -707,6 +707,29 @@ def _try_resolve_from_custom_pool(
return None
def _lift_model_capabilities(
entry: Dict[str, Any], model: Optional[str], result: Dict[str, Any]
) -> None:
"""Copy explicit boolean per-model capabilities into the runtime."""
capabilities = {
key: value
for key, value in (entry.get("capabilities") or {}).items()
if isinstance(key, str) and isinstance(value, bool)
}
models = entry.get("models")
model_config = models.get(model) if isinstance(models, dict) and model else None
if isinstance(model_config, dict):
capabilities.update(
{
key: value
for key, value in model_config.items()
if isinstance(key, str) and isinstance(value, bool)
}
)
if capabilities:
result["capabilities"] = capabilities
def _lift_max_output_tokens(entry: Dict[str, Any], result: Dict[str, Any]) -> None:
"""Propagate a per-provider output cap onto the resolved runtime dict.
@@ -832,6 +855,8 @@ def _get_named_custom_provider(requested_provider: str) -> Optional[Dict[str, An
if api_mode:
result["api_mode"] = api_mode
_lift_max_output_tokens(entry, result)
if isinstance(entry.get("capabilities"), dict):
result["capabilities"] = dict(entry["capabilities"])
return result
# Fall back to custom_providers: list (legacy format)
@@ -879,6 +904,8 @@ def _get_named_custom_provider(requested_provider: str) -> Optional[Dict[str, An
if model_name:
result["model"] = model_name
_lift_max_output_tokens(entry, result)
if isinstance(entry.get("capabilities"), dict):
result["capabilities"] = dict(entry["capabilities"])
return result
return None
@@ -1239,6 +1266,7 @@ def _resolve_named_custom_runtime(
model_name = target_model or custom_provider.get("model")
if model_name:
pool_result["model"] = model_name
_lift_model_capabilities(custom_provider, model_name, pool_result)
if isinstance(custom_provider.get("max_output_tokens"), int):
pool_result["max_output_tokens"] = custom_provider["max_output_tokens"]
request_overrides = _custom_provider_request_overrides(custom_provider)
@@ -1294,6 +1322,7 @@ def _resolve_named_custom_runtime(
"base_url": base_url,
"api_key": api_key or "no-key-required",
"source": f"custom_provider:{custom_provider.get('name', requested_provider)}",
"requested_provider": requested_provider,
}
# Propagate the model name so callers can override self.model when the
# provider name differs from the actual model string the API expects.
@@ -1305,6 +1334,9 @@ def _resolve_named_custom_runtime(
result["model"] = target_model
elif custom_provider.get("model"):
result["model"] = custom_provider["model"]
_lift_model_capabilities(
custom_provider, result.get("model"), result
)
if isinstance(custom_provider.get("max_output_tokens"), int):
result["max_output_tokens"] = custom_provider["max_output_tokens"]
# Per-provider extra HTTP headers (proxies, gateways, custom auth).
+2
View File
@@ -523,6 +523,7 @@ class AIAgent:
checkpoint_max_file_size_mb: int = 10,
pass_session_id: bool = False,
requested_provider: str = None,
capabilities: Dict[str, bool] | None = None,
):
"""Forwarder — see ``agent.agent_init.init_agent``."""
if tool_delay is not None:
@@ -539,6 +540,7 @@ class AIAgent:
api_key=api_key,
provider=provider,
requested_provider=requested_provider,
capabilities=capabilities,
api_mode=api_mode,
acp_command=acp_command,
acp_args=acp_args,
+10
View File
@@ -69,6 +69,16 @@ class TestAgentConfigSignature:
sig2 = GatewayRunner._agent_config_signature("claude-sonnet-4", rt2, ["hermes-telegram"], "")
assert sig1 != sig2
def test_capability_change_different_signature(self):
from gateway.run import GatewayRunner
runtime = {"api_key": "sk-test12345678", "base_url": "https://proxy.example/v1", "provider": "custom"}
native = {**runtime, "capabilities": {"openai_native_compaction": True}}
plain = {**runtime, "capabilities": {"openai_native_compaction": False}}
assert GatewayRunner._agent_config_signature("gpt-5.6", native, [], "") != (
GatewayRunner._agent_config_signature("gpt-5.6", plain, [], "")
)
# ---------------------------------------------------------------
# cache_keys (compression/context config cache-busting)
@@ -111,6 +111,9 @@ def test_runner_rehydrates_override_after_restart(store_factory):
"api_mode": "responses",
"base_url": "https://api.openai.example/v1",
"provider": "openai",
"requested_provider": "custom:chatgpt-tier",
"capabilities": {"openai_native_compaction": True},
"max_tokens": 32_768,
},
):
runner._rehydrate_session_model_override(session_key)
@@ -122,6 +125,20 @@ def test_runner_rehydrates_override_after_restart(store_factory):
# Credentials come from live resolution, never from disk.
assert override["api_key"] == "sk-fresh-from-keychain"
assert override["api_mode"] == "responses"
assert override["requested_provider"] == "custom:chatgpt-tier"
assert override["capabilities"] == {"openai_native_compaction": True}
assert override["max_tokens"] == 32_768
model, runtime = runner._resolve_session_agent_runtime(
session_key=session_key,
user_config={"model": {"default": "global-model"}},
)
assert model == "gpt-5o"
assert runtime["requested_provider"] == "custom:chatgpt-tier"
assert runtime["capabilities"] == {"openai_native_compaction": True}
assert runtime["max_tokens"] == 32_768
route = runner._resolve_turn_agent_config("", model, runtime)
assert route["runtime"]["capabilities"] == {"openai_native_compaction": True}
def test_sanitize_model_override():
@@ -632,6 +632,8 @@ def test_named_custom_provider_uses_saved_credentials(monkeypatch):
"name": "Local",
"base_url": "http://1.2.3.4:1234/v1",
"api_key": "local-provider-key",
"model": "gpt-5.6",
"capabilities": {"openai_native_compaction": True},
}
]
},
@@ -653,6 +655,7 @@ def test_named_custom_provider_uses_saved_credentials(monkeypatch):
assert resolved["base_url"] == "http://1.2.3.4:1234/v1"
assert resolved["api_key"] == "local-provider-key"
assert resolved["requested_provider"] == "local"
assert resolved["capabilities"] == {"openai_native_compaction": True}
assert resolved["source"] == "custom_provider:Local"
+12
View File
@@ -27,6 +27,7 @@ def _agent(
compression_enabled=True,
threshold: object = DEFAULT_COMPACT_THRESHOLD,
compressor=None,
capabilities=None,
):
return SimpleNamespace(
model=model,
@@ -35,6 +36,7 @@ def _agent(
compression_enabled=compression_enabled,
codex_responses_compact_threshold=threshold,
context_compressor=compressor,
capabilities=capabilities or {},
)
@@ -90,6 +92,16 @@ class TestRequestGate:
)
assert payload is not None
def test_trusted_proxy_capability_gets_payload(self):
payload = native_compaction_context_management(
_agent(
base_url="https://trusted-proxy.example/v1",
capabilities={"openai_native_compaction": True},
),
is_codex_backend=False,
)
assert payload is not None
def test_disabled_by_default_config_value(self):
assert (
native_compaction_context_management(