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:
@@ -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
|
||||
|
||||
@@ -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}."
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user