Skip to content
Closed
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
178 changes: 129 additions & 49 deletions run_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -1085,6 +1085,9 @@ def __init__(
_agent_cfg = _load_agent_config()
except Exception:
_agent_cfg = {}
if not isinstance(_agent_cfg, dict):
_agent_cfg = {}
self._agent_config = _agent_cfg

# Persistent memory (MEMORY.md + USER.md) -- loaded from disk
self._memory_store = None
Expand Down Expand Up @@ -1216,41 +1219,15 @@ def __init__(
compression_target_ratio = float(_compression_cfg.get("target_ratio", 0.20))
compression_protect_last = int(_compression_cfg.get("protect_last_n", 20))

# Read explicit context_length override from model config
_model_cfg = _agent_cfg.get("model", {})
if isinstance(_model_cfg, dict):
_config_context_length = _model_cfg.get("context_length")
else:
_config_context_length = None
if _config_context_length is not None:
try:
_config_context_length = int(_config_context_length)
except (TypeError, ValueError):
_config_context_length = None
_model_cfg = _agent_cfg.get("model", {}) if isinstance(_agent_cfg, dict) else {}
_config_context_length = self._resolve_context_length_override(
model=self.model,
base_url=self.base_url,
config=_agent_cfg,
)

# Store for reuse in switch_model (so config override persists across model switches)
# Store the active runtime override for reuse in switch_model.
self._config_context_length = _config_context_length

# Check custom_providers per-model context_length
if _config_context_length is None:
_custom_providers = _agent_cfg.get("custom_providers")
if isinstance(_custom_providers, list):
for _cp_entry in _custom_providers:
if not isinstance(_cp_entry, dict):
continue
_cp_url = (_cp_entry.get("base_url") or "").rstrip("/")
if _cp_url and _cp_url == self.base_url.rstrip("/"):
_cp_models = _cp_entry.get("models", {})
if isinstance(_cp_models, dict):
_cp_model_cfg = _cp_models.get(self.model, {})
if isinstance(_cp_model_cfg, dict):
_cp_ctx = _cp_model_cfg.get("context_length")
if _cp_ctx is not None:
try:
_config_context_length = int(_cp_ctx)
except (TypeError, ValueError):
pass
break

# Select context engine: config-driven (like memory providers).
# 1. Check config.yaml context.engine setting
Expand Down Expand Up @@ -1417,6 +1394,7 @@ def __init__(
"base_url": self.base_url,
"api_mode": self.api_mode,
"api_key": getattr(self, "api_key", ""),
"config_context_length": self._config_context_length,
"client_kwargs": dict(self._client_kwargs),
"use_prompt_caching": self._use_prompt_caching,
# Context engine state that _try_activate_fallback() overwrites.
Expand Down Expand Up @@ -1548,6 +1526,11 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod
or is_native_anthropic
)

self._config_context_length = self._resolve_context_length_override(
model=self.model,
base_url=self.base_url,
)

# ── Update context compressor ──
if hasattr(self, "context_compressor") and self.context_compressor:
from agent.model_metadata import get_model_context_length
Expand All @@ -1556,7 +1539,7 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod
base_url=self.base_url,
api_key=self.api_key,
provider=self.provider,
config_context_length=getattr(self, "_config_context_length", None),
config_context_length=self._config_context_length,
)
self.context_compressor.update_model(
model=self.model,
Expand All @@ -1578,6 +1561,7 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod
"base_url": self.base_url,
"api_mode": self.api_mode,
"api_key": getattr(self, "api_key", ""),
"config_context_length": self._config_context_length,
"client_kwargs": dict(self._client_kwargs),
"use_prompt_caching": self._use_prompt_caching,
"compressor_model": getattr(_cc, "model", self.model) if _cc else self.model,
Expand Down Expand Up @@ -1708,6 +1692,96 @@ def _current_main_runtime(self) -> Dict[str, str]:
"api_mode": getattr(self, "api_mode", "") or "",
}

@staticmethod
def _coerce_context_length_override(value: Any) -> Optional[int]:
"""Normalize a configured context length override to a positive int."""
if value is None:
return None
try:
value = int(value)
except (TypeError, ValueError):
return None
return value if value > 0 else None

def _resolve_context_length_override(
self,
model: str = "",
base_url: str = "",
config: Optional[Dict[str, Any]] = None,
) -> Optional[int]:
"""Resolve a configured context_length override for a runtime.

Resolution order mirrors the main context lookup entry points:
0. direct context_length on the provided config block
(e.g. auxiliary.compression.context_length)
1. top-level model.context_length for the configured default model
2. custom_providers[].models[model].context_length for matching base_url
"""
from agent.model_metadata import _strip_provider_prefix

cfg = config if isinstance(config, dict) else getattr(self, "_agent_config", {})
if not isinstance(cfg, dict):
return None

direct_override = self._coerce_context_length_override(cfg.get("context_length"))
if direct_override is not None:
return direct_override

target_model = model or getattr(self, "model", "") or ""
target_model_stripped = _strip_provider_prefix(target_model)
target_base_url = (base_url or getattr(self, "base_url", "") or "").rstrip("/")

model_cfg = cfg.get("model", {})
if isinstance(model_cfg, dict):
configured_model = model_cfg.get("default") or model_cfg.get("model") or ""
configured_model_stripped = _strip_provider_prefix(configured_model)
configured_base_url = (model_cfg.get("base_url") or "").rstrip("/")
if (
not target_model
or target_model == configured_model
or (target_model_stripped and target_model_stripped == configured_model_stripped)
) and (
not target_base_url
or not configured_base_url
or target_base_url == configured_base_url
):
override = self._coerce_context_length_override(
model_cfg.get("context_length")
)
if override is not None:
return override

if not target_base_url:
return None

custom_providers = cfg.get("custom_providers")
if not isinstance(custom_providers, list):
return None

for entry in custom_providers:
if not isinstance(entry, dict):
continue
entry_base_url = (entry.get("base_url") or "").rstrip("/")
if not entry_base_url or entry_base_url != target_base_url:
continue
models_cfg = entry.get("models", {})
if not isinstance(models_cfg, dict):
break

model_cfg = models_cfg.get(target_model)
if not isinstance(model_cfg, dict) and target_model_stripped:
model_cfg = models_cfg.get(target_model_stripped)

if isinstance(model_cfg, dict):
override = self._coerce_context_length_override(
model_cfg.get("context_length")
)
if override is not None:
return override
break

return None

def _check_compression_model_feasibility(self) -> None:
"""Warn at session start if the auxiliary compression model's context
window is smaller than the main model's compression threshold.
Expand Down Expand Up @@ -1748,25 +1822,24 @@ def _check_compression_model_feasibility(self) -> None:

aux_base_url = str(getattr(client, "base_url", ""))
aux_api_key = str(getattr(client, "api_key", ""))

# Read user-configured context_length for the compression model.
# Custom endpoints often don't support /models API queries so
# get_model_context_length() falls through to the 128K default,
# ignoring the explicit config value. Pass it as the highest-
# priority hint so the configured value is always respected.
_aux_cfg = (self.config or {}).get("auxiliary", {}).get("compression", {})
_aux_context_config = _aux_cfg.get("context_length") if isinstance(_aux_cfg, dict) else None
if _aux_context_config is not None:
try:
_aux_context_config = int(_aux_context_config)
except (TypeError, ValueError):
_aux_context_config = None

_aux_cfg = (getattr(self, "config", {}) or {}).get("auxiliary", {}).get("compression", {})
if not isinstance(_aux_cfg, dict):
_aux_cfg = {}
aux_context_override = self._resolve_context_length_override(
model=aux_model,
base_url=aux_base_url,
config=_aux_cfg,
)
if aux_context_override is None:
aux_context_override = self._resolve_context_length_override(
model=aux_model,
base_url=aux_base_url,
)
aux_context = get_model_context_length(
aux_model,
base_url=aux_base_url,
api_key=aux_api_key,
config_context_length=_aux_context_config,
config_context_length=aux_context_override,
)

threshold = self.context_compressor.threshold_tokens
Expand Down Expand Up @@ -5531,16 +5604,22 @@ def _try_activate_fallback(self) -> bool:
# causing oversized sessions to overflow the fallback.
if hasattr(self, 'context_compressor') and self.context_compressor:
from agent.model_metadata import get_model_context_length
self._config_context_length = self._resolve_context_length_override(
model=self.model,
base_url=self.base_url,
)
fb_context_length = get_model_context_length(
self.model, base_url=self.base_url,
api_key=self.api_key, provider=self.provider,
config_context_length=self._config_context_length,
)
self.context_compressor.update_model(
model=self.model,
context_length=fb_context_length,
base_url=self.base_url,
api_key=getattr(self, "api_key", ""),
provider=self.provider,
api_mode=self.api_mode,
)

self._emit_status(
Expand Down Expand Up @@ -5580,6 +5659,7 @@ def _restore_primary_runtime(self) -> bool:
self.base_url = rt["base_url"] # setter updates _base_url_lower
self.api_mode = rt["api_mode"]
self.api_key = rt["api_key"]
self._config_context_length = rt.get("config_context_length")
self._client_kwargs = dict(rt["client_kwargs"])
self._use_prompt_caching = rt["use_prompt_caching"]

Expand Down
Loading