chore: baseline WIP before unified model configuration implementation

Pre-existing uncommitted work (runtime snapshots, message budget middleware) preserved as baseline.
This commit is contained in:
m4
2026-07-20 20:15:38 +08:00
parent 38668c4ce5
commit 8a0ab17936
13 changed files with 1590 additions and 28 deletions
+12 -2
View File
@@ -683,6 +683,7 @@ def _get_default_middleware(
workspace_dir: str | Path | None = None,
cfg=None,
chat_model=None,
backend=None,
memory_source_agent: str = "EvoScientist",
):
"""Build the default middleware list.
@@ -713,6 +714,7 @@ def _get_default_middleware(
create_context_editing_middleware,
create_memory_lifecycle_middleware,
create_memory_middleware,
create_message_budget_middleware,
create_runtime_context_middleware,
create_scheduler_middleware,
create_tool_selector_middleware,
@@ -724,6 +726,13 @@ def _get_default_middleware(
if cfg.model_fallbacks:
load_fallback_chain(cfg.model_fallbacks)
model = chat_model if chat_model is not None else _ensure_chat_model()
if backend is None:
# Preserve the factory's pure path for callers that provide an
# explicit model/configuration (notably tests and subagent assembly).
# Production graph factories always pass their real composite backend.
from deepagents.backends import StateBackend
backend = StateBackend()
memory_dir = str(_paths_mod.MEMORIES_DIR)
source_type = (
MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN
@@ -770,6 +779,7 @@ def _get_default_middleware(
tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider)
mw = [
ConfigurableModelMiddleware(),
create_message_budget_middleware(model, backend),
create_context_editing_middleware(model),
ModelFallbackMiddleware(),
ContextOverflowMapperMiddleware(),
@@ -857,7 +867,7 @@ def _get_default_agent():
cfg = _ensure_config()
be = _get_default_backend()
mw = _get_default_middleware()
mw = _get_default_middleware(backend=be)
# HITL on main agent only (mirrors create_cli_agent). Use middleware,
# not interrupt_on= kwarg — the kwarg propagates to every subagent and
@@ -1022,7 +1032,7 @@ def create_cli_agent(
# CLI agent never drifts from the default chain. Anything CLI-specific
# (e.g. ``HumanInTheLoopMiddleware``) is appended below.
mw: list[AgentMiddleware] = _get_default_middleware(
workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model, backend=be
)
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
+240 -4
View File
@@ -22,7 +22,7 @@ import yaml
from .settings import get_config_dir
PROVIDER_PROFILES_VERSION = 2
PROVIDER_PROFILES_VERSION = 3
SUPPORTED_PROVIDER_ADAPTERS = (
"openai",
"anthropic",
@@ -79,12 +79,41 @@ class ProviderProfileError(ValueError):
"""Raised when a provider profile document or selection is invalid."""
@dataclass(frozen=True)
class ProviderRuntime:
timeout_seconds: int = 120
max_retries: int = 2
default_temperature: float | None = None
default_top_p: float | None = None
default_reasoning_effort: str = "auto"
@dataclass(frozen=True)
class ModelRuntime:
limit_mode: str | None = None
context_window_tokens: int | None = None
max_input_tokens: int | None = None
max_output_tokens: int = 4096
min_effective_input_tokens: int = 4096
limits_status: str = "needs_confirmation"
limits_source: str = "user"
temperature: float | None = None
top_p: float | None = None
reasoning_effort: str = "auto"
capabilities: tuple[tuple[str, bool | str], ...] = (
("tools", "auto"),
("vision", "auto"),
("structured_output", "auto"),
)
@dataclass(frozen=True)
class ProviderModel:
id: str
name: str
model_id: str
enabled: bool = True
runtime: ModelRuntime = ModelRuntime()
@dataclass(frozen=True)
@@ -97,6 +126,7 @@ class ProviderProfile:
enabled: bool
models: tuple[ProviderModel, ...]
auth_mode: str = "api_key"
runtime: ProviderRuntime = ProviderRuntime()
@dataclass(frozen=True)
@@ -151,6 +181,150 @@ def _validate_id(value: str, *, context: str) -> str:
return value
def _optional_number(
raw: dict[str, Any], key: str, *, context: str, minimum: float, maximum: float
) -> float | None:
value = raw.get(key)
if value is None:
return None
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ProviderProfileError(f"{context}.{key} must be a number or null.")
numeric = float(value)
if not minimum <= numeric <= maximum:
raise ProviderProfileError(
f"{context}.{key} must be between {minimum:g} and {maximum:g}."
)
return numeric
def _optional_positive_int(raw: dict[str, Any], key: str, *, context: str) -> int | None:
value = raw.get(key)
if value is None:
return None
if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
raise ProviderProfileError(f"{context}.{key} must be a positive integer or null.")
return value
def _parse_provider_runtime(raw: Any, *, context: str) -> ProviderRuntime:
if raw is None:
return ProviderRuntime()
if not isinstance(raw, dict):
raise ProviderProfileError(f"{context}.runtime must be an object.")
timeout_seconds = raw.get("timeout_seconds", 120)
max_retries = raw.get("max_retries", 2)
if isinstance(timeout_seconds, bool) or not isinstance(timeout_seconds, int):
raise ProviderProfileError(f"{context}.runtime.timeout_seconds must be an integer.")
if not 10 <= timeout_seconds <= 600:
raise ProviderProfileError(
f"{context}.runtime.timeout_seconds must be between 10 and 600."
)
if isinstance(max_retries, bool) or not isinstance(max_retries, int):
raise ProviderProfileError(f"{context}.runtime.max_retries must be an integer.")
if not 0 <= max_retries <= 5:
raise ProviderProfileError(
f"{context}.runtime.max_retries must be between 0 and 5."
)
effort = raw.get("default_reasoning_effort", "auto")
if effort not in {"auto", "low", "medium", "high"}:
raise ProviderProfileError(
f"{context}.runtime.default_reasoning_effort is invalid."
)
return ProviderRuntime(
timeout_seconds=timeout_seconds,
max_retries=max_retries,
default_temperature=_optional_number(
raw, "default_temperature", context=f"{context}.runtime", minimum=0, maximum=2
),
default_top_p=_optional_number(
raw, "default_top_p", context=f"{context}.runtime", minimum=0.000001, maximum=1
),
default_reasoning_effort=effort,
)
def resolve_model_input_limit(runtime: ModelRuntime) -> int:
"""Return the prompt limit declared by a confirmed model runtime."""
if runtime.limits_status != "confirmed" or runtime.limit_mode is None:
raise ProviderProfileError("Model limits require confirmation before use.")
if runtime.limit_mode == "combined":
if runtime.context_window_tokens is None:
raise ProviderProfileError("context_window_tokens is required for combined limits.")
return runtime.context_window_tokens - runtime.max_output_tokens
if runtime.limit_mode == "input_only":
if runtime.max_input_tokens is None:
raise ProviderProfileError("max_input_tokens is required for input_only limits.")
return runtime.max_input_tokens
raise ProviderProfileError("limit_mode must be combined or input_only.")
def _parse_model_runtime(raw: Any, *, context: str) -> ModelRuntime:
if raw is None:
return ModelRuntime()
if not isinstance(raw, dict):
raise ProviderProfileError(f"{context}.runtime must be an object.")
status = raw.get("limits_status", "needs_confirmation")
if status not in {"confirmed", "needs_confirmation"}:
raise ProviderProfileError(f"{context}.runtime.limits_status is invalid.")
source = raw.get("limits_source", "user")
if source not in {"catalog", "provider", "user"}:
raise ProviderProfileError(f"{context}.runtime.limits_source is invalid.")
limit_mode = raw.get("limit_mode")
if limit_mode is not None and limit_mode not in {"combined", "input_only"}:
raise ProviderProfileError(f"{context}.runtime.limit_mode is invalid.")
max_output = raw.get("max_output_tokens", 4096)
min_effective = raw.get("min_effective_input_tokens", 4096)
if isinstance(max_output, bool) or not isinstance(max_output, int) or max_output <= 0:
raise ProviderProfileError(f"{context}.runtime.max_output_tokens must be positive.")
if isinstance(min_effective, bool) or not isinstance(min_effective, int) or min_effective < 1024:
raise ProviderProfileError(
f"{context}.runtime.min_effective_input_tokens must be at least 1024."
)
effort = raw.get("reasoning_effort", "auto")
if effort not in {"auto", "low", "medium", "high"}:
raise ProviderProfileError(f"{context}.runtime.reasoning_effort is invalid.")
capabilities_raw = raw.get("capabilities", {})
if not isinstance(capabilities_raw, dict):
raise ProviderProfileError(f"{context}.runtime.capabilities must be an object.")
capabilities: list[tuple[str, bool | str]] = []
for key in ("tools", "vision", "structured_output"):
value = capabilities_raw.get(key, "auto")
if value not in {True, False, "auto"}:
raise ProviderProfileError(
f"{context}.runtime.capabilities.{key} must be true, false, or auto."
)
capabilities.append((key, value))
runtime = ModelRuntime(
limit_mode=limit_mode,
context_window_tokens=_optional_positive_int(
raw, "context_window_tokens", context=f"{context}.runtime"
),
max_input_tokens=_optional_positive_int(
raw, "max_input_tokens", context=f"{context}.runtime"
),
max_output_tokens=max_output,
min_effective_input_tokens=min_effective,
limits_status=status,
limits_source=source,
temperature=_optional_number(
raw, "temperature", context=f"{context}.runtime", minimum=0, maximum=2
),
top_p=_optional_number(
raw, "top_p", context=f"{context}.runtime", minimum=0.000001, maximum=1
),
reasoning_effort=effort,
capabilities=tuple(capabilities),
)
if status == "confirmed":
limit = resolve_model_input_limit(runtime)
safety = max(2048, int(limit * 0.10 + 0.999999))
if limit - safety < min_effective:
raise ProviderProfileError(
f"{context}.runtime leaves less than min_effective_input_tokens after safety reserve."
)
return runtime
def _parse_model(raw: Any, *, provider_id: str, index: int) -> ProviderModel:
context = f"providers[{provider_id}].models[{index}]"
if not isinstance(raw, dict):
@@ -164,6 +338,7 @@ def _parse_model(raw: Any, *, provider_id: str, index: int) -> ProviderModel:
name=_required_string(raw, "name", context=context, max_length=120),
model_id=_required_string(raw, "model_id", context=context, max_length=300),
enabled=bool(raw.get("enabled", True)),
runtime=_parse_model_runtime(raw.get("runtime"), context=context),
)
@@ -233,6 +408,20 @@ def _parse_profile(raw: Any, *, index: int, builtin: bool = False) -> ProviderPr
raise ProviderProfileError(
f"Provider {provider_id!r} contains duplicate model IDs."
)
portable_reasoning_adapters = {
"openai",
"openai-compatible",
"grok",
"antigravity",
"openrouter",
}
if adapter not in portable_reasoning_adapters:
for model in models:
if model.runtime.reasoning_effort != "auto":
raise ProviderProfileError(
f"{context}.models[{model.id}].runtime.reasoning_effort is "
f"not supported by adapter {adapter!r}."
)
return ProviderProfile(
id=provider_id,
@@ -243,6 +432,7 @@ def _parse_profile(raw: Any, *, index: int, builtin: bool = False) -> ProviderPr
enabled=bool(raw.get("enabled", True)),
models=models,
auth_mode=auth_mode,
runtime=_parse_provider_runtime(raw.get("runtime"), context=context),
)
@@ -251,9 +441,11 @@ def _parse_document(raw: Any) -> ProviderProfiles:
return ProviderProfiles()
if not isinstance(raw, dict):
raise ProviderProfileError("Provider profile document must be an object.")
version = raw.get("version", 1)
if version not in {1, PROVIDER_PROFILES_VERSION}:
raise ProviderProfileError(f"Unsupported provider profile version {version!r}.")
version = raw.get("version")
if version != PROVIDER_PROFILES_VERSION:
raise ProviderProfileError(
"PROVIDER_PROFILE_RESET_REQUIRED: only version 3 provider profiles are supported."
)
builtins_raw = raw.get("builtins", [])
if not isinstance(builtins_raw, list):
raise ProviderProfileError("builtins must be a list.")
@@ -290,12 +482,32 @@ def _profile_to_dict(profile: ProviderProfile) -> dict[str, Any]:
"api_key": profile.api_key,
"auth_mode": profile.auth_mode,
"enabled": profile.enabled,
"runtime": {
"timeout_seconds": profile.runtime.timeout_seconds,
"max_retries": profile.runtime.max_retries,
"default_temperature": profile.runtime.default_temperature,
"default_top_p": profile.runtime.default_top_p,
"default_reasoning_effort": profile.runtime.default_reasoning_effort,
},
"models": [
{
"id": model.id,
"name": model.name,
"model_id": model.model_id,
"enabled": model.enabled,
"runtime": {
"limit_mode": model.runtime.limit_mode,
"context_window_tokens": model.runtime.context_window_tokens,
"max_input_tokens": model.runtime.max_input_tokens,
"max_output_tokens": model.runtime.max_output_tokens,
"min_effective_input_tokens": model.runtime.min_effective_input_tokens,
"limits_status": model.runtime.limits_status,
"limits_source": model.runtime.limits_source,
"temperature": model.runtime.temperature,
"top_p": model.runtime.top_p,
"reasoning_effort": model.runtime.reasoning_effort,
"capabilities": dict(model.runtime.capabilities),
},
}
for model in profile.models
],
@@ -471,12 +683,32 @@ def _provider_profiles_public_payload(document: ProviderProfiles) -> dict[str, A
"enabled": profile.enabled,
"api_key_configured": bool(profile.api_key),
"api_key_hint": _api_key_hint(profile.api_key),
"runtime": {
"timeout_seconds": profile.runtime.timeout_seconds,
"max_retries": profile.runtime.max_retries,
"default_temperature": profile.runtime.default_temperature,
"default_top_p": profile.runtime.default_top_p,
"default_reasoning_effort": profile.runtime.default_reasoning_effort,
},
"models": [
{
"id": model.id,
"name": model.name,
"model_id": model.model_id,
"enabled": model.enabled,
"runtime": {
"limit_mode": model.runtime.limit_mode,
"context_window_tokens": model.runtime.context_window_tokens,
"max_input_tokens": model.runtime.max_input_tokens,
"max_output_tokens": model.runtime.max_output_tokens,
"min_effective_input_tokens": model.runtime.min_effective_input_tokens,
"limits_status": model.runtime.limits_status,
"limits_source": model.runtime.limits_source,
"temperature": model.runtime.temperature,
"top_p": model.runtime.top_p,
"reasoning_effort": model.runtime.reasoning_effort,
"capabilities": dict(model.runtime.capabilities),
},
}
for model in profile.models
],
@@ -565,6 +797,10 @@ def resolve_provider_model(
raise ProviderProfileError(
f"Model {model_name!r} in provider {provider_id!r} is disabled."
)
if model.runtime.limits_status != "confirmed":
raise ProviderProfileError(
f"RUNTIME_LIMITS_CONFIRMATION_REQUIRED: model {model_name!r} needs confirmed limits."
)
return profile, model
raise ProviderProfileError(
f"Model {model_name!r} is not configured for provider {provider_id!r}."
+45
View File
@@ -76,6 +76,7 @@ from EvoScientist.llm.provider_operations import (
discover_provider_models,
test_provider_model,
)
from EvoScientist.llm.runtime_snapshots import create_run_runtime_snapshot
from EvoScientist.sessions import (
MAIN_THREAD_FILTER_PARAMS,
MAIN_THREAD_FILTER_SQL,
@@ -655,6 +656,45 @@ async def provider_profiles_endpoint(request: Request) -> JSONResponse:
)
async def runtime_snapshot_endpoint(request: Request) -> JSONResponse:
"""Create an opaque server-side snapshot for one forthcoming graph run."""
auth_error = _provider_admin_error(request)
if auth_error is not None:
return auth_error
try:
payload = await request.json()
if not isinstance(payload, dict):
raise ProviderProfileError("Request body must be an object.")
snapshot_id = payload.get("snapshot_id")
model = payload.get("model")
provider = payload.get("provider")
if not isinstance(snapshot_id, str):
raise ProviderProfileError("snapshot_id is required.")
if not isinstance(model, str) or not isinstance(provider, str):
raise ProviderProfileError("model and provider are required.")
snapshot = await asyncio.to_thread(
create_run_runtime_snapshot,
snapshot_id,
model=model,
provider=provider,
)
return JSONResponse(
{"snapshot": snapshot.public_payload() if snapshot is not None else None}
)
except (json.JSONDecodeError, UnicodeDecodeError):
return JSONResponse(
{"error": "Request body must be valid JSON."}, status_code=400
)
except ProviderProfileError as exc:
return JSONResponse({"error": str(exc)}, status_code=400)
except Exception:
_logger.exception("Runtime snapshot request failed")
return JSONResponse(
{"error": "Runtime snapshot request failed."}, status_code=500
)
async def llm_config_endpoint(request: Request) -> JSONResponse:
"""Read or patch the LLM-related fields persisted in config.yaml."""
auth_error = _provider_admin_error(request)
@@ -1636,6 +1676,11 @@ app = Starlette(
provider_profiles_endpoint,
methods=["GET", "PUT"],
),
Route(
"/api/runtime-snapshots",
runtime_snapshot_endpoint,
methods=["POST"],
),
Route(
"/api/config",
llm_config_endpoint,
+167 -11
View File
@@ -15,6 +15,7 @@ import json
import os
import re
import warnings
from dataclasses import dataclass
from typing import Any
from langchain.chat_models import init_chat_model
@@ -29,6 +30,7 @@ from ..config.provider_profiles import (
get_provider_profiles_path,
list_configured_builtin_model_entries,
list_configured_model_entries,
resolve_model_input_limit,
resolve_provider_model,
)
from .context_window import apply_known_context_window
@@ -53,6 +55,20 @@ _DEEPSEEK_BASE_URL = "https://api.deepseek.com"
_MOONSHOT_BASE_URL = "https://api.moonshot.cn/v1"
_KIMI_CODING_BASE_URL = "https://api.kimi.com/coding/"
# ``glm`` is a common shorthand in existing configuration files. The
# supported provider identifier is ``zhipu`` because it describes the API
# vendor rather than a single model family. Normalize it at the boundary so
# legacy configuration reaches the same authenticated OpenAI-compatible route.
_PROVIDER_ALIASES: dict[str, str] = {
"glm": "zhipu",
}
def normalize_provider_id(provider: str) -> str:
"""Return the canonical provider ID accepted by the model registry."""
normalized = provider.strip().lower()
return _PROVIDER_ALIASES.get(normalized, normalized)
# Providers routed through the OpenAI provider with a custom base_url.
# Maps provider name → (base_url or None, env var for API key).
_OPENAI_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = {
@@ -84,6 +100,74 @@ _THINKING_CAPABLE_PROVIDERS: set[str] = {"minimax"}
_TRUTHY_ENV_VALUES = {"1", "true", "yes", "on"}
_FALSEY_ENV_VALUES = {"0", "false", "no", "off"}
@dataclass(frozen=True)
class ResolvedRuntimeOptions:
"""Provider-neutral parameters resolved once for a configured model."""
profile_revision: str | None
model_id: str
adapter_id: str
limit_mode: str
resolved_input_limit: int
max_output_tokens: int
min_effective_input_tokens: int
timeout_seconds: int
max_retries: int
temperature: float | None
top_p: float | None
reasoning_effort: str
limits_source: str
capabilities: tuple[tuple[str, bool | str], ...]
def resolve_runtime_options(
profile: ProviderProfile,
model: ProviderModel,
call_overrides: dict[str, Any] | None = None,
) -> ResolvedRuntimeOptions:
"""Resolve the bounded runtime options for one configured model.
Only sampling values and a lower output budget may be overridden by an
internal caller. Connection and capacity fields always come from the
persisted profile.
"""
overrides = call_overrides or {}
runtime = model.runtime
input_limit = resolve_model_input_limit(runtime)
output_limit = runtime.max_output_tokens
override_output = overrides.get("max_tokens")
if isinstance(override_output, int) and not isinstance(override_output, bool):
if override_output <= 0 or override_output > output_limit:
raise ValueError("max_tokens override must be positive and no greater than the model limit.")
output_limit = override_output
temperature = overrides.get("temperature", runtime.temperature)
if temperature is None:
temperature = profile.runtime.default_temperature
top_p = overrides.get("top_p", runtime.top_p)
if top_p is None:
top_p = profile.runtime.default_top_p
return ResolvedRuntimeOptions(
profile_revision=get_provider_profile_revision(profile.id),
model_id=model.model_id,
adapter_id=profile.adapter,
limit_mode=runtime.limit_mode or "input_only",
resolved_input_limit=input_limit,
max_output_tokens=output_limit,
min_effective_input_tokens=runtime.min_effective_input_tokens,
timeout_seconds=profile.runtime.timeout_seconds,
max_retries=profile.runtime.max_retries,
temperature=temperature,
top_p=top_p,
reasoning_effort=(
runtime.reasoning_effort
if runtime.reasoning_effort != "auto"
else profile.runtime.default_reasoning_effort
),
limits_source=runtime.limits_source,
capabilities=runtime.capabilities,
)
# Model registry: list of (short_name, model_id, provider)
# Allows same short_name across different providers.
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
@@ -261,6 +345,11 @@ _STATIC_PROVIDER_IDS = {provider for _, _, provider in _MODEL_ENTRIES} | {"ollam
_MODEL_CATALOG_ID_PATTERN = re.compile(r"^[a-z0-9][a-z0-9._-]{0,63}$")
def is_static_provider_id(provider: str) -> bool:
"""Return whether a provider is configured outside custom profiles."""
return normalize_provider_id(provider) in _STATIC_PROVIDER_IDS
def normalize_builtin_model_catalog(raw: Any) -> list[dict[str, Any]] | None:
"""Validate and normalize the optional config.yaml built-in model catalog."""
if raw is None:
@@ -341,7 +430,13 @@ def list_builtin_model_catalog_entries(
def _resolve_builtin_catalog_model(provider: str, model: str) -> str | None:
from ..config.settings import load_config
profile = get_builtin_provider_profile(provider)
try:
profile = get_builtin_provider_profile(provider)
except ProviderProfileError:
# Custom profiles use an explicit v3 reset path during development.
# Built-in model selection must still work from config.yaml while that
# reset has not happened, because it never consumes the old profile.
profile = None
if profile is not None:
if not profile.enabled:
return None
@@ -366,7 +461,10 @@ def get_model_runtime_revision(provider: str | None) -> str | None:
if not provider:
return None
if provider in _STATIC_PROVIDER_IDS:
managed_revision = get_provider_profile_revision(provider)
try:
managed_revision = get_provider_profile_revision(provider)
except ProviderProfileError:
managed_revision = None
if managed_revision is not None:
return managed_revision
from ..config.settings import load_config
@@ -398,7 +496,10 @@ def get_models_for_provider(provider: str) -> list[tuple[str, str]]:
List of (short_name, model_id) tuples for the provider.
"""
if provider in _STATIC_PROVIDER_IDS:
managed = get_builtin_provider_profile(provider)
try:
managed = get_builtin_provider_profile(provider)
except ProviderProfileError:
managed = None
if managed is not None:
if not managed.enabled:
return []
@@ -553,9 +654,12 @@ def get_chat_model(
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
"""
model = model or DEFAULT_MODEL
if provider is not None:
provider = normalize_provider_id(provider)
dynamic_profile = profile_override
dynamic_model = model_override
resolved_runtime: ResolvedRuntimeOptions | None = None
if dynamic_profile is not None:
provider = dynamic_profile.id
if dynamic_model is None:
@@ -575,10 +679,27 @@ def get_chat_model(
try:
dynamic_resolution = resolve_provider_model(provider, model)
except ProviderProfileError as exc:
raise ValueError(str(exc)) from exc
# An opaque provider override remains a valid LangChain escape
# hatch. A stale v1/v2 registry cannot be allowed to block it;
# configured custom profiles are rejected earlier by the snapshot
# endpoint and never silently selected from that document.
if str(exc).startswith("PROVIDER_PROFILE_RESET_REQUIRED"):
dynamic_resolution = None
else:
raise ValueError(str(exc)) from exc
if dynamic_resolution is not None:
dynamic_profile, dynamic_model = dynamic_resolution
if dynamic_profile is not None and dynamic_model is not None:
resolved_runtime = resolve_runtime_options(dynamic_profile, dynamic_model, kwargs)
kwargs.setdefault("timeout", resolved_runtime.timeout_seconds)
kwargs.setdefault("max_retries", resolved_runtime.max_retries)
kwargs.setdefault("max_tokens", resolved_runtime.max_output_tokens)
if resolved_runtime.temperature is not None:
kwargs.setdefault("temperature", resolved_runtime.temperature)
if resolved_runtime.top_p is not None:
kwargs.setdefault("top_p", resolved_runtime.top_p)
# Look up short name in the configured catalog, then the legacy registry.
model_id = dynamic_model.model_id if dynamic_model is not None else None
if model_id is None and provider in _STATIC_PROVIDER_IDS:
@@ -616,11 +737,13 @@ def get_chat_model(
)
_is_openai_proxy = False
_original_provider: str | None = None
managed_builtin_profile = (
get_builtin_provider_profile(provider)
if provider in _STATIC_PROVIDER_IDS
else None
)
if provider in _STATIC_PROVIDER_IDS:
try:
managed_builtin_profile = get_builtin_provider_profile(provider)
except ProviderProfileError:
managed_builtin_profile = None
else:
managed_builtin_profile = None
if managed_builtin_profile is not None:
if not managed_builtin_profile.enabled:
raise ValueError(f"Provider {provider!r} is disabled.")
@@ -679,6 +802,23 @@ def get_chat_model(
kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"})
_patch_openrouter_strip_responses_reasoning()
if resolved_runtime is not None and resolved_runtime.reasoning_effort != "auto":
if dynamic_profile.adapter in {
"openai",
"openai-compatible",
"grok",
"antigravity",
"openrouter",
}:
kwargs.setdefault(
"reasoning", {"effort": resolved_runtime.reasoning_effort}
)
else:
raise ValueError(
f"REASONING_EFFORT_UNSUPPORTED: adapter {dynamic_profile.adapter!r} "
"does not support a portable reasoning effort setting."
)
elif provider == "anthropic":
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
if base_url:
@@ -729,8 +869,10 @@ def get_chat_model(
if base_url:
kwargs.setdefault("base_url", base_url)
api_key = os.environ.get(api_key_env, "")
if api_key:
kwargs.setdefault("api_key", api_key)
# ChatOpenAI otherwise falls back to OPENAI_API_KEY when ``api_key``
# is omitted. That is unsafe for routed providers: a missing Zhipu
# credential must never send an unrelated OpenAI credential to Zhipu.
kwargs.setdefault("api_key", api_key or None)
# SiliconFlow: disable thinking — LangChain drops reasoning_content
# from history, causing error 20015 on multi-turn requests.
if provider == "siliconflow":
@@ -864,6 +1006,20 @@ def get_chat_model(
_patch_ccproxy_system_to_developer(chat_model)
apply_known_context_window(chat_model)
if resolved_runtime is not None:
profile_data = getattr(chat_model, "profile", None)
base_profile = profile_data if isinstance(profile_data, dict) else {}
try:
chat_model.profile = {
**base_profile,
"max_input_tokens": resolved_runtime.resolved_input_limit,
"min_effective_input_tokens": (
resolved_runtime.min_effective_input_tokens
),
"runtime_capabilities": dict(resolved_runtime.capabilities),
}
except Exception:
pass
return chat_model
+401
View File
@@ -0,0 +1,401 @@
"""Server-side, per-run snapshots for custom provider model configuration.
The browser and LangGraph run configuration carry only a random snapshot ID.
The SQLite record stores non-secret connection metadata and runtime limits.
Credentials stay in a process-local cache for the lifetime of a development
deployment, so a restarted deployment fails an old run explicitly instead of
silently using a newly edited credential.
"""
from __future__ import annotations
import json
import sqlite3
import threading
import time
from dataclasses import dataclass, replace
from pathlib import Path
from typing import Any
from ..config.provider_profiles import (
ModelRuntime,
ProviderModel,
ProviderProfile,
ProviderProfileError,
ProviderRuntime,
get_provider_profile_revision,
resolve_model_input_limit,
resolve_provider_model,
)
from ..config.settings import get_config_dir
from .models import (
ResolvedRuntimeOptions,
is_static_provider_id,
normalize_provider_id,
resolve_runtime_options,
)
_SNAPSHOT_TTL_SECONDS = 7 * 24 * 60 * 60
_MAX_SNAPSHOT_ID_LENGTH = 128
_DB_LOCK = threading.RLock()
_MODEL_CACHE_LOCK = threading.RLock()
_MODEL_CACHE: dict[str, Any] = {}
_SNAPSHOT_SECRETS: dict[str, str] = {}
def _validate_static_provider_credentials(provider: str) -> None:
"""Reject known static selections that lack their provider-specific key.
Runtime snapshots are created by the WebUI before a run is queued. This is
the latest point at which we can return a useful client error without
making the graph fail after the user has already submitted a message.
"""
if provider not in {"zhipu", "zhipu-code"}:
return
from ..config.settings import get_effective_config
from .provider_operations import build_builtin_provider_profile
config = get_effective_config()
try:
profile = build_builtin_provider_profile(config, provider)
except ProviderProfileError as exc:
# A stale custom registry must not hide configuration supplied through
# config.yaml or the environment for a built-in provider.
if not str(exc).startswith("PROVIDER_PROFILE_RESET_REQUIRED"):
raise
profile = build_builtin_provider_profile(
config, provider, load_saved=False
)
if not profile.api_key.strip():
raise ProviderProfileError(
"ZHIPU_API_KEY_NOT_CONFIGURED: GLM requires a valid Zhipu API key. "
"Configure ZHIPU_API_KEY before starting a conversation."
)
@dataclass(frozen=True)
class RunRuntimeSnapshot:
"""An immutable server-side configuration captured before one run starts."""
snapshot_id: str
created_at: int
expires_at: int
profile_revision: str | None
profile: ProviderProfile
model: ProviderModel
options: ResolvedRuntimeOptions
def public_payload(self) -> dict[str, Any]:
"""Return diagnostic fields that are safe to attach to a run."""
return {
"snapshot_id": self.snapshot_id,
"provider": self.profile.id,
"model": self.model.id,
"profile_revision": self.profile_revision,
"created_at": self.created_at,
"expires_at": self.expires_at,
"runtime": {
"max_input_tokens": self.options.resolved_input_limit,
"max_output_tokens": self.options.max_output_tokens,
"min_effective_input_tokens": self.options.min_effective_input_tokens,
"timeout_seconds": self.options.timeout_seconds,
"max_retries": self.options.max_retries,
"temperature": self.options.temperature,
"top_p": self.options.top_p,
"reasoning_effort": self.options.reasoning_effort,
"limits_source": self.options.limits_source,
"capabilities": dict(self.options.capabilities),
},
}
def _database_path() -> Path:
return get_config_dir() / "run-runtime-snapshots.sqlite3"
def _connect() -> sqlite3.Connection:
path = _database_path()
path.parent.mkdir(parents=True, exist_ok=True)
try:
path.parent.chmod(0o700)
except OSError:
pass
connection = sqlite3.connect(path)
connection.execute(
"""
CREATE TABLE IF NOT EXISTS run_runtime_snapshots (
snapshot_id TEXT PRIMARY KEY,
expires_at INTEGER NOT NULL,
payload_json TEXT NOT NULL
)
"""
)
try:
path.chmod(0o600)
except OSError:
pass
return connection
def _validate_snapshot_id(snapshot_id: str) -> str:
normalized = snapshot_id.strip()
if not normalized or len(normalized) > _MAX_SNAPSHOT_ID_LENGTH:
raise ProviderProfileError("runtime snapshot ID is invalid.")
if any(ord(character) < 33 or ord(character) > 126 for character in normalized):
raise ProviderProfileError("runtime snapshot ID is invalid.")
return normalized
def _provider_runtime_payload(runtime: ProviderRuntime) -> dict[str, Any]:
return {
"timeout_seconds": runtime.timeout_seconds,
"max_retries": runtime.max_retries,
"default_temperature": runtime.default_temperature,
"default_top_p": runtime.default_top_p,
"default_reasoning_effort": runtime.default_reasoning_effort,
}
def _model_runtime_payload(runtime: ModelRuntime) -> dict[str, Any]:
return {
"limit_mode": runtime.limit_mode,
"context_window_tokens": runtime.context_window_tokens,
"max_input_tokens": runtime.max_input_tokens,
"max_output_tokens": runtime.max_output_tokens,
"min_effective_input_tokens": runtime.min_effective_input_tokens,
"limits_status": runtime.limits_status,
"limits_source": runtime.limits_source,
"temperature": runtime.temperature,
"top_p": runtime.top_p,
"reasoning_effort": runtime.reasoning_effort,
"capabilities": dict(runtime.capabilities),
}
def _profile_payload(profile: ProviderProfile, model: ProviderModel) -> dict[str, Any]:
return {
"id": profile.id,
"name": profile.name,
"adapter": profile.adapter,
"base_url": profile.base_url,
# Credentials stay in the process-local secret cache. The persisted
# snapshot remains safe to inspect through diagnostics and never
# duplicates a Provider Profile API key.
"api_key": "",
"auth_mode": profile.auth_mode,
"enabled": profile.enabled,
"runtime": _provider_runtime_payload(profile.runtime),
"models": [
{
"id": model.id,
"name": model.name,
"model_id": model.model_id,
"enabled": model.enabled,
"runtime": _model_runtime_payload(model.runtime),
}
],
}
def _parse_snapshot_payload(
payload: dict[str, Any], *, require_secret: bool = True
) -> RunRuntimeSnapshot:
try:
profile_raw = payload["profile"]
if not isinstance(profile_raw, dict):
raise ProviderProfileError("runtime snapshot profile is invalid.")
runtime_raw = payload["runtime"]
if not isinstance(runtime_raw, dict):
raise ProviderProfileError("runtime snapshot options are invalid.")
# The source is a private snapshot, not a browser draft. Reuse the
# validated profile parser without looking up the current registry.
from ..config.provider_profiles import _parse_profile
profile = _parse_profile(profile_raw, index=0)
model = profile.models[0]
resolved_input_limit = int(runtime_raw["resolved_input_limit"])
if resolved_input_limit != resolve_model_input_limit(model.runtime):
raise ProviderProfileError("runtime snapshot limits are invalid.")
capabilities_raw = dict(runtime_raw["capabilities"])
options = ResolvedRuntimeOptions(
profile_revision=(
payload["profile_revision"]
if isinstance(payload.get("profile_revision"), str)
else None
),
model_id=str(runtime_raw["model_id"]),
adapter_id=str(runtime_raw["adapter_id"]),
limit_mode=str(runtime_raw["limit_mode"]),
resolved_input_limit=resolved_input_limit,
max_output_tokens=int(runtime_raw["max_output_tokens"]),
min_effective_input_tokens=int(
runtime_raw["min_effective_input_tokens"]
),
timeout_seconds=int(runtime_raw["timeout_seconds"]),
max_retries=int(runtime_raw["max_retries"]),
temperature=runtime_raw.get("temperature"),
top_p=runtime_raw.get("top_p"),
reasoning_effort=str(runtime_raw["reasoning_effort"]),
limits_source=str(runtime_raw["limits_source"]),
capabilities=tuple(
(key, capabilities_raw[key])
for key in ("tools", "vision", "structured_output")
if capabilities_raw.get(key) in {True, False, "auto"}
),
)
snapshot_id = _validate_snapshot_id(str(payload["snapshot_id"]))
secret = _SNAPSHOT_SECRETS.get(snapshot_id)
if require_secret and snapshot_id not in _SNAPSHOT_SECRETS:
raise ProviderProfileError(
"RUN_RUNTIME_SNAPSHOT_CREDENTIAL_UNAVAILABLE: the deployment was "
"restarted after this run began. Start the message again."
)
return RunRuntimeSnapshot(
snapshot_id=snapshot_id,
created_at=int(payload["created_at"]),
expires_at=int(payload["expires_at"]),
profile_revision=options.profile_revision,
profile=replace(profile, api_key=secret or ""),
model=model,
options=options,
)
except (KeyError, TypeError, ValueError) as exc:
raise ProviderProfileError("runtime snapshot is invalid.") from exc
def _snapshot_payload(snapshot: RunRuntimeSnapshot) -> dict[str, Any]:
return {
"snapshot_id": snapshot.snapshot_id,
"created_at": snapshot.created_at,
"expires_at": snapshot.expires_at,
"profile_revision": snapshot.profile_revision,
"profile": _profile_payload(snapshot.profile, snapshot.model),
"runtime": {
"model_id": snapshot.options.model_id,
"adapter_id": snapshot.options.adapter_id,
"limit_mode": snapshot.options.limit_mode,
"resolved_input_limit": snapshot.options.resolved_input_limit,
"max_output_tokens": snapshot.options.max_output_tokens,
"min_effective_input_tokens": snapshot.options.min_effective_input_tokens,
"timeout_seconds": snapshot.options.timeout_seconds,
"max_retries": snapshot.options.max_retries,
"temperature": snapshot.options.temperature,
"top_p": snapshot.options.top_p,
"reasoning_effort": snapshot.options.reasoning_effort,
"limits_source": snapshot.options.limits_source,
"capabilities": dict(snapshot.options.capabilities),
},
}
def create_run_runtime_snapshot(
snapshot_id: str,
*,
model: str,
provider: str,
) -> RunRuntimeSnapshot | None:
"""Freeze one custom provider model configuration for a forthcoming run.
Static built-in providers remain managed by their existing deployment
configuration and return ``None``. A missing custom model is an error so
a model selection never silently falls back to a different model.
"""
snapshot_id = _validate_snapshot_id(snapshot_id)
model = model.strip()
provider = normalize_provider_id(provider)
if not model or not provider:
raise ProviderProfileError("model and provider are required for a runtime snapshot.")
# Built-in providers are resolved from the deployment configuration and
# have no browser-managed secret or runtime profile to freeze. Avoid
# loading custom profiles at all, so a stale development registry cannot
# block a normal built-in run.
if is_static_provider_id(provider):
_validate_static_provider_credentials(provider)
return None
existing = get_run_runtime_snapshot(snapshot_id)
if existing is not None:
if existing.model.id != model or existing.profile.id != provider:
raise ProviderProfileError(
"runtime snapshot ID is already bound to a different model selection."
)
return existing
resolved = resolve_provider_model(provider, model)
if resolved is None:
return None
profile, provider_model = resolved
options = resolve_runtime_options(profile, provider_model)
now = int(time.time())
snapshot = RunRuntimeSnapshot(
snapshot_id=snapshot_id,
created_at=now,
expires_at=now + _SNAPSHOT_TTL_SECONDS,
profile_revision=get_provider_profile_revision(provider),
profile=profile,
model=provider_model,
options=options,
)
encoded = json.dumps(
_snapshot_payload(snapshot), sort_keys=True, separators=(",", ":")
)
with _DB_LOCK, _connect() as connection:
connection.execute(
"DELETE FROM run_runtime_snapshots WHERE expires_at <= ?", (now,)
)
connection.execute(
"INSERT OR REPLACE INTO run_runtime_snapshots(snapshot_id, expires_at, payload_json) "
"VALUES (?, ?, ?)",
(snapshot.snapshot_id, snapshot.expires_at, encoded),
)
connection.commit()
_SNAPSHOT_SECRETS[snapshot.snapshot_id] = profile.api_key
return snapshot
def get_run_runtime_snapshot(snapshot_id: str) -> RunRuntimeSnapshot | None:
"""Load a still-valid runtime snapshot by its opaque ID."""
snapshot_id = _validate_snapshot_id(snapshot_id)
now = int(time.time())
with _DB_LOCK, _connect() as connection:
connection.execute(
"DELETE FROM run_runtime_snapshots WHERE expires_at <= ?", (now,)
)
row = connection.execute(
"SELECT expires_at, payload_json FROM run_runtime_snapshots WHERE snapshot_id = ?",
(snapshot_id,),
).fetchone()
connection.commit()
if row is None or int(row[0]) <= now:
_SNAPSHOT_SECRETS.pop(snapshot_id, None)
return None
try:
payload = json.loads(str(row[1]))
except json.JSONDecodeError as exc:
raise ProviderProfileError("runtime snapshot is unreadable.") from exc
if not isinstance(payload, dict):
raise ProviderProfileError("runtime snapshot is unreadable.")
return _parse_snapshot_payload(payload)
def get_snapshot_chat_model(snapshot_id: str) -> Any:
"""Build (or reuse) the chat model bound to a frozen runtime snapshot."""
snapshot = get_run_runtime_snapshot(snapshot_id)
if snapshot is None:
raise ProviderProfileError(
"RUN_RUNTIME_SNAPSHOT_UNAVAILABLE: the run configuration snapshot "
"expired or is unavailable. Start the message again."
)
with _MODEL_CACHE_LOCK:
cached = _MODEL_CACHE.get(snapshot.snapshot_id)
if cached is not None:
return cached
from .models import get_profile_chat_model
model = get_profile_chat_model(snapshot.profile, snapshot.model)
with _MODEL_CACHE_LOCK:
_MODEL_CACHE[snapshot.snapshot_id] = model
return model
+8
View File
@@ -27,6 +27,11 @@ from .memory_lifecycle import (
create_memory_lifecycle_middleware,
default_memory_scheduler,
)
from .message_budget import (
MessageReservePolicy,
count_message_text_tokens,
create_message_budget_middleware,
)
from .model_fallback import ModelFallbackMiddleware, load_fallback_chain
from .runtime_context import RuntimeContextMiddleware, create_runtime_context_middleware
from .scheduler import (
@@ -46,16 +51,19 @@ __all__ = [
"ContextOverflowMapperMiddleware",
"EvoMemoryLifecycleMiddleware",
"EvoMemoryMiddleware",
"MessageReservePolicy",
"ModelFallbackMiddleware",
"Question",
"RuntimeContextMiddleware",
"SchedulerMiddleware",
"ToolErrorHandlerMiddleware",
"compute_context_editing_trigger",
"count_message_text_tokens",
"create_code_interpreter_middleware",
"create_context_editing_middleware",
"create_memory_lifecycle_middleware",
"create_memory_middleware",
"create_message_budget_middleware",
"create_runtime_context_middleware",
"create_scheduler_middleware",
"create_tool_selector_middleware",
+52 -3
View File
@@ -74,6 +74,23 @@ def _read_model_override() -> tuple[str | None, str | None]:
)
def _read_runtime_snapshot_id() -> str | None:
"""Return the opaque server-side runtime snapshot ID for this run."""
try:
from langgraph.config import get_config
cfg = get_config()
except Exception:
return None
if not isinstance(cfg, dict):
return None
configurable = cfg.get("configurable") or {}
if not isinstance(configurable, dict):
return None
snapshot_id = configurable.get("runtime_snapshot_id")
return snapshot_id if isinstance(snapshot_id, str) and snapshot_id else None
class ConfigurableModelMiddleware(AgentMiddleware):
"""Re-resolve the chat model from RunnableConfig.configurable on every call.
@@ -128,8 +145,32 @@ class ConfigurableModelMiddleware(AgentMiddleware):
provider,
)
def _resolve(self, model: str, provider: str | None) -> Any:
def _resolve(
self, model: str, provider: str | None, runtime_snapshot_id: str | None
) -> Any:
"""Return a cached or freshly-built chat model for ``(model, provider)``."""
if runtime_snapshot_id is not None:
from ..llm.models import get_profile_chat_model
from ..llm.runtime_snapshots import get_run_runtime_snapshot
snapshot = get_run_runtime_snapshot(runtime_snapshot_id)
if snapshot is None:
raise ValueError(
"RUN_RUNTIME_SNAPSHOT_UNAVAILABLE: the run configuration snapshot "
"expired or is unavailable. Start the message again."
)
if snapshot.model.id != model or snapshot.profile.id != provider:
raise ValueError("RUN_RUNTIME_SNAPSHOT_MISMATCH: run configuration is invalid.")
key = ("snapshot", runtime_snapshot_id, snapshot.profile_revision)
with self._lock:
cached = self._cache.get(key)
if cached is not None:
return cached
new_model = get_profile_chat_model(snapshot.profile, snapshot.model)
with self._lock:
self._cache[key] = new_model
return new_model
from ..llm.models import get_model_runtime_revision
revision = get_model_runtime_revision(provider)
@@ -156,9 +197,12 @@ class ConfigurableModelMiddleware(AgentMiddleware):
model_name, provider = _read_model_override()
if model_name is None:
return handler(request)
runtime_snapshot_id = _read_runtime_snapshot_id()
try:
new_model = self._resolve(model_name, provider)
new_model = self._resolve(model_name, provider, runtime_snapshot_id)
except Exception:
if runtime_snapshot_id is not None:
raise
logger.warning(
"ConfigurableModelMiddleware failed to resolve model=%r "
"provider=%r; falling back to compile-time model",
@@ -178,6 +222,7 @@ class ConfigurableModelMiddleware(AgentMiddleware):
model_name, provider = _read_model_override()
if model_name is None:
return await handler(request)
runtime_snapshot_id = _read_runtime_snapshot_id()
try:
# Offload first-call SDK init off the event loop. ``_resolve`` calls
# ``get_chat_model`` on a cache miss, which can spend hundreds of ms
@@ -185,8 +230,12 @@ class ConfigurableModelMiddleware(AgentMiddleware):
# ``async def`` would block every other coroutine on the same
# langgraph dev event loop. Cache hits are still fast (a dict
# lookup); the thread-pool overhead is irrelevant once warm.
new_model = await asyncio.to_thread(self._resolve, model_name, provider)
new_model = await asyncio.to_thread(
self._resolve, model_name, provider, runtime_snapshot_id
)
except Exception:
if runtime_snapshot_id is not None:
raise
logger.warning(
"ConfigurableModelMiddleware failed to resolve model=%r "
"provider=%r; falling back to compile-time model",
+327
View File
@@ -0,0 +1,327 @@
"""Message-only context budgeting and automatic conversation compaction.
The middleware deliberately does not attempt to count the whole provider
request. System instructions, memories, tools, and attachments instead consume
fixed conservative reserves. The measured value is only textual conversation
messages and tool-result text, which is the part compaction can actually
reduce.
"""
from __future__ import annotations
import threading
from collections.abc import Iterable, Mapping
from contextvars import ContextVar
from dataclasses import dataclass
from typing import Any
from langchain_core.messages import AnyMessage, SystemMessage
@dataclass(frozen=True)
class MessageReservePolicy:
"""Static non-message reserves used without full-request token counting."""
system_tokens: int = 4_096
memory_tokens: int = 8_192
tool_tokens: int = 8_192
attachment_tokens: int = 8_192
safety_fraction: float = 0.10
minimum_safety_tokens: int = 2_048
soft_fraction: float = 0.70
keep_fraction: float = 0.35
def fixed_tokens(
self, *, has_tools: bool = False, has_attachments: bool = False
) -> int:
"""Return only the fixed reserves applicable to this model request."""
return (
self.system_tokens
+ self.memory_tokens
+ (self.tool_tokens if has_tools else 0)
+ (self.attachment_tokens if has_attachments else 0)
)
def hard_budget(
self,
input_limit: int,
*,
has_tools: bool = False,
has_attachments: bool = False,
) -> int:
safety = max(
self.minimum_safety_tokens, int(input_limit * self.safety_fraction)
)
return max(
1_024,
input_limit
- self.fixed_tokens(
has_tools=has_tools, has_attachments=has_attachments
)
- safety,
)
def soft_budget(self, input_limit: int, **mode: bool) -> int:
return max(
1_024, int(self.hard_budget(input_limit, **mode) * self.soft_fraction)
)
def keep_budget(self, input_limit: int, **mode: bool) -> int:
return max(
1_024, int(self.hard_budget(input_limit, **mode) * self.keep_fraction)
)
@dataclass(frozen=True)
class MessageBudget:
"""The message budget selected for one model invocation."""
input_limit: int
hard_tokens: int
soft_tokens: int
keep_tokens: int
min_effective_input_tokens: int
has_tools: bool
has_attachments: bool
class ContextBudgetUnsatisfiableError(ValueError):
"""Raised before a call when static reserves leave no safe message budget."""
_ACTIVE_BUDGET: ContextVar[MessageBudget | None] = ContextVar(
"evoscientist_message_budget", default=None
)
def _content_text(content: Any) -> str:
"""Extract text without serializing image/file blocks or tool schemas."""
if isinstance(content, str):
return content
if not isinstance(content, list):
return ""
parts: list[str] = []
for block in content:
if isinstance(block, str):
parts.append(block)
continue
if not isinstance(block, Mapping):
continue
block_type = block.get("type")
if block_type in {"image", "image_url", "file", "document", "audio", "video"}:
continue
text = block.get("text")
if isinstance(text, str):
parts.append(text)
return "\n".join(parts)
def count_message_text_tokens(messages: Iterable[AnyMessage]) -> int:
"""Conservative, provider-neutral token estimate for compactable text only."""
characters = sum(len(_content_text(message.content)) for message in messages)
# A four-character estimate deliberately errs slightly high for Chinese and
# mixed code while remaining cheap enough to run before every model call.
return max(0, (characters + 3) // 4)
def _has_attachments(messages: Iterable[AnyMessage]) -> bool:
attachment_types = {"image", "image_url", "file", "document", "audio", "video"}
for message in messages:
if not isinstance(message.content, list):
continue
for block in message.content:
if isinstance(block, Mapping) and block.get("type") in attachment_types:
return True
return False
def _current_configurable() -> Mapping[str, Any]:
try:
from langgraph.config import get_config
config = get_config()
except Exception:
return {}
if not isinstance(config, Mapping):
return {}
configurable = config.get("configurable")
return configurable if isinstance(configurable, Mapping) else {}
class MessageBudgetMiddleware:
"""Factory namespace kept separate from DeepAgents' concrete middleware."""
@staticmethod
def create(model: Any, backend: Any, *, policy: MessageReservePolicy | None = None):
"""Create the runtime-aware DeepAgents summarization middleware."""
from deepagents.middleware.summarization import SummarizationMiddleware
class _RuntimeMessageBudgetMiddleware(SummarizationMiddleware):
name = "message_budget"
def __init__(self) -> None:
self._fallback_model = model
self._model_cache: dict[tuple[str, str | None], Any] = {}
self._model_cache_lock = threading.RLock()
self._policy = policy or MessageReservePolicy()
# Triggering and cutoff are overridden below. The base class is
# still used for safe AI/tool-pair handling, offloading, and
# persisted summarization events.
super().__init__(
model=model,
backend=backend,
trigger=("tokens", 1_000_000_000),
keep=("messages", 8),
token_counter=count_message_text_tokens,
trim_tokens_to_summarize=4_000,
truncate_args_settings={
"trigger": ("tokens", 2_048),
"keep": ("messages", 8),
"max_length": 2_000,
"truncation_text": "...(tool arguments compacted)",
},
)
@property
def model(self) -> Any: # type: ignore[override]
configurable = _current_configurable()
snapshot_id = configurable.get("runtime_snapshot_id")
if isinstance(snapshot_id, str) and snapshot_id:
from ..llm.runtime_snapshots import get_snapshot_chat_model
return get_snapshot_chat_model(snapshot_id)
model_name = configurable.get("model")
provider = configurable.get("model_provider")
if not isinstance(model_name, str) or not model_name:
return self._fallback_model
provider_name = provider if isinstance(provider, str) and provider else None
key = (model_name, provider_name)
with self._model_cache_lock:
cached = self._model_cache.get(key)
if cached is not None:
return cached
from ..llm.models import get_chat_model
resolved = get_chat_model(model=model_name, provider=provider_name)
with self._model_cache_lock:
self._model_cache[key] = resolved
return resolved
def _input_limit(self) -> int:
profile = getattr(self.model, "profile", None)
if isinstance(profile, Mapping):
value = profile.get("max_input_tokens")
if isinstance(value, int) and not isinstance(value, bool) and value > 0:
return value
# Existing static providers that lack profile data use a
# conservative lower fallback. Custom profiles always supply
# `max_input_tokens` through the frozen runtime options.
return 32_768
def _minimum_effective_input(self) -> int:
profile = getattr(self.model, "profile", None)
if isinstance(profile, Mapping):
value = profile.get("min_effective_input_tokens")
if isinstance(value, int) and not isinstance(value, bool) and value > 0:
return value
return 1_024
def _budget_for_request(self, request: Any) -> MessageBudget:
input_limit = self._input_limit()
has_tools = bool(getattr(request, "tools", None))
has_attachments = _has_attachments(getattr(request, "messages", []))
mode = {
"has_tools": has_tools,
"has_attachments": has_attachments,
}
budget = MessageBudget(
input_limit=input_limit,
hard_tokens=self._policy.hard_budget(input_limit, **mode),
soft_tokens=self._policy.soft_budget(input_limit, **mode),
keep_tokens=self._policy.keep_budget(input_limit, **mode),
min_effective_input_tokens=self._minimum_effective_input(),
has_tools=has_tools,
has_attachments=has_attachments,
)
if budget.hard_tokens < budget.min_effective_input_tokens:
raise ContextBudgetUnsatisfiableError(
"CONTEXT_BUDGET_UNSATISFIABLE: configured input limit "
f"{budget.input_limit:,} leaves only {budget.hard_tokens:,} "
"tokens after fixed reserves; increase the model window, reduce "
"the output budget, or disable tools/attachments."
)
return budget
def _active_budget(self) -> MessageBudget:
active = _ACTIVE_BUDGET.get()
if active is not None:
return active
input_limit = self._input_limit()
return MessageBudget(
input_limit=input_limit,
hard_tokens=self._policy.hard_budget(input_limit),
soft_tokens=self._policy.soft_budget(input_limit),
keep_tokens=self._policy.keep_budget(input_limit),
min_effective_input_tokens=self._minimum_effective_input(),
has_tools=False,
has_attachments=False,
)
def wrap_model_call(self, request: Any, handler: Any) -> Any:
token = _ACTIVE_BUDGET.set(self._budget_for_request(request))
try:
return super().wrap_model_call(request, handler)
finally:
_ACTIVE_BUDGET.reset(token)
async def awrap_model_call(self, request: Any, handler: Any) -> Any:
token = _ACTIVE_BUDGET.set(self._budget_for_request(request))
try:
return await super().awrap_model_call(request, handler)
finally:
_ACTIVE_BUDGET.reset(token)
def _get_profile_limits(self) -> int | None:
return self._active_budget().hard_tokens
def _count_tokens(
self,
messages: list[AnyMessage],
system_message: SystemMessage | None,
tools: list[Any] | None,
) -> int:
del system_message, tools
return count_message_text_tokens(messages)
def _should_summarize(
self, messages: list[AnyMessage], total_tokens: int
) -> bool:
del messages
return total_tokens >= self._active_budget().soft_tokens
def _determine_cutoff_index(self, messages: list[AnyMessage]) -> int:
keep_tokens = self._active_budget().keep_tokens
retained = 0
cutoff = len(messages)
for index in range(len(messages) - 1, -1, -1):
message_tokens = count_message_text_tokens([messages[index]])
if retained + message_tokens > keep_tokens:
# Keep at least the newest message even when it alone
# exceeds the retention budget. The safe-cutoff helper
# below expands that boundary when it is a tool result.
cutoff = index if retained == 0 else index + 1
break
retained += message_tokens
cutoff = index
if cutoff <= 0:
return 0
return self._lc_helper._find_safe_cutoff_point(messages, cutoff)
return _RuntimeMessageBudgetMiddleware()
def create_message_budget_middleware(
model: Any, backend: Any, *, policy: MessageReservePolicy | None = None
):
"""Construct automatic compaction middleware for a graph backend."""
return MessageBudgetMiddleware.create(model, backend, policy=policy)
+44
View File
@@ -124,12 +124,21 @@ def test_provider_profiles_api_round_trip_redacts_secret(tmp_path, monkeypatch):
"base_url": "https://llm.example.test/v1",
"api_key": "provider-secret",
"enabled": True,
"runtime": {"timeout_seconds": 90, "max_retries": 1},
"models": [
{
"id": "lab-model",
"name": "Lab Model",
"model_id": "vendor/model",
"enabled": True,
"runtime": {
"limit_mode": "combined",
"context_window_tokens": 32_768,
"max_output_tokens": 4_096,
"min_effective_input_tokens": 4_096,
"limits_status": "confirmed",
"limits_source": "provider",
},
}
],
}
@@ -147,6 +156,41 @@ def test_provider_profiles_api_round_trip_redacts_secret(tmp_path, monkeypatch):
assert get_response.json()["providers"][0]["models"][0]["id"] == "lab-model"
assert "provider-secret" not in get_response.text
snapshot_response = client.post(
"/api/runtime-snapshots",
headers=headers,
json={
"snapshot_id": "run-request-a",
"model": "lab-model",
"provider": "lab-openai",
},
)
assert snapshot_response.status_code == 200
snapshot = snapshot_response.json()["snapshot"]
assert snapshot["runtime"]["max_input_tokens"] == 28_672
assert "provider-secret" not in snapshot_response.text
def test_runtime_snapshot_reports_missing_zhipu_key_before_creating_a_run(
monkeypatch,
):
monkeypatch.setenv("EVOSCIENTIST_PROVIDER_ADMIN_TOKEN", "admin-secret")
monkeypatch.delenv("ZHIPU_API_KEY", raising=False)
monkeypatch.setenv("OPENAI_API_KEY", "unrelated-openai-key")
response = client.post(
"/api/runtime-snapshots",
headers={"X-EvoScientist-Admin-Token": "admin-secret"},
json={
"snapshot_id": "run-request-zhipu",
"model": "glm-5.2",
"provider": "glm",
},
)
assert response.status_code == 400
assert response.json()["error"].startswith("ZHIPU_API_KEY_NOT_CONFIGURED")
def test_llm_config_api_requires_admin_token_header(monkeypatch):
monkeypatch.delenv("EVOSCIENTIST_PROVIDER_ADMIN_TOKEN", raising=False)
+23
View File
@@ -218,6 +218,29 @@ class TestGetChatModel:
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model_provider"] == "custom_provider"
@patch("EvoScientist.llm.models.init_chat_model")
def test_glm_provider_alias_routes_to_zhipu(self, mock_init):
"""The historic ``glm`` shorthand uses the Zhipu compatible endpoint."""
mock_init.return_value = "mock_model"
get_chat_model("glm-5.2", provider="glm")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model"] == "glm-5.2"
assert call_kwargs["model_provider"] == "openai"
assert call_kwargs["base_url"] == "https://open.bigmodel.cn/api/paas/v4"
@patch("EvoScientist.llm.models.init_chat_model")
def test_zhipu_does_not_fall_back_to_openai_api_key(self, mock_init, monkeypatch):
"""A missing Zhipu key must not be sent from OPENAI_API_KEY instead."""
monkeypatch.setenv("OPENAI_API_KEY", "unrelated-openai-key")
monkeypatch.delenv("ZHIPU_API_KEY", raising=False)
mock_init.return_value = "mock_model"
get_chat_model("glm-5.2", provider="zhipu")
assert mock_init.call_args.kwargs["api_key"] is None
@patch("EvoScientist.llm.models.init_chat_model")
def test_passes_kwargs(self, mock_init):
"""Test that additional kwargs are passed through."""
+79
View File
@@ -0,0 +1,79 @@
"""Tests for message-only context budgeting."""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from EvoScientist.middleware.message_budget import (
ContextBudgetUnsatisfiableError,
MessageReservePolicy,
count_message_text_tokens,
create_message_budget_middleware,
)
def test_text_counter_excludes_attachment_payloads_and_counts_tool_results():
messages = [
HumanMessage(
content=[
{"type": "text", "text": "abcd"},
{"type": "image", "base64": "x" * 100_000},
{"type": "file", "data": "y" * 100_000},
]
),
ToolMessage(content="wxyz", tool_call_id="tool-1"),
]
assert count_message_text_tokens(messages) == 2
def test_reserve_policy_uses_fixed_overhead_without_request_token_counting():
policy = MessageReservePolicy()
assert policy.hard_budget(32_768) == 17_204
assert policy.soft_budget(32_768) == 12_042
assert policy.keep_budget(32_768) == 6_021
assert policy.hard_budget(32_768, has_tools=True) == 9_012
def test_budget_middleware_uses_message_threshold_and_safe_tool_cutoff():
model = MagicMock()
model.profile = {"max_input_tokens": 32_768}
middleware = create_message_budget_middleware(model, MagicMock())
messages = [
HumanMessage(content="a" * 30_000),
AIMessage(content="", tool_calls=[{"name": "read_file", "args": {}, "id": "tool-1"}]),
ToolMessage(content="b" * 30_000, tool_call_id="tool-1"),
HumanMessage(content="c" * 30_000),
]
total = count_message_text_tokens(messages)
assert middleware._should_summarize(messages, total) is True
cutoff = middleware._determine_cutoff_index(messages)
assert cutoff in {1, 3}
# A cutoff never leaves the tool response without its matching AI tool call.
if cutoff == 1:
assert isinstance(messages[cutoff], AIMessage)
def test_budget_rejects_tools_and_attachments_when_reserves_exceed_model_limit():
model = MagicMock()
model.profile = {
"max_input_tokens": 28_672,
"min_effective_input_tokens": 4_096,
}
middleware = create_message_budget_middleware(model, MagicMock())
request = SimpleNamespace(
messages=[
HumanMessage(content=[{"type": "image", "url": "https://example.test/a"}])
],
tools=[{"name": "read_file"}],
)
with pytest.raises(ContextBudgetUnsatisfiableError, match="CONTEXT_BUDGET_UNSATISFIABLE"):
middleware._budget_for_request(request)
+40 -8
View File
@@ -30,6 +30,7 @@ def provider_config_dir(tmp_path, monkeypatch):
def _document(api_key: str = "secret-key") -> dict:
return {
"version": 3,
"providers": [
{
"id": "lab-openai",
@@ -38,18 +39,41 @@ def _document(api_key: str = "secret-key") -> dict:
"base_url": "https://llm.example.test/v1/",
"api_key": api_key,
"enabled": True,
"runtime": {
"timeout_seconds": 90,
"max_retries": 1,
"default_temperature": 0.3,
"default_top_p": None,
"default_reasoning_effort": "auto",
},
"models": [
{
"id": "research-model",
"name": "Research Model",
"model_id": "vendor/research-1",
"enabled": True,
"runtime": {
"limit_mode": "combined",
"context_window_tokens": 32768,
"max_output_tokens": 4096,
"min_effective_input_tokens": 4096,
"limits_status": "confirmed",
"limits_source": "user",
},
},
{
"id": "disabled-model",
"name": "Disabled Model",
"model_id": "vendor/disabled",
"enabled": False,
"runtime": {
"limit_mode": "combined",
"context_window_tokens": 32768,
"max_output_tokens": 4096,
"min_effective_input_tokens": 4096,
"limits_status": "confirmed",
"limits_source": "user",
},
},
],
}
@@ -59,6 +83,7 @@ def _document(api_key: str = "secret-key") -> dict:
def _builtin_document(api_key: str = "builtin-secret") -> dict:
return {
"version": 3,
"builtins": [
{
"id": "openai",
@@ -68,12 +93,21 @@ def _builtin_document(api_key: str = "builtin-secret") -> dict:
"api_key": api_key,
"auth_mode": "api_key",
"enabled": True,
"runtime": {},
"models": [
{
"id": "chat-main",
"name": "Chat Main",
"model_id": "gpt-upstream",
"enabled": True,
"runtime": {
"limit_mode": "combined",
"context_window_tokens": 32768,
"max_output_tokens": 4096,
"min_effective_input_tokens": 4096,
"limits_status": "confirmed",
"limits_source": "user",
},
}
],
}
@@ -99,7 +133,7 @@ def test_replace_round_trip_and_redacts_api_key(provider_config_dir):
assert path.stat().st_mode & 0o777 == 0o600
def test_v1_document_loads_and_saves_as_v2(provider_config_dir):
def test_v1_document_requires_reset(provider_config_dir):
path = get_provider_profiles_path()
path.parent.mkdir(parents=True)
path.write_text(
@@ -114,13 +148,8 @@ def test_v1_document_loads_and_saves_as_v2(provider_config_dir):
encoding="utf-8",
)
loaded = load_provider_profiles()
assert loaded.version == 2
assert loaded.builtins == ()
assert loaded.providers[0].id == "lab-openai"
replace_provider_profiles({"providers": _document()["providers"]})
assert path.read_text(encoding="utf-8").startswith("version: 2\nbuiltins:")
with pytest.raises(ProviderProfileError, match="PROVIDER_PROFILE_RESET_REQUIRED"):
load_provider_profiles()
def test_builtin_profiles_round_trip_and_redact_secret(provider_config_dir):
@@ -300,6 +329,9 @@ def test_dynamic_profile_routes_through_selected_adapter(provider_config_dir):
assert kwargs["model_provider"] == "openai"
assert kwargs["base_url"] == "https://llm.example.test/v1"
assert kwargs["api_key"] == "secret-key"
assert kwargs["timeout"] == 90
assert kwargs["max_retries"] == 1
assert kwargs["max_tokens"] == 4096
assert kwargs["default_headers"]["User-Agent"] == "codex_cli_rs/0.0.0"
+152
View File
@@ -0,0 +1,152 @@
"""Tests for server-side run configuration snapshots."""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
from EvoScientist.config.provider_profiles import (
ProviderProfileError,
replace_provider_profiles,
)
from EvoScientist.llm.runtime_snapshots import (
create_run_runtime_snapshot,
get_run_runtime_snapshot,
get_snapshot_chat_model,
)
@pytest.fixture
def provider_config_dir(tmp_path, monkeypatch):
monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path))
return tmp_path / "evoscientist"
def _document(api_key: str = "first-secret", base_url: str = "https://one.test/v1"):
return {
"version": 3,
"providers": [
{
"id": "lab-openai",
"name": "Lab OpenAI",
"adapter": "openai-compatible",
"base_url": base_url,
"api_key": api_key,
"enabled": True,
"runtime": {"timeout_seconds": 90, "max_retries": 1},
"models": [
{
"id": "research-model",
"name": "Research Model",
"model_id": "vendor/research-1",
"enabled": True,
"runtime": {
"limit_mode": "combined",
"context_window_tokens": 32_768,
"max_output_tokens": 4_096,
"min_effective_input_tokens": 4_096,
"limits_status": "confirmed",
"limits_source": "provider",
},
}
],
}
],
}
def test_snapshot_keeps_connection_and_runtime_after_profile_changes(provider_config_dir):
replace_provider_profiles(_document())
created = create_run_runtime_snapshot(
"run-request-a", model="research-model", provider="lab-openai"
)
assert created is not None
assert created.options.resolved_input_limit == 28_672
assert created.profile.api_key == "first-secret"
assert "first-secret" not in str(created.public_payload())
assert "first-secret" not in (
provider_config_dir / "run-runtime-snapshots.sqlite3"
).read_bytes().decode("utf-8", errors="ignore")
replace_provider_profiles(
_document("second-secret", "https://two.test/v1")
)
loaded = get_run_runtime_snapshot("run-request-a")
assert loaded is not None
assert loaded.profile.api_key == "first-secret"
assert loaded.profile.base_url == "https://one.test/v1"
assert loaded.options.timeout_seconds == 90
def test_snapshot_id_is_idempotent_and_cannot_change_model(provider_config_dir):
replace_provider_profiles(_document())
first = create_run_runtime_snapshot(
"run-request-a", model="research-model", provider="lab-openai"
)
repeated = create_run_runtime_snapshot(
"run-request-a", model="research-model", provider="lab-openai"
)
assert repeated == first
with pytest.raises(ValueError, match="already bound"):
create_run_runtime_snapshot(
"run-request-a", model="another-model", provider="lab-openai"
)
def test_static_provider_does_not_need_a_custom_snapshot(provider_config_dir):
assert (
create_run_runtime_snapshot(
"run-request-a", model="gpt-5", provider="openai"
)
is None
)
def test_glm_alias_is_static_even_when_custom_registry_is_stale(
provider_config_dir, monkeypatch
):
monkeypatch.setenv("ZHIPU_API_KEY", "test-zhipu-key")
provider_config_dir.mkdir(parents=True)
(provider_config_dir / "providers.yaml").write_text(
"version: 2\nproviders: []\n", encoding="utf-8"
)
assert (
create_run_runtime_snapshot(
"run-request-glm", model="glm-5.2", provider="glm"
)
is None
)
def test_zhipu_snapshot_requires_its_own_api_key(provider_config_dir, monkeypatch):
monkeypatch.delenv("ZHIPU_API_KEY", raising=False)
monkeypatch.setenv("OPENAI_API_KEY", "unrelated-openai-key")
with pytest.raises(ProviderProfileError, match="ZHIPU_API_KEY_NOT_CONFIGURED"):
create_run_runtime_snapshot(
"run-request-zhipu", model="glm-5.2", provider="zhipu"
)
def test_snapshot_model_builder_uses_frozen_profile(provider_config_dir):
replace_provider_profiles(_document())
create_run_runtime_snapshot(
"run-request-a", model="research-model", provider="lab-openai"
)
model = MagicMock()
with patch(
"EvoScientist.llm.models.get_profile_chat_model", return_value=model
) as factory:
resolved = get_snapshot_chat_model("run-request-a")
assert resolved is model
profile, provider_model = factory.call_args.args
assert profile.api_key == "first-secret"
assert provider_model.model_id == "vendor/research-1"