fix(gateway): preserve native compaction capability on resume
This commit is contained in:
@@ -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 [])
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user