feat: add scoped model runtime configuration
Build / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
Build / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
Introduce provider, model, and invocation contracts with encrypted configuration persistence. Add web runtime fencing, route fallback, recovery middleware, workspace scoping, and comprehensive tests.
This commit is contained in:
+164
-25
@@ -308,6 +308,8 @@ def _inject_subagent_middleware(
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
ContextOverflowMapperMiddleware,
|
||||
ErrorNormalizationMiddleware,
|
||||
RecoverableMeteringMiddleware,
|
||||
RecoverableToolEffectMiddleware,
|
||||
RepetitiveToolCallGuardMiddleware,
|
||||
ToolErrorHandlerMiddleware,
|
||||
ToolProtocolGuardMiddleware,
|
||||
@@ -353,6 +355,8 @@ def _inject_subagent_middleware(
|
||||
# them into a non-dataclass envelope wrapper before
|
||||
# anything downstream sees them.
|
||||
ErrorNormalizationMiddleware(),
|
||||
RecoverableMeteringMiddleware(),
|
||||
RecoverableToolEffectMiddleware(),
|
||||
RepetitiveToolCallGuardMiddleware(
|
||||
threshold=repetitive_tool_call_threshold,
|
||||
max_consecutive_errors=max_consecutive_tool_errors,
|
||||
@@ -399,6 +403,42 @@ def _ensure_general_purpose_subagent(subs: list[dict]) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _apply_budgeted_skill_context(kwargs: dict, backend) -> dict:
|
||||
"""Replace DeepAgents' full-catalog skill prompts with bounded prompts."""
|
||||
|
||||
from .middleware import BudgetedSkillsMiddleware
|
||||
|
||||
updated = dict(kwargs)
|
||||
middleware = list(updated.get("middleware") or ())
|
||||
if not any(isinstance(item, BudgetedSkillsMiddleware) for item in middleware):
|
||||
middleware.append(
|
||||
BudgetedSkillsMiddleware(
|
||||
backend=backend,
|
||||
sources=list(DEFAULT_SKILL_SOURCES),
|
||||
)
|
||||
)
|
||||
updated["middleware"] = middleware
|
||||
updated["skills"] = None
|
||||
|
||||
subagents = []
|
||||
for spec in updated.get("subagents") or ():
|
||||
if not isinstance(spec, dict) or not spec.get("skills"):
|
||||
subagents.append(spec)
|
||||
continue
|
||||
child = dict(spec)
|
||||
raw_sources = child["skills"]
|
||||
sources = [raw_sources] if isinstance(raw_sources, str) else list(raw_sources)
|
||||
child["skills"] = None
|
||||
child_middleware = list(child.get("middleware") or ())
|
||||
child_middleware.append(
|
||||
BudgetedSkillsMiddleware(backend=backend, sources=sources)
|
||||
)
|
||||
child["middleware"] = child_middleware
|
||||
subagents.append(child)
|
||||
updated["subagents"] = subagents
|
||||
return updated
|
||||
|
||||
|
||||
def _maybe_swap_async_subagents(
|
||||
subs: list, middleware: list | None = None, *, cfg=None
|
||||
) -> list:
|
||||
@@ -459,7 +499,9 @@ def _maybe_swap_async_subagents(
|
||||
|
||||
from deepagents import AsyncSubAgent
|
||||
|
||||
port = int(getattr(cfg, "langgraph_dev_port", 6174))
|
||||
from .langgraph_dev.sdk import configured_langgraph_dev_url
|
||||
|
||||
runtime_url = configured_langgraph_dev_url()
|
||||
out = []
|
||||
agent_specs: dict[str, AsyncSubAgent] = {}
|
||||
# MCP tools routed to async sub-agents (via ``expose_to: <name>`` in
|
||||
@@ -474,7 +516,7 @@ def _maybe_swap_async_subagents(
|
||||
name=name,
|
||||
description=async_specs[name],
|
||||
graph_id=name,
|
||||
url=f"http://localhost:{port}",
|
||||
url=runtime_url,
|
||||
)
|
||||
agent_specs[name] = spec
|
||||
out.append(spec)
|
||||
@@ -619,8 +661,8 @@ def load_mcp_and_build_kwargs(
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def _get_default_backend():
|
||||
"""Build the default composite backend from current paths."""
|
||||
def _get_legacy_backend():
|
||||
"""Build the deployment-root backend used outside Web full deploy."""
|
||||
from deepagents.backends import CompositeBackend
|
||||
|
||||
from .backends import (
|
||||
@@ -662,6 +704,20 @@ def _get_default_backend():
|
||||
)
|
||||
|
||||
|
||||
def _get_default_backend():
|
||||
"""Use Origin's conversation-scoped backend for Web full deploy."""
|
||||
if os.environ.get("EVOSCIENTIST_DEPLOY_MODE", "").lower() != "full":
|
||||
return _get_legacy_backend()
|
||||
from .workspace_scope import create_workspace_backend_factory
|
||||
|
||||
cfg = _ensure_config()
|
||||
return create_workspace_backend_factory(
|
||||
_get_legacy_backend,
|
||||
dangerous=cfg.dangerous_mode,
|
||||
allow_unscoped_legacy=False,
|
||||
)
|
||||
|
||||
|
||||
def _get_default_middleware(
|
||||
*,
|
||||
for_async_subagent: bool = False,
|
||||
@@ -674,6 +730,11 @@ def _get_default_middleware(
|
||||
memory_max_inline_profile_chars: int | None = None,
|
||||
enable_background_execution: bool = True,
|
||||
enable_legacy_model_fallback: bool = True,
|
||||
tool_selector_model=None,
|
||||
include_configurable_model: bool = True,
|
||||
enable_scheduler: bool | None = None,
|
||||
enable_memory_workers: bool | None = None,
|
||||
install_subagent_guard: bool = False,
|
||||
):
|
||||
"""Build the default middleware list.
|
||||
|
||||
@@ -698,8 +759,11 @@ def _get_default_middleware(
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
ConfigurableModelMiddleware,
|
||||
ContextOverflowMapperMiddleware,
|
||||
DisableSubagentToolMiddleware,
|
||||
ErrorNormalizationMiddleware,
|
||||
ModelFallbackMiddleware,
|
||||
RecoverableMeteringMiddleware,
|
||||
RecoverableToolEffectMiddleware,
|
||||
RepetitiveToolCallGuardMiddleware,
|
||||
ToolErrorHandlerMiddleware,
|
||||
ToolProtocolGuardMiddleware,
|
||||
@@ -761,26 +825,30 @@ def _get_default_middleware(
|
||||
# keep the main model (they do real work, not a one-off helper call).
|
||||
# context_editing stays on the main model — its model only sizes the
|
||||
# context-window trigger for the main agent's own history.
|
||||
if for_async_subagent:
|
||||
tool_selector_model = model
|
||||
if tool_selector_model is not None:
|
||||
resolved_selector_model = tool_selector_model
|
||||
elif for_async_subagent:
|
||||
resolved_selector_model = model
|
||||
elif chat_model is None:
|
||||
tool_selector_model = _ensure_auxiliary_chat_model()
|
||||
resolved_selector_model = _ensure_auxiliary_chat_model()
|
||||
else:
|
||||
aux_model = cfg.auxiliary_model or cfg.model
|
||||
aux_provider = cfg.auxiliary_provider or cfg.provider
|
||||
if (aux_model, aux_provider) == (cfg.model, cfg.provider):
|
||||
tool_selector_model = model
|
||||
resolved_selector_model = model
|
||||
else:
|
||||
from .llm import get_chat_model
|
||||
|
||||
tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider)
|
||||
resolved_selector_model = get_chat_model(
|
||||
model=aux_model, provider=aux_provider
|
||||
)
|
||||
selector_middlewares = create_tool_selector_middleware(
|
||||
**(
|
||||
{"threshold": tool_selector_threshold}
|
||||
if tool_selector_threshold is not None
|
||||
else {}
|
||||
),
|
||||
model=tool_selector_model,
|
||||
model=resolved_selector_model,
|
||||
track_stream_selection=not for_async_subagent,
|
||||
)
|
||||
mw = [
|
||||
@@ -789,7 +857,8 @@ def _get_default_middleware(
|
||||
# middlewares) and normalizes them into a non-dataclass
|
||||
# envelope wrapper before anything downstream sees them.
|
||||
ErrorNormalizationMiddleware(),
|
||||
ConfigurableModelMiddleware(),
|
||||
RecoverableMeteringMiddleware(),
|
||||
RecoverableToolEffectMiddleware(),
|
||||
create_context_editing_middleware(model),
|
||||
*([ModelFallbackMiddleware()] if enable_legacy_model_fallback else []),
|
||||
RepetitiveToolCallGuardMiddleware(
|
||||
@@ -807,12 +876,18 @@ def _get_default_middleware(
|
||||
max_result_chars=cfg.code_interpreter_max_result_chars,
|
||||
),
|
||||
]
|
||||
if cfg.enable_scheduler and not for_async_subagent:
|
||||
if include_configurable_model:
|
||||
mw.insert(1, ConfigurableModelMiddleware())
|
||||
if enable_scheduler is None:
|
||||
enable_scheduler = bool(cfg.enable_scheduler)
|
||||
if enable_scheduler and not for_async_subagent:
|
||||
mw.append(create_scheduler_middleware())
|
||||
mw.append(create_runtime_context_middleware())
|
||||
if memory_controls.memory_enabled:
|
||||
mw.append(memory_middleware)
|
||||
if memory_controls.worker_needed(worker_target):
|
||||
if enable_memory_workers is not False and memory_controls.worker_needed(
|
||||
worker_target
|
||||
):
|
||||
mw.append(
|
||||
create_memory_lifecycle_middleware(
|
||||
memory_dir,
|
||||
@@ -837,6 +912,9 @@ def _get_default_middleware(
|
||||
|
||||
mw.append(BackgroundExecutionMiddleware())
|
||||
|
||||
if install_subagent_guard:
|
||||
mw.append(DisableSubagentToolMiddleware())
|
||||
|
||||
return mw
|
||||
|
||||
|
||||
@@ -898,6 +976,7 @@ def _get_default_agent():
|
||||
mw,
|
||||
workspace_dir=str(_paths_mod.WORKSPACE_ROOT),
|
||||
)
|
||||
kwargs = _apply_budgeted_skill_context(kwargs, be)
|
||||
|
||||
_EvoScientist_agent = create_deep_agent(
|
||||
**kwargs,
|
||||
@@ -938,6 +1017,8 @@ def create_cli_agent(
|
||||
enable_background_execution: bool = True,
|
||||
main_agent_outer_middlewares: Sequence[AgentMiddleware] | None = None,
|
||||
main_agent_route_middleware: AgentMiddleware | None = None,
|
||||
execution_profile=None,
|
||||
agent_model_set=None,
|
||||
) -> "CompiledStateGraph":
|
||||
"""Create agent with checkpointer for CLI multi-turn support.
|
||||
|
||||
@@ -1004,6 +1085,24 @@ def create_cli_agent(
|
||||
cfg = _ensure_config(config)
|
||||
chat_model = None
|
||||
|
||||
profile = execution_profile
|
||||
if agent_model_set is not None:
|
||||
chat_model = agent_model_set.main_agent
|
||||
if profile is not None:
|
||||
import copy
|
||||
|
||||
cfg = copy.copy(cfg)
|
||||
cfg.enable_async_subagents = bool(profile.async_subagents)
|
||||
cfg.enable_scheduler = bool(profile.scheduler)
|
||||
cfg.memory_workers_enabled = bool(profile.memory_workers)
|
||||
cfg.enable_ask_user = False
|
||||
cfg.auto_mode = True
|
||||
cfg.auto_approve = True
|
||||
enable_subagents = bool(enable_subagents and profile.subagents)
|
||||
enable_background_execution = bool(
|
||||
enable_background_execution and profile.background_execution
|
||||
)
|
||||
|
||||
if checkpointer is None:
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
|
||||
@@ -1056,15 +1155,53 @@ def create_cli_agent(
|
||||
# Delegate middleware construction to the single source of truth so the
|
||||
# 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,
|
||||
memory_dir=_mem_dir,
|
||||
cfg=cfg,
|
||||
chat_model=chat_model,
|
||||
tool_selector_threshold=tool_selector_threshold,
|
||||
memory_max_inline_profile_chars=memory_max_inline_profile_chars,
|
||||
enable_background_execution=enable_background_execution,
|
||||
enable_legacy_model_fallback=main_agent_route_middleware is None,
|
||||
mw: list[AgentMiddleware] = list(
|
||||
_get_default_middleware(
|
||||
workspace_dir=workspace_dir,
|
||||
memory_dir=_mem_dir,
|
||||
cfg=cfg,
|
||||
chat_model=chat_model,
|
||||
tool_selector_threshold=tool_selector_threshold,
|
||||
memory_max_inline_profile_chars=memory_max_inline_profile_chars,
|
||||
enable_background_execution=enable_background_execution,
|
||||
enable_legacy_model_fallback=(
|
||||
main_agent_route_middleware is None
|
||||
and not (
|
||||
profile is not None
|
||||
and getattr(profile, "name", "") in {"web_v1", "web_v3"}
|
||||
)
|
||||
),
|
||||
tool_selector_model=(
|
||||
agent_model_set.tool_selector if agent_model_set is not None else None
|
||||
),
|
||||
include_configurable_model=(
|
||||
bool(profile.configurable_model_override)
|
||||
if profile is not None
|
||||
else True
|
||||
),
|
||||
enable_scheduler=(bool(profile.scheduler) if profile is not None else None),
|
||||
enable_memory_workers=(
|
||||
bool(profile.memory_workers) if profile is not None else None
|
||||
),
|
||||
install_subagent_guard=(profile is not None and not profile.subagents),
|
||||
)
|
||||
)
|
||||
from .middleware import ProviderContextMediaMiddleware
|
||||
|
||||
# Keep assistant-generated binary output out of both provider history and
|
||||
# future checkpoints. The middleware persists media through the same
|
||||
# workspace backend before replacing it with a content-addressed reference.
|
||||
error_index = next(
|
||||
(
|
||||
index
|
||||
for index, middleware in enumerate(mw)
|
||||
if getattr(middleware, "name", "") == "error_normalization"
|
||||
),
|
||||
None,
|
||||
)
|
||||
mw.insert(
|
||||
(error_index + 1) if error_index is not None else 0,
|
||||
ProviderContextMediaMiddleware(be),
|
||||
)
|
||||
if main_agent_route_middleware is not None:
|
||||
configurable_index = next(
|
||||
@@ -1075,9 +1212,10 @@ def create_cli_agent(
|
||||
),
|
||||
None,
|
||||
)
|
||||
if configurable_index is None:
|
||||
raise RuntimeError("ConfigurableModelMiddleware route slot is unavailable")
|
||||
mw.insert(configurable_index + 1, main_agent_route_middleware)
|
||||
mw.insert(
|
||||
(configurable_index + 1) if configurable_index is not None else 1,
|
||||
main_agent_route_middleware,
|
||||
)
|
||||
if main_agent_outer_middlewares:
|
||||
mw = [*main_agent_outer_middlewares, *mw]
|
||||
|
||||
@@ -1106,6 +1244,7 @@ def create_cli_agent(
|
||||
)
|
||||
if not enable_subagents:
|
||||
kwargs = {**kwargs, "subagents": []}
|
||||
kwargs = _apply_budgeted_skill_context(kwargs, be)
|
||||
|
||||
return create_deep_agent(
|
||||
**kwargs,
|
||||
|
||||
@@ -30,6 +30,7 @@ _EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"MODELS": (".llm", "MODELS"),
|
||||
"list_models": (".llm", "list_models"),
|
||||
"DEFAULT_MODEL": (".llm", "DEFAULT_MODEL"),
|
||||
"EvoModelRuntime": (".llm.runtime", "EvoModelRuntime"),
|
||||
# Prompts
|
||||
"get_system_prompt": (".prompts", "get_system_prompt"),
|
||||
# Tools
|
||||
|
||||
@@ -20,6 +20,7 @@ from deepagents.backends.protocol import (
|
||||
LsResult,
|
||||
WriteResult,
|
||||
)
|
||||
from filelock import FileLock
|
||||
|
||||
from . import paths
|
||||
|
||||
@@ -835,6 +836,15 @@ class MemoryFilesystemBackend(FilesystemBackend):
|
||||
"/memories/profile/... files. Use memory tools for observations."
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
root_dir: str | Path | None = None,
|
||||
virtual_mode: bool | None = None,
|
||||
max_file_size_mb: int = 10,
|
||||
) -> None:
|
||||
super().__init__(root_dir, virtual_mode, max_file_size_mb)
|
||||
self._profile_write_lock = FileLock(str(self.cwd / ".profile-write.lock"))
|
||||
|
||||
@staticmethod
|
||||
def _is_profile_path(file_path: str) -> bool:
|
||||
normalized = posixpath.normpath("/" + file_path.strip().lstrip("/"))
|
||||
@@ -852,7 +862,8 @@ class MemoryFilesystemBackend(FilesystemBackend):
|
||||
) -> EditResult:
|
||||
if not self._is_profile_path(file_path):
|
||||
return EditResult(error=self._RAW_EDIT_ERROR)
|
||||
return super().edit(file_path, old_string, new_string, replace_all)
|
||||
with self._profile_write_lock:
|
||||
return super().edit(file_path, old_string, new_string, replace_all)
|
||||
|
||||
def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]:
|
||||
return [
|
||||
|
||||
@@ -90,11 +90,11 @@ def _step_langgraph_dev_port(config: EvoScientistConfig) -> int:
|
||||
"""
|
||||
if not getattr(config, "enable_async_subagents", True):
|
||||
# User has async disabled — port is irrelevant, no prompt.
|
||||
return getattr(config, "langgraph_dev_port", 6174)
|
||||
return getattr(config, "langgraph_dev_port", 3076)
|
||||
|
||||
from ...langgraph_dev.manager import _is_port_occupied, is_langgraph_dev_running
|
||||
|
||||
current_port = getattr(config, "langgraph_dev_port", 6174)
|
||||
current_port = getattr(config, "langgraph_dev_port", 3076)
|
||||
current_occupied = _is_port_occupied(current_port)
|
||||
if current_occupied and is_langgraph_dev_running(port=current_port):
|
||||
# Another EvoSci shell is already serving on this port — reuse, don't
|
||||
@@ -174,7 +174,7 @@ def _step_webui_port(config: EvoScientistConfig) -> int:
|
||||
from ...langgraph_dev.manager import _is_port_occupied
|
||||
|
||||
current_port = getattr(config, "webui_port", 4716)
|
||||
backend_port = getattr(config, "langgraph_dev_port", 6174)
|
||||
backend_port = getattr(config, "langgraph_dev_port", 3076)
|
||||
occupied = _is_port_occupied(current_port)
|
||||
conflicts_backend = current_port == backend_port
|
||||
|
||||
|
||||
@@ -156,6 +156,7 @@ class EvoScientistConfig:
|
||||
nvidia_api_key: NVIDIA API key for NVIDIA models.
|
||||
google_api_key: Google API key for Gemini models.
|
||||
tavily_api_key: Tavily API key for web search.
|
||||
semantic_scholar_api_key: Semantic Scholar API key for paper-navigator.
|
||||
provider: Default LLM provider ('anthropic', 'openai', 'google-genai', or 'nvidia').
|
||||
model: Default model name (short name or full ID).
|
||||
auxiliary_provider: Provider for auxiliary_model (empty = use main provider).
|
||||
@@ -189,6 +190,7 @@ class EvoScientistConfig:
|
||||
custom_anthropic_base_url: str = ""
|
||||
ollama_base_url: str = ""
|
||||
tavily_api_key: str = ""
|
||||
semantic_scholar_api_key: str = ""
|
||||
|
||||
# LLM Settings
|
||||
provider: str = "anthropic"
|
||||
@@ -212,15 +214,12 @@ class EvoScientistConfig:
|
||||
# synchronous sub-agents (planner / research / code / debug).
|
||||
enable_async_subagents: bool = True
|
||||
|
||||
# Port for the auto-started langgraph dev subprocess. 6174 is Kaprekar's
|
||||
# constant — a memorable EvoScientist-themed default that avoids collisions
|
||||
# with common dev ports (3000/5000/8000/8080) and the langgraph CLI default
|
||||
# 2024. Override if it conflicts with another local service.
|
||||
langgraph_dev_port: int = 6174
|
||||
# Port for the auto-started langgraph dev subprocess. Keep this aligned with
|
||||
# the Ai4Sci-Web Gateway's recoverable runtime URL.
|
||||
langgraph_dev_port: int = 3076
|
||||
|
||||
# Port for the WebUI front-end (Next.js server from @evoscientist/webui),
|
||||
# used only when ui_backend == "webui". 4716 is 6174 reversed — a memorable
|
||||
# pairing with the langgraph dev port that it connects to. The backend keeps
|
||||
# used only when ui_backend == "webui". The backend keeps
|
||||
# its own port (langgraph_dev_port); this is just the browser server.
|
||||
webui_port: int = 4716
|
||||
|
||||
@@ -256,11 +255,10 @@ class EvoScientistConfig:
|
||||
# (sessions.db), ContextEditingMiddleware (window management), and
|
||||
# EvoMemoryMiddleware (cross-turn memory).
|
||||
#
|
||||
# 1,000,000 is "effectively unlimited" — typical research turns use
|
||||
# 200-1000 steps; reaching 1M would cost ~$10K in tokens, by which point
|
||||
# rate limits, context overflow, or API quota errors would trip first.
|
||||
# Lower (e.g., 5000) if you want a tighter safety net against runaway loops.
|
||||
recursion_limit: int = 1_000_000
|
||||
# Typical research turns use 200-1000 steps. 5,000 leaves room for
|
||||
# legitimate long-running work while still stopping a runaway graph.
|
||||
# This is a control-flow limit, not a model-call or cost budget.
|
||||
recursion_limit: int = 5_000
|
||||
|
||||
# Number of consecutive model rounds with the same structured tool name and
|
||||
# arguments that activates provider-facing loop repair. Set 0 to disable.
|
||||
@@ -307,9 +305,6 @@ class EvoScientistConfig:
|
||||
# a deploy-style langgraph server instead of the in-terminal CLI/TUI.
|
||||
ui_backend: Literal["cli", "tui", "webui"] = "tui"
|
||||
log_level: str = "warning"
|
||||
# Empty means use the provider/model default. A non-empty value is an
|
||||
# explicit user override exported as EVOSCIENTIST_REASONING_EFFORT.
|
||||
reasoning_effort: str = ""
|
||||
# Anthropic prompt caching for OpenRouter anthropic/* models. Opt out if
|
||||
# cache-write costs outweigh the benefit for a workflow.
|
||||
openrouter_anthropic_prompt_cache: bool = True
|
||||
@@ -466,9 +461,6 @@ class EvoScientistConfig:
|
||||
# DM access control policy
|
||||
dm_policy: str = "allowlist"
|
||||
|
||||
# OpenAI API mode - "" = auto, "true" = force Responses, "false" = force Completions
|
||||
use_responses_api: str = ""
|
||||
|
||||
# ccproxy
|
||||
ccproxy_port: int = 8000
|
||||
|
||||
@@ -803,6 +795,7 @@ _ENV_MAPPINGS = {
|
||||
"custom_anthropic_base_url": "CUSTOM_ANTHROPIC_BASE_URL",
|
||||
"ollama_base_url": "OLLAMA_BASE_URL",
|
||||
"tavily_api_key": "TAVILY_API_KEY",
|
||||
"semantic_scholar_api_key": "S2_API_KEY",
|
||||
"default_mode": "EVOSCIENTIST_DEFAULT_MODE",
|
||||
"default_workdir": "EVOSCIENTIST_WORKSPACE_DIR",
|
||||
"ui_backend": "EVOSCIENTIST_UI_BACKEND",
|
||||
@@ -810,7 +803,6 @@ _ENV_MAPPINGS = {
|
||||
"model_fallbacks": "EVOSCIENTIST_MODEL_FALLBACKS",
|
||||
"auxiliary_provider": "EVOSCIENTIST_AUXILIARY_PROVIDER",
|
||||
"auxiliary_model": "EVOSCIENTIST_AUXILIARY_MODEL",
|
||||
"reasoning_effort": "EVOSCIENTIST_REASONING_EFFORT",
|
||||
"openrouter_anthropic_prompt_cache": (
|
||||
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"
|
||||
),
|
||||
@@ -820,7 +812,6 @@ _ENV_MAPPINGS = {
|
||||
"dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE",
|
||||
"channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING",
|
||||
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
|
||||
"use_responses_api": "EVOSCIENTIST_USE_RESPONSES_API",
|
||||
"checkpoint_keep_per_thread": "EVOSCIENTIST_CHECKPOINT_KEEP_PER_THREAD",
|
||||
"enable_async_subagents": "EVOSCIENTIST_ENABLE_ASYNC_SUBAGENTS",
|
||||
"langgraph_dev_port": "EVOSCIENTIST_LANGGRAPH_DEV_PORT",
|
||||
@@ -950,8 +941,8 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
|
||||
os.environ["OLLAMA_BASE_URL"] = config.ollama_base_url
|
||||
if config.tavily_api_key and not os.environ.get("TAVILY_API_KEY"):
|
||||
os.environ["TAVILY_API_KEY"] = config.tavily_api_key
|
||||
if config.reasoning_effort and not os.environ.get("EVOSCIENTIST_REASONING_EFFORT"):
|
||||
os.environ["EVOSCIENTIST_REASONING_EFFORT"] = config.reasoning_effort
|
||||
if config.semantic_scholar_api_key and not os.environ.get("S2_API_KEY"):
|
||||
os.environ["S2_API_KEY"] = config.semantic_scholar_api_key
|
||||
if config.openrouter_http_referer and not os.environ.get(
|
||||
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER"
|
||||
):
|
||||
@@ -982,7 +973,3 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
|
||||
os.environ["EVOSCIENTIST_DANGEROUS_MODE"] = "true"
|
||||
else:
|
||||
os.environ.pop("EVOSCIENTIST_DANGEROUS_MODE", None)
|
||||
if config.use_responses_api and not os.environ.get(
|
||||
"EVOSCIENTIST_USE_RESPONSES_API"
|
||||
):
|
||||
os.environ["EVOSCIENTIST_USE_RESPONSES_API"] = config.use_responses_api
|
||||
|
||||
@@ -44,7 +44,7 @@ def deploy(
|
||||
port: int | None = typer.Option(
|
||||
None,
|
||||
"--port",
|
||||
help="Port for langgraph dev (default: config.langgraph_dev_port = 6174)",
|
||||
help="Port for langgraph dev (default: config.langgraph_dev_port = 3076)",
|
||||
),
|
||||
tunnel: bool = typer.Option(
|
||||
False,
|
||||
|
||||
@@ -141,6 +141,7 @@ class BackgroundRun:
|
||||
run_id: str
|
||||
assistant_id: str
|
||||
metadata: Mapping[str, str]
|
||||
configurable: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -278,6 +279,7 @@ def _background_run_handle(
|
||||
run_id: str,
|
||||
payload: BackgroundRunPayload,
|
||||
) -> BackgroundRun:
|
||||
configurable = payload["config"].get("configurable")
|
||||
return BackgroundRun(
|
||||
name=request.name,
|
||||
url=url,
|
||||
@@ -286,6 +288,9 @@ def _background_run_handle(
|
||||
run_id=run_id,
|
||||
assistant_id=payload["assistant_id"],
|
||||
metadata=dict(payload["metadata"]),
|
||||
configurable=(
|
||||
dict(configurable) if isinstance(configurable, Mapping) else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -528,6 +533,7 @@ def spawn_background_run_status_thread(
|
||||
"graph_id": run.graph_id,
|
||||
"assistant_id": run.assistant_id,
|
||||
"metadata": run.metadata,
|
||||
"configurable": run.configurable,
|
||||
"name": run.name,
|
||||
"headers": headers,
|
||||
"hooks": hooks,
|
||||
@@ -547,6 +553,7 @@ def watch_background_run_sync(
|
||||
graph_id: str = "",
|
||||
assistant_id: str = "",
|
||||
metadata: Mapping[str, str] | None = None,
|
||||
configurable: Mapping[str, object] | None = None,
|
||||
name: str = "background run",
|
||||
headers: Mapping[str, str] | None = None,
|
||||
hooks: BackgroundRunHooks | None = None,
|
||||
@@ -565,6 +572,7 @@ def watch_background_run_sync(
|
||||
run_id=run_id,
|
||||
assistant_id=assistant_id,
|
||||
metadata=dict(metadata or {}),
|
||||
configurable=dict(configurable or {}),
|
||||
)
|
||||
failures = 0
|
||||
confirmed_finished = False
|
||||
@@ -632,6 +640,7 @@ def spawn_background_run_status_task(
|
||||
graph_id=run.graph_id,
|
||||
assistant_id=run.assistant_id,
|
||||
metadata=run.metadata,
|
||||
configurable=run.configurable,
|
||||
name=run.name,
|
||||
hooks=hooks,
|
||||
watcher_config=watcher_config,
|
||||
@@ -650,6 +659,7 @@ async def awatch_background_run(
|
||||
graph_id: str = "",
|
||||
assistant_id: str = "",
|
||||
metadata: Mapping[str, str] | None = None,
|
||||
configurable: Mapping[str, object] | None = None,
|
||||
name: str = "background run",
|
||||
hooks: BackgroundRunHooks | None = None,
|
||||
watcher_config: BackgroundRunWatcherConfig | None = None,
|
||||
@@ -665,6 +675,7 @@ async def awatch_background_run(
|
||||
run_id=run_id,
|
||||
assistant_id=assistant_id,
|
||||
metadata=dict(metadata or {}),
|
||||
configurable=dict(configurable or {}),
|
||||
)
|
||||
failures = 0
|
||||
confirmed_finished = False
|
||||
|
||||
@@ -23,6 +23,13 @@ memory.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import secrets
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Any
|
||||
from uuid import UUID
|
||||
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
@@ -32,6 +39,23 @@ from starlette.routing import Route
|
||||
from EvoScientist.config import get_effective_config
|
||||
from EvoScientist.llm.models import list_model_picker_entries
|
||||
|
||||
_recoverable_run_lock = asyncio.Lock()
|
||||
|
||||
|
||||
def _load_scope_service_token() -> str:
|
||||
configured = (
|
||||
os.getenv("EVOSCIENTIST_BACKEND_SERVICE_TOKEN", "").strip()
|
||||
or os.getenv("AI4SCI_EVO_RUNTIME_GRANT_SECRET", "").strip()
|
||||
)
|
||||
if configured:
|
||||
return configured
|
||||
from EvoScientist.scope_registry import get_scope_service_token
|
||||
|
||||
return get_scope_service_token()
|
||||
|
||||
|
||||
_SCOPE_SERVICE_TOKEN = _load_scope_service_token()
|
||||
|
||||
|
||||
async def get_models(_request: Request) -> JSONResponse:
|
||||
"""Return the model registry as ``{entries, default}``.
|
||||
@@ -75,8 +99,330 @@ async def get_models(_request: Request) -> JSONResponse:
|
||||
)
|
||||
|
||||
|
||||
async def recoverable_run_capabilities(_request: Request) -> JSONResponse:
|
||||
"""Capabilities required by Ai4Sci's durable dispatch outbox."""
|
||||
|
||||
return JSONResponse(
|
||||
{
|
||||
"version": 1,
|
||||
"deterministic_run_id": True,
|
||||
"stream_resumable": True,
|
||||
"durability_sync": True,
|
||||
"multitask_enqueue": True,
|
||||
"interrupt_resume": True,
|
||||
"pending_interrupt_state": True,
|
||||
"workspace_scope_v1": os.getenv("EVOSCIENTIST_DEPLOY_MODE", "").lower() == "full",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _scope_service_authorized(request: Request) -> JSONResponse | None:
|
||||
if not _SCOPE_SERVICE_TOKEN:
|
||||
return JSONResponse({"code": "WORKSPACE_SERVICE_UNAVAILABLE"}, status_code=503)
|
||||
header = request.headers.get("authorization", "")
|
||||
if not header.startswith("Bearer ") or not secrets.compare_digest(
|
||||
header[7:], _SCOPE_SERVICE_TOKEN
|
||||
):
|
||||
return JSONResponse({"code": "UNAUTHORIZED"}, status_code=401)
|
||||
return None
|
||||
|
||||
|
||||
def _scope_payload(record: Any) -> dict[str, Any]:
|
||||
return {
|
||||
"deployment_id": record.deployment_id,
|
||||
"scope_id": record.scope_id,
|
||||
"primary_thread_id": record.primary_thread_id,
|
||||
"primary_owner_id": record.primary_owner_id,
|
||||
"state": record.state,
|
||||
"revision": record.revision,
|
||||
}
|
||||
|
||||
|
||||
def _run_payload(run: Any) -> dict[str, Any]:
|
||||
return {
|
||||
"run_request_id": run.run_request_id,
|
||||
"turn_id": run.turn_id,
|
||||
"interrupt_key": run.interrupt_key,
|
||||
"request_hash": run.request_hash,
|
||||
"run_owner_id": run.run_owner_id,
|
||||
"run_id": run.run_id,
|
||||
"state": run.state,
|
||||
}
|
||||
|
||||
|
||||
def _registry_call(method: str, *args: Any, **kwargs: Any) -> Any:
|
||||
from EvoScientist.scope_registry import get_scope_registry
|
||||
from EvoScientist.workspace_scope import current_deployment_id
|
||||
|
||||
return getattr(get_scope_registry(), method)(current_deployment_id(), *args, **kwargs)
|
||||
|
||||
|
||||
def _provision_scope(thread_id: str) -> Any:
|
||||
from EvoScientist.workspace_scope import (
|
||||
current_deployment_id,
|
||||
provision_conversation_scope,
|
||||
)
|
||||
|
||||
return provision_conversation_scope(thread_id, deployment_id=current_deployment_id())
|
||||
|
||||
|
||||
async def provision_workspace_scope(request: Request) -> JSONResponse:
|
||||
if denied := _scope_service_authorized(request):
|
||||
return denied
|
||||
try:
|
||||
payload = await request.json()
|
||||
except json.JSONDecodeError:
|
||||
payload = None
|
||||
if not isinstance(payload, dict) or not isinstance(payload.get("thread_id"), str):
|
||||
return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400)
|
||||
try:
|
||||
record = await asyncio.to_thread(_provision_scope, payload["thread_id"])
|
||||
except Exception as exc:
|
||||
return JSONResponse({"code": "WORKSPACE_SCOPE_CONFLICT", "message": str(exc)}, status_code=409)
|
||||
return JSONResponse(_scope_payload(record), status_code=201)
|
||||
|
||||
|
||||
async def get_workspace_scope(request: Request) -> JSONResponse:
|
||||
if denied := _scope_service_authorized(request):
|
||||
return denied
|
||||
try:
|
||||
record = await asyncio.to_thread(
|
||||
_registry_call, "get_by_thread", str(request.path_params["thread_id"])
|
||||
)
|
||||
except Exception as exc:
|
||||
return JSONResponse({"code": "WORKSPACE_SCOPE_NOT_FOUND", "message": str(exc)}, status_code=404)
|
||||
return JSONResponse(_scope_payload(record))
|
||||
|
||||
|
||||
async def reserve_workspace_run(request: Request) -> JSONResponse:
|
||||
if denied := _scope_service_authorized(request):
|
||||
return denied
|
||||
try:
|
||||
payload = await request.json()
|
||||
except json.JSONDecodeError:
|
||||
payload = None
|
||||
if not isinstance(payload, dict):
|
||||
return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400)
|
||||
try:
|
||||
run = await asyncio.to_thread(
|
||||
_registry_call,
|
||||
"reserve_run",
|
||||
str(request.path_params["scope_id"]),
|
||||
str(payload["run_request_id"]),
|
||||
str(payload["turn_id"]),
|
||||
str(payload["request_hash"]),
|
||||
interrupt_key=(
|
||||
str(payload["interrupt_key"]) if payload.get("interrupt_key") else None
|
||||
),
|
||||
)
|
||||
except Exception as exc:
|
||||
code = (
|
||||
"INTERRUPT_ALREADY_RESOLVED"
|
||||
if type(exc).__name__ == "ScopeInterruptResolvedError"
|
||||
else "WORKSPACE_RUN_CONFLICT"
|
||||
)
|
||||
return JSONResponse({"code": code, "message": str(exc)}, status_code=409)
|
||||
return JSONResponse(_run_payload(run), status_code=201)
|
||||
|
||||
|
||||
async def bind_workspace_run(request: Request) -> JSONResponse:
|
||||
if denied := _scope_service_authorized(request):
|
||||
return denied
|
||||
try:
|
||||
payload = await request.json()
|
||||
except json.JSONDecodeError:
|
||||
payload = None
|
||||
if not isinstance(payload, dict) or not isinstance(payload.get("run_id"), str):
|
||||
return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400)
|
||||
try:
|
||||
run = await asyncio.to_thread(
|
||||
_registry_call,
|
||||
"bind_run",
|
||||
str(request.path_params["scope_id"]),
|
||||
str(request.path_params["run_request_id"]),
|
||||
payload["run_id"],
|
||||
)
|
||||
except Exception as exc:
|
||||
return JSONResponse({"code": "WORKSPACE_RUN_CONFLICT", "message": str(exc)}, status_code=409)
|
||||
return JSONResponse(_run_payload(run))
|
||||
|
||||
|
||||
def _materialize_target(scope_id: str, raw_path: str) -> Path:
|
||||
from EvoScientist.workspace_scope import conversation_files_dir
|
||||
|
||||
path = PurePosixPath(raw_path.replace("\\", "/"))
|
||||
if path.is_absolute() or not path.parts or path.parts[0] != "uploads":
|
||||
raise ValueError("only uploads/ paths are accepted")
|
||||
if any(part in {"", ".", ".."} for part in path.parts):
|
||||
raise ValueError("invalid upload path")
|
||||
root = conversation_files_dir(scope_id).resolve()
|
||||
target = root.joinpath(*path.parts)
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
target.parent.resolve().relative_to(root)
|
||||
except ValueError as exc:
|
||||
raise ValueError("upload path escapes workspace scope") from exc
|
||||
current = root
|
||||
for part in path.parts[:-1]:
|
||||
current = current / part
|
||||
if current.is_symlink():
|
||||
raise ValueError("symlink parents are rejected")
|
||||
if target.is_symlink():
|
||||
raise ValueError("symlink targets are rejected")
|
||||
return target
|
||||
|
||||
|
||||
async def materialize_workspace_file(request: Request) -> JSONResponse:
|
||||
if denied := _scope_service_authorized(request):
|
||||
return denied
|
||||
scope_id = str(request.path_params["scope_id"])
|
||||
try:
|
||||
await asyncio.to_thread(_registry_call, "get", scope_id)
|
||||
target = await asyncio.to_thread(
|
||||
_materialize_target, scope_id, str(request.path_params["path"])
|
||||
)
|
||||
except Exception as exc:
|
||||
return JSONResponse({"code": "WORKSPACE_PATH_INVALID", "message": str(exc)}, status_code=400)
|
||||
expected_hash = request.headers.get("x-content-sha256", "").lower()
|
||||
expected_size = int(request.headers.get("content-length") or 0)
|
||||
if expected_size > 100 * 1024 * 1024:
|
||||
return JSONResponse({"code": "WORKSPACE_FILE_TOO_LARGE"}, status_code=413)
|
||||
temporary = target.with_name(f".{target.name}.{secrets.token_hex(8)}.tmp")
|
||||
digest = hashlib.sha256()
|
||||
size = 0
|
||||
try:
|
||||
with temporary.open("xb") as handle:
|
||||
async for chunk in request.stream():
|
||||
size += len(chunk)
|
||||
if size > 100 * 1024 * 1024:
|
||||
raise ValueError("workspace file exceeds 100 MiB")
|
||||
digest.update(chunk)
|
||||
handle.write(chunk)
|
||||
handle.flush()
|
||||
os.fsync(handle.fileno())
|
||||
actual_hash = digest.hexdigest()
|
||||
if expected_hash and not secrets.compare_digest(actual_hash, expected_hash):
|
||||
raise ValueError("workspace file hash mismatch")
|
||||
os.replace(temporary, target)
|
||||
except Exception as exc:
|
||||
temporary.unlink(missing_ok=True)
|
||||
return JSONResponse({"code": "WORKSPACE_FILE_INVALID", "message": str(exc)}, status_code=409)
|
||||
return JSONResponse({"virtual_path": str(request.path_params["path"]), "size": size, "sha256": actual_hash})
|
||||
|
||||
|
||||
async def create_recoverable_run(request: Request) -> JSONResponse:
|
||||
"""Create a LangGraph Run with a caller-owned deterministic UUID.
|
||||
|
||||
LangGraph's public create endpoint always generates its own UUID. This
|
||||
adapter performs lookup and insertion while holding the process-wide run
|
||||
creation lock and passes the durable request UUID to ``create_valid_run``.
|
||||
Retrying after a lost HTTP response therefore cannot create another Run.
|
||||
"""
|
||||
|
||||
if request.headers.get("x-auth-scheme") != "langsmith":
|
||||
return JSONResponse({"code": "UNAUTHORIZED"}, status_code=401)
|
||||
value = await request.json()
|
||||
if not isinstance(value, dict):
|
||||
return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400)
|
||||
try:
|
||||
thread_id = str(UUID(str(value["thread_id"])))
|
||||
run_id = UUID(str(value["run_id"]))
|
||||
run_request_id = str(UUID(str(value["run_request_id"])))
|
||||
request_hash = str(value["request_hash"])
|
||||
assistant_id = str(value["assistant_id"])
|
||||
operation = str(value.get("operation") or "start")
|
||||
except (KeyError, TypeError, ValueError):
|
||||
return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400)
|
||||
if (
|
||||
str(run_id) != run_request_id
|
||||
or len(request_hash) != 64
|
||||
or operation not in {"start", "resume"}
|
||||
):
|
||||
return JSONResponse({"code": "INVALID_IDEMPOTENCY_KEY"}, status_code=400)
|
||||
command = value.get("command")
|
||||
if operation == "resume":
|
||||
if (
|
||||
value.get("input") is not None
|
||||
or not isinstance(command, dict)
|
||||
or set(command) != {"resume"}
|
||||
):
|
||||
return JSONResponse({"code": "INVALID_RESUME_REQUEST"}, status_code=400)
|
||||
elif command is not None:
|
||||
return JSONResponse({"code": "INVALID_START_REQUEST"}, status_code=400)
|
||||
|
||||
from langgraph_api.models.run import Runs, create_valid_run
|
||||
from langgraph_api.utils import fetchone
|
||||
from langgraph_runtime.database import connect
|
||||
|
||||
payload = {
|
||||
"assistant_id": assistant_id,
|
||||
"input": value.get("input"),
|
||||
"command": command,
|
||||
"metadata": value.get("metadata") or {},
|
||||
"config": value.get("config") or {},
|
||||
"stream_mode": value.get("stream_mode") or ["messages", "updates", "tasks", "custom"],
|
||||
"stream_resumable": True,
|
||||
"durability": "sync",
|
||||
"multitask_strategy": "enqueue",
|
||||
"if_not_exists": "create",
|
||||
}
|
||||
payload["metadata"] = {
|
||||
**payload["metadata"],
|
||||
"run_request_id": run_request_id,
|
||||
"request_hash": request_hash,
|
||||
}
|
||||
async with _recoverable_run_lock:
|
||||
async with connect() as conn:
|
||||
existing_iter = await Runs.get(conn, run_id, thread_id=UUID(thread_id))
|
||||
try:
|
||||
existing = await fetchone(existing_iter)
|
||||
except Exception as exc:
|
||||
if getattr(exc, "status_code", None) != 404:
|
||||
raise
|
||||
existing = None
|
||||
if existing is not None:
|
||||
metadata = existing.get("metadata") or {}
|
||||
if metadata.get("request_hash") != request_hash:
|
||||
return JSONResponse(
|
||||
{
|
||||
"code": "RUN_REQUEST_CONFLICT",
|
||||
"message": "run_request_id is bound to another request hash",
|
||||
},
|
||||
status_code=409,
|
||||
)
|
||||
return JSONResponse(
|
||||
{"run_id": str(existing["run_id"]), "status": existing["status"], "created": False}
|
||||
)
|
||||
created = await create_valid_run(
|
||||
conn,
|
||||
thread_id,
|
||||
payload,
|
||||
dict(request.headers),
|
||||
run_id=run_id,
|
||||
)
|
||||
return JSONResponse(
|
||||
{"run_id": str(created["run_id"]), "status": created["status"], "created": True},
|
||||
status_code=201,
|
||||
)
|
||||
|
||||
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Route("/api/models", get_models, methods=["GET"]),
|
||||
Route(
|
||||
"/api/ai4sci/recoverable-runs/capabilities",
|
||||
recoverable_run_capabilities,
|
||||
methods=["GET"],
|
||||
),
|
||||
Route(
|
||||
"/api/ai4sci/recoverable-runs/create",
|
||||
create_recoverable_run,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route("/internal/workspace-scopes/provision", provision_workspace_scope, methods=["POST"]),
|
||||
Route("/internal/workspace-scopes/by-thread/{thread_id}", get_workspace_scope, methods=["GET"]),
|
||||
Route("/internal/workspace-scopes/{scope_id}/runs/reserve", reserve_workspace_run, methods=["POST"]),
|
||||
Route("/internal/workspace-scopes/{scope_id}/runs/{run_request_id}", bind_workspace_run, methods=["PATCH"]),
|
||||
Route("/internal/workspace-scopes/{scope_id}/files/{path:path}", materialize_workspace_file, methods=["PUT"]),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -15,7 +15,7 @@
|
||||
"path": "EvoScientist.sessions.create_checkpointer_for_langgraph_api"
|
||||
},
|
||||
"config": {
|
||||
"recursion_limit": 1000000
|
||||
"recursion_limit": 5000
|
||||
},
|
||||
"http": {
|
||||
"app": "EvoScientist.langgraph_dev.http:app"
|
||||
|
||||
@@ -110,11 +110,11 @@ def needs_langgraph_dev(config: EvoScientistConfig) -> bool:
|
||||
_LOCK = threading.RLock()
|
||||
|
||||
|
||||
# Default port (Kaprekar's constant — see config/settings.py for the rationale).
|
||||
# Default port shared with the Ai4Sci-Web recoverable runtime.
|
||||
# Overridable per-call via ``start_langgraph_dev(port=...)`` /
|
||||
# ``ensure_langgraph_dev`` (which reads ``config.langgraph_dev_port``) and the
|
||||
# corresponding url= field on AsyncSubAgent specs.
|
||||
_DEFAULT_PORT = 6174
|
||||
_DEFAULT_PORT = 3076
|
||||
|
||||
|
||||
def _base_url(port: int = _DEFAULT_PORT) -> str:
|
||||
@@ -448,7 +448,7 @@ def _kill_owned_stale_process(port: int) -> bool:
|
||||
|
||||
Why this matters:
|
||||
1. ``net_connections`` may report any process bound to the port,
|
||||
including user-run dev servers that legitimately took 6174.
|
||||
including user-run dev servers that legitimately took 3076.
|
||||
SIGKILL'ing those is a data-loss event.
|
||||
2. Even with PID-file ownership, the OS may have recycled the PID
|
||||
to an unrelated process between sessions (e.g., after a SIGKILL'd
|
||||
@@ -556,7 +556,7 @@ def start_langgraph_dev(
|
||||
Determines where deployed agents' filesystem operations land
|
||||
(``CustomSandboxBackend`` derives its workspace root from cwd via
|
||||
``paths.WORKSPACE_ROOT``). Defaults to ``Path.cwd()``.
|
||||
port: TCP port to bind. Defaults to 6174 (Kaprekar's constant).
|
||||
port: TCP port to bind. Defaults to 3076.
|
||||
file_persistence: When True (default), langgraph dev writes its full
|
||||
``.langgraph_api/`` cache so async-task / Store / scheduler state
|
||||
survives subprocess restarts. Set False to suppress periodic
|
||||
|
||||
@@ -2,14 +2,18 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
|
||||
DEFAULT_LANGGRAPH_DEV_PORT = 6174
|
||||
DEFAULT_LANGGRAPH_DEV_PORT = 3076
|
||||
LANGGRAPH_DEV_AUTH_HEADERS = {"x-auth-scheme": "langsmith"}
|
||||
|
||||
|
||||
def langgraph_dev_url(config: object | None = None, *, port: int | None = None) -> str:
|
||||
"""Return the local langgraph-dev base URL for a config or explicit port."""
|
||||
runtime_url = os.environ.get("LANGGRAPH_SERVER_URL", "").strip().rstrip("/")
|
||||
if port is None and runtime_url:
|
||||
return runtime_url
|
||||
selected_port = (
|
||||
int(port)
|
||||
if port is not None
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
# Model Runtime Layout
|
||||
|
||||
The model runtime has three configuration and execution boundaries.
|
||||
|
||||
| Layer | Source | Owns | Must not own |
|
||||
| --- | --- | --- | --- |
|
||||
| Provider | `configuration/provider.py` | Adapter identity, credentials, endpoints, headers, connection defaults | Model capabilities, model token limits, derived tool transport |
|
||||
| Model | `configuration/model.py` | Provider model ID, capabilities, limits, canonical parameters, access, billing | Credentials, base URL, SDK client options, derived tool transport |
|
||||
| Invocation | `invocation/contract.py` | Immutable API mode, output parameter, tool transport, streaming flag, final SDK parameters | Admin persistence, credentials, routing decisions |
|
||||
|
||||
Supporting modules have narrower responsibilities:
|
||||
|
||||
- `model_config_v4.py` normalizes and persists the Provider + ModelProfile admin
|
||||
contract, then projects it to the stable runtime schema.
|
||||
- `model_config.py` parses and validates the runtime schema. It re-exports the
|
||||
provider and model contracts for compatibility with existing integrations.
|
||||
- `adapter_registry.py` declares provider/model-family support and converts
|
||||
canonical model parameters into provider SDK parameters.
|
||||
- `runtime.py` selects a frozen route, asks its adapter to compile parameters,
|
||||
compiles an `InvocationPlan`, and constructs the provider client from that
|
||||
plan only.
|
||||
|
||||
The call chain is fixed:
|
||||
|
||||
```text
|
||||
V4 Provider + ModelProfile
|
||||
-> normalize and validate
|
||||
-> V3 runtime projection
|
||||
-> select provider endpoint and model profile
|
||||
-> merge canonical model parameters
|
||||
-> provider adapter compilation
|
||||
-> immutable InvocationPlan validation
|
||||
-> provider SDK call
|
||||
```
|
||||
|
||||
Important invariants:
|
||||
|
||||
1. Environment variables may provide secrets, proxy settings, and timeouts;
|
||||
they cannot select an API protocol or rewrite a compiled invocation.
|
||||
2. `tool_call_transport` is not administrator configuration. It is derived as
|
||||
`native` when `capabilities.tools=true`, otherwise `disabled`.
|
||||
3. Exactly one provider output-limit parameter is allowed in a compiled plan:
|
||||
`max_output_tokens`, `max_completion_tokens`, or `max_tokens`.
|
||||
4. Provider-specific parameter names are selected by the adapter. Gateway,
|
||||
frontend, and generic runtime code must not guess them from model names.
|
||||
5. Runtime logs report the final non-secret plan and parameter names. They must
|
||||
never include credentials, authorization headers, or raw secret values.
|
||||
6. Provider input projection removes assistant history that has neither final
|
||||
text nor a tool call. A newly completed empty response receives one bounded
|
||||
same-route repair attempt, then fails as `MODEL_PROVIDER_RESPONSE_INVALID`.
|
||||
@@ -14,7 +14,20 @@ import lazy_loader as _lazy
|
||||
|
||||
__getattr__, __dir__, __all__ = _lazy.attach(
|
||||
__name__,
|
||||
submodules=["context_window", "models", "patches"],
|
||||
submodules=[
|
||||
"context_window",
|
||||
"models",
|
||||
"patches",
|
||||
"contracts",
|
||||
"config_admin",
|
||||
"configuration",
|
||||
"crypto",
|
||||
"invocation",
|
||||
"model_config",
|
||||
"runtime",
|
||||
"adapter_registry",
|
||||
"user_options",
|
||||
],
|
||||
submod_attrs={
|
||||
"context_window": [
|
||||
"DEFAULT_CONTEXT_WINDOW_FALLBACK",
|
||||
@@ -29,5 +42,45 @@ __getattr__, __dir__, __all__ = _lazy.attach(
|
||||
"get_models_for_provider",
|
||||
"list_models",
|
||||
],
|
||||
"contracts": [
|
||||
"AdmissionGrant",
|
||||
"AgentExecutionProfile",
|
||||
"AgentInputV3",
|
||||
"AgentModelSet",
|
||||
"EvoRuntimeError",
|
||||
"EvoRuntimeEvent",
|
||||
"HmacGrantAuthority",
|
||||
"PreparedRunQuote",
|
||||
"RoutePreparationGrant",
|
||||
"WebHostContext",
|
||||
],
|
||||
"model_config": [
|
||||
"EvoModelConfig",
|
||||
"FileEvoModelConfigStore",
|
||||
"SaveModelConfigCommand",
|
||||
],
|
||||
"config_admin": ["EvoModelConfigAdminService"],
|
||||
"configuration": [
|
||||
"EndpointConfig",
|
||||
"ModelConfig",
|
||||
"ProviderConfig",
|
||||
"ResolvedSecret",
|
||||
"SecretReference",
|
||||
"SecretResolver",
|
||||
],
|
||||
"invocation": [
|
||||
"InvocationPlan",
|
||||
"compile_invocation_plan",
|
||||
"derive_runtime_invocation",
|
||||
"derive_tool_call_transport",
|
||||
],
|
||||
"runtime": ["EvoModelRuntime"],
|
||||
"adapter_registry": ["AdapterRegistry", "get_adapter_registry"],
|
||||
"user_options": [
|
||||
"model_options_schema_hash",
|
||||
"project_user_options_for_purpose",
|
||||
"validate_parameter_constraints",
|
||||
"validate_user_model_options",
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,24 @@
|
||||
"""Canonical provider and model configuration contracts.
|
||||
|
||||
Parsing and persistence remain in :mod:`EvoScientist.llm.model_config` for
|
||||
backward compatibility. New code should import the owned contracts from this
|
||||
package so provider connection data and model capability data stay separate.
|
||||
"""
|
||||
|
||||
from .model import ModelConfig
|
||||
from .provider import (
|
||||
EndpointConfig,
|
||||
ProviderConfig,
|
||||
ResolvedSecret,
|
||||
SecretReference,
|
||||
SecretResolver,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"EndpointConfig",
|
||||
"ModelConfig",
|
||||
"ProviderConfig",
|
||||
"ResolvedSecret",
|
||||
"SecretReference",
|
||||
"SecretResolver",
|
||||
]
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Model-owned configuration.
|
||||
|
||||
This module contains model identity, capability, limit, access, billing, and
|
||||
canonical parameter declarations. It deliberately contains no credentials,
|
||||
endpoint URLs, SDK clients, or derived wire-transport fields.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from ..contracts import PricingQuote
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ModelConfig:
|
||||
"""Normalized configuration owned by one provider model profile."""
|
||||
|
||||
model_id: str
|
||||
params: Mapping[str, Any]
|
||||
supports_vision: bool
|
||||
supports_reasoning: bool
|
||||
allowed_reasoning_efforts: tuple[str, ...]
|
||||
context_window: int
|
||||
max_output_tokens: int
|
||||
reasoning_mode: str
|
||||
reasoning_enabled_params: Mapping[str, Any]
|
||||
reasoning_disabled_params: Mapping[str, Any]
|
||||
allowed_plans: tuple[str, ...]
|
||||
allowed_roles: tuple[str, ...]
|
||||
quote: PricingQuote
|
||||
model_key: str = ""
|
||||
display_name: str = ""
|
||||
description: str = ""
|
||||
tags: tuple[str, ...] = ()
|
||||
enabled: bool = True
|
||||
version_policy: str = "rolling"
|
||||
resolved_model_revision: str | None = None
|
||||
reproducible: bool = False
|
||||
capabilities: Mapping[str, bool] = field(default_factory=dict)
|
||||
purpose_overrides: Mapping[str, Mapping[str, Any]] = field(default_factory=dict)
|
||||
user_options: Mapping[str, Mapping[str, Any]] = field(default_factory=dict)
|
||||
parameter_constraints: tuple[Mapping[str, Any], ...] = ()
|
||||
access: Mapping[str, Any] = field(default_factory=dict)
|
||||
max_inflight_requests: int | None = None
|
||||
descriptor_parameters: Mapping[str, Any] = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def billing_sku(self) -> str:
|
||||
return self.quote.billing_sku
|
||||
@@ -0,0 +1,61 @@
|
||||
"""Provider-owned connection and adapter configuration.
|
||||
|
||||
Provider configuration owns credentials, endpoints, headers, connection
|
||||
defaults, and adapter identity. Model capabilities and model-level parameters
|
||||
are represented by :class:`ModelConfig`, not duplicated here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from .model import ModelConfig
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SecretReference:
|
||||
ref: str
|
||||
revision: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ResolvedSecret:
|
||||
value: str
|
||||
declared_revision: int
|
||||
authoritative_version: str | None
|
||||
runtime_fingerprint: str
|
||||
|
||||
|
||||
SecretResolver = Callable[[SecretReference], ResolvedSecret]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EndpointConfig:
|
||||
"""Provider connection endpoint; never contains model behavior."""
|
||||
|
||||
name: str
|
||||
base_url: str
|
||||
auth: SecretReference
|
||||
headers: Mapping[str, str]
|
||||
header_refs: Mapping[str, SecretReference]
|
||||
params: Mapping[str, Any]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProviderConfig:
|
||||
"""Provider adapter and connection configuration with owned model profiles."""
|
||||
|
||||
key: str
|
||||
protocol: str
|
||||
params: Mapping[str, Any]
|
||||
endpoints: Mapping[str, EndpointConfig]
|
||||
models: Mapping[str, ModelConfig]
|
||||
display_name: str = ""
|
||||
adapter_id: str = ""
|
||||
adapter_revision: str = ""
|
||||
wire_protocol: str = ""
|
||||
enabled: bool = True
|
||||
connection_defaults: Mapping[str, int] = field(default_factory=dict)
|
||||
implementation_fingerprint: str = ""
|
||||
@@ -0,0 +1,870 @@
|
||||
"""Public V3 contracts for the embedded EvoScientist Web model runtime.
|
||||
|
||||
The host can provide identity, storage and durable event services through
|
||||
these DTOs and protocols. Provider connection details never cross this
|
||||
boundary.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from dataclasses import asdict, dataclass, field, fields
|
||||
from typing import Any, Literal, Protocol, TypeVar
|
||||
|
||||
from .crypto import (
|
||||
HmacKeyRing,
|
||||
KeyMaterial,
|
||||
canonical_json_v1,
|
||||
hmac_id,
|
||||
sign_contract,
|
||||
verify_contract,
|
||||
)
|
||||
|
||||
CONTRACT_VERSION = 3
|
||||
ADMIN_CONTROL_VERSION = 2
|
||||
GATEWAY_ISSUER = "ai4sci-gateway"
|
||||
EVO_ISSUER = "evoscientist-runtime"
|
||||
GATEWAY_AUDIENCE = EVO_ISSUER
|
||||
EVO_AUDIENCE = GATEWAY_ISSUER
|
||||
MAX_CLOCK_SKEW_MS = 5_000
|
||||
MAX_CONTRACT_TTL_MS = 120_000
|
||||
MAX_ADMIN_TTL_MS = 60_000
|
||||
|
||||
_GATEWAY_GRANT_INFO = "ai4sci/gateway-to-evo-grant/v3"
|
||||
_EVO_QUOTE_INFO = "ai4sci/evo-to-gateway-quote/v3"
|
||||
_INPUT_DIGEST_INFO = "ai4sci/agent-input-digest/v3"
|
||||
_PREPARED_SNAPSHOT_INFO = "ai4sci/prepared-snapshot-digest/v3"
|
||||
_PREPARED_INPUT_INFO = "ai4sci/prepared-input-digest/v3"
|
||||
_TOOL_REGISTRY_INFO = "ai4sci/tool-registry-snapshot/v3"
|
||||
_ADMIN_CONFIG_INFO = "ai4sci/admin-config/v2"
|
||||
|
||||
|
||||
class EvoRuntimeError(RuntimeError):
|
||||
"""A stable error whose code may be projected to the host."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
code: str,
|
||||
message: str | None = None,
|
||||
*,
|
||||
details: Sequence[Mapping[str, Any]] = (),
|
||||
) -> None:
|
||||
super().__init__(message or code)
|
||||
self.code = code
|
||||
self.details = tuple(dict(item) for item in details)
|
||||
|
||||
|
||||
def now_ms() -> int:
|
||||
return time.time_ns() // 1_000_000
|
||||
|
||||
|
||||
def _unsigned(value: Any) -> dict[str, Any]:
|
||||
payload = asdict(value)
|
||||
payload.pop("signature", None)
|
||||
return payload
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RoutePreparationGrant:
|
||||
issuer: str
|
||||
audience: str
|
||||
grant_id: str
|
||||
request_id: str
|
||||
turn_id: str
|
||||
thread_id: str
|
||||
subject_id: str
|
||||
requested_model_ref: str
|
||||
plan: str
|
||||
roles: tuple[str, ...]
|
||||
requires_vision: bool
|
||||
reasoning_effort: str
|
||||
title_policy: Literal["disabled", "best_effort"]
|
||||
gateway_input_digest: str
|
||||
checkpoint_thread_id: str
|
||||
checkpoint_snapshot_id: str
|
||||
turn_fencing_token: int
|
||||
issued_at: int
|
||||
expires_at: int
|
||||
key_id: str
|
||||
signature: str
|
||||
contract_version: int = CONTRACT_VERSION
|
||||
|
||||
def unsigned_payload(self) -> dict[str, Any]:
|
||||
return _unsigned(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouteIdentity:
|
||||
config_revision: int
|
||||
config_identity_key_id: str
|
||||
purpose: str
|
||||
route_selector_id: str
|
||||
route_fingerprint: str
|
||||
provider_id: str
|
||||
endpoint_name: str
|
||||
model_id: str
|
||||
protocol: str
|
||||
api_mode: str
|
||||
tool_call_transport: str
|
||||
route_semantics_hash: str
|
||||
billing_sku: str
|
||||
pricing_revision: str
|
||||
quote_id: str
|
||||
|
||||
@property
|
||||
def route_key(self) -> str:
|
||||
return ":".join(
|
||||
(
|
||||
self.provider_id,
|
||||
self.endpoint_name,
|
||||
self.model_id,
|
||||
self.api_mode,
|
||||
self.tool_call_transport,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PricingQuote:
|
||||
billing_sku: str
|
||||
pricing_revision: str
|
||||
currency: str
|
||||
unit_scale: int
|
||||
input_microunits_per_million: int
|
||||
cached_input_microunits_per_million: int
|
||||
output_microunits_per_million: int
|
||||
quote_id: str
|
||||
multiplier: str = "1"
|
||||
|
||||
@property
|
||||
def cached_microunits_per_million(self) -> int:
|
||||
return self.cached_input_microunits_per_million
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouteCallBound:
|
||||
route_identity: RouteIdentity
|
||||
max_output_tokens: int
|
||||
payload_input_hard_cap: int
|
||||
billable_input_cap: int
|
||||
protocol_margin_tokens: int
|
||||
attempt_reserve_microunits: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PreparedRunQuote:
|
||||
issuer: str
|
||||
audience: str
|
||||
preparation_id: str
|
||||
request_id: str
|
||||
turn_id: str
|
||||
thread_id: str
|
||||
subject_id: str
|
||||
requested_model_ref: str
|
||||
plan: str
|
||||
roles: tuple[str, ...]
|
||||
requires_vision: bool
|
||||
reasoning_effort: str
|
||||
title_policy: Literal["disabled", "best_effort"]
|
||||
gateway_input_digest: str
|
||||
prepared_snapshot_digest: str
|
||||
prepared_input_digest: str
|
||||
config_revision: int
|
||||
catalog_revision: int
|
||||
enabled_purposes: tuple[str, ...]
|
||||
purpose_routes: Mapping[str, Mapping[str, Any]]
|
||||
purpose_route_call_bounds: Mapping[str, tuple[RouteCallBound, ...]]
|
||||
purpose_attempt_limits: Mapping[str, int]
|
||||
total_max_attempts: int
|
||||
pricing_quotes: Mapping[str, PricingQuote]
|
||||
quote_ids: tuple[str, ...]
|
||||
provider_run_reserve_microunits: int
|
||||
checkpoint_snapshot_id: str
|
||||
tool_registry_snapshot_id: str
|
||||
turn_fencing_token: int
|
||||
route_semantics_hashes: tuple[str, ...]
|
||||
issued_at: int
|
||||
expires_at: int
|
||||
key_id: str
|
||||
signature: str
|
||||
contract_version: int = CONTRACT_VERSION
|
||||
|
||||
def unsigned_payload(self) -> dict[str, Any]:
|
||||
return _unsigned(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AdmissionGrant:
|
||||
issuer: str
|
||||
audience: str
|
||||
grant_id: str
|
||||
preparation_id: str
|
||||
request_id: str
|
||||
turn_id: str
|
||||
thread_id: str
|
||||
subject_id: str
|
||||
requested_model_ref: str
|
||||
plan: str
|
||||
roles: tuple[str, ...]
|
||||
requires_vision: bool
|
||||
reasoning_effort: str
|
||||
title_policy: Literal["disabled", "best_effort"]
|
||||
gateway_input_digest: str
|
||||
prepared_snapshot_digest: str
|
||||
prepared_input_digest: str
|
||||
config_revision: int
|
||||
catalog_revision: int
|
||||
purpose_attempt_limits: Mapping[str, int]
|
||||
total_max_attempts: int
|
||||
checkpoint_snapshot_id: str
|
||||
tool_registry_snapshot_id: str
|
||||
turn_fencing_token: int
|
||||
admission_snapshot_id: str
|
||||
admission_id: str
|
||||
hold_id: str
|
||||
billing_fencing_token: int
|
||||
provider_run_reserve_microunits: int
|
||||
billing_policy_version: str
|
||||
issued_at: int
|
||||
expires_at: int
|
||||
key_id: str
|
||||
signature: str
|
||||
contract_version: int = CONTRACT_VERSION
|
||||
|
||||
def unsigned_payload(self) -> dict[str, Any]:
|
||||
return _unsigned(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class VerifiedModelSubject:
|
||||
issuer: str
|
||||
audience: str
|
||||
grant_id: str
|
||||
subject_id: str
|
||||
plan: str
|
||||
roles: tuple[str, ...]
|
||||
issued_at: int
|
||||
expires_at: int
|
||||
key_id: str
|
||||
signature: str
|
||||
contract_version: int = CONTRACT_VERSION
|
||||
|
||||
def unsigned_payload(self) -> dict[str, Any]:
|
||||
return _unsigned(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AdminConfigGrant:
|
||||
issuer: str
|
||||
audience: str
|
||||
grant_id: str
|
||||
subject_id: str
|
||||
action: str
|
||||
operation_id: str
|
||||
request_digest: str
|
||||
issued_at: int
|
||||
expires_at: int
|
||||
key_id: str
|
||||
signature: str
|
||||
contract_version: int = CONTRACT_VERSION
|
||||
|
||||
def unsigned_payload(self) -> dict[str, Any]:
|
||||
return _unsigned(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GetModelConfigRequest:
|
||||
operation_id: str
|
||||
admin_grant: AdminConfigGrant
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GetModelConfigResult:
|
||||
operation_id: str
|
||||
config_revision: int
|
||||
redacted_config: Mapping[str, Any]
|
||||
catalog_projection: Mapping[str, Any]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ConcreteRouteProposal:
|
||||
route: Mapping[str, str]
|
||||
route_semantics_hash: str
|
||||
endpoint_fingerprint: str
|
||||
required_probe_kinds: tuple[str, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ValidateCandidateConfigRequest:
|
||||
operation_id: str
|
||||
expected_revision: int
|
||||
payload: Mapping[str, Any]
|
||||
admin_grant: AdminConfigGrant
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ValidateCandidateConfigResult:
|
||||
operation_id: str
|
||||
target_revision: int
|
||||
config_identity_key_id: str
|
||||
proposal_hash: str
|
||||
concrete_routes: tuple[ConcreteRouteProposal, ...]
|
||||
expires_at: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProbeCandidateRouteRequest:
|
||||
operation_id: str
|
||||
proposal_hash: str
|
||||
route_semantics_hash: str
|
||||
probe_kind: Literal["connectivity", "tool_protocol"]
|
||||
admin_grant: AdminConfigGrant
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProbeCandidateRouteResult:
|
||||
operation_id: str
|
||||
evidence_id: str
|
||||
route_semantics_hash: str
|
||||
connectivity: str
|
||||
tool_capability: str
|
||||
adapter_revision: str
|
||||
fixture_digest: str
|
||||
expires_at: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CommitModelConfigRequest:
|
||||
operation_id: str
|
||||
expected_revision: int
|
||||
proposal_hash: str
|
||||
payload: Mapping[str, Any]
|
||||
evidence_ids: tuple[str, ...]
|
||||
admin_grant: AdminConfigGrant
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CommitModelConfigResult:
|
||||
operation_id: str
|
||||
config_revision: int
|
||||
committed_at: int
|
||||
redacted_diff: Mapping[str, Any]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CreateProposalRequest:
|
||||
operation_id: str
|
||||
expected_active_revision: int
|
||||
admin_grant: AdminConfigGrant
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class UpdateProposalRequest:
|
||||
operation_id: str
|
||||
proposal_id: str
|
||||
expected_state_version: int
|
||||
expected_draft_etag: str
|
||||
draft_payload: Mapping[str, Any]
|
||||
admin_grant: AdminConfigGrant
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ValidateProposalRequest:
|
||||
operation_id: str
|
||||
proposal_id: str
|
||||
expected_state_version: int
|
||||
expected_draft_etag: str
|
||||
admin_grant: AdminConfigGrant
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProbeProposalRequest:
|
||||
operation_id: str
|
||||
proposal_id: str
|
||||
validated_digest: str
|
||||
route_semantics_hash: str
|
||||
probe_kind: str
|
||||
admin_grant: AdminConfigGrant
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CommitProposalRequest:
|
||||
operation_id: str
|
||||
proposal_id: str
|
||||
expected_active_revision: int
|
||||
expected_state_version: int
|
||||
expected_draft_etag: str
|
||||
validated_digest: str
|
||||
evidence_ids: tuple[str, ...]
|
||||
admin_grant: AdminConfigGrant
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CancelProposalRequest:
|
||||
operation_id: str
|
||||
proposal_id: str
|
||||
expected_state_version: int
|
||||
admin_grant: AdminConfigGrant
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RollbackConfigRequest:
|
||||
operation_id: str
|
||||
expected_active_revision: int
|
||||
target_revision: int
|
||||
admin_grant: AdminConfigGrant
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DiscoverProviderModelsRequest:
|
||||
operation_id: str
|
||||
proposal_id: str
|
||||
provider_id: str
|
||||
admin_grant: AdminConfigGrant
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DiscoverProviderModelsResult:
|
||||
operation_id: str
|
||||
proposal_id: str
|
||||
provider_id: str
|
||||
source: str
|
||||
discovered_at: int
|
||||
models: tuple[Mapping[str, str], ...]
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AdminProposalResult:
|
||||
operation_id: str
|
||||
active_revision: int
|
||||
proposal_id: str
|
||||
state: str
|
||||
state_version: int
|
||||
draft_etag: str
|
||||
base_revision: int
|
||||
target_revision: int
|
||||
expires_at: int
|
||||
validated_digest: str | None = None
|
||||
routes: tuple[Mapping[str, Any], ...] = ()
|
||||
evidence_ids: tuple[str, ...] = ()
|
||||
committed_at: int | None = None
|
||||
draft_payload: Mapping[str, Any] | None = None
|
||||
migration_report: Mapping[str, Any] | None = None
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RollbackConfigResult:
|
||||
operation_id: str
|
||||
previous_active_revision: int
|
||||
active_revision: int
|
||||
target_revision: int
|
||||
rolled_back_at: int
|
||||
admin_control_version: int = ADMIN_CONTROL_VERSION
|
||||
|
||||
|
||||
class AdmissionGrantVerifier(Protocol):
|
||||
def require_preparation(self, grant: RoutePreparationGrant) -> None: ...
|
||||
|
||||
def require_admission(self, grant: AdmissionGrant) -> None: ...
|
||||
|
||||
def verify_subject(self, subject: VerifiedModelSubject) -> bool: ...
|
||||
|
||||
|
||||
class AdminConfigGrantVerifier(Protocol):
|
||||
def require_admin(self, grant: AdminConfigGrant) -> None: ...
|
||||
|
||||
|
||||
_ContractT = TypeVar(
|
||||
"_ContractT",
|
||||
RoutePreparationGrant,
|
||||
PreparedRunQuote,
|
||||
AdmissionGrant,
|
||||
VerifiedModelSubject,
|
||||
AdminConfigGrant,
|
||||
)
|
||||
|
||||
|
||||
class HmacGrantAuthority:
|
||||
"""Purpose-separated V3 contract signer with current/previous key support."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
secret: str | bytes,
|
||||
key_id: str = "runtime-current",
|
||||
*,
|
||||
previous_secret: str | bytes | None = None,
|
||||
previous_key_id: str | None = None,
|
||||
) -> None:
|
||||
previous = None
|
||||
if (previous_secret is None) != (previous_key_id is None):
|
||||
raise ValueError("previous secret and key id must be configured together")
|
||||
if previous_secret is not None and previous_key_id is not None:
|
||||
previous = KeyMaterial.create(previous_key_id, previous_secret)
|
||||
self.key_ring = HmacKeyRing(KeyMaterial.create(key_id, secret), previous)
|
||||
|
||||
def sign_preparation(self, **kwargs: Any) -> RoutePreparationGrant:
|
||||
return self._sign(
|
||||
RoutePreparationGrant, _GATEWAY_GRANT_INFO, gateway=True, **kwargs
|
||||
)
|
||||
|
||||
def sign_admission(self, **kwargs: Any) -> AdmissionGrant:
|
||||
return self._sign(AdmissionGrant, _GATEWAY_GRANT_INFO, gateway=True, **kwargs)
|
||||
|
||||
def sign_subject(self, **kwargs: Any) -> VerifiedModelSubject:
|
||||
return self._sign(
|
||||
VerifiedModelSubject, _GATEWAY_GRANT_INFO, gateway=True, **kwargs
|
||||
)
|
||||
|
||||
def sign_admin(self, **kwargs: Any) -> AdminConfigGrant:
|
||||
return self._sign(AdminConfigGrant, _ADMIN_CONFIG_INFO, gateway=True, **kwargs)
|
||||
|
||||
def sign_quote(self, **kwargs: Any) -> PreparedRunQuote:
|
||||
return self._sign(PreparedRunQuote, _EVO_QUOTE_INFO, gateway=False, **kwargs)
|
||||
|
||||
def agent_input_digest(self, payload: Any) -> str:
|
||||
_, key = self.key_ring.derive_current(_INPUT_DIGEST_INFO)
|
||||
return hmac_id(key, payload)
|
||||
|
||||
def prepared_input_digest(self, payload: Any) -> str:
|
||||
_, key = self.key_ring.derive_current(_PREPARED_INPUT_INFO)
|
||||
return hmac_id(key, payload)
|
||||
|
||||
def prepared_snapshot_digest(self, payload: Any) -> str:
|
||||
_, key = self.key_ring.derive_current(_PREPARED_SNAPSHOT_INFO)
|
||||
return hmac_id(key, payload)
|
||||
|
||||
def tool_registry_snapshot_id(self, payload: Any) -> str:
|
||||
_, key = self.key_ring.derive_current(_TOOL_REGISTRY_INFO)
|
||||
return hmac_id(key, payload)
|
||||
|
||||
def require_preparation(self, grant: RoutePreparationGrant) -> None:
|
||||
if not isinstance(grant, RoutePreparationGrant):
|
||||
raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID")
|
||||
self._require(grant, _GATEWAY_GRANT_INFO, gateway=True)
|
||||
|
||||
def require_admission(self, grant: AdmissionGrant) -> None:
|
||||
if not isinstance(grant, AdmissionGrant):
|
||||
raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID")
|
||||
self._require(grant, _GATEWAY_GRANT_INFO, gateway=True)
|
||||
|
||||
def require_admin(self, grant: AdminConfigGrant) -> None:
|
||||
if not isinstance(grant, AdminConfigGrant):
|
||||
raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID")
|
||||
self._require(
|
||||
grant, _ADMIN_CONFIG_INFO, gateway=True, max_ttl_ms=MAX_ADMIN_TTL_MS
|
||||
)
|
||||
|
||||
def require_quote(self, quote: PreparedRunQuote) -> None:
|
||||
if not isinstance(quote, PreparedRunQuote):
|
||||
raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID")
|
||||
self._require(quote, _EVO_QUOTE_INFO, gateway=False)
|
||||
|
||||
def verify_admission(self, grant: AdmissionGrant) -> bool:
|
||||
try:
|
||||
self.require_admission(grant)
|
||||
except EvoRuntimeError:
|
||||
return False
|
||||
return True
|
||||
|
||||
def verify_subject(self, subject: VerifiedModelSubject) -> bool:
|
||||
try:
|
||||
self._require(subject, _GATEWAY_GRANT_INFO, gateway=True)
|
||||
except EvoRuntimeError:
|
||||
return False
|
||||
return True
|
||||
|
||||
def verify_admin(self, grant: AdminConfigGrant) -> bool:
|
||||
try:
|
||||
self.require_admin(grant)
|
||||
except EvoRuntimeError:
|
||||
return False
|
||||
return True
|
||||
|
||||
def _sign(
|
||||
self,
|
||||
contract_class: type[_ContractT],
|
||||
info: str,
|
||||
*,
|
||||
gateway: bool,
|
||||
**kwargs: Any,
|
||||
) -> _ContractT:
|
||||
issued_at = int(kwargs.pop("issued_at", now_ms()))
|
||||
ttl_ms = int(kwargs.pop("ttl_ms", 60_000))
|
||||
key_id, key = self.key_ring.derive_current(info)
|
||||
defaults = {
|
||||
"contract_version": CONTRACT_VERSION,
|
||||
"issuer": GATEWAY_ISSUER if gateway else EVO_ISSUER,
|
||||
"audience": GATEWAY_AUDIENCE if gateway else EVO_AUDIENCE,
|
||||
"issued_at": issued_at,
|
||||
"expires_at": int(kwargs.pop("expires_at", issued_at + ttl_ms)),
|
||||
"key_id": key_id,
|
||||
}
|
||||
field_names = {item.name for item in fields(contract_class)}
|
||||
if "grant_id" in field_names:
|
||||
defaults["grant_id"] = str(kwargs.pop("grant_id", uuid.uuid4()))
|
||||
payload = {**defaults, **kwargs}
|
||||
if "roles" in payload:
|
||||
payload["roles"] = tuple(sorted({str(role) for role in payload["roles"]}))
|
||||
unsigned = {
|
||||
key_name: value
|
||||
for key_name, value in payload.items()
|
||||
if key_name != "signature"
|
||||
}
|
||||
signature = sign_contract(contract_class.__name__, unsigned, key)
|
||||
return contract_class(signature=signature, **payload)
|
||||
|
||||
def _require(
|
||||
self,
|
||||
contract: Any,
|
||||
info: str,
|
||||
*,
|
||||
gateway: bool,
|
||||
max_ttl_ms: int = MAX_CONTRACT_TTL_MS,
|
||||
) -> None:
|
||||
if int(contract.contract_version) != CONTRACT_VERSION:
|
||||
raise EvoRuntimeError("CONTRACT_VERSION_UNSUPPORTED")
|
||||
expected_issuer = GATEWAY_ISSUER if gateway else EVO_ISSUER
|
||||
expected_audience = GATEWAY_AUDIENCE if gateway else EVO_AUDIENCE
|
||||
if contract.issuer != expected_issuer or contract.audience != expected_audience:
|
||||
raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID")
|
||||
try:
|
||||
key = self.key_ring.derive(contract.key_id, info)
|
||||
except KeyError as exc:
|
||||
raise EvoRuntimeError("CONTRACT_KEY_UNKNOWN") from exc
|
||||
if not verify_contract(
|
||||
type(contract).__name__,
|
||||
contract.unsigned_payload(),
|
||||
contract.signature,
|
||||
key,
|
||||
):
|
||||
raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID")
|
||||
current = now_ms()
|
||||
issued_at = int(contract.issued_at)
|
||||
expires_at = int(contract.expires_at)
|
||||
if expires_at <= issued_at or expires_at - issued_at > max_ttl_ms:
|
||||
raise EvoRuntimeError("CONTRACT_EXPIRED")
|
||||
if (
|
||||
issued_at > current + MAX_CLOCK_SKEW_MS
|
||||
or expires_at < current - MAX_CLOCK_SKEW_MS
|
||||
):
|
||||
raise EvoRuntimeError("CONTRACT_EXPIRED")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MediaDescriptor:
|
||||
media_id: str
|
||||
media_type: str
|
||||
content_hash: str
|
||||
version: str
|
||||
byte_length: int
|
||||
token_bound: int
|
||||
locator: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AgentInputV3:
|
||||
message: Any
|
||||
checkpoint_thread_id: str
|
||||
metadata: Mapping[str, Any] = field(default_factory=dict)
|
||||
media: Sequence[MediaDescriptor | Mapping[str, Any] | str] = ()
|
||||
content_blocks: Sequence[Mapping[str, Any]] = ()
|
||||
|
||||
def projection(self) -> Mapping[str, Any]:
|
||||
allowed_metadata = {
|
||||
key: self.metadata[key]
|
||||
for key in sorted(self.metadata)
|
||||
if key
|
||||
in {
|
||||
"force_context_repair",
|
||||
"public_thread_id",
|
||||
"source",
|
||||
"user_id",
|
||||
"model_options",
|
||||
}
|
||||
}
|
||||
media = [
|
||||
asdict(item) if isinstance(item, MediaDescriptor) else item
|
||||
for item in self.media
|
||||
]
|
||||
return {
|
||||
"message": self.message,
|
||||
"checkpoint_thread_id": self.checkpoint_thread_id,
|
||||
"metadata": allowed_metadata,
|
||||
"media": media,
|
||||
"content_blocks": list(self.content_blocks),
|
||||
}
|
||||
|
||||
def canonical_bytes(self) -> bytes:
|
||||
return canonical_json_v1(self.projection())
|
||||
|
||||
|
||||
AgentInput = AgentInputV3
|
||||
|
||||
|
||||
class RuntimeEventSink(Protocol):
|
||||
async def commit(self, event: EvoRuntimeEvent) -> str: ...
|
||||
|
||||
async def confirm(self, event_id: str, payload_digest: str) -> str: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WebHostContext:
|
||||
workspace_dir: str
|
||||
memory_dir: str
|
||||
workspace_backend: Any
|
||||
checkpointer: Any
|
||||
tool_selector_threshold: int | None = None
|
||||
memory_max_inline_profile_chars: int | None = None
|
||||
on_mcp_progress: Any = None
|
||||
runtime_event_sink: RuntimeEventSink | None = None
|
||||
tool_registry: Sequence[Any] = ()
|
||||
tool_registry_revision: str = "static"
|
||||
tool_registry_provider: Any = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AgentExecutionProfile:
|
||||
name: str
|
||||
configurable_model_override: bool = True
|
||||
subagents: bool = True
|
||||
async_subagents: bool = True
|
||||
memory_workers: bool = True
|
||||
scheduler: bool = True
|
||||
background_execution: bool = True
|
||||
|
||||
@classmethod
|
||||
def web_v3(cls) -> AgentExecutionProfile:
|
||||
return cls(
|
||||
name="web_v3",
|
||||
configurable_model_override=False,
|
||||
subagents=False,
|
||||
async_subagents=False,
|
||||
memory_workers=False,
|
||||
scheduler=False,
|
||||
background_execution=False,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def web_v1(cls) -> AgentExecutionProfile:
|
||||
return cls.web_v3()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AgentModelSet:
|
||||
main_agent: Any
|
||||
tool_selector: Any
|
||||
deepagents_summarizer: Any
|
||||
title: Any | None = None
|
||||
main_fallbacks: tuple[Any, ...] = ()
|
||||
route_health: Any = None
|
||||
capacity: Any = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ModelCatalogEntry:
|
||||
alias: str
|
||||
provider: str
|
||||
supports_vision: bool
|
||||
supports_reasoning: bool
|
||||
allowed_reasoning_efforts: tuple[str, ...]
|
||||
context_window: int
|
||||
max_output_tokens: int
|
||||
reasoning_mode: str
|
||||
billing_sku: str
|
||||
quote: PricingQuote
|
||||
default_reasoning_effort: str = ""
|
||||
health: Literal["closed", "open", "half_open"] = "closed"
|
||||
display_name: str = ""
|
||||
provider_display_name: str = ""
|
||||
description: str = ""
|
||||
version_policy: str = "rolling"
|
||||
resolved_model_revision: str | None = None
|
||||
reproducible: bool = False
|
||||
capabilities: tuple[str, ...] = ()
|
||||
user_options: Mapping[str, Mapping[str, Any]] = field(default_factory=dict)
|
||||
parameter_constraints: tuple[Mapping[str, Any], ...] = ()
|
||||
options_schema_hash: str = ""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ModelCatalog:
|
||||
catalog_revision: int
|
||||
default_alias: str
|
||||
entries: tuple[ModelCatalogEntry, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ModelAttemptEvent:
|
||||
request_id: str
|
||||
turn_id: str
|
||||
run_id: str
|
||||
preparation_id: str
|
||||
admission_snapshot_id: str
|
||||
admission_id: str
|
||||
hold_id: str
|
||||
prepared_snapshot_digest: str
|
||||
prepared_input_digest: str
|
||||
billing_fencing_token: int
|
||||
turn_fencing_token: int
|
||||
logical_call_id: str
|
||||
attempt_id: str
|
||||
purpose: Literal["main_agent", "tool_selector", "deepagents_summarizer", "title"]
|
||||
attempt_index: int
|
||||
identity: RouteIdentity
|
||||
quote_id: str
|
||||
billing_intent: Literal["user_charge", "platform_cost"]
|
||||
outcome: Literal["rejected", "started", "succeeded", "failed", "usage_unconfirmed"]
|
||||
provider_request_started: bool
|
||||
provider_input_bound_tokens: int
|
||||
provider_reserved_microunits: int
|
||||
usage: Mapping[str, int | str | None] | None = None
|
||||
usage_available: bool = False
|
||||
error_code: str | None = None
|
||||
health_state: Literal["closed", "open", "half_open"] = "closed"
|
||||
fallback_from_attempt_id: str | None = None
|
||||
timestamp: int = field(default_factory=now_ms)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EvoRuntimeEvent:
|
||||
event_id: str
|
||||
runtime_instance_id: str
|
||||
run_id: str
|
||||
request_id: str
|
||||
turn_id: str
|
||||
sequence: int
|
||||
kind: Literal["agent", "model_attempt", "title", "run"]
|
||||
payload: Mapping[str, Any]
|
||||
emitted_at: int = field(default_factory=now_ms)
|
||||
schema_version: int = CONTRACT_VERSION
|
||||
|
||||
|
||||
class EvoWebRun(Protocol):
|
||||
@property
|
||||
def run_id(self) -> str: ...
|
||||
|
||||
async def stream(
|
||||
self, after_sequence: int | None = None
|
||||
) -> AsyncIterator[EvoRuntimeEvent]: ...
|
||||
|
||||
async def cancel(self, reason: str) -> str: ...
|
||||
|
||||
|
||||
def event_payload(value: Any) -> dict[str, Any]:
|
||||
if isinstance(value, ModelAttemptEvent):
|
||||
return asdict(value)
|
||||
if isinstance(value, Mapping):
|
||||
return dict(value)
|
||||
return {"value": str(value)}
|
||||
@@ -0,0 +1,159 @@
|
||||
"""Cryptographic primitives for the Evo Web runtime V3 contracts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import unicodedata
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import asdict, dataclass, is_dataclass
|
||||
from typing import Any
|
||||
|
||||
import rfc8785
|
||||
|
||||
_HKDF_SALT = b"ai4sci-evo-runtime-v3"
|
||||
_MIN_SECRET_BYTES = 32
|
||||
|
||||
|
||||
class CanonicalJsonError(ValueError):
|
||||
"""Raised when a value cannot be represented by canonical JSON."""
|
||||
|
||||
|
||||
def _normalize(value: Any) -> Any:
|
||||
if is_dataclass(value) and not isinstance(value, type):
|
||||
return _normalize(asdict(value))
|
||||
if value is None or isinstance(value, bool | int | float):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
return unicodedata.normalize("NFC", value)
|
||||
if isinstance(value, Mapping):
|
||||
normalized: dict[str, Any] = {}
|
||||
for key, item in value.items():
|
||||
if not isinstance(key, str):
|
||||
raise CanonicalJsonError("JSON object keys must be strings")
|
||||
normalized_key = unicodedata.normalize("NFC", key)
|
||||
if normalized_key in normalized:
|
||||
raise CanonicalJsonError("duplicate JSON key after NFC normalization")
|
||||
normalized[normalized_key] = _normalize(item)
|
||||
return normalized
|
||||
if isinstance(value, Sequence) and not isinstance(
|
||||
value, bytes | bytearray | memoryview
|
||||
):
|
||||
return [_normalize(item) for item in value]
|
||||
raise CanonicalJsonError(f"unsupported canonical JSON type: {type(value).__name__}")
|
||||
|
||||
|
||||
def canonical_json_v1(value: Any) -> bytes:
|
||||
"""Encode NFC-normalized data with RFC 8785 JSON canonicalization."""
|
||||
|
||||
try:
|
||||
return rfc8785.dumps(_normalize(value))
|
||||
except (rfc8785.CanonicalizationError, rfc8785.FloatDomainError, TypeError) as exc:
|
||||
raise CanonicalJsonError("value is not valid RFC 8785 JSON") from exc
|
||||
|
||||
|
||||
def sha256_id(value: Any) -> str:
|
||||
return f"sha256:{hashlib.sha256(canonical_json_v1(value)).hexdigest()}"
|
||||
|
||||
|
||||
def hkdf_sha256(root_key: bytes, *, info: str, length: int = 32) -> bytes:
|
||||
"""RFC 5869 HKDF-SHA256 with the protocol's fixed salt."""
|
||||
|
||||
if length < 1 or length > 255 * hashlib.sha256().digest_size:
|
||||
raise ValueError("invalid HKDF output length")
|
||||
prk = hmac.new(_HKDF_SALT, root_key, hashlib.sha256).digest()
|
||||
output = bytearray()
|
||||
previous = b""
|
||||
counter = 1
|
||||
info_bytes = info.encode("utf-8")
|
||||
while len(output) < length:
|
||||
previous = hmac.new(
|
||||
prk,
|
||||
previous + info_bytes + bytes((counter,)),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
output.extend(previous)
|
||||
counter += 1
|
||||
return bytes(output[:length])
|
||||
|
||||
|
||||
def hmac_id(key: bytes, value: Any) -> str:
|
||||
digest = hmac.new(key, canonical_json_v1(value), hashlib.sha256).hexdigest()
|
||||
return f"hmac-sha256:{digest}"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class KeyMaterial:
|
||||
key_id: str
|
||||
secret: bytes
|
||||
|
||||
@classmethod
|
||||
def create(cls, key_id: str, secret: str | bytes) -> KeyMaterial:
|
||||
normalized_id = str(key_id or "").strip()
|
||||
encoded = secret.encode("utf-8") if isinstance(secret, str) else bytes(secret)
|
||||
if not normalized_id:
|
||||
raise ValueError("key id is required")
|
||||
if len(encoded) < _MIN_SECRET_BYTES:
|
||||
raise ValueError("signing secret must contain at least 32 bytes")
|
||||
return cls(normalized_id, encoded)
|
||||
|
||||
|
||||
class HmacKeyRing:
|
||||
"""Versioned root keys with purpose-separated HKDF children.
|
||||
|
||||
``previous`` remains as a compatibility property for callers that still
|
||||
expose a two-key deployment contract. New persistence code uses
|
||||
``retained``/``key_ids`` so immutable revisions can outlive one rotation.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
current: KeyMaterial,
|
||||
previous: KeyMaterial | None = None,
|
||||
*,
|
||||
retained: Sequence[KeyMaterial] = (),
|
||||
) -> None:
|
||||
if previous is not None and previous.key_id == current.key_id:
|
||||
raise ValueError("current and previous key ids must differ")
|
||||
self.current = current
|
||||
self.previous = previous
|
||||
self._roots = {current.key_id: current.secret}
|
||||
if previous is not None:
|
||||
self._roots[previous.key_id] = previous.secret
|
||||
for item in retained:
|
||||
existing = self._roots.get(item.key_id)
|
||||
if existing is not None and existing != item.secret:
|
||||
raise ValueError("duplicate signing key id has different material")
|
||||
self._roots[item.key_id] = item.secret
|
||||
|
||||
@property
|
||||
def key_ids(self) -> frozenset[str]:
|
||||
return frozenset(self._roots)
|
||||
|
||||
def contains(self, key_id: str) -> bool:
|
||||
return key_id in self._roots
|
||||
|
||||
def derive_current(self, info: str) -> tuple[str, bytes]:
|
||||
return self.current.key_id, hkdf_sha256(self.current.secret, info=info)
|
||||
|
||||
def derive(self, key_id: str, info: str) -> bytes:
|
||||
try:
|
||||
root = self._roots[key_id]
|
||||
except KeyError as exc:
|
||||
raise KeyError("unknown signing key") from exc
|
||||
return hkdf_sha256(root, info=info)
|
||||
|
||||
|
||||
def sign_contract(contract_type: str, payload: Mapping[str, Any], key: bytes) -> str:
|
||||
message = contract_type.encode("utf-8") + b"\0" + canonical_json_v1(payload)
|
||||
return hmac.new(key, message, hashlib.sha256).hexdigest()
|
||||
|
||||
|
||||
def verify_contract(
|
||||
contract_type: str,
|
||||
payload: Mapping[str, Any],
|
||||
signature: str,
|
||||
key: bytes,
|
||||
) -> bool:
|
||||
expected = sign_contract(contract_type, payload, key)
|
||||
return hmac.compare_digest(expected, str(signature))
|
||||
@@ -60,10 +60,24 @@ class AgentControlError(Exception):
|
||||
"retryable": self.retryable,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def model_construct(cls, **payload: Any) -> AgentControlError:
|
||||
"""Rebuild the allowlisted checkpoint form without trusting extra fields."""
|
||||
return cls(
|
||||
str(payload.get("code") or "MODEL_REQUEST_REJECTED"),
|
||||
str(payload.get("message") or "Model request rejected."),
|
||||
status_code=int(payload.get("status_code") or 403),
|
||||
retryable=bool(payload.get("retryable", False)),
|
||||
)
|
||||
|
||||
|
||||
class ModelToolProtocolError(AgentControlError):
|
||||
"""A completed model response contained an invalid tool-call protocol."""
|
||||
|
||||
# Unlike authorization and admission control errors, a malformed model
|
||||
# response is safe to retry before the agent executes any tool.
|
||||
non_fallbackable = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reason: str,
|
||||
@@ -82,7 +96,7 @@ class ModelToolProtocolError(AgentControlError):
|
||||
"MODEL_TOOL_PROTOCOL_INVALID",
|
||||
"The model returned an invalid structured tool call.",
|
||||
status_code=502,
|
||||
retryable=False,
|
||||
retryable=True,
|
||||
)
|
||||
self.reason = reason
|
||||
self.provider = provider
|
||||
@@ -123,6 +137,56 @@ class ModelToolProtocolError(AgentControlError):
|
||||
payload[key] = value
|
||||
return payload
|
||||
|
||||
@classmethod
|
||||
def model_construct(cls, **payload: Any) -> ModelToolProtocolError:
|
||||
"""Rebuild only the public, redacted checkpoint projection."""
|
||||
|
||||
def optional_text(name: str) -> str | None:
|
||||
value = payload.get(name)
|
||||
return str(value) if value is not None else None
|
||||
|
||||
generation = payload.get("config_generation")
|
||||
return cls(
|
||||
str(payload.get("reason") or "invalid_tool_protocol"),
|
||||
provider=optional_text("provider"),
|
||||
model=optional_text("model"),
|
||||
route_key=optional_text("route_key"),
|
||||
config_generation=(int(generation) if generation is not None else None),
|
||||
api_mode=optional_text("api_mode"),
|
||||
endpoint=optional_text("endpoint"),
|
||||
tool_call_transport=optional_text("tool_call_transport"),
|
||||
call_id=optional_text("call_id"),
|
||||
)
|
||||
|
||||
|
||||
class ModelProviderResponseError(AgentControlError):
|
||||
"""A completed provider response had no final text or tool call."""
|
||||
|
||||
non_fallbackable = False
|
||||
|
||||
def __init__(self, reason: str = "empty_assistant_response") -> None:
|
||||
super().__init__(
|
||||
"MODEL_PROVIDER_RESPONSE_INVALID",
|
||||
"The model returned no final text or structured tool call.",
|
||||
status_code=502,
|
||||
retryable=True,
|
||||
)
|
||||
self.reason = reason
|
||||
self.fallbackable = True
|
||||
self.recoverable = True
|
||||
|
||||
def model_dump(self) -> dict[str, Any]:
|
||||
return {
|
||||
**super().model_dump(),
|
||||
"reason": self.reason,
|
||||
"fallbackable": self.fallbackable,
|
||||
"recoverable": self.recoverable,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def model_construct(cls, **payload: Any) -> ModelProviderResponseError:
|
||||
return cls(str(payload.get("reason") or "empty_assistant_response"))
|
||||
|
||||
|
||||
class ProviderStreamError(Exception):
|
||||
"""Envelope-shaped wrapper for a provider SDK exception raised
|
||||
@@ -194,6 +258,27 @@ class ProviderStreamError(Exception):
|
||||
"""
|
||||
return self.as_envelope()
|
||||
|
||||
@classmethod
|
||||
def model_construct(cls, **payload: Any) -> ProviderStreamError:
|
||||
"""Rebuild the redacted provider envelope stored in a checkpoint."""
|
||||
|
||||
def optional_text(name: str) -> str | None:
|
||||
value = payload.get(name)
|
||||
return str(value) if value is not None else None
|
||||
|
||||
status = payload.get("status_code")
|
||||
return cls(
|
||||
provider=str(payload.get("provider") or "unknown"),
|
||||
class_qualname=str(
|
||||
payload.get("class") or payload.get("error") or "ProviderError"
|
||||
),
|
||||
message=str(payload.get("message") or "Provider request failed."),
|
||||
status_code=int(status) if status is not None else None,
|
||||
code=optional_text("code"),
|
||||
err_type=optional_text("type"),
|
||||
request_id=optional_text("request_id"),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# API-key redaction — env-driven, prefix-only
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Secretless ChatModel proxy for Ai4Sci Graph-native runs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from langchain_core.language_models.chat_models import BaseChatModel
|
||||
from langchain_core.messages import BaseMessage, messages_from_dict, messages_to_dict
|
||||
from langchain_core.outputs import ChatGeneration, ChatResult
|
||||
from langchain_core.tools import BaseTool
|
||||
from langchain_core.utils.function_calling import convert_to_openai_tool
|
||||
from pydantic import Field
|
||||
|
||||
|
||||
class GatewayProxyChatModel(BaseChatModel):
|
||||
gateway_url: str
|
||||
run_id: str
|
||||
envelope_signature: str
|
||||
provider_id: str = ""
|
||||
model_id: str = ""
|
||||
bound_tools: list[dict[str, Any]] = Field(default_factory=list)
|
||||
bound_tool_choice: Any = None
|
||||
|
||||
@property
|
||||
def _llm_type(self) -> str:
|
||||
return "ai4sci-gateway-proxy"
|
||||
|
||||
@property
|
||||
def _identifying_params(self) -> dict[str, Any]:
|
||||
return {"provider_id": self.provider_id, "model_id": self.model_id}
|
||||
|
||||
def bind_tools(
|
||||
self,
|
||||
tools: Sequence[dict[str, Any] | type | BaseTool],
|
||||
*,
|
||||
tool_choice: str | dict[str, Any] | bool | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
del kwargs
|
||||
serialized = [convert_to_openai_tool(tool) for tool in tools]
|
||||
return self.model_copy(
|
||||
update={"bound_tools": serialized, "bound_tool_choice": tool_choice}
|
||||
)
|
||||
|
||||
def _generate(self, *args: Any, **kwargs: Any) -> ChatResult:
|
||||
del args, kwargs
|
||||
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH")
|
||||
|
||||
async def _agenerate(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
stop: list[str] | None = None,
|
||||
run_manager: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> ChatResult:
|
||||
del stop, kwargs
|
||||
attempt_id = str(getattr(run_manager, "run_id", None) or self.run_id)
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(660.0, connect=5.0)) as client:
|
||||
response = await client.post(
|
||||
f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/invoke",
|
||||
json={
|
||||
"run_id": self.run_id,
|
||||
"attempt_id": attempt_id,
|
||||
"envelope_signature": self.envelope_signature,
|
||||
"messages": messages_to_dict(messages),
|
||||
"tools": self.bound_tools,
|
||||
"tool_choice": self.bound_tool_choice,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
value = response.json()
|
||||
parsed = messages_from_dict([value["message"]])
|
||||
if len(parsed) != 1:
|
||||
raise RuntimeError("AI4SCI_MODEL_PROXY_RESPONSE_INVALID")
|
||||
return ChatResult(generations=[ChatGeneration(message=parsed[0])])
|
||||
|
||||
|
||||
def proxy_from_config(
|
||||
value: Mapping[str, Any], *, provider_id: str = "", model_id: str = ""
|
||||
) -> GatewayProxyChatModel:
|
||||
required = {
|
||||
name: str(value.get(name) or "")
|
||||
for name in ("gateway_url", "run_id", "envelope_signature")
|
||||
}
|
||||
if not all(required.values()):
|
||||
raise RuntimeError("AI4SCI_MODEL_PROXY_CONFIG_INVALID")
|
||||
return GatewayProxyChatModel(
|
||||
**required,
|
||||
provider_id=provider_id,
|
||||
model_id=model_id,
|
||||
)
|
||||
@@ -0,0 +1,379 @@
|
||||
"""LangChain chat-model bridge for the stateless Gemini Interactions API."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langchain_core.language_models.chat_models import BaseChatModel
|
||||
from langchain_core.messages import (
|
||||
AIMessage,
|
||||
AIMessageChunk,
|
||||
BaseMessage,
|
||||
HumanMessage,
|
||||
SystemMessage,
|
||||
ToolMessage,
|
||||
)
|
||||
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
|
||||
from langchain_core.tools import BaseTool
|
||||
from langchain_core.utils.function_calling import convert_to_openai_tool
|
||||
from pydantic import Field, SecretStr
|
||||
|
||||
|
||||
class GeminiInteractionsChatModel(BaseChatModel):
|
||||
"""Minimal native bridge that preserves signed Provider content blocks."""
|
||||
|
||||
model_name: str
|
||||
api_key: SecretStr
|
||||
base_url: str = "https://generativelanguage.googleapis.com"
|
||||
max_output_tokens: int = 8192
|
||||
temperature: float | None = None
|
||||
top_p: float | None = None
|
||||
thinking: bool | None = None
|
||||
store: bool = False
|
||||
bound_tools: tuple[dict[str, Any], ...] = Field(default_factory=tuple)
|
||||
|
||||
@property
|
||||
def _llm_type(self) -> str:
|
||||
return "google-gemini-interactions"
|
||||
|
||||
@property
|
||||
def _identifying_params(self) -> dict[str, Any]:
|
||||
return {"model_name": self.model_name, "api_mode": "interactions"}
|
||||
|
||||
def bind_tools(
|
||||
self,
|
||||
tools: Sequence[dict[str, Any] | type | BaseTool | Any],
|
||||
*,
|
||||
tool_choice: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
_ = tool_choice, kwargs
|
||||
compiled = []
|
||||
for tool in tools:
|
||||
value = convert_to_openai_tool(tool)
|
||||
function = value.get("function", value)
|
||||
compiled.append(
|
||||
{
|
||||
"type": "function",
|
||||
"name": function["name"],
|
||||
"description": function.get("description", ""),
|
||||
"parameters": function.get(
|
||||
"parameters", {"type": "object", "properties": {}}
|
||||
),
|
||||
}
|
||||
)
|
||||
return self.model_copy(update={"bound_tools": tuple(compiled)})
|
||||
|
||||
def _client(self) -> Any:
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
|
||||
return genai.Client(
|
||||
api_key=self.api_key.get_secret_value(),
|
||||
http_options=types.HttpOptions(base_url=self.base_url.rstrip("/")),
|
||||
)
|
||||
|
||||
def _request(self, messages: Sequence[BaseMessage]) -> dict[str, Any]:
|
||||
turns, system_instruction = _compile_messages(messages)
|
||||
generation_config: dict[str, Any] = {
|
||||
"max_output_tokens": self.max_output_tokens,
|
||||
}
|
||||
if self.temperature is not None:
|
||||
generation_config["temperature"] = self.temperature
|
||||
if self.top_p is not None:
|
||||
generation_config["top_p"] = self.top_p
|
||||
if self.thinking is not None:
|
||||
generation_config["thinking_level"] = "high" if self.thinking else "minimal"
|
||||
generation_config["thinking_summaries"] = (
|
||||
"auto" if self.thinking else "none"
|
||||
)
|
||||
return {
|
||||
"model": self.model_name,
|
||||
"input": turns,
|
||||
"system_instruction": system_instruction or "",
|
||||
"generation_config": generation_config,
|
||||
"tools": list(self.bound_tools),
|
||||
"store": False,
|
||||
"stream": False,
|
||||
}
|
||||
|
||||
def _generate(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
stop: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ChatResult:
|
||||
request = self._request(messages)
|
||||
if stop:
|
||||
request["generation_config"]["stop_sequences"] = stop
|
||||
response = self._client().interactions.create(**request)
|
||||
return _chat_result(response)
|
||||
|
||||
async def _agenerate(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
stop: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ChatResult:
|
||||
request = self._request(messages)
|
||||
if stop:
|
||||
request["generation_config"]["stop_sequences"] = stop
|
||||
response = await self._client().aio.interactions.create(**request)
|
||||
return _chat_result(response)
|
||||
|
||||
async def _astream(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
stop: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[ChatGenerationChunk]:
|
||||
request = self._request(messages)
|
||||
request["stream"] = True
|
||||
if stop:
|
||||
request["generation_config"]["stop_sequences"] = stop
|
||||
stream = await self._client().aio.interactions.create(**request)
|
||||
blocks: dict[int, dict[str, Any]] = {}
|
||||
async for event in stream:
|
||||
payload = _dump(event)
|
||||
event_type = payload.get("event_type")
|
||||
if event_type == "content.start":
|
||||
blocks[int(payload["index"])] = dict(payload.get("content") or {})
|
||||
continue
|
||||
if event_type == "content.delta":
|
||||
index = int(payload["index"])
|
||||
delta = dict(payload.get("delta") or {})
|
||||
block = blocks.setdefault(index, {})
|
||||
_merge_stream_delta(block, delta)
|
||||
if delta.get("type") == "text" and delta.get("text"):
|
||||
yield ChatGenerationChunk(
|
||||
message=AIMessageChunk(content=str(delta["text"]))
|
||||
)
|
||||
continue
|
||||
if event_type == "content.stop":
|
||||
block = blocks.get(int(payload["index"]), {})
|
||||
if block.get("type") == "function_call":
|
||||
yield ChatGenerationChunk(
|
||||
message=AIMessageChunk(
|
||||
content="",
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"id": str(block.get("id") or ""),
|
||||
"name": str(block.get("name") or ""),
|
||||
"args": json.dumps(
|
||||
block.get("arguments") or {},
|
||||
separators=(",", ":"),
|
||||
),
|
||||
"index": int(payload["index"]),
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
)
|
||||
)
|
||||
continue
|
||||
if event_type == "error":
|
||||
raise RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR")
|
||||
if event_type == "interaction.complete":
|
||||
interaction = payload.get("interaction") or {}
|
||||
ordered_blocks = [blocks[index] for index in sorted(blocks)]
|
||||
yield ChatGenerationChunk(
|
||||
message=AIMessageChunk(
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"gemini_interaction_content": ordered_blocks
|
||||
},
|
||||
usage_metadata=_usage_metadata(
|
||||
interaction.get("usage"),
|
||||
provider_request_id=interaction.get("id"),
|
||||
),
|
||||
response_metadata={
|
||||
"model_name": str(
|
||||
(interaction.get("model") or {}).get("id") or ""
|
||||
),
|
||||
"finish_reason": str(
|
||||
interaction.get("status") or "unknown"
|
||||
),
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _message_content(message: BaseMessage) -> list[dict[str, Any]]:
|
||||
if isinstance(message, AIMessage):
|
||||
preserved = message.additional_kwargs.get("gemini_interaction_content")
|
||||
if isinstance(preserved, list):
|
||||
return [dict(item) for item in preserved]
|
||||
content = message.content
|
||||
if isinstance(content, str):
|
||||
blocks: list[dict[str, Any]] = [{"type": "text", "text": content}]
|
||||
elif isinstance(content, list):
|
||||
blocks = []
|
||||
for item in content:
|
||||
if isinstance(item, str):
|
||||
blocks.append({"type": "text", "text": item})
|
||||
elif isinstance(item, dict) and item.get("type") in {
|
||||
"text",
|
||||
"thought",
|
||||
"function_call",
|
||||
"function_result",
|
||||
"image",
|
||||
"audio",
|
||||
"video",
|
||||
"document",
|
||||
}:
|
||||
blocks.append(dict(item))
|
||||
else:
|
||||
raise ValueError("MODEL_CONTENT_BLOCK_UNSUPPORTED")
|
||||
else:
|
||||
blocks = [{"type": "text", "text": str(content)}]
|
||||
if isinstance(message, AIMessage):
|
||||
for call in message.tool_calls:
|
||||
blocks.append(
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": str(call["id"]),
|
||||
"name": str(call["name"]),
|
||||
"arguments": dict(call.get("args") or {}),
|
||||
}
|
||||
)
|
||||
return blocks
|
||||
|
||||
|
||||
def _compile_messages(
|
||||
messages: Sequence[BaseMessage],
|
||||
) -> tuple[list[dict[str, Any]], str]:
|
||||
turns: list[dict[str, Any]] = []
|
||||
system_parts: list[str] = []
|
||||
for message in messages:
|
||||
if isinstance(message, SystemMessage):
|
||||
system_parts.append(str(message.content))
|
||||
continue
|
||||
if isinstance(message, ToolMessage):
|
||||
turns.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "function_result",
|
||||
"call_id": str(message.tool_call_id),
|
||||
"name": str(getattr(message, "name", "") or ""),
|
||||
"result": message.content,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
continue
|
||||
role = "model" if isinstance(message, AIMessage) else "user"
|
||||
if not isinstance(message, (AIMessage, HumanMessage)):
|
||||
role = "user"
|
||||
turns.append({"role": role, "content": _message_content(message)})
|
||||
return turns, "\n\n".join(system_parts)
|
||||
|
||||
|
||||
def _chat_result(response: Any) -> ChatResult:
|
||||
blocks = [
|
||||
item.model_dump(mode="json", by_alias=True, exclude_none=True)
|
||||
if hasattr(item, "model_dump")
|
||||
else dict(item)
|
||||
for item in (getattr(response, "outputs", None) or [])
|
||||
]
|
||||
tool_calls = [
|
||||
{
|
||||
"id": str(item.get("id") or ""),
|
||||
"name": str(item.get("name") or ""),
|
||||
"args": dict(item.get("arguments") or {}),
|
||||
"type": "tool_call",
|
||||
}
|
||||
for item in blocks
|
||||
if item.get("type") == "function_call"
|
||||
]
|
||||
usage_metadata = _usage_metadata(
|
||||
_dump(getattr(response, "usage", None)),
|
||||
provider_request_id=getattr(response, "id", None),
|
||||
)
|
||||
message = AIMessage(
|
||||
content=blocks,
|
||||
tool_calls=tool_calls,
|
||||
additional_kwargs={"gemini_interaction_content": blocks},
|
||||
usage_metadata=usage_metadata,
|
||||
response_metadata={
|
||||
"model_name": str(getattr(getattr(response, "model", None), "id", "")),
|
||||
"finish_reason": str(getattr(response, "status", "unknown")),
|
||||
},
|
||||
)
|
||||
return ChatResult(generations=[ChatGeneration(message=message)])
|
||||
|
||||
|
||||
def _dump(value: Any) -> dict[str, Any]:
|
||||
if value is None:
|
||||
return {}
|
||||
if hasattr(value, "model_dump"):
|
||||
return value.model_dump(mode="json", by_alias=True, exclude_none=True)
|
||||
if isinstance(value, Mapping):
|
||||
return dict(value)
|
||||
return {}
|
||||
|
||||
|
||||
def _merge_stream_delta(block: dict[str, Any], delta: Mapping[str, Any]) -> None:
|
||||
kind = str(delta.get("type") or "")
|
||||
if kind == "text":
|
||||
block["type"] = "text"
|
||||
block["text"] = str(block.get("text") or "") + str(delta.get("text") or "")
|
||||
elif kind == "thought_signature":
|
||||
block.setdefault("type", "thought")
|
||||
block["signature"] = delta.get("signature")
|
||||
elif kind == "thought_summary":
|
||||
block.setdefault("type", "thought")
|
||||
content = delta.get("content")
|
||||
if content is not None:
|
||||
block.setdefault("summary", []).append(content)
|
||||
elif kind == "text_annotation":
|
||||
block.setdefault("annotations", []).extend(delta.get("annotations") or [])
|
||||
else:
|
||||
block.update(delta)
|
||||
|
||||
|
||||
def _usage_metadata(
|
||||
usage: Mapping[str, Any] | None, *, provider_request_id: Any
|
||||
) -> dict[str, Any] | None:
|
||||
if not usage:
|
||||
return None
|
||||
required = (
|
||||
usage.get("total_input_tokens"),
|
||||
usage.get("total_cached_tokens"),
|
||||
usage.get("total_output_tokens"),
|
||||
)
|
||||
if any(value is None for value in required):
|
||||
return None
|
||||
result: dict[str, Any] = {
|
||||
"input_tokens": int(required[0]),
|
||||
"cached_input_tokens": int(required[1]),
|
||||
"output_tokens": int(required[2]),
|
||||
"total_tokens": int(
|
||||
usage.get("total_tokens")
|
||||
if usage.get("total_tokens") is not None
|
||||
else int(required[0]) + int(required[2])
|
||||
),
|
||||
"usage_finality": "confirmed",
|
||||
}
|
||||
optional = {
|
||||
"reasoning_tokens": usage.get("total_thought_tokens"),
|
||||
"provider_request_id": provider_request_id,
|
||||
}
|
||||
result.update({key: value for key, value in optional.items() if value is not None})
|
||||
return result
|
||||
|
||||
|
||||
def create_gemini_interactions_model(**kwargs: Any) -> GeminiInteractionsChatModel:
|
||||
return GeminiInteractionsChatModel(
|
||||
model_name=str(kwargs["model"]),
|
||||
api_key=SecretStr(str(kwargs["api_key"])),
|
||||
base_url=str(
|
||||
kwargs.get("base_url") or "https://generativelanguage.googleapis.com"
|
||||
),
|
||||
max_output_tokens=int(kwargs.get("max_output_tokens") or 8192),
|
||||
temperature=kwargs.get("temperature"),
|
||||
top_p=kwargs.get("top_p"),
|
||||
thinking=kwargs.get("thinking"),
|
||||
)
|
||||
@@ -0,0 +1,20 @@
|
||||
"""Compilation of normalized configuration into provider invocation plans."""
|
||||
|
||||
from .contract import (
|
||||
InvocationPlan,
|
||||
ToolCallTransport,
|
||||
compile_invocation_plan,
|
||||
derive_runtime_invocation,
|
||||
derive_tool_call_transport,
|
||||
)
|
||||
from .messages import assistant_message_has_output, project_provider_messages
|
||||
|
||||
__all__ = [
|
||||
"InvocationPlan",
|
||||
"ToolCallTransport",
|
||||
"assistant_message_has_output",
|
||||
"compile_invocation_plan",
|
||||
"derive_runtime_invocation",
|
||||
"derive_tool_call_transport",
|
||||
"project_provider_messages",
|
||||
]
|
||||
@@ -0,0 +1,133 @@
|
||||
"""Pure invocation contract shared by configuration and runtime.
|
||||
|
||||
Provider configuration owns connection/authentication and adapter selection.
|
||||
Model configuration owns capabilities, limits, and canonical parameters. This
|
||||
module is the only layer that derives the effective wire invocation from those
|
||||
inputs; the derived fields are runtime state, not administrator configuration.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Literal
|
||||
|
||||
from ..contracts import EvoRuntimeError
|
||||
from ..crypto import canonical_json_v1
|
||||
|
||||
ToolCallTransport = Literal["native", "disabled"]
|
||||
|
||||
|
||||
def derive_tool_call_transport(
|
||||
capabilities: Mapping[str, Any],
|
||||
) -> ToolCallTransport:
|
||||
"""Derive the wire transport from the authoritative tools capability."""
|
||||
|
||||
return "native" if bool(capabilities.get("tools", False)) else "disabled"
|
||||
|
||||
|
||||
def derive_runtime_invocation(
|
||||
api_mode: str,
|
||||
capabilities: Mapping[str, Any],
|
||||
) -> dict[str, str]:
|
||||
"""Build the internal invocation projection for a normalized model."""
|
||||
|
||||
return {
|
||||
"api_mode": str(api_mode),
|
||||
"tool_call_transport": derive_tool_call_transport(capabilities),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InvocationPlan:
|
||||
"""Complete, immutable, non-secret wire contract for one model call."""
|
||||
|
||||
api_mode: str
|
||||
output_token_parameter: str
|
||||
output_token_limit: int
|
||||
tool_call_transport: ToolCallTransport
|
||||
reasoning_effort: str
|
||||
streaming: bool
|
||||
sdk_params: Mapping[str, Any]
|
||||
plan_hash: str
|
||||
|
||||
def model_kwargs(self) -> dict[str, Any]:
|
||||
return dict(self.sdk_params)
|
||||
|
||||
def projection(self) -> dict[str, Any]:
|
||||
return {
|
||||
"api_mode": self.api_mode,
|
||||
"output_token_parameter": self.output_token_parameter,
|
||||
"output_token_limit": self.output_token_limit,
|
||||
"tool_call_transport": self.tool_call_transport,
|
||||
"reasoning_effort": self.reasoning_effort,
|
||||
"streaming": self.streaming,
|
||||
"sdk_params": dict(self.sdk_params),
|
||||
"plan_hash": self.plan_hash,
|
||||
}
|
||||
|
||||
|
||||
def compile_invocation_plan(
|
||||
*,
|
||||
api_mode: str,
|
||||
declared_tool_call_transport: str,
|
||||
supports_tools: bool,
|
||||
purpose: str,
|
||||
output_token_limit: int,
|
||||
reasoning_effort: str,
|
||||
runtime_provider: str,
|
||||
sdk_params: Mapping[str, Any],
|
||||
) -> InvocationPlan:
|
||||
"""Validate and freeze adapter output before constructing a provider SDK."""
|
||||
|
||||
params = dict(sdk_params)
|
||||
streaming = purpose == "main_agent"
|
||||
params["streaming"] = streaming
|
||||
if runtime_provider == "openai" and streaming:
|
||||
params["stream_usage"] = True
|
||||
token_fields = tuple(
|
||||
key
|
||||
for key in ("max_output_tokens", "max_completion_tokens", "max_tokens")
|
||||
if key in params
|
||||
)
|
||||
if len(token_fields) != 1 or params[token_fields[0]] != output_token_limit:
|
||||
raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED")
|
||||
output_token_parameter = token_fields[0]
|
||||
if api_mode == "responses" and output_token_parameter != "max_output_tokens":
|
||||
raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED")
|
||||
if api_mode == "chat_completions" and output_token_parameter == "max_output_tokens":
|
||||
raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED")
|
||||
if runtime_provider == "openai" and api_mode in {
|
||||
"responses",
|
||||
"chat_completions",
|
||||
}:
|
||||
if params.get("use_responses_api") is not (api_mode == "responses"):
|
||||
raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED")
|
||||
|
||||
tool_call_transport: ToolCallTransport = (
|
||||
"native" if supports_tools else "disabled"
|
||||
)
|
||||
if declared_tool_call_transport != tool_call_transport:
|
||||
raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED")
|
||||
|
||||
plan_payload = {
|
||||
"api_mode": api_mode,
|
||||
"output_token_parameter": output_token_parameter,
|
||||
"output_token_limit": output_token_limit,
|
||||
"tool_call_transport": tool_call_transport,
|
||||
"reasoning_effort": reasoning_effort,
|
||||
"streaming": streaming,
|
||||
"sdk_params": params,
|
||||
}
|
||||
return InvocationPlan(
|
||||
api_mode=api_mode,
|
||||
output_token_parameter=output_token_parameter,
|
||||
output_token_limit=output_token_limit,
|
||||
tool_call_transport=tool_call_transport,
|
||||
reasoning_effort=reasoning_effort,
|
||||
streaming=streaming,
|
||||
sdk_params=MappingProxyType(params),
|
||||
plan_hash=hashlib.sha256(canonical_json_v1(plan_payload)).hexdigest(),
|
||||
)
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Provider-facing message projection and response validation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
_TOOL_BLOCK_TYPES = frozenset({"tool_call", "tool_use", "function_call"})
|
||||
_NON_FINAL_BLOCK_TYPES = frozenset({"reasoning", "thinking"})
|
||||
|
||||
|
||||
def assistant_message_has_output(message: AIMessage) -> bool:
|
||||
"""Return whether an assistant message has provider-visible output."""
|
||||
|
||||
if (
|
||||
getattr(message, "tool_calls", None)
|
||||
or getattr(message, "invalid_tool_calls", None)
|
||||
or _additional_tool_calls(message)
|
||||
):
|
||||
return True
|
||||
content = getattr(message, "content", None)
|
||||
if isinstance(content, str):
|
||||
return bool(content.strip())
|
||||
if not isinstance(content, Sequence) or isinstance(content, str | bytes):
|
||||
return content is not None
|
||||
for block in content:
|
||||
if isinstance(block, str):
|
||||
if block.strip():
|
||||
return True
|
||||
continue
|
||||
if not isinstance(block, Mapping):
|
||||
return True
|
||||
block_type = str(block.get("type") or "").strip()
|
||||
if block_type in _TOOL_BLOCK_TYPES:
|
||||
return True
|
||||
if block_type in _NON_FINAL_BLOCK_TYPES:
|
||||
continue
|
||||
text = block.get("text")
|
||||
if isinstance(text, str):
|
||||
if text.strip():
|
||||
return True
|
||||
continue
|
||||
# Unknown non-reasoning blocks may carry multimodal or refusal output.
|
||||
if block:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def project_provider_messages(messages: Sequence[Any]) -> tuple[list[Any], int]:
|
||||
"""Drop unusable assistant history without mutating checkpoint objects."""
|
||||
|
||||
projected: list[Any] = []
|
||||
dropped = 0
|
||||
for message in messages:
|
||||
if isinstance(message, AIMessage) and not assistant_message_has_output(message):
|
||||
dropped += 1
|
||||
continue
|
||||
projected.append(message)
|
||||
return projected, dropped
|
||||
|
||||
|
||||
def _additional_tool_calls(message: AIMessage) -> bool:
|
||||
additional = getattr(message, "additional_kwargs", None)
|
||||
if not isinstance(additional, Mapping):
|
||||
return False
|
||||
return bool(additional.get("tool_calls") or additional.get("function_call"))
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+24
-83
@@ -84,11 +84,6 @@ def _resolve_codex_client_version() -> str:
|
||||
return _CODEX_CLIENT_VERSION_FALLBACK
|
||||
|
||||
|
||||
def _resolve_reasoning_effort(default: str) -> str:
|
||||
"""Return the configured reasoning effort or a provider-specific default."""
|
||||
return os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or default
|
||||
|
||||
|
||||
# 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]] = {
|
||||
@@ -439,8 +434,10 @@ def _apply_auto_config(
|
||||
)
|
||||
else "high"
|
||||
)
|
||||
_eff = _resolve_reasoning_effort(_default_effort)
|
||||
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
|
||||
# An explicit API envelope belongs to the compiled invocation plan.
|
||||
# Do not add a legacy Responses-style reasoning object to a Chat plan.
|
||||
if "use_responses_api" not in kwargs:
|
||||
kwargs["reasoning"] = {"effort": _default_effort, "summary": "auto"}
|
||||
|
||||
# Google GenAI: surface thinking traces
|
||||
if provider == "google-genai" and not disable_reasoning:
|
||||
@@ -474,46 +471,7 @@ def get_chat_model(
|
||||
>>> model = get_chat_model("gpt-4o") # OpenAI model
|
||||
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
|
||||
"""
|
||||
skip_runtime_resolver = bool(kwargs.pop("_skip_runtime_model_resolver", False))
|
||||
runtime_provider_name: str | None = None
|
||||
runtime_supports_reasoning: bool | None = None
|
||||
runtime_resolved = None
|
||||
if not skip_runtime_resolver:
|
||||
from EvoScientist.runtime_integrations import resolve_runtime_model
|
||||
|
||||
runtime_resolved = resolve_runtime_model(model, provider)
|
||||
|
||||
if runtime_resolved is not None:
|
||||
resolved_params = dict(getattr(runtime_resolved, "params", {}) or {})
|
||||
extra_body = resolved_params.pop("_extra_body", None)
|
||||
default_headers = resolved_params.pop("_default_headers", None)
|
||||
if extra_body:
|
||||
resolved_params["extra_body"] = extra_body
|
||||
if default_headers:
|
||||
resolved_params["default_headers"] = default_headers
|
||||
resolved_params.update(kwargs)
|
||||
kwargs = resolved_params
|
||||
|
||||
resolved_api_key = str(getattr(runtime_resolved, "api_key", "") or "")
|
||||
resolved_base_url = str(getattr(runtime_resolved, "base_url", "") or "")
|
||||
if resolved_api_key:
|
||||
kwargs.setdefault("api_key", resolved_api_key)
|
||||
if resolved_base_url:
|
||||
kwargs.setdefault("base_url", resolved_base_url.rstrip("/"))
|
||||
|
||||
runtime_provider_name = str(
|
||||
getattr(runtime_resolved, "provider_name", "") or ""
|
||||
)
|
||||
runtime_supports_reasoning = bool(
|
||||
getattr(runtime_resolved, "supports_reasoning", False)
|
||||
)
|
||||
if not runtime_supports_reasoning:
|
||||
kwargs.setdefault("_disable_reasoning", True)
|
||||
kwargs.setdefault("_disable_thinking", True)
|
||||
model = str(runtime_resolved.model_id)
|
||||
provider = str(runtime_resolved.protocol)
|
||||
else:
|
||||
model = model or DEFAULT_MODEL
|
||||
model = model or DEFAULT_MODEL
|
||||
|
||||
# Look up short name in registry (provider-aware)
|
||||
model_id = None
|
||||
@@ -548,19 +506,15 @@ def get_chat_model(
|
||||
_is_third_party = (
|
||||
provider in _OPENAI_ROUTED_PROVIDERS or provider in _ANTHROPIC_ROUTED_PROVIDERS
|
||||
)
|
||||
if runtime_provider_name and runtime_provider_name != provider:
|
||||
_is_third_party = True
|
||||
explicit_base_url = str(kwargs.get("base_url") or "")
|
||||
if (
|
||||
runtime_resolved is not None
|
||||
and provider == "openai"
|
||||
and resolved_base_url
|
||||
and "api.openai.com" not in resolved_base_url.lower()
|
||||
provider == "openai"
|
||||
and explicit_base_url
|
||||
and "api.openai.com" not in explicit_base_url.lower()
|
||||
):
|
||||
_is_third_party = True
|
||||
_is_openai_proxy = False
|
||||
_original_provider: str | None = (
|
||||
runtime_provider_name if runtime_provider_name != provider else None
|
||||
)
|
||||
_original_provider: str | None = None
|
||||
if provider == "anthropic":
|
||||
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
|
||||
if base_url:
|
||||
@@ -578,17 +532,6 @@ def get_chat_model(
|
||||
kwargs.get("base_url"), kwargs.get("api_key")
|
||||
)
|
||||
if _is_openai_proxy:
|
||||
# Use Responses API for ccproxy: bypasses the format chain
|
||||
# converter (Chat→Responses→Chat) which returns 502 on
|
||||
# complex responses. System messages are converted to
|
||||
# developer role by _patch_ccproxy_system_to_developer().
|
||||
kwargs.setdefault("use_responses_api", True)
|
||||
# Streaming must stay ON for Responses API: ccproxy's
|
||||
# StreamingBufferService loses output when assembling
|
||||
# non-streaming responses. (The old streaming=False was
|
||||
# for Chat Completions tool_call duplication — not an issue
|
||||
# with the Responses API SSE format.)
|
||||
kwargs.pop("streaming", None) # remove if set elsewhere
|
||||
# ccproxy forwards client headers upstream and only
|
||||
# gap-fills its own, so the Codex backend sees this
|
||||
# client's identity. Without Codex-CLI-shaped headers it
|
||||
@@ -650,8 +593,7 @@ def get_chat_model(
|
||||
# passback (OpenRouter's `/responses` beta is stateless, store=false —
|
||||
# "Item with id 'rs_...' not found"); the patch strips them on passback,
|
||||
# so enabling `summary` is safe. See langchain-ai/langchain#37777.
|
||||
effort = _resolve_reasoning_effort("high")
|
||||
kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"})
|
||||
kwargs.setdefault("reasoning", {"effort": "high", "summary": "auto"})
|
||||
# App attribution (issue #339): identify EvoScientist to OpenRouter so
|
||||
# usage is credited to the project (app rankings, model app tabs,
|
||||
# analytics) rather than langchain-openrouter's LangChain-branded
|
||||
@@ -736,19 +678,6 @@ def get_chat_model(
|
||||
_apply_auto_config(provider, model_id, _is_third_party, kwargs, _original_provider)
|
||||
_apply_openrouter_anthropic_prompt_cache(provider, model_id, kwargs)
|
||||
|
||||
# User-level override for the OpenAI Responses API vs Chat Completions.
|
||||
# When "false", force Chat Completions and drop reasoning (which triggers
|
||||
# the Responses API path in langchain-openai). Only applies to OpenAI.
|
||||
if provider == "openai":
|
||||
_responses_api_setting = (
|
||||
os.environ.get("EVOSCIENTIST_USE_RESPONSES_API", "").strip().lower()
|
||||
)
|
||||
if _responses_api_setting == "false":
|
||||
kwargs["use_responses_api"] = False
|
||||
kwargs.pop("reasoning", None)
|
||||
elif _responses_api_setting == "true":
|
||||
kwargs["use_responses_api"] = True
|
||||
|
||||
anthropic_auth_token = None
|
||||
if provider == "anthropic" and kwargs.get("api_key"):
|
||||
anthropic_auth_token = os.environ.pop("ANTHROPIC_AUTH_TOKEN", None)
|
||||
@@ -769,7 +698,19 @@ def get_chat_model(
|
||||
# Anthropic-routed providers accept media in tool results natively;
|
||||
# only OpenAI-compatible providers need tool-media hoisting.
|
||||
_hoist = _original_provider not in _ANTHROPIC_ROUTED_PROVIDERS
|
||||
_patch_openai_compat_content(chat_model, hoist_tool_media=_hoist)
|
||||
_patch_openai_compat_content(
|
||||
chat_model,
|
||||
hoist_tool_media=_hoist,
|
||||
# Generic OpenAI-compatible proxies must not receive hidden
|
||||
# reasoning traces emitted by a different provider. DeepSeek has
|
||||
# its own explicit passback patch below, so preserve that path.
|
||||
drop_reasoning_metadata=(
|
||||
_is_third_party
|
||||
and provider == "openai"
|
||||
and _original_provider is None
|
||||
and not _is_openai_proxy
|
||||
),
|
||||
)
|
||||
|
||||
# DeepSeek thinking mode requires reasoning_content passback in multi-turn
|
||||
# + tool_use scenarios.
|
||||
|
||||
@@ -11,6 +11,8 @@ Patches:
|
||||
- _patch_openai_capture_reasoning_content: capture provider
|
||||
reasoning_content into AIMessage.additional_kwargs (module-level,
|
||||
applied at import)
|
||||
- _patch_openai_empty_sse_keepalive: ignore blank SSE keepalive events
|
||||
emitted by some OpenAI-compatible Responses endpoints
|
||||
- _patch_deepseek_reasoning_passback: re-inject reasoning_content into
|
||||
outgoing DeepSeek assistant messages for thinking-mode multi-turn /
|
||||
tool_use scenarios
|
||||
@@ -70,6 +72,65 @@ def _patch_anthropic_proxy_compat() -> None:
|
||||
_patch_anthropic_proxy_compat()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Patch (module-level): some OpenAI-compatible endpoints emit SSE keepalive
|
||||
# frames with an empty ``data`` field. OpenAI SDK 2.x unconditionally passes
|
||||
# every event to ``json.loads``, so an otherwise harmless keepalive becomes a
|
||||
# JSONDecodeError and aborts the stream. Filter only blank data frames before
|
||||
# the SDK's parser; JSON events, error events, and [DONE] are unchanged.
|
||||
# ---------------------------------------------------------------------------
|
||||
_openai_empty_sse_keepalive_patched = False
|
||||
|
||||
|
||||
def _is_blank_sse_keepalive(event: Any) -> bool:
|
||||
"""Return whether an SSE event has no JSON payload to parse."""
|
||||
|
||||
data = getattr(event, "data", None)
|
||||
return data is None or (isinstance(data, str) and not data.strip())
|
||||
|
||||
|
||||
def _patch_openai_empty_sse_keepalive() -> None:
|
||||
global _openai_empty_sse_keepalive_patched
|
||||
if _openai_empty_sse_keepalive_patched:
|
||||
return
|
||||
try:
|
||||
import functools
|
||||
|
||||
from openai._streaming import AsyncStream as _AsyncStream
|
||||
from openai._streaming import Stream as _Stream
|
||||
|
||||
original_async = _AsyncStream._iter_events
|
||||
if not getattr(original_async, "_evoscientist_skips_blank_sse", False):
|
||||
|
||||
@functools.wraps(original_async)
|
||||
async def _filtered_async_events(self: Any) -> Any:
|
||||
async for event in original_async(self):
|
||||
if not _is_blank_sse_keepalive(event):
|
||||
yield event
|
||||
|
||||
_filtered_async_events._evoscientist_skips_blank_sse = True # type: ignore[attr-defined]
|
||||
_AsyncStream._iter_events = _filtered_async_events
|
||||
|
||||
original_sync = _Stream._iter_events
|
||||
if not getattr(original_sync, "_evoscientist_skips_blank_sse", False):
|
||||
|
||||
@functools.wraps(original_sync)
|
||||
def _filtered_sync_events(self: Any) -> Any:
|
||||
for event in original_sync(self):
|
||||
if not _is_blank_sse_keepalive(event):
|
||||
yield event
|
||||
|
||||
_filtered_sync_events._evoscientist_skips_blank_sse = True # type: ignore[attr-defined]
|
||||
_Stream._iter_events = _filtered_sync_events
|
||||
_openai_empty_sse_keepalive_patched = True
|
||||
except Exception:
|
||||
# The patch is only needed when the optional OpenAI SDK is available.
|
||||
pass
|
||||
|
||||
|
||||
_patch_openai_empty_sse_keepalive()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Patch: ccproxy-api 0.2.7 Codex compatibility.
|
||||
#
|
||||
@@ -207,6 +268,9 @@ def _is_ccproxy_codex(
|
||||
# preserved, not flattened away.
|
||||
# ---------------------------------------------------------------------------
|
||||
_SKIP_CONTENT_TYPES = frozenset({"thinking", "reasoning", "reasoning_content"})
|
||||
_NONPORTABLE_REASONING_METADATA = frozenset(
|
||||
{"reasoning_content", "reasoning_details"}
|
||||
)
|
||||
|
||||
# Media block types preserved when flattening (positive allowlist;
|
||||
# thinking/reasoning still dropped). Images + files (PDF/documents): both
|
||||
@@ -407,6 +471,8 @@ def _copy_ai_message_with_tool_pairs(
|
||||
# LangChain content blocks use id; the Responses converter later
|
||||
# maps it to call_id.
|
||||
block["id"] = call_id
|
||||
if "call_id" in block:
|
||||
block["call_id"] = call_id
|
||||
block["name"] = call_name
|
||||
if isinstance(block.get("function"), dict):
|
||||
block["function"] = {**block["function"], "name": call_name}
|
||||
@@ -540,7 +606,7 @@ def _validate_openai_tool_history(messages: list[Any]) -> None:
|
||||
}:
|
||||
continue
|
||||
block_id = str(
|
||||
block.get("id") or block.get("call_id") or ""
|
||||
block.get("call_id") or block.get("id") or ""
|
||||
).strip()
|
||||
block_name = block.get("name") or block.get("tool_name")
|
||||
function = block.get("function")
|
||||
@@ -566,7 +632,12 @@ def _validate_openai_tool_history(messages: list[Any]) -> None:
|
||||
raise ValueError("assistant tool call is missing its tool result")
|
||||
|
||||
|
||||
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
|
||||
def _sanitize_messages(
|
||||
messages: list[Any],
|
||||
hoist_tool_media: bool = True,
|
||||
*,
|
||||
drop_reasoning_metadata: bool = False,
|
||||
) -> list[Any]:
|
||||
"""Flatten list content for OpenAI-compatible APIs, preserving media.
|
||||
|
||||
Text/reasoning content is flattened to a string; image blocks are
|
||||
@@ -593,6 +664,15 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
|
||||
pending_media.clear()
|
||||
|
||||
for msg in messages:
|
||||
if drop_reasoning_metadata:
|
||||
additional_kwargs = getattr(msg, "additional_kwargs", None) or {}
|
||||
if set(additional_kwargs) & _NONPORTABLE_REASONING_METADATA:
|
||||
msg = copy.copy(msg)
|
||||
msg.additional_kwargs = {
|
||||
key: value
|
||||
for key, value in additional_kwargs.items()
|
||||
if key not in _NONPORTABLE_REASONING_METADATA
|
||||
}
|
||||
is_tool = getattr(msg, "type", None) == "tool"
|
||||
if not is_tool:
|
||||
_flush() # emit hoisted media before any non-tool message
|
||||
@@ -714,7 +794,12 @@ def _strip_media_types(messages: list[Any], types: set[str]) -> list[Any]:
|
||||
return out
|
||||
|
||||
|
||||
def _patch_openai_compat_content(model: Any, hoist_tool_media: bool = True) -> None:
|
||||
def _patch_openai_compat_content(
|
||||
model: Any,
|
||||
hoist_tool_media: bool = True,
|
||||
*,
|
||||
drop_reasoning_metadata: bool = False,
|
||||
) -> None:
|
||||
"""Flatten list content to strings before OpenAI-compatible API calls.
|
||||
|
||||
Wraps ``_generate`` / ``_agenerate`` / ``_stream`` / ``_astream`` to prevent
|
||||
@@ -749,11 +834,17 @@ def _patch_openai_compat_content(model: Any, hoist_tool_media: bool = True) -> N
|
||||
|
||||
def _prepare(messages: list[BaseMessage]) -> list[BaseMessage]:
|
||||
msgs = _strip_media_types(messages, blocked) if blocked else messages
|
||||
return _sanitize_messages(msgs, hoist_tool_media)
|
||||
return _sanitize_messages(
|
||||
msgs,
|
||||
hoist_tool_media,
|
||||
drop_reasoning_metadata=drop_reasoning_metadata,
|
||||
)
|
||||
|
||||
def _stripped(messages: list[BaseMessage], suspects: set[str]) -> list[BaseMessage]:
|
||||
return _sanitize_messages(
|
||||
_strip_media_types(messages, blocked | suspects), hoist_tool_media
|
||||
_strip_media_types(messages, blocked | suspects),
|
||||
hoist_tool_media,
|
||||
drop_reasoning_metadata=drop_reasoning_metadata,
|
||||
)
|
||||
|
||||
orig_generate = getattr(model, "_generate", None)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,472 @@
|
||||
"""Encrypted, versioned storage for model-provider credentials.
|
||||
|
||||
Only secret references cross the model-route configuration boundary. The
|
||||
store deliberately has no API for reading a plaintext credential after it was
|
||||
written; callers receive an opaque reference and masked metadata instead.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import os
|
||||
import re
|
||||
import sqlite3
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from cryptography.fernet import Fernet, InvalidToken
|
||||
from filelock import FileLock
|
||||
|
||||
from ..config.settings import get_config_dir
|
||||
from .configuration import ResolvedSecret, SecretReference
|
||||
from .contracts import EvoRuntimeError
|
||||
|
||||
_SECRET_ID_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._/-]{0,127}$")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SecretMetadata:
|
||||
secret_id: str
|
||||
version: int
|
||||
masked_value: str
|
||||
created_by: str
|
||||
created_at: str
|
||||
status: str = "active"
|
||||
retired_at: str | None = None
|
||||
revoked_at: str | None = None
|
||||
revoked_by: str | None = None
|
||||
revoke_reason: str | None = None
|
||||
|
||||
@property
|
||||
def ref(self) -> str:
|
||||
return f"secret://{self.secret_id}#{self.version}"
|
||||
|
||||
|
||||
class EncryptedModelSecretStore:
|
||||
"""SQLite-backed Fernet store scoped to the Evo configuration directory."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: Path | None = None,
|
||||
*,
|
||||
master_secret: str | None = None,
|
||||
) -> None:
|
||||
self.path = path or (get_config_dir() / "model_secrets.sqlite")
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._lock = FileLock(str(self.path) + ".lock")
|
||||
material = master_secret or os.environ.get(
|
||||
"AI4SCI_EVO_MODEL_SECRET_MASTER_KEY", ""
|
||||
)
|
||||
if not material:
|
||||
# The identity secret is already mandatory for the V3 runtime.
|
||||
# Operators can configure a dedicated secret to separate rotation.
|
||||
material = os.environ.get("AI4SCI_EVO_CONFIG_IDENTITY_SECRET", "")
|
||||
if len(material.encode("utf-8")) < 32:
|
||||
raise RuntimeError(
|
||||
"AI4SCI_EVO_MODEL_SECRET_MASTER_KEY or "
|
||||
"AI4SCI_EVO_CONFIG_IDENTITY_SECRET must contain at least 32 bytes"
|
||||
)
|
||||
key = base64.urlsafe_b64encode(hashlib.sha256(material.encode()).digest())
|
||||
self._fernet = Fernet(key)
|
||||
self._init_schema()
|
||||
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
connection = sqlite3.connect(self.path)
|
||||
connection.row_factory = sqlite3.Row
|
||||
connection.execute("PRAGMA foreign_keys=ON")
|
||||
connection.execute("PRAGMA busy_timeout=5000")
|
||||
connection.execute("PRAGMA synchronous=FULL")
|
||||
connection.execute("PRAGMA journal_mode=WAL")
|
||||
return connection
|
||||
|
||||
def _init_schema(self) -> None:
|
||||
with self._lock, self._connect() as connection:
|
||||
connection.executescript(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS model_secret_versions (
|
||||
secret_id TEXT NOT NULL,
|
||||
version INTEGER NOT NULL CHECK (version > 0),
|
||||
ciphertext BLOB NOT NULL,
|
||||
masked_value TEXT NOT NULL,
|
||||
created_by TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
PRIMARY KEY (secret_id, version)
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_model_secret_versions_latest
|
||||
ON model_secret_versions (secret_id, version DESC);
|
||||
"""
|
||||
)
|
||||
columns = {
|
||||
str(row["name"])
|
||||
for row in connection.execute(
|
||||
"PRAGMA table_info(model_secret_versions)"
|
||||
).fetchall()
|
||||
}
|
||||
additions = {
|
||||
"status": "TEXT NOT NULL DEFAULT 'active'",
|
||||
"retired_at": "TEXT",
|
||||
"revoked_at": "TEXT",
|
||||
"revoked_by": "TEXT",
|
||||
"revoke_reason": "TEXT",
|
||||
"transition_operation_id": "TEXT",
|
||||
}
|
||||
for name, definition in additions.items():
|
||||
if name not in columns:
|
||||
connection.execute(
|
||||
f"ALTER TABLE model_secret_versions ADD COLUMN {name} {definition}"
|
||||
)
|
||||
# Legacy stores treated every version as active. Keep only the latest
|
||||
# active version before installing the partial unique index.
|
||||
connection.execute(
|
||||
"""UPDATE model_secret_versions AS current
|
||||
SET status='retired', retired_at=COALESCE(retired_at, CURRENT_TIMESTAMP)
|
||||
WHERE status='active' AND version < (
|
||||
SELECT MAX(newer.version) FROM model_secret_versions AS newer
|
||||
WHERE newer.secret_id=current.secret_id
|
||||
)"""
|
||||
)
|
||||
connection.execute(
|
||||
"""CREATE UNIQUE INDEX IF NOT EXISTS uq_model_secret_active
|
||||
ON model_secret_versions(secret_id) WHERE status='active'"""
|
||||
)
|
||||
connection.execute(
|
||||
"""CREATE TABLE IF NOT EXISTS model_secret_operations (
|
||||
operation_id TEXT PRIMARY KEY,
|
||||
action TEXT NOT NULL,
|
||||
request_digest TEXT NOT NULL,
|
||||
secret_id TEXT NOT NULL,
|
||||
version INTEGER NOT NULL,
|
||||
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
)"""
|
||||
)
|
||||
try:
|
||||
os.chmod(self.path, 0o600)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _validate_secret_id(secret_id: str) -> str:
|
||||
normalized = str(secret_id or "").strip()
|
||||
if not _SECRET_ID_RE.fullmatch(normalized):
|
||||
raise EvoRuntimeError("LLM_SECRET_INVALID")
|
||||
return normalized
|
||||
|
||||
@staticmethod
|
||||
def _mask(value: str) -> str:
|
||||
if len(value) <= 8:
|
||||
return "*" * len(value)
|
||||
return f"{value[:4]}...{value[-4:]}"
|
||||
|
||||
def put(
|
||||
self,
|
||||
secret_id: str,
|
||||
value: str,
|
||||
*,
|
||||
created_by: str,
|
||||
status: str = "active",
|
||||
operation_id: str | None = None,
|
||||
) -> SecretMetadata:
|
||||
secret_id = self._validate_secret_id(secret_id)
|
||||
value = str(value or "")
|
||||
if not value or any(char in value for char in "\r\n\0"):
|
||||
raise EvoRuntimeError("LLM_SECRET_INVALID")
|
||||
actor = str(created_by or "unknown")[:256]
|
||||
if status not in {"pending", "active"}:
|
||||
raise EvoRuntimeError("LLM_SECRET_INVALID")
|
||||
ciphertext = self._fernet.encrypt(value.encode("utf-8"))
|
||||
with self._lock, self._connect() as connection:
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
row = connection.execute(
|
||||
"SELECT COALESCE(MAX(version), 0) AS version "
|
||||
"FROM model_secret_versions WHERE secret_id=?",
|
||||
(secret_id,),
|
||||
).fetchone()
|
||||
version = int(row["version"]) + 1
|
||||
if status == "active":
|
||||
connection.execute(
|
||||
"""UPDATE model_secret_versions
|
||||
SET status='retired', retired_at=CURRENT_TIMESTAMP,
|
||||
transition_operation_id=?
|
||||
WHERE secret_id=? AND status='active'""",
|
||||
(operation_id, secret_id),
|
||||
)
|
||||
connection.execute(
|
||||
"""INSERT INTO model_secret_versions
|
||||
(secret_id, version, ciphertext, masked_value, created_by,
|
||||
status, transition_operation_id)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)""",
|
||||
(
|
||||
secret_id,
|
||||
version,
|
||||
ciphertext,
|
||||
self._mask(value),
|
||||
actor,
|
||||
status,
|
||||
operation_id,
|
||||
),
|
||||
)
|
||||
saved = connection.execute(
|
||||
"""SELECT secret_id, version, masked_value, created_by, created_at,
|
||||
status, retired_at, revoked_at, revoked_by, revoke_reason
|
||||
FROM model_secret_versions WHERE secret_id=? AND version=?""",
|
||||
(secret_id, version),
|
||||
).fetchone()
|
||||
return SecretMetadata(
|
||||
secret_id=str(saved["secret_id"]),
|
||||
version=int(saved["version"]),
|
||||
masked_value=str(saved["masked_value"]),
|
||||
created_by=str(saved["created_by"]),
|
||||
created_at=str(saved["created_at"]),
|
||||
status=str(saved["status"]),
|
||||
retired_at=saved["retired_at"],
|
||||
revoked_at=saved["revoked_at"],
|
||||
revoked_by=saved["revoked_by"],
|
||||
revoke_reason=saved["revoke_reason"],
|
||||
)
|
||||
|
||||
def create_pending(
|
||||
self,
|
||||
provider_id: str,
|
||||
value: str,
|
||||
*,
|
||||
created_by: str,
|
||||
operation_id: str,
|
||||
) -> SecretMetadata:
|
||||
secret_id = f"model-providers/{self._validate_secret_id(provider_id)}"
|
||||
digest = hashlib.sha256(
|
||||
(secret_id + "\0" + str(value)).encode("utf-8")
|
||||
).hexdigest()
|
||||
with self._lock:
|
||||
with self._connect() as connection:
|
||||
replay = connection.execute(
|
||||
"SELECT * FROM model_secret_operations WHERE operation_id=?",
|
||||
(operation_id,),
|
||||
).fetchone()
|
||||
if replay is not None:
|
||||
if (
|
||||
str(replay["action"]) != "create_pending"
|
||||
or str(replay["request_digest"]) != digest
|
||||
or str(replay["secret_id"]) != secret_id
|
||||
):
|
||||
raise EvoRuntimeError("IDEMPOTENCY_CONFLICT")
|
||||
return self._metadata(secret_id, int(replay["version"]))
|
||||
result = self.put(
|
||||
secret_id,
|
||||
value,
|
||||
created_by=created_by,
|
||||
status="pending",
|
||||
operation_id=operation_id,
|
||||
)
|
||||
with self._connect() as connection:
|
||||
connection.execute(
|
||||
"""INSERT INTO model_secret_operations
|
||||
(operation_id, action, request_digest, secret_id, version)
|
||||
VALUES (?, 'create_pending', ?, ?, ?)""",
|
||||
(operation_id, digest, secret_id, result.version),
|
||||
)
|
||||
return result
|
||||
|
||||
def list_metadata(self) -> list[SecretMetadata]:
|
||||
with self._connect() as connection:
|
||||
rows = connection.execute(
|
||||
"""SELECT secret_id, version, masked_value, created_by, created_at,
|
||||
status, retired_at, revoked_at, revoked_by, revoke_reason
|
||||
FROM model_secret_versions
|
||||
ORDER BY secret_id ASC, version DESC"""
|
||||
).fetchall()
|
||||
return [
|
||||
SecretMetadata(
|
||||
secret_id=str(row["secret_id"]),
|
||||
version=int(row["version"]),
|
||||
masked_value=str(row["masked_value"]),
|
||||
created_by=str(row["created_by"]),
|
||||
created_at=str(row["created_at"]),
|
||||
status=str(row["status"]),
|
||||
retired_at=row["retired_at"],
|
||||
revoked_at=row["revoked_at"],
|
||||
revoked_by=row["revoked_by"],
|
||||
revoke_reason=row["revoke_reason"],
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
|
||||
def current_provider_metadata(self, provider_id: str) -> SecretMetadata | None:
|
||||
secret_id = f"model-providers/{self._validate_secret_id(provider_id)}"
|
||||
with self._connect() as connection:
|
||||
row = connection.execute(
|
||||
"""SELECT secret_id, version, masked_value, created_by, created_at,
|
||||
status, retired_at, revoked_at, revoked_by, revoke_reason
|
||||
FROM model_secret_versions
|
||||
WHERE secret_id=? AND status='active'
|
||||
ORDER BY version DESC LIMIT 1""",
|
||||
(secret_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
return SecretMetadata(
|
||||
secret_id=str(row["secret_id"]),
|
||||
version=int(row["version"]),
|
||||
masked_value=str(row["masked_value"]),
|
||||
created_by=str(row["created_by"]),
|
||||
created_at=str(row["created_at"]),
|
||||
status=str(row["status"]),
|
||||
retired_at=row["retired_at"],
|
||||
revoked_at=row["revoked_at"],
|
||||
revoked_by=row["revoked_by"],
|
||||
revoke_reason=row["revoke_reason"],
|
||||
)
|
||||
|
||||
def resolve(self, reference: SecretReference) -> ResolvedSecret:
|
||||
if reference.ref.startswith("provider://"):
|
||||
provider_id = self._validate_secret_id(reference.ref[11:])
|
||||
secret_id = f"model-providers/{provider_id}"
|
||||
with self._connect() as connection:
|
||||
row = connection.execute(
|
||||
"""SELECT ciphertext, version FROM model_secret_versions
|
||||
WHERE secret_id=? AND status='active'
|
||||
ORDER BY version DESC LIMIT 1""",
|
||||
(secret_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE")
|
||||
try:
|
||||
value = self._fernet.decrypt(bytes(row["ciphertext"])).decode("utf-8")
|
||||
except (InvalidToken, UnicodeDecodeError) as exc:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc
|
||||
fingerprint = hashlib.sha256(value.encode("utf-8")).hexdigest()
|
||||
return ResolvedSecret(
|
||||
value, reference.revision, str(row["version"]), fingerprint
|
||||
)
|
||||
if not reference.ref.startswith("secret://") or "#" not in reference.ref:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE")
|
||||
secret_id, version_text = reference.ref[9:].rsplit("#", 1)
|
||||
try:
|
||||
version = int(version_text)
|
||||
except ValueError as exc:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc
|
||||
if version != reference.revision:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE")
|
||||
secret_id = self._validate_secret_id(secret_id)
|
||||
with self._connect() as connection:
|
||||
row = connection.execute(
|
||||
"""SELECT ciphertext, status FROM model_secret_versions
|
||||
WHERE secret_id=? AND version=?""",
|
||||
(secret_id, version),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE")
|
||||
if str(row["status"]) in {"revoked", "destroyed"}:
|
||||
raise EvoRuntimeError("MODEL_CREDENTIAL_REVOKED")
|
||||
try:
|
||||
value = self._fernet.decrypt(bytes(row["ciphertext"])).decode("utf-8")
|
||||
except (InvalidToken, UnicodeDecodeError) as exc:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc
|
||||
fingerprint = hashlib.sha256(value.encode("utf-8")).hexdigest()
|
||||
return ResolvedSecret(value, version, str(version), fingerprint)
|
||||
|
||||
def activate(
|
||||
self, secret_id: str, version: int, *, operation_id: str
|
||||
) -> SecretMetadata:
|
||||
secret_id = self._validate_secret_id(secret_id)
|
||||
with self._lock, self._connect() as connection:
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
row = connection.execute(
|
||||
"SELECT status FROM model_secret_versions WHERE secret_id=? AND version=?",
|
||||
(secret_id, version),
|
||||
).fetchone()
|
||||
if row is None or str(row["status"]) in {"revoked", "destroyed"}:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE")
|
||||
connection.execute(
|
||||
"""UPDATE model_secret_versions SET status='retired',
|
||||
retired_at=COALESCE(retired_at, CURRENT_TIMESTAMP),
|
||||
transition_operation_id=?
|
||||
WHERE secret_id=? AND status='active' AND version<>?""",
|
||||
(operation_id, secret_id, version),
|
||||
)
|
||||
connection.execute(
|
||||
"""UPDATE model_secret_versions SET status='active', retired_at=NULL,
|
||||
transition_operation_id=? WHERE secret_id=? AND version=?""",
|
||||
(operation_id, secret_id, version),
|
||||
)
|
||||
return self._metadata(secret_id, version)
|
||||
|
||||
def retire(
|
||||
self, secret_id: str, version: int, *, operation_id: str
|
||||
) -> SecretMetadata:
|
||||
secret_id = self._validate_secret_id(secret_id)
|
||||
with self._lock, self._connect() as connection:
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
row = connection.execute(
|
||||
"SELECT status FROM model_secret_versions WHERE secret_id=? AND version=?",
|
||||
(secret_id, version),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE")
|
||||
if str(row["status"]) not in {"revoked", "destroyed"}:
|
||||
connection.execute(
|
||||
"""UPDATE model_secret_versions SET status='retired',
|
||||
retired_at=COALESCE(retired_at, CURRENT_TIMESTAMP),
|
||||
transition_operation_id=? WHERE secret_id=? AND version=?""",
|
||||
(operation_id, secret_id, version),
|
||||
)
|
||||
return self._metadata(secret_id, version)
|
||||
|
||||
def revoke(
|
||||
self,
|
||||
secret_id: str,
|
||||
version: int,
|
||||
*,
|
||||
revoked_by: str,
|
||||
reason: str,
|
||||
operation_id: str,
|
||||
) -> SecretMetadata:
|
||||
secret_id = self._validate_secret_id(secret_id)
|
||||
clean_reason = str(reason or "").strip()
|
||||
if not clean_reason:
|
||||
raise EvoRuntimeError("LLM_SECRET_INVALID")
|
||||
with self._lock, self._connect() as connection:
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
row = connection.execute(
|
||||
"SELECT status FROM model_secret_versions WHERE secret_id=? AND version=?",
|
||||
(secret_id, version),
|
||||
).fetchone()
|
||||
if row is None or str(row["status"]) == "destroyed":
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE")
|
||||
if str(row["status"]) != "revoked":
|
||||
connection.execute(
|
||||
"""UPDATE model_secret_versions SET status='revoked',
|
||||
revoked_at=CURRENT_TIMESTAMP, revoked_by=?, revoke_reason=?,
|
||||
transition_operation_id=? WHERE secret_id=? AND version=?""",
|
||||
(
|
||||
revoked_by[:256],
|
||||
clean_reason[:1024],
|
||||
operation_id,
|
||||
secret_id,
|
||||
version,
|
||||
),
|
||||
)
|
||||
return self._metadata(secret_id, version)
|
||||
|
||||
def _metadata(self, secret_id: str, version: int) -> SecretMetadata:
|
||||
with self._connect() as connection:
|
||||
row = connection.execute(
|
||||
"""SELECT secret_id, version, masked_value, created_by, created_at,
|
||||
status, retired_at, revoked_at, revoked_by, revoke_reason
|
||||
FROM model_secret_versions WHERE secret_id=? AND version=?""",
|
||||
(secret_id, version),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE")
|
||||
return SecretMetadata(
|
||||
secret_id=str(row["secret_id"]),
|
||||
version=int(row["version"]),
|
||||
masked_value=str(row["masked_value"]),
|
||||
created_by=str(row["created_by"]),
|
||||
created_at=str(row["created_at"]),
|
||||
status=str(row["status"]),
|
||||
retired_at=row["retired_at"],
|
||||
revoked_at=row["revoked_at"],
|
||||
revoked_by=row["revoked_by"],
|
||||
revoke_reason=row["revoke_reason"],
|
||||
)
|
||||
@@ -0,0 +1,258 @@
|
||||
"""Canonical, provider-neutral validation for user-adjustable model options."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import math
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from .contracts import EvoRuntimeError
|
||||
from .crypto import canonical_json_v1
|
||||
|
||||
_REASONING_VALUES = frozenset({"off", "on", "low", "medium", "high", "max"})
|
||||
_MAX_OPTION_COUNT = 16
|
||||
_MAX_CANONICAL_BYTES = 4_096
|
||||
|
||||
|
||||
def project_user_options_for_purpose(
|
||||
*,
|
||||
values: Mapping[str, Any],
|
||||
user_options: Mapping[str, Mapping[str, Any]],
|
||||
purpose: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Return values whose user-option contract permits the target purpose.
|
||||
|
||||
Parameters outside ``user_options`` are runtime-owned and remain untouched;
|
||||
downstream adapter validation is still authoritative for those fields.
|
||||
"""
|
||||
|
||||
return {
|
||||
name: value
|
||||
for name, value in values.items()
|
||||
if name not in user_options
|
||||
or purpose
|
||||
in set(user_options[name].get("applies_to") or ("main_agent",))
|
||||
}
|
||||
|
||||
|
||||
def model_options_schema_hash(
|
||||
*,
|
||||
model_profile_id: str,
|
||||
user_options: Mapping[str, Mapping[str, Any]],
|
||||
supports_reasoning: bool,
|
||||
reasoning_mode: str,
|
||||
allowed_reasoning_efforts: Sequence[str],
|
||||
parameter_constraints: Sequence[Mapping[str, Any]],
|
||||
adapter_id: str = "",
|
||||
adapter_revision: str = "",
|
||||
) -> str:
|
||||
"""Hash only the public semantics that determine valid user input."""
|
||||
|
||||
public_options = {
|
||||
str(name): {
|
||||
key: (
|
||||
sorted(str(item) for item in value)
|
||||
if key in {"choices", "applies_to"}
|
||||
else value
|
||||
)
|
||||
for key, value in dict(option).items()
|
||||
if key
|
||||
in {
|
||||
"type",
|
||||
"applies_to",
|
||||
"minimum",
|
||||
"maximum",
|
||||
"minimum_exclusive",
|
||||
"maximum_exclusive",
|
||||
"choices",
|
||||
}
|
||||
}
|
||||
for name, option in sorted(user_options.items())
|
||||
}
|
||||
normalized_constraints = [
|
||||
{
|
||||
key: (
|
||||
sorted(str(item) for item in value)
|
||||
if key == "at_most_one_of"
|
||||
else value
|
||||
)
|
||||
for key, value in dict(constraint).items()
|
||||
}
|
||||
for constraint in parameter_constraints
|
||||
]
|
||||
normalized_constraints.sort(key=canonical_json_v1)
|
||||
payload = {
|
||||
"model_profile_id": str(model_profile_id),
|
||||
"user_options": public_options,
|
||||
"supports_reasoning": bool(supports_reasoning),
|
||||
"reasoning_mode": str(reasoning_mode),
|
||||
"allowed_reasoning_efforts": sorted(
|
||||
str(value) for value in allowed_reasoning_efforts
|
||||
),
|
||||
"parameter_constraints": normalized_constraints,
|
||||
"adapter_id": str(adapter_id),
|
||||
"adapter_revision": str(adapter_revision),
|
||||
}
|
||||
return "sha256:" + hashlib.sha256(canonical_json_v1(payload)).hexdigest()
|
||||
|
||||
|
||||
def validate_parameter_constraints(
|
||||
constraints: Sequence[Mapping[str, Any]], *, allowed_names: set[str]
|
||||
) -> tuple[dict[str, Any], ...]:
|
||||
"""Validate and normalize the supported public constraint grammar."""
|
||||
|
||||
normalized: list[dict[str, Any]] = []
|
||||
for constraint in constraints:
|
||||
values = dict(constraint)
|
||||
if set(values) != {"at_most_one_of"}:
|
||||
raise EvoRuntimeError(
|
||||
"MODEL_PARAMETER_INVALID", "unsupported user option constraint"
|
||||
)
|
||||
names = values["at_most_one_of"]
|
||||
if not isinstance(names, list | tuple) or len(names) < 2:
|
||||
raise EvoRuntimeError(
|
||||
"MODEL_PARAMETER_INVALID", "at_most_one_of must contain two names"
|
||||
)
|
||||
clean = tuple(str(name) for name in names)
|
||||
if len(set(clean)) != len(clean) or not set(clean) <= allowed_names:
|
||||
raise EvoRuntimeError(
|
||||
"MODEL_PARAMETER_INVALID", "constraint references invalid options"
|
||||
)
|
||||
normalized.append({"at_most_one_of": clean})
|
||||
return tuple(normalized)
|
||||
|
||||
|
||||
def validate_user_model_options(
|
||||
*,
|
||||
supplied: Mapping[str, Any] | None,
|
||||
user_options: Mapping[str, Mapping[str, Any]],
|
||||
supports_reasoning: bool,
|
||||
reasoning_mode: str,
|
||||
allowed_reasoning_efforts: Sequence[str],
|
||||
default_reasoning_effort: str | None,
|
||||
parameter_constraints: Sequence[Mapping[str, Any]],
|
||||
purpose: str = "main_agent",
|
||||
allow_reasoning: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
"""Return canonical user options or raise a stable contract error."""
|
||||
|
||||
values = dict(supplied or {})
|
||||
if len(values) > _MAX_OPTION_COUNT:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "too many model options")
|
||||
try:
|
||||
encoded = canonical_json_v1(values)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise EvoRuntimeError(
|
||||
"MODEL_PARAMETER_INVALID", "model options must be canonical JSON"
|
||||
) from exc
|
||||
if len(encoded) > _MAX_CANONICAL_BYTES:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "model options are too large")
|
||||
|
||||
allowed = set(user_options)
|
||||
if allow_reasoning and supports_reasoning:
|
||||
allowed.add("reasoning")
|
||||
unknown = set(values) - allowed
|
||||
if unknown:
|
||||
raise EvoRuntimeError(
|
||||
"MODEL_PARAMETER_INVALID",
|
||||
f"unsupported model options: {', '.join(sorted(unknown))}",
|
||||
)
|
||||
|
||||
result: dict[str, Any] = {}
|
||||
for name, value in values.items():
|
||||
if name == "reasoning":
|
||||
result[name] = _normalize_reasoning(
|
||||
value,
|
||||
supports_reasoning=supports_reasoning,
|
||||
reasoning_mode=reasoning_mode,
|
||||
allowed_reasoning_efforts=allowed_reasoning_efforts,
|
||||
default_reasoning_effort=default_reasoning_effort,
|
||||
)
|
||||
continue
|
||||
option = dict(user_options[name])
|
||||
applies_to = set(option.get("applies_to") or ("main_agent",))
|
||||
if purpose not in applies_to:
|
||||
raise EvoRuntimeError(
|
||||
"MODEL_PARAMETER_INVALID",
|
||||
f"model option {name} cannot apply to {purpose}",
|
||||
)
|
||||
_validate_option_value(name, value, option)
|
||||
result[name] = value
|
||||
|
||||
constraints = validate_parameter_constraints(
|
||||
parameter_constraints,
|
||||
allowed_names=set(user_options)
|
||||
| ({"reasoning"} if supports_reasoning else set()),
|
||||
)
|
||||
for constraint in constraints:
|
||||
names = tuple(constraint["at_most_one_of"])
|
||||
if sum(name in result for name in names) > 1:
|
||||
raise EvoRuntimeError(
|
||||
"MODEL_PARAMETER_CONFLICT",
|
||||
f"at most one of {', '.join(names)} may be supplied",
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def _normalize_reasoning(
|
||||
value: Any,
|
||||
*,
|
||||
supports_reasoning: bool,
|
||||
reasoning_mode: str,
|
||||
allowed_reasoning_efforts: Sequence[str],
|
||||
default_reasoning_effort: str | None,
|
||||
) -> str:
|
||||
if not supports_reasoning:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "reasoning is unsupported")
|
||||
clean = str(value or "").strip().lower()
|
||||
if clean == "disabled":
|
||||
clean = "off"
|
||||
if clean not in _REASONING_VALUES:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "reasoning value is invalid")
|
||||
if clean == "off":
|
||||
return clean
|
||||
if reasoning_mode == "boolean":
|
||||
return "on"
|
||||
if reasoning_mode != "effort":
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "reasoning is unsupported")
|
||||
if clean == "on":
|
||||
clean = str(default_reasoning_effort or "")
|
||||
allowed = {str(item) for item in allowed_reasoning_efforts}
|
||||
if clean == "medium" and clean not in allowed and "high" in allowed:
|
||||
clean = "high"
|
||||
if clean not in allowed:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "reasoning effort is invalid")
|
||||
return clean
|
||||
|
||||
|
||||
def _validate_option_value(name: str, value: Any, option: Mapping[str, Any]) -> None:
|
||||
kind = str(option.get("type") or "")
|
||||
valid_type = (
|
||||
(kind == "boolean" and isinstance(value, bool))
|
||||
or (
|
||||
kind == "integer" and isinstance(value, int) and not isinstance(value, bool)
|
||||
)
|
||||
or (
|
||||
kind == "number"
|
||||
and isinstance(value, int | float)
|
||||
and not isinstance(value, bool)
|
||||
and math.isfinite(float(value))
|
||||
)
|
||||
or (kind in {"string", "enum"} and isinstance(value, str))
|
||||
or (kind == "object" and isinstance(value, Mapping))
|
||||
)
|
||||
if not valid_type:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} has invalid type")
|
||||
if "minimum" in option and value < option["minimum"]:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} is below minimum")
|
||||
if "minimum_exclusive" in option and value <= option["minimum_exclusive"]:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} is below minimum")
|
||||
if "maximum" in option and value > option["maximum"]:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} exceeds maximum")
|
||||
if "maximum_exclusive" in option and value >= option["maximum_exclusive"]:
|
||||
raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} exceeds maximum")
|
||||
if option.get("choices") and value not in option["choices"]:
|
||||
raise EvoRuntimeError(
|
||||
"MODEL_PARAMETER_INVALID", f"{name} is not an allowed value"
|
||||
)
|
||||
@@ -98,6 +98,11 @@ def build_memory_agent_graph(
|
||||
workspace_dir=workspace_dir,
|
||||
memory_dir=memory_dir,
|
||||
)
|
||||
middleware = list(middleware)
|
||||
if skills:
|
||||
from ...middleware import BudgetedSkillsMiddleware
|
||||
|
||||
middleware.append(BudgetedSkillsMiddleware(backend=backend, sources=skills))
|
||||
|
||||
agent = create_deep_agent(
|
||||
name=name,
|
||||
@@ -105,9 +110,9 @@ def build_memory_agent_graph(
|
||||
system_prompt=system_prompt,
|
||||
tools=list(tools),
|
||||
backend=backend,
|
||||
middleware=list(middleware),
|
||||
middleware=middleware,
|
||||
subagents=[],
|
||||
skills=skills,
|
||||
skills=None,
|
||||
**kwargs,
|
||||
)
|
||||
return agent.with_config({"recursion_limit": recursion_limit})
|
||||
|
||||
@@ -428,8 +428,10 @@ def _memory_worker_middleware(
|
||||
enable_observation_memory: bool = True,
|
||||
):
|
||||
"""Build middleware for memory workers, excluding task execution tools."""
|
||||
from ...middleware.configurable_model import ConfigurableModelMiddleware
|
||||
from ...middleware.error_normalization import ErrorNormalizationMiddleware
|
||||
from ...middleware.memory import create_memory_middleware
|
||||
from ...middleware.recoverable_metering import RecoverableMeteringMiddleware
|
||||
|
||||
memory_controls = MemoryControls(
|
||||
profile_enabled=enable_profile_memory,
|
||||
@@ -444,6 +446,8 @@ def _memory_worker_middleware(
|
||||
# Outermost — normalize provider-SDK exceptions from the
|
||||
# auxiliary model call before any inner middleware sees them.
|
||||
ErrorNormalizationMiddleware(),
|
||||
ConfigurableModelMiddleware(),
|
||||
RecoverableMeteringMiddleware(),
|
||||
*memory_agent_middleware(
|
||||
create_memory_middleware(
|
||||
str(memory_dir),
|
||||
|
||||
@@ -71,7 +71,9 @@ def build_observation_linker_graph(
|
||||
workspace_dir: str | Path | None = None,
|
||||
) -> CompiledStateGraph:
|
||||
"""Build the registered LangGraph observation linker."""
|
||||
from ...middleware.configurable_model import ConfigurableModelMiddleware
|
||||
from ...middleware.error_normalization import ErrorNormalizationMiddleware
|
||||
from ...middleware.recoverable_metering import RecoverableMeteringMiddleware
|
||||
|
||||
agent_paths = resolve_memory_agent_paths(
|
||||
memory_dir=memory_dir,
|
||||
@@ -89,5 +91,10 @@ def build_observation_linker_graph(
|
||||
workspace_dir=agent_paths.workspace_dir,
|
||||
# Outermost — normalize provider-SDK exceptions from the
|
||||
# auxiliary model call before any inner middleware sees them.
|
||||
middleware=[ErrorNormalizationMiddleware(), *memory_agent_middleware()],
|
||||
middleware=[
|
||||
ErrorNormalizationMiddleware(),
|
||||
ConfigurableModelMiddleware(),
|
||||
RecoverableMeteringMiddleware(),
|
||||
*memory_agent_middleware(),
|
||||
],
|
||||
)
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
|
||||
@@ -42,6 +42,59 @@ OBSERVATION_LINKER_GRAPH_ID = "evomemory-observation-linker"
|
||||
MemoryWorkerFinishedHook = Callable[[BackgroundRun, MemoryOutputDelta | None], None]
|
||||
MemoryWorkerAbortedHook = Callable[[BackgroundRun, MemoryOutputDelta | None], None]
|
||||
|
||||
_INHERITED_RUNTIME_KEYS = (
|
||||
"model",
|
||||
"model_provider",
|
||||
"ai4sci_metering",
|
||||
"ai4sci_model_proxy",
|
||||
)
|
||||
|
||||
|
||||
def _current_runtime_context() -> tuple[str | None, dict[str, object]]:
|
||||
"""Capture the signed parent Run context before launching a child graph."""
|
||||
try:
|
||||
from langgraph.config import get_config
|
||||
|
||||
config = get_config()
|
||||
except Exception:
|
||||
return None, {}
|
||||
if not isinstance(config, Mapping):
|
||||
return None, {}
|
||||
configurable = config.get("configurable")
|
||||
inherited = {
|
||||
key: configurable[key]
|
||||
for key in _INHERITED_RUNTIME_KEYS
|
||||
if isinstance(configurable, Mapping) and key in configurable
|
||||
}
|
||||
metadata = config.get("metadata")
|
||||
runtime_url = (
|
||||
str(metadata.get("langgraph_api_url") or "")
|
||||
if isinstance(metadata, Mapping)
|
||||
else ""
|
||||
)
|
||||
return runtime_url or None, inherited
|
||||
|
||||
|
||||
def _with_runtime_context(
|
||||
payload: BackgroundRunPayload,
|
||||
*,
|
||||
inherited: Mapping[str, object] | None,
|
||||
source_type: str,
|
||||
) -> BackgroundRunPayload:
|
||||
"""Attach a child billing scope without changing the signed envelope."""
|
||||
normalized = cast("BackgroundRunPayload", dict(payload))
|
||||
config = dict(normalized.get("config") or {})
|
||||
configurable = dict(config.get("configurable") or {})
|
||||
configurable.update(dict(inherited or {}))
|
||||
metering = configurable.get("ai4sci_metering")
|
||||
if isinstance(metering, Mapping):
|
||||
configurable["ai4sci_metering"] = {
|
||||
**dict(metering),
|
||||
"source_type": source_type,
|
||||
}
|
||||
config["configurable"] = configurable
|
||||
return cast("BackgroundRunPayload", {**normalized, "config": config})
|
||||
|
||||
|
||||
def _observation_linking_enabled() -> bool:
|
||||
return MemoryControls.from_config(get_effective_config()).observations_enabled
|
||||
@@ -104,6 +157,7 @@ def _memory_worker_run_payload(
|
||||
*,
|
||||
context: MemorySourceContext,
|
||||
thread_id: str,
|
||||
inherited: Mapping[str, object] | None = None,
|
||||
) -> BackgroundRunPayload:
|
||||
"""Build the LangGraph SDK run payload for a memory worker."""
|
||||
metadata = _memory_worker_metadata(context)
|
||||
@@ -121,7 +175,17 @@ def _memory_worker_run_payload(
|
||||
}
|
||||
},
|
||||
}
|
||||
return _runs_create_kwargs(payload)
|
||||
payload = _runs_create_kwargs(payload)
|
||||
source_type = (
|
||||
"evomemory_turn_worker"
|
||||
if context.source_type == MemorySourceType.TURN
|
||||
else "evomemory_subagent_worker"
|
||||
)
|
||||
return _with_runtime_context(
|
||||
payload,
|
||||
inherited=inherited,
|
||||
source_type=source_type,
|
||||
)
|
||||
|
||||
|
||||
def memory_worker_launch_request(
|
||||
@@ -129,14 +193,20 @@ def memory_worker_launch_request(
|
||||
) -> BackgroundRunRequest:
|
||||
"""Build the background run request for a memory worker."""
|
||||
metadata = _memory_worker_metadata(context)
|
||||
runtime_url, inherited = _current_runtime_context()
|
||||
|
||||
def run_payload(thread_id: str) -> BackgroundRunPayload:
|
||||
return _memory_worker_run_payload(context=context, thread_id=thread_id)
|
||||
return _memory_worker_run_payload(
|
||||
context=context,
|
||||
thread_id=thread_id,
|
||||
inherited=inherited,
|
||||
)
|
||||
|
||||
return BackgroundRunRequest(
|
||||
graph_id=_memory_worker_graph_id(context.source_type),
|
||||
run_payload=run_payload,
|
||||
thread_metadata=metadata,
|
||||
url=runtime_url,
|
||||
name="EvoMemory worker",
|
||||
)
|
||||
|
||||
@@ -192,7 +262,12 @@ def _observation_linker_run_payload(
|
||||
}
|
||||
},
|
||||
}
|
||||
return _runs_create_kwargs(payload)
|
||||
payload = _runs_create_kwargs(payload)
|
||||
return _with_runtime_context(
|
||||
payload,
|
||||
inherited=context.runtime_configurable,
|
||||
source_type="evomemory_linker",
|
||||
)
|
||||
|
||||
|
||||
def observation_linker_launch_request(
|
||||
@@ -210,6 +285,7 @@ def observation_linker_launch_request(
|
||||
graph_id=OBSERVATION_LINKER_GRAPH_ID,
|
||||
run_payload=run_payload,
|
||||
thread_metadata=_observation_linker_metadata(context),
|
||||
url=context.runtime_url,
|
||||
name="EvoMemory observation linker",
|
||||
)
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ from .index import (
|
||||
build_observation_index_context,
|
||||
build_observation_linker_index_context,
|
||||
)
|
||||
from .relations import link_observation_files
|
||||
from .relations import archive_observation_file, link_observation_files
|
||||
from .store import (
|
||||
OBSERVATION_DIR,
|
||||
ObservationFrontmatter,
|
||||
@@ -51,6 +51,7 @@ __all__ = [
|
||||
"RecordObservationArgs",
|
||||
"RelatedObservationEntry",
|
||||
"SearchObservationsArgs",
|
||||
"archive_observation_file",
|
||||
"build_observation_index_context",
|
||||
"build_observation_linker_index_context",
|
||||
"create_link_observations_tool",
|
||||
|
||||
@@ -2,21 +2,23 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import os
|
||||
import posixpath
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
|
||||
from filelock import FileLock
|
||||
|
||||
from ..types import ObservationRelation
|
||||
from .store import (
|
||||
ObservationFrontmatter,
|
||||
RelatedObservationEntry,
|
||||
observation_document_by_id,
|
||||
read_observation_document,
|
||||
related_observation_entries,
|
||||
write_observation_document,
|
||||
)
|
||||
|
||||
_link_write_lock = threading.Lock()
|
||||
|
||||
|
||||
def _relation_value(value: ObservationRelation | str) -> str:
|
||||
try:
|
||||
@@ -87,7 +89,8 @@ def link_observation_files(
|
||||
if not reason_text:
|
||||
raise ValueError("reason must not be empty")
|
||||
|
||||
with _link_write_lock:
|
||||
relation_lock = Path(memory_dir).expanduser() / ".relation-write.lock"
|
||||
with FileLock(str(relation_lock), timeout=30):
|
||||
source_document = observation_document_by_id(
|
||||
memory_dir=memory_dir,
|
||||
project_id=project_id,
|
||||
@@ -154,3 +157,84 @@ def link_observation_files(
|
||||
],
|
||||
"missing_observation_ids": [],
|
||||
}
|
||||
|
||||
|
||||
def archive_observation_file(
|
||||
*,
|
||||
memory_dir: str | Path,
|
||||
observation_id: str,
|
||||
observation_path: str,
|
||||
) -> dict[str, object]:
|
||||
"""Archive one observation and remove references to it from live memory."""
|
||||
|
||||
requested_id = observation_id.strip()
|
||||
raw_path = observation_path.strip().replace("\\", "/").lstrip("/")
|
||||
if ".." in raw_path.split("/"):
|
||||
raise ValueError("invalid observation archive target")
|
||||
normalized = posixpath.normpath("/" + raw_path)
|
||||
parts = Path(normalized.lstrip("/")).parts
|
||||
if (
|
||||
not requested_id.startswith("O-")
|
||||
or len(parts) < 3
|
||||
or parts[0] != "observations"
|
||||
or parts[-1] != f"{requested_id}.md"
|
||||
or ".." in parts
|
||||
):
|
||||
raise ValueError("invalid observation archive target")
|
||||
|
||||
root = Path(memory_dir).expanduser().resolve()
|
||||
source = root.joinpath(*parts)
|
||||
if source.is_symlink() or not source.is_file():
|
||||
raise FileNotFoundError(observation_path)
|
||||
|
||||
observation_lock = FileLock(str(root / ".observation-write.lock"), timeout=30)
|
||||
relation_lock = FileLock(str(root / ".relation-write.lock"), timeout=30)
|
||||
with observation_lock, relation_lock:
|
||||
document = read_observation_document(source)
|
||||
if document is None or document[0].id != requested_id:
|
||||
raise FileNotFoundError(observation_path)
|
||||
target_metadata, _target_body = document
|
||||
|
||||
archive_id = datetime.now(UTC).strftime("%Y%m%dT%H%M%S%fZ")
|
||||
archived = root / "trash" / "observations" / archive_id / Path(*parts[1:])
|
||||
archived.parent.mkdir(parents=True, exist_ok=True)
|
||||
os.replace(source, archived)
|
||||
|
||||
observations_root = root / "observations"
|
||||
if target_metadata.scope.value == "global":
|
||||
candidates = observations_root.rglob("*.md")
|
||||
else:
|
||||
project_id = str(target_metadata.project_id or "")
|
||||
candidates = iter(
|
||||
[
|
||||
*(observations_root / "global").glob("*.md"),
|
||||
*(observations_root / "projects" / project_id).glob("*.md"),
|
||||
]
|
||||
)
|
||||
|
||||
updated_observation_ids: list[str] = []
|
||||
for candidate in sorted(candidates):
|
||||
if candidate.is_symlink() or not candidate.is_file():
|
||||
continue
|
||||
related_document = read_observation_document(candidate)
|
||||
if related_document is None:
|
||||
continue
|
||||
metadata, body = related_document
|
||||
retained = [
|
||||
entry
|
||||
for entry in metadata.related_observations
|
||||
if entry.id != requested_id
|
||||
]
|
||||
if len(retained) == len(metadata.related_observations):
|
||||
continue
|
||||
metadata.related_observations = retained
|
||||
write_observation_document(candidate, metadata=metadata, body=body)
|
||||
updated_observation_ids.append(metadata.id)
|
||||
|
||||
return {
|
||||
"removed": True,
|
||||
"observation_id": requested_id,
|
||||
"archive_id": archive_id,
|
||||
"archive_path": archived.relative_to(root).as_posix(),
|
||||
"updated_observation_ids": updated_observation_ids,
|
||||
}
|
||||
|
||||
@@ -9,11 +9,14 @@ from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import replace
|
||||
from datetime import UTC, date, datetime
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
from filelock import FileLock
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator
|
||||
|
||||
from ..search import (
|
||||
@@ -227,7 +230,25 @@ def write_observation_document(
|
||||
allow_unicode=True,
|
||||
sort_keys=False,
|
||||
)
|
||||
Path(path).write_text(f"---\n{frontmatter}---\n{body}", encoding="utf-8")
|
||||
_atomic_write_text(Path(path), f"---\n{frontmatter}---\n{body}")
|
||||
|
||||
|
||||
def _atomic_write_text(path: Path, content: str) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
fd, temporary = tempfile.mkstemp(
|
||||
prefix=f".{path.name}.", suffix=".tmp", dir=path.parent
|
||||
)
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as handle:
|
||||
handle.write(content)
|
||||
handle.flush()
|
||||
os.fsync(handle.fileno())
|
||||
os.replace(temporary, path)
|
||||
finally:
|
||||
try:
|
||||
os.unlink(temporary)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
|
||||
def read_observation_id_from_path(path: str | Path) -> str | None:
|
||||
@@ -605,25 +626,26 @@ def record_observation_file(
|
||||
)
|
||||
path = Path(memory_dir).expanduser() / memory_path.lstrip("/")
|
||||
created = False
|
||||
if not path.exists():
|
||||
created_at = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ")
|
||||
content = _format_observation_markdown(
|
||||
observation_id=observation_id,
|
||||
created_at=created_at,
|
||||
memory_type=memory_type,
|
||||
summary=summary_text,
|
||||
observation=observation_text,
|
||||
why_it_matters=why_text,
|
||||
evidence=evidence.strip() if evidence else None,
|
||||
scope=scope,
|
||||
source_type=source_type,
|
||||
source_agent=source_agent,
|
||||
source_session_id=source_session_id,
|
||||
project_id=project_id,
|
||||
)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(content, encoding="utf-8")
|
||||
created = True
|
||||
memory_root = Path(memory_dir).expanduser()
|
||||
with FileLock(str(memory_root / ".observation-write.lock"), timeout=30):
|
||||
if not path.exists():
|
||||
created_at = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ")
|
||||
content = _format_observation_markdown(
|
||||
observation_id=observation_id,
|
||||
created_at=created_at,
|
||||
memory_type=memory_type,
|
||||
summary=summary_text,
|
||||
observation=observation_text,
|
||||
why_it_matters=why_text,
|
||||
evidence=evidence.strip() if evidence else None,
|
||||
scope=scope,
|
||||
source_type=source_type,
|
||||
source_agent=source_agent,
|
||||
source_session_id=source_session_id,
|
||||
project_id=project_id,
|
||||
)
|
||||
_atomic_write_text(path, content)
|
||||
created = True
|
||||
|
||||
result: ObservationRecordResult = {
|
||||
"observation_id": observation_id,
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import NamedTuple
|
||||
@@ -29,6 +29,8 @@ class ObservationLinkerContext:
|
||||
workspace_dir: Path
|
||||
project_id: str
|
||||
observation_ids: tuple[str, ...]
|
||||
runtime_url: str | None = None
|
||||
runtime_configurable: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
ObservationLinkerLauncher = Callable[[ObservationLinkerContext], BackgroundRun | None]
|
||||
@@ -77,6 +79,9 @@ class MemoryScheduler:
|
||||
self._launch_linker = launch_linker
|
||||
self._has_active_workers = has_active_workers
|
||||
self._pending: dict[_BatchKey, set[str]] = {}
|
||||
self._runtime_context: dict[
|
||||
_BatchKey, tuple[str | None, Mapping[str, object] | None]
|
||||
] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def _launch_ready(
|
||||
@@ -121,6 +126,11 @@ class MemoryScheduler:
|
||||
)
|
||||
with self._lock:
|
||||
self._pending.setdefault(key, set()).update(context.observation_ids)
|
||||
if context.runtime_url or context.runtime_configurable:
|
||||
self._runtime_context[key] = (
|
||||
context.runtime_url,
|
||||
context.runtime_configurable,
|
||||
)
|
||||
|
||||
def flush_ready(self) -> None:
|
||||
"""Launch any pending linker batches that are no longer blocked."""
|
||||
@@ -144,19 +154,39 @@ class MemoryScheduler:
|
||||
contexts = self._record_finished_and_drain_ready(run=run, delta=delta)
|
||||
self._launch_ready(contexts)
|
||||
|
||||
def _ready_batches_locked(self) -> list[tuple[_BatchKey, set[str]]]:
|
||||
def _ready_batches_locked(
|
||||
self,
|
||||
) -> list[
|
||||
tuple[
|
||||
_BatchKey,
|
||||
set[str],
|
||||
tuple[str | None, Mapping[str, object] | None],
|
||||
]
|
||||
]:
|
||||
ready_batches = []
|
||||
for key in list(self._pending):
|
||||
if not self._has_active_workers(key.memory_dir):
|
||||
ready_batches.append((key, self._pending.pop(key)))
|
||||
ready_batches.append(
|
||||
(
|
||||
key,
|
||||
self._pending.pop(key),
|
||||
self._runtime_context.pop(key, (None, None)),
|
||||
)
|
||||
)
|
||||
return ready_batches
|
||||
|
||||
def _contexts_for_batches(
|
||||
self,
|
||||
ready_batches: list[tuple[_BatchKey, set[str]]],
|
||||
ready_batches: list[
|
||||
tuple[
|
||||
_BatchKey,
|
||||
set[str],
|
||||
tuple[str | None, Mapping[str, object] | None],
|
||||
]
|
||||
],
|
||||
) -> tuple[ObservationLinkerContext, ...]:
|
||||
ready_contexts = []
|
||||
for key, observation_ids in ready_batches:
|
||||
for key, observation_ids, runtime_context in ready_batches:
|
||||
if not observation_ids:
|
||||
continue
|
||||
ready_contexts.append(
|
||||
@@ -165,6 +195,8 @@ class MemoryScheduler:
|
||||
workspace_dir=Path(key.workspace_dir),
|
||||
project_id=key.project_id,
|
||||
observation_ids=tuple(sorted(observation_ids)),
|
||||
runtime_url=runtime_context[0],
|
||||
runtime_configurable=runtime_context[1],
|
||||
)
|
||||
)
|
||||
|
||||
@@ -194,6 +226,10 @@ class MemoryScheduler:
|
||||
with self._lock:
|
||||
if key is not None and observation_ids:
|
||||
self._pending.setdefault(key, set()).update(observation_ids)
|
||||
if isinstance(run.configurable, Mapping) and isinstance(
|
||||
run.configurable.get("ai4sci_metering"), Mapping
|
||||
):
|
||||
self._runtime_context[key] = (run.url, run.configurable)
|
||||
|
||||
ready_batches = self._ready_batches_locked()
|
||||
|
||||
|
||||
@@ -23,6 +23,7 @@ DEFAULT_MATCH_LINES = 3
|
||||
DEFAULT_MATCH_CHARS = 240
|
||||
|
||||
_TOKEN_RE = re.compile(r"[a-z0-9_]+")
|
||||
_CJK_RUN_RE = re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff]+")
|
||||
|
||||
|
||||
def _compile_query_pattern(query: str) -> re.Pattern[str]:
|
||||
@@ -35,11 +36,20 @@ def _compile_query_pattern(query: str) -> re.Pattern[str]:
|
||||
|
||||
def _tokens(text: str) -> list[str]:
|
||||
"""Return simple lowercase search tokens."""
|
||||
return [
|
||||
normalized = text.casefold()
|
||||
tokens = [
|
||||
token
|
||||
for token in _TOKEN_RE.findall(text.casefold())
|
||||
for token in _TOKEN_RE.findall(normalized)
|
||||
if len(token) >= MIN_TOKEN_CHARS
|
||||
]
|
||||
for run in _CJK_RUN_RE.findall(normalized):
|
||||
tokens.append(run)
|
||||
for size in (2, 3):
|
||||
tokens.extend(
|
||||
run[index : index + size]
|
||||
for index in range(max(0, len(run) - size + 1))
|
||||
)
|
||||
return tokens
|
||||
|
||||
|
||||
def _document_tokens(document: ObservationSearchDocument) -> set[str]:
|
||||
|
||||
@@ -18,6 +18,7 @@ from .context_editing import (
|
||||
create_context_editing_middleware,
|
||||
)
|
||||
from .context_overflow import ContextOverflowMapperMiddleware
|
||||
from .disable_subagent import DisableSubagentToolMiddleware
|
||||
from .error_normalization import ErrorNormalizationMiddleware
|
||||
from .memory import (
|
||||
EvoMemoryMiddleware,
|
||||
@@ -29,6 +30,9 @@ from .memory_lifecycle import (
|
||||
default_memory_scheduler,
|
||||
)
|
||||
from .model_fallback import ModelFallbackMiddleware, load_fallback_chain
|
||||
from .provider_context import ProviderContextMediaMiddleware
|
||||
from .recoverable_metering import RecoverableMeteringMiddleware
|
||||
from .recoverable_tools import RecoverableToolEffectMiddleware
|
||||
from .repetitive_tool_guard import (
|
||||
DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS,
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
@@ -40,6 +44,12 @@ from .scheduler import (
|
||||
SchedulerMiddleware,
|
||||
create_scheduler_middleware,
|
||||
)
|
||||
from .skill_context import (
|
||||
DEFAULT_MAX_DESCRIPTION_BYTES,
|
||||
DEFAULT_MAX_SKILLS,
|
||||
DEFAULT_MAX_SKILLS_BYTES,
|
||||
BudgetedSkillsMiddleware,
|
||||
)
|
||||
from .tool_error_handler import ToolErrorHandlerMiddleware
|
||||
from .tool_protocol_guard import ToolProtocolGuardMiddleware
|
||||
from .tool_selector import create_tool_selector_middleware
|
||||
@@ -47,18 +57,26 @@ from .utils import disable_thinking
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS",
|
||||
"DEFAULT_MAX_DESCRIPTION_BYTES",
|
||||
"DEFAULT_MAX_SKILLS",
|
||||
"DEFAULT_MAX_SKILLS_BYTES",
|
||||
"DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD",
|
||||
"AskUserMiddleware",
|
||||
"AskUserRequest",
|
||||
"AskUserWidgetResult",
|
||||
"BudgetedSkillsMiddleware",
|
||||
"Choice",
|
||||
"ConfigurableModelMiddleware",
|
||||
"ContextOverflowMapperMiddleware",
|
||||
"DisableSubagentToolMiddleware",
|
||||
"ErrorNormalizationMiddleware",
|
||||
"EvoMemoryLifecycleMiddleware",
|
||||
"EvoMemoryMiddleware",
|
||||
"ModelFallbackMiddleware",
|
||||
"ProviderContextMediaMiddleware",
|
||||
"Question",
|
||||
"RecoverableMeteringMiddleware",
|
||||
"RecoverableToolEffectMiddleware",
|
||||
"RepetitiveToolCallGuardMiddleware",
|
||||
"RuntimeContextMiddleware",
|
||||
"SchedulerMiddleware",
|
||||
|
||||
@@ -74,6 +74,28 @@ def _read_model_override() -> tuple[str | None, str | None]:
|
||||
)
|
||||
|
||||
|
||||
def _read_proxy_override() -> tuple[dict[str, Any] | None, str, str]:
|
||||
try:
|
||||
from langgraph.config import get_config
|
||||
|
||||
cfg = get_config()
|
||||
except Exception:
|
||||
return None, "", ""
|
||||
configurable = cfg.get("configurable") if isinstance(cfg, dict) else None
|
||||
if not isinstance(configurable, dict):
|
||||
return None, "", ""
|
||||
proxy = configurable.get("ai4sci_model_proxy")
|
||||
metering = configurable.get("ai4sci_metering")
|
||||
if not isinstance(proxy, dict):
|
||||
return None, "", ""
|
||||
metering = metering if isinstance(metering, dict) else {}
|
||||
return (
|
||||
dict(proxy),
|
||||
str(metering.get("provider_id") or ""),
|
||||
str(metering.get("model_id") or ""),
|
||||
)
|
||||
|
||||
|
||||
class ConfigurableModelMiddleware(AgentMiddleware):
|
||||
"""Re-resolve the chat model from RunnableConfig.configurable on every call.
|
||||
|
||||
@@ -150,6 +172,17 @@ class ConfigurableModelMiddleware(AgentMiddleware):
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
proxy, provider_id, model_id = _read_proxy_override()
|
||||
if proxy is not None:
|
||||
from ..llm.gateway_proxy import proxy_from_config
|
||||
|
||||
return handler(
|
||||
request.override(
|
||||
model=proxy_from_config(
|
||||
proxy, provider_id=provider_id, model_id=model_id
|
||||
)
|
||||
)
|
||||
)
|
||||
model_name, provider = _read_model_override()
|
||||
if model_name is None:
|
||||
return handler(request)
|
||||
@@ -172,6 +205,17 @@ class ConfigurableModelMiddleware(AgentMiddleware):
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
proxy, provider_id, model_id = _read_proxy_override()
|
||||
if proxy is not None:
|
||||
from ..llm.gateway_proxy import proxy_from_config
|
||||
|
||||
return await handler(
|
||||
request.override(
|
||||
model=proxy_from_config(
|
||||
proxy, provider_id=provider_id, model_id=model_id
|
||||
)
|
||||
)
|
||||
)
|
||||
model_name, provider = _read_model_override()
|
||||
if model_name is None:
|
||||
return await handler(request)
|
||||
|
||||
@@ -67,6 +67,8 @@ class ContextOverflowMapperMiddleware(AgentMiddleware):
|
||||
|
||||
It triggers when there's an error 400 raised and one of specified patterns exists in the error message.
|
||||
"""
|
||||
if getattr(exc, "code", None) == "MODEL_CONTEXT_WINDOW_EXCEEDED":
|
||||
return True
|
||||
err_msg = str(exc).lower()
|
||||
|
||||
patterns = [
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
"""Web profile guard that makes DeepAgents subagent delegation unavailable."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
ModelRequest,
|
||||
ModelResponse,
|
||||
ToolCallRequest,
|
||||
)
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langgraph.types import Command
|
||||
|
||||
|
||||
def _tool_name(tool: Any) -> str:
|
||||
if isinstance(tool, dict):
|
||||
return str(tool.get("name") or tool.get("function", {}).get("name") or "")
|
||||
return str(getattr(tool, "name", "") or "")
|
||||
|
||||
|
||||
class DisableSubagentToolMiddleware(AgentMiddleware):
|
||||
"""Hide and reject the built-in ``task`` tool for the Web profile.
|
||||
|
||||
DeepAgents adds a general-purpose task tool even when an empty subagent
|
||||
list is supplied. Filtering the model request alone is therefore not a
|
||||
sufficient safety boundary; the tool-call guard protects restored or
|
||||
malformed checkpoints as well.
|
||||
"""
|
||||
|
||||
name = "disable_subagent_tool"
|
||||
|
||||
def wrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
tools = [tool for tool in request.tools if _tool_name(tool) != "task"]
|
||||
return handler(request.override(tools=tools))
|
||||
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
tools = [tool for tool in request.tools if _tool_name(tool) != "task"]
|
||||
return await handler(request.override(tools=tools))
|
||||
|
||||
def wrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]],
|
||||
) -> ToolMessage | Command[Any]:
|
||||
if str(request.tool_call.get("name") or "") == "task":
|
||||
return self._rejected(request)
|
||||
return handler(request)
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]],
|
||||
) -> ToolMessage | Command[Any]:
|
||||
if str(request.tool_call.get("name") or "") == "task":
|
||||
return self._rejected(request)
|
||||
return await handler(request)
|
||||
|
||||
@staticmethod
|
||||
def _rejected(request: ToolCallRequest) -> ToolMessage:
|
||||
return ToolMessage(
|
||||
content="SUBAGENTS_DISABLED",
|
||||
tool_call_id=str(request.tool_call.get("id") or "subagents_disabled"),
|
||||
name="task",
|
||||
status="error",
|
||||
)
|
||||
|
||||
@@ -152,7 +152,6 @@ def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError
|
||||
_extract_provider_code,
|
||||
_extract_status_code,
|
||||
_provider_from_model,
|
||||
_redact_api_keys,
|
||||
)
|
||||
|
||||
# Already normalized (e.g. by ModelFallbackMiddleware wrapping against
|
||||
@@ -191,11 +190,24 @@ def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError
|
||||
else None
|
||||
)
|
||||
|
||||
status_code = _extract_status_code(exc)
|
||||
safe_message = {
|
||||
400: "Provider rejected the request.",
|
||||
401: "Provider authentication failed.",
|
||||
403: "Provider authorization failed.",
|
||||
404: "Provider model or endpoint was not found.",
|
||||
408: "Provider request timed out.",
|
||||
429: "Provider rate limit was exceeded.",
|
||||
500: "Provider request failed.",
|
||||
502: "Provider gateway failed.",
|
||||
503: "Provider is temporarily unavailable.",
|
||||
504: "Provider gateway timed out.",
|
||||
}.get(status_code, "Provider request failed.")
|
||||
return ProviderStreamError(
|
||||
provider=provider,
|
||||
class_qualname=class_qualname,
|
||||
message=_redact_api_keys(str(exc)),
|
||||
status_code=_extract_status_code(exc),
|
||||
message=safe_message,
|
||||
status_code=status_code,
|
||||
code=_extract_provider_code(exc),
|
||||
err_type=_extract_error_type(exc),
|
||||
request_id=request_id,
|
||||
|
||||
@@ -0,0 +1,455 @@
|
||||
"""Evo-owned fallback middleware for frozen Web model-route snapshots."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
ModelRequest,
|
||||
ModelResponse,
|
||||
)
|
||||
from langchain_core.messages import AIMessage, SystemMessage, ToolMessage
|
||||
|
||||
from ..llm.contracts import EvoRuntimeError
|
||||
from ..llm.errors import ModelProviderResponseError, ModelToolProtocolError
|
||||
from ..llm.invocation import (
|
||||
assistant_message_has_output,
|
||||
project_provider_messages,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_TOOL_RESULT_CONTEXT_LIMIT = 12_000
|
||||
_TOOL_PROTOCOL_METADATA = frozenset(
|
||||
{"tool_calls", "tool_call_chunks", "function_call", "tool_call_id", "call_id"}
|
||||
)
|
||||
|
||||
|
||||
class EvoRouteFallbackMiddleware(AgentMiddleware):
|
||||
"""Retry only configured, billing-equivalent Evo fallback route models."""
|
||||
|
||||
name = "evo_route_fallback"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
fallback_models: Iterable[Any],
|
||||
route_health: Any = None,
|
||||
capacity: Any = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._fallback_models = tuple(fallback_models)
|
||||
self._route_health = route_health
|
||||
self._capacity = capacity
|
||||
|
||||
def wrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
if not self._available(request.model):
|
||||
return self._try_sync(
|
||||
request,
|
||||
handler,
|
||||
EvoRuntimeError("MODEL_ROUTE_UNAVAILABLE"),
|
||||
force_fallback=True,
|
||||
)
|
||||
try:
|
||||
return self._invoke_sync(
|
||||
self._request_for_model(request, request.model), handler
|
||||
)
|
||||
except Exception as primary_error:
|
||||
return self._try_sync(request, handler, primary_error)
|
||||
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
if not self._available(request.model):
|
||||
return await self._try_async_fallbacks(
|
||||
request,
|
||||
handler,
|
||||
EvoRuntimeError("MODEL_ROUTE_UNAVAILABLE"),
|
||||
)
|
||||
try:
|
||||
return await self._invoke_async(
|
||||
self._request_for_model(request, request.model), handler
|
||||
)
|
||||
except Exception as primary_error:
|
||||
if not _fallbackable(primary_error, request.model):
|
||||
raise
|
||||
return await self._try_async_fallbacks(request, handler, primary_error)
|
||||
|
||||
async def _try_async_fallbacks(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
primary_error: Exception,
|
||||
) -> ModelResponse:
|
||||
last_error = primary_error
|
||||
for model in self._retry_models(request, primary_error):
|
||||
if not self._available(model):
|
||||
continue
|
||||
try:
|
||||
return await self._invoke_async(
|
||||
self._request_for_model(
|
||||
request,
|
||||
model,
|
||||
repair_error=last_error,
|
||||
),
|
||||
handler,
|
||||
)
|
||||
except Exception as error:
|
||||
if not _fallbackable(error, model):
|
||||
raise
|
||||
last_error = error
|
||||
raise last_error from primary_error
|
||||
|
||||
async def _invoke_async(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
lease = await self._capacity.acquire(request.model) if self._capacity else ()
|
||||
try:
|
||||
metadata = getattr(request.model, "metadata", None) or {}
|
||||
async with asyncio.timeout(
|
||||
int(metadata.get("attempt_timeout_seconds") or 600)
|
||||
):
|
||||
response = await handler(request)
|
||||
_require_valid_response(response, request)
|
||||
return response
|
||||
finally:
|
||||
if self._capacity:
|
||||
self._capacity.release(lease)
|
||||
|
||||
def _try_sync(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
primary_error: Exception,
|
||||
*,
|
||||
force_fallback: bool = False,
|
||||
) -> ModelResponse:
|
||||
if not force_fallback and not _fallbackable(primary_error, request.model):
|
||||
raise primary_error
|
||||
last_error = primary_error
|
||||
models = (
|
||||
self._fallback_models
|
||||
if force_fallback
|
||||
else self._retry_models(request, primary_error)
|
||||
)
|
||||
for model in models:
|
||||
if not self._available(model):
|
||||
continue
|
||||
try:
|
||||
return self._invoke_sync(
|
||||
self._request_for_model(
|
||||
request,
|
||||
model,
|
||||
repair_error=last_error,
|
||||
),
|
||||
handler,
|
||||
)
|
||||
except Exception as error:
|
||||
if not _fallbackable(error, model):
|
||||
raise
|
||||
last_error = error
|
||||
raise last_error
|
||||
|
||||
@staticmethod
|
||||
def _invoke_sync(
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
response = handler(request)
|
||||
_require_valid_response(response, request)
|
||||
return response
|
||||
|
||||
def _retry_models(
|
||||
self, request: ModelRequest, error: Exception
|
||||
) -> tuple[Any, ...]:
|
||||
candidates = (
|
||||
(request.model, *self._fallback_models)
|
||||
if isinstance(error, ModelProviderResponseError)
|
||||
else self._fallback_models
|
||||
)
|
||||
result: list[Any] = []
|
||||
for model in candidates:
|
||||
if all(model is not existing for existing in result):
|
||||
result.append(model)
|
||||
return tuple(result)
|
||||
|
||||
def _available(self, model: Any) -> bool:
|
||||
if self._route_health is None:
|
||||
return True
|
||||
metadata = getattr(model, "metadata", None) or {}
|
||||
route_key = str(metadata.get("route_key") or "")
|
||||
return not route_key or not self._route_health.is_open(route_key)
|
||||
|
||||
@staticmethod
|
||||
def _request_for_model(
|
||||
request: ModelRequest,
|
||||
model: Any,
|
||||
*,
|
||||
repair_error: Exception | None = None,
|
||||
) -> ModelRequest:
|
||||
metadata = getattr(model, "metadata", None) or {}
|
||||
supports_tools = metadata.get("route_supports_tools")
|
||||
tools = [] if supports_tools is False else request.tools
|
||||
messages = (
|
||||
_without_tool_protocol(request.messages)
|
||||
if supports_tools is False
|
||||
else request.messages
|
||||
)
|
||||
messages, dropped_empty_assistants = project_provider_messages(messages)
|
||||
if dropped_empty_assistants:
|
||||
logger.info(
|
||||
"model_input_repaired route=%s dropped_empty_assistant=%s",
|
||||
str(metadata.get("route_key") or ""),
|
||||
dropped_empty_assistants,
|
||||
)
|
||||
if isinstance(repair_error, ModelToolProtocolError) and tools:
|
||||
tool_names = sorted(
|
||||
{
|
||||
name
|
||||
for tool in tools
|
||||
if (name := _tool_name(tool)) is not None
|
||||
}
|
||||
)
|
||||
allowed = ", ".join(tool_names[:64])
|
||||
if len(tool_names) > 64:
|
||||
allowed += ", ..."
|
||||
repair = SystemMessage(
|
||||
content=(
|
||||
"Retry the previous response because its structured tool call "
|
||||
f"was invalid ({repair_error.reason}). If a tool is needed, use "
|
||||
"exactly one of the supplied tool names, include a non-empty call "
|
||||
"ID, and emit arguments as one valid JSON object. Otherwise answer "
|
||||
"normally."
|
||||
+ (f" Supplied tool names: {allowed}." if allowed else "")
|
||||
)
|
||||
)
|
||||
leading_system_messages = 0
|
||||
for message in messages:
|
||||
if getattr(message, "type", "") not in {"system", "developer"}:
|
||||
break
|
||||
leading_system_messages += 1
|
||||
messages = [
|
||||
*messages[:leading_system_messages],
|
||||
repair,
|
||||
*messages[leading_system_messages:],
|
||||
]
|
||||
elif isinstance(repair_error, ModelProviderResponseError):
|
||||
logger.warning(
|
||||
"model_response_repair_retry route=%s model=%s api_mode=%s "
|
||||
"tool_transport=%s reason=%s",
|
||||
str(metadata.get("route_key") or ""),
|
||||
str(metadata.get("route_model") or ""),
|
||||
str(metadata.get("route_api_mode") or ""),
|
||||
str(metadata.get("route_tool_call_transport") or ""),
|
||||
repair_error.reason,
|
||||
)
|
||||
repair = SystemMessage(
|
||||
content=(
|
||||
"Retry the response because the previous attempt completed "
|
||||
"without final text or a structured tool call. Complete the "
|
||||
"request with a user-visible final answer, or emit a valid "
|
||||
"tool call when a tool is required."
|
||||
)
|
||||
)
|
||||
leading_system_messages = 0
|
||||
for message in messages:
|
||||
if getattr(message, "type", "") not in {"system", "developer"}:
|
||||
break
|
||||
leading_system_messages += 1
|
||||
messages = [
|
||||
*messages[:leading_system_messages],
|
||||
repair,
|
||||
*messages[leading_system_messages:],
|
||||
]
|
||||
overrides: dict[str, Any] = {
|
||||
"model": model,
|
||||
"tools": tools,
|
||||
"messages": messages,
|
||||
# The invocation plan has already compiled every provider SDK
|
||||
# parameter. Agent-level settings must not mutate it afterwards.
|
||||
"model_settings": {},
|
||||
}
|
||||
if supports_tools is False:
|
||||
# A no-tools route cannot accept an inherited forced tool choice or
|
||||
# structured-output contract from a prior model invocation.
|
||||
overrides["tool_choice"] = None
|
||||
overrides["response_format"] = None
|
||||
return request.override(**overrides)
|
||||
|
||||
|
||||
def _without_tool_protocol(messages: list[Any]) -> list[Any]:
|
||||
"""Make a checkpoint replayable by a route that does not support tools.
|
||||
|
||||
A model switch can leave completed ToolMessage/AI tool-call pairs in the
|
||||
checkpointer. Sending those pairs while omitting ``tools`` is rejected by
|
||||
strict OpenAI-compatible providers. Keep normal assistant text, remove
|
||||
protocol-only fields, and retain a bounded plain-text transcript of tool
|
||||
results so a text-only model does not lose work completed before a switch.
|
||||
"""
|
||||
|
||||
sanitized: list[Any] = []
|
||||
for message in messages:
|
||||
if isinstance(message, ToolMessage) or getattr(message, "type", "") == "tool":
|
||||
transcript = _tool_result_transcript(message)
|
||||
if transcript:
|
||||
sanitized.append(AIMessage(content=transcript))
|
||||
continue
|
||||
if not isinstance(message, AIMessage):
|
||||
sanitized.append(message)
|
||||
continue
|
||||
additional = dict(getattr(message, "additional_kwargs", {}) or {})
|
||||
has_tool_protocol = bool(
|
||||
getattr(message, "tool_calls", None)
|
||||
or getattr(message, "invalid_tool_calls", None)
|
||||
or additional.get("tool_calls")
|
||||
or _contains_tool_content(getattr(message, "content", ""))
|
||||
)
|
||||
if not has_tool_protocol:
|
||||
sanitized.append(message)
|
||||
continue
|
||||
portable_content = _portable_assistant_content(getattr(message, "content", ""))
|
||||
# An empty assistant turn only represented a tool call. Its following
|
||||
# ToolMessage becomes a transcript entry, so there is nothing useful to
|
||||
# send for this message itself.
|
||||
if not portable_content:
|
||||
continue
|
||||
for key in _TOOL_PROTOCOL_METADATA:
|
||||
additional.pop(key, None)
|
||||
sanitized.append(
|
||||
message.model_copy(
|
||||
update={
|
||||
"content": portable_content,
|
||||
"tool_calls": [],
|
||||
"invalid_tool_calls": [],
|
||||
"additional_kwargs": additional,
|
||||
}
|
||||
)
|
||||
)
|
||||
return sanitized
|
||||
|
||||
|
||||
def _portable_assistant_content(content: Any) -> str:
|
||||
"""Keep textual assistant content while removing embedded tool-call blocks."""
|
||||
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if not isinstance(content, list):
|
||||
return str(content or "")
|
||||
parts: list[str] = []
|
||||
for block in content:
|
||||
if isinstance(block, str):
|
||||
parts.append(block)
|
||||
continue
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
block_type = str(block.get("type") or "")
|
||||
if block_type in {"tool_call", "tool_use", "function_call"}:
|
||||
continue
|
||||
text = block.get("text")
|
||||
if isinstance(text, str):
|
||||
parts.append(text)
|
||||
return "\n".join(part for part in parts if part)
|
||||
|
||||
|
||||
def _contains_tool_content(content: Any) -> bool:
|
||||
"""Identify provider content blocks that encode a tool call."""
|
||||
|
||||
return isinstance(content, list) and any(
|
||||
isinstance(block, dict)
|
||||
and str(block.get("type") or "")
|
||||
in {"tool_call", "tool_use", "function_call"}
|
||||
for block in content
|
||||
)
|
||||
|
||||
|
||||
def _tool_result_transcript(message: Any) -> str:
|
||||
"""Project a completed tool result to bounded text for a no-tools route."""
|
||||
|
||||
content = getattr(message, "content", "")
|
||||
if isinstance(content, list):
|
||||
parts = []
|
||||
for block in content:
|
||||
if isinstance(block, str):
|
||||
parts.append(block)
|
||||
elif isinstance(block, dict) and isinstance(block.get("text"), str):
|
||||
parts.append(block["text"])
|
||||
text = "\n".join(parts)
|
||||
else:
|
||||
text = str(content or "")
|
||||
text = text.strip()
|
||||
if not text:
|
||||
return ""
|
||||
if len(text) > _TOOL_RESULT_CONTEXT_LIMIT:
|
||||
text = text[:_TOOL_RESULT_CONTEXT_LIMIT] + "\n[Tool result truncated]"
|
||||
name = str(getattr(message, "name", "") or "tool")[:128]
|
||||
return f"[Completed tool result: {name}]\n{text}"
|
||||
|
||||
|
||||
def _tool_name(tool: Any) -> str | None:
|
||||
if isinstance(tool, dict):
|
||||
value = tool.get("name")
|
||||
function = tool.get("function")
|
||||
if not value and isinstance(function, dict):
|
||||
value = function.get("name")
|
||||
else:
|
||||
value = getattr(tool, "name", None)
|
||||
return value.strip() if isinstance(value, str) and value.strip() else None
|
||||
|
||||
|
||||
def _fallbackable(error: Exception, model: Any) -> bool:
|
||||
if isinstance(error, EvoRuntimeError):
|
||||
return False
|
||||
explicit = getattr(error, "fallbackable", None)
|
||||
if explicit is not None:
|
||||
return bool(explicit)
|
||||
if getattr(error, "non_fallbackable", False):
|
||||
return False
|
||||
metadata = getattr(model, "metadata", None) or {}
|
||||
adapter_id = str(metadata.get("route_adapter_id") or "")
|
||||
adapter_revision = str(metadata.get("route_adapter_revision") or "")
|
||||
if not adapter_id or not adapter_revision:
|
||||
return isinstance(error, (ConnectionError, TimeoutError))
|
||||
from ..llm.adapter_registry import get_adapter_registry
|
||||
|
||||
return (
|
||||
get_adapter_registry()
|
||||
.get(adapter_id, adapter_revision)
|
||||
.classify_error(error)
|
||||
.retryable
|
||||
)
|
||||
|
||||
|
||||
def _require_valid_response(response: Any, request: ModelRequest) -> None:
|
||||
"""Reject completed assistant responses that cannot advance the agent."""
|
||||
|
||||
messages: list[AIMessage]
|
||||
if isinstance(response, AIMessage):
|
||||
messages = [response]
|
||||
else:
|
||||
result = getattr(response, "result", None)
|
||||
messages = (
|
||||
[message for message in result if isinstance(message, AIMessage)]
|
||||
if isinstance(result, list | tuple)
|
||||
else []
|
||||
)
|
||||
if messages and not any(assistant_message_has_output(message) for message in messages):
|
||||
metadata = getattr(request.model, "metadata", None) or {}
|
||||
logger.warning(
|
||||
"model_provider_response_invalid route=%s model=%s api_mode=%s "
|
||||
"tool_transport=%s reason=empty_assistant_response",
|
||||
str(metadata.get("route_key") or ""),
|
||||
str(metadata.get("route_model") or ""),
|
||||
str(metadata.get("route_api_mode") or ""),
|
||||
str(metadata.get("route_tool_call_transport") or ""),
|
||||
)
|
||||
raise ModelProviderResponseError()
|
||||
@@ -14,6 +14,7 @@ can deal with them.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import threading
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
@@ -63,6 +64,20 @@ _AUTH_ERROR_PATTERNS: list[str] = [
|
||||
These are intentionally *not* treated as non-fallbackable because a different
|
||||
provider in the chain may have valid credentials."""
|
||||
|
||||
_SAFE_ERROR_CODE = re.compile(r"^[A-Z][A-Z0-9_]{1,63}$")
|
||||
|
||||
|
||||
def _safe_error_label(exc: BaseException) -> str:
|
||||
"""Describe an error without rendering a provider-controlled response body."""
|
||||
label = type(exc).__name__
|
||||
status = getattr(exc, "status_code", None)
|
||||
if isinstance(status, int) and 100 <= status <= 599:
|
||||
label = f"{label} status={status}"
|
||||
code = getattr(exc, "code", None)
|
||||
if isinstance(code, str) and _SAFE_ERROR_CODE.fullmatch(code):
|
||||
label = f"{label} code={code}"
|
||||
return label
|
||||
|
||||
|
||||
def set_ui_emit(fn: Callable[[str, str], None] | None) -> None:
|
||||
"""Register (or clear) the UI callback for fallback status messages.
|
||||
@@ -260,13 +275,9 @@ async def _try_fallbacks(
|
||||
"""
|
||||
from ..llm.models import get_chat_model
|
||||
|
||||
_emit(
|
||||
f"Primary model failed: {type(primary_exc).__name__}: {primary_exc}",
|
||||
style="yellow",
|
||||
)
|
||||
logger.warning(
|
||||
"Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc
|
||||
)
|
||||
primary_label = _safe_error_label(primary_exc)
|
||||
_emit(f"Primary model failed: {primary_label}", style="yellow")
|
||||
logger.warning("Primary model failed: %s", primary_label)
|
||||
|
||||
# Track the request whose model actually raised ``last_exc`` so we
|
||||
# can attribute the exception to the failing model, not the
|
||||
@@ -280,7 +291,7 @@ async def _try_fallbacks(
|
||||
for model_name, provider in get_fallback_chain():
|
||||
_emit(
|
||||
f" -> Falling back to {model_name} ({provider}) "
|
||||
f"due to: {type(last_exc).__name__}: {last_exc}",
|
||||
f"due to: {_safe_error_label(last_exc)}",
|
||||
style="yellow",
|
||||
)
|
||||
try:
|
||||
@@ -305,15 +316,14 @@ async def _try_fallbacks(
|
||||
last_exc = fb_exc
|
||||
last_failing_request = fb_request
|
||||
_emit(
|
||||
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
|
||||
f" x {model_name} also failed: {_safe_error_label(fb_exc)}",
|
||||
style="red",
|
||||
)
|
||||
logger.warning(
|
||||
"Fallback %s (provider=%s) failed: %s: %s",
|
||||
"Fallback %s (provider=%s) failed: %s",
|
||||
model_name,
|
||||
provider,
|
||||
type(fb_exc).__name__,
|
||||
fb_exc,
|
||||
_safe_error_label(fb_exc),
|
||||
)
|
||||
|
||||
_emit(" All fallbacks exhausted -- re-raising last error", style="red")
|
||||
|
||||
@@ -0,0 +1,376 @@
|
||||
"""Bound provider context by externalizing assistant-generated inline media."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import logging
|
||||
import mimetypes
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import replace
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
ExtendedModelResponse,
|
||||
ModelRequest,
|
||||
ModelResponse,
|
||||
)
|
||||
from langchain_core.messages import AIMessage, BaseMessage
|
||||
from langgraph.types import Overwrite
|
||||
|
||||
from ..llm.contracts import EvoRuntimeError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_MAX_INLINE_MEDIA_BYTES = 16_777_216
|
||||
_MEDIA_PREFIX = "/artifacts/model-output"
|
||||
|
||||
|
||||
def _decode_base64_block(block: Mapping[str, Any]) -> tuple[bytes, str] | None:
|
||||
payload = block.get("base64")
|
||||
mime = str(block.get("mime_type") or "application/octet-stream")
|
||||
if not isinstance(payload, str) or not payload:
|
||||
return None
|
||||
try:
|
||||
return base64.b64decode(payload, validate=True), mime
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise EvoRuntimeError(
|
||||
"MEDIA_PERSIST_FAILED",
|
||||
details=({"reason": "invalid_base64", "media_type": mime},),
|
||||
) from exc
|
||||
|
||||
|
||||
def _extension(mime: str) -> str:
|
||||
return (mimetypes.guess_extension(mime) or ".bin").lstrip(".")
|
||||
|
||||
|
||||
def _artifact_reference(path: str, mime: str, digest: str) -> dict[str, Any]:
|
||||
major = mime.split("/", 1)[0]
|
||||
if major in {"image", "audio", "video"}:
|
||||
return {
|
||||
"type": major,
|
||||
"url": path,
|
||||
"mime_type": mime,
|
||||
"media_id": f"sha256:{digest}",
|
||||
}
|
||||
return {
|
||||
"type": "text",
|
||||
"text": f'<generated_file path="{path}" media_type="{mime}" sha256="{digest}" />',
|
||||
}
|
||||
|
||||
|
||||
def _provider_reference(block: Mapping[str, Any]) -> dict[str, str] | None:
|
||||
block_type = str(block.get("type") or "")
|
||||
path = block.get("url")
|
||||
if (
|
||||
block_type not in {"image", "audio", "video"}
|
||||
or not isinstance(path, str)
|
||||
or not path.startswith(f"{_MEDIA_PREFIX}/")
|
||||
):
|
||||
return None
|
||||
mime = str(block.get("mime_type") or f"{block_type}/unknown")
|
||||
media_id = str(block.get("media_id") or "")
|
||||
media_id_attribute = f' media_id="{media_id}"' if media_id else ""
|
||||
return {
|
||||
"type": "text",
|
||||
"text": (
|
||||
f'<generated_{block_type} path="{path}" media_type="{mime}"'
|
||||
f"{media_id_attribute} />"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _assistant_messages(messages: Sequence[BaseMessage]) -> list[AIMessage]:
|
||||
return [message for message in messages if isinstance(message, AIMessage)]
|
||||
|
||||
|
||||
def _collect_inline_media(
|
||||
messages: Sequence[BaseMessage],
|
||||
*,
|
||||
max_inline_media_bytes: int,
|
||||
) -> dict[str, tuple[str, bytes]]:
|
||||
collected: dict[str, tuple[str, bytes]] = {}
|
||||
for message in _assistant_messages(messages):
|
||||
content = message.content
|
||||
if not isinstance(content, list):
|
||||
continue
|
||||
for value in content:
|
||||
if not isinstance(value, Mapping):
|
||||
continue
|
||||
decoded = _decode_base64_block(value)
|
||||
if decoded is None:
|
||||
continue
|
||||
raw, mime = decoded
|
||||
if len(raw) > max_inline_media_bytes:
|
||||
raise EvoRuntimeError(
|
||||
"MEDIA_PERSIST_FAILED",
|
||||
details=(
|
||||
{
|
||||
"reason": "media_too_large",
|
||||
"media_type": mime,
|
||||
"media_bytes": len(raw),
|
||||
},
|
||||
),
|
||||
)
|
||||
digest = hashlib.sha256(raw).hexdigest()
|
||||
collected.setdefault(digest, (mime, raw))
|
||||
return collected
|
||||
|
||||
|
||||
def _rewrite_messages(
|
||||
messages: Sequence[BaseMessage],
|
||||
paths: Mapping[str, tuple[str, str]],
|
||||
*,
|
||||
for_provider: bool,
|
||||
) -> list[BaseMessage]:
|
||||
rewritten: list[BaseMessage] = []
|
||||
for message in messages:
|
||||
if not isinstance(message, AIMessage) or not isinstance(message.content, list):
|
||||
rewritten.append(message)
|
||||
continue
|
||||
modified = False
|
||||
content: list[Any] = []
|
||||
for value in message.content:
|
||||
if not isinstance(value, Mapping):
|
||||
content.append(value)
|
||||
continue
|
||||
if for_provider:
|
||||
reference = _provider_reference(value)
|
||||
if reference is not None:
|
||||
content.append(reference)
|
||||
modified = True
|
||||
continue
|
||||
decoded = _decode_base64_block(value)
|
||||
if decoded is None:
|
||||
content.append(value)
|
||||
continue
|
||||
raw, mime = decoded
|
||||
digest = hashlib.sha256(raw).hexdigest()
|
||||
path_entry = paths.get(digest)
|
||||
if path_entry is None:
|
||||
raise EvoRuntimeError(
|
||||
"MEDIA_PERSIST_FAILED",
|
||||
details=({"reason": "artifact_path_missing", "media_type": mime},),
|
||||
)
|
||||
path, stored_mime = path_entry
|
||||
reference = _artifact_reference(path, stored_mime, digest)
|
||||
content.append(
|
||||
_provider_reference(reference) if for_provider else reference
|
||||
)
|
||||
modified = True
|
||||
if modified:
|
||||
copy = message.model_copy()
|
||||
copy.content = content
|
||||
rewritten.append(copy)
|
||||
else:
|
||||
rewritten.append(message)
|
||||
return rewritten
|
||||
|
||||
|
||||
def _response_messages(
|
||||
response: Any,
|
||||
) -> tuple[list[BaseMessage], Callable[[list[BaseMessage]], Any]]:
|
||||
if isinstance(response, ExtendedModelResponse):
|
||||
return response.model_response.result, lambda result: replace(
|
||||
response,
|
||||
model_response=replace(response.model_response, result=result),
|
||||
)
|
||||
if isinstance(response, ModelResponse):
|
||||
return response.result, lambda result: replace(response, result=result)
|
||||
if isinstance(response, AIMessage):
|
||||
return [response], lambda result: result[0]
|
||||
return [], lambda _result: response
|
||||
|
||||
|
||||
class ProviderContextMediaMiddleware(AgentMiddleware):
|
||||
"""Persist assistant media and keep base64 out of later provider calls."""
|
||||
|
||||
name = "provider_context_media"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
backend: Any,
|
||||
*,
|
||||
max_inline_media_bytes: int = _DEFAULT_MAX_INLINE_MEDIA_BYTES,
|
||||
) -> None:
|
||||
self.backend = backend
|
||||
self.max_inline_media_bytes = max(1, int(max_inline_media_bytes))
|
||||
|
||||
@staticmethod
|
||||
def _paths_for(
|
||||
media: Mapping[str, tuple[str, bytes]],
|
||||
) -> dict[str, tuple[str, str]]:
|
||||
return {
|
||||
digest: (f"{_MEDIA_PREFIX}/{digest[:24]}.{_extension(mime)}", mime)
|
||||
for digest, (mime, _raw) in media.items()
|
||||
}
|
||||
|
||||
def _persist(self, messages: Sequence[BaseMessage]) -> dict[str, tuple[str, str]]:
|
||||
media = _collect_inline_media(
|
||||
messages,
|
||||
max_inline_media_bytes=self.max_inline_media_bytes,
|
||||
)
|
||||
paths = self._paths_for(media)
|
||||
for digest, (mime, raw) in media.items():
|
||||
path = paths[digest][0]
|
||||
responses = self.backend.upload_files([(path, raw)])
|
||||
error = (
|
||||
getattr(responses[0], "error", None)
|
||||
if responses
|
||||
else "missing upload response"
|
||||
)
|
||||
if error:
|
||||
raise EvoRuntimeError(
|
||||
"MEDIA_PERSIST_FAILED",
|
||||
details=({"reason": "artifact_upload_failed", "media_type": mime},),
|
||||
)
|
||||
logger.info(
|
||||
"provider_context_media_persisted path=%s media_type=%s bytes=%s",
|
||||
path,
|
||||
mime,
|
||||
len(raw),
|
||||
)
|
||||
return paths
|
||||
|
||||
async def _apersist(
|
||||
self, messages: Sequence[BaseMessage]
|
||||
) -> dict[str, tuple[str, str]]:
|
||||
media = _collect_inline_media(
|
||||
messages,
|
||||
max_inline_media_bytes=self.max_inline_media_bytes,
|
||||
)
|
||||
paths = self._paths_for(media)
|
||||
for digest, (mime, raw) in media.items():
|
||||
path = paths[digest][0]
|
||||
responses = await self.backend.aupload_files([(path, raw)])
|
||||
error = (
|
||||
getattr(responses[0], "error", None)
|
||||
if responses
|
||||
else "missing upload response"
|
||||
)
|
||||
if error:
|
||||
raise EvoRuntimeError(
|
||||
"MEDIA_PERSIST_FAILED",
|
||||
details=({"reason": "artifact_upload_failed", "media_type": mime},),
|
||||
)
|
||||
logger.info(
|
||||
"provider_context_media_persisted path=%s media_type=%s bytes=%s",
|
||||
path,
|
||||
mime,
|
||||
len(raw),
|
||||
)
|
||||
return paths
|
||||
|
||||
def _prepare_request(self, request: ModelRequest) -> ModelRequest:
|
||||
paths = self._persist(request.messages)
|
||||
messages = _rewrite_messages(
|
||||
request.messages,
|
||||
paths,
|
||||
for_provider=True,
|
||||
)
|
||||
return request.override(messages=messages)
|
||||
|
||||
async def _aprepare_request(self, request: ModelRequest) -> ModelRequest:
|
||||
paths = await self._apersist(request.messages)
|
||||
messages = _rewrite_messages(
|
||||
request.messages,
|
||||
paths,
|
||||
for_provider=True,
|
||||
)
|
||||
return request.override(messages=messages)
|
||||
|
||||
def _prepare_response(self, response: Any) -> Any:
|
||||
messages, rebuild = _response_messages(response)
|
||||
if not messages:
|
||||
return response
|
||||
paths = self._persist(messages)
|
||||
return rebuild(_rewrite_messages(messages, paths, for_provider=False))
|
||||
|
||||
async def _aprepare_response(self, response: Any) -> Any:
|
||||
messages, rebuild = _response_messages(response)
|
||||
if not messages:
|
||||
return response
|
||||
paths = await self._apersist(messages)
|
||||
return rebuild(_rewrite_messages(messages, paths, for_provider=False))
|
||||
|
||||
def before_model(self, state: Any, runtime: Any) -> dict[str, Any] | None:
|
||||
_ = runtime
|
||||
try:
|
||||
messages = state.get("messages") if isinstance(state, Mapping) else None
|
||||
if not isinstance(messages, Sequence) or isinstance(messages, str | bytes):
|
||||
return None
|
||||
original = list(messages)
|
||||
paths = self._persist(original)
|
||||
rewritten = _rewrite_messages(original, paths, for_provider=False)
|
||||
if all(
|
||||
left is right for left, right in zip(original, rewritten, strict=True)
|
||||
):
|
||||
return None
|
||||
logger.info(
|
||||
"provider_context_media_checkpoint_repaired messages=%s",
|
||||
len(rewritten),
|
||||
)
|
||||
return {"messages": Overwrite(rewritten)}
|
||||
except EvoRuntimeError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise self._middleware_failure("before_model", exc) from exc
|
||||
|
||||
async def abefore_model(self, state: Any, runtime: Any) -> dict[str, Any] | None:
|
||||
_ = runtime
|
||||
try:
|
||||
messages = state.get("messages") if isinstance(state, Mapping) else None
|
||||
if not isinstance(messages, Sequence) or isinstance(messages, str | bytes):
|
||||
return None
|
||||
original = list(messages)
|
||||
paths = await self._apersist(original)
|
||||
rewritten = _rewrite_messages(original, paths, for_provider=False)
|
||||
if all(
|
||||
left is right for left, right in zip(original, rewritten, strict=True)
|
||||
):
|
||||
return None
|
||||
logger.info(
|
||||
"provider_context_media_checkpoint_repaired messages=%s",
|
||||
len(rewritten),
|
||||
)
|
||||
return {"messages": Overwrite(rewritten)}
|
||||
except EvoRuntimeError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise self._middleware_failure("before_model", exc) from exc
|
||||
|
||||
@staticmethod
|
||||
def _middleware_failure(node: str, exc: Exception) -> EvoRuntimeError:
|
||||
return EvoRuntimeError(
|
||||
"AGENT_MIDDLEWARE_FAILED",
|
||||
details=(
|
||||
{
|
||||
"failure_stage": "agent_middleware",
|
||||
"middleware": ProviderContextMediaMiddleware.name,
|
||||
"middleware_node": (
|
||||
f"{ProviderContextMediaMiddleware.name}.{node}"
|
||||
),
|
||||
"agent_error_type": type(exc).__name__[:128],
|
||||
"agent_error_module": type(exc).__module__[:128],
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
def wrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
return self._prepare_response(handler(self._prepare_request(request)))
|
||||
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
prepared = await self._aprepare_request(request)
|
||||
return await self._aprepare_response(await handler(prepared))
|
||||
|
||||
|
||||
__all__ = ["ProviderContextMediaMiddleware"]
|
||||
@@ -0,0 +1,272 @@
|
||||
"""Durable model-attempt metering for Graph-native Ai4Sci runs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
ModelRequest,
|
||||
ModelResponse,
|
||||
)
|
||||
from langchain_core.callbacks import AsyncCallbackHandler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_clients: dict[str, httpx.AsyncClient] = {}
|
||||
|
||||
|
||||
def _config() -> dict[str, Any] | None:
|
||||
try:
|
||||
from langgraph.config import get_config
|
||||
|
||||
value = get_config()
|
||||
except Exception:
|
||||
return None
|
||||
return value if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _metering_config(config: Mapping[str, Any] | None) -> dict[str, str] | None:
|
||||
if not config:
|
||||
return None
|
||||
configurable = config.get("configurable")
|
||||
if not isinstance(configurable, Mapping):
|
||||
return None
|
||||
value = configurable.get("ai4sci_metering")
|
||||
if not isinstance(value, Mapping):
|
||||
return None
|
||||
required = ("gateway_url", "run_id", "envelope_signature")
|
||||
normalized = {
|
||||
name: str(value.get(name) or "")
|
||||
for name in (*required, "provider_id", "model_id", "source_type")
|
||||
}
|
||||
return normalized if all(normalized[name] for name in required) else None
|
||||
|
||||
|
||||
def _first(mapping: Mapping[str, Any], *keys: str) -> Any:
|
||||
for key in keys:
|
||||
value = mapping.get(key)
|
||||
if value not in (None, ""):
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _source_type(metadata: Mapping[str, Any], tags: list[str]) -> str:
|
||||
explicit = str(metadata.get("metering_scope") or "").lower()
|
||||
aliases = {
|
||||
"main": "main_agent",
|
||||
"main_agent": "main_agent",
|
||||
"subagent": "subagent",
|
||||
"tool_selector": "tool_selector",
|
||||
"summarizer": "summarizer",
|
||||
"title": "title",
|
||||
"memory": "evomemory_turn_worker",
|
||||
"evomemory_turn_worker": "evomemory_turn_worker",
|
||||
"evomemory_subagent_worker": "evomemory_subagent_worker",
|
||||
"evomemory_linker": "evomemory_linker",
|
||||
}
|
||||
if explicit in aliases:
|
||||
return aliases[explicit]
|
||||
hint = " ".join([*tags, *(str(value) for value in metadata.values())]).lower()
|
||||
if "selector" in hint:
|
||||
return "tool_selector"
|
||||
if "summar" in hint or "compact" in hint:
|
||||
return "summarizer"
|
||||
if "title" in hint:
|
||||
return "title"
|
||||
if "memory" in hint or "observation" in hint:
|
||||
return "evomemory_turn_worker"
|
||||
if "subagent" in hint or "sub_agent" in hint or "task:" in hint:
|
||||
return "subagent"
|
||||
return "main_agent"
|
||||
|
||||
|
||||
def _usage(response: Any) -> dict[str, Any] | None:
|
||||
llm_output = dict(getattr(response, "llm_output", None) or {})
|
||||
candidates: list[Mapping[str, Any]] = [
|
||||
dict(llm_output.get("token_usage") or llm_output.get("usage") or {})
|
||||
]
|
||||
for group in getattr(response, "generations", None) or []:
|
||||
for generation in group or []:
|
||||
message = getattr(generation, "message", None)
|
||||
if message is not None:
|
||||
candidates.append(dict(getattr(message, "usage_metadata", None) or {}))
|
||||
for value in candidates:
|
||||
input_tokens = _first(value, "input_tokens", "prompt_tokens", "input_token_count")
|
||||
output_tokens = _first(value, "output_tokens", "completion_tokens", "output_token_count")
|
||||
if input_tokens is None or output_tokens is None:
|
||||
continue
|
||||
input_details = dict(
|
||||
value.get("input_token_details") or value.get("prompt_tokens_details") or {}
|
||||
)
|
||||
output_details = dict(
|
||||
value.get("output_token_details") or value.get("completion_tokens_details") or {}
|
||||
)
|
||||
normalized_input = max(0, int(input_tokens))
|
||||
normalized_output = max(0, int(output_tokens))
|
||||
return {
|
||||
"input_tokens": normalized_input,
|
||||
"output_tokens": normalized_output,
|
||||
"cached_input_tokens": max(
|
||||
0,
|
||||
int(
|
||||
_first(
|
||||
input_details,
|
||||
"cache_read",
|
||||
"cached_tokens",
|
||||
"cache_read_input_tokens",
|
||||
)
|
||||
or 0
|
||||
),
|
||||
),
|
||||
"reasoning_tokens": max(
|
||||
0, int(_first(output_details, "reasoning", "reasoning_tokens") or 0)
|
||||
),
|
||||
"total_tokens": normalized_input + normalized_output,
|
||||
}
|
||||
return None
|
||||
|
||||
|
||||
async def _post(config: Mapping[str, str], phase: str, payload: dict[str, Any]) -> None:
|
||||
base_url = config["gateway_url"]
|
||||
client = _clients.get(base_url)
|
||||
if client is None:
|
||||
client = httpx.AsyncClient(timeout=httpx.Timeout(15.0, connect=3.0))
|
||||
_clients[base_url] = client
|
||||
response = await client.post(
|
||||
f"{base_url}/api/internal/recoverable-runs/metering/{phase}",
|
||||
json={
|
||||
**payload,
|
||||
"run_id": config["run_id"],
|
||||
"envelope_signature": config["envelope_signature"],
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
|
||||
class RecoverableMeteringCallback(AsyncCallbackHandler):
|
||||
raise_error = True
|
||||
|
||||
def __init__(self, config: Mapping[str, str]) -> None:
|
||||
super().__init__()
|
||||
self.config = dict(config)
|
||||
self._attempts: set[uuid.UUID] = set()
|
||||
|
||||
async def on_chat_model_start(
|
||||
self,
|
||||
serialized: dict[str, Any],
|
||||
messages: list[list[Any]],
|
||||
*,
|
||||
run_id: uuid.UUID,
|
||||
tags: list[str] | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
del messages
|
||||
metadata = dict(metadata or {})
|
||||
invocation = dict(kwargs.get("invocation_params") or {})
|
||||
serialized_kwargs = dict(serialized.get("kwargs") or {})
|
||||
provider = self.config.get("provider_id") or str(
|
||||
_first(metadata, "route_provider", "ls_provider", "provider")
|
||||
or _first(invocation, "provider", "model_provider")
|
||||
or serialized_kwargs.get("model_provider")
|
||||
or "unknown"
|
||||
)
|
||||
model = self.config.get("model_id") or str(
|
||||
_first(metadata, "route_model", "ls_model_name", "model")
|
||||
or _first(invocation, "model", "model_name")
|
||||
or _first(serialized_kwargs, "model", "model_name")
|
||||
or serialized.get("name")
|
||||
or ""
|
||||
)
|
||||
await _post(
|
||||
self.config,
|
||||
"start",
|
||||
{
|
||||
"attempt_id": str(run_id),
|
||||
"source_type": self.config.get("source_type")
|
||||
or _source_type(metadata, list(tags or [])),
|
||||
"provider_id": provider,
|
||||
"model_id": model,
|
||||
},
|
||||
)
|
||||
self._attempts.add(run_id)
|
||||
|
||||
async def on_llm_end(self, response: Any, *, run_id: uuid.UUID, **kwargs: Any) -> None:
|
||||
del kwargs
|
||||
await self._terminal(run_id, "succeeded", _usage(response))
|
||||
|
||||
async def on_llm_error(
|
||||
self, error: BaseException, *, run_id: uuid.UUID, **kwargs: Any
|
||||
) -> None:
|
||||
del error, kwargs
|
||||
await self._terminal(run_id, "failed", None)
|
||||
|
||||
async def _terminal(
|
||||
self, run_id: uuid.UUID, outcome: str, usage: dict[str, Any] | None
|
||||
) -> None:
|
||||
if run_id not in self._attempts:
|
||||
return
|
||||
last_error: BaseException | None = None
|
||||
for attempt in range(3):
|
||||
try:
|
||||
await _post(
|
||||
self.config,
|
||||
"terminal",
|
||||
{"attempt_id": str(run_id), "outcome": outcome, "usage": usage},
|
||||
)
|
||||
self._attempts.discard(run_id)
|
||||
return
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
await asyncio.sleep(0.1 * (attempt + 1))
|
||||
raise RuntimeError("AI4SCI_METERING_TERMINAL_FAILED") from last_error
|
||||
|
||||
|
||||
class RecoverableMeteringMiddleware(AgentMiddleware):
|
||||
"""Attach one durable callback manager to every model invoked by this Run."""
|
||||
|
||||
name = "recoverable_metering"
|
||||
|
||||
@staticmethod
|
||||
def _install() -> None:
|
||||
config = _config()
|
||||
metering = _metering_config(config)
|
||||
if config is None or metering is None:
|
||||
return
|
||||
callbacks = config.get("callbacks")
|
||||
if callbacks is None:
|
||||
config["callbacks"] = [RecoverableMeteringCallback(metering)]
|
||||
return
|
||||
if isinstance(callbacks, list):
|
||||
if not any(isinstance(item, RecoverableMeteringCallback) for item in callbacks):
|
||||
callbacks.append(RecoverableMeteringCallback(metering))
|
||||
return
|
||||
if hasattr(callbacks, "add_handler"):
|
||||
handlers = list(getattr(callbacks, "handlers", []) or [])
|
||||
if not any(isinstance(item, RecoverableMeteringCallback) for item in handlers):
|
||||
callbacks.add_handler(RecoverableMeteringCallback(metering), inherit=True)
|
||||
return
|
||||
logger.error("Unsupported LangChain callback container for recoverable metering")
|
||||
raise RuntimeError("AI4SCI_METERING_CALLBACKS_UNAVAILABLE")
|
||||
|
||||
def wrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
if _metering_config(_config()) is None:
|
||||
return handler(request)
|
||||
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH")
|
||||
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
self._install()
|
||||
return await handler(request)
|
||||
@@ -0,0 +1,205 @@
|
||||
"""At-least-once tool-effect recovery for Ai4Sci Graph runs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import httpx
|
||||
from langchain.agents.middleware.types import AgentMiddleware
|
||||
from langchain_core.messages import ToolMessage, message_to_dict, messages_from_dict
|
||||
from langgraph.types import Command
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain.agents.middleware.types import ToolCallRequest
|
||||
|
||||
_READ_ONLY_PREFIXES = ("read_", "get_", "list_", "search_", "find_", "check_")
|
||||
_READ_ONLY_NAMES = {
|
||||
"web_search",
|
||||
"glob",
|
||||
"grep",
|
||||
"ls",
|
||||
"view_image",
|
||||
"fetch_url",
|
||||
}
|
||||
_IDEMPOTENT_NAMES = {
|
||||
"write_file",
|
||||
"create_directory",
|
||||
"mkdir",
|
||||
"update_file",
|
||||
}
|
||||
|
||||
|
||||
def _canonical(value: Any) -> str:
|
||||
return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=str)
|
||||
|
||||
|
||||
def _hash(value: Any) -> str:
|
||||
return hashlib.sha256(_canonical(value).encode()).hexdigest()
|
||||
|
||||
|
||||
def _context() -> tuple[dict[str, str] | None, dict[str, Any]]:
|
||||
try:
|
||||
from langgraph.config import get_config
|
||||
|
||||
config = get_config()
|
||||
except Exception:
|
||||
return None, {}
|
||||
configurable = config.get("configurable") if isinstance(config, dict) else None
|
||||
if not isinstance(configurable, Mapping):
|
||||
return None, {}
|
||||
proxy = configurable.get("ai4sci_tool_effect")
|
||||
metadata = dict(config.get("metadata") or {})
|
||||
if not isinstance(proxy, Mapping):
|
||||
# Compatibility for Runs dispatched before the tool-effect grant was
|
||||
# split from the model proxy. Detached EvoMemory graphs must never use
|
||||
# the parent conversation's tool-effect authority.
|
||||
run_kind = str(metadata.get("run_kind") or "")
|
||||
if not run_kind.startswith("evomemory_"):
|
||||
proxy = configurable.get("ai4sci_model_proxy")
|
||||
if not isinstance(proxy, Mapping):
|
||||
return None, metadata
|
||||
normalized = {
|
||||
name: str(proxy.get(name) or "")
|
||||
for name in ("gateway_url", "run_id", "envelope_signature")
|
||||
}
|
||||
return (
|
||||
normalized if all(normalized.values()) else None,
|
||||
metadata,
|
||||
)
|
||||
|
||||
|
||||
def _effect_class(tool_name: str) -> str:
|
||||
lowered = tool_name.lower()
|
||||
if lowered in _READ_ONLY_NAMES or lowered.startswith(_READ_ONLY_PREFIXES):
|
||||
return "read_only"
|
||||
if lowered in _IDEMPOTENT_NAMES:
|
||||
return "idempotent"
|
||||
return "non_idempotent"
|
||||
|
||||
|
||||
async def _post(proxy: Mapping[str, str], phase: str, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(30.0, connect=3.0)) as client:
|
||||
response = await client.post(
|
||||
f"{proxy['gateway_url'].rstrip('/')}/api/internal/recoverable-runs/tool-effect/{phase}",
|
||||
json={
|
||||
**payload,
|
||||
"run_id": proxy["run_id"],
|
||||
"attempt_id": proxy["run_id"],
|
||||
"envelope_signature": proxy["envelope_signature"],
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
return dict(response.json())
|
||||
|
||||
|
||||
class RecoverableToolEffectMiddleware(AgentMiddleware):
|
||||
name = "recoverable_tool_effect"
|
||||
|
||||
def wrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]],
|
||||
) -> ToolMessage | Command[Any]:
|
||||
proxy, _ = _context()
|
||||
if proxy is None:
|
||||
return handler(request)
|
||||
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_TOOL_PATH")
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]],
|
||||
) -> ToolMessage | Command[Any]:
|
||||
proxy, metadata = _context()
|
||||
if proxy is None:
|
||||
return await handler(request)
|
||||
tool_call = dict(request.tool_call)
|
||||
tool_name = str(tool_call.get("name") or "unknown_tool")
|
||||
tool_call_id = str(tool_call.get("id") or "")
|
||||
arguments = tool_call.get("args") or {}
|
||||
request_hash = _hash(arguments)
|
||||
checkpoint_ns = str(metadata.get("checkpoint_ns") or metadata.get("langgraph_checkpoint_ns") or "")
|
||||
task_path = ":".join(
|
||||
str(metadata.get(name) or "")
|
||||
for name in ("langgraph_step", "langgraph_node", "langgraph_task_idx")
|
||||
)
|
||||
effect_id = _hash(
|
||||
{
|
||||
"run_id": proxy["run_id"],
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"task_path": task_path,
|
||||
"tool_call_id": tool_call_id,
|
||||
"tool_name": tool_name,
|
||||
"request_hash": request_hash,
|
||||
}
|
||||
)
|
||||
prepared = await _post(
|
||||
proxy,
|
||||
"prepare",
|
||||
{
|
||||
"effect_id": effect_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"task_path": task_path,
|
||||
"tool_call_id": tool_call_id,
|
||||
"tool_name": tool_name,
|
||||
"effect_class": _effect_class(tool_name),
|
||||
"request_hash": request_hash,
|
||||
},
|
||||
)
|
||||
if prepared.get("action") == "manual_reconcile":
|
||||
return ToolMessage(
|
||||
content=(
|
||||
f"Tool '{tool_name}' may already have produced an external side effect. "
|
||||
"Automatic retry is blocked; user confirmation is required."
|
||||
),
|
||||
tool_call_id=tool_call_id,
|
||||
name=tool_name,
|
||||
status="error",
|
||||
)
|
||||
if prepared.get("action") == "cached":
|
||||
result = prepared.get("result")
|
||||
if isinstance(result, dict) and isinstance(result.get("message"), dict):
|
||||
messages = messages_from_dict([result["message"]])
|
||||
if len(messages) == 1 and isinstance(messages[0], ToolMessage):
|
||||
return messages[0]
|
||||
return ToolMessage(
|
||||
content="Cached tool result is unavailable; user confirmation is required.",
|
||||
tool_call_id=tool_call_id,
|
||||
name=tool_name,
|
||||
status="error",
|
||||
)
|
||||
fencing_token = int(prepared["fencing_token"])
|
||||
try:
|
||||
result = await handler(request)
|
||||
except BaseException:
|
||||
await _post(
|
||||
proxy,
|
||||
"terminal",
|
||||
{
|
||||
"effect_id": effect_id,
|
||||
"fencing_token": fencing_token,
|
||||
"outcome": "failed",
|
||||
"result": {},
|
||||
},
|
||||
)
|
||||
raise
|
||||
successful = isinstance(result, ToolMessage) and result.status != "error"
|
||||
payload = (
|
||||
{"message": message_to_dict(result)}
|
||||
if isinstance(result, ToolMessage)
|
||||
else {"command_result": True}
|
||||
)
|
||||
await _post(
|
||||
proxy,
|
||||
"terminal",
|
||||
{
|
||||
"effect_id": effect_id,
|
||||
"fencing_token": fencing_token,
|
||||
"outcome": "succeeded" if successful else "failed",
|
||||
"result": payload,
|
||||
},
|
||||
)
|
||||
return result
|
||||
@@ -0,0 +1,161 @@
|
||||
"""Bounded skill discovery for model prompts.
|
||||
|
||||
DeepAgents' stock ``SkillsMiddleware`` keeps the full skill catalog in agent
|
||||
state and renders every description into each model request. That is suitable
|
||||
for a small catalog but makes large global catalogs consume the whole context
|
||||
window. This subclass preserves catalog loading and file access while exposing
|
||||
only a relevant, byte-bounded subset in the prompt.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from collections.abc import Iterable, Sequence
|
||||
from typing import Any
|
||||
|
||||
from deepagents.middleware._utils import append_to_system_message
|
||||
from deepagents.middleware.skills import SkillsMiddleware
|
||||
from langchain.agents.middleware.types import ModelRequest
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
DEFAULT_MAX_SKILLS = 16
|
||||
DEFAULT_MAX_SKILLS_BYTES = 12 * 1024
|
||||
DEFAULT_MAX_DESCRIPTION_BYTES = 320
|
||||
|
||||
_WORD_RE = re.compile(r"[a-z0-9][a-z0-9_-]{1,}", re.IGNORECASE)
|
||||
_CJK_RUN_RE = re.compile(r"[\u4e00-\u9fff]{2,}")
|
||||
|
||||
|
||||
def _truncate_utf8(value: str, limit: int) -> str:
|
||||
"""Truncate text on a UTF-8 boundary, reserving room for an ellipsis."""
|
||||
|
||||
encoded = value.encode("utf-8")
|
||||
if len(encoded) <= limit:
|
||||
return value
|
||||
if limit <= 3:
|
||||
return ""
|
||||
return encoded[: limit - 3].decode("utf-8", errors="ignore") + "..."
|
||||
|
||||
|
||||
def _query_terms(value: str) -> set[str]:
|
||||
"""Extract ASCII words and CJK n-grams without a tokenizer dependency."""
|
||||
|
||||
normalized = value.lower()
|
||||
terms = set(_WORD_RE.findall(normalized))
|
||||
for run in _CJK_RUN_RE.findall(normalized):
|
||||
for width in range(2, min(4, len(run)) + 1):
|
||||
terms.update(
|
||||
run[index : index + width] for index in range(len(run) - width + 1)
|
||||
)
|
||||
return terms
|
||||
|
||||
|
||||
def _message_text(messages: Iterable[Any]) -> str:
|
||||
"""Return recent user text only; tool output must not drive skill ranking."""
|
||||
|
||||
for message in reversed(list(messages)):
|
||||
if not isinstance(message, HumanMessage):
|
||||
continue
|
||||
content = message.content
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, Sequence):
|
||||
return " ".join(str(item) for item in content)
|
||||
return ""
|
||||
|
||||
|
||||
class BudgetedSkillsMiddleware(SkillsMiddleware):
|
||||
"""Expose only relevant skills within a fixed system-prompt byte budget."""
|
||||
|
||||
name = "budgeted_skills"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
backend: Any,
|
||||
sources: Sequence[str | tuple[str, str]] | str,
|
||||
max_skills: int = DEFAULT_MAX_SKILLS,
|
||||
max_skills_bytes: int = DEFAULT_MAX_SKILLS_BYTES,
|
||||
max_description_bytes: int = DEFAULT_MAX_DESCRIPTION_BYTES,
|
||||
) -> None:
|
||||
resolved_sources = [sources] if isinstance(sources, str) else list(sources)
|
||||
super().__init__(backend=backend, sources=resolved_sources)
|
||||
if max_skills < 1 or max_skills_bytes < 1 or max_description_bytes < 1:
|
||||
raise ValueError("skill context limits must be positive")
|
||||
self.max_skills = max_skills
|
||||
self.max_skills_bytes = max_skills_bytes
|
||||
self.max_description_bytes = max_description_bytes
|
||||
|
||||
def _select_skills(
|
||||
self, skills: Sequence[dict[str, Any]], query: str
|
||||
) -> list[dict[str, Any]]:
|
||||
terms = _query_terms(query)
|
||||
if not terms:
|
||||
return []
|
||||
|
||||
scored: list[tuple[int, str, dict[str, Any]]] = []
|
||||
for skill in skills:
|
||||
name = str(skill.get("name") or "")
|
||||
description = str(skill.get("description") or "")
|
||||
name_text = name.lower()
|
||||
description_text = description.lower()
|
||||
score = 0
|
||||
for term in terms:
|
||||
if term == name_text:
|
||||
score += 100
|
||||
elif term in name_text:
|
||||
score += 24
|
||||
if term in description_text:
|
||||
score += 4
|
||||
if score:
|
||||
scored.append((score, name, skill))
|
||||
|
||||
scored.sort(key=lambda item: (-item[0], item[1]))
|
||||
return [skill for _, _, skill in scored[: self.max_skills]]
|
||||
|
||||
def _format_budgeted_skills(self, skills: Sequence[dict[str, Any]]) -> str:
|
||||
lines: list[str] = []
|
||||
used = 0
|
||||
for skill in skills:
|
||||
name = str(skill.get("name") or "unnamed")
|
||||
path = str(skill.get("path") or "")
|
||||
description = _truncate_utf8(
|
||||
str(skill.get("description") or ""), self.max_description_bytes
|
||||
)
|
||||
item = (
|
||||
f"- **{name}**: {description}\n -> Read `{path}` for full instructions"
|
||||
)
|
||||
item_bytes = len(item.encode("utf-8"))
|
||||
separator = 1 if lines else 0
|
||||
if used + separator + item_bytes > self.max_skills_bytes:
|
||||
continue
|
||||
lines.append(item)
|
||||
used += separator + item_bytes
|
||||
return "\n".join(lines)
|
||||
|
||||
def modify_request(self, request: ModelRequest) -> ModelRequest:
|
||||
if self.system_prompt_template is None:
|
||||
return request
|
||||
|
||||
state = request.state or {}
|
||||
skills = state.get("skills_metadata", [])
|
||||
selected = self._select_skills(skills, _message_text(request.messages))
|
||||
skills_list = self._format_budgeted_skills(selected)
|
||||
if len(selected) < len(skills):
|
||||
discovery_note = (
|
||||
"\n\nThis is a query-relevant subset of the installed skills. "
|
||||
"Use skill_manager(action='list') to discover the full catalog."
|
||||
)
|
||||
skills_list += discovery_note
|
||||
skills_section = self.system_prompt_template.format(
|
||||
skills_locations=self._format_skills_locations(),
|
||||
skills_load_warnings=self._format_skills_load_warnings(
|
||||
state.get("skills_load_errors", [])
|
||||
),
|
||||
skills_list=skills_list,
|
||||
)
|
||||
return request.override(
|
||||
system_message=append_to_system_message(
|
||||
request.system_message, skills_section
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,464 @@
|
||||
"""Normalize provider tool-call shapes into LangChain's canonical call form.
|
||||
|
||||
Provider adapters are allowed to disagree about mechanical fields such as
|
||||
``tool_call.id`` and whether arguments are delivered as a JSON object or a
|
||||
JSON string. The agent execution layer is not. This module is deliberately
|
||||
limited to deterministic protocol translation: it never infers a tool name,
|
||||
repairs incomplete JSON, or extracts calls from ordinary model text.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import ExtendedModelResponse, ModelResponse
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call", "tool_use"})
|
||||
_ID_SECRET_ENVS = (
|
||||
"AI4SCI_EVO_TOOL_CALL_ID_SECRET",
|
||||
"EVOSCI_TOOL_CALL_ID_SECRET",
|
||||
"AI4SCI_EVO_RUNTIME_GRANT_SECRET",
|
||||
)
|
||||
_DEFAULT_ID_SECRET = b"evoscientist-tool-call-normalizer-v1"
|
||||
_RAW_PROVIDER_CALL_FIELDS = (
|
||||
"tool_calls",
|
||||
"function_call",
|
||||
"functionCall",
|
||||
"tool_use",
|
||||
)
|
||||
|
||||
|
||||
class ToolCallNormalizationError(ValueError):
|
||||
"""A provider response could not be translated without guessing."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reason: str,
|
||||
*,
|
||||
call_index: int,
|
||||
source: str,
|
||||
) -> None:
|
||||
super().__init__(reason)
|
||||
self.reason = reason
|
||||
self.call_index = call_index
|
||||
self.source = source
|
||||
|
||||
def diagnostic(self) -> dict[str, Any]:
|
||||
return {
|
||||
"source": self.source,
|
||||
"call_index": self.call_index,
|
||||
"normalization_failure": self.reason,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CanonicalToolCall:
|
||||
"""Provider-neutral tool-call representation accepted by the ToolNode."""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
arguments: Any
|
||||
index: int
|
||||
source: str
|
||||
id_origin: str
|
||||
|
||||
def as_langchain_call(self) -> dict[str, Any]:
|
||||
args = dict(self.arguments) if isinstance(self.arguments, Mapping) else self.arguments
|
||||
return {"id": self.id, "name": self.name, "args": args, "type": "tool_call"}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _DecodedToolCall:
|
||||
call_id: str
|
||||
name: str
|
||||
arguments: Any
|
||||
arguments_present: bool
|
||||
index: int
|
||||
source: str
|
||||
|
||||
|
||||
def _text(value: Any) -> str:
|
||||
return value.strip() if isinstance(value, str) else ""
|
||||
|
||||
|
||||
def _call_fields(call: Mapping[str, Any]) -> tuple[str, str, Any, bool]:
|
||||
"""Read the common OpenAI, Anthropic and Gemini call field variants."""
|
||||
function = call.get("function")
|
||||
function = function if isinstance(function, Mapping) else {}
|
||||
# Responses function-call items expose two identifiers: ``id`` identifies
|
||||
# the output item (for example ``fc_*``), while ``call_id`` links the
|
||||
# eventual tool result (for example ``call_*``). Chat Completions only
|
||||
# has ``id``. Prefer the transport-level call identifier when present so
|
||||
# the parsed LangChain call and its content block describe the same call.
|
||||
call_id = _text(call.get("call_id") or call.get("id"))
|
||||
name = _text(call.get("name") or call.get("tool_name") or function.get("name"))
|
||||
if "args" in call:
|
||||
return call_id, name, call.get("args"), True
|
||||
if "arguments" in call:
|
||||
return call_id, name, call.get("arguments"), True
|
||||
if "input" in call:
|
||||
return call_id, name, call.get("input"), True
|
||||
if "arguments" in function:
|
||||
return call_id, name, function.get("arguments"), True
|
||||
return call_id, name, None, False
|
||||
|
||||
|
||||
def _tool_blocks(message: AIMessage) -> list[Mapping[str, Any]]:
|
||||
content = getattr(message, "content", None)
|
||||
if not isinstance(content, list):
|
||||
return []
|
||||
return [
|
||||
block
|
||||
for block in content
|
||||
if isinstance(block, Mapping) and block.get("type") in _TOOL_BLOCK_TYPES
|
||||
]
|
||||
|
||||
|
||||
class AdapterToolCallDecoder:
|
||||
"""Decode provider-shaped calls without exposing provider data downstream.
|
||||
|
||||
LangChain has already decoded most Provider wire formats to ``AIMessage``.
|
||||
This adapter deliberately accepts those canonical calls plus the three
|
||||
remaining lossless sources: OpenAI-compatible raw calls, Anthropic/Gemini
|
||||
content blocks and old ``function_call`` fields. Adapter-specific
|
||||
decoders can replace this class later without changing the normalizer or
|
||||
guard contract.
|
||||
"""
|
||||
|
||||
def __init__(self, adapter_id: str | None = None) -> None:
|
||||
self.adapter_id = adapter_id or "generic"
|
||||
|
||||
def decode(self, message: AIMessage) -> list[_DecodedToolCall]:
|
||||
sources = self._sources(message)
|
||||
if not sources:
|
||||
return []
|
||||
primary_name, primary_calls = sources[0]
|
||||
decoded = [
|
||||
self._decode_one(call, index=index, source=primary_name)
|
||||
for index, call in enumerate(primary_calls)
|
||||
]
|
||||
for source_name, source_calls in sources[1:]:
|
||||
# A second source can only be used as field-level evidence when it
|
||||
# preserves the same call ordering. Anything else is ambiguous.
|
||||
if len(source_calls) != len(decoded):
|
||||
raise ToolCallNormalizationError(
|
||||
"inconsistent_source_count",
|
||||
call_index=0,
|
||||
source=source_name,
|
||||
)
|
||||
decoded = [
|
||||
self._merge(
|
||||
primary,
|
||||
self._decode_one(raw, index=index, source=source_name),
|
||||
)
|
||||
for index, (primary, raw) in enumerate(
|
||||
zip(decoded, source_calls, strict=True)
|
||||
)
|
||||
]
|
||||
return decoded
|
||||
|
||||
@staticmethod
|
||||
def _sources(message: AIMessage) -> list[tuple[str, list[Mapping[str, Any]]]]:
|
||||
sources: list[tuple[str, list[Mapping[str, Any]]]] = []
|
||||
|
||||
def add_source(source: str, value: Any) -> None:
|
||||
if value is None:
|
||||
return
|
||||
calls = list(value) if isinstance(value, list | tuple) else [value]
|
||||
if not calls:
|
||||
return
|
||||
if any(not isinstance(call, Mapping) for call in calls):
|
||||
raise ToolCallNormalizationError(
|
||||
"invalid_call_shape", call_index=0, source=source
|
||||
)
|
||||
sources.append((source, calls))
|
||||
|
||||
add_source("parsed_tool_calls", getattr(message, "tool_calls", None) or [])
|
||||
|
||||
additional = getattr(message, "additional_kwargs", None)
|
||||
additional = additional if isinstance(additional, Mapping) else {}
|
||||
add_source("provider_raw_tool_calls", additional.get("tool_calls"))
|
||||
for legacy_field in _RAW_PROVIDER_CALL_FIELDS[1:]:
|
||||
add_source(f"provider_{legacy_field}", additional.get(legacy_field))
|
||||
|
||||
blocks = _tool_blocks(message)
|
||||
if blocks:
|
||||
sources.append(("content_blocks", blocks))
|
||||
return sources
|
||||
|
||||
@staticmethod
|
||||
def _decode_one(
|
||||
call: Mapping[str, Any], *, index: int, source: str
|
||||
) -> _DecodedToolCall:
|
||||
call_id, name, arguments, arguments_present = _call_fields(call)
|
||||
return _DecodedToolCall(
|
||||
call_id=call_id,
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
arguments_present=arguments_present,
|
||||
index=index,
|
||||
source=source,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _merge(
|
||||
primary: _DecodedToolCall, evidence: _DecodedToolCall
|
||||
) -> _DecodedToolCall:
|
||||
def choose(field: str, first: Any, second: Any) -> Any:
|
||||
first_present = bool(first) if field in {"call_id", "name"} else first is not None
|
||||
second_present = bool(second) if field in {"call_id", "name"} else second is not None
|
||||
if first_present and second_present and first != second:
|
||||
raise ToolCallNormalizationError(
|
||||
"inconsistent_source",
|
||||
call_index=primary.index,
|
||||
source=evidence.source,
|
||||
)
|
||||
return first if first_present else second
|
||||
|
||||
call_id = choose("call_id", primary.call_id, evidence.call_id)
|
||||
name = choose("name", primary.name, evidence.name)
|
||||
if primary.arguments_present and evidence.arguments_present:
|
||||
if _parse_json_object(primary.arguments) != _parse_json_object(
|
||||
evidence.arguments
|
||||
):
|
||||
raise ToolCallNormalizationError(
|
||||
"inconsistent_source",
|
||||
call_index=primary.index,
|
||||
source=evidence.source,
|
||||
)
|
||||
arguments = primary.arguments
|
||||
elif primary.arguments_present:
|
||||
arguments = primary.arguments
|
||||
else:
|
||||
arguments = evidence.arguments
|
||||
return _DecodedToolCall(
|
||||
call_id=call_id,
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
arguments_present=(
|
||||
primary.arguments_present or evidence.arguments_present
|
||||
),
|
||||
index=primary.index,
|
||||
source=primary.source,
|
||||
)
|
||||
|
||||
|
||||
def _parse_json_object(value: Any) -> Any:
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
try:
|
||||
return json.loads(value)
|
||||
except (TypeError, ValueError, json.JSONDecodeError):
|
||||
return value
|
||||
|
||||
|
||||
class ToolCallNormalizer:
|
||||
"""Create canonical calls and immutable normalized AI messages."""
|
||||
|
||||
version = "v1"
|
||||
|
||||
def __init__(self, *, id_secret: bytes | None = None) -> None:
|
||||
env_secret = next(
|
||||
(os.environ[name] for name in _ID_SECRET_ENVS if os.environ.get(name)),
|
||||
None,
|
||||
)
|
||||
self._id_secret = id_secret or (
|
||||
env_secret.encode("utf-8") if env_secret else _DEFAULT_ID_SECRET
|
||||
)
|
||||
|
||||
def normalize_message(
|
||||
self,
|
||||
message: AIMessage,
|
||||
*,
|
||||
adapter_id: str | None,
|
||||
request_scope: str,
|
||||
) -> AIMessage:
|
||||
decoded = AdapterToolCallDecoder(adapter_id).decode(message)
|
||||
if not decoded:
|
||||
return message
|
||||
|
||||
canonical_calls = [
|
||||
self._canonicalize(call, request_scope=request_scope) for call in decoded
|
||||
]
|
||||
canonical_dicts = [call.as_langchain_call() for call in canonical_calls]
|
||||
if self._already_canonical(message, canonical_dicts):
|
||||
return message
|
||||
return self._copy_message(message, canonical_dicts)
|
||||
|
||||
def _canonicalize(
|
||||
self, call: _DecodedToolCall, *, request_scope: str
|
||||
) -> CanonicalToolCall:
|
||||
arguments = self._arguments(call)
|
||||
call_id = call.call_id
|
||||
id_origin = "provider"
|
||||
if not call_id and call.name and _is_json_object(arguments):
|
||||
call_id = self._generated_id(
|
||||
request_scope=request_scope,
|
||||
call=call,
|
||||
arguments=arguments,
|
||||
)
|
||||
id_origin = "gateway_generated"
|
||||
return CanonicalToolCall(
|
||||
id=call_id,
|
||||
name=call.name,
|
||||
arguments=arguments,
|
||||
index=call.index,
|
||||
source=call.source,
|
||||
id_origin=id_origin,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _arguments(call: _DecodedToolCall) -> Any:
|
||||
if not call.arguments_present:
|
||||
return None
|
||||
return _parse_json_object(call.arguments)
|
||||
|
||||
def _generated_id(
|
||||
self,
|
||||
*,
|
||||
request_scope: str,
|
||||
call: _DecodedToolCall,
|
||||
arguments: Mapping[str, Any],
|
||||
) -> str:
|
||||
canonical_args = json.dumps(
|
||||
arguments,
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
material = "\x1f".join(
|
||||
(self.version, request_scope, str(call.index), call.name, canonical_args)
|
||||
).encode("utf-8")
|
||||
return "call_" + hmac.new(
|
||||
self._id_secret, material, hashlib.sha256
|
||||
).hexdigest()[:32]
|
||||
|
||||
@staticmethod
|
||||
def _already_canonical(
|
||||
message: AIMessage, canonical_calls: list[dict[str, Any]]
|
||||
) -> bool:
|
||||
current_calls = list(getattr(message, "tool_calls", None) or [])
|
||||
if current_calls != canonical_calls:
|
||||
return False
|
||||
if getattr(message, "invalid_tool_calls", None):
|
||||
return False
|
||||
additional = getattr(message, "additional_kwargs", None) or {}
|
||||
if any(additional.get(key) for key in _RAW_PROVIDER_CALL_FIELDS):
|
||||
return False
|
||||
blocks = _tool_blocks(message)
|
||||
if not blocks:
|
||||
return True
|
||||
if len(blocks) != len(canonical_calls):
|
||||
return False
|
||||
return all(
|
||||
_text(block.get("call_id") or block.get("id")) == call["id"]
|
||||
and _text(
|
||||
block.get("name")
|
||||
or block.get("tool_name")
|
||||
or (
|
||||
block.get("function", {}).get("name")
|
||||
if isinstance(block.get("function"), Mapping)
|
||||
else ""
|
||||
)
|
||||
)
|
||||
== call["name"]
|
||||
for block, call in zip(blocks, canonical_calls, strict=True)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _copy_message(
|
||||
message: AIMessage, canonical_calls: list[dict[str, Any]]
|
||||
) -> AIMessage:
|
||||
copied = copy.copy(message)
|
||||
copied.tool_calls = canonical_calls
|
||||
# A repaired call cannot make a second malformed call safe. Preserve
|
||||
# invalid calls so the Guard still rejects the entire response.
|
||||
copied.invalid_tool_calls = list(
|
||||
getattr(message, "invalid_tool_calls", None) or []
|
||||
)
|
||||
additional = dict(getattr(message, "additional_kwargs", None) or {})
|
||||
# Parsed calls are now canonical; retaining a raw provider copy can
|
||||
# reintroduce the missing ID when history is replayed.
|
||||
for field in _RAW_PROVIDER_CALL_FIELDS:
|
||||
additional.pop(field, None)
|
||||
copied.additional_kwargs = additional
|
||||
|
||||
content = getattr(message, "content", None)
|
||||
if isinstance(content, list):
|
||||
call_index = 0
|
||||
normalized_content: list[Any] = []
|
||||
for value in content:
|
||||
if not isinstance(value, Mapping) or value.get("type") not in _TOOL_BLOCK_TYPES:
|
||||
normalized_content.append(value)
|
||||
continue
|
||||
if call_index >= len(canonical_calls):
|
||||
normalized_content.append(value)
|
||||
continue
|
||||
call = canonical_calls[call_index]
|
||||
block = dict(value)
|
||||
block["id"] = call["id"]
|
||||
block["name"] = call["name"]
|
||||
if isinstance(block.get("function"), Mapping):
|
||||
block["function"] = {
|
||||
**block["function"],
|
||||
"name": call["name"],
|
||||
}
|
||||
normalized_content.append(block)
|
||||
call_index += 1
|
||||
copied.content = normalized_content
|
||||
return copied
|
||||
|
||||
def normalize_response(
|
||||
self,
|
||||
response: Any,
|
||||
*,
|
||||
adapter_id: str | None,
|
||||
request_scope: str,
|
||||
) -> Any:
|
||||
if isinstance(response, AIMessage):
|
||||
return self.normalize_message(
|
||||
response, adapter_id=adapter_id, request_scope=request_scope
|
||||
)
|
||||
if isinstance(response, ExtendedModelResponse):
|
||||
normalized = self.normalize_response(
|
||||
response.model_response,
|
||||
adapter_id=adapter_id,
|
||||
request_scope=request_scope,
|
||||
)
|
||||
return (
|
||||
response
|
||||
if normalized is response.model_response
|
||||
else replace(response, model_response=normalized)
|
||||
)
|
||||
if isinstance(response, ModelResponse):
|
||||
result = list(response.result)
|
||||
normalized_result = [
|
||||
self.normalize_message(
|
||||
message,
|
||||
adapter_id=adapter_id,
|
||||
request_scope=request_scope,
|
||||
)
|
||||
if isinstance(message, AIMessage)
|
||||
else message
|
||||
for message in result
|
||||
]
|
||||
return response if normalized_result == result else replace(response, result=normalized_result)
|
||||
return response
|
||||
|
||||
|
||||
def _is_json_object(value: Any) -> bool:
|
||||
if not isinstance(value, Mapping):
|
||||
return False
|
||||
try:
|
||||
json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
return True
|
||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
@@ -17,8 +18,11 @@ from langchain_core.messages import AIMessage
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
from ..llm.errors import ModelToolProtocolError, _provider_from_model
|
||||
from .tool_call_normalizer import ToolCallNormalizationError, ToolCallNormalizer
|
||||
|
||||
_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call"})
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call", "tool_use"})
|
||||
_MAX_DIAGNOSTIC_KEYS = 16
|
||||
_MAX_DIAGNOSTIC_KEY_CHARS = 64
|
||||
|
||||
@@ -54,7 +58,10 @@ def _ai_messages(response: Any) -> list[AIMessage]:
|
||||
|
||||
|
||||
def _block_identity(block: Mapping[str, Any]) -> tuple[str, str]:
|
||||
call_id = str(block.get("id") or block.get("call_id") or "").strip()
|
||||
# Responses output items use ``id`` for the item and ``call_id`` for the
|
||||
# tool-result correlation key. The latter is the identity ToolMessage
|
||||
# must carry; Chat Completions continues to fall back to ``id``.
|
||||
call_id = str(block.get("call_id") or block.get("id") or "").strip()
|
||||
name = block.get("name") or block.get("tool_name")
|
||||
function = block.get("function")
|
||||
if not name and isinstance(function, Mapping):
|
||||
@@ -104,12 +111,22 @@ def _argument_diagnostic(value: Any, *, present: bool) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _is_json_object(value: Any) -> bool:
|
||||
if not isinstance(value, Mapping):
|
||||
return False
|
||||
try:
|
||||
json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _summarize_call(call: Any) -> dict[str, Any]:
|
||||
if not isinstance(call, Mapping):
|
||||
return {"call_type": type(call).__name__}
|
||||
function = call.get("function")
|
||||
function = function if isinstance(function, Mapping) else {}
|
||||
call_id = str(call.get("id") or call.get("call_id") or "").strip()
|
||||
call_id = str(call.get("call_id") or call.get("id") or "").strip()
|
||||
name = call.get("name") or call.get("tool_name") or function.get("name")
|
||||
name = str(name or "").strip()
|
||||
if "args" in call:
|
||||
@@ -205,6 +222,28 @@ def _route_metadata(request: ModelRequest) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _normalization_scope(request: ModelRequest) -> str:
|
||||
"""Return non-sensitive per-turn material for generated call identifiers."""
|
||||
route = _route_metadata(request)
|
||||
runtime = getattr(request, "runtime", None)
|
||||
config = getattr(runtime, "config", None)
|
||||
config = config if isinstance(config, Mapping) else {}
|
||||
configurable = config.get("configurable")
|
||||
configurable = configurable if isinstance(configurable, Mapping) else {}
|
||||
thread_id = str(configurable.get("thread_id") or "")
|
||||
messages = getattr(request, "messages", None)
|
||||
message_count = len(messages) if isinstance(messages, Sequence) else 0
|
||||
return "|".join(
|
||||
(
|
||||
str(route.get("route_key") or ""),
|
||||
str(route.get("config_generation") or ""),
|
||||
str(route.get("api_mode") or ""),
|
||||
thread_id,
|
||||
str(message_count),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _raise_protocol_error(
|
||||
request: ModelRequest,
|
||||
reason: str,
|
||||
@@ -212,11 +251,24 @@ def _raise_protocol_error(
|
||||
call_id: str | None = None,
|
||||
call_diagnostic: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
route = _route_metadata(request)
|
||||
logger.warning(
|
||||
"model_tool_protocol_invalid reason=%s provider=%s model=%s route_key=%s "
|
||||
"api_mode=%s transport=%s call_id_present=%s diagnostic=%s",
|
||||
reason,
|
||||
route["provider"],
|
||||
route["model"],
|
||||
route["route_key"],
|
||||
route["api_mode"],
|
||||
route["tool_call_transport"],
|
||||
bool(call_id),
|
||||
json.dumps(call_diagnostic or {}, ensure_ascii=True, sort_keys=True),
|
||||
)
|
||||
raise ModelToolProtocolError(
|
||||
reason,
|
||||
call_id=call_id or None,
|
||||
call_diagnostic=call_diagnostic,
|
||||
**_route_metadata(request),
|
||||
**route,
|
||||
)
|
||||
|
||||
|
||||
@@ -282,7 +334,7 @@ def _validate_message(
|
||||
call_diagnostic=diagnostic,
|
||||
)
|
||||
args = raw_call.get("args")
|
||||
if not isinstance(args, Mapping):
|
||||
if not _is_json_object(args):
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
"invalid_args",
|
||||
@@ -330,10 +382,30 @@ def _validate_message(
|
||||
|
||||
|
||||
class ToolProtocolGuardMiddleware(AgentMiddleware):
|
||||
"""Fail closed on malformed final tool calls using the actual request tools."""
|
||||
"""Normalize then fail closed on malformed final tool calls."""
|
||||
|
||||
name = "tool_protocol_guard"
|
||||
|
||||
def __init__(self, *, normalizer: ToolCallNormalizer | None = None) -> None:
|
||||
super().__init__()
|
||||
self._normalizer = normalizer or ToolCallNormalizer()
|
||||
|
||||
def _normalize(self, response: Any, request: ModelRequest) -> Any:
|
||||
metadata = getattr(request.model, "metadata", None)
|
||||
metadata = metadata if isinstance(metadata, Mapping) else {}
|
||||
try:
|
||||
return self._normalizer.normalize_response(
|
||||
response,
|
||||
adapter_id=str(metadata.get("route_adapter_id") or "") or None,
|
||||
request_scope=_normalization_scope(request),
|
||||
)
|
||||
except ToolCallNormalizationError as error:
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
error.reason,
|
||||
call_diagnostic=error.diagnostic(),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate(response: Any, request: ModelRequest) -> None:
|
||||
allowed_names = frozenset(
|
||||
@@ -348,6 +420,7 @@ class ToolProtocolGuardMiddleware(AgentMiddleware):
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
response = handler(request)
|
||||
response = self._normalize(response, request)
|
||||
self._validate(response, request)
|
||||
return response
|
||||
|
||||
@@ -357,5 +430,6 @@ class ToolProtocolGuardMiddleware(AgentMiddleware):
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
response = await handler(request)
|
||||
response = self._normalize(response, request)
|
||||
self._validate(response, request)
|
||||
return response
|
||||
|
||||
@@ -16,7 +16,6 @@ from typing import Any
|
||||
AsyncProvider = Callable[[], Awaitable[Any]]
|
||||
AsyncFileHandler = Callable[[Path], Awaitable[Any]]
|
||||
AsyncUsageRecorder = Callable[[str, str], Awaitable[Any]]
|
||||
ModelResolver = Callable[[str | None, str | None], Any | None]
|
||||
|
||||
|
||||
class RuntimeIntegrationUnavailable(RuntimeError):
|
||||
@@ -33,7 +32,6 @@ class RuntimeIntegrations:
|
||||
knowledge_file_handler: AsyncFileHandler | None = None
|
||||
usage_recorder: AsyncUsageRecorder | None = None
|
||||
image_backend_factory: Callable[[], Any] | None = None
|
||||
model_resolver: ModelResolver | None = None
|
||||
|
||||
|
||||
_integrations = RuntimeIntegrations()
|
||||
@@ -89,12 +87,6 @@ def resolve_user_storage_root(user_id: str) -> Path | None:
|
||||
return provider(user_id) if provider is not None else None
|
||||
|
||||
|
||||
def resolve_runtime_model(model: str | None, provider: str | None = None) -> Any | None:
|
||||
"""Resolve a host-managed model configuration when one is registered."""
|
||||
resolver = _integrations.model_resolver
|
||||
return resolver(model, provider) if resolver is not None else None
|
||||
|
||||
|
||||
async def handle_knowledge_file(path: Path) -> None:
|
||||
handler = _integrations.knowledge_file_handler
|
||||
if handler is not None:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+248
-29
@@ -29,8 +29,10 @@ WebUI / langgraph-dev checkpointer:
|
||||
|
||||
import asyncio
|
||||
import atexit
|
||||
import hashlib
|
||||
import logging
|
||||
import math
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from contextlib import asynccontextmanager
|
||||
@@ -76,6 +78,19 @@ MAIN_THREAD_FILTER_SQL = (
|
||||
" OR json_extract(metadata, '$.graph_id') = ?)"
|
||||
)
|
||||
MAIN_THREAD_FILTER_PARAMS = (AGENT_NAME, AGENT_NAME)
|
||||
_CHECKPOINT_MSGPACK_MODULES = frozenset(
|
||||
{
|
||||
("EvoScientist.llm.errors", "AgentControlError"),
|
||||
("EvoScientist.llm.errors", "ModelToolProtocolError"),
|
||||
("EvoScientist.llm.errors", "ProviderStreamError"),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _checkpoint_serde() -> JsonPlusSerializer:
|
||||
"""Return the checkpoint serializer with app-owned types allowlisted."""
|
||||
|
||||
return JsonPlusSerializer(allowed_msgpack_modules=_CHECKPOINT_MSGPACK_MODULES)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -182,7 +197,9 @@ class PruningCheckpointer(AsyncSqliteSaver):
|
||||
keep_per_ns: int = _DEFAULT_KEEP_PER_NS,
|
||||
serde: Any = None,
|
||||
) -> None:
|
||||
super().__init__(conn, serde=serde)
|
||||
super().__init__(
|
||||
conn, serde=serde if serde is not None else _checkpoint_serde()
|
||||
)
|
||||
self._keep_per_ns = max(0, int(keep_per_ns))
|
||||
# Outer lock guarantees ``super().aput()`` and ``_prune_after_put()``
|
||||
# are atomic *as a pair*. Without this, a concurrent ``aput()`` on a
|
||||
@@ -321,12 +338,7 @@ class PruningCheckpointer(AsyncSqliteSaver):
|
||||
checkpoint_ns: str,
|
||||
oldest_anchor_id: str,
|
||||
) -> set[str]:
|
||||
"""Walk parent chain until hitting a ``messages`` seed.
|
||||
|
||||
Returns the set of ancestor ids to preserve (inclusive of the
|
||||
snapshot ancestor). On chain-break or deserialization failure,
|
||||
returns what was visited so far — the safe side is over-preserve.
|
||||
"""
|
||||
"""Walk the parent chain until a messages seed is found."""
|
||||
extra: set[str] = set()
|
||||
cursor = await self._fetch_parent_checkpoint_id(
|
||||
thread_id, checkpoint_ns, oldest_anchor_id
|
||||
@@ -336,7 +348,7 @@ class PruningCheckpointer(AsyncSqliteSaver):
|
||||
steps += 1
|
||||
blob = await self._fetch_checkpoint_blob(thread_id, checkpoint_ns, cursor)
|
||||
if blob is None:
|
||||
break # chain broken (legacy DB); preserve what we have
|
||||
break
|
||||
extra.add(cursor)
|
||||
try:
|
||||
ck = self.serde.loads_typed(blob)
|
||||
@@ -348,10 +360,10 @@ class PruningCheckpointer(AsyncSqliteSaver):
|
||||
thread_id,
|
||||
exc,
|
||||
)
|
||||
break # safe-side: preserve everything visited so far
|
||||
break
|
||||
cv = ck.get("channel_values") or {}
|
||||
if _unwrap_messages_seed(cv.get("messages")) is not None:
|
||||
break # found seed; this ancestor anchors reconstruction
|
||||
break
|
||||
cursor = await self._fetch_parent_checkpoint_id(
|
||||
thread_id, checkpoint_ns, cursor
|
||||
)
|
||||
@@ -392,25 +404,12 @@ class PruningCheckpointer(AsyncSqliteSaver):
|
||||
agent: str,
|
||||
kept_ids: set[str],
|
||||
) -> None:
|
||||
"""DELETE rows whose ``checkpoint_id`` is NOT in ``kept_ids``.
|
||||
|
||||
Writes deleted first to preserve referential ordering — if we
|
||||
dropped checkpoints first, surviving writes' ``checkpoint_id``
|
||||
would become orphans.
|
||||
|
||||
Empty ``kept_ids`` is a no-op rather than "delete everything" —
|
||||
a defensive check; the caller always passes anchor_ids which is
|
||||
non-empty by construction (already checked ``len >= keep`` in
|
||||
the caller).
|
||||
"""
|
||||
"""Delete checkpoints outside the retained head and seed chain."""
|
||||
if not kept_ids:
|
||||
return
|
||||
kept_list = list(kept_ids)
|
||||
placeholders = ",".join("?" * len(kept_list))
|
||||
|
||||
# Writes DELETE only runs if the ``writes`` table exists. Legacy
|
||||
# DBs from pre-DeltaChannel builds may have only ``checkpoints`` —
|
||||
# we still want to prune those, just skipping the writes step.
|
||||
if await _table_exists(self.conn, "writes"):
|
||||
del_writes = (
|
||||
"DELETE FROM writes "
|
||||
@@ -446,6 +445,226 @@ class PruningCheckpointer(AsyncSqliteSaver):
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TurnLease:
|
||||
thread_id: str
|
||||
owner_id: str
|
||||
fencing_token: int
|
||||
expires_at_ms: int
|
||||
checkpoint_snapshot_id: str
|
||||
|
||||
|
||||
class FencedPruningCheckpointer(PruningCheckpointer):
|
||||
"""Single-worker linearizable turn lease around LangGraph checkpoint writes."""
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self._turn_fence_lock = asyncio.Lock()
|
||||
|
||||
@classmethod
|
||||
@asynccontextmanager
|
||||
async def from_conn_string_with_keep(
|
||||
cls, conn_string: str, keep_per_ns: int = _DEFAULT_KEEP_PER_NS
|
||||
) -> AsyncIterator["FencedPruningCheckpointer"]:
|
||||
async with aiosqlite.connect(conn_string) as conn:
|
||||
saver = cls(conn, keep_per_ns=keep_per_ns)
|
||||
await saver.setup_fencing()
|
||||
yield saver
|
||||
|
||||
async def setup_fencing(self) -> None:
|
||||
await super().setup()
|
||||
async with self.lock:
|
||||
await self.conn.executescript(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS thread_turn_fences (
|
||||
thread_id TEXT PRIMARY KEY,
|
||||
generation INTEGER NOT NULL CHECK (generation > 0),
|
||||
owner_id TEXT,
|
||||
expires_at_ms INTEGER,
|
||||
updated_at_ms INTEGER NOT NULL
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS thread_checkpoint_versions (
|
||||
thread_id TEXT PRIMARY KEY,
|
||||
sequence INTEGER NOT NULL DEFAULT 0 CHECK (sequence >= 0),
|
||||
checkpoint_id TEXT,
|
||||
updated_at_ms INTEGER NOT NULL
|
||||
);
|
||||
"""
|
||||
)
|
||||
await self.conn.commit()
|
||||
|
||||
async def acquire_turn_lease(
|
||||
self,
|
||||
thread_id: str,
|
||||
owner_id: str,
|
||||
*,
|
||||
ttl_seconds: int,
|
||||
) -> TurnLease:
|
||||
clean_thread = str(thread_id or "").strip()
|
||||
clean_owner = str(owner_id or "").strip()
|
||||
if not clean_thread or not clean_owner or ttl_seconds < 1:
|
||||
raise ValueError("thread, owner and positive TTL are required")
|
||||
async with self._turn_fence_lock, self.lock:
|
||||
now = time.time_ns() // 1_000_000
|
||||
await self.conn.execute("BEGIN IMMEDIATE")
|
||||
try:
|
||||
row = await (
|
||||
await self.conn.execute(
|
||||
"SELECT generation, owner_id, expires_at_ms FROM thread_turn_fences WHERE thread_id=?",
|
||||
(clean_thread,),
|
||||
)
|
||||
).fetchone()
|
||||
if (
|
||||
row is not None
|
||||
and row[1]
|
||||
and row[1] != clean_owner
|
||||
and int(row[2] or 0) >= now
|
||||
):
|
||||
raise RuntimeError("TURN_LEASE_BUSY")
|
||||
generation = int(row[0]) + 1 if row is not None else 1
|
||||
expires_at = now + ttl_seconds * 1000
|
||||
await self.conn.execute(
|
||||
"""INSERT INTO thread_turn_fences
|
||||
(thread_id, generation, owner_id, expires_at_ms, updated_at_ms)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
ON CONFLICT(thread_id) DO UPDATE SET
|
||||
generation=excluded.generation,
|
||||
owner_id=excluded.owner_id,
|
||||
expires_at_ms=excluded.expires_at_ms,
|
||||
updated_at_ms=excluded.updated_at_ms""",
|
||||
(clean_thread, generation, clean_owner, expires_at, now),
|
||||
)
|
||||
version = await (
|
||||
await self.conn.execute(
|
||||
"SELECT sequence, checkpoint_id FROM thread_checkpoint_versions WHERE thread_id=?",
|
||||
(clean_thread,),
|
||||
)
|
||||
).fetchone()
|
||||
sequence = int(version[0]) if version is not None else 0
|
||||
checkpoint_id = (
|
||||
str(version[1] or "root") if version is not None else "root"
|
||||
)
|
||||
if version is None:
|
||||
await self.conn.execute(
|
||||
"""INSERT INTO thread_checkpoint_versions
|
||||
(thread_id, sequence, checkpoint_id, updated_at_ms)
|
||||
VALUES (?, 0, NULL, ?)""",
|
||||
(clean_thread, now),
|
||||
)
|
||||
await self.conn.commit()
|
||||
except Exception:
|
||||
await self.conn.rollback()
|
||||
raise
|
||||
snapshot = hashlib.sha256(
|
||||
f"{clean_thread}\0{sequence}\0{checkpoint_id}".encode()
|
||||
).hexdigest()
|
||||
return TurnLease(
|
||||
clean_thread, clean_owner, generation, expires_at, f"sha256:{snapshot}"
|
||||
)
|
||||
|
||||
async def renew_turn_lease(
|
||||
self, lease: TurnLease, *, ttl_seconds: int
|
||||
) -> TurnLease:
|
||||
async with self._turn_fence_lock, self.lock:
|
||||
now = time.time_ns() // 1_000_000
|
||||
expires_at = now + ttl_seconds * 1000
|
||||
cursor = await self.conn.execute(
|
||||
"""UPDATE thread_turn_fences
|
||||
SET expires_at_ms=?, updated_at_ms=?
|
||||
WHERE thread_id=? AND generation=? AND owner_id=?
|
||||
AND expires_at_ms>=?""",
|
||||
(
|
||||
expires_at,
|
||||
now,
|
||||
lease.thread_id,
|
||||
lease.fencing_token,
|
||||
lease.owner_id,
|
||||
now,
|
||||
),
|
||||
)
|
||||
await self.conn.commit()
|
||||
if cursor.rowcount != 1:
|
||||
raise RuntimeError("TURN_LEASE_LOST")
|
||||
return TurnLease(
|
||||
lease.thread_id,
|
||||
lease.owner_id,
|
||||
lease.fencing_token,
|
||||
expires_at,
|
||||
lease.checkpoint_snapshot_id,
|
||||
)
|
||||
|
||||
async def release_turn_lease(self, lease: TurnLease) -> bool:
|
||||
async with self._turn_fence_lock, self.lock:
|
||||
now = time.time_ns() // 1_000_000
|
||||
cursor = await self.conn.execute(
|
||||
"""UPDATE thread_turn_fences
|
||||
SET owner_id=NULL, expires_at_ms=NULL, updated_at_ms=?
|
||||
WHERE thread_id=? AND generation=? AND owner_id=?""",
|
||||
(now, lease.thread_id, lease.fencing_token, lease.owner_id),
|
||||
)
|
||||
await self.conn.commit()
|
||||
return cursor.rowcount == 1
|
||||
|
||||
async def _require_write_lease(self, config: Any) -> tuple[str, int]:
|
||||
configurable = dict(config.get("configurable") or {})
|
||||
thread_id = str(configurable.get("thread_id") or "")
|
||||
if not thread_id.startswith("web:"):
|
||||
return thread_id, 0
|
||||
owner_id = str(configurable.get("turn_lease_owner") or "")
|
||||
token = int(configurable.get("turn_fencing_token") or 0)
|
||||
now = time.time_ns() // 1_000_000
|
||||
row = await (
|
||||
await self.conn.execute(
|
||||
"""SELECT 1 FROM thread_turn_fences
|
||||
WHERE thread_id=? AND generation=? AND owner_id=?
|
||||
AND expires_at_ms>=?""",
|
||||
(thread_id, token, owner_id, now),
|
||||
)
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise RuntimeError("TURN_FENCED")
|
||||
return thread_id, token
|
||||
|
||||
async def aput(
|
||||
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
|
||||
) -> Any:
|
||||
async with self._turn_fence_lock:
|
||||
async with self.lock:
|
||||
thread_id, token = await self._require_write_lease(config)
|
||||
result = await super().aput(config, checkpoint, metadata, new_versions)
|
||||
if token == 0:
|
||||
return result
|
||||
checkpoint_id = str(
|
||||
result.get("configurable", {}).get("checkpoint_id") or ""
|
||||
)
|
||||
async with self.lock:
|
||||
now = time.time_ns() // 1_000_000
|
||||
await self.conn.execute(
|
||||
"""INSERT INTO thread_checkpoint_versions
|
||||
(thread_id, sequence, checkpoint_id, updated_at_ms)
|
||||
VALUES (?, 1, ?, ?)
|
||||
ON CONFLICT(thread_id) DO UPDATE SET
|
||||
sequence=thread_checkpoint_versions.sequence+1,
|
||||
checkpoint_id=excluded.checkpoint_id,
|
||||
updated_at_ms=excluded.updated_at_ms""",
|
||||
(thread_id, checkpoint_id, now),
|
||||
)
|
||||
await self.conn.commit()
|
||||
return result
|
||||
|
||||
async def aput_writes(
|
||||
self,
|
||||
config: Any,
|
||||
writes: Any,
|
||||
task_id: str,
|
||||
task_path: str = "",
|
||||
) -> None:
|
||||
async with self._turn_fence_lock:
|
||||
async with self.lock:
|
||||
await self._require_write_lease(config)
|
||||
await super().aput_writes(config, writes, task_id, task_path)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Checkpointer context manager
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -462,7 +681,7 @@ def _resolve_keep_per_ns() -> int:
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def get_checkpointer() -> AsyncIterator[PruningCheckpointer]:
|
||||
async def get_checkpointer() -> AsyncIterator[FencedPruningCheckpointer]:
|
||||
"""Yield a pruning-enabled checkpointer connected to the sessions DB.
|
||||
|
||||
Wraps ``AsyncSqliteSaver`` with ``PruningCheckpointer`` so every
|
||||
@@ -479,7 +698,7 @@ async def get_checkpointer() -> AsyncIterator[PruningCheckpointer]:
|
||||
On failure ``user_version`` is NOT bumped, so the next launch retries.
|
||||
"""
|
||||
keep = _resolve_keep_per_ns()
|
||||
async with PruningCheckpointer.from_conn_string_with_keep(
|
||||
async with FencedPruningCheckpointer.from_conn_string_with_keep(
|
||||
str(get_db_path()), keep_per_ns=keep
|
||||
) as saver:
|
||||
# The whole gate is wrapped in a broad try/except: any unexpected
|
||||
@@ -919,7 +1138,7 @@ async def list_threads(
|
||||
|
||||
if (include_message_count or include_preview) and threads:
|
||||
# Share one saver across all threads so ``setup()`` runs once.
|
||||
serde = JsonPlusSerializer()
|
||||
serde = _checkpoint_serde()
|
||||
saver = AsyncSqliteSaver(conn, serde=serde)
|
||||
for t in threads:
|
||||
msgs = await _load_checkpoint_messages(saver, t["thread_id"])
|
||||
@@ -1086,7 +1305,7 @@ async def get_thread_messages(thread_id: str) -> list:
|
||||
async with conn.execute(check, (thread_id, *MAIN_THREAD_FILTER_PARAMS)) as cur:
|
||||
if not await cur.fetchone():
|
||||
return []
|
||||
serde = JsonPlusSerializer()
|
||||
serde = _checkpoint_serde()
|
||||
saver = AsyncSqliteSaver(conn, serde=serde)
|
||||
return await _load_checkpoint_messages(saver, thread_id)
|
||||
|
||||
@@ -1646,7 +1865,7 @@ async def _restore_webui_threads_to_global_store() -> None:
|
||||
# message (stubs carry values=None, so the WebUI would otherwise
|
||||
# render every restored thread as "Untitled Thread").
|
||||
if sqlite_data:
|
||||
saver = AsyncSqliteSaver(conn, serde=JsonPlusSerializer())
|
||||
saver = AsyncSqliteSaver(conn, serde=_checkpoint_serde())
|
||||
for thread_uuid in sqlite_data:
|
||||
try:
|
||||
msgs = await _load_checkpoint_messages(saver, str(thread_uuid))
|
||||
|
||||
@@ -541,7 +541,10 @@ class _V3EventProcessor:
|
||||
) -> str:
|
||||
matches = [
|
||||
call_id
|
||||
for (candidate_scope, call_id), candidate in self._pending_tool_calls.items()
|
||||
for (
|
||||
candidate_scope,
|
||||
call_id,
|
||||
), candidate in self._pending_tool_calls.items()
|
||||
if candidate_scope == scope and candidate == (name, args)
|
||||
]
|
||||
return matches[0] if len(matches) == 1 else ""
|
||||
@@ -849,7 +852,9 @@ class _V3EventProcessor:
|
||||
call_id = str(
|
||||
request_map.get("id") or request_map.get("tool_call_id") or ""
|
||||
)
|
||||
name = str(request_map.get("name") or request_map.get("tool_name") or "")
|
||||
name = str(
|
||||
request_map.get("name") or request_map.get("tool_name") or ""
|
||||
)
|
||||
args_map = _as_raw_map(
|
||||
request_map.get("args")
|
||||
if "args" in request_map
|
||||
@@ -1007,6 +1012,7 @@ async def stream_agent_events(
|
||||
metadata: dict[str, Any] | None = None,
|
||||
media: list[str] | None = None,
|
||||
callbacks: list[Any] | None = None,
|
||||
configurable: dict[str, Any] | None = None,
|
||||
error_mode: str = "emit",
|
||||
) -> AsyncGenerator[dict[str, Any], None]:
|
||||
"""Stream events from a DeepAgents/LangGraph v3 run.
|
||||
@@ -1032,7 +1038,9 @@ async def stream_agent_events(
|
||||
subagent_start, subagent_tool_call, subagent_tool_result, subagent_end,
|
||||
done, error
|
||||
"""
|
||||
config: dict[str, Any] = {"configurable": {"thread_id": thread_id}}
|
||||
config_values = dict(configurable or {})
|
||||
config_values["thread_id"] = thread_id
|
||||
config: dict[str, Any] = {"configurable": config_values}
|
||||
if metadata:
|
||||
config["metadata"] = metadata
|
||||
if callbacks:
|
||||
|
||||
@@ -47,6 +47,7 @@ def build_async_subagent_graph(name: str) -> Any:
|
||||
_get_default_middleware,
|
||||
_inject_subagent_middleware,
|
||||
)
|
||||
from EvoScientist.middleware import BudgetedSkillsMiddleware
|
||||
from EvoScientist.tools import skill_manager, tavily_search, think_tool
|
||||
from EvoScientist.utils import load_subagents
|
||||
|
||||
@@ -113,13 +114,19 @@ def build_async_subagent_graph(name: str) -> Any:
|
||||
_ensure_auxiliary_chat_model() if name == "scheduler" else _ensure_chat_model()
|
||||
)
|
||||
|
||||
backend = _get_default_backend()
|
||||
if skill_sources := spec.get("skills"):
|
||||
middleware.append(
|
||||
BudgetedSkillsMiddleware(backend=backend, sources=skill_sources)
|
||||
)
|
||||
|
||||
return create_deep_agent(
|
||||
name=name,
|
||||
model=model,
|
||||
system_prompt=spec.get("system_prompt", ""),
|
||||
tools=spec.get("tools", []) + agent_mcp_tools,
|
||||
skills=spec.get("skills"),
|
||||
backend=_get_default_backend(),
|
||||
skills=None,
|
||||
backend=backend,
|
||||
middleware=middleware,
|
||||
subagents=subagents,
|
||||
).with_config({"recursion_limit": cfg.recursion_limit})
|
||||
|
||||
@@ -0,0 +1,170 @@
|
||||
"""Web-specific agent construction owned by EvoScientist.
|
||||
|
||||
Gateway supplies a workspace/checkpointer host context only. Model routes,
|
||||
provider clients, and the profile semantics remain entirely inside the Evo
|
||||
runtime.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import AgentMiddleware, ToolCallRequest
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langgraph.types import Command
|
||||
|
||||
from .config.settings import load_config
|
||||
from .llm.contracts import (
|
||||
AgentExecutionProfile,
|
||||
AgentModelSet,
|
||||
EvoRuntimeError,
|
||||
WebHostContext,
|
||||
)
|
||||
|
||||
|
||||
def web_tool_registry_manifest() -> tuple[tuple[dict[str, Any], ...], str]:
|
||||
"""Return the current bounded Web profile and main-agent MCP tools."""
|
||||
|
||||
names = (
|
||||
"think_tool",
|
||||
"execute",
|
||||
"read_file",
|
||||
"write_file",
|
||||
"edit_file",
|
||||
"ls",
|
||||
"glob",
|
||||
"grep",
|
||||
"write_todos",
|
||||
"web_search",
|
||||
"parse_documents",
|
||||
"use_skill",
|
||||
)
|
||||
if os.environ.get("TAVILY_API_KEY"):
|
||||
names = (*names, "tavily_search")
|
||||
schema = {
|
||||
"type": "object",
|
||||
"additionalProperties": True,
|
||||
"maxProperties": 32,
|
||||
}
|
||||
manifest: tuple[dict[str, Any], ...] = tuple(
|
||||
{
|
||||
"name": name,
|
||||
"description": "EvoScientist Web runtime tool",
|
||||
"schema": schema,
|
||||
}
|
||||
for name in names
|
||||
)
|
||||
from .EvoScientist import _load_mcp_config_once, _load_mcp_tools_cached
|
||||
|
||||
mcp_tools = _load_mcp_tools_cached().get("main", [])
|
||||
dynamic = []
|
||||
for tool in mcp_tools:
|
||||
args_schema = getattr(tool, "args_schema", None)
|
||||
if hasattr(args_schema, "model_json_schema"):
|
||||
args_schema = args_schema.model_json_schema()
|
||||
dynamic.append(
|
||||
{
|
||||
"name": str(getattr(tool, "name", type(tool).__name__)),
|
||||
"description": str(getattr(tool, "description", "")),
|
||||
"schema": args_schema or {},
|
||||
}
|
||||
)
|
||||
static_names = {item["name"] for item in manifest}
|
||||
dynamic_names = [item["name"] for item in dynamic]
|
||||
if static_names & set(dynamic_names) or len(dynamic_names) != len(
|
||||
set(dynamic_names)
|
||||
):
|
||||
raise EvoRuntimeError("TOOL_REGISTRY_CONFLICT")
|
||||
manifest = tuple(sorted((*manifest, *dynamic), key=lambda item: item["name"]))
|
||||
mcp_config_signature, _mcp_config = _load_mcp_config_once()
|
||||
mcp_config_revision = hashlib.sha256(
|
||||
mcp_config_signature.encode("utf-8")
|
||||
).hexdigest()
|
||||
encoded = json.dumps(
|
||||
{
|
||||
"tools": manifest,
|
||||
"mcp_config_revision": mcp_config_revision,
|
||||
"builtin_revision": "evoscientist-web-tools-v1",
|
||||
},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
).encode()
|
||||
return manifest, f"sha256:{hashlib.sha256(encoded).hexdigest()}"
|
||||
|
||||
|
||||
class _ToolRegistryFenceMiddleware(AgentMiddleware):
|
||||
name = "web_tool_registry_fence"
|
||||
|
||||
def __init__(self, expected_revision: str) -> None:
|
||||
super().__init__()
|
||||
self.expected_revision = expected_revision
|
||||
|
||||
def _require_current(self) -> None:
|
||||
_manifest, revision = web_tool_registry_manifest()
|
||||
if revision != self.expected_revision:
|
||||
raise EvoRuntimeError("TOOL_REGISTRY_STALE")
|
||||
|
||||
def wrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]],
|
||||
) -> ToolMessage | Command[Any]:
|
||||
self._require_current()
|
||||
return handler(request)
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[
|
||||
[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]
|
||||
],
|
||||
) -> ToolMessage | Command[Any]:
|
||||
self._require_current()
|
||||
return await handler(request)
|
||||
|
||||
|
||||
def create_web_agent(
|
||||
*, snapshot: Any, host: WebHostContext, model_set: AgentModelSet
|
||||
) -> Any:
|
||||
"""Create a `web_v3` agent without importing Gateway types."""
|
||||
|
||||
from .EvoScientist import create_cli_agent
|
||||
from .middleware.evo_route_fallback import EvoRouteFallbackMiddleware
|
||||
|
||||
config = copy.copy(load_config())
|
||||
config.auto_approve = True
|
||||
config.auto_mode = True
|
||||
config.enable_ask_user = False
|
||||
config.enable_async_subagents = False
|
||||
config.enable_scheduler = False
|
||||
config.memory_workers_enabled = False
|
||||
route_middleware = EvoRouteFallbackMiddleware(
|
||||
model_set.main_fallbacks,
|
||||
route_health=model_set.route_health,
|
||||
capacity=model_set.capacity,
|
||||
)
|
||||
tool_registry_fence = _ToolRegistryFenceMiddleware(
|
||||
host.tool_registry_revision
|
||||
)
|
||||
return create_cli_agent(
|
||||
workspace_dir=host.workspace_dir,
|
||||
memory_dir=host.memory_dir,
|
||||
workspace_backend=host.workspace_backend,
|
||||
checkpointer=host.checkpointer,
|
||||
config=config,
|
||||
chat_model=model_set.main_agent,
|
||||
on_mcp_progress=host.on_mcp_progress,
|
||||
tool_selector_threshold=host.tool_selector_threshold,
|
||||
memory_max_inline_profile_chars=host.memory_max_inline_profile_chars,
|
||||
enable_subagents=False,
|
||||
enable_background_execution=False,
|
||||
main_agent_outer_middlewares=[tool_registry_fence],
|
||||
main_agent_route_middleware=route_middleware,
|
||||
execution_profile=AgentExecutionProfile.web_v3(),
|
||||
agent_model_set=model_set,
|
||||
)
|
||||
@@ -0,0 +1,541 @@
|
||||
"""Conversation-scoped workspace resolution and DeepAgents backend factory."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from deepagents.backends.protocol import (
|
||||
EditResult,
|
||||
ExecuteResponse,
|
||||
FileDownloadResponse,
|
||||
FileUploadResponse,
|
||||
GlobResult,
|
||||
GrepResult,
|
||||
LsResult,
|
||||
ReadResult,
|
||||
SandboxBackendProtocol,
|
||||
WriteResult,
|
||||
)
|
||||
from langchain.tools import ToolRuntime
|
||||
|
||||
from . import paths
|
||||
from .scope_registry import (
|
||||
ScopeAccessError,
|
||||
ScopeRecord,
|
||||
deployment_id_for_workspace,
|
||||
get_scope_registry,
|
||||
)
|
||||
|
||||
IsolationMode = str
|
||||
|
||||
|
||||
def workspace_isolation_mode() -> IsolationMode:
|
||||
value = os.getenv("EVOSCIENTIST_WORKSPACE_ISOLATION", "optional").strip().lower()
|
||||
if value not in {"legacy", "optional", "required"}:
|
||||
raise RuntimeError("EVOSCIENTIST_WORKSPACE_ISOLATION must be legacy, optional or required")
|
||||
return value
|
||||
|
||||
|
||||
def is_required() -> bool:
|
||||
return workspace_isolation_mode() == "required"
|
||||
|
||||
|
||||
def verify_required_executor() -> None:
|
||||
"""Fail startup unless the pinned scope executor is locally usable."""
|
||||
|
||||
docker = shutil.which("docker")
|
||||
if not docker:
|
||||
raise RuntimeError("required workspace isolation needs the docker OCI runtime")
|
||||
image = os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_IMAGE", "").strip()
|
||||
if "@sha256:" not in image:
|
||||
raise RuntimeError("required workspace isolation needs an OCI image pinned by digest")
|
||||
try:
|
||||
probe = subprocess.run(
|
||||
[docker, "image", "inspect", image],
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
except (OSError, subprocess.TimeoutExpired) as exc:
|
||||
raise RuntimeError("required workspace isolation cannot verify the OCI executor") from exc
|
||||
if probe.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"required workspace isolation needs local OCI image {image!r}"
|
||||
)
|
||||
|
||||
|
||||
def current_deployment_id() -> str:
|
||||
return deployment_id_for_workspace(paths.WORKSPACE_ROOT)
|
||||
|
||||
|
||||
def conversation_root(scope_id: str, workspace_root: Path | None = None) -> Path:
|
||||
scope = str(uuid.UUID(scope_id))
|
||||
# The deploy process supplies an absolute workspace root. This helper is
|
||||
# called from synchronous DeepAgents backend factories on the ASGI loop.
|
||||
root = (workspace_root or paths.WORKSPACE_ROOT).expanduser()
|
||||
return root / ".evoscientist" / "conversations" / scope
|
||||
|
||||
|
||||
def conversation_files_dir(scope_id: str, workspace_root: Path | None = None) -> Path:
|
||||
return conversation_root(scope_id, workspace_root) / "files"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScopeContext:
|
||||
deployment_id: str
|
||||
scope_id: str
|
||||
owner_id: str
|
||||
thread_id: str
|
||||
revision: int
|
||||
files_dir: Path
|
||||
runtime_dir: Path
|
||||
primary_thread_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RuntimeScopeConfig:
|
||||
"""Untrusted runtime identifiers parsed without filesystem or Registry I/O."""
|
||||
|
||||
scope_id: str
|
||||
owner_id: str
|
||||
thread_id: str
|
||||
deployment_id: str | None
|
||||
|
||||
|
||||
class ScopedContainerBackend:
|
||||
"""Filesystem backend whose shell commands execute in a scope-only OCI container."""
|
||||
|
||||
def __init__(self, root_dir: Path, *, timeout: int) -> None:
|
||||
from .backends import CustomSandboxBackend
|
||||
|
||||
# Reuse the hardened filesystem operations; only ``execute`` is
|
||||
# replaced so no agent shell runs in the host process.
|
||||
self._filesystem = CustomSandboxBackend(
|
||||
root_dir=str(root_dir), virtual_mode=True, timeout=timeout, dangerous=False
|
||||
)
|
||||
self._root_dir = root_dir
|
||||
self._timeout = timeout
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._filesystem, name)
|
||||
|
||||
def execute(self, command: str, *, timeout: int | None = None) -> Any:
|
||||
from .backends import ExecuteResponse, prepare_sandbox_command
|
||||
|
||||
command, error = prepare_sandbox_command(
|
||||
command, self._filesystem.cwd, virtual_mode=True, dangerous=False
|
||||
)
|
||||
if error:
|
||||
return ExecuteResponse(output=error, exit_code=1, truncated=False)
|
||||
image = os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_IMAGE", "").strip()
|
||||
if "@sha256:" not in image:
|
||||
return ExecuteResponse(
|
||||
output="Required workspace isolation needs an OCI image pinned by digest.",
|
||||
exit_code=125,
|
||||
truncated=False,
|
||||
)
|
||||
effective_timeout = max(1, min(timeout or self._timeout, 3600))
|
||||
invocation = [
|
||||
"docker",
|
||||
"run",
|
||||
"--rm",
|
||||
"--network",
|
||||
"none",
|
||||
"--read-only",
|
||||
"--tmpfs",
|
||||
"/tmp:rw,noexec,nosuid,size=64m",
|
||||
"--cap-drop",
|
||||
"ALL",
|
||||
"--security-opt",
|
||||
"no-new-privileges",
|
||||
"--pids-limit",
|
||||
os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_PIDS", "128"),
|
||||
"--memory",
|
||||
os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_MEMORY", "1g"),
|
||||
"--cpus",
|
||||
os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_CPUS", "1"),
|
||||
"--mount",
|
||||
f"type=bind,src={self._root_dir},dst=/workspace",
|
||||
"--workdir",
|
||||
"/workspace",
|
||||
image,
|
||||
"sh",
|
||||
"-lc",
|
||||
command,
|
||||
]
|
||||
try:
|
||||
completed = subprocess.run(
|
||||
invocation,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=effective_timeout,
|
||||
)
|
||||
except FileNotFoundError:
|
||||
return ExecuteResponse(
|
||||
output="Required workspace isolation needs an OCI runtime (docker was not found).",
|
||||
exit_code=127,
|
||||
truncated=False,
|
||||
)
|
||||
except subprocess.TimeoutExpired as exc:
|
||||
output = (exc.stdout or "") + (exc.stderr or "")
|
||||
return ExecuteResponse(output=output, exit_code=124, truncated=False)
|
||||
output = completed.stdout + completed.stderr
|
||||
return ExecuteResponse(output=output, exit_code=completed.returncode, truncated=False)
|
||||
|
||||
|
||||
def _configurable(runtime: ToolRuntime[Any, Any] | Any | None) -> dict[str, Any]:
|
||||
"""Return the active runnable config, with a non-graph fallback.
|
||||
|
||||
``ToolRuntime`` deliberately does not expose ``RunnableConfig`` during a
|
||||
graph execution. LangGraph keeps it in a context variable instead. The
|
||||
fallback preserves direct callers and unit tests that supply a lightweight
|
||||
runtime object outside a runnable context.
|
||||
"""
|
||||
|
||||
config: Any = None
|
||||
try:
|
||||
from langgraph.config import get_config
|
||||
|
||||
config = get_config()
|
||||
except (ImportError, LookupError, RuntimeError):
|
||||
pass
|
||||
if not isinstance(config, dict) and runtime is not None:
|
||||
config = getattr(runtime, "config", None) or {}
|
||||
if not isinstance(config, dict):
|
||||
return {}
|
||||
configurable = config.get("configurable") or {}
|
||||
return dict(configurable) if isinstance(configurable, dict) else {}
|
||||
|
||||
|
||||
def _required_string(configurable: dict[str, Any], key: str) -> str:
|
||||
value = configurable.get(key)
|
||||
if not isinstance(value, str) or not value:
|
||||
raise ScopeAccessError(f"missing {key}")
|
||||
return value
|
||||
|
||||
|
||||
def _runtime_scope_config(
|
||||
runtime: ToolRuntime[Any, Any] | Any | None,
|
||||
*,
|
||||
kind: str,
|
||||
) -> _RuntimeScopeConfig | None:
|
||||
"""Parse scope identifiers without treating config as an authorization grant."""
|
||||
|
||||
configurable = _configurable(runtime)
|
||||
scope_id = configurable.get("workspace_scope_id")
|
||||
owner_id = configurable.get("workspace_scope_owner_id")
|
||||
thread_id = configurable.get("thread_id")
|
||||
|
||||
if scope_id is None and owner_id is None:
|
||||
if workspace_isolation_mode() == "required":
|
||||
raise ScopeAccessError(f"{kind} requires a workspace scope")
|
||||
return None
|
||||
if (
|
||||
not isinstance(scope_id, str)
|
||||
or not isinstance(owner_id, str)
|
||||
or not isinstance(thread_id, str)
|
||||
):
|
||||
raise ScopeAccessError(f"{kind} has an incomplete workspace scope")
|
||||
try:
|
||||
canonical_scope_id = str(uuid.UUID(scope_id))
|
||||
canonical_owner_id = str(uuid.UUID(owner_id))
|
||||
except ValueError as exc:
|
||||
raise ScopeAccessError(f"{kind} has an invalid workspace scope") from exc
|
||||
deployment_id = configurable.get("workspace_deployment_id")
|
||||
if deployment_id is not None and (
|
||||
not isinstance(deployment_id, str) or not deployment_id
|
||||
):
|
||||
raise ScopeAccessError(f"{kind} has an invalid workspace deployment")
|
||||
return _RuntimeScopeConfig(
|
||||
scope_id=canonical_scope_id,
|
||||
owner_id=canonical_owner_id,
|
||||
thread_id=thread_id,
|
||||
deployment_id=deployment_id,
|
||||
)
|
||||
|
||||
|
||||
def _validated_scope_directories(scope_id: str) -> tuple[Path, Path]:
|
||||
"""Return canonical private directories after preventing symlink escape.
|
||||
|
||||
This function intentionally resolves paths and must run only from a
|
||||
filesystem-operation worker, never from the runtime backend factory.
|
||||
"""
|
||||
|
||||
conversations_dir = (
|
||||
paths.WORKSPACE_ROOT.expanduser() / ".evoscientist" / "conversations"
|
||||
).resolve(strict=True)
|
||||
scope_root = (conversations_dir / scope_id).resolve(strict=True)
|
||||
files_dir = (scope_root / "files").resolve(strict=True)
|
||||
runtime_dir = (scope_root / "runtime").resolve(strict=True)
|
||||
if (
|
||||
scope_root.parent != conversations_dir
|
||||
or files_dir.parent != scope_root
|
||||
or runtime_dir.parent != scope_root
|
||||
):
|
||||
raise ScopeAccessError("workspace directory escapes its scope")
|
||||
if not files_dir.is_dir() or not runtime_dir.is_dir():
|
||||
raise ScopeAccessError("workspace directory is missing")
|
||||
return files_dir, runtime_dir
|
||||
|
||||
|
||||
def _resolve_scope_context(config: _RuntimeScopeConfig | None) -> ScopeContext | None:
|
||||
"""Validate parsed scope identifiers against the active registry."""
|
||||
|
||||
if config is None:
|
||||
return None
|
||||
deployment_id = config.deployment_id or current_deployment_id()
|
||||
registry = get_scope_registry(paths.WORKSPACE_ROOT)
|
||||
if registry.active_lock(deployment_id, "workspace-cutover") is not None:
|
||||
raise ScopeAccessError("workspace cutover is in progress")
|
||||
record = registry.assert_runtime(
|
||||
deployment_id, config.scope_id, config.thread_id, config.owner_id
|
||||
)
|
||||
files_dir, runtime_dir = _validated_scope_directories(record.scope_id)
|
||||
return ScopeContext(
|
||||
deployment_id=deployment_id,
|
||||
scope_id=record.scope_id,
|
||||
owner_id=config.owner_id,
|
||||
thread_id=config.thread_id,
|
||||
revision=record.revision,
|
||||
files_dir=files_dir,
|
||||
runtime_dir=runtime_dir,
|
||||
primary_thread_id=record.primary_thread_id,
|
||||
)
|
||||
|
||||
|
||||
def require_scoped_runtime(
|
||||
runtime: ToolRuntime[Any, Any] | Any | None,
|
||||
*,
|
||||
kind: str = "tool",
|
||||
) -> ScopeContext | None:
|
||||
"""Resolve and validate a runtime scope.
|
||||
|
||||
``optional`` retains legacy CLI compatibility when no scope has been
|
||||
injected. ``required`` never falls back to ``WORKSPACE_ROOT``.
|
||||
"""
|
||||
|
||||
config = _runtime_scope_config(runtime, kind=kind)
|
||||
return _resolve_scope_context(config)
|
||||
|
||||
|
||||
def provision_conversation_scope(
|
||||
thread_id: str,
|
||||
*,
|
||||
deployment_id: str | None = None,
|
||||
scope_id: str | None = None,
|
||||
workspace_root: Path | None = None,
|
||||
lock_operation_id: str | None = None,
|
||||
) -> ScopeRecord:
|
||||
"""Create the registry mapping and private directory for a primary thread."""
|
||||
|
||||
root = (workspace_root or paths.WORKSPACE_ROOT).expanduser()
|
||||
deployment_id = deployment_id or deployment_id_for_workspace(root)
|
||||
registry = get_scope_registry(root)
|
||||
record = registry.provision(
|
||||
deployment_id,
|
||||
thread_id,
|
||||
scope_id=scope_id,
|
||||
lock_operation_id=lock_operation_id,
|
||||
)
|
||||
root = conversation_root(record.scope_id, root)
|
||||
try:
|
||||
(root / "files").mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
(root / "runtime").mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
for directory in (root, root / "files", root / "runtime"):
|
||||
try:
|
||||
directory.chmod(0o700)
|
||||
except OSError:
|
||||
pass
|
||||
except OSError:
|
||||
# Keep the durable reservation for the recovery job; it is safer than
|
||||
# silently falling back to the shared deployment root.
|
||||
raise
|
||||
return record
|
||||
|
||||
|
||||
def _build_backend(root_dir: Path, *, dangerous: bool) -> Any:
|
||||
from deepagents.backends import CompositeBackend
|
||||
|
||||
from .backends import (
|
||||
CustomSandboxBackend,
|
||||
MemoryFilesystemBackend,
|
||||
MergedSkillsBackend,
|
||||
)
|
||||
from .EvoScientist import SKILLS_DIR
|
||||
|
||||
cfg_timeout = int(os.getenv("EVOSCIENTIST_SANDBOX_EXECUTE_TIMEOUT", "300"))
|
||||
ws_backend: Any
|
||||
if is_required():
|
||||
ws_backend = ScopedContainerBackend(root_dir, timeout=cfg_timeout)
|
||||
else:
|
||||
ws_backend = CustomSandboxBackend(
|
||||
root_dir=str(root_dir),
|
||||
virtual_mode=True,
|
||||
timeout=cfg_timeout,
|
||||
dangerous=dangerous,
|
||||
)
|
||||
return CompositeBackend(
|
||||
default=ws_backend,
|
||||
routes={
|
||||
"/skills/": MergedSkillsBackend(
|
||||
primary_dir=str(paths.USER_SKILLS_DIR),
|
||||
global_dir=str(paths.GLOBAL_SKILLS_DIR),
|
||||
secondary_dir=SKILLS_DIR,
|
||||
),
|
||||
"/memories/": MemoryFilesystemBackend(
|
||||
root_dir=str(paths.MEMORIES_DIR), virtual_mode=True
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class DeferredScopedBackend(SandboxBackendProtocol):
|
||||
"""Resolve the scoped filesystem backend only from a worker thread.
|
||||
|
||||
DeepAgents invokes its deprecated backend factory from async middleware.
|
||||
Its concrete filesystem backends synchronously call ``Path.resolve()`` in
|
||||
their constructors, so doing that work in the factory makes every run fail
|
||||
under LangGraph's blocking-call detector. This proxy itself is I/O-free;
|
||||
the inherited async methods dispatch the synchronous operations to a
|
||||
thread, where Registry validation and concrete backend construction occur.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: _RuntimeScopeConfig,
|
||||
*,
|
||||
dangerous: bool,
|
||||
) -> None:
|
||||
self._config = config
|
||||
self._dangerous = dangerous
|
||||
self._backend: Any | None = None
|
||||
self._backend_key: tuple[str, str, str, int] | None = None
|
||||
self._lock = threading.RLock()
|
||||
|
||||
@property
|
||||
def id(self) -> str:
|
||||
# This is queried while composing the model request; do not initialize
|
||||
# the real backend or touch the Registry here.
|
||||
return f"scope-{self._config.scope_id[:8]}-{self._config.owner_id[:8]}"
|
||||
|
||||
def _delegate(self) -> Any:
|
||||
"""Validate the current scope and return a concrete backend.
|
||||
|
||||
Every operation enters here, so a deleted scope or stale owner cannot
|
||||
keep using a backend constructed before the lifecycle transition.
|
||||
"""
|
||||
|
||||
# Async backend methods run this code in a worker thread. LangGraph's
|
||||
# RunnableConfig context variable is not available there, so validate
|
||||
# the immutable scope parsed by the factory on the graph thread.
|
||||
context = _resolve_scope_context(self._config)
|
||||
if context is None:
|
||||
raise ScopeAccessError("scoped backend lost its workspace scope")
|
||||
if is_required() and self._dangerous:
|
||||
raise ScopeAccessError(
|
||||
"dangerous_mode is incompatible with required isolation"
|
||||
)
|
||||
key = (context.scope_id, context.owner_id, context.thread_id, context.revision)
|
||||
with self._lock:
|
||||
if self._backend is None or self._backend_key != key:
|
||||
self._backend = _build_backend(
|
||||
context.files_dir, dangerous=self._dangerous
|
||||
)
|
||||
self._backend_key = key
|
||||
return self._backend
|
||||
|
||||
def ls(self, path: str) -> LsResult:
|
||||
return self._delegate().ls(path)
|
||||
|
||||
def read(self, file_path: str, offset: int = 0, limit: int = 2000) -> ReadResult:
|
||||
return self._delegate().read(file_path, offset, limit)
|
||||
|
||||
def grep(
|
||||
self, pattern: str, path: str | None = None, glob: str | None = None
|
||||
) -> GrepResult:
|
||||
return self._delegate().grep(pattern, path, glob)
|
||||
|
||||
def glob(self, pattern: str, path: str | None = None) -> GlobResult:
|
||||
return self._delegate().glob(pattern, path)
|
||||
|
||||
def write(self, file_path: str, content: str) -> WriteResult:
|
||||
return self._delegate().write(file_path, content)
|
||||
|
||||
def edit(
|
||||
self,
|
||||
file_path: str,
|
||||
old_string: str,
|
||||
new_string: str,
|
||||
replace_all: bool = False,
|
||||
) -> EditResult:
|
||||
return self._delegate().edit(file_path, old_string, new_string, replace_all)
|
||||
|
||||
def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]:
|
||||
return self._delegate().upload_files(files)
|
||||
|
||||
def download_files(self, paths: list[str]) -> list[FileDownloadResponse]:
|
||||
return self._delegate().download_files(paths)
|
||||
|
||||
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
||||
return self._delegate().execute(command, timeout=timeout)
|
||||
|
||||
|
||||
def create_workspace_backend(
|
||||
runtime: ToolRuntime[Any, Any],
|
||||
*,
|
||||
legacy_backend: Callable[[], Any],
|
||||
dangerous: bool = False,
|
||||
allow_unscoped_legacy: bool = True,
|
||||
) -> Any:
|
||||
"""Return a backend handle without blocking the Agent event loop."""
|
||||
|
||||
config = _runtime_scope_config(runtime, kind="filesystem backend")
|
||||
if config is None:
|
||||
if not allow_unscoped_legacy:
|
||||
raise ScopeAccessError(
|
||||
"deployed graph runs require a workspace scope"
|
||||
)
|
||||
return legacy_backend()
|
||||
if is_required() and dangerous:
|
||||
raise ScopeAccessError("dangerous_mode is incompatible with required isolation")
|
||||
return DeferredScopedBackend(config, dangerous=dangerous)
|
||||
|
||||
|
||||
def create_workspace_backend_factory(
|
||||
legacy_backend: Callable[[], Any],
|
||||
*,
|
||||
dangerous: bool = False,
|
||||
allow_unscoped_legacy: bool = True,
|
||||
) -> Callable[[ToolRuntime[Any, Any]], Any]:
|
||||
def factory(runtime: ToolRuntime[Any, Any]) -> Any:
|
||||
return create_workspace_backend(
|
||||
runtime,
|
||||
legacy_backend=legacy_backend,
|
||||
dangerous=dangerous,
|
||||
allow_unscoped_legacy=allow_unscoped_legacy,
|
||||
)
|
||||
|
||||
return factory
|
||||
|
||||
|
||||
def workspace_metadata(record: ScopeRecord) -> dict[str, str | int]:
|
||||
"""Metadata mirrored onto the LangGraph primary thread by trusted callers."""
|
||||
|
||||
return {
|
||||
"workspace_schema_version": 1,
|
||||
"workspace_scope_id": record.scope_id,
|
||||
"workspace_status": record.state,
|
||||
"workspace_scope_owner_id": record.primary_owner_id,
|
||||
"workspace_scope_revision": record.revision,
|
||||
"workspace_deployment_id": record.deployment_id,
|
||||
}
|
||||
@@ -26,6 +26,7 @@ dependencies = [
|
||||
"langchain-openrouter>=0.2.5",
|
||||
"tavily-python>=0.7",
|
||||
"pyyaml>=6.0",
|
||||
"rfc8785==0.1.4",
|
||||
"rich>=15.0",
|
||||
"prompt-toolkit>=3.0",
|
||||
"questionary>=2.1",
|
||||
|
||||
Executable
+30
@@ -0,0 +1,30 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
PROJECT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)"
|
||||
LANGGRAPH_CONFIG="${PROJECT_DIR}/EvoScientist/langgraph_dev/langgraph.json"
|
||||
HOST="${EVOSCIENTIST_LANGGRAPH_HOST:-127.0.0.1}"
|
||||
PORT="${EVOSCIENTIST_LANGGRAPH_DEV_PORT:-3076}"
|
||||
WEB_ENV="${PROJECT_DIR}/../Ai4Sci-Web/.env"
|
||||
|
||||
if [[ ! -x "${PROJECT_DIR}/.venv/bin/langgraph" ]]; then
|
||||
echo "LangGraph executable not found: ${PROJECT_DIR}/.venv/bin/langgraph" >&2
|
||||
echo "Run 'uv sync' in ${PROJECT_DIR} first." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
cd "${PROJECT_DIR}"
|
||||
|
||||
# The Web Gateway always supplies a verified conversation workspace scope.
|
||||
export EVOSCIENTIST_DEPLOY_MODE="${EVOSCIENTIST_DEPLOY_MODE:-full}"
|
||||
export EVOSCIENTIST_WORKSPACE_DIR="${EVOSCIENTIST_WORKSPACE_DIR:-${PROJECT_DIR}/../.ai4sci/workspace}"
|
||||
|
||||
exec uv run --env-file "${WEB_ENV}" langgraph dev \
|
||||
--config "${LANGGRAPH_CONFIG}" \
|
||||
--host "${HOST}" \
|
||||
--port "${PORT}" \
|
||||
--no-browser \
|
||||
--no-reload \
|
||||
--allow-blocking \
|
||||
--n-jobs-per-worker 1
|
||||
@@ -0,0 +1 @@
|
||||
"""EvoScientist test package."""
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
from dataclasses import fields
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.config_admin import EvoModelConfigAdminService
|
||||
from EvoScientist.llm.contracts import (
|
||||
ADMIN_CONTROL_VERSION,
|
||||
CommitProposalRequest,
|
||||
CreateProposalRequest,
|
||||
HmacGrantAuthority,
|
||||
ProbeProposalRequest,
|
||||
UpdateProposalRequest,
|
||||
ValidateProposalRequest,
|
||||
)
|
||||
from EvoScientist.llm.crypto import HmacKeyRing, KeyMaterial, sha256_id
|
||||
from EvoScientist.llm.model_config import FileEvoModelConfigStore
|
||||
from EvoScientist.llm.secret_store import EncryptedModelSecretStore
|
||||
from tests.test_provider_model_config_v3 import v3_payload
|
||||
|
||||
|
||||
def _request(authority, cls, action: str, **values):
|
||||
payload = {"admin_control_version": ADMIN_CONTROL_VERSION, **values}
|
||||
operation_id = str(values["operation_id"])
|
||||
grant = authority.sign_admin(
|
||||
subject_id="admin-1",
|
||||
action=action,
|
||||
operation_id=operation_id,
|
||||
request_digest=sha256_id(payload),
|
||||
)
|
||||
valid = {item.name for item in fields(cls)}
|
||||
complete = {**values, "admin_grant": grant, "admin_control_version": 2}
|
||||
return cls(**{key: complete[key] for key in valid})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_control_v2_proposal_commit(tmp_path) -> None:
|
||||
authority = HmacGrantAuthority("g" * 32, key_id="grant-v1")
|
||||
ring = HmacKeyRing(KeyMaterial.create("identity-v1", "i" * 32))
|
||||
store = FileEvoModelConfigStore(
|
||||
tmp_path / "model_routes.yaml",
|
||||
admin_verifier=authority,
|
||||
ops_path=tmp_path / "model_config_ops.sqlite",
|
||||
)
|
||||
secrets = EncryptedModelSecretStore(
|
||||
tmp_path / "model_secrets.sqlite", master_secret="s" * 32
|
||||
)
|
||||
payload = v3_payload()
|
||||
for provider in payload["providers"]:
|
||||
item = secrets.create_pending(
|
||||
provider["provider_id"],
|
||||
"sk-" + provider["provider_id"],
|
||||
created_by="admin-1",
|
||||
operation_id="secret-" + provider["provider_id"],
|
||||
)
|
||||
provider["connection"]["credential_ref"] = item.ref
|
||||
service = EvoModelConfigAdminService(
|
||||
store,
|
||||
grant_authority=authority,
|
||||
identity_key_ring=ring,
|
||||
secret_resolver=secrets.resolve,
|
||||
secret_store=secrets,
|
||||
probe_runner=lambda _config, _route, _kind: True,
|
||||
)
|
||||
|
||||
created = service.create_proposal(
|
||||
_request(
|
||||
authority,
|
||||
CreateProposalRequest,
|
||||
"model_config:proposal:create",
|
||||
operation_id="create-1",
|
||||
expected_active_revision=0,
|
||||
)
|
||||
)
|
||||
updated = service.update_proposal(
|
||||
_request(
|
||||
authority,
|
||||
UpdateProposalRequest,
|
||||
"model_config:proposal:update",
|
||||
operation_id="update-1",
|
||||
proposal_id=created.proposal_id,
|
||||
expected_state_version=created.state_version,
|
||||
expected_draft_etag=created.draft_etag,
|
||||
draft_payload=payload,
|
||||
)
|
||||
)
|
||||
validated = service.validate_proposal(
|
||||
_request(
|
||||
authority,
|
||||
ValidateProposalRequest,
|
||||
"model_config:proposal:validate",
|
||||
operation_id="validate-1",
|
||||
proposal_id=created.proposal_id,
|
||||
expected_state_version=updated.state_version,
|
||||
expected_draft_etag=updated.draft_etag,
|
||||
)
|
||||
)
|
||||
current = validated
|
||||
for index, route in enumerate(validated.routes):
|
||||
for kind in route["required_probe_kinds"]:
|
||||
current = await service.probe_proposal(
|
||||
_request(
|
||||
authority,
|
||||
ProbeProposalRequest,
|
||||
"model_config:proposal:probe",
|
||||
operation_id=f"probe-{index}-{kind}",
|
||||
proposal_id=created.proposal_id,
|
||||
validated_digest=validated.validated_digest,
|
||||
route_semantics_hash=route["route_semantics_hash"],
|
||||
probe_kind=kind,
|
||||
)
|
||||
)
|
||||
assert current.state == "READY"
|
||||
committed = service.commit_proposal(
|
||||
_request(
|
||||
authority,
|
||||
CommitProposalRequest,
|
||||
"model_config:proposal:commit",
|
||||
operation_id="commit-1",
|
||||
proposal_id=created.proposal_id,
|
||||
expected_active_revision=0,
|
||||
expected_state_version=current.state_version,
|
||||
expected_draft_etag=current.draft_etag,
|
||||
validated_digest=validated.validated_digest,
|
||||
evidence_ids=current.evidence_ids,
|
||||
)
|
||||
)
|
||||
assert committed.state == "COMMITTED"
|
||||
assert committed.active_revision == 1
|
||||
active = store.load()
|
||||
assert active.schema_version == 3
|
||||
assert len(active.providers) == 4
|
||||
assert all(item.status == "active" for item in secrets.list_metadata())
|
||||
|
||||
with sqlite3.connect(store.ops_path) as connection:
|
||||
connection.execute(
|
||||
"UPDATE config_commit_operations SET stage='CONFIG_COMMITTED' WHERE operation_id='commit-1'"
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE config_proposals SET state='COMMITTING' WHERE proposal_id=?",
|
||||
(created.proposal_id,),
|
||||
)
|
||||
EvoModelConfigAdminService(
|
||||
store,
|
||||
grant_authority=authority,
|
||||
identity_key_ring=ring,
|
||||
secret_resolver=secrets.resolve,
|
||||
secret_store=secrets,
|
||||
probe_runner=lambda _config, _route, _kind: True,
|
||||
)
|
||||
with sqlite3.connect(store.ops_path) as connection:
|
||||
stage = connection.execute(
|
||||
"SELECT stage FROM config_commit_operations WHERE operation_id='commit-1'"
|
||||
).fetchone()[0]
|
||||
state = connection.execute(
|
||||
"SELECT state FROM config_proposals WHERE proposal_id=?",
|
||||
(created.proposal_id,),
|
||||
).fetchone()[0]
|
||||
assert stage == "COMPLETED"
|
||||
assert state == "COMMITTED"
|
||||
@@ -68,10 +68,13 @@ def test_create_cli_agent_accepts_host_backend_and_memory_options(
|
||||
assert calls["middleware_kwargs"]["tool_selector_threshold"] == 8
|
||||
assert calls["middleware_kwargs"]["memory_max_inline_profile_chars"] == 1000
|
||||
assert calls["middleware_kwargs"]["enable_background_execution"] is False
|
||||
assert calls["middleware_kwargs"]["enable_legacy_model_fallback"] is True
|
||||
assert calls["agent_config"] == {"recursion_limit": 321}
|
||||
|
||||
|
||||
def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, tmp_path):
|
||||
def test_create_cli_agent_installs_route_middleware_in_fixed_slot(
|
||||
monkeypatch, tmp_path
|
||||
):
|
||||
import EvoScientist.EvoScientist as agent_module
|
||||
from EvoScientist.config.settings import EvoScientistConfig
|
||||
|
||||
@@ -106,11 +109,14 @@ def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, t
|
||||
monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _Backend)
|
||||
monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _Backend)
|
||||
monkeypatch.setattr(agent_module, "set_active_workspace", lambda _path: None)
|
||||
|
||||
def fake_default_middleware(**kwargs):
|
||||
calls["middleware_kwargs"] = kwargs
|
||||
return list(default_chain)
|
||||
|
||||
monkeypatch.setattr(agent_module, "_get_default_middleware", fake_default_middleware)
|
||||
monkeypatch.setattr(
|
||||
agent_module, "_get_default_middleware", fake_default_middleware
|
||||
)
|
||||
|
||||
def fake_load(_backend, middleware, **_kwargs):
|
||||
calls["middleware"] = middleware
|
||||
@@ -128,10 +134,68 @@ def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, t
|
||||
)
|
||||
|
||||
assert calls["middleware_kwargs"]["enable_legacy_model_fallback"] is False
|
||||
assert [middleware.name for middleware in calls["middleware"][:5]] == [
|
||||
assert [middleware.name for middleware in calls["middleware"][:6]] == [
|
||||
"error_normalization",
|
||||
"provider_context_media",
|
||||
"configurable_model",
|
||||
"gateway_route_fallback",
|
||||
"context_editing",
|
||||
"tool_protocol_guard",
|
||||
]
|
||||
|
||||
|
||||
def test_create_cli_agent_replaces_framework_skill_injection(monkeypatch, tmp_path):
|
||||
import EvoScientist.EvoScientist as agent_module
|
||||
from EvoScientist.config.settings import EvoScientistConfig
|
||||
from EvoScientist.middleware import BudgetedSkillsMiddleware
|
||||
|
||||
calls = {}
|
||||
|
||||
class _Backend:
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
class _CompositeBackend:
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
class _Agent:
|
||||
def with_config(self, _config):
|
||||
return self
|
||||
|
||||
def _create_deep_agent(**kwargs):
|
||||
calls["kwargs"] = kwargs
|
||||
return _Agent()
|
||||
|
||||
monkeypatch.setattr("deepagents.backends.CompositeBackend", _CompositeBackend)
|
||||
monkeypatch.setattr("deepagents.create_deep_agent", _create_deep_agent)
|
||||
monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _Backend)
|
||||
monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _Backend)
|
||||
monkeypatch.setattr(agent_module, "set_active_workspace", lambda _path: None)
|
||||
monkeypatch.setattr(agent_module, "_get_default_middleware", lambda **_kwargs: [])
|
||||
monkeypatch.setattr(
|
||||
agent_module,
|
||||
"load_mcp_and_build_kwargs",
|
||||
lambda *_args, **_kwargs: {
|
||||
"skills": ["/skills/"],
|
||||
"middleware": [],
|
||||
"subagents": [{"name": "research", "skills": ["/skills/"]}],
|
||||
},
|
||||
)
|
||||
|
||||
agent_module.create_cli_agent(
|
||||
workspace_dir=str(tmp_path),
|
||||
checkpointer=object(),
|
||||
config=EvoScientistConfig(auto_approve=True),
|
||||
chat_model=object(),
|
||||
workspace_backend=object(),
|
||||
)
|
||||
|
||||
assert calls["kwargs"]["skills"] is None
|
||||
assert any(
|
||||
isinstance(middleware, BudgetedSkillsMiddleware)
|
||||
for middleware in calls["kwargs"]["middleware"]
|
||||
)
|
||||
subagent = calls["kwargs"]["subagents"][0]
|
||||
assert subagent["skills"] is None
|
||||
assert isinstance(subagent["middleware"][0], BudgetedSkillsMiddleware)
|
||||
|
||||
@@ -78,7 +78,7 @@ def test_factory_requests_async_safe_middleware(
|
||||
"name": "writing-agent",
|
||||
"system_prompt": "",
|
||||
"tools": [],
|
||||
"skills": None,
|
||||
"skills": ["/skills/"],
|
||||
}
|
||||
]
|
||||
# ``create_deep_agent(...).with_config({...})`` chain — return something
|
||||
@@ -94,6 +94,12 @@ def test_factory_requests_async_safe_middleware(
|
||||
for_async_subagent=True,
|
||||
memory_source_agent="writing-agent",
|
||||
)
|
||||
from EvoScientist.middleware import BudgetedSkillsMiddleware
|
||||
|
||||
assert mock_create.call_args.kwargs["skills"] is None
|
||||
assert isinstance(
|
||||
mock_create.call_args.kwargs["middleware"][-1], BudgetedSkillsMiddleware
|
||||
)
|
||||
subagents = mock_create.call_args.kwargs["subagents"]
|
||||
assert subagents[0]["name"] == "general-purpose"
|
||||
_assert_subagent_memory_middleware(
|
||||
|
||||
+13
-5
@@ -25,6 +25,10 @@ from EvoScientist.config import (
|
||||
set_config_value,
|
||||
)
|
||||
|
||||
|
||||
def test_langgraph_dev_port_defaults_to_ai4sci_runtime_port():
|
||||
assert EvoScientistConfig().langgraph_dev_port == 3076
|
||||
|
||||
# =============================================================================
|
||||
# Fixtures
|
||||
# =============================================================================
|
||||
@@ -60,6 +64,7 @@ def temp_config_dir(tmp_path, monkeypatch):
|
||||
"ANTHROPIC_API_KEY",
|
||||
"OPENAI_API_KEY",
|
||||
"TAVILY_API_KEY",
|
||||
"S2_API_KEY",
|
||||
"EVOSCIENTIST_DEFAULT_MODE",
|
||||
"EVOSCIENTIST_WORKSPACE_DIR",
|
||||
"EVOSCIENTIST_UI_BACKEND",
|
||||
@@ -90,6 +95,7 @@ def clean_env(monkeypatch):
|
||||
"ANTHROPIC_API_KEY",
|
||||
"OPENAI_API_KEY",
|
||||
"TAVILY_API_KEY",
|
||||
"S2_API_KEY",
|
||||
"EVOSCIENTIST_DEFAULT_MODE",
|
||||
"EVOSCIENTIST_WORKSPACE_DIR",
|
||||
"EVOSCIENTIST_UI_BACKEND",
|
||||
@@ -125,6 +131,7 @@ class TestEvoScientistConfig:
|
||||
assert config.anthropic_api_key == ""
|
||||
assert config.openai_api_key == ""
|
||||
assert config.tavily_api_key == ""
|
||||
assert config.semantic_scholar_api_key == ""
|
||||
assert config.provider == "anthropic"
|
||||
assert config.model == "claude-sonnet-4-6"
|
||||
assert config.default_mode == "daemon"
|
||||
@@ -132,7 +139,6 @@ class TestEvoScientistConfig:
|
||||
assert config.show_thinking is True
|
||||
assert config.ui_backend == "tui"
|
||||
assert config.log_level == "warning"
|
||||
assert config.reasoning_effort == ""
|
||||
assert config.openrouter_anthropic_prompt_cache is True
|
||||
assert config.openrouter_http_referer == (
|
||||
"https://github.com/EvoScientist/EvoScientist"
|
||||
@@ -599,12 +605,12 @@ class TestPriorityChain:
|
||||
config = get_effective_config()
|
||||
assert config.log_level == "DEBUG"
|
||||
|
||||
def test_env_reasoning_effort_override(self, temp_config_dir, monkeypatch):
|
||||
"""Reasoning effort can be selected via environment variable."""
|
||||
save_config(EvoScientistConfig(reasoning_effort="medium"))
|
||||
def test_env_reasoning_effort_is_not_a_config_override(self, temp_config_dir, monkeypatch):
|
||||
"""Reasoning is selected by the invocation plan, never deployment env."""
|
||||
save_config(EvoScientistConfig())
|
||||
monkeypatch.setenv("EVOSCIENTIST_REASONING_EFFORT", "high")
|
||||
config = get_effective_config()
|
||||
assert config.reasoning_effort == "high"
|
||||
assert not hasattr(config, "reasoning_effort")
|
||||
|
||||
def test_env_channel_debug_tracing_override(self, temp_config_dir, monkeypatch):
|
||||
"""Channel tracing can be enabled via environment variable."""
|
||||
@@ -750,6 +756,7 @@ class TestApplyConfigToEnv:
|
||||
anthropic_api_key="config-ant-key",
|
||||
openai_api_key="config-oai-key",
|
||||
tavily_api_key="config-tav-key",
|
||||
semantic_scholar_api_key="config-s2-key",
|
||||
)
|
||||
|
||||
apply_config_to_env(config)
|
||||
@@ -757,6 +764,7 @@ class TestApplyConfigToEnv:
|
||||
assert os.environ.get("ANTHROPIC_API_KEY") == "config-ant-key"
|
||||
assert os.environ.get("OPENAI_API_KEY") == "config-oai-key"
|
||||
assert os.environ.get("TAVILY_API_KEY") == "config-tav-key"
|
||||
assert os.environ.get("S2_API_KEY") == "config-s2-key"
|
||||
|
||||
def test_does_not_override_existing_env(self, monkeypatch):
|
||||
"""Test that existing env vars are not overridden."""
|
||||
|
||||
@@ -7,6 +7,7 @@ from langchain.agents.middleware.types import ModelRequest
|
||||
from langchain_core.exceptions import ContextOverflowError
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError
|
||||
from EvoScientist.middleware.context_overflow import ContextOverflowMapperMiddleware
|
||||
|
||||
|
||||
@@ -24,6 +25,15 @@ def test_is_context_limit_error_anthropic():
|
||||
assert mw._is_context_limit_error(exc) is True
|
||||
|
||||
|
||||
def test_is_context_limit_error_for_runtime_admission_guard():
|
||||
mw = ContextOverflowMapperMiddleware()
|
||||
|
||||
assert (
|
||||
mw._is_context_limit_error(EvoRuntimeError("MODEL_CONTEXT_WINDOW_EXCEEDED"))
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_is_not_context_limit_error_without_400():
|
||||
mw = ContextOverflowMapperMiddleware()
|
||||
exc = Exception("context_length_exceeded, but no status code")
|
||||
|
||||
@@ -295,6 +295,8 @@ class TestMiddleware:
|
||||
self._run_awrap(mw, req, handler)
|
||||
assert excinfo.value.provider == "openrouter"
|
||||
assert excinfo.value.__cause__ is raised
|
||||
assert excinfo.value.message == "Provider request failed."
|
||||
assert "boom" not in str(excinfo.value.model_dump())
|
||||
|
||||
def test_awrap_passes_through_non_provider_model_exception(self):
|
||||
"""If the model isn't a recognized provider SDK, the exception
|
||||
|
||||
@@ -0,0 +1,328 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from langchain.agents.middleware.types import ModelRequest
|
||||
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage
|
||||
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError
|
||||
from EvoScientist.llm.errors import (
|
||||
AgentControlError,
|
||||
ModelProviderResponseError,
|
||||
ModelToolProtocolError,
|
||||
)
|
||||
from EvoScientist.middleware.evo_route_fallback import EvoRouteFallbackMiddleware
|
||||
|
||||
|
||||
class _Model:
|
||||
def __init__(self, route_key: str, *, supports_tools: bool | None = None) -> None:
|
||||
self.metadata = {"route_key": route_key}
|
||||
if supports_tools is not None:
|
||||
self.metadata["route_supports_tools"] = supports_tools
|
||||
|
||||
|
||||
class _Health:
|
||||
def __init__(self, open_routes: set[str] | None = None) -> None:
|
||||
self.open_routes = open_routes or set()
|
||||
|
||||
def is_open(self, route_key: str) -> bool:
|
||||
return route_key in self.open_routes
|
||||
|
||||
|
||||
def _request(model: _Model) -> ModelRequest:
|
||||
return ModelRequest(model=model, messages=[], tools=[])
|
||||
|
||||
|
||||
def test_route_strips_tools_when_frozen_model_capability_is_false():
|
||||
model = _Model("text-only", supports_tools=False)
|
||||
request = ModelRequest(
|
||||
model=model,
|
||||
messages=[],
|
||||
tools=[{"type": "function", "function": {"name": "search"}}],
|
||||
)
|
||||
|
||||
def handler(routed: ModelRequest):
|
||||
return routed.tools
|
||||
|
||||
assert EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler) == []
|
||||
|
||||
|
||||
def test_tools_disabled_route_removes_checkpoint_tool_protocol():
|
||||
model = _Model("text-only", supports_tools=False)
|
||||
request = ModelRequest(
|
||||
model=model,
|
||||
messages=[
|
||||
SystemMessage("Keep the answer concise."),
|
||||
AIMessage(
|
||||
content="",
|
||||
tool_calls=[{"name": "search", "args": {"q": "K3"}, "id": "call_1"}],
|
||||
),
|
||||
ToolMessage(content="tool result", tool_call_id="call_1", name="search"),
|
||||
AIMessage(
|
||||
content="The tool found a result.",
|
||||
additional_kwargs={"tool_calls": [{"id": "call_2"}]},
|
||||
),
|
||||
HumanMessage("Continue."),
|
||||
],
|
||||
tools=[{"type": "function", "function": {"name": "search"}}],
|
||||
tool_choice="required",
|
||||
response_format={"type": "json_object"},
|
||||
model_settings={"temperature": 0, "max_tokens": 1},
|
||||
)
|
||||
|
||||
def handler(routed: ModelRequest):
|
||||
return routed
|
||||
|
||||
routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler)
|
||||
|
||||
assert routed.tools == []
|
||||
assert routed.tool_choice is None
|
||||
assert routed.response_format is None
|
||||
assert routed.model_settings == {}
|
||||
assert [message.type for message in routed.messages] == [
|
||||
"system",
|
||||
"ai",
|
||||
"ai",
|
||||
"human",
|
||||
]
|
||||
assistant = routed.messages[1]
|
||||
assert isinstance(assistant, AIMessage)
|
||||
assert assistant.content == "[Completed tool result: search]\ntool result"
|
||||
assert routed.messages[2].content == "The tool found a result."
|
||||
assert routed.messages[2].tool_calls == []
|
||||
assert "tool_calls" not in routed.messages[2].additional_kwargs
|
||||
|
||||
|
||||
def test_tools_disabled_route_projects_tool_result_and_content_blocks_to_text():
|
||||
model = _Model("text-only", supports_tools=False)
|
||||
request = ModelRequest(
|
||||
model=model,
|
||||
messages=[
|
||||
AIMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "I checked the source."},
|
||||
{"type": "tool_call", "id": "call_1", "name": "read"},
|
||||
],
|
||||
),
|
||||
ToolMessage(
|
||||
content="x" * 12_100,
|
||||
tool_call_id="call_1",
|
||||
name="read",
|
||||
),
|
||||
],
|
||||
tools=[{"type": "function", "function": {"name": "read"}}],
|
||||
)
|
||||
|
||||
routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, lambda item: item)
|
||||
|
||||
assert routed.tools == []
|
||||
assert len(routed.messages) == 2
|
||||
assert routed.messages[0].content == "I checked the source."
|
||||
assert routed.messages[0].tool_calls == []
|
||||
assert routed.messages[1].content.startswith("[Completed tool result: read]\n")
|
||||
assert routed.messages[1].content.endswith("[Tool result truncated]")
|
||||
|
||||
|
||||
def test_route_keeps_tools_when_frozen_model_capability_is_true():
|
||||
model = _Model("tools", supports_tools=True)
|
||||
tools = [{"type": "function", "function": {"name": "search"}}]
|
||||
request = ModelRequest(
|
||||
model=model,
|
||||
messages=[],
|
||||
tools=tools,
|
||||
model_settings={"temperature": 0},
|
||||
)
|
||||
|
||||
def handler(routed: ModelRequest):
|
||||
return routed
|
||||
|
||||
routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler)
|
||||
assert routed.tools == tools
|
||||
assert routed.model_settings == {}
|
||||
|
||||
|
||||
def test_route_drops_empty_assistant_history_but_keeps_tool_calls():
|
||||
model = _Model("tools", supports_tools=True)
|
||||
tool_message = AIMessage(
|
||||
content="",
|
||||
tool_calls=[{"name": "search", "args": {"q": "K3"}, "id": "call_1"}],
|
||||
)
|
||||
request = ModelRequest(
|
||||
model=model,
|
||||
messages=[
|
||||
HumanMessage("First"),
|
||||
AIMessage(content="", additional_kwargs={"reasoning_content": "hidden"}),
|
||||
HumanMessage("Continue"),
|
||||
tool_message,
|
||||
],
|
||||
tools=[],
|
||||
)
|
||||
|
||||
routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, lambda item: item)
|
||||
|
||||
assert routed.messages == [request.messages[0], request.messages[2], tool_message]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_provider_response_retries_same_route_once():
|
||||
primary = _Model("primary", supports_tools=True)
|
||||
request = ModelRequest(model=primary, messages=[HumanMessage("Answer")], tools=[])
|
||||
seen: list[ModelRequest] = []
|
||||
|
||||
async def handler(routed: ModelRequest):
|
||||
seen.append(routed)
|
||||
if len(seen) == 1:
|
||||
return AIMessage(
|
||||
content="", additional_kwargs={"reasoning_content": "hidden"}
|
||||
)
|
||||
return AIMessage(content="Final answer")
|
||||
|
||||
response = await EvoRouteFallbackMiddleware([]).awrap_model_call(request, handler)
|
||||
|
||||
assert isinstance(response, AIMessage)
|
||||
assert response.content == "Final answer"
|
||||
assert len(seen) == 2
|
||||
assert isinstance(seen[1].messages[0], SystemMessage)
|
||||
assert "without final text" in str(seen[1].messages[0].content)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_repeated_empty_provider_response_fails_with_specific_code():
|
||||
primary = _Model("primary", supports_tools=True)
|
||||
attempts = 0
|
||||
|
||||
async def handler(_routed: ModelRequest):
|
||||
nonlocal attempts
|
||||
attempts += 1
|
||||
return AIMessage(content=[{"type": "reasoning", "summary": []}])
|
||||
|
||||
with pytest.raises(ModelProviderResponseError) as captured:
|
||||
await EvoRouteFallbackMiddleware([]).awrap_model_call(
|
||||
ModelRequest(model=primary, messages=[], tools=[]), handler
|
||||
)
|
||||
|
||||
assert captured.value.code == "MODEL_PROVIDER_RESPONSE_INVALID"
|
||||
assert attempts == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_applies_its_own_tool_capability():
|
||||
primary = _Model("primary", supports_tools=True)
|
||||
fallback = _Model("fallback", supports_tools=False)
|
||||
tools = [{"type": "function", "function": {"name": "search"}}]
|
||||
request = ModelRequest(model=primary, messages=[], tools=tools)
|
||||
seen = []
|
||||
|
||||
async def handler(routed: ModelRequest):
|
||||
seen.append((routed.model.metadata["route_key"], list(routed.tools)))
|
||||
if routed.model is primary:
|
||||
raise ConnectionError("upstream unavailable")
|
||||
return "ok"
|
||||
|
||||
assert await EvoRouteFallbackMiddleware([fallback]).awrap_model_call(
|
||||
request, handler
|
||||
) == "ok"
|
||||
assert seen == [("primary", tools), ("fallback", [])]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_fallback_uses_evo_frozen_fallback_model():
|
||||
primary = _Model("primary")
|
||||
fallback = _Model("fallback")
|
||||
middleware = EvoRouteFallbackMiddleware([fallback], _Health())
|
||||
seen: list[str] = []
|
||||
|
||||
async def handler(request: ModelRequest):
|
||||
route = request.model.metadata["route_key"]
|
||||
seen.append(route)
|
||||
if route == "primary":
|
||||
raise ConnectionError("upstream unavailable")
|
||||
return "fallback-response"
|
||||
|
||||
result = await middleware.awrap_model_call(_request(primary), handler)
|
||||
|
||||
assert result == "fallback-response"
|
||||
assert seen == ["primary", "fallback"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_protocol_error_retries_with_a_bounded_repair_instruction():
|
||||
primary = _Model("primary", supports_tools=True)
|
||||
request = ModelRequest(
|
||||
model=primary,
|
||||
messages=[],
|
||||
tools=[{"type": "function", "function": {"name": "search"}}],
|
||||
)
|
||||
seen: list[ModelRequest] = []
|
||||
|
||||
async def handler(routed: ModelRequest):
|
||||
seen.append(routed)
|
||||
if len(seen) == 1:
|
||||
raise ModelToolProtocolError("unknown_name")
|
||||
return "repaired-response"
|
||||
|
||||
result = await EvoRouteFallbackMiddleware([primary]).awrap_model_call(
|
||||
request, handler
|
||||
)
|
||||
|
||||
assert result == "repaired-response"
|
||||
assert len(seen) == 2
|
||||
assert seen[0].messages == []
|
||||
assert isinstance(seen[1].messages[0], SystemMessage)
|
||||
assert "unknown_name" in str(seen[1].messages[0].content)
|
||||
assert "search" in str(seen[1].messages[0].content)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_other_agent_control_errors_remain_non_fallbackable():
|
||||
primary = _Model("primary")
|
||||
fallback = _Model("fallback")
|
||||
seen: list[str] = []
|
||||
|
||||
async def handler(routed: ModelRequest):
|
||||
seen.append(routed.model.metadata["route_key"])
|
||||
raise AgentControlError("MODEL_REQUEST_REJECTED", "rejected")
|
||||
|
||||
with pytest.raises(AgentControlError):
|
||||
await EvoRouteFallbackMiddleware([fallback]).awrap_model_call(
|
||||
_request(primary), handler
|
||||
)
|
||||
|
||||
assert seen == ["primary"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_open_primary_is_skipped_and_control_errors_do_not_fallback():
|
||||
primary = _Model("primary")
|
||||
fallback = _Model("fallback")
|
||||
middleware = EvoRouteFallbackMiddleware([fallback], _Health({"primary"}))
|
||||
seen: list[str] = []
|
||||
|
||||
async def handler(request: ModelRequest):
|
||||
seen.append(request.model.metadata["route_key"])
|
||||
return "fallback-response"
|
||||
|
||||
assert await middleware.awrap_model_call(_request(primary), handler) == "fallback-response"
|
||||
assert seen == ["fallback"]
|
||||
|
||||
async def controlled(_request: ModelRequest):
|
||||
raise EvoRuntimeError("ADMISSION_EXHAUSTED")
|
||||
|
||||
with pytest.raises(EvoRuntimeError, match="ADMISSION_EXHAUSTED"):
|
||||
await EvoRouteFallbackMiddleware([fallback]).awrap_model_call(
|
||||
_request(primary),
|
||||
controlled,
|
||||
)
|
||||
|
||||
|
||||
def test_sync_open_primary_is_skipped():
|
||||
primary = _Model("primary")
|
||||
fallback = _Model("fallback")
|
||||
middleware = EvoRouteFallbackMiddleware([fallback], _Health({"primary"}))
|
||||
seen: list[str] = []
|
||||
|
||||
def handler(request: ModelRequest):
|
||||
seen.append(request.model.metadata["route_key"])
|
||||
return "fallback-response"
|
||||
|
||||
assert middleware.wrap_model_call(_request(primary), handler) == "fallback-response"
|
||||
assert seen == ["fallback"]
|
||||
@@ -127,6 +127,7 @@ def test_launch_background_run_submits_run_and_invokes_hooks(monkeypatch):
|
||||
assert handle is not None
|
||||
assert handle.thread_id == "thread-1"
|
||||
assert handle.run_id == "run-1"
|
||||
assert handle.configurable == {"thread_id": "thread-1"}
|
||||
assert payload_calls == ["thread-1"]
|
||||
assert before_calls == ["thread-1"]
|
||||
assert started == [handle]
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from EvoScientist.llm.errors import AgentControlError
|
||||
from EvoScientist.middleware.model_fallback import _is_non_fallbackable
|
||||
from EvoScientist.middleware.recoverable_metering import _metering_config, _source_type
|
||||
|
||||
|
||||
def test_agent_control_error_is_non_fallbackable():
|
||||
@@ -11,3 +12,24 @@ def test_agent_control_error_is_non_fallbackable():
|
||||
|
||||
assert "platform control error" in (_is_non_fallbackable(error) or "")
|
||||
assert error.model_dump()["code"] == "INSUFFICIENT_BALANCE"
|
||||
|
||||
|
||||
def test_recoverable_metering_reads_explicit_evomemory_scope():
|
||||
metering = _metering_config(
|
||||
{
|
||||
"configurable": {
|
||||
"ai4sci_metering": {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "run-parent",
|
||||
"envelope_signature": "signed-parent",
|
||||
"source_type": "evomemory_linker",
|
||||
}
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
assert metering is not None
|
||||
assert metering["source_type"] == "evomemory_linker"
|
||||
assert _source_type({"metering_scope": "evomemory_subagent_worker"}, []) == (
|
||||
"evomemory_subagent_worker"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.configuration import (
|
||||
EndpointConfig,
|
||||
ModelConfig,
|
||||
ProviderConfig,
|
||||
)
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError
|
||||
from EvoScientist.llm.invocation import (
|
||||
compile_invocation_plan,
|
||||
derive_runtime_invocation,
|
||||
derive_tool_call_transport,
|
||||
)
|
||||
from EvoScientist.llm.model_config import (
|
||||
EndpointConfig as LegacyEndpointConfig,
|
||||
)
|
||||
from EvoScientist.llm.model_config import ModelConfig as LegacyModelConfig
|
||||
from EvoScientist.llm.model_config import ProviderConfig as LegacyProviderConfig
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("capabilities", "expected"),
|
||||
[
|
||||
({"text": True, "tools": True}, "native"),
|
||||
({"text": True, "tools": False}, "disabled"),
|
||||
({"text": True}, "disabled"),
|
||||
],
|
||||
)
|
||||
def test_tool_transport_is_derived_only_from_model_capabilities(
|
||||
capabilities, expected
|
||||
):
|
||||
assert derive_tool_call_transport(capabilities) == expected
|
||||
|
||||
|
||||
def test_runtime_invocation_combines_model_api_mode_and_derived_transport():
|
||||
assert derive_runtime_invocation(
|
||||
"chat_completions", {"text": True, "tools": False}
|
||||
) == {
|
||||
"api_mode": "chat_completions",
|
||||
"tool_call_transport": "disabled",
|
||||
}
|
||||
|
||||
|
||||
def test_k3_chat_plan_freezes_adapter_compiled_parameters():
|
||||
plan = compile_invocation_plan(
|
||||
api_mode="chat_completions",
|
||||
declared_tool_call_transport="disabled",
|
||||
supports_tools=False,
|
||||
purpose="main_agent",
|
||||
output_token_limit=65_000,
|
||||
reasoning_effort="high",
|
||||
runtime_provider="openai",
|
||||
sdk_params={
|
||||
"max_completion_tokens": 65_000,
|
||||
"reasoning_effort": "high",
|
||||
"use_responses_api": False,
|
||||
},
|
||||
)
|
||||
|
||||
assert plan.output_token_parameter == "max_completion_tokens"
|
||||
assert plan.tool_call_transport == "disabled"
|
||||
assert plan.streaming is True
|
||||
assert plan.sdk_params["streaming"] is True
|
||||
with pytest.raises(TypeError):
|
||||
plan.sdk_params["max_completion_tokens"] = 1
|
||||
|
||||
|
||||
def test_responses_plan_accepts_native_tools_when_capability_is_enabled():
|
||||
plan = compile_invocation_plan(
|
||||
api_mode="responses",
|
||||
declared_tool_call_transport="native",
|
||||
supports_tools=True,
|
||||
purpose="title",
|
||||
output_token_limit=1_024,
|
||||
reasoning_effort="disabled",
|
||||
runtime_provider="openai",
|
||||
sdk_params={"max_output_tokens": 1_024, "use_responses_api": True},
|
||||
)
|
||||
|
||||
assert plan.tool_call_transport == "native"
|
||||
assert plan.streaming is False
|
||||
|
||||
|
||||
def test_plan_rejects_runtime_projection_that_disagrees_with_capabilities():
|
||||
with pytest.raises(EvoRuntimeError, match="MODEL_ADAPTER_COMPILE_FAILED"):
|
||||
compile_invocation_plan(
|
||||
api_mode="chat_completions",
|
||||
declared_tool_call_transport="native",
|
||||
supports_tools=False,
|
||||
purpose="main_agent",
|
||||
output_token_limit=1_024,
|
||||
reasoning_effort="disabled",
|
||||
runtime_provider="openai",
|
||||
sdk_params={"max_tokens": 1_024, "use_responses_api": False},
|
||||
)
|
||||
|
||||
|
||||
def test_provider_and_model_contracts_have_separate_ownership():
|
||||
model_fields = ModelConfig.__dataclass_fields__
|
||||
endpoint_fields = EndpointConfig.__dataclass_fields__
|
||||
provider_fields = ProviderConfig.__dataclass_fields__
|
||||
|
||||
assert "capabilities" in model_fields
|
||||
assert "max_output_tokens" in model_fields
|
||||
assert "base_url" not in model_fields
|
||||
assert "auth" not in model_fields
|
||||
assert "tool_call_transport" not in model_fields
|
||||
assert "base_url" in endpoint_fields
|
||||
assert "auth" in endpoint_fields
|
||||
assert "adapter_id" in provider_fields
|
||||
assert "connection_defaults" in provider_fields
|
||||
|
||||
|
||||
def test_legacy_model_config_imports_reexport_canonical_contracts():
|
||||
assert LegacyEndpointConfig is EndpointConfig
|
||||
assert LegacyModelConfig is ModelConfig
|
||||
assert LegacyProviderConfig is ProviderConfig
|
||||
@@ -6,7 +6,9 @@ langgraph dev.
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from EvoScientist.config import EvoScientistConfig
|
||||
@@ -161,3 +163,55 @@ def test_ollama_discovery_skipped_when_base_url_absent():
|
||||
{"name": n, "model_id": m, "provider": p}
|
||||
for n, m, p in list_models_by_provider()
|
||||
]
|
||||
|
||||
|
||||
def test_recoverable_capabilities_include_resume_and_workspace(monkeypatch):
|
||||
monkeypatch.setenv("EVOSCIENTIST_DEPLOY_MODE", "full")
|
||||
response = client.get("/api/ai4sci/recoverable-runs/capabilities")
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["interrupt_resume"] is True
|
||||
assert body["pending_interrupt_state"] is True
|
||||
assert body["workspace_scope_v1"] is True
|
||||
|
||||
|
||||
def test_recoverable_capabilities_fail_closed_without_full_deploy(monkeypatch):
|
||||
monkeypatch.delenv("EVOSCIENTIST_DEPLOY_MODE", raising=False)
|
||||
response = client.get("/api/ai4sci/recoverable-runs/capabilities")
|
||||
assert response.status_code == 200
|
||||
assert response.json()["workspace_scope_v1"] is False
|
||||
|
||||
|
||||
def test_recoverable_resume_rejects_human_input(monkeypatch):
|
||||
run_id = str(uuid4())
|
||||
response = client.post(
|
||||
"/api/ai4sci/recoverable-runs/create",
|
||||
headers={"x-auth-scheme": "langsmith"},
|
||||
json={
|
||||
"operation": "resume",
|
||||
"thread_id": str(uuid4()),
|
||||
"run_id": run_id,
|
||||
"run_request_id": run_id,
|
||||
"request_hash": "a" * 64,
|
||||
"assistant_id": "EvoScientist",
|
||||
"input": {"messages": [{"role": "user", "content": "continue"}]},
|
||||
"command": {"resume": {"decisions": [{"type": "approve"}]}},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert response.json()["code"] == "INVALID_RESUME_REQUEST"
|
||||
|
||||
|
||||
def test_workspace_scope_routes_require_service_token():
|
||||
response = client.post(
|
||||
"/internal/workspace-scopes/provision",
|
||||
json={"thread_id": str(uuid4())},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_workspace_materialize_rejects_path_escape():
|
||||
from EvoScientist.langgraph_dev.http import _materialize_target
|
||||
|
||||
with pytest.raises(ValueError, match="uploads"):
|
||||
_materialize_target(str(uuid4()), "../secret.txt")
|
||||
|
||||
+104
-132
@@ -1,6 +1,5 @@
|
||||
"""Tests for EvoScientist LLM module."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
@@ -161,67 +160,24 @@ class TestGetModelInfo:
|
||||
|
||||
|
||||
class TestGetChatModel:
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_uses_host_model_resolver(self, mock_init):
|
||||
"""An embedding host can provide model routing without a core dependency."""
|
||||
from EvoScientist.runtime_integrations import (
|
||||
configure_runtime_integrations,
|
||||
reset_runtime_integrations,
|
||||
)
|
||||
|
||||
mock_init.return_value = "mock_model"
|
||||
resolved = SimpleNamespace(
|
||||
provider_name="relay-a",
|
||||
model_id="model-a",
|
||||
protocol="openai",
|
||||
api_key="sk-host",
|
||||
base_url="https://relay.example/v1/",
|
||||
params={"max_tokens": 8192, "_default_headers": {"X-Relay": "a"}},
|
||||
supports_reasoning=False,
|
||||
)
|
||||
configure_runtime_integrations(model_resolver=lambda model, provider: resolved)
|
||||
try:
|
||||
assert get_chat_model("alias-a") == "mock_model"
|
||||
finally:
|
||||
reset_runtime_integrations()
|
||||
|
||||
call_kwargs = mock_init.call_args.kwargs
|
||||
assert call_kwargs["model"] == "model-a"
|
||||
assert call_kwargs["model_provider"] == "openai"
|
||||
assert call_kwargs["api_key"] == "sk-host"
|
||||
assert call_kwargs["base_url"] == "https://relay.example/v1"
|
||||
assert call_kwargs["max_tokens"] == 8192
|
||||
assert call_kwargs["default_headers"] == {"X-Relay": "a"}
|
||||
assert "reasoning" not in call_kwargs
|
||||
|
||||
@patch("EvoScientist.llm.models._patch_openai_compat_content")
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_host_openai_provider_with_custom_base_uses_compat_patch(
|
||||
self, mock_init, mock_compat
|
||||
):
|
||||
from EvoScientist.runtime_integrations import (
|
||||
configure_runtime_integrations,
|
||||
reset_runtime_integrations,
|
||||
)
|
||||
|
||||
def test_openai_custom_base_url_uses_compat_patch(self, mock_init, mock_compat):
|
||||
model_instance = object()
|
||||
mock_init.return_value = model_instance
|
||||
resolved = SimpleNamespace(
|
||||
provider_name="openai",
|
||||
model_id="gpt-5.5",
|
||||
protocol="openai",
|
||||
|
||||
get_chat_model(
|
||||
"gpt-5.5",
|
||||
provider="openai",
|
||||
api_key="sk-host",
|
||||
base_url="https://relay.example/v1",
|
||||
params={},
|
||||
supports_reasoning=True,
|
||||
)
|
||||
configure_runtime_integrations(model_resolver=lambda model, provider: resolved)
|
||||
try:
|
||||
get_chat_model("gpt-5.5", provider="openai")
|
||||
finally:
|
||||
reset_runtime_integrations()
|
||||
|
||||
mock_compat.assert_called_once_with(model_instance, hoist_tool_media=True)
|
||||
mock_compat.assert_called_once_with(
|
||||
model_instance,
|
||||
hoist_tool_media=True,
|
||||
drop_reasoning_metadata=True,
|
||||
)
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_uses_default_model_when_none(self, mock_init):
|
||||
@@ -498,9 +454,6 @@ class TestThirdPartyRouting:
|
||||
"""OpenRouter should use native 'openrouter' provider via init_chat_model."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key-456")
|
||||
# Assert the DEFAULT effort, so isolate from any leaked env override.
|
||||
monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False)
|
||||
|
||||
get_chat_model("x-ai/grok-4.3", provider="openrouter")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
@@ -524,8 +477,10 @@ class TestThirdPartyRouting:
|
||||
assert call_kwargs["reasoning"] == {"effort": "low"}
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_openrouter_reasoning_effort_from_env(self, mock_init, monkeypatch):
|
||||
"""Reasoning effort should be configurable via env var."""
|
||||
def test_openrouter_reasoning_effort_environment_is_ignored(
|
||||
self, mock_init, monkeypatch
|
||||
):
|
||||
"""The deployment environment cannot alter an invocation parameter."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
|
||||
monkeypatch.setenv("EVOSCIENTIST_REASONING_EFFORT", "medium")
|
||||
@@ -533,7 +488,7 @@ class TestThirdPartyRouting:
|
||||
get_chat_model("x-ai/grok-4.3", provider="openrouter")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["reasoning"] == {"effort": "medium", "summary": "auto"}
|
||||
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
|
||||
|
||||
# --- OpenRouter app attribution (issue #339) ---
|
||||
|
||||
@@ -634,7 +589,6 @@ class TestThirdPartyRouting:
|
||||
"""Attribution must not disturb reasoning or Anthropic prompt caching."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
|
||||
monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False)
|
||||
monkeypatch.delenv(
|
||||
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE", raising=False
|
||||
)
|
||||
@@ -1233,6 +1187,57 @@ class TestFlattenMessageContent:
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestOpenAIEmptySSEKeepalivePatch:
|
||||
def test_blank_sse_keepalive_is_skipped(self):
|
||||
from EvoScientist.llm.patches import _is_blank_sse_keepalive
|
||||
|
||||
class Event:
|
||||
def __init__(self, data):
|
||||
self.data = data
|
||||
|
||||
assert _is_blank_sse_keepalive(Event(""))
|
||||
assert _is_blank_sse_keepalive(Event(" \t\n"))
|
||||
assert _is_blank_sse_keepalive(Event(None))
|
||||
assert not _is_blank_sse_keepalive(Event('{"type":"response.output_text"}'))
|
||||
assert not _is_blank_sse_keepalive(Event("[DONE]"))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_stream_filters_blank_keepalive_before_json_parse(self):
|
||||
from openai._streaming import AsyncStream, ServerSentEvent
|
||||
|
||||
class Decoder:
|
||||
async def aiter_bytes(self, _bytes):
|
||||
yield ServerSentEvent(data="")
|
||||
yield ServerSentEvent(data='{"type":"response.created"}')
|
||||
|
||||
class Response:
|
||||
async def aiter_bytes(self):
|
||||
if False:
|
||||
yield b""
|
||||
|
||||
stream = type("Stream", (), {"_decoder": Decoder(), "response": Response()})()
|
||||
events = [event async for event in AsyncStream._iter_events(stream)]
|
||||
|
||||
assert [event.data for event in events] == ['{"type":"response.created"}']
|
||||
|
||||
def test_sync_stream_filters_blank_keepalive_before_json_parse(self):
|
||||
from openai._streaming import ServerSentEvent, Stream
|
||||
|
||||
class Decoder:
|
||||
def iter_bytes(self, _bytes):
|
||||
yield ServerSentEvent(data="")
|
||||
yield ServerSentEvent(data='{"type":"response.created"}')
|
||||
|
||||
class Response:
|
||||
def iter_bytes(self):
|
||||
return iter(())
|
||||
|
||||
stream = type("Stream", (), {"_decoder": Decoder(), "response": Response()})()
|
||||
events = list(Stream._iter_events(stream))
|
||||
|
||||
assert [event.data for event in events] == ['{"type":"response.created"}']
|
||||
|
||||
|
||||
class TestPatchOpenAICompatContent:
|
||||
"""Verify content flattening covers _generate, _agenerate, _stream, _astream."""
|
||||
|
||||
@@ -1305,6 +1310,7 @@ class TestPatchOpenAICompatContent:
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "wrong-id",
|
||||
"call_id": "wrong-call-id",
|
||||
"name": "wrong-name",
|
||||
"args": {},
|
||||
}
|
||||
@@ -1316,6 +1322,7 @@ class TestPatchOpenAICompatContent:
|
||||
)
|
||||
|
||||
assert normalized[0].content[0]["id"] == "call-1"
|
||||
assert normalized[0].content[0]["call_id"] == "call-1"
|
||||
assert normalized[0].content[0]["name"] == "execute"
|
||||
|
||||
def test_invalid_tool_call_is_not_replayed_to_responses_api(self):
|
||||
@@ -1394,6 +1401,25 @@ class TestPatchOpenAICompatContent:
|
||||
assert [message.type for message in normalized] == ["ai", "human"]
|
||||
assert normalized[0].tool_calls == []
|
||||
|
||||
def test_nonportable_reasoning_metadata_is_removed_for_cross_model_replay(self):
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
from EvoScientist.llm.patches import _sanitize_messages
|
||||
|
||||
message = AIMessage(
|
||||
content="portable answer",
|
||||
additional_kwargs={
|
||||
"reasoning_content": "provider-specific trace",
|
||||
"reasoning_details": [{"type": "reasoning"}],
|
||||
"safe_field": "preserved",
|
||||
},
|
||||
)
|
||||
|
||||
normalized = _sanitize_messages([message], drop_reasoning_metadata=True)
|
||||
|
||||
assert normalized[0].additional_kwargs == {"safe_field": "preserved"}
|
||||
assert "reasoning_content" in message.additional_kwargs
|
||||
|
||||
def test_generate_flattened(self):
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
@@ -2675,11 +2701,6 @@ class TestPatchOpenrouterStripResponsesReasoning:
|
||||
|
||||
|
||||
class TestAutoConfig:
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_reasoning_effort_env(self, monkeypatch):
|
||||
"""Keep auto-config tests independent of the developer environment."""
|
||||
monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False)
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_internal_sentinels_disable_auto_reasoning(self, mock_init, monkeypatch):
|
||||
"""Internal callers can disable reasoning without leaking sentinels."""
|
||||
@@ -2796,8 +2817,6 @@ class TestAutoConfig:
|
||||
"""gpt-5.4+ and codex models get xhigh reasoning."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False)
|
||||
|
||||
get_chat_model("gpt-5.4", provider="openai")
|
||||
assert mock_init.call_args[1]["reasoning"] == {
|
||||
"effort": "xhigh",
|
||||
@@ -2823,8 +2842,10 @@ class TestAutoConfig:
|
||||
}
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_openai_reasoning_effort_from_env(self, mock_init, monkeypatch):
|
||||
"""Native OpenAI reasoning effort should be configurable via env var."""
|
||||
def test_openai_reasoning_effort_environment_is_ignored(
|
||||
self, mock_init, monkeypatch
|
||||
):
|
||||
"""The deployment environment cannot alter an invocation parameter."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
monkeypatch.setenv("EVOSCIENTIST_REASONING_EFFORT", "high")
|
||||
@@ -2832,7 +2853,7 @@ class TestAutoConfig:
|
||||
get_chat_model("gpt-5.5", provider="openai")
|
||||
|
||||
assert mock_init.call_args[1]["reasoning"] == {
|
||||
"effort": "high",
|
||||
"effort": "xhigh",
|
||||
"summary": "auto",
|
||||
}
|
||||
|
||||
@@ -2867,10 +2888,10 @@ class TestAutoConfig:
|
||||
assert call_kwargs["model_provider"] == "openai"
|
||||
assert call_kwargs["base_url"] == "http://127.0.0.1:8000/codex/v1"
|
||||
assert call_kwargs["api_key"] == "ccproxy-oauth"
|
||||
# ccproxy uses the Responses API, so reasoning configuration is valid.
|
||||
# Endpoint detection may add compatible client headers, but cannot
|
||||
# select an API protocol. The compiled invocation plan owns that.
|
||||
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
|
||||
# Proxy mode: Responses API (bypasses format chain), streaming ON
|
||||
assert call_kwargs["use_responses_api"] is True
|
||||
assert "use_responses_api" not in call_kwargs
|
||||
assert "streaming" not in call_kwargs
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
@@ -3054,72 +3075,23 @@ class TestAutoConfig:
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["include_thoughts"] is True
|
||||
|
||||
@pytest.mark.parametrize("env_value", ["false", "true", " TRUE "])
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_use_responses_api_false(self, mock_init, monkeypatch):
|
||||
"""use_responses_api=false forces Chat Completions and drops reasoning."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
monkeypatch.setenv("EVOSCIENTIST_USE_RESPONSES_API", "false")
|
||||
|
||||
get_chat_model("gpt-5-nano", provider="openai")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["use_responses_api"] is False
|
||||
assert "reasoning" not in call_kwargs
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_use_responses_api_true(self, mock_init, monkeypatch):
|
||||
"""use_responses_api=true explicitly enables the Responses API."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
monkeypatch.setenv("EVOSCIENTIST_USE_RESPONSES_API", "true")
|
||||
|
||||
get_chat_model("gpt-5-nano", provider="openai")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["use_responses_api"] is True
|
||||
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_use_responses_api_default_unchanged(self, mock_init, monkeypatch):
|
||||
"""Empty use_responses_api preserves default behavior (no kwarg set)."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
monkeypatch.delenv("EVOSCIENTIST_USE_RESPONSES_API", raising=False)
|
||||
|
||||
get_chat_model("gpt-5-nano", provider="openai")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert "use_responses_api" not in call_kwargs
|
||||
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
|
||||
|
||||
@pytest.mark.parametrize("env_value", ["FALSE", " false ", "False"])
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_use_responses_api_false_normalization(
|
||||
def test_response_api_environment_cannot_override_an_explicit_call_plan(
|
||||
self, mock_init, monkeypatch, env_value
|
||||
):
|
||||
"""Case/whitespace variants of 'false' are normalized correctly."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
monkeypatch.setenv("EVOSCIENTIST_USE_RESPONSES_API", env_value)
|
||||
|
||||
get_chat_model("gpt-5-nano", provider="openai")
|
||||
get_chat_model(
|
||||
"gpt-5-nano",
|
||||
provider="openai",
|
||||
use_responses_api=False,
|
||||
reasoning_effort="high",
|
||||
)
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["use_responses_api"] is False
|
||||
assert call_kwargs["reasoning_effort"] == "high"
|
||||
assert "reasoning" not in call_kwargs
|
||||
|
||||
@pytest.mark.parametrize("env_value", ["TRUE", " true ", "True"])
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
def test_use_responses_api_true_normalization(
|
||||
self, mock_init, monkeypatch, env_value
|
||||
):
|
||||
"""Case/whitespace variants of 'true' are normalized correctly."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
monkeypatch.setenv("EVOSCIENTIST_USE_RESPONSES_API", env_value)
|
||||
|
||||
get_chat_model("gpt-5-nano", provider="openai")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["use_responses_api"] is True
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
"""Tests for bounded skill context in background memory-agent graphs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def test_memory_agent_factory_replaces_framework_skill_injection(monkeypatch, tmp_path):
|
||||
from EvoScientist.memory.agents import _factory
|
||||
from EvoScientist.middleware import BudgetedSkillsMiddleware
|
||||
|
||||
captured = {}
|
||||
|
||||
class _Agent:
|
||||
def with_config(self, _config):
|
||||
return self
|
||||
|
||||
def _create_deep_agent(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return _Agent()
|
||||
|
||||
monkeypatch.setattr("deepagents.create_deep_agent", _create_deep_agent)
|
||||
monkeypatch.setattr(
|
||||
"EvoScientist.EvoScientist._ensure_auxiliary_chat_model", lambda: object()
|
||||
)
|
||||
|
||||
_factory.build_memory_agent_graph(
|
||||
name="test-memory-agent",
|
||||
system_prompt="Test prompt",
|
||||
memory_dir=tmp_path / "memory",
|
||||
workspace_dir=tmp_path / "workspace",
|
||||
tools=[],
|
||||
middleware=[],
|
||||
skills=["/skills/"],
|
||||
backend=object(),
|
||||
)
|
||||
|
||||
assert captured["skills"] is None
|
||||
assert isinstance(captured["middleware"][-1], BudgetedSkillsMiddleware)
|
||||
@@ -0,0 +1,460 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
from EvoScientist.llm.config_admin import EvoModelConfigAdminService
|
||||
from EvoScientist.llm.contracts import (
|
||||
CommitModelConfigRequest,
|
||||
EvoRuntimeError,
|
||||
HmacGrantAuthority,
|
||||
ProbeCandidateRouteRequest,
|
||||
ValidateCandidateConfigRequest,
|
||||
)
|
||||
from EvoScientist.llm.crypto import canonical_json_v1, sha256_id
|
||||
from EvoScientist.llm.model_config import EvoModelConfig, FileEvoModelConfigStore
|
||||
from tests.v3_fixtures import (
|
||||
RUNTIME_KEY_ID,
|
||||
RUNTIME_SECRET,
|
||||
identity_ring,
|
||||
v3_payload,
|
||||
)
|
||||
|
||||
|
||||
def _grant(authority, *, action, operation_id, payload):
|
||||
return authority.sign_admin(
|
||||
subject_id="admin",
|
||||
action=action,
|
||||
operation_id=operation_id,
|
||||
request_digest=sha256_id(payload),
|
||||
ttl_ms=60_000,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_probe_commit_bootstraps_v2_config(monkeypatch, tmp_path):
|
||||
monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret")
|
||||
authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID)
|
||||
store = FileEvoModelConfigStore(
|
||||
tmp_path / "model_routes.yaml", admin_verifier=authority
|
||||
)
|
||||
service = EvoModelConfigAdminService(
|
||||
store,
|
||||
grant_authority=authority,
|
||||
identity_key_ring=identity_ring(),
|
||||
probe_runner=lambda *_args: True,
|
||||
)
|
||||
candidate = v3_payload()
|
||||
candidate.pop("capability_evidence")
|
||||
validate_operation = str(uuid.uuid4())
|
||||
validate_payload = {
|
||||
"operation_id": validate_operation,
|
||||
"expected_revision": 0,
|
||||
"payload": candidate,
|
||||
}
|
||||
result = service.validate_candidate(
|
||||
ValidateCandidateConfigRequest(
|
||||
validate_operation,
|
||||
0,
|
||||
candidate,
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:validate",
|
||||
operation_id=validate_operation,
|
||||
payload=validate_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
evidence_ids = []
|
||||
for route in result.concrete_routes:
|
||||
for probe_kind in route.required_probe_kinds:
|
||||
operation_id = str(uuid.uuid4())
|
||||
payload = {
|
||||
"operation_id": operation_id,
|
||||
"proposal_hash": result.proposal_hash,
|
||||
"route_semantics_hash": route.route_semantics_hash,
|
||||
"probe_kind": probe_kind,
|
||||
}
|
||||
probe = await service.probe_candidate(
|
||||
ProbeCandidateRouteRequest(
|
||||
operation_id,
|
||||
result.proposal_hash,
|
||||
route.route_semantics_hash,
|
||||
probe_kind,
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:probe",
|
||||
operation_id=operation_id,
|
||||
payload=payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
evidence_ids.append(probe.evidence_id)
|
||||
commit_operation = str(uuid.uuid4())
|
||||
commit_payload = {
|
||||
"operation_id": commit_operation,
|
||||
"expected_revision": 0,
|
||||
"proposal_hash": result.proposal_hash,
|
||||
"payload": candidate,
|
||||
"evidence_ids": sorted(set(evidence_ids)),
|
||||
}
|
||||
committed = service.commit(
|
||||
CommitModelConfigRequest(
|
||||
commit_operation,
|
||||
0,
|
||||
result.proposal_hash,
|
||||
candidate,
|
||||
tuple(commit_payload["evidence_ids"]),
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:commit",
|
||||
operation_id=commit_operation,
|
||||
payload=commit_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
assert committed.config_revision == 1
|
||||
assert store.load().config_revision == 1
|
||||
|
||||
|
||||
def test_v2_parser_rejects_old_alias_and_manual_capability_boundary():
|
||||
payload = v3_payload()
|
||||
payload["providers"]["custom-openai"]["models"][0]["alias"] = "legacy"
|
||||
with pytest.raises(EvoRuntimeError, match="unknown fields"):
|
||||
EvoModelConfig.parse(payload)
|
||||
|
||||
payload = v3_payload()
|
||||
payload["capability_requirements"] = []
|
||||
with pytest.raises(EvoRuntimeError, match="unknown fields"):
|
||||
EvoModelConfig.parse(payload)
|
||||
|
||||
|
||||
def test_model_output_capability_rejects_oversized_output_token_limit():
|
||||
payload = v3_payload()
|
||||
payload["providers"]["custom-openai"]["models"][0]["params"] = {
|
||||
"output_token_limit": 4096
|
||||
}
|
||||
|
||||
with pytest.raises(
|
||||
EvoRuntimeError, match="output_token_limit exceeds model capability"
|
||||
):
|
||||
EvoModelConfig.parse(payload)
|
||||
|
||||
|
||||
def test_admin_grant_is_bound_to_action_and_payload(monkeypatch, tmp_path):
|
||||
monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret")
|
||||
authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID)
|
||||
service = EvoModelConfigAdminService(
|
||||
FileEvoModelConfigStore(
|
||||
tmp_path / "model_routes.yaml", admin_verifier=authority
|
||||
),
|
||||
grant_authority=authority,
|
||||
identity_key_ring=identity_ring(),
|
||||
)
|
||||
candidate = v3_payload()
|
||||
candidate.pop("capability_evidence")
|
||||
operation_id = str(uuid.uuid4())
|
||||
wrong_payload = {
|
||||
"operation_id": operation_id,
|
||||
"expected_revision": 1,
|
||||
"payload": candidate,
|
||||
}
|
||||
with pytest.raises(EvoRuntimeError, match="ADMIN_CONFIG_FORBIDDEN"):
|
||||
service.validate_candidate(
|
||||
ValidateCandidateConfigRequest(
|
||||
operation_id,
|
||||
0,
|
||||
candidate,
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:validate",
|
||||
operation_id=operation_id,
|
||||
payload=wrong_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_validate_operation_replays_and_rejects_changed_payload(monkeypatch, tmp_path):
|
||||
monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret")
|
||||
authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID)
|
||||
service = EvoModelConfigAdminService(
|
||||
FileEvoModelConfigStore(tmp_path / "model_routes.yaml"),
|
||||
grant_authority=authority,
|
||||
identity_key_ring=identity_ring(),
|
||||
)
|
||||
candidate = v3_payload()
|
||||
candidate.pop("capability_evidence")
|
||||
operation_id = str(uuid.uuid4())
|
||||
payload = {
|
||||
"operation_id": operation_id,
|
||||
"expected_revision": 0,
|
||||
"payload": candidate,
|
||||
}
|
||||
request = ValidateCandidateConfigRequest(
|
||||
operation_id,
|
||||
0,
|
||||
candidate,
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:validate",
|
||||
operation_id=operation_id,
|
||||
payload=payload,
|
||||
),
|
||||
)
|
||||
first = service.validate_candidate(request)
|
||||
replay = service.validate_candidate(request)
|
||||
assert replay == first
|
||||
|
||||
changed = {**candidate, "runtime_defaults": {"max_retries": 1}}
|
||||
changed_payload = {**payload, "payload": changed}
|
||||
with pytest.raises(EvoRuntimeError, match="IDEMPOTENCY_CONFLICT"):
|
||||
service.validate_candidate(
|
||||
ValidateCandidateConfigRequest(
|
||||
operation_id,
|
||||
0,
|
||||
changed,
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:validate",
|
||||
operation_id=operation_id,
|
||||
payload=changed_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_provider_defaults_accept_values_up_to_v3_bounds():
|
||||
from EvoScientist.llm.model_config import _parse_v3_provider_defaults
|
||||
|
||||
runtime = {
|
||||
"connect_timeout_seconds": 10,
|
||||
"first_event_timeout_seconds": 60,
|
||||
"stream_idle_timeout_seconds": 60,
|
||||
"attempt_timeout_seconds": 600,
|
||||
}
|
||||
defaults = _parse_v3_provider_defaults(
|
||||
{"connect_timeout_seconds": 60, "attempt_timeout_seconds": 600}, runtime
|
||||
)
|
||||
assert defaults["connect_timeout_seconds"] == 60
|
||||
assert defaults["attempt_timeout_seconds"] == 600
|
||||
|
||||
|
||||
def test_provider_defaults_reject_values_beyond_v3_bounds():
|
||||
from EvoScientist.llm.model_config import _parse_v3_provider_defaults
|
||||
|
||||
runtime = {
|
||||
"connect_timeout_seconds": 10,
|
||||
"first_event_timeout_seconds": 60,
|
||||
"stream_idle_timeout_seconds": 60,
|
||||
"attempt_timeout_seconds": 600,
|
||||
}
|
||||
with pytest.raises(EvoRuntimeError):
|
||||
_parse_v3_provider_defaults({"connect_timeout_seconds": 61}, runtime)
|
||||
|
||||
|
||||
def test_store_recovers_renamed_target_from_preparing_journal(tmp_path):
|
||||
path = tmp_path / "model_routes.yaml"
|
||||
ops_path = tmp_path / "model_config_ops.sqlite"
|
||||
store = FileEvoModelConfigStore(path, ops_path=ops_path)
|
||||
store.bootstrap_for_development(v3_payload(), operation_id="bootstrap")
|
||||
|
||||
target = v3_payload(revision=2)
|
||||
payload_hash = sha256_id(target)
|
||||
path.write_text(
|
||||
yaml.safe_dump(target, allow_unicode=True, sort_keys=False),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with sqlite3.connect(ops_path) as connection:
|
||||
connection.execute(
|
||||
"""INSERT INTO config_operations
|
||||
(subject_id, action, operation_id, request_digest, status,
|
||||
expected_revision, target_revision, payload_hash,
|
||||
canonical_payload, created_at, updated_at)
|
||||
VALUES ('admin', 'model_config:commit', 'recover-op', ?,
|
||||
'PREPARING', 1, 2, ?, ?, 1, 1)""",
|
||||
(
|
||||
payload_hash,
|
||||
payload_hash,
|
||||
canonical_json_v1(target).decode("utf-8"),
|
||||
),
|
||||
)
|
||||
|
||||
recovered = FileEvoModelConfigStore(path, ops_path=ops_path)
|
||||
assert recovered.load().config_revision == 2
|
||||
with sqlite3.connect(ops_path) as connection:
|
||||
status = connection.execute(
|
||||
"SELECT status FROM config_operations WHERE operation_id='recover-op'"
|
||||
).fetchone()[0]
|
||||
assert status == "COMMITTED"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_probe_dispatch_is_durable_and_operation_replay_is_free(
|
||||
monkeypatch, tmp_path
|
||||
):
|
||||
monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret")
|
||||
authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID)
|
||||
sink_events = []
|
||||
provider_calls = 0
|
||||
|
||||
async def sink(payload):
|
||||
sink_events.append(dict(payload))
|
||||
if payload["outcome"] == "started":
|
||||
return "committed"
|
||||
return f"terminal:{payload['probe_result']}"
|
||||
|
||||
async def runner(*_args):
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
return True
|
||||
|
||||
service = EvoModelConfigAdminService(
|
||||
FileEvoModelConfigStore(tmp_path / "model_routes.yaml"),
|
||||
grant_authority=authority,
|
||||
identity_key_ring=identity_ring(),
|
||||
probe_runner=runner,
|
||||
probe_event_sink=sink,
|
||||
)
|
||||
candidate = v3_payload()
|
||||
candidate.pop("capability_evidence")
|
||||
validate_id = str(uuid.uuid4())
|
||||
validate_payload = {
|
||||
"operation_id": validate_id,
|
||||
"expected_revision": 0,
|
||||
"payload": candidate,
|
||||
}
|
||||
validated = service.validate_candidate(
|
||||
ValidateCandidateConfigRequest(
|
||||
validate_id,
|
||||
0,
|
||||
candidate,
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:validate",
|
||||
operation_id=validate_id,
|
||||
payload=validate_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
route = validated.concrete_routes[0]
|
||||
operation_id = str(uuid.uuid4())
|
||||
payload = {
|
||||
"operation_id": operation_id,
|
||||
"proposal_hash": validated.proposal_hash,
|
||||
"route_semantics_hash": route.route_semantics_hash,
|
||||
"probe_kind": route.required_probe_kinds[0],
|
||||
}
|
||||
request = ProbeCandidateRouteRequest(
|
||||
operation_id,
|
||||
validated.proposal_hash,
|
||||
route.route_semantics_hash,
|
||||
route.required_probe_kinds[0],
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:probe",
|
||||
operation_id=operation_id,
|
||||
payload=payload,
|
||||
),
|
||||
)
|
||||
first = await service.probe_candidate(request)
|
||||
replay = await service.probe_candidate(request)
|
||||
|
||||
assert replay == first
|
||||
assert provider_calls == 1
|
||||
assert [event["outcome"] for event in sink_events] == [
|
||||
"started",
|
||||
"usage_unconfirmed",
|
||||
]
|
||||
assert all(event["billing_intent"] == "platform_cost" for event in sink_events)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_rejects_probe_evidence_after_runtime_secret_changes(
|
||||
monkeypatch, tmp_path
|
||||
):
|
||||
monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "secret-before-probe")
|
||||
authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID)
|
||||
store = FileEvoModelConfigStore(tmp_path / "model_routes.yaml")
|
||||
service = EvoModelConfigAdminService(
|
||||
store,
|
||||
grant_authority=authority,
|
||||
identity_key_ring=identity_ring(),
|
||||
probe_runner=lambda *_args: True,
|
||||
)
|
||||
candidate = v3_payload()
|
||||
candidate.pop("capability_evidence")
|
||||
validate_id = str(uuid.uuid4())
|
||||
validate_payload = {
|
||||
"operation_id": validate_id,
|
||||
"expected_revision": 0,
|
||||
"payload": candidate,
|
||||
}
|
||||
validated = service.validate_candidate(
|
||||
ValidateCandidateConfigRequest(
|
||||
validate_id,
|
||||
0,
|
||||
candidate,
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:validate",
|
||||
operation_id=validate_id,
|
||||
payload=validate_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
evidence_ids = []
|
||||
for route in validated.concrete_routes:
|
||||
for probe_kind in route.required_probe_kinds:
|
||||
operation_id = str(uuid.uuid4())
|
||||
probe_payload = {
|
||||
"operation_id": operation_id,
|
||||
"proposal_hash": validated.proposal_hash,
|
||||
"route_semantics_hash": route.route_semantics_hash,
|
||||
"probe_kind": probe_kind,
|
||||
}
|
||||
result = await service.probe_candidate(
|
||||
ProbeCandidateRouteRequest(
|
||||
operation_id,
|
||||
validated.proposal_hash,
|
||||
route.route_semantics_hash,
|
||||
probe_kind,
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:probe",
|
||||
operation_id=operation_id,
|
||||
payload=probe_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
evidence_ids.append(result.evidence_id)
|
||||
|
||||
monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "secret-after-probe")
|
||||
commit_id = str(uuid.uuid4())
|
||||
commit_payload = {
|
||||
"operation_id": commit_id,
|
||||
"expected_revision": 0,
|
||||
"proposal_hash": validated.proposal_hash,
|
||||
"payload": candidate,
|
||||
"evidence_ids": sorted(evidence_ids),
|
||||
}
|
||||
with pytest.raises(EvoRuntimeError, match="CAPABILITY_EVIDENCE_STALE"):
|
||||
service.commit(
|
||||
CommitModelConfigRequest(
|
||||
commit_id,
|
||||
0,
|
||||
validated.proposal_hash,
|
||||
candidate,
|
||||
tuple(commit_payload["evidence_ids"]),
|
||||
_grant(
|
||||
authority,
|
||||
action="model_config:commit",
|
||||
operation_id=commit_id,
|
||||
payload=commit_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,822 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError, HmacGrantAuthority
|
||||
from EvoScientist.llm.crypto import HmacKeyRing, KeyMaterial, sha256_id
|
||||
from EvoScientist.llm.model_config_v4 import (
|
||||
EvoModelConfig,
|
||||
UnifiedModelConfigStore,
|
||||
build_supported_v3_evidence,
|
||||
model_profile_identity_map,
|
||||
normalize_v4_config,
|
||||
project_v4_to_v3,
|
||||
)
|
||||
from EvoScientist.llm.runtime import EvoModelRuntime
|
||||
|
||||
|
||||
def _payload(*, provider_id: str | None = None, profile_id: str | None = None) -> dict:
|
||||
provider_id = provider_id or str(uuid.uuid4())
|
||||
profile_id = profile_id or str(uuid.uuid4())
|
||||
return {
|
||||
"schema_version": 4,
|
||||
"default_model_profile_id": profile_id,
|
||||
"purpose_defaults": {},
|
||||
"purpose_call_limits": {},
|
||||
"providers": [
|
||||
{
|
||||
"provider_id": provider_id,
|
||||
"display_name": "OpenAI Production",
|
||||
"adapter_id": "openai",
|
||||
"adapter_revision": "openai-v1",
|
||||
"enabled": True,
|
||||
"connection": {"base_url": "https://api.openai.com/v1"},
|
||||
"models": [
|
||||
{
|
||||
"model_profile_id": profile_id,
|
||||
"provider_model_id": "gpt-test",
|
||||
"display_name": "General",
|
||||
"enabled": True,
|
||||
"version_policy": "rolling",
|
||||
"invocation": {
|
||||
"api_mode": "responses",
|
||||
"tool_call_transport": "native",
|
||||
},
|
||||
"capabilities": {"text": True, "tools": True},
|
||||
"limits": {
|
||||
"context_tokens": 128_000,
|
||||
"max_output_tokens": 8_192,
|
||||
},
|
||||
"parameters": {},
|
||||
"billing": {
|
||||
"sku": "internal/gpt-test",
|
||||
"pricing_revision": "configured-v1",
|
||||
"currency": "CNY",
|
||||
"unit_scale": 1_000_000,
|
||||
"input_microunits_per_million": 1_000_000,
|
||||
"cached_microunits_per_million": 200_000,
|
||||
"output_microunits_per_million": 4_000_000,
|
||||
"multiplier": 1,
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _ring() -> HmacKeyRing:
|
||||
return HmacKeyRing(KeyMaterial.create("identity-v1", "i" * 32))
|
||||
|
||||
|
||||
def test_v4_normalization_uses_stable_profile_identity():
|
||||
payload = normalize_v4_config(_payload())
|
||||
identities = model_profile_identity_map(payload)
|
||||
|
||||
assert payload["schema_version"] == 4
|
||||
assert payload["default_model_profile_id"] in identities
|
||||
assert identities[payload["default_model_profile_id"]][1] == "gpt-test"
|
||||
|
||||
|
||||
def test_v4_rejects_enabled_model_without_billing_configuration():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0].pop("billing")
|
||||
|
||||
with pytest.raises(EvoRuntimeError) as raised:
|
||||
normalize_v4_config(raw, lenient=True, warnings=[])
|
||||
|
||||
assert raised.value.details[0]["code"] == "CONFIG_PRICING_REQUIRED"
|
||||
|
||||
|
||||
def test_v4_accepts_explicit_zero_price_for_free_model():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0]["billing"].update(
|
||||
input_microunits_per_million=0,
|
||||
cached_microunits_per_million=0,
|
||||
output_microunits_per_million=0,
|
||||
)
|
||||
|
||||
payload = normalize_v4_config(raw, lenient=True, warnings=[])
|
||||
|
||||
assert payload["providers"][0]["models"][0]["billing"][
|
||||
"input_microunits_per_million"
|
||||
] == 0
|
||||
|
||||
|
||||
def test_v4_normalization_preserves_model_output_token_limit():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0]["parameters"] = {
|
||||
"defaults": {"output_token_limit": 2048, "temperature": 0.2},
|
||||
"purpose_overrides": {
|
||||
"main_agent": {"output_token_limit": 1024, "temperature": 0.3}
|
||||
},
|
||||
"user_options": {
|
||||
"output_token_limit": {
|
||||
"default": 512,
|
||||
"applies_to": ["main_agent"],
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
payload = normalize_v4_config(raw)
|
||||
parameters = payload["providers"][0]["models"][0]["parameters"]
|
||||
|
||||
assert parameters["defaults"] == {"output_token_limit": 2048, "temperature": 0.2}
|
||||
assert parameters["purpose_overrides"]["main_agent"] == {
|
||||
"output_token_limit": 1024,
|
||||
"temperature": 0.3,
|
||||
}
|
||||
assert parameters["user_options"] == {}
|
||||
for limit in payload["purpose_call_limits"].values():
|
||||
assert "max_output_tokens" not in limit
|
||||
assert payload["purpose_call_limits"]["main_agent"] == {
|
||||
"max_attempts_per_run": 4
|
||||
}
|
||||
|
||||
|
||||
def test_v4_rejects_output_token_limit_above_model_capability():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0]["parameters"] = {
|
||||
"defaults": {"output_token_limit": 16_384}
|
||||
}
|
||||
|
||||
with pytest.raises(EvoRuntimeError) as raised:
|
||||
normalize_v4_config(raw)
|
||||
|
||||
assert raised.value.code == "LLM_ROUTE_CONFIGURATION_REQUIRED"
|
||||
assert raised.value.details[0]["code"] == "CONFIG_LIMIT_EXCEEDED"
|
||||
|
||||
|
||||
def test_v4_lenient_clamps_output_token_limit_to_model_capability():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0]["parameters"] = {
|
||||
"defaults": {"output_token_limit": 16_384}
|
||||
}
|
||||
warnings: list[dict] = []
|
||||
|
||||
payload = normalize_v4_config(raw, lenient=True, warnings=warnings)
|
||||
|
||||
defaults = payload["providers"][0]["models"][0]["parameters"]["defaults"]
|
||||
assert defaults["output_token_limit"] == 8_192
|
||||
assert any(item["code"] == "CONFIG_LIMIT_EXCEEDED" for item in warnings)
|
||||
|
||||
|
||||
def test_v4_removes_legacy_tool_transport_from_admin_config():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0]["capabilities"]["tools"] = False
|
||||
|
||||
payload = normalize_v4_config(raw)
|
||||
|
||||
model = payload["providers"][0]["models"][0]
|
||||
assert model["invocation"] == {"api_mode": "responses"}
|
||||
|
||||
|
||||
def test_v4_runtime_projection_derives_tool_transport_from_capabilities():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0]["capabilities"]["tools"] = False
|
||||
warnings: list[dict] = []
|
||||
|
||||
payload = normalize_v4_config(raw, lenient=True, warnings=warnings)
|
||||
|
||||
model = payload["providers"][0]["models"][0]
|
||||
assert "tool_call_transport" not in model["invocation"]
|
||||
assert not any(
|
||||
item["code"] == "CONFIG_INVOCATION_CONTRACT_INVALID" for item in warnings
|
||||
)
|
||||
|
||||
projection = project_v4_to_v3(
|
||||
payload,
|
||||
revision=1,
|
||||
identity_key_id=_ring().current.key_id,
|
||||
bindings={payload["providers"][0]["provider_id"]: 1},
|
||||
)
|
||||
config = EvoModelConfig.parse(projection, require_evidence=False)
|
||||
route = config.concrete_routes(next(iter(config.route_selectors)))[0]
|
||||
assert route.tool_call_transport == "disabled"
|
||||
|
||||
|
||||
def test_v4_rejects_non_uuid_business_ids():
|
||||
raw = _payload(provider_id="openai-main")
|
||||
|
||||
with pytest.raises(EvoRuntimeError) as raised:
|
||||
normalize_v4_config(raw)
|
||||
|
||||
assert raised.value.code == "LLM_ROUTE_CONFIGURATION_REQUIRED"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"base_url",
|
||||
[
|
||||
"",
|
||||
"not a url",
|
||||
"ftp://user:pass@example.test/path?key=value#fragment",
|
||||
"http://127.0.0.1:11434/v1",
|
||||
"https://169.254.169.254/latest/meta-data",
|
||||
"custom://model.internal:70000/path",
|
||||
],
|
||||
)
|
||||
def test_v4_base_url_is_not_validated(base_url):
|
||||
raw = _payload()
|
||||
raw["providers"][0]["connection"]["base_url"] = base_url
|
||||
|
||||
payload = normalize_v4_config(raw)
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
projection = project_v4_to_v3(
|
||||
payload,
|
||||
revision=1,
|
||||
identity_key_id=_ring().current.key_id,
|
||||
bindings={provider_id: 1},
|
||||
)
|
||||
|
||||
assert payload["providers"][0]["connection"]["base_url"] == base_url
|
||||
assert projection["providers"][0]["connection"]["base_url"] == base_url
|
||||
|
||||
|
||||
def test_unified_store_commits_config_and_secret_without_evidence_gate(tmp_path):
|
||||
ring = _ring()
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
payload = normalize_v4_config(_payload())
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
profile_id = payload["default_model_profile_id"]
|
||||
plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"})
|
||||
projection = project_v4_to_v3(
|
||||
payload,
|
||||
revision=1,
|
||||
identity_key_id=ring.current.key_id,
|
||||
bindings=plan.bindings,
|
||||
)
|
||||
evidence = build_supported_v3_evidence(
|
||||
projection,
|
||||
identity_key_ring=ring,
|
||||
verified_profiles={provider_id: frozenset({profile_id})},
|
||||
)
|
||||
request_hash = sha256_id({"payload": payload, "fingerprints": plan.fingerprints})
|
||||
store.begin_operation(
|
||||
"save-1",
|
||||
request_hash=request_hash,
|
||||
expected_revision=0,
|
||||
actor="admin",
|
||||
lease_owner="worker-1",
|
||||
)
|
||||
result = store.commit(
|
||||
payload,
|
||||
expected_revision=0,
|
||||
operation_id="save-1",
|
||||
actor="admin",
|
||||
request_hash=request_hash,
|
||||
credential_plan=plan,
|
||||
evidence=evidence,
|
||||
)
|
||||
|
||||
assert result.config_revision == 1
|
||||
assert result.changed is True
|
||||
runtime = store.load()
|
||||
assert runtime.main_routes.default_alias == profile_id
|
||||
assert store.resolve(
|
||||
runtime.providers[provider_id].endpoints[provider_id].auth
|
||||
).value == "sk-secret-value"
|
||||
revision, admin_payload, credentials = store.get_admin_config()
|
||||
assert revision == 1
|
||||
assert admin_payload == payload
|
||||
assert credentials[provider_id].configured is True
|
||||
assert "sk-secret-value" not in str(credentials)
|
||||
with sqlite3.connect(store.path) as connection:
|
||||
assert connection.execute("SELECT COUNT(*) FROM model_config_revisions").fetchone()[0] == 1
|
||||
assert connection.execute("SELECT COUNT(*) FROM provider_secret_versions").fetchone()[0] == 1
|
||||
assert connection.execute("SELECT COUNT(*) FROM capability_evidence").fetchone()[0] == 0
|
||||
store.record_runtime_observation(
|
||||
config_revision=1,
|
||||
provider_id=provider_id,
|
||||
model_profile_id=profile_id,
|
||||
provider_model_id="gpt-test",
|
||||
api_mode="responses",
|
||||
purpose="main_agent",
|
||||
outcome="succeeded",
|
||||
error_code=None,
|
||||
strategy={"structured_output": False},
|
||||
)
|
||||
summary = store.runtime_observation_summary()
|
||||
assert len(summary) == 1
|
||||
assert summary[0]["model_profile_id"] == profile_id
|
||||
assert summary[0]["provider_id"] == provider_id
|
||||
assert summary[0]["latest_outcome"] == "succeeded"
|
||||
assert summary[0]["latest_error_code"] is None
|
||||
assert summary[0]["total_calls"] == 1
|
||||
assert summary[0]["successful_calls"] == 1
|
||||
assert summary[0]["strategy"] == {"structured_output": False}
|
||||
assert summary[0]["latest_at"]
|
||||
|
||||
|
||||
def test_unified_store_idempotency_rejects_same_id_different_request(tmp_path):
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=_ring(),
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
store.begin_operation(
|
||||
"save-1",
|
||||
request_hash="hmac:one",
|
||||
expected_revision=0,
|
||||
actor="admin",
|
||||
lease_owner="worker-1",
|
||||
)
|
||||
|
||||
with pytest.raises(EvoRuntimeError) as raised:
|
||||
store.begin_operation(
|
||||
"save-1",
|
||||
request_hash="hmac:two",
|
||||
expected_revision=0,
|
||||
actor="admin",
|
||||
lease_owner="worker-2",
|
||||
)
|
||||
|
||||
assert raised.value.code == "IDEMPOTENCY_CONFLICT"
|
||||
|
||||
|
||||
def test_operation_failure_preserves_last_progress(tmp_path):
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=_ring(),
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
store.begin_operation(
|
||||
"save-1", request_hash="request-1", expected_revision=0,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
store.set_operation_status(
|
||||
"save-1", "PROBING", progress={"completed": 2, "total": 3}
|
||||
)
|
||||
store.set_operation_status(
|
||||
"save-1",
|
||||
"FAILED_RETRYABLE",
|
||||
stage="PROBING",
|
||||
error_code="MODEL_PROBE_FAILED",
|
||||
error_details=({"path": "models.test", "code": "MODEL_TIMEOUT"},),
|
||||
)
|
||||
|
||||
operation = store.get_operation("save-1")
|
||||
assert operation is not None
|
||||
assert operation["progress"] == {"completed": 2, "total": 3}
|
||||
assert operation["error_details"] == [
|
||||
{"path": "models.test", "code": "MODEL_TIMEOUT"}
|
||||
]
|
||||
|
||||
|
||||
|
||||
def test_unified_store_load_ignores_historical_evidence(tmp_path):
|
||||
ring = _ring()
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
payload = normalize_v4_config(_payload())
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
profile_id = payload["default_model_profile_id"]
|
||||
plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"})
|
||||
projection = project_v4_to_v3(
|
||||
payload, revision=1, identity_key_id=ring.current.key_id, bindings=plan.bindings
|
||||
)
|
||||
evidence = build_supported_v3_evidence(
|
||||
projection,
|
||||
identity_key_ring=ring,
|
||||
verified_profiles={provider_id: frozenset({profile_id})},
|
||||
)
|
||||
store.begin_operation(
|
||||
"save-1", request_hash="request-1", expected_revision=0,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
store.commit(
|
||||
payload, expected_revision=0, operation_id="save-1", actor="admin",
|
||||
request_hash="request-1", credential_plan=plan, evidence=evidence,
|
||||
)
|
||||
# The supplied evidence is intentionally not persisted or read by load.
|
||||
# A legacy/tampered evidence row therefore cannot make the revision stale.
|
||||
assert store.load_revision(1).config_revision == 1
|
||||
|
||||
|
||||
def test_rollback_publishes_a_new_monotonic_revision(tmp_path):
|
||||
ring = _ring()
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
first = normalize_v4_config(_payload())
|
||||
provider_id = first["providers"][0]["provider_id"]
|
||||
profile_id = first["default_model_profile_id"]
|
||||
def publish(payload: dict, revision: int, operation_id: str, request_hash: str):
|
||||
plan = store.plan_credentials(
|
||||
payload, {provider_id: "sk-secret-value"} if revision == 1 else {}
|
||||
)
|
||||
projection = project_v4_to_v3(
|
||||
payload,
|
||||
revision=revision,
|
||||
identity_key_id=ring.current.key_id,
|
||||
bindings=plan.bindings,
|
||||
)
|
||||
evidence = build_supported_v3_evidence(
|
||||
projection,
|
||||
identity_key_ring=ring,
|
||||
verified_profiles={provider_id: frozenset({profile_id})},
|
||||
)
|
||||
store.begin_operation(
|
||||
operation_id, request_hash=request_hash, expected_revision=revision - 1,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
return store.commit(
|
||||
payload, expected_revision=revision - 1, operation_id=operation_id,
|
||||
actor="admin", request_hash=request_hash, credential_plan=plan,
|
||||
evidence=evidence,
|
||||
)
|
||||
|
||||
publish(first, 1, "save-1", "request-1")
|
||||
second = json.loads(json.dumps(first))
|
||||
second["providers"][0]["display_name"] = "Changed display name"
|
||||
publish(second, 2, "save-2", "request-2")
|
||||
store.begin_operation(
|
||||
"rollback-1", request_hash="rollback-request", expected_revision=2,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
result = store.rollback(
|
||||
target_revision=1,
|
||||
expected_revision=2,
|
||||
operation_id="rollback-1",
|
||||
actor="admin",
|
||||
request_hash="rollback-request",
|
||||
)
|
||||
|
||||
assert result.config_revision == 3
|
||||
assert store.current_revision() == 3
|
||||
assert store.get_admin_config()[1]["providers"][0]["display_name"] == "OpenAI Production"
|
||||
assert store.load().config_revision == 3
|
||||
|
||||
|
||||
def test_identity_key_rotation_does_not_rotate_unchanged_credential(tmp_path):
|
||||
first_ring = _ring()
|
||||
path = tmp_path / "model-config.sqlite"
|
||||
store = UnifiedModelConfigStore(
|
||||
path,
|
||||
identity_key_ring=first_ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
payload = normalize_v4_config(_payload())
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
profile_id = payload["default_model_profile_id"]
|
||||
plan = store.plan_credentials(payload, {provider_id: "sk-same-secret"})
|
||||
projection = project_v4_to_v3(
|
||||
payload, revision=1, identity_key_id=first_ring.current.key_id,
|
||||
bindings=plan.bindings,
|
||||
)
|
||||
evidence = build_supported_v3_evidence(
|
||||
projection,
|
||||
identity_key_ring=first_ring,
|
||||
verified_profiles={provider_id: frozenset({profile_id})},
|
||||
)
|
||||
store.begin_operation(
|
||||
"save-1", request_hash="request-1", expected_revision=0,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
store.commit(
|
||||
payload, expected_revision=0, operation_id="save-1", actor="admin",
|
||||
request_hash="request-1", credential_plan=plan, evidence=evidence,
|
||||
)
|
||||
rotated_ring = HmacKeyRing(
|
||||
KeyMaterial.create("identity-v2", "j" * 32), retained=(first_ring.current,)
|
||||
)
|
||||
rotated_store = UnifiedModelConfigStore(
|
||||
path,
|
||||
identity_key_ring=rotated_ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
|
||||
rotated_plan = rotated_store.plan_credentials(
|
||||
payload, {provider_id: "sk-same-secret"}
|
||||
)
|
||||
|
||||
assert rotated_plan.new_versions == frozenset()
|
||||
assert rotated_plan.bindings[provider_id] == 1
|
||||
assert rotated_plan.fingerprints[provider_id] == plan.fingerprints[provider_id]
|
||||
|
||||
|
||||
def test_evidence_requires_observed_results_for_every_declared_capability():
|
||||
ring = _ring()
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0]["capabilities"] = {
|
||||
"text": True,
|
||||
"vision": True,
|
||||
"tools": True,
|
||||
}
|
||||
payload = normalize_v4_config(raw)
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
profile_id = payload["default_model_profile_id"]
|
||||
projection = project_v4_to_v3(
|
||||
payload,
|
||||
revision=1,
|
||||
identity_key_id=ring.current.key_id,
|
||||
bindings={provider_id: 1},
|
||||
)
|
||||
|
||||
with pytest.raises(EvoRuntimeError) as raised:
|
||||
build_supported_v3_evidence(
|
||||
projection,
|
||||
identity_key_ring=ring,
|
||||
probe_results={
|
||||
provider_id: {profile_id: {"connectivity": "supported"}}
|
||||
},
|
||||
)
|
||||
|
||||
assert raised.value.code == "CAPABILITY_EVIDENCE_STALE"
|
||||
|
||||
|
||||
def test_historical_evidence_is_not_reused_for_route_admission(tmp_path):
|
||||
ring = _ring()
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
payload = normalize_v4_config(_payload())
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
profile_id = payload["default_model_profile_id"]
|
||||
plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"})
|
||||
projection = project_v4_to_v3(
|
||||
payload,
|
||||
revision=1,
|
||||
identity_key_id=ring.current.key_id,
|
||||
bindings=plan.bindings,
|
||||
)
|
||||
evidence = build_supported_v3_evidence(
|
||||
projection,
|
||||
identity_key_ring=ring,
|
||||
verified_profiles={provider_id: frozenset({profile_id})},
|
||||
)
|
||||
store.begin_operation(
|
||||
"save-1", request_hash="request-1", expected_revision=0,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
store.commit(
|
||||
payload, expected_revision=0, operation_id="save-1", actor="admin",
|
||||
request_hash="request-1", credential_plan=plan, evidence=evidence,
|
||||
)
|
||||
|
||||
next_payload = json.loads(json.dumps(payload))
|
||||
next_payload["providers"][0]["display_name"] = "Renamed Provider"
|
||||
next_plan = store.plan_credentials(next_payload, {})
|
||||
next_projection = project_v4_to_v3(
|
||||
next_payload,
|
||||
revision=2,
|
||||
identity_key_id=ring.current.key_id,
|
||||
bindings=next_plan.bindings,
|
||||
)
|
||||
observed, windows, missing = store.reusable_probe_results(
|
||||
next_projection,
|
||||
fingerprints=next_plan.fingerprints,
|
||||
versions=next_plan.bindings,
|
||||
)
|
||||
|
||||
assert observed == {}
|
||||
assert windows == {}
|
||||
assert missing == frozenset({(provider_id, profile_id)})
|
||||
|
||||
|
||||
def test_availability_reflects_enabled_state(tmp_path):
|
||||
ring = _ring()
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
raw = _payload()
|
||||
disabled_profile_id = str(uuid.uuid4())
|
||||
disabled = json.loads(json.dumps(raw["providers"][0]["models"][0]))
|
||||
disabled.update(
|
||||
{
|
||||
"model_profile_id": disabled_profile_id,
|
||||
"provider_model_id": "disabled-model",
|
||||
"display_name": "Disabled",
|
||||
"enabled": False,
|
||||
}
|
||||
)
|
||||
raw["providers"][0]["models"].append(disabled)
|
||||
payload = normalize_v4_config(raw)
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
profile_id = payload["default_model_profile_id"]
|
||||
plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"})
|
||||
store.begin_operation(
|
||||
"save-1", request_hash="request-1", expected_revision=0,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
store.commit(
|
||||
payload, expected_revision=0, operation_id="save-1", actor="admin",
|
||||
request_hash="request-1", credential_plan=plan, evidence=[],
|
||||
)
|
||||
|
||||
availability = store.get_model_availability()
|
||||
|
||||
assert availability[profile_id] == {
|
||||
"model_profile_id": profile_id,
|
||||
"enabled": True,
|
||||
"selectable": True,
|
||||
}
|
||||
assert availability[disabled_profile_id] == {
|
||||
"model_profile_id": disabled_profile_id,
|
||||
"enabled": False,
|
||||
"selectable": False,
|
||||
}
|
||||
|
||||
|
||||
def test_schema_upgrade_adds_availability_and_operation_columns(tmp_path):
|
||||
path = tmp_path / "model-config.sqlite"
|
||||
UnifiedModelConfigStore(
|
||||
path,
|
||||
identity_key_ring=_ring(),
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
with sqlite3.connect(path) as connection:
|
||||
connection.execute("DROP INDEX idx_capability_evidence_semantics")
|
||||
connection.execute("ALTER TABLE active_model_config DROP COLUMN evidence_epoch")
|
||||
connection.execute("ALTER TABLE capability_evidence DROP COLUMN credential_version")
|
||||
connection.execute("ALTER TABLE model_config_operations DROP COLUMN stage")
|
||||
connection.execute("ALTER TABLE model_config_operations DROP COLUMN progress_json")
|
||||
connection.execute("ALTER TABLE model_config_operations DROP COLUMN error_details_json")
|
||||
|
||||
UnifiedModelConfigStore(
|
||||
path,
|
||||
identity_key_ring=_ring(),
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
|
||||
with sqlite3.connect(path) as connection:
|
||||
assert "evidence_epoch" in {
|
||||
row[1] for row in connection.execute("PRAGMA table_info(active_model_config)")
|
||||
}
|
||||
assert "credential_version" in {
|
||||
row[1] for row in connection.execute("PRAGMA table_info(capability_evidence)")
|
||||
}
|
||||
operation_columns = {
|
||||
row[1] for row in connection.execute("PRAGMA table_info(model_config_operations)")
|
||||
}
|
||||
assert {"stage", "progress_json", "error_details_json"} <= operation_columns
|
||||
|
||||
|
||||
def test_v3_runtime_accepts_expired_historical_evidence(tmp_path):
|
||||
ring = _ring()
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
payload = normalize_v4_config(_payload())
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
profile_id = payload["default_model_profile_id"]
|
||||
plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"})
|
||||
projection = project_v4_to_v3(
|
||||
payload, revision=1, identity_key_id=ring.current.key_id,
|
||||
bindings=plan.bindings,
|
||||
)
|
||||
now = datetime.now(UTC)
|
||||
evidence = build_supported_v3_evidence(
|
||||
projection,
|
||||
identity_key_ring=ring,
|
||||
verified_profiles={provider_id: frozenset({profile_id})},
|
||||
evidence_windows={
|
||||
provider_id: {
|
||||
profile_id: (
|
||||
(now - timedelta(hours=2)).isoformat(timespec="microseconds"),
|
||||
(now - timedelta(hours=1)).isoformat(timespec="microseconds"),
|
||||
)
|
||||
}
|
||||
},
|
||||
)
|
||||
store.begin_operation(
|
||||
"save-1", request_hash="request-1", expected_revision=0,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
store.commit(
|
||||
payload, expected_revision=0, operation_id="save-1", actor="admin",
|
||||
request_hash="request-1", credential_plan=plan, evidence=evidence,
|
||||
)
|
||||
config = store.load()
|
||||
authority = HmacGrantAuthority("r" * 32)
|
||||
runtime = EvoModelRuntime(
|
||||
store,
|
||||
admission_verifier=authority,
|
||||
quote_authority=authority,
|
||||
identity_key_ring=ring,
|
||||
secret_resolver=store.resolve,
|
||||
)
|
||||
selector = config.resolve_main_selector(None)
|
||||
route = config.concrete_routes(selector.selector_id)[0]
|
||||
|
||||
resolved = runtime._resolve_route(config, route, "main_agent")
|
||||
|
||||
assert resolved.identity.model_id == "gpt-test"
|
||||
|
||||
|
||||
def test_lenient_normalization_ignores_unknown_fields_with_warning():
|
||||
raw = _payload()
|
||||
raw["unexpected_top"] = True
|
||||
raw["providers"][0]["unexpected_provider"] = "x"
|
||||
raw["providers"][0]["models"][0]["unexpected_model"] = 1
|
||||
warnings: list[dict] = []
|
||||
|
||||
payload = normalize_v4_config(raw, lenient=True, warnings=warnings)
|
||||
|
||||
assert payload["schema_version"] == 4
|
||||
assert [item["code"] for item in warnings] == ["CONFIG_UNKNOWN_FIELDS"] * 3
|
||||
|
||||
|
||||
def test_lenient_normalization_clamps_out_of_range_and_timeout_conflict():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["connection"] = {
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
"connect_timeout_seconds": 120,
|
||||
"attempt_timeout_seconds": 30,
|
||||
}
|
||||
warnings: list[dict] = []
|
||||
|
||||
payload = normalize_v4_config(raw, lenient=True, warnings=warnings)
|
||||
connection = payload["providers"][0]["connection"]
|
||||
|
||||
assert connection["connect_timeout_seconds"] == 60
|
||||
assert connection["attempt_timeout_seconds"] == 60
|
||||
codes = {item["code"] for item in warnings}
|
||||
assert "CONFIG_VALUE_OUT_OF_RANGE" in codes
|
||||
assert "CONFIG_VALUE_CONFLICT" in codes
|
||||
|
||||
|
||||
def test_lenient_normalization_forces_text_capability_and_keeps_unknown_purpose():
|
||||
raw = _payload()
|
||||
raw["providers"][0]["models"][0]["capabilities"] = {"text": False, "vision": True}
|
||||
raw["providers"][0]["models"][0]["parameters"] = {
|
||||
"purpose_overrides": {"custom_purpose": {"temperature": 0.5}}
|
||||
}
|
||||
warnings: list[dict] = []
|
||||
|
||||
payload = normalize_v4_config(raw, lenient=True, warnings=warnings)
|
||||
model = payload["providers"][0]["models"][0]
|
||||
|
||||
assert model["capabilities"]["text"] is True
|
||||
assert model["parameters"]["purpose_overrides"]["custom_purpose"] == {"temperature": 0.5}
|
||||
assert any(item["code"] == "CONFIG_REQUIRED" for item in warnings)
|
||||
assert any(item["code"] == "CONFIG_REFERENCE_INVALID" for item in warnings)
|
||||
|
||||
|
||||
def test_strict_normalization_still_rejects_unknown_fields():
|
||||
raw = _payload()
|
||||
raw["unexpected_top"] = True
|
||||
|
||||
with pytest.raises(EvoRuntimeError):
|
||||
normalize_v4_config(raw)
|
||||
|
||||
|
||||
def test_unverified_enabled_model_is_selectable(tmp_path):
|
||||
ring = _ring()
|
||||
store = UnifiedModelConfigStore(
|
||||
tmp_path / "model-config.sqlite",
|
||||
identity_key_ring=ring,
|
||||
encryption_keys={"enc-v1": "e" * 32},
|
||||
current_encryption_key_id="enc-v1",
|
||||
)
|
||||
payload = normalize_v4_config(_payload())
|
||||
provider_id = payload["providers"][0]["provider_id"]
|
||||
profile_id = payload["default_model_profile_id"]
|
||||
plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"})
|
||||
store.begin_operation(
|
||||
"save-1", request_hash="request-1", expected_revision=0,
|
||||
actor="admin", lease_owner="worker",
|
||||
)
|
||||
store.commit(
|
||||
payload, expected_revision=0, operation_id="save-1", actor="admin",
|
||||
request_hash="request-1", credential_plan=plan, evidence=[],
|
||||
)
|
||||
|
||||
availability = store.get_model_availability()[profile_id]
|
||||
|
||||
assert availability == {
|
||||
"model_profile_id": profile_id,
|
||||
"enabled": True,
|
||||
"selectable": True,
|
||||
}
|
||||
@@ -161,6 +161,26 @@ class TestTryFallbacks:
|
||||
invoke.assert_awaited_once()
|
||||
mock_gcm.assert_called_once_with(model="fb-model", provider="fb-provider")
|
||||
|
||||
async def test_provider_error_details_are_not_emitted(self):
|
||||
"""Fallback diagnostics must not expose provider bodies or credentials."""
|
||||
add_fallback("fb-model", "fb-provider")
|
||||
req = _fake_request()
|
||||
invoke = AsyncMock(return_value=AI_RESPONSE)
|
||||
emitted: list[str] = []
|
||||
set_ui_emit(lambda message, _style: emitted.append(message))
|
||||
|
||||
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
|
||||
mock_gcm.return_value = MagicMock()
|
||||
await _try_fallbacks(
|
||||
req,
|
||||
invoke,
|
||||
Exception("provider body contains sk-live-do-not-log"),
|
||||
)
|
||||
|
||||
output = "\n".join(emitted)
|
||||
assert "sk-live-do-not-log" not in output
|
||||
assert "Exception" in output
|
||||
|
||||
async def test_skips_failing_fallback_tries_next(self):
|
||||
"""When the first fallback fails, try the second."""
|
||||
add_fallback("fb-bad", "prov-a")
|
||||
@@ -270,7 +290,8 @@ class TestTryFallbacks:
|
||||
# Attribution flipped to moonshot (the failing fallback), not
|
||||
# openai (the original request's model).
|
||||
assert exc_info.value.provider == "moonshot"
|
||||
assert "quota exceeded" in exc_info.value.message
|
||||
assert exc_info.value.message == "Provider request failed."
|
||||
assert "quota exceeded" not in exc_info.value.message
|
||||
|
||||
async def test_langgraph_error_at_fallback_raise_point_passes_through(self):
|
||||
"""Regression: ``_raise_normalized`` calls ``_normalize``
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError
|
||||
from EvoScientist.llm.model_config import SecretReference
|
||||
from EvoScientist.llm.secret_store import EncryptedModelSecretStore
|
||||
|
||||
|
||||
def test_secret_store_versions_masks_and_resolves(tmp_path):
|
||||
store = EncryptedModelSecretStore(
|
||||
tmp_path / "secrets.sqlite",
|
||||
master_secret="test-master-secret-that-is-at-least-32-bytes",
|
||||
)
|
||||
|
||||
first = store.put("dashscope/primary", "sk-first-secret-value", created_by="admin")
|
||||
second = store.put(
|
||||
"dashscope/primary", "sk-second-secret-value", created_by="admin"
|
||||
)
|
||||
|
||||
assert first.version == 1
|
||||
assert second.version == 2
|
||||
assert "second-secret" not in second.masked_value
|
||||
assert second.ref == "secret://dashscope/primary#2"
|
||||
metadata = store.list_metadata()
|
||||
assert [item.version for item in metadata] == [2, 1]
|
||||
resolved = store.resolve(SecretReference(second.ref, second.version))
|
||||
assert resolved.value == "sk-second-secret-value"
|
||||
assert resolved.authoritative_version == "2"
|
||||
|
||||
|
||||
def test_secret_store_rejects_wrong_revision_and_master_key(tmp_path):
|
||||
path = tmp_path / "secrets.sqlite"
|
||||
store = EncryptedModelSecretStore(
|
||||
path,
|
||||
master_secret="test-master-secret-that-is-at-least-32-bytes",
|
||||
)
|
||||
item = store.put("openai/primary", "sk-secret", created_by="admin")
|
||||
|
||||
with pytest.raises(EvoRuntimeError, match="ROUTE_SECRET_UNAVAILABLE"):
|
||||
store.resolve(SecretReference(item.ref, item.version + 1))
|
||||
|
||||
wrong_key = EncryptedModelSecretStore(
|
||||
path,
|
||||
master_secret="different-master-secret-that-is-at-least-32-bytes",
|
||||
)
|
||||
with pytest.raises(EvoRuntimeError, match="ROUTE_SECRET_UNAVAILABLE"):
|
||||
wrong_key.resolve(SecretReference(item.ref, item.version))
|
||||
@@ -38,6 +38,7 @@ from EvoScientist.memory.observations import (
|
||||
MemorySourceType,
|
||||
MemoryType,
|
||||
ObservationSearchMode,
|
||||
archive_observation_file,
|
||||
create_link_observations_tool,
|
||||
create_read_memory_tool,
|
||||
create_search_observations_tool,
|
||||
@@ -145,6 +146,7 @@ def _memory_worker_run(
|
||||
source_agent: str = "EvoScientist",
|
||||
source_session_id: str = "thread-1",
|
||||
trajectory_digest: str = "digest-1",
|
||||
configurable: dict[str, object] | None = None,
|
||||
) -> background_runs.BackgroundRun:
|
||||
return background_runs.BackgroundRun(
|
||||
name="EvoMemory worker",
|
||||
@@ -160,6 +162,7 @@ def _memory_worker_run(
|
||||
"source_session_id": source_session_id,
|
||||
"trajectory_digest": trajectory_digest,
|
||||
},
|
||||
configurable=configurable,
|
||||
)
|
||||
|
||||
|
||||
@@ -345,6 +348,26 @@ def test_record_observation_file_writes_contract_and_dedupes(tmp_path):
|
||||
}
|
||||
|
||||
|
||||
def test_record_observation_file_serializes_concurrent_deduplication(tmp_path):
|
||||
memories = tmp_path / "memories"
|
||||
barrier = threading.Barrier(2)
|
||||
results: list[dict[str, Any]] = []
|
||||
|
||||
def record() -> None:
|
||||
barrier.wait()
|
||||
results.append(_record_test_observation(memories))
|
||||
|
||||
threads = [threading.Thread(target=record) for _ in range(2)]
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
|
||||
assert sorted(result["created"] for result in results) == [False, True]
|
||||
observation_path = memories / _memory_relative_path(results[0])
|
||||
assert read_observation_document(observation_path) is not None
|
||||
|
||||
|
||||
def test_link_observation_files_writes_frontmatter_and_dedupes(tmp_path):
|
||||
memories = tmp_path / "memories"
|
||||
first = record_observation_file(
|
||||
@@ -384,6 +407,7 @@ def test_link_observation_files_writes_frontmatter_and_dedupes(tmp_path):
|
||||
relation=ObservationRelation.COMPLEMENTS,
|
||||
reason="Both observations describe the durable background-memory flow.",
|
||||
)
|
||||
|
||||
duplicate = link_observation_files(
|
||||
memory_dir=memories,
|
||||
project_id="P-project",
|
||||
@@ -432,6 +456,59 @@ def test_link_observation_files_writes_frontmatter_and_dedupes(tmp_path):
|
||||
datetime.strptime(first_links[0]["linked_at"], "%Y-%m-%dT%H:%M:%SZ")
|
||||
|
||||
|
||||
def test_archive_observation_removes_live_file_and_related_links(tmp_path):
|
||||
memories = tmp_path / "memories"
|
||||
source = _record_test_observation(
|
||||
memories,
|
||||
summary="Global source memory.",
|
||||
observation="A global memory linked to one project observation.",
|
||||
)
|
||||
target = _record_test_observation(
|
||||
memories,
|
||||
summary="Project target memory.",
|
||||
observation="A project memory that can be archived safely.",
|
||||
scope=MemoryScope.PROJECT,
|
||||
)
|
||||
link_observation_files(
|
||||
memory_dir=memories,
|
||||
project_id="P-project",
|
||||
source_observation_id=source["observation_id"],
|
||||
target_observation_id=target["observation_id"],
|
||||
reason="Archive cleanup regression test.",
|
||||
)
|
||||
|
||||
result = archive_observation_file(
|
||||
memory_dir=memories,
|
||||
observation_id=target["observation_id"],
|
||||
observation_path=_memory_relative_path(target),
|
||||
)
|
||||
|
||||
target_path = memories / _memory_relative_path(target)
|
||||
source_document = read_observation_document(
|
||||
memories / _memory_relative_path(source)
|
||||
)
|
||||
assert result["removed"] is True
|
||||
assert not target_path.exists()
|
||||
assert (memories / str(result["archive_path"])).is_file()
|
||||
assert source_document is not None
|
||||
assert all(
|
||||
relation.id != target["observation_id"]
|
||||
for relation in source_document[0].related_observations
|
||||
)
|
||||
|
||||
|
||||
def test_archive_observation_rejects_path_mismatch(tmp_path):
|
||||
memories = tmp_path / "memories"
|
||||
observation = _record_test_observation(memories)
|
||||
|
||||
with pytest.raises(ValueError, match="invalid observation archive target"):
|
||||
archive_observation_file(
|
||||
memory_dir=memories,
|
||||
observation_id=observation["observation_id"],
|
||||
observation_path="../observations/global/other.md",
|
||||
)
|
||||
|
||||
|
||||
def test_link_observation_files_serializes_concurrent_frontmatter_updates(tmp_path):
|
||||
memories = tmp_path / "memories"
|
||||
source = record_observation_file(
|
||||
@@ -1333,6 +1410,43 @@ def test_search_observation_files_returns_no_low_confidence_fallback(tmp_path):
|
||||
assert hits == []
|
||||
|
||||
|
||||
def test_search_observation_files_ranks_chinese_queries(tmp_path):
|
||||
memories = tmp_path / "memories"
|
||||
relevant = record_observation_file(
|
||||
memory_dir=memories,
|
||||
project_id="P-project",
|
||||
memory_type=MemoryType.PROCEDURAL,
|
||||
summary="切换对话时保留用户级长期记忆",
|
||||
observation="用户画像跨对话共享,项目记忆跟随当前工作区。",
|
||||
why_it_matters="不同对话需要共享稳定偏好,但不能混合项目约束。",
|
||||
scope=MemoryScope.GLOBAL,
|
||||
source_type=MemorySourceType.TURN,
|
||||
source_session_id="thread-1",
|
||||
source_agent="EvoScientist",
|
||||
)
|
||||
record_observation_file(
|
||||
memory_dir=memories,
|
||||
project_id="P-project",
|
||||
memory_type=MemoryType.SEMANTIC,
|
||||
summary="模型计费快照",
|
||||
observation="模型调用使用不可变计费快照。",
|
||||
why_it_matters="结算需要稳定证据。",
|
||||
scope=MemoryScope.GLOBAL,
|
||||
source_type=MemorySourceType.TURN,
|
||||
source_session_id="thread-1",
|
||||
source_agent="EvoScientist",
|
||||
)
|
||||
|
||||
hits = search_observation_files(
|
||||
memory_dir=memories,
|
||||
project_id="P-project",
|
||||
query="跨对话共享记忆",
|
||||
)
|
||||
|
||||
assert hits
|
||||
assert hits[0]["observation_id"] == relevant["observation_id"]
|
||||
|
||||
|
||||
def test_record_observation_tool_can_use_worker_config_source(tmp_path):
|
||||
from EvoScientist.middleware.memory import create_memory_middleware
|
||||
|
||||
@@ -1825,6 +1939,54 @@ def test_memory_worker_run_payload_use_server_thread_id_and_source_metadata(
|
||||
}
|
||||
|
||||
|
||||
def test_memory_worker_inherits_signed_runtime_and_scopes_metering(monkeypatch):
|
||||
metering = {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "run-parent",
|
||||
"envelope_signature": "signed-parent",
|
||||
"provider_id": "provider-1",
|
||||
"model_id": "model-1",
|
||||
}
|
||||
proxy = {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "run-parent",
|
||||
"envelope_signature": "signed-parent",
|
||||
}
|
||||
monkeypatch.setattr(
|
||||
"langgraph.config.get_config",
|
||||
lambda: {
|
||||
"configurable": {
|
||||
"model": "model-1",
|
||||
"model_provider": "provider-1",
|
||||
"ai4sci_metering": metering,
|
||||
"ai4sci_model_proxy": proxy,
|
||||
"untrusted_extra": "do-not-copy",
|
||||
},
|
||||
"metadata": {"langgraph_api_url": "http://127.0.0.1:6176"},
|
||||
},
|
||||
)
|
||||
context = _memory_source_context(
|
||||
memory_dir="/memories",
|
||||
workspace_dir="/workspace",
|
||||
source_type=MemorySourceType.TURN,
|
||||
trajectory=[{"role": "human", "content": "hi"}],
|
||||
)
|
||||
|
||||
request = memory_launch.memory_worker_launch_request(context)
|
||||
payload = request.run_payload("worker-thread")
|
||||
configurable = payload["config"]["configurable"]
|
||||
|
||||
assert request.url == "http://127.0.0.1:6176"
|
||||
assert configurable["ai4sci_model_proxy"] == proxy
|
||||
assert "ai4sci_tool_effect" not in configurable
|
||||
assert configurable["ai4sci_metering"] == {
|
||||
**metering,
|
||||
"source_type": "evomemory_turn_worker",
|
||||
}
|
||||
assert "untrusted_extra" not in configurable
|
||||
assert "source_type" not in metering
|
||||
|
||||
|
||||
def test_memory_worker_finish_launches_linker_for_new_observations(
|
||||
tmp_path,
|
||||
):
|
||||
@@ -2154,6 +2316,66 @@ def test_observation_linker_launch_request_encodes_batch_context(tmp_path):
|
||||
]
|
||||
|
||||
|
||||
def test_observation_linker_inherits_worker_runtime_scope(tmp_path):
|
||||
context = memory_scheduler.ObservationLinkerContext(
|
||||
memory_dir=tmp_path / "memories",
|
||||
workspace_dir=tmp_path / "workspace",
|
||||
project_id="P-project",
|
||||
observation_ids=("O-1",),
|
||||
runtime_url="http://127.0.0.1:6176",
|
||||
runtime_configurable={
|
||||
"ai4sci_metering": {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "run-parent",
|
||||
"envelope_signature": "signed-parent",
|
||||
"source_type": "evomemory_turn_worker",
|
||||
},
|
||||
"ai4sci_model_proxy": {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "run-parent",
|
||||
"envelope_signature": "signed-parent",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
request = memory_launch.observation_linker_launch_request(context)
|
||||
configurable = request.run_payload("linker-thread")["config"]["configurable"]
|
||||
|
||||
assert request.url == "http://127.0.0.1:6176"
|
||||
assert configurable["ai4sci_metering"]["source_type"] == "evomemory_linker"
|
||||
assert configurable["ai4sci_model_proxy"]["run_id"] == "run-parent"
|
||||
|
||||
|
||||
def test_memory_scheduler_propagates_signed_runtime_to_linker(tmp_path):
|
||||
launched: list[memory_scheduler.ObservationLinkerContext] = []
|
||||
scheduler = memory_scheduler.MemoryScheduler(launch_linker=launched.append)
|
||||
memory_dir = tmp_path / "memories"
|
||||
observation = _record_test_observation(memory_dir)
|
||||
delta = worker_activity.MemoryOutputDelta(
|
||||
memory_dir=memory_dir,
|
||||
observation_paths=(_memory_relative_path(observation),),
|
||||
)
|
||||
configurable = {
|
||||
"ai4sci_metering": {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "run-parent",
|
||||
"envelope_signature": "signed-parent",
|
||||
}
|
||||
}
|
||||
|
||||
scheduler.record_worker_finished(
|
||||
_memory_worker_run(
|
||||
workspace_dir=str(tmp_path / "workspace"),
|
||||
configurable=configurable,
|
||||
),
|
||||
delta,
|
||||
)
|
||||
|
||||
assert len(launched) == 1
|
||||
assert launched[0].runtime_url == "http://x"
|
||||
assert launched[0].runtime_configurable == configurable
|
||||
|
||||
|
||||
def test_observation_linker_does_not_launch_when_observations_disabled(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
@@ -2298,7 +2520,12 @@ def test_memory_worker_observation_writer_modes(
|
||||
# exceptions from the auxiliary model call get normalized before
|
||||
# the tool-error handler sees them.
|
||||
assert type(middleware[0]).__name__ == "ErrorNormalizationMiddleware"
|
||||
assert type(middleware[1]).__name__ == "ToolErrorHandlerMiddleware"
|
||||
assert [type(item).__name__ for item in middleware[:4]] == [
|
||||
"ErrorNormalizationMiddleware",
|
||||
"ConfigurableModelMiddleware",
|
||||
"RecoverableMeteringMiddleware",
|
||||
"ToolErrorHandlerMiddleware",
|
||||
]
|
||||
assert _memory_tool_names(middleware) == expected_tools
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from langchain.agents import create_agent
|
||||
from langchain.agents.middleware.types import ModelRequest, ModelResponse
|
||||
from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from langgraph.types import Overwrite
|
||||
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError
|
||||
from EvoScientist.middleware.provider_context import ProviderContextMediaMiddleware
|
||||
|
||||
|
||||
class _Backend:
|
||||
def __init__(self, *, error: str | None = None) -> None:
|
||||
self.error = error
|
||||
self.uploads: list[tuple[str, bytes]] = []
|
||||
|
||||
def upload_files(self, files):
|
||||
self.uploads.extend(files)
|
||||
return [SimpleNamespace(path=path, error=self.error) for path, _data in files]
|
||||
|
||||
async def aupload_files(self, files):
|
||||
return self.upload_files(files)
|
||||
|
||||
|
||||
class _FakeModel(FakeMessagesListChatModel):
|
||||
def bind_tools(self, _tools, *, tool_choice=None, **_kwargs):
|
||||
return self
|
||||
|
||||
|
||||
def _request(messages):
|
||||
return ModelRequest(
|
||||
messages=list(messages),
|
||||
model=MagicMock(),
|
||||
state={},
|
||||
runtime=MagicMock(),
|
||||
system_message=MagicMock(),
|
||||
)
|
||||
|
||||
|
||||
def _image_block(raw: bytes) -> dict:
|
||||
return {
|
||||
"type": "image",
|
||||
"base64": base64.b64encode(raw).decode("ascii"),
|
||||
"mime_type": "image/png",
|
||||
}
|
||||
|
||||
|
||||
def test_externalizes_historical_and_new_assistant_media() -> None:
|
||||
backend = _Backend()
|
||||
middleware = ProviderContextMediaMiddleware(backend)
|
||||
old_raw = b"old-png"
|
||||
new_raw = b"new-png"
|
||||
historical = AIMessage(content=[_image_block(old_raw)])
|
||||
captured = {}
|
||||
|
||||
def handler(request):
|
||||
captured["messages"] = request.messages
|
||||
return ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "done"},
|
||||
_image_block(new_raw),
|
||||
]
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
response = middleware.wrap_model_call(_request([historical]), handler)
|
||||
|
||||
provider_content = captured["messages"][0].content
|
||||
assert all("base64" not in block for block in provider_content)
|
||||
assert "generated_image" in provider_content[0]["text"]
|
||||
stored_content = response.result[0].content
|
||||
assert stored_content[1]["type"] == "image"
|
||||
assert stored_content[1]["url"].startswith("/artifacts/model-output/")
|
||||
assert "base64" not in stored_content[1]
|
||||
assert {data for _path, data in backend.uploads} == {old_raw, new_raw}
|
||||
|
||||
|
||||
def test_before_model_durably_replaces_historical_inline_media() -> None:
|
||||
middleware = ProviderContextMediaMiddleware(_Backend())
|
||||
|
||||
update = middleware.before_model(
|
||||
{"messages": [AIMessage(content=[_image_block(b"old-png")])]}, None
|
||||
)
|
||||
|
||||
assert update is not None
|
||||
assert isinstance(update["messages"], Overwrite)
|
||||
content = update["messages"].value[0].content
|
||||
assert content[0]["url"].startswith("/artifacts/model-output/")
|
||||
assert "base64" not in content[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_real_langgraph_injects_runtime_and_repairs_checkpoint_media() -> None:
|
||||
backend = _Backend()
|
||||
raw = b"x" * 1_349_952
|
||||
agent = create_agent(
|
||||
model=_FakeModel(responses=[AIMessage(content="done")]),
|
||||
tools=[],
|
||||
middleware=[ProviderContextMediaMiddleware(backend)],
|
||||
)
|
||||
|
||||
result = await agent.ainvoke(
|
||||
{
|
||||
"messages": [
|
||||
AIMessage(content=[_image_block(raw)]),
|
||||
HumanMessage(content="continue"),
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
repaired = result["messages"][0].content[0]
|
||||
assert repaired["url"].startswith("/artifacts/model-output/")
|
||||
assert "base64" not in repaired
|
||||
assert any(data == raw for _path, data in backend.uploads)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_before_model_normalizes_internal_middleware_failure() -> None:
|
||||
class _BrokenBackend(_Backend):
|
||||
async def aupload_files(self, _files):
|
||||
raise TypeError("sensitive internal detail")
|
||||
|
||||
middleware = ProviderContextMediaMiddleware(_BrokenBackend())
|
||||
|
||||
with pytest.raises(EvoRuntimeError, match="AGENT_MIDDLEWARE_FAILED") as exc:
|
||||
await middleware.abefore_model(
|
||||
{"messages": [AIMessage(content=[_image_block(b"png")])]}, None
|
||||
)
|
||||
|
||||
assert exc.value.details == (
|
||||
{
|
||||
"failure_stage": "agent_middleware",
|
||||
"middleware": "provider_context_media",
|
||||
"middleware_node": "provider_context_media.before_model",
|
||||
"agent_error_type": "TypeError",
|
||||
"agent_error_module": "builtins",
|
||||
},
|
||||
)
|
||||
assert "sensitive internal detail" not in str(exc.value.details)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_externalization_is_content_addressed() -> None:
|
||||
backend = _Backend()
|
||||
middleware = ProviderContextMediaMiddleware(backend)
|
||||
raw = b"same-png"
|
||||
|
||||
async def handler(_request):
|
||||
return ModelResponse(result=[AIMessage(content=[_image_block(raw)])])
|
||||
|
||||
first = await middleware.awrap_model_call(_request([]), handler)
|
||||
second = await middleware.awrap_model_call(_request(first.result), handler)
|
||||
|
||||
assert first.result[0].content[0]["url"] == second.result[0].content[0]["url"]
|
||||
assert all(len(data) < 100 for _path, data in backend.uploads)
|
||||
|
||||
|
||||
def test_upload_failure_is_terminal_and_does_not_return_inline_media() -> None:
|
||||
middleware = ProviderContextMediaMiddleware(_Backend(error="disk full"))
|
||||
|
||||
with pytest.raises(EvoRuntimeError, match="MEDIA_PERSIST_FAILED"):
|
||||
middleware.wrap_model_call(
|
||||
_request([]),
|
||||
lambda _request: ModelResponse(
|
||||
result=[AIMessage(content=[_image_block(b"png")])]
|
||||
),
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,53 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from EvoScientist.middleware import recoverable_tools
|
||||
|
||||
|
||||
def test_tool_effect_context_prefers_dedicated_grant(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"langgraph.config.get_config",
|
||||
lambda: {
|
||||
"configurable": {
|
||||
"ai4sci_model_proxy": {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "model-run",
|
||||
"envelope_signature": "model-signature",
|
||||
},
|
||||
"ai4sci_tool_effect": {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "tool-run",
|
||||
"envelope_signature": "tool-signature",
|
||||
},
|
||||
},
|
||||
"metadata": {},
|
||||
},
|
||||
)
|
||||
|
||||
proxy, _ = recoverable_tools._context()
|
||||
|
||||
assert proxy == {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "tool-run",
|
||||
"envelope_signature": "tool-signature",
|
||||
}
|
||||
|
||||
|
||||
def test_evomemory_never_falls_back_to_parent_model_proxy(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"langgraph.config.get_config",
|
||||
lambda: {
|
||||
"configurable": {
|
||||
"ai4sci_model_proxy": {
|
||||
"gateway_url": "http://gateway",
|
||||
"run_id": "parent-run",
|
||||
"envelope_signature": "parent-signature",
|
||||
},
|
||||
},
|
||||
"metadata": {"run_kind": "evomemory_turn_worker"},
|
||||
},
|
||||
)
|
||||
|
||||
proxy, metadata = recoverable_tools._context()
|
||||
|
||||
assert proxy is None
|
||||
assert metadata["run_kind"] == "evomemory_turn_worker"
|
||||
@@ -15,7 +15,6 @@ from EvoScientist.runtime_integrations import (
|
||||
handle_knowledge_file,
|
||||
record_service_usage,
|
||||
reset_runtime_integrations,
|
||||
resolve_runtime_model,
|
||||
)
|
||||
|
||||
|
||||
@@ -101,17 +100,3 @@ async def test_host_can_register_runtime_integrations(tmp_path):
|
||||
assert get_image_backend() is image_backend
|
||||
assert knowledge_paths == [path]
|
||||
assert usage == [("mineru", "parse")]
|
||||
|
||||
|
||||
def test_host_can_register_model_resolver():
|
||||
resolved = object()
|
||||
calls = []
|
||||
|
||||
def resolve_model(model, provider):
|
||||
calls.append((model, provider))
|
||||
return resolved
|
||||
|
||||
configure_runtime_integrations(model_resolver=resolve_model)
|
||||
|
||||
assert resolve_runtime_model("model-a", "provider-a") is resolved
|
||||
assert calls == [("model-a", "provider-a")]
|
||||
|
||||
@@ -16,6 +16,7 @@ from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
||||
|
||||
from EvoScientist.sessions import (
|
||||
AGENT_NAME,
|
||||
_checkpoint_serde,
|
||||
_format_relative_time,
|
||||
_reduce_messages_delta,
|
||||
delete_thread,
|
||||
@@ -76,6 +77,59 @@ class TestGetDbPath(unittest.TestCase):
|
||||
assert ".evoscientist" in long_form or "evoscientist" in long_form.lower()
|
||||
|
||||
|
||||
def test_checkpoint_serde_allows_app_owned_errors():
|
||||
allowed = _checkpoint_serde()._allowed_msgpack_modules
|
||||
assert ("EvoScientist.llm.errors", "AgentControlError") in allowed
|
||||
assert ("EvoScientist.llm.errors", "ModelToolProtocolError") in allowed
|
||||
assert ("EvoScientist.llm.errors", "ProviderStreamError") in allowed
|
||||
|
||||
|
||||
def test_checkpoint_serde_roundtrips_model_tool_protocol_error():
|
||||
from EvoScientist.llm.errors import ModelToolProtocolError
|
||||
|
||||
error = ModelToolProtocolError(
|
||||
"missing_name",
|
||||
provider="openai",
|
||||
model="gpt-example",
|
||||
route_key="openai:primary:gpt-example",
|
||||
config_generation=7,
|
||||
call_id="call-1",
|
||||
call_diagnostic={"raw": "must-not-be-checkpointed"},
|
||||
)
|
||||
serde = _checkpoint_serde()
|
||||
restored = serde.loads_typed(serde.dumps_typed({"error": error}))["error"]
|
||||
|
||||
assert isinstance(restored, ModelToolProtocolError)
|
||||
assert restored.code == "MODEL_TOOL_PROTOCOL_INVALID"
|
||||
assert restored.reason == "missing_name"
|
||||
assert restored.provider == "openai"
|
||||
assert restored.config_generation == 7
|
||||
assert restored.call_id == "call-1"
|
||||
assert restored.call_diagnostic == {}
|
||||
|
||||
|
||||
def test_checkpoint_serde_roundtrips_provider_stream_error():
|
||||
from EvoScientist.llm.errors import ProviderStreamError
|
||||
|
||||
error = ProviderStreamError(
|
||||
provider="openai",
|
||||
class_qualname="openai.BadRequestError",
|
||||
message="Provider rejected the request.",
|
||||
status_code=400,
|
||||
code="invalid_request",
|
||||
request_id="request-1",
|
||||
)
|
||||
serde = _checkpoint_serde()
|
||||
restored = serde.loads_typed(serde.dumps_typed({"error": error}))["error"]
|
||||
|
||||
assert isinstance(restored, ProviderStreamError)
|
||||
assert restored.provider == "openai"
|
||||
assert restored.class_qualname == "openai.BadRequestError"
|
||||
assert restored.status_code == 400
|
||||
assert restored.code == "invalid_request"
|
||||
assert restored.request_id == "request-1"
|
||||
|
||||
|
||||
class TestFormatRelativeTime(unittest.TestCase):
|
||||
def test_none(self):
|
||||
assert _format_relative_time(None) == ""
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from langchain.agents.middleware.types import ModelRequest
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
from EvoScientist.middleware.skill_context import BudgetedSkillsMiddleware
|
||||
|
||||
|
||||
def _skill(name: str, description: str) -> dict[str, object]:
|
||||
return {
|
||||
"name": name,
|
||||
"description": description,
|
||||
"path": f"/skills/{name}/SKILL.md",
|
||||
"allowed_tools": [],
|
||||
}
|
||||
|
||||
|
||||
def _request(query: str, skills: list[dict[str, object]]) -> ModelRequest:
|
||||
return ModelRequest(
|
||||
model=object(),
|
||||
messages=[HumanMessage(content=query)],
|
||||
system_prompt="Base system prompt.",
|
||||
state={"skills_metadata": skills},
|
||||
)
|
||||
|
||||
|
||||
def test_skill_context_prefers_relevant_skill_and_omits_irrelevant_catalog():
|
||||
middleware = BudgetedSkillsMiddleware(
|
||||
backend=object(), sources=["/skills/"], max_skills=2, max_skills_bytes=1024
|
||||
)
|
||||
result = middleware.modify_request(
|
||||
_request(
|
||||
"请分析蛋白质结构预测结果",
|
||||
[
|
||||
_skill("protein-structure", "蛋白质结构预测和结果分析。"),
|
||||
_skill("frontend-design", "Build polished user interfaces."),
|
||||
],
|
||||
)
|
||||
)
|
||||
|
||||
prompt = str(result.system_message.content)
|
||||
assert "protein-structure" in prompt
|
||||
assert "frontend-design" not in prompt
|
||||
assert "query-relevant subset" in prompt
|
||||
|
||||
|
||||
def test_skill_context_enforces_count_and_utf8_budget():
|
||||
middleware = BudgetedSkillsMiddleware(
|
||||
backend=object(),
|
||||
sources=["/skills/"],
|
||||
max_skills=16,
|
||||
max_skills_bytes=2048,
|
||||
max_description_bytes=128,
|
||||
)
|
||||
skills = [
|
||||
_skill(f"analysis-{index}", "analysis " + "x" * 1_000) for index in range(300)
|
||||
]
|
||||
|
||||
selected = middleware._select_skills(skills, "analysis")
|
||||
rendered = middleware._format_budgeted_skills(selected)
|
||||
|
||||
assert len(selected) == 16
|
||||
assert len(rendered.encode("utf-8")) <= 2048
|
||||
assert "analysis-299" not in rendered
|
||||
|
||||
|
||||
def test_skill_context_does_not_fall_back_to_all_skills_without_a_match():
|
||||
middleware = BudgetedSkillsMiddleware(backend=object(), sources=["/skills/"])
|
||||
result = middleware.modify_request(
|
||||
_request(
|
||||
"unrelated request",
|
||||
[_skill("protein-structure", "Protein folding workflow.")],
|
||||
)
|
||||
)
|
||||
|
||||
prompt = str(result.system_message.content)
|
||||
assert "protein-structure" not in prompt
|
||||
assert "query-relevant subset" in prompt
|
||||
|
||||
|
||||
def test_skill_context_accepts_a_single_skill_source():
|
||||
middleware = BudgetedSkillsMiddleware(backend=object(), sources="/skills/")
|
||||
|
||||
assert middleware.sources == ["/skills/"]
|
||||
@@ -35,7 +35,6 @@ def _call(call_id: str = "call-1", name: str = "search", args: Any = None):
|
||||
(_call(name=""), "missing_name"),
|
||||
(_call(name=" "), "missing_name"),
|
||||
(_call(name="missing"), "unknown_name"),
|
||||
(_call(call_id=""), "missing_id"),
|
||||
],
|
||||
)
|
||||
def test_invalid_final_tool_call_fails_closed(call, reason):
|
||||
@@ -46,14 +45,28 @@ def test_invalid_final_tool_call_fails_closed(call, reason):
|
||||
middleware.wrap_model_call(request, lambda _request: _response(call))
|
||||
|
||||
assert caught.value.reason == reason
|
||||
assert caught.value.retryable is False
|
||||
assert caught.value.retryable is True
|
||||
assert caught.value.fallbackable is True
|
||||
assert caught.value.non_fallbackable is False
|
||||
|
||||
|
||||
def test_non_mapping_args_are_rejected_if_adapter_bypasses_message_validation():
|
||||
def test_json_string_args_are_normalized_to_an_object():
|
||||
message = AIMessage(content="", tool_calls=[_call()])
|
||||
message.tool_calls[0]["args"] = "{}"
|
||||
|
||||
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
|
||||
assert result.result[0].tool_calls[0]["args"] == {}
|
||||
assert message.tool_calls[0]["args"] == "{}"
|
||||
|
||||
|
||||
def test_malformed_json_args_are_rejected():
|
||||
message = AIMessage(content="", tool_calls=[_call()])
|
||||
message.tool_calls[0]["args"] = "{"
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
@@ -121,7 +134,30 @@ def test_content_block_must_match_parsed_call():
|
||||
lambda _request: response,
|
||||
)
|
||||
|
||||
assert caught.value.reason == "inconsistent_block"
|
||||
assert caught.value.reason == "inconsistent_source"
|
||||
|
||||
|
||||
def test_responses_content_block_uses_call_id_over_output_item_id():
|
||||
"""Responses item IDs are not the identifiers used for tool results."""
|
||||
|
||||
response = _response(
|
||||
_call(call_id="call-result-1", name="search", args={}),
|
||||
content=[
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "fc-output-item-1",
|
||||
"call_id": "call-result-1",
|
||||
"name": "search",
|
||||
"arguments": "{}",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]), lambda _request: response
|
||||
)
|
||||
|
||||
assert result.result[0].tool_calls[0]["id"] == "call-result-1"
|
||||
|
||||
|
||||
def test_parsed_only_valid_call_and_extended_response_pass():
|
||||
@@ -172,37 +208,50 @@ def test_error_carries_safe_route_metadata():
|
||||
assert "args" not in payload
|
||||
|
||||
|
||||
def test_missing_id_carries_redacted_call_diagnostic_only_for_internal_logging():
|
||||
def test_missing_id_is_generated_without_mutating_the_provider_message():
|
||||
call = _call(call_id="", name="search", args={"query": "private search text"})
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: _response(call),
|
||||
)
|
||||
provider_response = _response(call)
|
||||
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: provider_response,
|
||||
)
|
||||
|
||||
diagnostic = caught.value.call_diagnostic
|
||||
assert diagnostic == {
|
||||
"source": "parsed_tool_calls",
|
||||
"call_index": 0,
|
||||
"call_count": 1,
|
||||
"call_type": "object",
|
||||
"name": "search",
|
||||
"id_present": False,
|
||||
"args_present": True,
|
||||
"args_type": "object",
|
||||
"args_key_count": 1,
|
||||
"args_keys": ["query"],
|
||||
"args_keys_truncated": False,
|
||||
"args_digest": diagnostic["args_digest"],
|
||||
"raw_openai_call_available": False,
|
||||
}
|
||||
assert diagnostic["args_digest"].startswith("sha256:")
|
||||
assert "private search text" not in str(diagnostic)
|
||||
assert "call_diagnostic" not in caught.value.model_dump()
|
||||
normalized = result.result[0].tool_calls[0]
|
||||
assert normalized["id"].startswith("call_")
|
||||
assert normalized["args"] == {"query": "private search text"}
|
||||
assert provider_response.result[0].tool_calls[0]["id"] == ""
|
||||
|
||||
|
||||
def test_diagnostic_compares_parsed_and_preserved_raw_openai_call_shapes():
|
||||
def test_generic_openai_compatible_parallel_calls_without_ids_are_stable():
|
||||
model = SimpleNamespace(
|
||||
metadata={
|
||||
"route_adapter_id": "generic-openai-compatible",
|
||||
"route_key": "qwen-primary",
|
||||
"route_config_generation": 7,
|
||||
"route_api_mode": "chat_completions",
|
||||
}
|
||||
)
|
||||
response = _response(
|
||||
_call(call_id="", name="think_tool", args={"reflection": "plan"}),
|
||||
_call(call_id="", name="execute", args={"command": "pwd"}),
|
||||
)
|
||||
|
||||
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(
|
||||
tools=[{"name": "think_tool"}, {"name": "execute"}], model=model
|
||||
),
|
||||
lambda _request: response,
|
||||
)
|
||||
|
||||
calls = result.result[0].tool_calls
|
||||
assert [call["name"] for call in calls] == ["think_tool", "execute"]
|
||||
assert all(call["id"].startswith("call_") for call in calls)
|
||||
assert calls[0]["id"] != calls[1]["id"]
|
||||
assert response.result[0].tool_calls[0]["id"] == ""
|
||||
|
||||
|
||||
def test_raw_openai_id_is_merged_into_the_canonical_call():
|
||||
parsed = _call(call_id="", name="search", args={"query": "secret"})
|
||||
raw = {
|
||||
"id": "provider-call-id",
|
||||
@@ -215,19 +264,62 @@ def test_diagnostic_compares_parsed_and_preserved_raw_openai_call_shapes():
|
||||
additional_kwargs={"tool_calls": [raw]},
|
||||
)
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
|
||||
diagnostic = caught.value.call_diagnostic
|
||||
assert diagnostic["id_present"] is False
|
||||
assert diagnostic["raw_openai_call_available"] is True
|
||||
assert diagnostic["raw_openai_call"]["id_present"] is True
|
||||
assert diagnostic["raw_openai_call"]["name"] == "search"
|
||||
assert "provider-call-id" not in str(diagnostic)
|
||||
assert "secret" not in str(diagnostic)
|
||||
normalized = result.result[0]
|
||||
assert normalized.tool_calls == [
|
||||
{
|
||||
"id": "provider-call-id",
|
||||
"name": "search",
|
||||
"args": {"query": "secret"},
|
||||
"type": "tool_call",
|
||||
}
|
||||
]
|
||||
assert "tool_calls" not in normalized.additional_kwargs
|
||||
assert message.additional_kwargs["tool_calls"] == [raw]
|
||||
|
||||
|
||||
def test_content_only_function_call_is_decoded_and_normalized():
|
||||
message = AIMessage(
|
||||
content=[
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "",
|
||||
"name": "search",
|
||||
"arguments": '{"query":"x"}',
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
|
||||
normalized = result.result[0]
|
||||
call_id = normalized.tool_calls[0]["id"]
|
||||
assert call_id.startswith("call_")
|
||||
assert normalized.tool_calls[0]["args"] == {"query": "x"}
|
||||
assert normalized.content[0]["id"] == call_id
|
||||
|
||||
|
||||
def test_legacy_function_call_is_decoded_and_removed_from_replay_metadata():
|
||||
legacy = {"name": "search", "arguments": '{"query":"x"}'}
|
||||
message = AIMessage(content="", additional_kwargs={"function_call": legacy})
|
||||
|
||||
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
|
||||
normalized = result.result[0]
|
||||
assert normalized.tool_calls[0]["id"].startswith("call_")
|
||||
assert normalized.tool_calls[0]["args"] == {"query": "x"}
|
||||
assert "function_call" not in normalized.additional_kwargs
|
||||
assert message.additional_kwargs["function_call"] == legacy
|
||||
|
||||
|
||||
def test_diagnostic_failure_cannot_mask_the_protocol_error():
|
||||
@@ -244,3 +336,18 @@ def test_diagnostic_failure_cannot_mask_the_protocol_error():
|
||||
|
||||
assert caught.value.reason == "missing_id"
|
||||
assert caught.value.call_diagnostic["args_digest"].startswith("sha256:")
|
||||
|
||||
|
||||
def test_protocol_failure_logs_only_redacted_call_diagnostic(caplog):
|
||||
call = _call(name="", args={"query": "private search text"})
|
||||
|
||||
with caplog.at_level("WARNING"), pytest.raises(ModelToolProtocolError):
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: _response(call),
|
||||
)
|
||||
|
||||
record = caplog.records[-1].getMessage()
|
||||
assert "reason=missing_name" in record
|
||||
assert '"args_keys": ["query"]' in record
|
||||
assert "private search text" not in record
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError
|
||||
from EvoScientist.llm.user_options import (
|
||||
model_options_schema_hash,
|
||||
project_user_options_for_purpose,
|
||||
validate_user_model_options,
|
||||
)
|
||||
|
||||
|
||||
def _validate(options):
|
||||
return validate_user_model_options(
|
||||
supplied=options,
|
||||
user_options={
|
||||
"temperature": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"maximum": 2,
|
||||
"applies_to": ["main_agent"],
|
||||
},
|
||||
"top_p": {
|
||||
"type": "number",
|
||||
"minimum_exclusive": 0,
|
||||
"maximum": 1,
|
||||
"applies_to": ["main_agent"],
|
||||
},
|
||||
},
|
||||
supports_reasoning=True,
|
||||
reasoning_mode="effort",
|
||||
allowed_reasoning_efforts=("low", "high"),
|
||||
default_reasoning_effort="high",
|
||||
parameter_constraints=({"at_most_one_of": ("temperature", "top_p")},),
|
||||
)
|
||||
|
||||
|
||||
def test_user_options_validate_canonical_values():
|
||||
assert _validate({"temperature": 0.4, "reasoning": "high"}) == {
|
||||
"temperature": 0.4,
|
||||
"reasoning": "high",
|
||||
}
|
||||
|
||||
|
||||
def test_purpose_projection_filters_user_options_but_preserves_runtime_parameters():
|
||||
projected = project_user_options_for_purpose(
|
||||
values={"temperature": 1, "structured_output": True},
|
||||
user_options={
|
||||
"temperature": {
|
||||
"type": "number",
|
||||
"applies_to": ["main_agent"],
|
||||
}
|
||||
},
|
||||
purpose="tool_selector",
|
||||
)
|
||||
|
||||
assert projected == {"structured_output": True}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"options",
|
||||
[
|
||||
{"temperature": 3},
|
||||
{"top_p": 0},
|
||||
{"unknown": True},
|
||||
{"temperature": 0.4, "top_p": 0.8},
|
||||
],
|
||||
)
|
||||
def test_user_options_reject_invalid_or_conflicting_values(options):
|
||||
with pytest.raises(EvoRuntimeError):
|
||||
_validate(options)
|
||||
|
||||
|
||||
def test_options_schema_hash_tracks_semantics_not_defaults():
|
||||
common = {
|
||||
"model_profile_id": "profile-1",
|
||||
"supports_reasoning": False,
|
||||
"reasoning_mode": "none",
|
||||
"allowed_reasoning_efforts": (),
|
||||
"parameter_constraints": (),
|
||||
"adapter_id": "openai-compatible",
|
||||
"adapter_revision": "3",
|
||||
}
|
||||
first = model_options_schema_hash(
|
||||
**common,
|
||||
user_options={"temperature": {"type": "number", "maximum": 2, "default": 1}},
|
||||
)
|
||||
default_changed = model_options_schema_hash(
|
||||
**common,
|
||||
user_options={"temperature": {"type": "number", "maximum": 2, "default": 0}},
|
||||
)
|
||||
constraint_changed = model_options_schema_hash(
|
||||
**common,
|
||||
user_options={"temperature": {"type": "number", "maximum": 1, "default": 0}},
|
||||
)
|
||||
adapter_changed = model_options_schema_hash(
|
||||
**{**common, "adapter_revision": "4"},
|
||||
user_options={"temperature": {"type": "number", "maximum": 2, "default": 0}},
|
||||
)
|
||||
|
||||
assert first == default_changed
|
||||
assert first != constraint_changed
|
||||
assert first != adapter_changed
|
||||
@@ -0,0 +1,146 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import replace
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.configuration.provider import ResolvedSecret
|
||||
from EvoScientist.llm.contracts import (
|
||||
AgentInputV3,
|
||||
EvoRuntimeError,
|
||||
HmacGrantAuthority,
|
||||
WebHostContext,
|
||||
)
|
||||
from EvoScientist.llm.model_config import EvoModelConfig
|
||||
from EvoScientist.llm.runtime import EvoModelRuntime
|
||||
from EvoScientist.llm.user_options import model_options_schema_hash
|
||||
from tests.test_provider_model_config_v3 import v3_payload as provider_v3_payload
|
||||
from tests.test_web_model_runtime import (
|
||||
_input,
|
||||
_preparation,
|
||||
_runtime,
|
||||
_Sink,
|
||||
)
|
||||
from tests.v3_fixtures import (
|
||||
IDENTITY_KEY_ID,
|
||||
RUNTIME_KEY_ID,
|
||||
RUNTIME_SECRET,
|
||||
identity_ring,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_rejects_gateway_catalog_schema_that_became_stale(
|
||||
monkeypatch, tmp_path
|
||||
):
|
||||
runtime, authority = _runtime(tmp_path, monkeypatch)
|
||||
original = _input()
|
||||
agent_input = replace(
|
||||
original,
|
||||
metadata={
|
||||
**dict(original.metadata),
|
||||
"model_options_schema_hash": "sha256:stale",
|
||||
},
|
||||
)
|
||||
host = WebHostContext(
|
||||
"/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()
|
||||
)
|
||||
|
||||
with pytest.raises(EvoRuntimeError, match="MODEL_OPTIONS_STALE"):
|
||||
await runtime.prepare_model_run(
|
||||
_preparation(authority, agent_input),
|
||||
agent_input,
|
||||
host,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_filters_main_only_options_from_inherited_auxiliary_routes():
|
||||
payload = provider_v3_payload()
|
||||
payload["config_identity_key_id"] = IDENTITY_KEY_ID
|
||||
provider = next(
|
||||
item for item in payload["providers"] if item["adapter_id"] == "openai"
|
||||
)
|
||||
model_payload = provider["models"][0]
|
||||
model_payload["capabilities"]["thinking"] = True
|
||||
model_payload["parameters"]["reasoning_policy"] = {
|
||||
"mode": "effort",
|
||||
"allowed_efforts": ["low", "medium", "high"],
|
||||
"default_effort": "high",
|
||||
}
|
||||
model_payload["parameters"]["user_options"]["temperature"] = {
|
||||
"default": 1,
|
||||
"applies_to": ["main_agent"],
|
||||
"minimum": 0,
|
||||
"maximum_exclusive": 2,
|
||||
}
|
||||
alias_payload = next(
|
||||
item for item in payload["aliases"] if item["alias"] == "openai-prod-general"
|
||||
)
|
||||
alias_payload["defaults"] = {"temperature": 1}
|
||||
config = EvoModelConfig.parse(payload, require_evidence=False)
|
||||
authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID)
|
||||
|
||||
def resolve_secret(reference):
|
||||
return ResolvedSecret(
|
||||
"test-secret", reference.revision, "1", "test-fingerprint"
|
||||
)
|
||||
|
||||
runtime = EvoModelRuntime(
|
||||
SimpleNamespace(load=lambda: config),
|
||||
admission_verifier=authority,
|
||||
quote_authority=authority,
|
||||
identity_key_ring=identity_ring(),
|
||||
secret_resolver=resolve_secret,
|
||||
)
|
||||
selector = config.resolve_main_selector("openai-prod-general")
|
||||
selected_model = config.providers[selector.provider].models[selector.model]
|
||||
selected_provider = config.providers[selector.provider]
|
||||
schema_hash = model_options_schema_hash(
|
||||
model_profile_id="openai-prod-general",
|
||||
user_options=selected_model.user_options,
|
||||
supports_reasoning=selected_model.supports_reasoning,
|
||||
reasoning_mode=selected_model.reasoning_mode,
|
||||
allowed_reasoning_efforts=selected_model.allowed_reasoning_efforts,
|
||||
parameter_constraints=selected_model.parameter_constraints,
|
||||
adapter_id=selected_provider.adapter_id,
|
||||
adapter_revision=selected_provider.adapter_revision,
|
||||
)
|
||||
agent_input = AgentInputV3(
|
||||
"hello",
|
||||
"web:user:thread",
|
||||
metadata={
|
||||
"source": "web",
|
||||
"model_options": {"temperature": 0.7},
|
||||
"model_options_schema_hash": schema_hash,
|
||||
},
|
||||
)
|
||||
grant = authority.sign_preparation(
|
||||
request_id="11111111-1111-4111-8111-111111111111",
|
||||
turn_id="22222222-2222-4222-8222-222222222222",
|
||||
thread_id="thread",
|
||||
subject_id="user",
|
||||
requested_model_ref="openai-prod-general",
|
||||
plan="starter",
|
||||
roles=("user",),
|
||||
requires_vision=False,
|
||||
reasoning_effort="high",
|
||||
title_policy="best_effort",
|
||||
gateway_input_digest=authority.agent_input_digest(agent_input.projection()),
|
||||
checkpoint_thread_id=agent_input.checkpoint_thread_id,
|
||||
checkpoint_snapshot_id="sha256:checkpoint",
|
||||
turn_fencing_token=1,
|
||||
ttl_ms=60_000,
|
||||
)
|
||||
|
||||
await runtime.prepare_model_run(
|
||||
grant,
|
||||
agent_input,
|
||||
WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()),
|
||||
)
|
||||
|
||||
snapshot = next(iter(runtime._prepared.values())).snapshot
|
||||
assert snapshot.purpose_routes["main_agent"][0].params["temperature"] == 0.7
|
||||
assert "temperature" not in snapshot.purpose_routes["tool_selector"][0].params
|
||||
assert "temperature" not in snapshot.purpose_routes["deepagents_summarizer"][0].params
|
||||
@@ -0,0 +1,64 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import replace
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError, HmacGrantAuthority
|
||||
from EvoScientist.llm.crypto import canonical_json_v1
|
||||
from EvoScientist.sessions import FencedPruningCheckpointer
|
||||
from tests.v3_fixtures import RUNTIME_KEY_ID, RUNTIME_SECRET
|
||||
|
||||
|
||||
def test_canonical_json_normalizes_unicode_and_orders_keys():
|
||||
assert canonical_json_v1({"z": "e\u0301", "a": 1}) == canonical_json_v1(
|
||||
{"a": 1, "z": "\u00e9"}
|
||||
)
|
||||
|
||||
|
||||
def test_previous_runtime_key_verifies_during_rotation():
|
||||
old_secret = "old-runtime-secret-for-tests-000000000000000000000"
|
||||
old = HmacGrantAuthority(old_secret, "old-key")
|
||||
grant = old.sign_subject(
|
||||
subject_id="user", plan="starter", roles=("user",), ttl_ms=60_000
|
||||
)
|
||||
rotated = HmacGrantAuthority(
|
||||
RUNTIME_SECRET,
|
||||
RUNTIME_KEY_ID,
|
||||
previous_secret=old_secret,
|
||||
previous_key_id="old-key",
|
||||
)
|
||||
assert rotated.verify_subject(grant)
|
||||
with pytest.raises(EvoRuntimeError, match="CONTRACT_SIGNATURE_INVALID"):
|
||||
rotated.require_admin(replace(grant, audience="evoscientist-runtime"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_turn_token_cannot_write_or_release_new_lease(tmp_path):
|
||||
database = tmp_path / "sessions.sqlite"
|
||||
async with FencedPruningCheckpointer.from_conn_string_with_keep(
|
||||
str(database), keep_per_ns=10
|
||||
) as saver:
|
||||
first = await saver.acquire_turn_lease(
|
||||
"web:user:thread", "owner-1", ttl_seconds=30
|
||||
)
|
||||
assert await saver.release_turn_lease(first)
|
||||
current = await saver.acquire_turn_lease(
|
||||
"web:user:thread", "owner-2", ttl_seconds=30
|
||||
)
|
||||
assert current.fencing_token == first.fencing_token + 1
|
||||
assert not await saver.release_turn_lease(first)
|
||||
with pytest.raises(RuntimeError, match="TURN_FENCED"):
|
||||
await saver.aput_writes(
|
||||
{
|
||||
"configurable": {
|
||||
"thread_id": first.thread_id,
|
||||
"checkpoint_id": "checkpoint",
|
||||
"turn_lease_owner": first.owner_id,
|
||||
"turn_fencing_token": first.fencing_token,
|
||||
}
|
||||
},
|
||||
[("channel", "value")],
|
||||
"task",
|
||||
)
|
||||
assert await saver.release_turn_lease(current)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,17 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.contracts import EvoRuntimeError
|
||||
from EvoScientist.web_runtime import _ToolRegistryFenceMiddleware
|
||||
|
||||
|
||||
def test_tool_dispatch_fence_rejects_changed_registry(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"EvoScientist.web_runtime.web_tool_registry_manifest",
|
||||
lambda: ((), "registry-v2"),
|
||||
)
|
||||
middleware = _ToolRegistryFenceMiddleware("registry-v1")
|
||||
|
||||
with pytest.raises(EvoRuntimeError, match="TOOL_REGISTRY_STALE"):
|
||||
middleware._require_current()
|
||||
@@ -0,0 +1,161 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from EvoScientist.llm.crypto import HmacKeyRing, KeyMaterial
|
||||
from EvoScientist.llm.model_config import (
|
||||
EvoModelConfig,
|
||||
adapter_revision,
|
||||
endpoint_fingerprint,
|
||||
route_semantics_hash,
|
||||
)
|
||||
|
||||
RUNTIME_SECRET = "runtime-secret-for-tests-000000000000000000000000"
|
||||
IDENTITY_SECRET = "identity-secret-for-tests-00000000000000000000000"
|
||||
RUNTIME_KEY_ID = "runtime-test"
|
||||
IDENTITY_KEY_ID = "identity-test"
|
||||
|
||||
|
||||
def identity_ring() -> HmacKeyRing:
|
||||
return HmacKeyRing(KeyMaterial.create(IDENTITY_KEY_ID, IDENTITY_SECRET))
|
||||
|
||||
|
||||
def v3_payload(*, revision: int = 1) -> dict[str, Any]:
|
||||
payload: dict[str, Any] = {
|
||||
"schema_version": 2,
|
||||
"config_revision": revision,
|
||||
"config_identity_key_id": IDENTITY_KEY_ID,
|
||||
"runtime_defaults": {"max_retries": 0},
|
||||
"purpose_defaults": {
|
||||
"main_agent": {"reasoning_effort": "medium"},
|
||||
"tool_selector": {"reasoning_effort": "disabled"},
|
||||
"deepagents_summarizer": {"reasoning_effort": "disabled"},
|
||||
"title": {"reasoning_effort": "disabled"},
|
||||
},
|
||||
"providers": {
|
||||
"custom-openai": {
|
||||
"protocol": "custom-openai",
|
||||
"params": {},
|
||||
"endpoints": [
|
||||
{
|
||||
"name": "primary",
|
||||
"base_url": "https://provider.example/v1",
|
||||
"auth": {"ref": "env://WEB_RUNTIME_TEST_KEY", "revision": 1},
|
||||
"headers": {"User-Agent": "Ai4Sci-Test"},
|
||||
"header_refs": {},
|
||||
"params": {"extra_body": {}},
|
||||
}
|
||||
],
|
||||
"models": [
|
||||
{
|
||||
"id": "model-id",
|
||||
"params": {"output_token_limit": 512},
|
||||
"supports_vision": True,
|
||||
"supports_reasoning": True,
|
||||
"allowed_reasoning_efforts": ["low", "medium", "high"],
|
||||
"context_window": 8192,
|
||||
"max_output_tokens": 2048,
|
||||
"access": {"allowed_plans": [], "allowed_roles": []},
|
||||
"billing": {
|
||||
"sku": "visible-model",
|
||||
"pricing_revision": "test-1",
|
||||
"currency": "CNY",
|
||||
"unit_scale": 1_000_000,
|
||||
"input_microunits_per_million": 1_000_000,
|
||||
"output_microunits_per_million": 2_000_000,
|
||||
"cached_microunits_per_million": 500_000,
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
},
|
||||
"endpoint_pools": {
|
||||
"default": {
|
||||
"provider": "custom-openai",
|
||||
"strategy": "smooth_weighted_round_robin",
|
||||
"endpoints": [{"name": "primary", "weight": 1}],
|
||||
}
|
||||
},
|
||||
"route_health": {
|
||||
"failure_threshold": 3,
|
||||
"cooldown_seconds": 30,
|
||||
"half_open_max_inflight": 1,
|
||||
"counted_error_codes": ["PROVIDER_5XX", "PROVIDER_TIMEOUT"],
|
||||
"open_immediately_error_codes": ["PROVIDER_AUTH_INVALID"],
|
||||
},
|
||||
"route_selectors": {
|
||||
"visible-main": {
|
||||
"provider": "custom-openai",
|
||||
"endpoint_pool": "default",
|
||||
"model": "model-id",
|
||||
"api_mode": "chat_completions",
|
||||
"tool_call_transport": "non_streaming",
|
||||
}
|
||||
},
|
||||
"purpose_routes": {
|
||||
"main_agent": {
|
||||
"default_alias": "visible-model",
|
||||
"selectable": {"visible-model": "visible-main"},
|
||||
},
|
||||
"title": {"default": "visible-main"},
|
||||
},
|
||||
"purpose_call_limits": {
|
||||
"main_agent": {"max_attempts_per_run": 4},
|
||||
"tool_selector": {"max_attempts_per_run": 2},
|
||||
"deepagents_summarizer": {"max_attempts_per_run": 1},
|
||||
"title": {"max_attempts_per_run": 1},
|
||||
},
|
||||
"web_runtime": {
|
||||
"title_start_timeout_seconds": 30,
|
||||
"prepare_ttl_seconds": 30,
|
||||
"turn_lease_grace_seconds": 10,
|
||||
"active_run_timeout_seconds": 60,
|
||||
"max_run_journal_events": 1000,
|
||||
"max_run_journal_bytes": 1_048_576,
|
||||
"max_prepared_runs_per_subject": 2,
|
||||
"max_prepared_runs_total": 100,
|
||||
},
|
||||
"capability_evidence": [],
|
||||
"tool_protocol_fallbacks": [{"primary": "visible-main", "fallbacks": []}],
|
||||
}
|
||||
candidate = EvoModelConfig.parse(payload, require_evidence=False)
|
||||
ring = identity_ring()
|
||||
semantics_key = ring.derive_current("ai4sci/route-semantics-hash/v3")[1]
|
||||
endpoint_key = ring.derive_current("ai4sci/endpoint-fingerprint/v3")[1]
|
||||
evidence = []
|
||||
seen = set()
|
||||
for selector_id in candidate.route_selectors:
|
||||
for route in candidate.concrete_routes(selector_id):
|
||||
if route.key() in seen:
|
||||
continue
|
||||
seen.add(route.key())
|
||||
evidence.append(
|
||||
{
|
||||
"route": {
|
||||
"provider": route.provider,
|
||||
"endpoint": route.endpoint,
|
||||
"model": route.model,
|
||||
"api_mode": route.api_mode,
|
||||
"tool_call_transport": route.tool_call_transport,
|
||||
},
|
||||
"connectivity": "supported",
|
||||
"tool_capability": "supported",
|
||||
"probe": {
|
||||
"route_semantics_hash": route_semantics_hash(
|
||||
candidate, route, semantics_key
|
||||
),
|
||||
"endpoint_fingerprint": endpoint_fingerprint(
|
||||
candidate, route, endpoint_key
|
||||
),
|
||||
"config_identity_key_id": IDENTITY_KEY_ID,
|
||||
"adapter_revision": adapter_revision(
|
||||
candidate.providers[route.provider].protocol,
|
||||
route.api_mode,
|
||||
),
|
||||
"fixture_digest": "sha256:test-fixture-v3",
|
||||
"verified_at": "2026-07-20T00:00:00+00:00",
|
||||
},
|
||||
}
|
||||
)
|
||||
payload["capability_evidence"] = evidence
|
||||
return payload
|
||||
@@ -968,6 +968,7 @@ dependencies = [
|
||||
{ name = "python-dotenv" },
|
||||
{ name = "pyyaml" },
|
||||
{ name = "questionary" },
|
||||
{ name = "rfc8785" },
|
||||
{ name = "rich" },
|
||||
{ name = "tavily-python" },
|
||||
{ name = "textual" },
|
||||
@@ -1090,6 +1091,7 @@ requires-dist = [
|
||||
{ name = "qrcode", marker = "extra == 'qq'", specifier = ">=7.4" },
|
||||
{ name = "qrcode", marker = "extra == 'wechat'", specifier = ">=7.4" },
|
||||
{ name = "questionary", specifier = ">=2.1" },
|
||||
{ name = "rfc8785", specifier = "==0.1.4" },
|
||||
{ name = "rich", specifier = ">=15.0" },
|
||||
{ name = "ruff", marker = "extra == 'dev'", specifier = ">=0.5" },
|
||||
{ name = "slack-sdk", marker = "extra == 'all-channels'", specifier = ">=3.27" },
|
||||
@@ -3871,6 +3873,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/3f/51/d4db610ef29373b879047326cbf6fa98b6c1969d6f6dc423279de2b1be2c/requests_toolbelt-1.0.0-py2.py3-none-any.whl", hash = "sha256:cccfdd665f0a24fcf4726e690f65639d272bb0637b9b92dfd91a5568ccf6bd06", size = 54481, upload-time = "2023-05-01T04:11:28.427Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rfc8785"
|
||||
version = "0.1.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/ef/2f/fa1d2e740c490191b572d33dbca5daa180cb423c24396b856f5886371d8b/rfc8785-0.1.4.tar.gz", hash = "sha256:e545841329fe0eee4f6a3b44e7034343100c12b4ec566dc06ca9735681deb4da", size = 14321, upload-time = "2024-09-27T16:33:31.206Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/4d/78/119878110660b2ad709888c8a1614fce7e2fab39080ab960656dc8605bf6/rfc8785-0.1.4-py3-none-any.whl", hash = "sha256:520d690b448ecf0703691c76e1a34a24ddcd4fc5bc41d589cb7c58ec651bcd48", size = 9240, upload-time = "2024-09-27T16:33:29.683Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rich"
|
||||
version = "15.0.0"
|
||||
|
||||
Reference in New Issue
Block a user