55 Commits

Author SHA1 Message Date
m4 96c380caa3 fix(release): idempotent bump/tag, push current branch, ignore state file
Docker / build (push) Has been cancelled
2026-08-13 11:03:15 +08:00
m4 59c77c65b1 feat(release): allow release branches via RELEASE_BRANCHES env (default main) 2026-08-13 10:35:55 +08:00
m4 ae461d1c2d chore: bump version to 0.2.3 2026-08-13 10:18:33 +08:00
m4 2a7cccd598 feat(model-registry): add admin config export endpoints with plaintext secrets 2026-08-12 20:07:59 +08:00
m4 aae8d0a379 feat: workspace file references, read-file-images middleware, image model enabled flag
In-progress work committed to unblock the config import/export plan:
- prompts: FILE_REFERENCES section for workspace-relative file citation
- backends: resolve quoted virtual absolute paths onto the sandbox workspace
- middleware: read_file_images middleware; message_budget extensions
- image_gen/model_registry: image model 'enabled' flag refactor
- memory/launch, gateway/background_runs, tools/image follow-ons
- scripts: dev_backend.sh, release.sh
- tests for the above
2026-08-12 19:43:35 +08:00
m4 f3ca381ab3 feat(release): local release script with unified version bump and Gitea publish
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-08-12 17:25:43 +08:00
m4 8f6a568646 feat(update): system update/status/rollback HTTP endpoints with system:write scope 2026-08-12 16:42:10 +08:00
m4 91e2a87be5 feat(update): rollback version list with BREAKING-DB truncation 2026-08-12 16:34:16 +08:00
m4 03771df508 feat(update): standalone updater runner (wait/install/restart/result) 2026-08-12 16:30:49 +08:00
m4 90e773b30b feat(update): plan builder, plan lock and updater spawn 2026-08-12 16:27:39 +08:00
m4 40b896bcb5 feat(update): deployment detection and install command builder 2026-08-12 16:24:17 +08:00
m4 174f03b92d feat(update): stage artifacts under config dir and reuse verified local files 2026-08-12 16:21:45 +08:00
m4 194402fc88 fix(usage): stop sharing EVOSCIENTIST_DEPLOYMENT_ID with scope partitioning
The usage identity exported the same variable the scope registry reads to
partition workspace scopes, so a backend started with the usage environment
(61d1b61b) could not see scopes provisioned under the workspace-derived id
(5a882492) and every scope lookup 404'd. Usage attribution now reads
EVOSCIENTIST_USAGE_DEPLOYMENT_ID; the scope side keeps the original variable.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-30 22:08:21 +08:00
m4 c6efdaa13f docs(webui): thinking-timer design for the chat message area
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-30 11:11:33 +08:00
m4 08aa0d0e05 feat(registry): mutually-exclusive sampling_override frozen into run snapshots 2026-07-30 10:08:26 +08:00
m4 8eff551bb6 docs(registry): implementation plan for mutually-exclusive sampling_override
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-30 09:48:34 +08:00
m4 862c1e9743 docs(registry): supersede flat temperature/top_p overrides with mutually-exclusive sampling_override
temperature and top_p cannot be set together; the flat-fields design
allowed both. Replaced by a discriminated union (default | temperature
| top_p) where overriding one omits the other from the request entirely.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-30 09:29:06 +08:00
m4 d7711484ef feat(registry): expose generation defaults on selectable models 2026-07-28 17:42:16 +08:00
m4 ba7d908276 feat(registry): freeze per-thread temperature/top_p overrides into run snapshots
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-28 17:23:26 +08:00
m4 e57ecd4588 docs(plan): per-thread temperature/top_p override implementation plan
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-28 16:59:37 +08:00
m4 261845830d docs(spec): per-thread temperature/top_p override design
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-28 16:17:16 +08:00
m4 8efd4ad0ab fix(image-gen): wrap corrupt-config GET in the 422 error envelope 2026-07-24 09:45:57 +08:00
m4 01e674e1bc feat(image-gen): add /api/image-generation config endpoints with masked keys 2026-07-24 09:36:29 +08:00
m4 7f26ecc19a fix(image-gen): sanitize config validation errors; widen expected_revision type
load_image_generation_settings now re-raises pydantic ValidationError as a
sanitized ImageGenError carrying only field locations and error types, so a
mis-indented config.yaml can never echo a literal API key into agent-visible
errors. Also widen save_registry's expected_revision annotation to
int | None to match the http_api caller (value remains ignored).

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-24 07:51:36 +08:00
m4 395baab7d5 test(registry): add IMAGE_MODEL_NOT_CHAT_MODEL to error taxonomy mirror 2026-07-24 07:44:37 +08:00
m4 ccf4173990 feat(registry): reject image-only models in chat model saves 2026-07-24 07:36:34 +08:00
m4 a57c52c676 feat(image-gen): add image-artist skill for generation workflow 2026-07-24 07:26:23 +08:00
m4 384bc13a5b feat(image-gen): add generate_image/edit_image agent tools
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-23 22:36:06 +08:00
m4 5622b40cf3 feat(image-gen): add service layer with safe artifact saving 2026-07-23 22:10:14 +08:00
m4 802f71bd46 feat(image-gen): add Gemini (Imagen) image adapter 2026-07-23 21:53:40 +08:00
m4 cc9dfb1cc9 fix(image-gen): address review findings in OpenAI image adapter
Scope download Authorization header to the provider origin, translate
httpx errors in _download/edit into safe ImageGenError messages, and
prevent entry params from clobbering core payload keys; also harden
strip_data_uri against malformed values and drop dead code.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-23 21:44:57 +08:00
m4 12e4f34005 feat(image-gen): add OpenAI-compatible image adapter with responses fallback 2026-07-23 21:28:38 +08:00
m4 b0f9a8d785 feat(image-gen): add image_generation config section and model detection
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-23 21:03:50 +08:00
m4 15cc389b3d feat(registry): make expected_revision optional and ignored in save API
Regenerate the checked-in OpenAPI export to match the relaxed schema.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-23 20:55:52 +08:00
m4 3e67e64067 refactor(registry): make saves last-write-wins, ignore expected_revision
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-23 20:47:08 +08:00
m4 bb9bed82e1 feat(runtime)!: remove auxiliary model role, resolve all roles from run snapshot
ModelRole collapses to "primary": every role (main, tool selector, memory
agents, subagents, summarizer) resolves to the snapshot's frozen primary
model, per design 6.1/8.3 — users typically configure a single usable LLM,
so compile-time auxiliary bindings were bypassing run snapshots and
mis-attributing usage. Legacy auxiliary keys in stored snapshots, registry
JSON, and thread metadata are tolerated on read and dropped.

BREAKING CHANGE: ThreadModelSelection no longer carries an auxiliary ref;
snapshot selection_hash is computed over {primary, reasoning_effort} only;
ConfigurableModelMiddleware(role="auxiliary") is rejected.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-23 10:18:44 +08:00
m4 1a01fb5d74 docs(env): drop stale LLM-key and provider-admin-token entries from .env.example
Provider credentials now live exclusively in the Model Registry and the
x-evoscientist-admin-token / provider-admin-token mechanism was removed;
no code reads these variables anymore.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-21 21:14:31 +08:00
m4 087781556b fix(runtime): serve Config API in bootstrap and verify snapshot issuer by registered deployment set
Graph construction no longer raises on a bootstrap registry: build paths
bind a shared RegistryNotReadyChatModel placeholder that fails every call
with MODEL_REGISTRY_NOT_READY, so langgraph dev serves the Config API for
first-time configuration while run creation stays forbidden.

Run snapshot binding no longer compares configurable
'workspace_deployment_id' (the workspace-isolation scope id) against the
snapshot's issuing deployment — a mismatch that made every BFF run fail
with SNAPSHOT_NOT_FOUND. SnapshotService.get_for_run verifies thread_id
equality plus membership in the platform-registered deployment set
(local_deployment_id + webui_delegation_public_keys entries).

Blocking I/O moved off the event loop for langgraph dev's blockbuster:
Config API authentication (store mkdir/chmod, config.yaml read, jti
registration) and the message-budget snapshot read now run in threads,
with the immutable snapshot cached per run.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-21 20:21:32 +08:00
m4 a1bfbd92ca chore: untrack accidentally committed docx
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-21 18:10:40 +08:00
m4 421a664336 feat(runtime)!: complete legacy removal, local snapshot entries, and TTL cleanup
- Remove legacy provider profiles, admin-token auth, /model command,
  model picker widget, and config.yaml LLM fields (design doc section 10)
- Wire CLI/channels/cron and async sub-agents through the local snapshot
  entry; run creation rejects model config outside runtime_snapshot_id
- Add periodic run-snapshot TTL cleanup to the config service lifespan
- Isolate tests from the real config dir and activate the registry where
  run/model paths fail closed in bootstrap

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-21 18:10:23 +08:00
m4 57176b359a feat(runtime)!: switch middleware and agent factories to snapshot-driven models
Replace config.yaml-driven model selection with registry snapshot resolution
across the runtime chain:

- ConfigurableModelMiddleware reads configurable["runtime_snapshot_id"] only;
  model/model_provider overrides are rejected with MODEL_CONFIG_OUTSIDE_SNAPSHOT
- MessageBudgetMiddleware derives budgets from snapshot reserves
  (system/tools/attachments) and re-resolves the summarizer per snapshot
- Agent factory and subagent factory resolve models via SnapshotRuntime
  (auxiliary/tool_selector/scheduler -> defaults.auxiliary ?? defaults.primary)
- Remove ModelFallbackMiddleware, /model-fallback command, and fallback chain
- Add model_registry/runtime.py SnapshotRuntime glue layer

Legacy config.yaml LLM fields, /model command, and llm/models.py remain for
Task 7. Report: .superpowers/sdd/briefs/task-6-report.md
2026-07-21 12:54:35 +08:00
m4 b2660fc38c feat(model-registry): add provider test API and guarded verification recording
POST /api/model-registry/test (model_config:test) runs the section 9.4
flow: resolve_for_test, per-test credential resolution, build_chat_model
with both safe clients, one minimal chat call, and per-capability probes
(tools/structured_output/vision) whose failures only mark that capability
unverified. Results upsert the model_verifications five-tuple inside a
BEGIN IMMEDIATE transaction that re-checks the registry revision and
configuration hash, returning 409 MODEL_CONFIGURATION_CHANGED on any
concurrent change. resolve_for_test now also relaxes the passing-
verification gate, which the provider test itself produces.
effective_request_options reuses the redacted adapter.build_request
output; the OpenAPI contract and checked-in openapi.json are updated.
2026-07-21 11:39:43 +08:00
m4 940db565b3 fix(model-registry): close review gaps in save-time validation
- add the missing rule-8 counterexample test: declared capabilities
  exceeding the adapter protocol are rejected with
  CAPABILITY_UNSUPPORTED_BY_ADAPTER (all eight section 9.2 checks now
  have at least one negative test)
- raise CREDENTIAL_NOT_CONFIGURED explicitly in _check_enabled_model
  when a required credential reference is null instead of relying on
  resolve_parameters call ordering
2026-07-21 09:38:37 +08:00
m4 dbb6b7abde feat(model-registry): add delegation-JWT auth and config/snapshot HTTP API
- BFF service token (constant-time, plaintext or SHA-256 hash) plus
  X-Evo-Actor delegation JWT verification (ES256/RS256, iss/aud, <=60s
  lifetime, required claims, thread binding) with atomic jti anti-replay
- Config API: GET/PUT /api/model-registry, credential rotation endpoint,
  GET /api/models selector; PUT runs the section 9.2 save-time checks
  inside the registry write transaction after credential writes
- Snapshot API: create/bind/delete routes delegating to SnapshotService
  with thread/deployment binding checks and 9.5 unified error payloads
- Platform security config loader (config.yaml fields), OpenAPI export
  (scripts/export_model_registry_schema.py -> model_registry/openapi.json)
- Mount new routes in langgraph_dev/http.py; retire the legacy
  GET /api/models and POST /api/runtime-snapshots handlers
- Declare PyJWT>=2.8 (previously transitive); extend the 9.5 error code
  table with the HTTP-layer codes (400/401/403/422/500)
2026-07-21 09:28:55 +08:00
m4 c8c46eab16 fix(model-registry): make snapshot abort atomic against concurrent bind
The abort path was read-then-write with an unconditional UPDATE, so a bind
committing between the two calls was clobbered back to aborted, losing its
langgraph_run_id. Add a conditional store-level abort_run_snapshot
(prepared-only UPDATE, rowcount-checked) and re-read on a lost race, matching
the bind loop. Also pin the inherit selection_hash test to a hardcoded
SHA-256 literal instead of reimplementing the serialization in the test.
2026-07-21 08:45:33 +08:00
m4 0cc995eb80 docs(model-registry): add task 4 resolver and snapshot service report 2026-07-21 08:34:50 +08:00
m4 b1233d42dc feat(model-registry): add ModelRegistryResolver and run snapshot service
Resolver (8.1): validates provider/model/credential/capability/limits and
the 6.5 four-mode input budget, freezes ResolvedModelConfig; resolve_for_test
relaxes only the enabled-visibility check (9.4); compute_availability is the
single 4.3 six-state judgement (stale beats configured, selectable only when
enabled).

SnapshotService (8.2, shared by the Task 5 HTTP API and Task 7 local entry):
freezes both roles' full ResolvedModelConfig with adapter spec revision,
fixed reserves, capabilities, and credential revisions; selection-hash
idempotency with pre-resolution semantics; prepared(15min)/bound(+24h)/
expired/aborted lifecycle with atomic bind; binding-checked reads that
revalidate frozen spec revisions; per-call credential resolution against the
frozen revision with no in-process secret cache (5.2); public diagnostic
view limited to the 8.2 safe subset.

Store gains additive helpers (credential pointer lookup, verification
listing, active-triplet lookup, conditional bind, due-expiry sweep) and the
taxonomy gains SNAPSHOT_NOT_FOUND (404) for missing snapshots.
2026-07-21 08:33:12 +08:00
m4 b2e28249fd fix(model-registry): inject async safe clients and split ollama transports
Review fixes for the Task 3 contract layer:

- build_chat_model now accepts http_async_client alongside http_client
  (at least one required) and wires it into ChatOpenAI
  (http_async_client), ChatAnthropic (seeded _async_client), and
  ChatOllama (async_client_kwargs transport), closing the unsafe
  default-async-client gap.
- ChatOllama safe transports move from the shared client_kwargs to
  sync_client_kwargs/async_client_kwargs; langchain-ollama merges shared
  kwargs into both clients, which poisoned the async client with a sync
  transport and crashed ainvoke.
- Unsupported parameters now actually execute the contract-declared
  normalizer (reject_non_auto) instead of a hardcoded raise, with a
  fallback rejection if a normalizer would let a value through.
- build_chat_model rejects overlapping client_options/request_options
  keys instead of silently overwriting.
2026-07-20 22:40:33 +08:00
m4 af4ae1aef5 feat(model-registry): add adapter parameter contracts and build_chat_model factory
Add the Task 3 parameter contract layer (design doc 6.1-6.4):

- adapters.py: versioned built-in contracts for the five phase-1
  adapters plus the openai-compatible/glm-5.2 model-specific contract
  (verbatim section 6.2 values); exact > longest glob > generic
  matching with spec_revision pinning; resolve_parameters implementing
  the section 6.1 inherit/omit semantics, contract validation with
  stable error codes, and named normalizers (identity,
  clamp_to_model_limit, omit_when_none, omit_when_auto,
  reject_non_auto); Adapter.build_request as the single entry point
  mapping ResolvedModelConfig to {client_options, request_options};
  compute_effective_capabilities (protocol AND declared AND verified).
- factory.py: build_chat_model(resolved_config, http_client, *,
  credential=None) with no **kwargs and no setdefault merging; injects
  the safe HTTP client into ChatOpenAI/ChatAnthropic/ChatOllama, never
  reads provider API-key environment variables, and strips the
  OLLAMA_API_KEY authorization header for mode=none adapters.
- tests: per-adapter request-capturing fakes plus an httpx.MockTransport
  outbound capture proving registry resolution matches the wire request.
2026-07-20 22:17:11 +08:00
m4 c46ae17084 feat(model-registry): add EndpointPolicy and SafeHttpTransport SSRF defenses
EndpointPolicy validates provider base URLs (section 4.3): public https
endpoints with hostname and optional port pass; loopback, private,
link-local, multicast, unspecified, and cloud-metadata addresses are
denied unless the normalized URL exactly matches a registered
development_endpoints entry (no prefix or wildcard matching). URLs with
user info, fragments, or non-http(s) schemes are rejected with the new
stable 422 code ENDPOINT_NOT_ALLOWED.

SafeHttpTransport is the single network egress for adapters: a custom
httpcore NetworkBackend resolves DNS under control on every connect
(retries included), filters denied ranges, and connects directly to the
selected IP, while TLS SNI/certificate checks and the HTTP Host header
keep the original hostname. Redirects and env proxies are disabled;
every request origin re-passes URL-layer validation before any I/O.
2026-07-20 21:25:15 +08:00
m4 c21fc0a272 feat(model-registry): add RegistryV4 schema, SQLite store, and unified error codes
New independent subpackage EvoScientist/model_registry implementing the
frozen unified model configuration design (v1.1.0, sections 4.2, 4.3,
5.1, 8.2, 9.5):

- schemas.py: single Pydantic v2 RegistryV4 schema (ModelRef, ProviderConfig,
  ModelConfig, AuthConfig) plus AdapterParameterSpec/AuthSpec/ParameterRule,
  ModelAvailability, ResolvedModelConfig (no secrets), CredentialStatus
- errors.py: all 23 stable error codes from the 9.5 table with HTTP status
  mapping and the unified {code, message, details, request_id} payload
- store.py: ModelRuntimeStore over model-runtime.sqlite3 (0700 dir, 0600
  file, WAL, foreign keys, busy_timeout) with BEGIN IMMEDIATE revision CAS,
  bootstrap->active atomic transition, immutable credential versions with
  masked status, model_verifications upsert, run_runtime_snapshots partial
  unique index, delegation_jtis, and shared-storage lock probe
- hashing.py: normalized SHA-256 configuration_hash

No existing module behavior changed. 89 new tests; full suite passes
(2922 passed, 10 skipped).
2026-07-20 20:43:28 +08:00
m4 8a0ab17936 chore: baseline WIP before unified model configuration implementation
Pre-existing uncommitted work (runtime snapshots, message budget middleware) preserved as baseline.
2026-07-20 20:15:38 +08:00
m4 38668c4ce5 feat: add workspace isolation and provider administration 2026-07-19 12:17:18 +08:00
m4 7a3fcc7c8e Merge remote-tracking branch 'upstream/main'
# Conflicts:
#	README.md
#	uv.lock
2026-07-13 09:46:07 +08:00
m4 e0acc6155e feat: improve WebUI run recovery
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
2026-07-10 17:35:44 +08:00
199 changed files with 34979 additions and 10334 deletions
+37 -28
View File
@@ -1,32 +1,41 @@
# EvoScientist — cp .env.example .env && fill in your keys
# LLM provider (pick at least one)
ANTHROPIC_API_KEY= # console.anthropic.com
OPENAI_API_KEY= # platform.openai.com
GOOGLE_API_KEY= # aistudio.google.com/api-keys
NVIDIA_API_KEY= # build.nvidia.com
# Direct providers (optional)
MINIMAX_API_KEY= # platform.minimaxi.com (China, default) or platform.minimax.io (Global)
MINIMAX_BASE_URL= # https://api.minimaxi.com/anthropic (default) or https://api.minimax.io/anthropic
ZHIPU_API_KEY= # open.bigmodel.cn (智谱)
VOLCENGINE_API_KEY= # volcengine.com (火山引擎)
DASHSCOPE_API_KEY= # dashscope.aliyuncs.com (阿里云)
MOONSHOT_API_KEY= # platform.moonshot.cn (月之暗面)
KIMI_API_KEY= # kimi.com/code (Kimi 代码计划)
# Aggregator platforms (optional)
SILICONFLOW_API_KEY= # siliconflow.cn
OPENROUTER_API_KEY= # openrouter.ai
# Custom endpoints (optional)
CUSTOM_OPENAI_API_KEY= # OpenAI-compatible endpoint
CUSTOM_OPENAI_BASE_URL=
CUSTOM_ANTHROPIC_API_KEY= # Anthropic-compatible endpoint
CUSTOM_ANTHROPIC_BASE_URL=
# Local models (optional)
OLLAMA_BASE_URL= # http://localhost:11434 (default)
#
# LLM providers, models, and API keys are managed exclusively through the
# Model Registry (WebUI 大模型配置 / Config API); no provider credential is
# read from environment variables. See
# docs/unified-model-configuration-architecture.md.
# Web search (optional)
TAVILY_API_KEY= # app.tavily.com
# WebUI conversation workspace policy. EVOSCIENTIST_WORKSPACE_DIR is the
# deployment root, not a per-conversation directory. In isolated modes each
# conversation is stored under <root>/.evoscientist/conversations/<scope-id>/.
#
# EVOSCIENTIST_WORKSPACE_ISOLATION accepts exactly:
# - legacy: all WebUI conversations share the deployment root. Compatibility
# rollback only; files are visible to every conversation using this deployment.
# - optional: default. New WebUI conversations receive isolated scope folders;
# missing Registry/token/scope fails the request instead of silently sharing.
# - required: isolated scopes plus strict runtime validation. It requires a
# completed cutover and a verified OCI executor; no legacy fallback exists.
#
# This is a deployment-startup security setting. Change it only during a
# maintenance window, restart backend and WebUI afterwards, and never use it to
# convert an existing conversation between shared and isolated directories.
EVOSCIENTIST_WORKSPACE_DIR=
EVOSCIENTIST_WORKSPACE_ISOLATION=optional
# Required mode supports only a single-host Registry topology in v1.
EVOSCIENTIST_SCOPE_REGISTRY_TOPOLOGY=single-host
# Required mode: use a pinned image digest, preserve single-host topology, and
# keep the Code Interpreter disabled unless its scoped implementation is enabled.
# Do not put EVOSCIENTIST_BACKEND_SERVICE_TOKEN here for a same-host `EvoSci
# deploy`: it is generated and passed privately at startup.
# EVOSCIENTIST_WORKSPACE_ISOLATION=required
# EVOSCIENTIST_STRICT_EXECUTOR=oci
# EVOSCIENTIST_STRICT_EXECUTOR_IMAGE=registry.example/evoscientist-runtime@sha256:replace-with-verified-digest
# EVOSCIENTIST_STRICT_CODE_INTERPRETER=disabled
# Conversation workspace isolation retention defaults (used by workspace_maintenance.py).
EVOSCIENTIST_DRAFT_WORKSPACE_TTL_HOURS=24
EVOSCIENTIST_WORKSPACE_TRASH_RETENTION_DAYS=7
+26 -1
View File
@@ -16,11 +16,16 @@ jobs:
strategy:
fail-fast: false
matrix:
os: [ubuntu-latest, windows-latest]
os: [ubuntu-latest, macos-latest, windows-latest]
python-version: ["3.11", "3.12"]
runs-on: ${{ matrix.os }}
steps:
- uses: actions/checkout@v5
- name: Check out shared usage fixtures
uses: actions/checkout@v5
with:
repository: EvoScientist/EvoScientist-WebUI
path: EvoScientist-WebUI
- uses: astral-sh/setup-uv@v6
with:
python-version: ${{ matrix.python-version }}
@@ -29,3 +34,23 @@ jobs:
run: uv sync --dev
- name: Run pytest
run: uv run pytest -v --timeout=30
env:
EVOSCIENTIST_USAGE_FIXTURES: ${{ github.workspace }}/EvoScientist-WebUI/docs/schemas/fixtures
usage-spool-benchmark:
timeout-minutes: 15
strategy:
fail-fast: false
matrix:
os: [ubuntu-latest, macos-latest, windows-latest]
runs-on: ${{ matrix.os }}
steps:
- uses: actions/checkout@v5
- uses: astral-sh/setup-uv@v6
with:
python-version: "3.11"
cache-dependency-glob: "**/pyproject.toml"
- name: Install dependencies
run: uv sync --dev
- name: Verify durable spool latency
run: uv run python scripts/benchmark_usage_spool.py
+5
View File
@@ -48,3 +48,8 @@ conversation_history/
*meals/
botpy.log
large_tool_results/
runs/
# local runtime artifacts (scope tokens, control DBs)
.evoscientist/
.release-state.json
+132
View File
@@ -0,0 +1,132 @@
# Task 1 报告 — RegistryV4 schema + SQLite 存储层 + 统一错误码
## 实现摘要
按简报与设计文档 v1.1.0(4.2、4.3、5.1、8.2、9.5 节)在仓库内新建独立子包
`EvoScientist/model_registry/`,未改动任何现有文件(顶层 `EvoScientist/__init__.py`
采用惰性导出,无需修改)。
1. **`schemas.py` — 唯一 RegistryV4 schema(Pydantic v2)**
- `ProviderId`/`ModelKey`/`CredentialId`/`AdapterId` 共用锚定全匹配模式
`\A[a-z0-9][a-z0-9._-]{0,63}\z`(pydantic 的 pattern 约束默认是子串搜索,必须锚定)。
- `upstream_model_id` 仅限长 1–300、保留大小写;`ValidatedEndpoint` 限长 2048 且要求
http(s) scheme(EndpointPolicy 属后续任务)。
- `RegistryV4 { version: Literal[4], revision: PositiveInt, state, defaults, providers }`;
schema 级校验:Provider ID 唯一、Provider 内模型 key 唯一、defaults 引用必须存在、
`active` 状态下 primary 非空且引用已启用模型(auxiliary 同理)。
- `ProviderConfig.runtime`:`timeout_seconds` [10,600](缺省 120)、`max_retries` [0,5]
(缺省 2)、`default_temperature` [0,2]|null、`default_top_p` (0,1]|null、
`default_reasoning_effort` 缺省 `auto`。
- `ModelConfig.runtime`:`limit_mode` combined 要求 `context_window_tokens`、input_only
要求 `max_input_tokens`(model_validator);`min_effective_input_tokens` 缺省 4096、
下限 1024;三个 `fixed_*_reserve_tokens` 非负;temperature/top_p/reasoning_effort 与
`declared_capabilities` 按 4.3 定义。
- `AuthConfig`:`mode=none` 时 `credential_id` 必须为 null(model_validator)。
- 6.2/6.4/9.1 结构:`AdapterParameterSpec`(含 `connection` 可选块,对应 6.2 示例
`chat_model`/`model_field`/`base_url_field`)、`AuthSpec`、`ParameterRule`、
`ModelAvailability`(`VerificationInfo`)、`ResolvedModelConfig`
(`auth_ref={mode, credential_id?, credential_revision?}`,无任何 secret 字段)、
`CredentialStatus`、`CredentialWrite`(9.2 的 `operation: replace`)。
2. **`errors.py` — 统一错误码与载荷**
- 23 个稳定错误码常量 + `ERROR_HTTP_STATUS` 映射,覆盖 9.5 总表全部 17 组
(409×6、401×1、404×1、422×15)。
- `ErrorPayload {code, message, details:[{path, code}], request_id}` 与 9.5 结构一致;
`ModelRegistryError` 携带 code/`http_status`/`payload()`,未知 code 直接拒绝。
3. **`store.py` — `ModelRuntimeStore`**
- 数据库 `<config_dir>/model-runtime.sqlite3`(默认 `~/.config/evoscientist`,可注入);
目录 0700、文件 0600、WAL、外键、`busy_timeout=30000`。
- 手写 DDL(`CREATE TABLE IF NOT EXISTS`,无 alembic):`registry_state`(单行)、
`credential_pointers`、`credential_versions`((credential_id, revision) 主键)、
`model_verifications`(五元组主键,upsert 只留最近一次)、`run_runtime_snapshots`
(含部分唯一索引 `UNIQUE(deployment_id, thread_id, run_request_id)
WHERE status IN ('prepared','bound')`)、`delegation_jtis`。
- `load_registry()` 无行时返回 bootstrap/revision=1 空 RegistryV4;
`save_registry(expected_revision=, registry=, credential_writes=)` 在
`BEGIN IMMEDIATE` 事务内校验 revision(不符抛 `REGISTRY_REVISION_CONFLICT`)、
写入不可变凭据版本、registry revision+1;首次同时具备已启用模型+有效 primary+已配置
凭据(或 `mode=none`)时原子转为 `active`;任一失败整体回滚。每次保存强制执行
9.2 第 7 条(defaults 必须引用已启用模型,违反抛 `MODEL_DISABLED`)。
- 凭据:`write_credential_version`(递增 revision;重写相同当前密钥幂等返回原
revision)、`resolve_credential`(不存在/已销毁抛
`RUN_CREDENTIAL_REVISION_UNAVAILABLE`)、`retire_credential_version`、
`credential_status`(hint 末 4 位 `...abcd`,短于 4 字符的密钥 hint 为 null 绝不泄露;
绝不返回明文)。
- 快照:`insert_run_snapshot`/`set_run_snapshot_status`/`get_run_snapshot`
(状态机 prepared|bound|expired|aborted,部分唯一索引行为由测试覆盖)。
- `check_shared_storage()`:探测 `BEGIN IMMEDIATE` 写锁能力,失败抛
`SharedStorageError`(多节点不共享持久卷时启动失败)。
4. **`hashing.py` — `configuration_hash(provider, model)`**
- 覆盖 adapter、base_url、upstream_model_id、Provider 与 Model 全部运行参数(含声明
能力与限制),`json.dumps(sort_keys=True, separators=(",", ":"))` 规范化后 SHA-256。
## 文件清单
新增(无修改既有文件):
- `EvoScientist/model_registry/__init__.py`
- `EvoScientist/model_registry/schemas.py`
- `EvoScientist/model_registry/errors.py`
- `EvoScientist/model_registry/store.py`
- `EvoScientist/model_registry/hashing.py`
- `tests/test_model_registry_schemas.py`
- `tests/test_model_registry_store.py`
- `.superpowers/sdd/briefs/task-1-report.md`(本文件)
## 测试命令与输出
TDD 流程:先写两个测试文件并确认失败(`ModuleNotFoundError: No module named
'EvoScientist.model_registry'`),再实现。
```
$ .venv/bin/python -m pytest tests/test_model_registry_schemas.py tests/test_model_registry_store.py -x -q
........................................................................ [ 80%]
................. [100%]
89 passed in 0.22s
```
全量回归(无既有失败,无回归):
```
$ .venv/bin/python -m pytest tests/ -x -q
........sssss....................................... [100%]
2922 passed, 10 skipped, 1 warning in 73.00s
```
(warning 为 `test_langgraph_dev_http.py` 的 StarletteDeprecationWarning,既有、与本任务无关。)
lint 与格式:
```
$ .venv/bin/ruff check EvoScientist/model_registry tests/test_model_registry_schemas.py tests/test_model_registry_store.py
All checks passed!
$ .venv/bin/ruff format --check ... # 已格式化
```
## 自我审查发现(已处理)
1. **测试副作用污染真实配置目录**:初版 `test_default_config_dir` 未注入路径,运行时在
真实 `~/.config/evoscientist/` 创建了空的 `model-runtime.sqlite3`。已确认该库所有表
为空(确为测试副产物)后删除(含 -wal/-shm),并把测试改为 monkeypatch
`store.DEFAULT_CONFIG_DIR` 到 tmp_path,此后测试不再触碰真实 home。
2. **共享 `Field()` 实例**:初版 `_TEMPERATURE`/`_TOP_P` 在三个模型间复用同一 FieldInfo,
已改为各字段独立 `Field(...)`,规避 pydantic 共享元数据的潜在风险。
3. **docstring 混入中文**:`save_registry` 一处 docstring 误用中文,已改为英文以符合
仓库注释惯例。
4. **并发 CAS 断言**:`sorted()` 大小写排序导致误报,改为不区分大小写排序。
5. ruff 修复:`datetime.UTC` 别名、导入排序、`pytest.raises` 增加 `match=`。
## 遗留疑虑
1. **`connection` 块为可选**:6.2 正文的"至少包含"清单未列 `connection`,但 glm-5.2
示例契约包含它。schema 将其建模为可选字段,Task 2 落地五种 Adapter 契约时若确认
每个契约都有 connection,可考虑收紧为必填。
2. **激活就绪判定中的"满足 AuthSpec"**:4.3 要求激活时认证状态满足 AuthSpec,但
Adapter 契约属 Task 2。当前 store 仅做存储层判定(`mode=none` 或凭据已配置);
AuthSpec 级别校验(如 credential_kind 匹配)需在 Task 2/3 的 API 层补充。
3. **幂等语义解释**:简报称 `write_credential_version` 为"幂等 prepare",实现为"重写
与当前版本完全相同的密钥时返回现有 revision";不同密钥轮换仍产生新 revision。若
后续任务对幂等键有不同约定(如客户端提供 request id),需再对齐。
4. **`save_registry` 参数为 keyword-only**:与简报签名
`save_registry(expected_revision, registry, credential_writes=[])` 语义一致,仅调用
形式略异。
5. `MODEL_DISABLED` 用于 defaults 引用未启用模型的保存错误(9.2 第 7 条属 422 校验,
总表无更贴切码);引用不存在模型由 schema 层先行拒绝,store 内同名分支仅作防御。
+89
View File
@@ -0,0 +1,89 @@
# Task 4 报告 — ModelRegistryResolver + 运行快照服务
- 状态:DONE
- 分支:`feature/unified-model-config`
- Commit:`b1233d4` `feat(model-registry): add ModelRegistryResolver and run snapshot service`
- 规格来源:`docs/unified-model-configuration-architecture.md` v1.1.1,章节 4.3 / 5.2 / 6.1 / 6.4 / 6.5 / 8.1 / 8.2
## 交付物
### 新建 `EvoScientist/model_registry/resolver.py`
`ModelRegistryResolver(store, *, specs=None)`:
- `resolve(model_ref, role="primary", *, registry=None) -> ResolvedModelConfig`,校验顺序:
1. Provider 存在(否则 404 `MODEL_NOT_FOUND`)且 enabled(否则 422 `MODEL_NOT_AVAILABLE`);
2. 模型存在(`MODEL_NOT_FOUND`)且 enabled(否则 422 `MODEL_DISABLED`);
3. `find_adapter_spec` 匹配契约:未开放 Adapter 抛 `ADAPTER_NOT_SUPPORTED`,无匹配契约抛 `MODEL_NOT_AVAILABLE`(只可保持 configured);
4. `resolve_parameters` 重放保存期契约校验(AuthSpec、能力、参数范围/互斥);
5. 凭据:`credential_required` 时读 `credential_pointers.current_revision` 写入 `auth_ref`(指针缺失 → `CREDENTIAL_NOT_CONFIGURED`),**从不读 `secret_value`**;`mode=none` 冻结 `{mode: none}`;
6. 当前验证五元组(`configuration_hash` + 当前凭据 revision + 当前 `adapter_spec_revision`)必须存在且 `passed`,否则 `MODEL_NOT_AVAILABLE`;
7. `limits_status=confirmed`,否则 `MODEL_LIMITS_UNCONFIRMED`;
8. 预算(6.5):`combined` → `context_window - max_output`;`input_only` → `max_input`;四种工具/附件组合扣除固定预留后均须 ≥ `min_effective_input_tokens`,否则 `CONTEXT_BUDGET_UNSATISFIABLE`。冻结 `resolved_input_limit`、三个固定预留与 base `message_budget`(仅扣系统预留;逐次调用的预算由 MessageBudgetMiddleware 按 6.5 公式重算,Task 6 接线);
9. 能力按唯一规则 `protocol AND declared AND verified` 计算并冻结。
- `resolve_for_test(model_ref)`:仅放宽"模型必须 enabled",其余校验(含 Provider enabled、凭据、验证记录、限制、预算)全部执行(9.4)。
- `compute_availability(registry, verifications, *, credential_revisions=None) -> list[ModelAvailability]`:4.3 判定顺序 `unavailable → enabled → verified → verification_failed → verification_stale → configured`,首个命中生效;`verification_stale` 优先于 `configured`;`selectable` 仅 `enabled` 为 true。`verification` 报告当前记录的 passed/failed 或最新旧记录的 stale;`effective_capabilities` 仅在当前记录 passed 时非全 false。reason_code 取值:`PROVIDER_DISABLED` / `NO_ADAPTER_CONTRACT` / `MODEL_DISABLED` / 记录的错误码 / `VERIFICATION_STALE`。
- 6.1 角色映射实现为 `snapshots.config_for_role`(映射对象是快照而非 Registry,故放在快照模块):`primary → snapshot.primary`;`auxiliary`/`summary`/`tool_selector → snapshot.auxiliary ?? snapshot.primary`;未知角色 `ValueError`。Resolver 不猜测其他映射。
### 新建 `EvoScientist/model_registry/snapshots.py`
`SnapshotService(store, resolver)`(Task 5 HTTP API 与 Task 7 本地入口共用):
- `create(SnapshotCreateRequest)`:bootstrap → 422 `MODEL_REGISTRY_NOT_READY`;`primary=null`(inherit)解析 Registry `defaults.primary`,`auxiliary=null` 解析 `defaults.auxiliary`;冻结两个完整 `ResolvedModelConfig`(含 `adapter_spec_revision`、固定预留、能力、`auth_ref` 凭据版本)与 `registry_revision`、`model_selection_revision` 写入 `payload_json`,绝无 secret。`selection_hash` = 解析前 `{primary, auxiliary}`(inherit 以 null 参与)固定字段序 JSON 的 SHA-256;`model_selection_revision` 仅审计、不参与哈希。
- 幂等:同一三元组哈希相同返回原快照(`created=False`,200 语义),不同抛 409 `RUN_REQUEST_CONFLICT`;并发创建撞部分唯一索引时回退到同一幂等比较。`expired`/`aborted` 不占三元组,同三元组可重建(`created=True`,201)。
- `bind(snapshot_id, langgraph_run_id)`:`prepared→bound` 一次(条件 UPDATE 保证原子),并把 `expires_at` 延长到 +24h;相同 run id 重复 bind 幂等成功,不同值抛 `SNAPSHOT_ALREADY_BOUND`;终态抛 `SNAPSHOT_EXPIRED`;不存在抛 `SNAPSHOT_NOT_FOUND`。
- `abort(snapshot_id)`:仅 `prepared→aborted`;已 bound 抛 `SNAPSHOT_ALREADY_BOUND`;expired 抛 `SNAPSHOT_EXPIRED`;重复 abort 幂等成功。
- `get(snapshot_id, *, deployment_id, thread_id)`:绑定关系不匹配按 `SNAPSHOT_NOT_FOUND` 失败(不跨线程/部署泄露存在性);终态抛 `SNAPSHOT_EXPIRED`;读取时对两个冻结配置逐一重校验 `adapter_spec_revision` 仍存在(经 Resolver 的 spec 列表),已移除抛 `ADAPTER_NOT_SUPPORTED`,绝不静默替换。
- `cleanup_expired(now)`:到期 prepared(创建时 TTL **15 分钟**)与 bound(bind 时 **+24 小时**)置为终态 `expired`,返回迁移的 ID。
- `resolve_snapshot_credential(snapshot, role) -> str`:经 6.1 角色映射取冻结 `auth_ref`,按冻结 `credential_revision` 每次从凭据存储解析(**无进程内密钥缓存**);版本销毁抛 `RUN_CREDENTIAL_REVISION_UNAVAILABLE`;`mode=none` 返回 `""`。
- `public_snapshot_view(snapshot)`:仅 8.2 示例字段(`snapshot_id`、`registry_revision`、primary/auxiliary 的 `provider_id`/`model_key`/`adapter_spec_revision`/runtime 五项),无 base_url、无 secret。
### Task 1-3 文件的增补(均为纯新增,未改动任何既有对外行为)
- `errors.py`:新增 `SNAPSHOT_NOT_FOUND`(404)。9.5 承认"404 资源不存在"类别但总表只有 `MODEL_NOT_FOUND`(语义为 ModelRef);快照缺失需要独立稳定码,属对总表的增补,已同步更新 Task 1 的 `ALL_ERROR_CODES` 测试(仅加一行)。
- `store.py`:新增 `current_credential_revision`、`list_model_verifications`(可用性判定的输入)、`find_active_run_snapshot`(三元组幂等查询)、`bind_run_snapshot`(`WHERE status='prepared'` 条件绑定 + 延长 expires_at)、`expire_due_run_snapshots`;行→dict 映射提取为共享私有helper,既有方法签名不变。
- `__init__.py`:导出新符号。
- **未删除** `EvoScientist/llm/runtime_snapshots.py`(Task 6 切换)。
### Task 3 评审接线要求的遵守
本任务不构造任何 HTTP client / ChatModel:`resolve`/`resolve_for_test` 只产出 `ResolvedModelConfig`,`SnapshotService` 只做冻结与读取。因此 `build_chat_model` 双传 sync+async client、ollama `max_retries` 经 client builder `retries=` 执行这两条接线要求在本任务无适用点,也未被绕过;Task 5/7 构造 client 时仍须遵守(`build_safe_http_client(policy, retries=resolved.client_options.max_retries)` 双传)。
## 测试(TDD)
先写 `tests/test_resolver.py`(37 例)与 `tests/test_snapshots.py`(34 例)并确认红灯(模块不存在),再实现至全绿。覆盖简报全部要求:
- resolve 正/反例:完整冻结断言(含预算数值 1048576-32768-4096)、Provider disabled→`MODEL_NOT_AVAILABLE`、模型 disabled→`MODEL_DISABLED`、无匹配契约→不可用、验证缺失/失败/五元组不一致(配置哈希、凭据轮换)→不可用、`MODEL_LIMITS_UNCONFIRMED`、`CONTEXT_BUDGET_UNSATISFIABLE`、四种角色戳记、未开放 Adapter、`AUTH_MODE_UNSUPPORTED`、mode=none 无凭据解析、resolved config 不含 secret;
- resolve_for_test:放宽模型 enabled、其余校验不放宽;
- ModelAvailability 六态判定顺序、stale 优先 configured、凭据轮换致 stale、selectable 仅 enabled、全 Provider 全模型覆盖;
- 快照:inherit 解析 defaults(含 auxiliary default null 两条路径)、显式选择冻结双角色、bootstrap 422、selection_hash 解析前语义(显式等于默认仍不同哈希)、revision 不参与哈希、幂等 200/冲突 409、expired/aborted 后重建 201、payload 冻结凭据版本且无 secret、prepared TTL 15min 断言;
- bind 一次/同 id 幂等/异 id 409/expired/aborted/未知 404;abort 规则全集;get 绑定校验(跨线程/跨部署拒绝)、expired 拒绝、spec_revision 移除后读取 `ADAPTER_NOT_SUPPORTED`;cleanup 到期迁移;bind 后 +24h 断言;
- 凭据:冻结 revision 解析、轮换后旧版本仍可用、销毁后 `RUN_CREDENTIAL_REVISION_UNAVAILABLE`、新 store 实例(无进程缓存)仍可解析、辅助角色解析到 auxiliary 凭据、mode=none 返回空串;
- 6.1 角色映射(auxiliary 冻结/缺省两路径、未知角色)与 `public_snapshot_view` 形状(无 secret、无 base_url)。
## 验证结果
- `.venv/bin/python -m pytest tests/test_resolver.py tests/test_snapshots.py -x -q` → **71 passed**
- `.venv/bin/python -m pytest tests/ -x -q` → **3131 passed, 10 skipped**(无回归;期间修复一处:Task 1 错误码表测试因新增 `SNAPSHOT_NOT_FOUND` 需增补一行期望值)
- `ruff check` 与 `ruff format --check`(model_registry 全包 + 涉及测试)→ 全净
## 疑虑 / 后续注意
1. `SNAPSHOT_NOT_FOUND` 是对设计文档 9.5 错误码总表的增补(404 类别文档已承认,但总表未列快照缺失码);Task 5 HTTP 层应直接复用。
2. 终态(expired/aborted)快照的 bind/get 统一抛 `SNAPSHOT_EXPIRED`;文档只明文规定 expired 的情形,aborted 按同一终态语义处理。
3. 冻结的 `budget.message_budget` 取 base(仅扣系统预留);逐次调用的 has_tools/has_attachments 重算属 Task 6 的 MessageBudgetMiddleware。
4. `resolve` 支持 `registry=` 参数供 `create` 传入同一份 Registry,保证快照 `registry_revision` 与解析所用文档一致。
## 评审修复(2026-07-21,commit 见下)
1. **Important:`abort` read-then-write 竞态**。原实现先读后写且 `set_run_snapshot_status` 为无条件 UPDATE,并发 bind 在两次调用间提交时会把 bound 改写为 aborted 并丢失 `langgraph_run_id`。修复:store 层新增 `abort_run_snapshot`(`UPDATE ... SET status='aborted' WHERE snapshot_id=? AND status='prepared'`,按 rowcount 判定),`SnapshotService.abort` 改为与 `bind` 相同的读-条件写-失败重读循环;已 aborted 重复调用保持幂等成功。
2. **Minor:selection_hash 期望值自证**。`test_selection_hash_uses_pre_resolution_semantics` 原先用测试内重复实现的同一序列化逻辑计算期望值(两侧同变不红)。改为钉死离线算出的 SHA-256 字面值(`{"auxiliary":null,"primary":null}` → `697c0462...55abc`,常量 `INHERIT_SELECTION_HASH`),锁住对外契约;删除测试内的重复实现。
新增测试(先红后绿):
- `test_abort_run_snapshot_store_update_is_conditional`:store 层条件 UPDATE 的 rowcount 语义(prepared→True,重复/bound→False 且行不被改写)。
- `test_abort_losing_bind_race_keeps_bound_state`:monkeypatch `get_run_snapshot` 在 abort 读与写之间插入并发 bind,断言 abort 抛 `SNAPSHOT_ALREADY_BOUND` 且行保持 bound、`langgraph_run_id` 完好(旧实现此测试必红)。
验证:
- `.venv/bin/python -m pytest tests/test_snapshots.py -x -q` → **42 passed**
- `.venv/bin/python -m pytest tests/ -x -q` → **3133 passed, 10 skipped**(无回归)
- `ruff check` / `ruff format --check`(涉及文件)→ 全净
+184 -209
View File
@@ -55,19 +55,11 @@ DEFAULT_SKILL_SOURCES = ("/skills/",)
_config = None
_chat_model = None
# Track the (model, provider) binding of _chat_model so cache invalidates
# when config.model/provider change (e.g. via /model). Without this,
# _ensure_chat_model() returns the stale cached instance even after
# _ensure_config(new_cfg) has overwritten the active config — causing
# /model switch to lag one step (see issue #179).
_chat_model_key: tuple[str | None, str | None] | None = None
# Auxiliary model for background/helper LLM calls (memory workers + main-agent
# tool selector). Cached separately from the main model; falls back to the main
# instance when the auxiliary_* config fields are empty (see
# _ensure_auxiliary_chat_model).
_auxiliary_chat_model = None
_auxiliary_chat_model_key: tuple[str | None, str | None] | None = None
# Track the (provider_id, model_key, registry_revision) binding of
# _chat_model so the cache invalidates when the registry default changes.
# The compile-time binding is only a placeholder — per-run resolution
# always comes from the run snapshot via ConfigurableModelMiddleware.
_chat_model_key: tuple[str, str, int] | None = None
# Cache MCP tools by the effective config signature to avoid reconnecting
# to MCP servers on every `/new` when config is unchanged.
@@ -88,10 +80,10 @@ _EvoScientist_agent = None
def set_active_config(cfg) -> None:
"""Commit *cfg* as the active module config.
Public commit path for callers (e.g. ``/model``) that built an agent on
the pure ``create_cli_agent(config=..., chat_model=...)`` path and now
want it to become the session-wide active config. This is the write half
of ``_ensure_config(cfg)`` extracted so the pure path can defer the commit
Public commit path for callers that built an agent on the pure
``create_cli_agent(config=..., chat_model=...)`` path and now want it
to become the session-wide active config. This is the write half of
``_ensure_config(cfg)`` extracted so the pure path can defer the commit
until the agent has been built successfully.
"""
global _config
@@ -118,25 +110,11 @@ def _ensure_config(config=None):
return _config
def _build_chat_model(cfg):
"""Build a chat model from *cfg* without writing any module globals.
Pure-construction counterpart to ``_ensure_chat_model``: used by ``/model``
to verify a switch before committing, and threaded into
``create_cli_agent(chat_model=...)`` so the new agent binds the requested
model without touching the cached ``_chat_model``.
"""
from .llm import get_chat_model
return get_chat_model(model=cfg.model, provider=cfg.provider)
def _replace_chat_model(instance, key: tuple[str | None, str | None]) -> None:
def _replace_chat_model(instance, key: tuple[str, str, int]) -> None:
"""Install a new chat model and propagate the related invariants.
Single write point for ``_chat_model`` / ``_chat_model_key`` /
``_EvoScientist_agent``: both ``_ensure_chat_model`` (cache-miss
rebuild) and ``set_chat_model`` (explicit switch via ``/model``)
``_EvoScientist_agent``: ``_ensure_chat_model`` cache-miss rebuilds
funnel through here so the three globals can never drift.
"""
global _chat_model, _chat_model_key, _EvoScientist_agent
@@ -148,81 +126,50 @@ def _replace_chat_model(instance, key: tuple[str | None, str | None]) -> None:
def _ensure_chat_model():
"""Return cached chat model, rebuilding if cfg.model/provider changed.
"""Return the compile-time main chat model, built from registry defaults.
The cache key is the current config's ``(model, provider)``. If it
differs from the key that built ``_chat_model``, rebuild — this makes
``create_cli_agent(config=temp_cfg)`` bind the freshly requested model
into the new agent without requiring callers to interleave
``set_chat_model()`` calls in any particular order.
The model is resolved from the active registry's ``defaults.primary``
and cached under its ``(provider_id, model_key, registry_revision)``
key. It is only the compile-time placeholder: every run re-resolves its
model from the run snapshot via ``ConfigurableModelMiddleware``.
Raises:
ModelRegistryError: ``MODEL_REGISTRY_NOT_READY`` when the registry
is still in bootstrap (no enabled primary model configured).
"""
cfg = _ensure_config()
key = (cfg.model, cfg.provider)
from .model_registry.runtime import get_snapshot_runtime
runtime = get_snapshot_runtime()
primary_ref, revision = runtime.registry_default()
key = (primary_ref.provider_id, primary_ref.model_key, revision)
if _chat_model is None or _chat_model_key != key:
_replace_chat_model(_build_chat_model(cfg), key)
_replace_chat_model(runtime.build_default_role_model("primary"), key)
return _chat_model
def _ensure_auxiliary_chat_model():
"""Return the auxiliary chat model for background/helper LLM calls.
def _compile_time_role_model(role: str = "primary"):
"""Return the compile-time model binding for graph construction.
Resolves ``(cfg.auxiliary_model or cfg.model, cfg.auxiliary_provider or
cfg.provider)``. When the auxiliary fields are empty — or resolve to the same
``(model, provider)`` pair as the main model — returns the main
``_ensure_chat_model()`` instance directly, so no second client is built.
Otherwise it is cached separately under its own key. Onboard sets the
provider alongside the model, so the ``or cfg.provider`` fallback only
matters for a model set without an explicit auxiliary provider.
Unlike ``_ensure_chat_model()`` this never raises on a bootstrap
registry: graphs must still materialize so the Config API can serve
(design doc section 10 — only run creation is forbidden in bootstrap),
so a ``RegistryNotReadyChatModel`` placeholder is bound instead. It
raises ``MODEL_REGISTRY_NOT_READY`` on the first model call; per-run
resolution still comes from the run snapshot via
``ConfigurableModelMiddleware``. Nothing is cached in module globals —
once the registry becomes active, rebuilt graphs get the real model.
"""
global _auxiliary_chat_model, _auxiliary_chat_model_key
from .llm import get_chat_model
from .model_registry.errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError
from .model_registry.placeholder import RegistryNotReadyChatModel
from .model_registry.runtime import get_snapshot_runtime
cfg = _ensure_config()
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):
return _ensure_chat_model()
key = (aux_model, aux_provider)
if _auxiliary_chat_model is None or _auxiliary_chat_model_key != key:
_auxiliary_chat_model = get_chat_model(model=aux_model, provider=aux_provider)
_auxiliary_chat_model_key = key
return _auxiliary_chat_model
def set_chat_model(model: str, provider: str | None = None):
"""Replace the cached chat model with a new one.
Called by ``/model`` to switch the LLM mid-session. No-op when the
cache already holds the requested ``(model, provider)`` — avoids
spawning a second ``get_chat_model`` instance (and its HTTP client)
under the ``/model`` flow where ``_ensure_chat_model`` has already
rebuilt ``_chat_model`` during the preceding ``_load_agent`` call.
Returns the current chat model instance.
"""
from .llm import get_chat_model
# Invalidate the auxiliary cache too: when auxiliary_* is empty it mirrors
# the main model, so a /model switch must let it re-resolve to the new main.
global _auxiliary_chat_model, _auxiliary_chat_model_key
_auxiliary_chat_model = None
_auxiliary_chat_model_key = None
key = (model, provider)
if _chat_model is None or _chat_model_key != key:
_replace_chat_model(get_chat_model(model=model, provider=provider), key)
return _chat_model
def set_chat_model_instance(instance, key: tuple[str | None, str | None]) -> None:
"""Commit an already-built chat model *instance* as the active model.
Companion to ``set_active_config`` for the pure path: installs a model that
``_build_chat_model`` already constructed (e.g. during a ``/model`` verify)
without rebuilding it, keeping ``_chat_model`` / ``_chat_model_key`` /
``_EvoScientist_agent`` in sync via ``_replace_chat_model``. Unlike
``set_chat_model``, the caller owns the ``(model, provider)`` *key*.
"""
_replace_chat_model(instance, key)
runtime = get_snapshot_runtime()
try:
return runtime.build_default_role_model(role)
except ModelRegistryError as exc:
if exc.code != MODEL_REGISTRY_NOT_READY:
raise
return RegistryNotReadyChatModel(detail=str(exc))
# =============================================================================
@@ -341,9 +288,12 @@ def _inject_subagent_middleware(
ToolErrorHandlerMiddleware(),
ContextOverflowMapperMiddleware(),
]
if memory_controls.memory_enabled:
if memory_controls.memory_enabled and cfg.workspace_isolation != "required":
middleware.append(memory_middleware)
if memory_controls.worker_needed(MemoryObservationTarget.SUBAGENT_WORKER):
if (
memory_controls.worker_needed(MemoryObservationTarget.SUBAGENT_WORKER)
and cfg.workspace_isolation != "required"
):
middleware.append(
create_memory_lifecycle_middleware(
memory_dir,
@@ -463,9 +413,10 @@ def _maybe_swap_async_subagents(
middleware.append(AsyncWatcherMiddleware(agent_specs))
# Forward the CLI's live (model, provider) into deepagents'
# start/update_async_task tool calls so the deployed graph can
# re-resolve its chat model per run via ConfigurableModelMiddleware.
# Wrap deepagents' start/update_async_task tool calls so workspace scope
# and usage correlation metadata reach the deployed graph's runs. Model
# configuration is NOT forwarded: the deployed graph resolves its model
# per run from ``runtime_snapshot_id`` via ConfigurableModelMiddleware.
# Idempotent — safe to call on every CLI startup.
if agent_specs:
from .llm.patches import _patch_deepagents_model_passthrough
@@ -479,14 +430,25 @@ def _build_base_kwargs(
base_backend, base_middleware, *, cfg=None, chat_model=None, workspace_dir=None
):
"""Build agent kwargs *without* MCP (fast, no subprocess spawning)."""
from .tools import skill_manager, tavily_search, think_tool
from .tools import (
edit_image,
generate_image,
refresh_image_tool_descriptions,
skill_manager,
tavily_search,
think_tool,
)
from .utils import load_subagents
cfg = cfg if cfg is not None else _ensure_config()
model = chat_model if chat_model is not None else _compile_time_role_model("primary")
tool_registry = {"think_tool": think_tool}
if os.environ.get("TAVILY_API_KEY"):
tool_registry["tavily_search"] = tavily_search
base_tools = [think_tool, skill_manager]
refresh_image_tool_descriptions()
base_tools = [think_tool, generate_image, edit_image]
if cfg.workspace_isolation != "required":
base_tools.append(skill_manager)
subs = load_subagents(
SUBAGENTS_CONFIG,
@@ -494,12 +456,12 @@ def _build_base_kwargs(
)
_ensure_general_purpose_subagent(subs)
_inject_subagent_middleware(
subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=model
)
subs = _maybe_swap_async_subagents(subs, base_middleware, cfg=cfg)
return {
"name": "EvoScientist",
"model": chat_model if chat_model is not None else _ensure_chat_model(),
"model": model,
"tools": list(base_tools),
"backend": base_backend,
"subagents": subs,
@@ -531,24 +493,35 @@ def load_mcp_and_build_kwargs(
chat_model: Explicit chat model to bind instead of
``_ensure_chat_model()`` (which would write module globals).
"""
from .tools import skill_manager, tavily_search, think_tool
from .tools import (
edit_image,
generate_image,
refresh_image_tool_descriptions,
skill_manager,
tavily_search,
think_tool,
)
from .utils import load_subagents
cfg = cfg if cfg is not None else _ensure_config()
model = chat_model if chat_model is not None else _compile_time_role_model("primary")
mcp_by_agent = _load_mcp_tools_cached(on_progress=on_mcp_progress)
if not mcp_by_agent:
return _build_base_kwargs(
base_backend,
base_middleware,
cfg=cfg,
chat_model=chat_model,
chat_model=model,
workspace_dir=workspace_dir,
)
tool_registry = {"think_tool": think_tool}
if os.environ.get("TAVILY_API_KEY"):
tool_registry["tavily_search"] = tavily_search
base_tools = [think_tool, skill_manager]
refresh_image_tool_descriptions()
base_tools = [think_tool, generate_image, edit_image]
if cfg.workspace_isolation != "required":
base_tools.append(skill_manager)
# Fresh tool registry — start from base tools + MCP tools
registry = dict(tool_registry)
@@ -565,7 +538,7 @@ def load_mcp_and_build_kwargs(
_ensure_general_purpose_subagent(subs)
_inject_subagent_middleware(
subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=model
)
# Inject MCP tools into subagents by name
@@ -579,7 +552,7 @@ def load_mcp_and_build_kwargs(
return {
"name": "EvoScientist",
"model": chat_model if chat_model is not None else _ensure_chat_model(),
"model": model,
"tools": base_tools + mcp_main,
"backend": base_backend,
"subagents": subs,
@@ -594,8 +567,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 by CLI and legacy mode only."""
from deepagents.backends import CompositeBackend
from .backends import (
@@ -637,17 +610,45 @@ def _get_default_backend():
)
def _get_default_backend():
"""Return a runtime-scoped backend factory when isolation is enabled."""
from .workspace_cutover import verify_required_cutover
from .workspace_scope import (
create_workspace_backend_factory,
verify_required_executor,
)
cfg = _ensure_config()
if cfg.workspace_isolation == "legacy":
return _get_legacy_backend()
if cfg.workspace_isolation == "required" and cfg.dangerous_mode:
raise RuntimeError(
"dangerous_mode is incompatible with required workspace isolation"
)
if cfg.workspace_isolation == "required":
verify_required_cutover(_paths_mod.WORKSPACE_ROOT)
verify_required_executor()
return create_workspace_backend_factory(
_get_legacy_backend,
dangerous=cfg.dangerous_mode,
# The CLI and its stripped async-subagent service keep their configured
# shared workspace. The WebUI deployment must receive a scope from the
# trusted WebUI/API boundary instead.
allow_unscoped_legacy=os.environ.get("EVOSCIENTIST_DEPLOY_MODE", "").lower()
!= "full",
)
def _get_default_middleware(
*,
for_async_subagent: bool = False,
workspace_dir: str | Path | None = None,
memory_dir: str | Path | None = None,
cfg=None,
chat_model=None,
backend=None,
memory_source_agent: str = "EvoScientist",
tool_selector_threshold: int | None = None,
memory_max_inline_profile_chars: int | None = None,
enable_background_execution: bool = True,
snapshot_role: str = "primary",
):
"""Build the default middleware list.
@@ -667,28 +668,35 @@ def _get_default_middleware(
(avoids writing module globals on the pure path).
memory_source_agent: Attribution name for profile/observation writes.
Async sub-agent factories pass their deployed agent name here.
snapshot_role: The model role ``ConfigurableModelMiddleware``
resolves from the run snapshot; every role maps to the
snapshot's frozen primary (design doc 6.1).
"""
from .middleware import (
ConfigurableModelMiddleware,
ContextOverflowMapperMiddleware,
ModelFallbackMiddleware,
ToolErrorHandlerMiddleware,
create_code_interpreter_middleware,
create_context_editing_middleware,
create_memory_lifecycle_middleware,
create_memory_middleware,
create_message_budget_middleware,
create_runtime_context_middleware,
create_scheduler_middleware,
create_tool_selector_middleware,
default_memory_scheduler,
load_fallback_chain,
)
cfg = cfg if cfg is not None else _ensure_config()
if cfg.model_fallbacks:
load_fallback_chain(cfg.model_fallbacks)
model = chat_model if chat_model is not None else _ensure_chat_model()
memory_dir = str(memory_dir or _paths_mod.MEMORIES_DIR)
model = chat_model if chat_model is not None else _compile_time_role_model("primary")
if backend is None:
# Preserve the factory's pure path for callers that provide an
# explicit model/configuration (notably tests and subagent assembly).
# Production graph factories always pass their real composite backend.
from deepagents.backends import StateBackend
backend = StateBackend()
memory_dir = str(_paths_mod.MEMORIES_DIR)
source_type = (
MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN
)
@@ -699,69 +707,59 @@ def _get_default_middleware(
if for_async_subagent
else MemoryObservationTarget.TURN_WORKER
)
# ``ConfigurableModelMiddleware`` is placed first so it wraps
# ``ModelFallbackMiddleware``: a configurable.model override sets the
# PRIMARY model only, leaving the fallback chain free to try its own
# alternatives instead of re-overriding every retry to the same model.
memory_kwargs = {
"workspace_dir": workspace_dir,
"source_type": source_type,
"source_agent": memory_source_agent,
"enable_profile_memory": memory_controls.profile_enabled,
"enable_observation_memory": memory_controls.observations_enabled,
"enable_observation_tool": memory_controls.observation_tool_enabled(
# ``ConfigurableModelMiddleware`` sits first so the snapshot-driven model
# override applies before any other middleware inspects the request.
memory_middleware = create_memory_middleware(
memory_dir,
workspace_dir=workspace_dir,
source_type=source_type,
source_agent=memory_source_agent,
enable_profile_memory=memory_controls.profile_enabled,
enable_observation_memory=memory_controls.observations_enabled,
enable_observation_tool=memory_controls.observation_tool_enabled(
MemoryObservationTarget.AGENT
),
"memory_scheduler": memory_scheduler,
}
if memory_max_inline_profile_chars is not None:
memory_kwargs["max_inline_profile_chars"] = memory_max_inline_profile_chars
memory_middleware = create_memory_middleware(memory_dir, **memory_kwargs)
# Main-agent tool selection may use the auxiliary model; async sub-agents
# keep the main model (they do real work, not a one-off helper call).
memory_scheduler=memory_scheduler,
)
# Main-agent tool selection resolves its helper model from the run
# snapshot on every call (model=None below); async sub-agents and the
# pure path (explicit model + config) keep their threaded model.
# 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
elif chat_model is None:
tool_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
else:
from .llm import get_chat_model
tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider)
tool_selector_model = None if (not for_async_subagent and chat_model is None) else model
mw = [
ConfigurableModelMiddleware(),
ConfigurableModelMiddleware(role=snapshot_role),
create_message_budget_middleware(model, backend, snapshot_role=snapshot_role),
create_context_editing_middleware(model),
ModelFallbackMiddleware(),
ContextOverflowMapperMiddleware(),
ToolErrorHandlerMiddleware(),
*create_tool_selector_middleware(
**(
{"threshold": tool_selector_threshold}
if tool_selector_threshold is not None
else {}
),
model=tool_selector_model,
track_stream_selection=not for_async_subagent,
),
# Interpreter prompt must land before runtime/memory context, so this
# middleware sits ahead of runtime_context in the stack.
create_code_interpreter_middleware(
timeout=cfg.code_interpreter_timeout,
max_result_chars=cfg.code_interpreter_max_result_chars,
*(
[]
if cfg.workspace_isolation == "required"
and cfg.strict_code_interpreter == "disabled"
else [
create_code_interpreter_middleware(
timeout=cfg.code_interpreter_timeout,
max_result_chars=cfg.code_interpreter_max_result_chars,
)
]
),
]
if cfg.enable_scheduler and not for_async_subagent:
mw.append(create_scheduler_middleware())
mw.append(create_runtime_context_middleware())
if memory_controls.memory_enabled:
if memory_controls.memory_enabled and cfg.workspace_isolation != "required":
mw.append(memory_middleware)
if memory_controls.worker_needed(worker_target):
if (
memory_controls.worker_needed(worker_target)
and cfg.workspace_isolation != "required"
):
mw.append(
create_memory_lifecycle_middleware(
memory_dir,
@@ -781,7 +779,7 @@ def _get_default_middleware(
# Background-process tools (run_in_background / check_process / stop_process /
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
# must not spawn local OS processes.
if not for_async_subagent and enable_background_execution:
if not for_async_subagent and cfg.workspace_isolation != "required":
from .middleware.background import BackgroundExecutionMiddleware
mw.append(BackgroundExecutionMiddleware())
@@ -818,7 +816,7 @@ def _get_default_agent():
cfg = _ensure_config()
be = _get_default_backend()
mw = _get_default_middleware()
mw = _get_default_middleware(backend=be)
# HITL on main agent only (mirrors create_cli_agent). Use middleware,
# not interrupt_on= kwarg — the kwarg propagates to every subagent and
@@ -835,7 +833,10 @@ def _get_default_agent():
)
)
if os.environ.get("EVOSCIENTIST_DEPLOY_MODE", "").lower() == "stripped":
if (
os.environ.get("EVOSCIENTIST_DEPLOY_MODE", "").lower() == "stripped"
or cfg.workspace_isolation == "required"
):
kwargs = _build_base_kwargs(
be,
mw,
@@ -879,12 +880,6 @@ def create_cli_agent(
chat_model=None,
*,
on_mcp_progress=None,
workspace_backend=None,
memory_dir: str | Path | None = None,
tool_selector_threshold: int | None = None,
memory_max_inline_profile_chars: int | None = None,
enable_subagents: bool = True,
enable_background_execution: bool = True,
) -> "CompiledStateGraph":
"""Create agent with checkpointer for CLI multi-turn support.
@@ -895,10 +890,10 @@ def create_cli_agent(
**Pure path:** when *both* ``config`` and ``chat_model`` are explicit, this
writes none of the cached config/model module globals (``_config``,
``_chat_model``, ``_chat_model_key``, ``_EvoScientist_agent``) — the agent
is built purely from the passed-in locals. The caller commits the switch
on success via ``set_active_config`` / ``set_chat_model_instance`` (see
``/model``). Otherwise the existing module-global path runs (langgraph
dev, notebooks, and CLI startup, which pass ``config=`` only).
is built purely from the passed-in locals. The caller commits the config
on success via ``set_active_config``. Otherwise the existing
module-global path runs (langgraph dev, notebooks, and CLI startup, which
pass ``config=`` only).
Args:
workspace_dir: Per-session workspace directory. If ``None``,
@@ -911,16 +906,6 @@ def create_cli_agent(
chat_model: Optional pre-built chat model. Only triggers the pure
path when ``config`` is also explicit; otherwise it is ignored in
favor of the ``_ensure_chat_model()`` fallback.
workspace_backend: Optional host-provided backend for the workspace
route. The default remains ``CustomSandboxBackend``.
memory_dir: Optional memory root used by both the backend route and
memory middleware.
tool_selector_threshold: Optional adaptive tool-selection threshold.
memory_max_inline_profile_chars: Optional memory profile injection cap.
enable_subagents: Whether configured subagents are available to the agent.
enable_background_execution: Whether local background-process tools are
installed. Embedding hosts should disable this when process execution
is provided by an external backend.
"""
import os as _os
@@ -962,21 +947,19 @@ def create_cli_agent(
workspace_dir = str(_paths.WORKSPACE_ROOT)
# Read paths dynamically so runtime set_workspace_root() changes are picked up
_mem_dir = str(memory_dir or _paths.MEMORIES_DIR)
_mem_dir = str(_paths.MEMORIES_DIR)
_usr_skills_dir = str(_paths.USER_SKILLS_DIR)
_global_skills_dir = str(_paths.GLOBAL_SKILLS_DIR)
# Always construct fresh backends from current paths (avoids stale
# module-level backend when workspace root changed at runtime).
set_active_workspace(workspace_dir)
ws_backend = workspace_backend
if ws_backend is None:
ws_backend = CustomSandboxBackend(
root_dir=workspace_dir,
virtual_mode=True,
timeout=cfg.sandbox_execute_timeout,
dangerous=cfg.dangerous_mode,
)
ws_backend = CustomSandboxBackend(
root_dir=workspace_dir,
virtual_mode=True,
timeout=cfg.sandbox_execute_timeout,
dangerous=cfg.dangerous_mode,
)
sk_backend = MergedSkillsBackend(
primary_dir=_usr_skills_dir,
global_dir=_global_skills_dir,
@@ -998,13 +981,7 @@ def create_cli_agent(
# CLI agent never drifts from the default chain. Anything CLI-specific
# (e.g. ``HumanInTheLoopMiddleware``) is appended below.
mw: list[AgentMiddleware] = _get_default_middleware(
workspace_dir=workspace_dir,
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,
workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model, backend=be
)
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
@@ -1030,8 +1007,6 @@ def create_cli_agent(
chat_model=chat_model,
workspace_dir=workspace_dir,
)
if not enable_subagents:
kwargs = {**kwargs, "subagents": []}
return create_deep_agent(
**kwargs,
-7
View File
@@ -9,8 +9,6 @@ from __future__ import annotations
from importlib import import_module
__version__ = "0.2.2"
_EXPORTS: dict[str, tuple[str, str]] = {
# Agent graph (lazy to avoid expensive initialization at import time)
"EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"),
@@ -25,11 +23,6 @@ _EXPORTS: dict[str, tuple[str, str]] = {
"save_config": (".config", "save_config"),
"get_effective_config": (".config", "get_effective_config"),
"get_config_path": (".config", "get_config_path"),
# LLM
"get_chat_model": (".llm", "get_chat_model"),
"MODELS": (".llm", "MODELS"),
"list_models": (".llm", "list_models"),
"DEFAULT_MODEL": (".llm", "DEFAULT_MODEL"),
# Prompts
"get_system_prompt": (".prompts", "get_system_prompt"),
# Tools
+47 -9
View File
@@ -682,6 +682,27 @@ def _guard_bare_absolute(result: str | None) -> str | None:
return result
def _resolve_workspace_quoted_path(path: str, workspace_root: str) -> str | None:
"""Map a quoted virtual absolute path onto the sandbox workspace.
Returns the ``./...`` form of *path* when its parent directory exists
under *workspace_root*, else ``None``. Parent-existence is what
disambiguates workspace paths (``/微信图片.jpg``, ``/artifacts/new.png``
— including files not written yet) from real system paths
(``/etc/hosts``, ``/bin/echo``), which stay untouched.
"""
rel = posixpath.normpath(path.lstrip("/"))
if rel in ("", "."):
return "."
if rel == ".." or rel.startswith("../"):
return None
parent_rel = posixpath.dirname(rel)
parent = Path(workspace_root) / parent_rel if parent_rel else Path(workspace_root)
if not parent.is_dir():
return None
return "./" + rel
def _rewrite_quoted_path(
path: str,
workspace_name: str | None,
@@ -720,6 +741,7 @@ def _rewrite_quoted_path(
def convert_virtual_paths_in_command(
command: str,
workspace_name: str | None = None,
workspace_root: str | None = None,
) -> str:
"""Convert virtual paths (starting with ``/``) in commands to relative paths.
@@ -730,21 +752,36 @@ def convert_virtual_paths_in_command(
mount (``/skills/...``, ``/memories/...``) or a workspace-prefixed
system path are rewritten as a single shell token — this fixes #237
where ``python "/skills/my skill/main.py"`` was truncated at the
embedded space. Bare quoted ``/...`` paths (e.g. ``echo "/hi"``)
are left untouched since their semantics are ambiguous.
embedded space. When *workspace_root* is given, quoted ``/...`` paths
whose parent directory exists under the workspace (e.g. an uploaded
file the WebUI references as ``/<name>``) are rewritten to ``./...``
as well; quoted paths whose parent is absent from the workspace
(``/etc/hosts``, ``/bin/echo``) are left untouched.
After pre-processing, the original regex handles unquoted
paths and workspace-name correction as before.
"""
# Pre-process: rewrite quoted paths whose decoded content starts with /
def _rewrite_quoted(match: re.Match[str]) -> str:
quote = match.group(1)
decoded = re.sub(r"\\(.)", r"\1", match.group(2))
# Quoted workspace paths (e.g. an uploaded file the WebUI references
# as ``/<name>``) would hit the host root inside scripts. Map them
# onto the workspace, keeping the original quote char — inside
# heredoc code the quotes are syntax, and shlex.quote's bare form
# would corrupt the code.
if workspace_root is not None and decoded.startswith("/"):
mapped = _resolve_workspace_quoted_path(decoded, workspace_root)
if mapped == ".":
return "."
if mapped is not None and quote not in mapped:
return quote + mapped + quote
return (
_rewrite_quoted_path(decoded, workspace_name) or match.group(0)
)
command = re.sub(
r'(["\'])((?:\\.|(?!\1).)*?)\1',
lambda m: (
_rewrite_quoted_path(
re.sub(r"\\(.)", r"\1", m.group(2)),
workspace_name,
)
or m.group(0)
),
_rewrite_quoted,
command,
)
@@ -1103,6 +1140,7 @@ def prepare_sandbox_command(
command = convert_virtual_paths_in_command(
command=command,
workspace_name=Path(cwd_str).name,
workspace_root=cwd_str,
)
# Skills/memory dirs must be allowlisted: the workspace-literal replace above runs
# before the resolver, so any absolute path it later injects reaches validate unstripped.
+4 -5
View File
@@ -150,11 +150,10 @@ class BackgroundAgentLoader(Generic[AgentT]):
def adopt(self, agent: AgentT) -> None:
"""Install an externally-built agent and supersede any in-flight load.
Used by ``/model`` (and any other caller that constructs a
replacement agent directly): bumps the generation token so a
late-arriving background load can't clobber ``self.agent`` via
the done-callback, cancels the in-flight wrapper, and seats the
new agent immediately.
Used by any caller that constructs a replacement agent directly:
bumps the generation token so a late-arriving background load can't
clobber ``self.agent`` via the done-callback, cancels the in-flight
wrapper, and seats the new agent immediately.
"""
prev = self._task
if prev is not None and not prev.done():
+17
View File
@@ -62,6 +62,23 @@ def _create_session_workspace(name: str | None = None) -> str:
return workspace_dir
def current_model_label() -> str:
"""Return the display label for the registry's primary default model.
The CLI no longer carries ``config.model``/``config.provider`` free
strings; the status bar and run metadata label the active model as
``provider_id/model_key`` from the Model Registry defaults. Returns
``"unconfigured"`` when the registry is still in bootstrap.
"""
try:
from ..model_registry.runtime import get_snapshot_runtime
primary, _ = get_snapshot_runtime().registry_default()
except Exception:
return "unconfigured"
return f"{primary.provider_id}/{primary.model_key}"
def _load_agent(
workspace_dir: str | None = None,
checkpointer=None,
+1 -1
View File
@@ -306,7 +306,7 @@ async def dispatch_channel_slash_command(
resolution, or the dispatcher's input agent when no resolver is
supplied. Callers can compare ``ctx.agent`` with
``original_agent`` to detect command-driven swaps. Used by Rich
CLI to (a) adopt an agent swap (``/model``) back into the
CLI to (a) adopt an agent swap back into the
running session and (b) refresh the status snapshot for
commands that mutate session-level state (``/new``,
``/compact``) — mirrors the REPL dispatch at
+62 -94
View File
@@ -74,16 +74,9 @@ if TYPE_CHECKING:
@app.command()
def onboard(
skip_validation: bool = typer.Option(
False, "--skip-validation", help="Skip API key validation during setup"
False, "--skip-validation", help="Skip Tavily key validation during setup"
),
# ---- Pre-fill answers (any subset; remaining prompts stay interactive)
provider: str | None = typer.Option(
None, "--provider", help="Pre-set LLM provider (e.g. anthropic, openai)"
),
model: str | None = typer.Option(None, "--model", help="Pre-set model name"),
api_key: str | None = typer.Option(
None, "--api-key", help="Pre-set API key for the chosen --provider"
),
tavily_key: str | None = typer.Option(
None, "--tavily-key", help="Pre-set Tavily API key"
),
@@ -120,16 +113,16 @@ def onboard(
):
"""Interactive setup wizard for EvoScientist.
Guides you through configuring API keys, model selection,
workspace settings, and agent parameters.
Guides you through workspace settings, channels, and agent parameters.
Models and provider credentials are configured through the Model
Registry (WebUI configuration page or the model-registry API), not
here.
Any answer can be pre-set via a flag (``--provider anthropic
--model claude-sonnet-4-6 ...``); prompts for unset answers stay
Any answer can be pre-set via a flag; prompts for unset answers stay
interactive unless ``--non-interactive`` is passed, in which case any
missing required answer aborts the wizard.
"""
from ..config.onboard.constants import (
VALID_PROVIDERS,
VALID_UI_BACKENDS,
VALID_WORKSPACE_MODES,
)
@@ -149,11 +142,6 @@ def onboard(
f"--workspace-mode must be one of {sorted(VALID_WORKSPACE_MODES)}",
param_hint="--workspace-mode",
)
if provider is not None and provider not in VALID_PROVIDERS:
raise typer.BadParameter(
f"--provider must be one of {sorted(VALID_PROVIDERS)}",
param_hint="--provider",
)
# Match the interactive prompt's range (1024 < port < 65536). Without
# this check, --port 80 or --port 99999 would land in config and break
# the langgraph dev server on startup.
@@ -169,12 +157,6 @@ def onboard(
answers["ui"] = ui
if port is not None:
answers["port"] = str(port)
if provider is not None:
answers["provider"] = provider
if model is not None:
answers["model"] = model
if api_key is not None:
answers["api_key"] = api_key
if tavily_key is not None:
answers["tavily_key"] = tavily_key
if workspace_mode is not None:
@@ -212,8 +194,6 @@ def onboard(
_CONFIGURE_SECTIONS = {
"ui": "UI backend",
"port": "LangGraph server port",
"provider": "LLM provider + auth + API key",
"model": "Model + reasoning effort",
"tavily": "Tavily search key",
"workspace": "Workspace mode",
"thinking": "Thinking panel",
@@ -263,29 +243,6 @@ def configure_port():
_configure_section("port")
@configure_app.command("provider")
def configure_provider(
skip_validation: bool = typer.Option(False, "--skip-validation"),
):
"""Re-run LLM provider, auth mode, and API key prompts.
Model selection is automatically re-run after provider — the model list
depends on the provider, and silently leaving e.g. ``model="claude-...""``
when the provider was switched to ``openai`` would break the first
request. Press Enter on the model picker to keep the current default.
"""
_run_onboard_cli(
skip_validation=skip_validation,
only_sections={"provider", "model"},
)
@configure_app.command("model")
def configure_model():
"""Re-run model selection (and reasoning effort for OpenRouter)."""
_configure_section("model")
@configure_app.command("tavily")
def configure_tavily(
skip_validation: bool = typer.Option(False, "--skip-validation"),
@@ -1022,7 +979,7 @@ def _make_serve_cmd_completed_hook(
):
"""Build the ``on_cmd_completed`` hook used by serve mode.
Adopts ``/model`` agent swaps and ``/resume`` thread/workspace
Adopts command-driven agent swaps and ``/resume`` thread/workspace
swaps back into ``runtime_state`` so the outer poll loop picks up
the new handles on subsequent messages. Also keeps
``channel_runtime`` in sync so the bus sees the new values.
@@ -1389,6 +1346,47 @@ def _serve_drain_notifications(
_notif_loop.close()
def _startup_gates() -> None:
"""Refuse agent startup on legacy artifacts or a bootstrap registry.
Design doc section 10 step 4 and section 8.1:
1. Legacy model-configuration artifacts (old ``providers.yaml``, old
``run-runtime-snapshots.sqlite3``, leftover LLM fields in
``config.yaml``) abort startup with an explicit reset guide —
nothing is read partially.
2. A bootstrap registry aborts startup with
``MODEL_REGISTRY_NOT_READY`` — there is no implicit default model
anymore; configure and enable a primary model through the Model
Registry first.
"""
from ..config.legacy_artifacts import (
LegacyArtifactsError,
assert_no_legacy_artifacts,
)
from ..model_registry.errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError
try:
assert_no_legacy_artifacts()
except LegacyArtifactsError as exc:
console.print(f"[red]{exc}[/red]")
raise typer.Exit(1) from exc
try:
from ..model_registry.runtime import get_snapshot_runtime
get_snapshot_runtime().registry_default()
except ModelRegistryError as exc:
if exc.code != MODEL_REGISTRY_NOT_READY:
raise
console.print(
f"[red]{exc.code}: {exc}[/red]\n"
"Configure a provider/model through the Model Registry (WebUI "
"configuration page or the model-registry API), run the provider "
"test, enable the model, and set it as the primary default."
)
raise typer.Exit(1) from exc
@app.command()
def serve(
no_thinking: bool = typer.Option(
@@ -1452,20 +1450,7 @@ def serve(
if debug:
_configure_logging()
# Auto-start ccproxy if any provider uses OAuth mode
_ccproxy_proc_serve = None
if config.anthropic_auth_mode == "oauth" or config.openai_auth_mode == "oauth":
try:
from ..ccproxy_manager import maybe_start_ccproxy, stop_ccproxy
_ccproxy_proc_serve = maybe_start_ccproxy(config)
if _ccproxy_proc_serve:
import atexit
atexit.register(stop_ccproxy, _ccproxy_proc_serve)
except RuntimeError as exc:
console.print(f"[red]{exc}[/red]")
raise typer.Exit(1) from exc
_startup_gates()
if not config.channel_enabled:
console.print("[red]No channels configured.[/red]")
@@ -1565,6 +1550,9 @@ def serve(
_orig_sigterm = signal.signal(signal.SIGTERM, _handle_shutdown)
try:
from .agent import current_model_label
model_label = current_model_label()
while not shutdown_event.is_set():
try:
msg = _message_queue.get(timeout=0.5)
@@ -1577,7 +1565,7 @@ def serve(
_serve_process_message(
msg,
runtime_state=runtime_state,
model=config.model,
model=model_label,
workspace_dir=ws,
show_thinking=effective_channel_thinking,
on_cmd_completed=_serve_on_cmd_completed,
@@ -1593,7 +1581,7 @@ def serve(
if async_notifier.has_pending_notifications(runtime_state.thread_id):
_serve_drain_notifications(
runtime_state=runtime_state,
model=config.model,
model=model_label,
workspace_dir=ws,
show_thinking=effective_channel_thinking,
)
@@ -2084,11 +2072,6 @@ def _main_callback(
"--dangerous",
help="DANGEROUS: real-filesystem access (no workspace confinement); implies --auto-approve",
),
auth_mode: str | None = typer.Option(
None,
"--auth-mode",
help="Auth mode for Anthropic/OpenAI: api_key (default) or oauth (ccproxy).",
),
ui: str | None = typer.Option(
None,
"--ui",
@@ -2168,29 +2151,11 @@ def _main_callback(
cli_overrides["enable_ask_user"] = True
if dangerous:
cli_overrides["dangerous_mode"] = True
if auth_mode:
if auth_mode not in ("api_key", "oauth"):
raise typer.BadParameter("--auth-mode must be 'api_key' or 'oauth'")
cli_overrides["anthropic_auth_mode"] = auth_mode
cli_overrides["openai_auth_mode"] = auth_mode
config = get_effective_config(cli_overrides)
apply_config_to_env(config)
# Auto-start ccproxy if any provider uses OAuth mode
_ccproxy_proc = None
if config.anthropic_auth_mode == "oauth" or config.openai_auth_mode == "oauth":
try:
from ..ccproxy_manager import maybe_start_ccproxy, stop_ccproxy
_ccproxy_proc = maybe_start_ccproxy(config)
if _ccproxy_proc:
import atexit
atexit.register(stop_ccproxy, _ccproxy_proc)
except RuntimeError as exc:
console.print(f"[red]{exc}[/red]")
raise typer.Exit(1) from exc
_startup_gates()
show_thinking = config.show_thinking if not no_thinking else False
effective_channel_thinking = config.channel_send_thinking and (not no_thinking)
@@ -2356,6 +2321,9 @@ def _main_callback(
config=config,
)
try:
from .agent import current_model_label
model_label = current_model_label()
if effective_output_format == "stream-json":
# Headless JSONL path: drive the sink through the gateway
# directly. We are already inside the async single-shot
@@ -2364,7 +2332,7 @@ def _main_callback(
request = RunRequest(
message=prompt,
thread_id=tid,
metadata=build_metadata(workspace_dir, config.model),
metadata=build_metadata(workspace_dir, model_label),
target=GraphTarget(
local_graph=agent, workspace_dir=workspace_dir
),
@@ -2388,7 +2356,7 @@ def _main_callback(
thread_id=tid,
show_thinking=show_thinking,
workspace_dir=workspace_dir,
model=config.model,
model=model_label,
ui_backend=config.ui_backend,
runtime_gateways=runtime_gateways,
)
@@ -2403,6 +2371,7 @@ def _main_callback(
nest_asyncio.apply()
asyncio.get_event_loop().run_until_complete(_single_shot())
else:
from .agent import current_model_label
from .interactive import cmd_interactive
# Interactive mode (default) — checkpointer managed inside cmd_interactive
@@ -2412,8 +2381,7 @@ def _main_callback(
workspace_dir=workspace_dir,
workspace_fixed=workspace_fixed,
mode=effective_mode,
model=config.model,
provider=config.provider,
model=current_model_label(),
run_name=name,
thread_id=thread_id,
ui_backend=config.ui_backend,
+8 -10
View File
@@ -961,7 +961,7 @@ def cmd_interactive(
ctx: CommandContext, original_agent: Any, cmd: Command
) -> None:
"""Mirror the REPL adoption block at
``interactive.py:1005-1030`` so ``/model`` and similar
``interactive.py:1005-1030`` so agent-swap and similar
state-mutating commands invoked via a channel actually
rebind the running session and keep the status bar
in sync."""
@@ -970,11 +970,10 @@ def cmd_interactive(
ctx.agent is not None and ctx.agent is not original_agent
)
if agent_swapped:
from ..EvoScientist import _ensure_config
from .agent import current_model_label
agent_loader.adopt(ctx.agent)
cfg = _ensure_config()
model = cfg.model
model = current_model_label()
state["status_base_snapshot"] = make_empty_status_snapshot(
model
)
@@ -1326,7 +1325,7 @@ def cmd_interactive(
if not state["running"]:
break
# Agent swap (e.g. /model successfully built a
# Agent swap (a command successfully built a
# new agent): adopt into loader + reset status
# snapshot + sync channel runtime.
agent_swapped = (
@@ -1334,11 +1333,10 @@ def cmd_interactive(
and ctx.agent is not _agent_for_ctx
)
if agent_swapped:
from ..EvoScientist import _ensure_config
from .agent import current_model_label
agent_loader.adopt(ctx.agent)
cfg = _ensure_config()
model = cfg.model
model = current_model_label()
state["status_base_snapshot"] = (
make_empty_status_snapshot(model)
)
@@ -1361,8 +1359,8 @@ def cmd_interactive(
# Commands that mutate status fields need an
# async refresh here (/compact + /new use sync
# callbacks; /model swaps the agent). /resume
# awaits its own refresh inline inside the
# callbacks; an agent swap rebuilds the snapshot).
# /resume awaits its own refresh inline inside the
# async callback.
if agent_swapped or _cmd.name in ("/compact", "/new"):
await _refresh_status_snapshot(
-35
View File
@@ -19,7 +19,6 @@ from collections.abc import Awaitable, Callable
from typing import Any
from rich.console import Console
from rich.table import Table
from ..commands.base import CommandUI
@@ -73,40 +72,6 @@ class RichCLICommandUI(CommandUI):
# Rich console flushes synchronously; nothing to await.
return
# ── /model interactive picker fallback ──────────────────
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None:
"""Print the model table and return ``None``; user re-runs with
``/model <name>`` since the CLI has no interactive picker."""
table = Table(
title="Available Models",
show_header=True,
header_style="bold cyan",
)
table.add_column("Name", style="bold")
table.add_column("Provider", style="dim")
for name, _mid, prov in entries:
marker = " *" if name == current_model and prov == current_provider else ""
table.add_row(f"{name}{marker}", prov)
self.console.print(table)
self.console.print(
"[dim]Usage: /model <name> [provider] [--save] — "
"provider is optional, auto-detected from model name[/dim]"
)
return None
def update_status_after_model_change(
self, new_model: str, new_provider: str | None = None
) -> None:
"""No-op; the CLI REPL refreshes status itself after detecting an
``ctx.agent`` change post-``cmd_manager.execute``."""
return
# ── Interactive pickers ────────────────────────────────
async def wait_for_thread_pick(
+3 -118
View File
@@ -222,13 +222,12 @@ async def _sync_tui_command_completion(
"""Adopt successful command-side state changes back into the TUI app."""
agent_swapped = ctx.agent is not None and ctx.agent is not original_agent
if agent_swapped:
from ..EvoScientist import _ensure_config
from .agent import current_model_label
app._agent_loader.adopt(ctx.agent)
cfg = _ensure_config()
update_model = getattr(app, "update_status_after_model_change", None)
if callable(update_model):
update_model(cfg.model, cfg.provider)
update_model(current_model_label())
# Rebind the runtime whenever the agent OR thread_id may have moved
# — ``/new`` and ``/resume`` rotate ``app._conversation_tid``
@@ -456,7 +455,6 @@ def run_textual_interactive(
self._picker_future: asyncio.Future | None = None
self._browser_future: asyncio.Future | None = None
self._mcp_browser_future: asyncio.Future | None = None
self._model_picker_future: asyncio.Future | None = None
self._history_suggester = HistorySuggester(DATA_DIR / "history")
self._history_index: int = -1 # -1 = not browsing history
self._history_saved_input: str = "" # saved current input before browsing
@@ -627,26 +625,6 @@ def run_textual_interactive(
return await self._wait_for_mcp_browse(browser)
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None:
from .widgets.model_picker import ModelPickerWidget
container = self.query_one("#chat", VerticalScroll)
picker = ModelPickerWidget(
entries,
current_model=current_model,
current_provider=current_provider,
)
await container.mount(picker)
self._anchor_chat(container)
picker.focus()
return await self._wait_for_model_pick(picker)
def clear_chat(self) -> None:
container = self.query_one("#chat", VerticalScroll)
welcome = self.query_one("#welcome", Static)
@@ -794,12 +772,6 @@ def run_textual_interactive(
yield Static("", id="status")
def on_mount(self) -> None:
# Register fallback middleware UI callback so messages appear
# as SystemMessage widgets in the chat container.
from ..middleware.model_fallback import set_ui_emit
set_ui_emit(lambda text, style: self._append_system(text, style))
self._render_welcome()
self._render_status()
self.set_interval(1.0, self._render_status)
@@ -1249,34 +1221,6 @@ def run_textual_interactive(
if self._mcp_browser_future and not self._mcp_browser_future.done():
self._mcp_browser_future.set_result(None)
async def _wait_for_model_pick(self, picker_widget) -> tuple[str, str] | None:
"""Wait for user to pick a model from ModelPickerWidget.
Returns ``(name, provider)`` or ``None`` on cancel/timeout.
"""
self._model_picker_future = asyncio.get_event_loop().create_future()
try:
return await asyncio.wait_for(self._model_picker_future, timeout=120)
except (TimeoutError, asyncio.CancelledError):
return None
finally:
self._model_picker_future = None
try:
picker_widget.remove()
except Exception:
_channel_logger.debug("model picker cleanup failed", exc_info=True)
self.query_one("#prompt", ChatTextArea).focus()
def on_model_picker_widget_picked(self, event) -> None: # type: ignore[override]
"""Handle ModelPickerWidget.Picked message."""
if self._model_picker_future and not self._model_picker_future.done():
self._model_picker_future.set_result((event.name, event.provider))
def on_model_picker_widget_cancelled(self, event) -> None: # type: ignore[override]
"""Handle ModelPickerWidget.Cancelled message."""
if self._model_picker_future and not self._model_picker_future.done():
self._model_picker_future.set_result(None)
# ── Streaming core ─────────────────────────────────────
async def _stream_with_widgets(
@@ -2495,7 +2439,6 @@ def run_textual_interactive(
if focused is not None:
from .widgets.approval_widget import ApprovalWidget
from .widgets.mcp_browser import MCPBrowserWidget
from .widgets.model_picker import ModelPickerWidget
from .widgets.skill_browser import SkillBrowserWidget
from .widgets.thread_selector import ThreadPickerWidget
@@ -2511,33 +2454,6 @@ def run_textual_interactive(
if isinstance(focused, MCPBrowserWidget):
focused.action_cancel()
return
# ModelPickerWidget: when in "Custom Ollama" input mode, its
# child Input widget owns focus, so ``focused`` isn't the
# picker itself. Walk the parent chain to find it, then let
# the widget's own action_cancel decide whether to close the
# picker (list mode) or just exit input mode.
picker: ModelPickerWidget | None = None
if isinstance(focused, ModelPickerWidget):
picker = focused
else:
node = focused.parent
while node is not None and not isinstance(node, ModelPickerWidget):
node = node.parent
picker = node
if picker is not None:
prev_mode = getattr(picker, "_mode", "list")
picker.action_cancel()
# In list mode action_cancel posted Cancelled; resolve the
# future immediately to avoid a frame of lag. In input
# mode action_cancel flipped back to list — keep picker
# open, do NOT close the future.
if prev_mode == "list":
if (
self._model_picker_future
and not self._model_picker_future.done()
):
self._model_picker_future.set_result(None)
return
if self._queued_messages:
self._queued_messages.pop()
self._render_queue_indicator()
@@ -2561,7 +2477,6 @@ def run_textual_interactive(
from .widgets.approval_widget import ApprovalWidget
from .widgets.ask_user_widget import AskUserWidget
from .widgets.mcp_browser import MCPBrowserWidget
from .widgets.model_picker import ModelPickerWidget
from .widgets.skill_browser import SkillBrowserWidget
from .widgets.thread_selector import ThreadPickerWidget
@@ -2580,20 +2495,6 @@ def run_textual_interactive(
if isinstance(focused, MCPBrowserWidget):
focused.action_move_up()
return
# ModelPickerWidget: Up from the Custom Ollama Input child
# must reach the picker (to exit input mode). See the Esc
# handler above for the parent-walk rationale.
picker_up: ModelPickerWidget | None = None
if isinstance(focused, ModelPickerWidget):
picker_up = focused
else:
node = focused.parent
while node is not None and not isinstance(node, ModelPickerWidget):
node = node.parent
picker_up = node
if picker_up is not None:
picker_up.action_move_up()
return
if self._queued_messages:
last = self._queued_messages.pop()
prompt = self.query_one("#prompt", ChatTextArea)
@@ -2629,7 +2530,6 @@ def run_textual_interactive(
from .widgets.approval_widget import ApprovalWidget
from .widgets.ask_user_widget import AskUserWidget
from .widgets.mcp_browser import MCPBrowserWidget
from .widgets.model_picker import ModelPickerWidget
from .widgets.skill_browser import SkillBrowserWidget
from .widgets.thread_selector import ThreadPickerWidget
@@ -2648,18 +2548,6 @@ def run_textual_interactive(
if isinstance(focused, MCPBrowserWidget):
focused.action_move_down()
return
# Same parent-walk rationale as action_edit_queued / cancel.
picker_down: ModelPickerWidget | None = None
if isinstance(focused, ModelPickerWidget):
picker_down = focused
else:
node = focused.parent
while node is not None and not isinstance(node, ModelPickerWidget):
node = node.parent
picker_down = node
if picker_down is not None:
picker_down.action_move_down()
return
# History browsing (down key)
if self._history_index >= 0:
@@ -2970,9 +2858,6 @@ def run_textual_interactive(
def _do_exit(self) -> None:
"""Clean up channels, unregister callbacks, and exit."""
from ..middleware.model_fallback import set_ui_emit
set_ui_emit(None)
if self._channel_timer is not None:
self._channel_timer.stop()
self._channel_timer = None
@@ -3078,7 +2963,7 @@ def run_textual_interactive(
def update_status_after_model_change(
self, new_model: str, new_provider: str | None = None
) -> None:
"""Update the status bar and welcome banner after /model switches the LLM."""
"""Update the status bar and welcome banner after an agent swap."""
self._current_model = new_model
if new_provider is not None:
self._current_provider = new_provider
-390
View File
@@ -1,390 +0,0 @@
"""Inline model picker widget for /model command in TUI.
Keyboard-driven widget mounted directly into the chat container.
Models are grouped by provider with a search/filter input.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, ClassVar, Literal
from rich.text import Text
from textual.binding import Binding, BindingType
from textual.containers import Container
from textual.message import Message
from textual.widget import Widget
from textual.widgets import Input, Static
if TYPE_CHECKING:
from textual import events
from textual.app import ComposeResult
# Sentinel ``model_id`` used for the "Custom Ollama model..." pseudo-row.
# Selecting this row switches the widget into free-text input mode instead
# of posting ``Picked`` — the user types a model name, Enter confirms.
_CUSTOM_OLLAMA_ID = "__custom_ollama__"
def _build_items(
entries: list[tuple[str, str, str]],
current_model: str | None = None,
current_provider: str | None = None,
filter_text: str = "",
) -> list[dict]:
"""Build the flat item list rendered by ModelPickerWidget.
Returns a list of::
{"type": "header", "label": str}
{"type": "model", "name": str, "model_id": str, "provider": str, "current": bool}
"""
# Apply filter. The Custom Ollama sentinel is the user's escape hatch
# when no local models match; it must remain visible regardless of filter.
if filter_text:
ft = filter_text.lower()
entries = [
(n, mid, p)
for n, mid, p in entries
if mid == _CUSTOM_OLLAMA_ID or ft in n.lower() or ft in p.lower()
]
# Group by provider preserving order. Deduplicate the Custom Ollama
# sentinel defensively — if callers somehow pass two sentinel rows
# (state reuse, stale merges), collapse them into one to avoid
# rendering duplicate "Custom Ollama model..." rows in the picker.
groups: dict[str, list[tuple[str, str, str]]] = {}
seen_sentinel = False
for name, model_id, provider in entries:
if model_id == _CUSTOM_OLLAMA_ID:
if seen_sentinel:
continue
seen_sentinel = True
if provider not in groups:
groups[provider] = []
groups[provider].append((name, model_id, provider))
items: list[dict] = []
for provider, models in groups.items():
items.append({"type": "header", "label": provider})
for name, model_id, prov in models:
is_current = name == current_model and prov == current_provider
items.append(
{
"type": "model",
"name": name,
"model_id": model_id,
"provider": prov,
"current": is_current,
}
)
return items
class ModelPickerWidget(Widget):
"""Inline model picker -- mounts in chat, keyboard-driven.
Posts ``Picked(name, provider)`` on Enter, ``Cancelled()`` on Esc.
Type to filter models.
"""
can_focus = True
# Required so the Custom Ollama ``Input`` child can hold focus when the
# user is typing a model name.
can_focus_children = True
DEFAULT_CSS = """
ModelPickerWidget {
height: auto;
max-height: 30;
margin: 1 0;
padding: 0 1;
background: $surface;
border: solid $primary;
}
ModelPickerWidget .picker-custom-input {
height: 3;
margin: 1 0 0 0;
}
ModelPickerWidget .picker-title {
height: 1;
text-style: bold;
color: $primary;
}
ModelPickerWidget .picker-filter {
height: 1;
padding: 0 1;
color: $text;
}
ModelPickerWidget .picker-rows {
height: auto;
max-height: 22;
overflow-y: auto;
}
ModelPickerWidget .picker-header {
height: 1;
padding: 0 1;
margin-top: 1;
}
ModelPickerWidget .picker-row {
height: 1;
padding: 0 1;
}
ModelPickerWidget .picker-row-selected {
background: $primary;
text-style: bold;
}
ModelPickerWidget .picker-help {
height: 1;
color: $text-muted;
text-style: italic;
}
"""
BINDINGS: ClassVar[list[BindingType]] = [
Binding("up", "move_up", "Up", show=False),
Binding("down", "move_down", "Down", show=False),
Binding("enter", "select", "Select", show=False),
Binding("escape", "cancel", "Cancel", show=False),
Binding("backspace", "backspace", "Backspace", show=False),
]
class Picked(Message):
def __init__(self, name: str, provider: str) -> None:
super().__init__()
self.name = name
self.provider = provider
class Cancelled(Message):
"""Posted when user cancels selection."""
def __init__(
self,
entries: list[tuple[str, str, str]],
*,
current_model: str | None = None,
current_provider: str | None = None,
title: str = ">>> Select model <<<",
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self._entries = entries
self._current_model = current_model
self._current_provider = current_provider
self._title = title
self._filter_text = ""
self._items = _build_items(
entries,
current_model=current_model,
current_provider=current_provider,
)
self._selected = self._first_model_index()
self._row_widgets: list[Static] = []
self._filter_widget: Static | None = None
# "list" = arrow-key selection over models; "input" = free-text entry
# for Custom Ollama model name. Transitions: selecting the sentinel
# row enters input mode; Esc or Up arrow inside input mode returns to
# list mode without closing the picker.
self._mode: Literal["list", "input"] = "list"
self._custom_input: Input | None = None
def _first_model_index(self) -> int:
for i, item in enumerate(self._items):
if item["type"] == "model":
return i
return 0
def _move(self, direction: int) -> None:
if not self._items:
return
i = (self._selected + direction) % len(self._items)
steps = 0
while self._items[i]["type"] != "model" and steps < len(self._items):
i = (i + direction) % len(self._items)
steps += 1
if self._items[i]["type"] == "model":
self._selected = i
self._update_rows()
def _rebuild(self) -> None:
"""Rebuild items from filter and re-render."""
self._items = _build_items(
self._entries,
current_model=self._current_model,
current_provider=self._current_provider,
filter_text=self._filter_text,
)
self._selected = self._first_model_index()
# Re-mount rows
rows_container = self.query_one(".picker-rows", Container)
for w in list(rows_container.children):
w.remove()
self._row_widgets.clear()
for item in self._items:
css = "picker-header" if item["type"] == "header" else "picker-row"
widget = Static("", classes=css)
self._row_widgets.append(widget)
rows_container.mount(widget)
self._update_rows()
self._update_filter()
def compose(self) -> ComposeResult:
yield Static(self._title, classes="picker-title")
self._filter_widget = Static("", classes="picker-filter")
yield self._filter_widget
with Container(classes="picker-rows"):
for item in self._items:
css = "picker-header" if item["type"] == "header" else "picker-row"
widget = Static("", classes=css)
self._row_widgets.append(widget)
yield widget
# Hidden until the user selects "Custom Ollama model..." \u2014 then shown
# and focused for free-text entry of an Ollama model name.
self._custom_input = Input(
placeholder="Type Ollama model name (e.g. llama3.3)...",
classes="picker-custom-input",
)
self._custom_input.display = False
yield self._custom_input
yield Static(
"\u2191/\u2193 navigate \u00b7 Enter select \u00b7 Type to filter \u00b7 Esc cancel",
classes="picker-help",
)
def on_mount(self) -> None:
self._update_rows()
self._update_filter()
self.call_later(self.focus)
def _update_filter(self) -> None:
if self._filter_widget is not None:
if self._filter_text:
t = Text()
t.append(" Filter: ", style="dim")
t.append(self._filter_text, style="bold")
t.append("\u2588", style="blink")
self._filter_widget.update(t)
else:
self._filter_widget.update(
Text(" Type to filter...", style="dim italic")
)
def _update_rows(self) -> None:
for i, (item, widget) in enumerate(
zip(self._items, self._row_widgets, strict=False)
):
widget.remove_class("picker-row-selected")
if item["type"] == "header":
t = Text()
t.append("\u2500\u2500 ", style="bold cyan")
t.append(item["label"], style="bold cyan")
widget.update(t)
else:
is_selected = i == self._selected
t = Text()
cursor = "\u25b8 " if is_selected else " "
t.append(cursor, style="bold cyan" if is_selected else "dim")
t.append(item["name"], style="bold" if is_selected else "")
if item["current"]:
t.append(" *", style="bold green")
t.append(f" ({item['provider']})", style="dim italic")
widget.update(t)
if is_selected:
widget.add_class("picker-row-selected")
widget.scroll_visible()
def on_key(self, event: events.Key) -> None:
# In input mode, the Input child owns printable keys + backspace.
if self._mode == "input":
return
# Let bindings handle special keys
if event.key in ("up", "down", "enter", "escape", "backspace"):
return
# Printable characters -> filter
if event.character and event.character.isprintable():
self._filter_text += event.character
self._rebuild()
event.prevent_default()
def action_backspace(self) -> None:
if self._mode == "input":
# Input widget handles its own backspace.
return
if self._filter_text:
self._filter_text = self._filter_text[:-1]
self._rebuild()
def action_move_up(self) -> None:
if self._mode == "input":
# Up from the Input field escapes back to list selection.
self._exit_input_mode()
return
self._move(-1)
def action_move_down(self) -> None:
if self._mode == "input":
# Down in input mode is ambiguous; absorb rather than toggle.
return
self._move(1)
def action_select(self) -> None:
if self._mode == "input":
self._submit_custom_input()
return
if not self._items or self._selected >= len(self._items):
self.post_message(self.Cancelled())
return
item = self._items[self._selected]
if item["type"] != "model":
self.post_message(self.Cancelled())
return
if item["provider"] == "ollama" and item["model_id"] == _CUSTOM_OLLAMA_ID:
self._enter_input_mode()
return
self.post_message(self.Picked(item["name"], item["provider"]))
def action_cancel(self) -> None:
if self._mode == "input":
# Esc returns to list selection; does NOT close the picker.
self._exit_input_mode()
return
self.post_message(self.Cancelled())
def on_blur(self, event: events.Blur) -> None:
# When the Input child has focus we must NOT steal it back.
if self._mode == "input":
return
self.call_after_refresh(self.focus)
def on_input_submitted(self, event: Input.Submitted) -> None:
"""Safety net: Enter fired inside the Input widget rather than
bubbling to ``action_select``. Route to the same submit path."""
if event.input is self._custom_input:
event.stop()
self._submit_custom_input()
def _enter_input_mode(self) -> None:
"""Show the Custom Ollama Input and move focus into it."""
self._mode = "input"
if self._custom_input is not None:
self._custom_input.display = True
# Carry any filter text over as a nice touch — user may have
# started typing a model name thinking it would filter.
self._custom_input.value = self._filter_text
self._custom_input.focus()
def _exit_input_mode(self) -> None:
"""Hide the Input and return focus to the list."""
self._mode = "list"
if self._custom_input is not None:
self._custom_input.display = False
self._custom_input.value = ""
self.focus()
def _submit_custom_input(self) -> None:
"""Confirm the typed Ollama model name. Empty input is a no-op —
user can Esc out or keep typing."""
typed = (self._custom_input.value if self._custom_input else "").strip()
if not typed:
return
self.post_message(self.Picked(typed, "ollama"))
+1 -1
View File
@@ -12,7 +12,7 @@ class CompletionKind(StrEnum):
EMPTY = "empty"
_CATEGORY_ORDER = ["Session", "Skills", "MCP", "Channels", "Model", "General"]
_CATEGORY_ORDER = ["Session", "Skills", "MCP", "Channels", "General"]
@dataclass(frozen=True)
-6
View File
@@ -47,12 +47,6 @@ class CommandUI(Protocol):
async def wait_for_mcp_browse(
self, servers: list, installed_names: set[str], pre_filter_tag: str
) -> list | None: ...
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None: ...
def clear_chat(self) -> None: ...
def request_quit(self) -> None: ...
def force_quit(self) -> None: ...
@@ -5,8 +5,6 @@ from . import (
channel,
general,
mcp,
model,
model_fallback,
schedule,
session,
skills,
@@ -17,8 +15,6 @@ __all__ = [
"channel",
"general",
"mcp",
"model",
"model_fallback",
"schedule",
"session",
"skills",
@@ -1,204 +0,0 @@
from __future__ import annotations
from typing import ClassVar
from ..base import Argument, Command, CommandContext
from ..manager import manager
def extract_model_and_provider(args: list[str]) -> tuple[str, str]:
"""Parse model name and provider from argument list.
Args:
args: Non-empty argument list (model_name [provider]).
Returns:
``(model_name, provider)`` tuple.
Raises:
ValueError: If the model is not in the registry. Skipped when
``provider_override == "ollama"``, since Ollama models are
locally-installed and never appear in ``MODELS``.
"""
from ...llm.models import MODELS
model_name = args[0]
provider_override = args[1] if len(args) > 1 else None
# Ollama models are locally-installed — not in the registry. Pass the name
# through verbatim; get_chat_model's "Assume full model ID" fallback
# (models.py) accepts them.
if provider_override == "ollama":
return model_name, "ollama"
if model_name not in MODELS:
raise ValueError(f"Unknown model '{model_name}'")
if provider_override:
provider = provider_override
else:
_, provider = MODELS[model_name]
return model_name, provider
class ModelCommand(Command):
"""Switch the LLM model for the current session."""
name = "/model"
description = "Switch model (--save to persist)"
category = "Model"
# ``--save`` is parsed manually in ``execute`` via ``"--save" in args``;
# ``type=bool`` below is declarative metadata, not enforced by the manager.
arguments: ClassVar[list[Argument]] = [
Argument(
name="model_name",
type=str,
description="Model short name (e.g. claude-sonnet-4-6). Opens picker if omitted.",
required=False,
),
Argument(
name="--save",
type=bool,
description="Save the choice to config file",
required=False,
),
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...EvoScientist import _ensure_config
from ...llm.models import list_model_picker_entries
cfg = _ensure_config()
current_model = cfg.model
current_provider = cfg.provider
# Parse --save flag
save = "--save" in args
args = [a for a in args if a != "--save"]
if args:
try:
model_name, provider = extract_model_and_provider(args)
except ValueError:
ctx.ui.append_system(
f"Unknown model '{args[0]}'. Use /model to browse available models.",
style="red",
)
return
await self._apply_model(ctx, model_name, provider, save=save)
return
# Interactive picker
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model <name> [provider] [--save]",
style="yellow",
)
return
entries = await list_model_picker_entries(
getattr(cfg, "ollama_base_url", None),
include_custom_ollama=True,
)
result = await ctx.ui.wait_for_model_pick(
entries,
current_model=current_model,
current_provider=current_provider,
)
if result is None:
return
name, provider = result
# Defense-in-depth: the widget should have replaced the sentinel with
# the user-typed name. If it didn't, treat as cancel rather than try
# to switch to a literal "__custom_ollama__" model.
if provider == "ollama" and name in (
"Custom Ollama model...",
"__custom_ollama__",
):
return
await self._apply_model(ctx, name, provider, save=save)
async def _apply_model(
self,
ctx: CommandContext,
model_name: str,
provider: str,
*,
save: bool = False,
) -> None:
import copy
from ...cli.agent import _load_agent
from ...EvoScientist import (
_build_chat_model,
_ensure_config,
set_active_config,
set_chat_model_instance,
)
cfg = _ensure_config()
# Build a temporary config + its chat model and verify the agent can be
# built before committing anything. ``create_cli_agent(config=...,
# chat_model=...)`` is pure (issue #183) — it writes none of the cached
# config/model module globals — so a failure below leaves the session
# on the original model with no snapshot/restore needed.
temp_cfg = copy.copy(cfg)
temp_cfg.model = model_name
temp_cfg.provider = provider
try:
new_chat_model = _build_chat_model(temp_cfg)
new_agent = _load_agent(
workspace_dir=ctx.workspace_dir,
checkpointer=ctx.checkpointer,
config=temp_cfg,
chat_model=new_chat_model,
)
except Exception as e:
ctx.ui.append_system(f"Failed to switch model: {e}", style="red")
return
# Agent built with no global mutation — commit the switch atomically.
# These are pure assignments and cannot fail, so the session can never
# be left half-switched. Apply the switch to the LIVE ``cfg`` in place
# (the active config object) instead of rebinding ``_config`` to the
# fresh ``temp_cfg`` — callers that hold the active config by reference
# (e.g. serve's ``agent_holder["config"]`` and its workspace-changing
# ``/resume`` reload) must observe the new model/provider. The verify
# build above used the ``temp_cfg`` copy, so a failed build never reaches
# here and the live ``cfg`` stays untouched (failure still no-ops).
cfg.model = model_name
cfg.provider = provider
set_active_config(cfg)
set_chat_model_instance(new_chat_model, (model_name, provider))
ctx.agent = new_agent
# Persist to config file if --save was given
if save:
from ...config.settings import set_config_value
set_config_value("model", model_name)
set_config_value("provider", provider)
# Propagate to the channel runtime if channels are running so the
# bus picks up the new agent on the next inbound message.
if ctx.channel_runtime is not None and ctx.channel_runtime.agent is not None:
ctx.channel_runtime.agent = new_agent
# Update status bar if available
update_model_fn = getattr(ctx.ui, "update_status_after_model_change", None)
if callable(update_model_fn):
update_model_fn(model_name, provider)
saved_note = " (saved to config)" if save else ""
ctx.ui.append_system(
f"Switched to {model_name} ({provider}){saved_note}", style="green"
)
manager.register(ModelCommand())
@@ -1,304 +0,0 @@
"""Slash command for managing the model fallback chain.
Provides ``/model-fallback`` (alias ``/fallback``) with subcommands to
add, remove, list, clear, save, and display help for fallback models.
"""
from __future__ import annotations
from typing import ClassVar
from ..base import Argument, Command, CommandContext, SubCommand
from ..manager import manager
class ModelFallbackCommand(Command):
"""Manage the model fallback chain."""
name = "/model-fallback"
alias: ClassVar[list[str]] = ["/fallback"]
description = "Manage fallback models (add/remove/list/clear)"
category = "Model"
arguments: ClassVar[list[Argument]] = [
Argument(
name="action",
type=str,
description="add|remove|list|clear|save|help",
required=False,
),
]
subcommands: ClassVar[list[SubCommand]] = [
SubCommand("list", "Display the current fallback chain"),
SubCommand("add", "Append a model to the fallback chain"),
SubCommand("remove", "Remove a model by position"),
SubCommand("clear", "Remove all fallback entries"),
SubCommand("save", "Persist the chain to config"),
SubCommand("help", "Show subcommand reference"),
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...llm.models import MODELS
from ...middleware.model_fallback import (
add_fallback,
clear_fallbacks,
get_fallback_chain,
remove_fallback_at,
serialize_fallback_chain,
)
save = "--save" in args
args = [a for a in args if a != "--save"]
if not args:
await self._show_list(ctx, get_fallback_chain())
return
action = args[0].lower()
if action == "list":
await self._show_list(ctx, get_fallback_chain())
elif action == "add":
if len(args) >= 2:
model_name = args[1]
provider = args[2] if len(args) > 2 else None
if provider is None:
if model_name in MODELS:
_, provider = MODELS[model_name]
else:
ctx.ui.append_system(
f"Unknown model '{model_name}'. Specify provider explicitly: "
f"/model-fallback add {model_name} <provider>",
style="red",
)
return
else:
picked = await self._pick_model(ctx)
if picked is None:
return
model_name, provider = picked
if add_fallback(model_name, provider):
ctx.ui.append_system(
f"Added {model_name} ({provider}) to fallback chain", style="green"
)
else:
ctx.ui.append_system(
f"{model_name} ({provider}) is already in the fallback chain",
style="yellow",
)
return
if save:
self._save_to_config(serialize_fallback_chain())
elif action == "remove":
chain = get_fallback_chain()
if not chain:
ctx.ui.append_system("Fallback chain is empty", style="yellow")
return
if len(args) >= 2:
arg = args[1]
try:
idx = int(arg) - 1
except ValueError:
ctx.ui.append_system(
f"Expected a position number (1-{len(chain)}), got '{arg}'. "
"Use /model-fallback list to see positions.",
style="red",
)
return
removed = remove_fallback_at(idx)
if removed is None:
ctx.ui.append_system(
f"Invalid position {arg}. "
f"Use a number between 1 and {len(chain)}.",
style="red",
)
return
model_name, provider = removed
else:
picked = await self._pick_fallback_to_remove(ctx, chain)
if picked is None:
return
model_name, provider = picked
live_chain = get_fallback_chain()
try:
idx = live_chain.index((model_name, provider))
except ValueError:
ctx.ui.append_system(
f"{model_name} ({provider}) is no longer in the fallback chain",
style="yellow",
)
return
remove_fallback_at(idx)
ctx.ui.append_system(
f"Removed {model_name} ({provider}) from fallback chain",
style="green",
)
if save:
self._save_to_config(serialize_fallback_chain())
elif action == "clear":
clear_fallbacks()
ctx.ui.append_system("Cleared all fallback models", style="green")
if save:
self._save_to_config("")
elif action == "save":
self._save_to_config(serialize_fallback_chain())
ctx.ui.append_system("Fallback chain saved to config", style="green")
elif action == "help":
self._show_help(ctx)
else:
self._show_help(ctx)
async def _pick_model(self, ctx: CommandContext) -> tuple[str, str] | None:
"""Open the interactive model picker to select a fallback model.
Falls back to a usage hint when the UI does not support interactive
widgets (CLI mode without a model argument).
Args:
ctx: Current command context.
Returns:
``(model_name, provider)`` tuple, or ``None`` if cancelled.
"""
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model-fallback add <model> [provider]", style="yellow"
)
return None
from ...EvoScientist import _ensure_config
from ...llm.models import list_model_picker_entries
cfg = _ensure_config()
entries = await list_model_picker_entries(
getattr(cfg, "ollama_base_url", None),
include_custom_ollama=True,
)
result = await ctx.ui.wait_for_model_pick(
entries,
current_model=cfg.model,
current_provider=cfg.provider,
)
if result is None:
return None
name, provider = result
if provider == "ollama" and name in (
"Custom Ollama model...",
"__custom_ollama__",
):
return None
return name, provider
async def _pick_fallback_to_remove(
self, ctx: CommandContext, chain: list[tuple[str, str]]
) -> tuple[str, str] | None:
"""Open the model picker populated with the current fallback chain.
Falls back to a usage hint in CLI mode.
Args:
ctx: Current command context.
chain: The current fallback chain to choose from.
Returns:
``(model_name, provider)`` tuple, or ``None`` if cancelled.
"""
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model-fallback remove <position> "
"(use /model-fallback list to see positions)",
style="yellow",
)
return None
entries = [(m, m, p) for m, p in chain]
result = await ctx.ui.wait_for_model_pick(
entries, current_model=None, current_provider=None
)
if result is None:
return None
return result
def _show_help(self, ctx: CommandContext) -> None:
"""Render the subcommand reference table.
Args:
ctx: Current command context.
"""
from rich.text import Text
text = Text("/model-fallback subcommands:\n", style="bold")
for cmd, desc in (
(
"add [model] [provider]",
"Add a fallback model (opens picker if omitted)",
),
(
"remove [position]",
"Remove a fallback by position (opens picker in TUI)",
),
("list", "Show the current fallback chain"),
("clear", "Remove all fallback models"),
("save", "Save current fallback chain to config file"),
("help", "Show this help message"),
):
text.append(f" {cmd:<26}", style="cyan")
text.append(f"{desc}\n", style="dim")
text.append(
"\nAdd --save to add/remove/clear to persist the change immediately.\n",
style="dim",
)
ctx.ui.mount_renderable(text)
async def _show_list(
self, ctx: CommandContext, chain: list[tuple[str, str]]
) -> None:
"""Display the current fallback chain as a numbered list.
Args:
ctx: Current command context.
chain: The fallback chain to display.
"""
if not chain:
ctx.ui.append_system("No fallback models configured", style="dim")
ctx.ui.append_system(
"Use /model-fallback add <model> [provider] to add one",
style="dim",
)
return
from rich.text import Text
text = Text("Fallback chain:\n", style="bold")
for idx, (model, provider) in enumerate(chain, 1):
text.append(f" {idx}. ", style="dim")
text.append(model, style="cyan")
text.append(f" ({provider})\n", style="dim")
ctx.ui.mount_renderable(text)
def _save_to_config(self, value: str) -> None:
"""Persist the fallback chain string to the config file.
Args:
value: Serialized chain (``"model:provider,..."``).
"""
from ...config.settings import set_config_value
set_config_value("model_fallbacks", value)
manager.register(ModelFallbackCommand())
+4
View File
@@ -19,7 +19,9 @@ from .settings import (
get_config_dir,
get_config_path,
get_config_value,
get_default_workspace_dir,
get_effective_config,
is_config_applied_env,
list_config,
load_config,
reset_config,
@@ -39,7 +41,9 @@ __all__ = [
"get_config_dir",
"get_config_path",
"get_config_value",
"get_default_workspace_dir",
"get_effective_config",
"is_config_applied_env",
"list_config",
"load_config",
"reset_config",
+132
View File
@@ -0,0 +1,132 @@
"""Startup detection of pre-Registry legacy model configuration artifacts.
Design doc section 10 step 4: this project is in development and keeps no
historical compatibility. When the config service or the CLI finds legacy
artifacts — an old ``providers.yaml``, an old
``run-runtime-snapshots.sqlite3``, or LLM fields left behind in
``config.yaml`` — it must refuse to start and log an explicit reset guide
instead of partially reading them.
"""
from __future__ import annotations
import logging
from pathlib import Path
import yaml
from .settings import get_config_dir
logger = logging.getLogger(__name__)
#: LLM configuration keys that ``config.yaml`` must no longer carry
#: (section 10 step 2). Platform configuration (workspace, MCP, ports,
#: storage, scheduling, security) stays; everything model/provider/credential
#: related lives in the Model Registry (``model-runtime.sqlite3``) only.
LEGACY_CONFIG_YAML_KEYS = frozenset(
{
"provider",
"model",
"model_catalog",
"model_fallbacks",
"auxiliary_provider",
"auxiliary_model",
"anthropic_api_key",
"anthropic_base_url",
"anthropic_auth_mode",
"openai_api_key",
"openai_auth_mode",
"nvidia_api_key",
"google_api_key",
"minimax_api_key",
"minimax_base_url",
"siliconflow_api_key",
"openrouter_api_key",
"deepseek_api_key",
"zhipu_api_key",
"volcengine_api_key",
"dashscope_api_key",
"moonshot_api_key",
"kimi_api_key",
"custom_openai_api_key",
"custom_openai_base_url",
"custom_anthropic_api_key",
"custom_anthropic_base_url",
"ollama_base_url",
"use_responses_api",
"openrouter_anthropic_prompt_cache",
}
)
_LEGACY_PROVIDERS_FILE = "providers.yaml"
_LEGACY_SNAPSHOTS_DB = "run-runtime-snapshots.sqlite3"
_RESET_GUIDANCE = """\
EvoScientist no longer reads legacy model configuration (design doc §10).
To reset the development environment:
1. Delete {config_dir}/providers.yaml (Provider Profiles are superseded
by the Model Registry).
2. Delete {config_dir}/run-runtime-snapshots.sqlite3 (the old snapshot
store; run snapshots now live in model-runtime.sqlite3).
3. Remove the leftover LLM fields listed above from {config_path} —
platform fields (workspace, MCP, ports, scheduling, security) stay.
4. Configure providers/models through the Model Registry (WebUI
configuration page or the model-registry API), run the provider test,
then enable the models you need.
Startup is refused — no legacy artifact is read, even partially.\
"""
class LegacyArtifactsError(RuntimeError):
"""Raised at startup when pre-Registry configuration artifacts remain."""
def find_legacy_artifacts(config_dir: Path | None = None) -> list[str]:
"""Return human-readable descriptions of every legacy artifact found."""
config_dir = config_dir if config_dir is not None else get_config_dir()
found: list[str] = []
providers_yaml = config_dir / _LEGACY_PROVIDERS_FILE
if providers_yaml.exists():
found.append(f"legacy Provider Profiles file: {providers_yaml}")
snapshots_db = config_dir / _LEGACY_SNAPSHOTS_DB
if snapshots_db.exists():
found.append(f"legacy run snapshot database: {snapshots_db}")
config_path = config_dir / "config.yaml"
if config_path.exists():
try:
with open(config_path, encoding="utf-8") as handle:
data = yaml.safe_load(handle) or {}
except yaml.YAMLError:
data = {}
if isinstance(data, dict):
leftover = sorted(LEGACY_CONFIG_YAML_KEYS & data.keys())
if leftover:
found.append(
f"leftover LLM fields in {config_path}: {', '.join(leftover)}"
)
return found
def assert_no_legacy_artifacts(config_dir: Path | None = None) -> None:
"""Refuse startup when any legacy model configuration artifact remains.
Logs the findings plus the explicit reset guide (section 10 step 4) and
raises :class:`LegacyArtifactsError`. Nothing is read partially: the
caller must not catch-and-continue.
"""
found = find_legacy_artifacts(config_dir)
if not found:
return
config_dir = config_dir if config_dir is not None else get_config_dir()
guidance = _RESET_GUIDANCE.format(
config_dir=config_dir,
config_path=config_dir / "config.yaml",
)
message = "Legacy model configuration artifacts detected:\n" + "\n".join(
f" - {item}" for item in found
)
logger.error("%s\n%s", message, guidance)
raise LegacyArtifactsError(f"{message}\n\n{guidance}")
+2 -3
View File
@@ -7,13 +7,12 @@ Everything else lives in submodules — import directly from them:
``STEPS``, ``render_progress``
- :mod:`EvoScientist.config.onboard.steps` — per-step functions
- :mod:`EvoScientist.config.onboard.channels` — channel selection + setup
- :mod:`EvoScientist.config.onboard.helpers` — API-key prompt, ccproxy,
- :mod:`EvoScientist.config.onboard.helpers` — API-key prompt,
npx/node, LaTeX, iMessage helpers
- :mod:`EvoScientist.config.onboard.style` — Rich styles + ``_checkbox_ask``
- :mod:`EvoScientist.config.onboard.validators` — input validators
- :mod:`EvoScientist.config.onboard.prompter` — ``NonInteractivePrompter``
(CLI-answer container) + ``select_navigation_active`` / ``GoBack`` for
keyboard nav
(CLI-answer container) + ``select_navigation_active`` for keyboard nav
- :mod:`EvoScientist.config.onboard.constants` — canonical valid-value sets
This module used to re-export every symbol from every submodule for
+3 -27
View File
@@ -2,46 +2,22 @@
The interactive ``_step_*`` functions in ``steps.py`` use these for ``Choice``
construction (or are checked against them by tests). The CLI ``onboard``
command in ``cli/commands.py`` uses them to validate ``--provider`` /
``--ui`` / ``--workspace-mode`` flag inputs.
command in ``cli/commands.py`` uses them to validate ``--ui`` /
``--workspace-mode`` flag inputs.
Single source of truth — adding a new provider here AND to the corresponding
Single source of truth — adding a new value here AND to the corresponding
``Choice(value=...)`` in ``steps.py`` is required; a drift test in
``tests/test_onboard.py`` keeps both sides in sync.
"""
from __future__ import annotations
VALID_PROVIDERS: frozenset[str] = frozenset(
{
"anthropic",
"openai",
"google-genai",
"minimax",
"zhipu",
"zhipu-code",
"volcengine",
"dashscope",
"dashscope-code",
"deepseek",
"moonshot",
"kimi-coding",
"ollama",
"nvidia",
"siliconflow",
"openrouter",
"custom-openai",
"custom-anthropic",
}
)
VALID_UI_BACKENDS: frozenset[str] = frozenset({"tui", "cli", "webui"})
VALID_WORKSPACE_MODES: frozenset[str] = frozenset({"daemon", "run"})
__all__ = [
"VALID_PROVIDERS",
"VALID_UI_BACKENDS",
"VALID_WORKSPACE_MODES",
]
+1 -198
View File
@@ -1,4 +1,4 @@
"""Cross-step helpers: API key prompt loop, ccproxy login, npx/node bootstrapping,
"""Cross-step helpers: API key prompt loop, npx/node bootstrapping,
LaTeX detection/install, iMessage setup.
"""
@@ -11,126 +11,7 @@ import sys
import questionary
from ..settings import EvoScientistConfig
from .style import QMARK, WIZARD_STYLE, console
from .validators import (
validate_anthropic_key,
validate_dashscope_code_key,
validate_dashscope_key,
validate_deepseek_key,
validate_google_key,
validate_kimi_key,
validate_minimax_key,
validate_moonshot_key,
validate_nvidia_key,
validate_openai_key,
validate_openrouter_key,
validate_siliconflow_key,
validate_volcengine_key,
validate_zhipu_key,
)
def _provider_key_info(config: EvoScientistConfig, provider: str):
"""Return (display_name, current_value, validate_fn) for a provider."""
mapping = {
"anthropic": (
"Anthropic",
config.anthropic_api_key or os.environ.get("ANTHROPIC_API_KEY", ""),
validate_anthropic_key,
),
"minimax": (
"MiniMax",
config.minimax_api_key or os.environ.get("MINIMAX_API_KEY", ""),
lambda key: validate_minimax_key(
key,
base_url=config.minimax_base_url
or os.environ.get(
"MINIMAX_BASE_URL", "https://api.minimaxi.com/anthropic"
),
),
),
"nvidia": (
"NVIDIA",
config.nvidia_api_key or os.environ.get("NVIDIA_API_KEY", ""),
validate_nvidia_key,
),
"google-genai": (
"Google",
config.google_api_key or os.environ.get("GOOGLE_API_KEY", ""),
validate_google_key,
),
"siliconflow": (
"SiliconFlow",
config.siliconflow_api_key or os.environ.get("SILICONFLOW_API_KEY", ""),
validate_siliconflow_key,
),
"openrouter": (
"OpenRouter",
config.openrouter_api_key or os.environ.get("OPENROUTER_API_KEY", ""),
validate_openrouter_key,
),
"deepseek": (
"DeepSeek",
config.deepseek_api_key or os.environ.get("DEEPSEEK_API_KEY", ""),
validate_deepseek_key,
),
"zhipu": (
"ZhipuAI",
config.zhipu_api_key or os.environ.get("ZHIPU_API_KEY", ""),
validate_zhipu_key,
),
"zhipu-code": (
"ZhipuAI CodePlan",
config.zhipu_api_key or os.environ.get("ZHIPU_API_KEY", ""),
validate_zhipu_key,
),
"volcengine": (
"Volcengine",
config.volcengine_api_key or os.environ.get("VOLCENGINE_API_KEY", ""),
validate_volcengine_key,
),
"dashscope": (
"DashScope",
config.dashscope_api_key or os.environ.get("DASHSCOPE_API_KEY", ""),
validate_dashscope_key,
),
"dashscope-code": (
"DashScope Coding Plan",
config.dashscope_api_key or os.environ.get("DASHSCOPE_API_KEY", ""),
validate_dashscope_code_key,
),
"moonshot": (
"Moonshot",
config.moonshot_api_key or os.environ.get("MOONSHOT_API_KEY", ""),
validate_moonshot_key,
),
"kimi-coding": (
"Kimi Coding Plan",
config.kimi_api_key or os.environ.get("KIMI_API_KEY", ""),
validate_kimi_key,
),
"custom-openai": (
"OpenAI-compatible",
config.custom_openai_api_key or os.environ.get("CUSTOM_OPENAI_API_KEY", ""),
None,
),
"custom-anthropic": (
"Custom Anthropic",
config.custom_anthropic_api_key
or os.environ.get("CUSTOM_ANTHROPIC_API_KEY", ""),
None,
),
"ollama": ("Ollama", "__no_key__", None),
}
return mapping.get(
provider,
(
"OpenAI",
config.openai_api_key or os.environ.get("OPENAI_API_KEY", ""),
validate_openai_key,
),
)
def _prompt_and_validate_api_key(
@@ -192,66 +73,6 @@ def _prompt_and_validate_api_key(
return new_key or None
def _prompt_ccproxy_port(config: EvoScientistConfig) -> None:
"""Prompt the user for a ccproxy port and save it to config."""
def valid_port(value: str) -> bool:
if not value: # empty = keep default
return True
try:
return 0 < int(value) < 2**16
except (ValueError, TypeError):
return False
current_port = getattr(config, "ccproxy_port", 8000)
try:
raw = questionary.text(
f"Enter port number for ccproxy to run on (Current: {current_port}, Enter to keep):",
validate=valid_port,
style=WIZARD_STYLE,
qmark=QMARK,
).ask()
if raw is None:
raise KeyboardInterrupt()
raw = raw.strip()
ccproxy_port = int(raw) if raw else current_port
except (ValueError, TypeError):
ccproxy_port = current_port
console.print(f" [dim]Using default port: {ccproxy_port}[/dim]")
config.ccproxy_port = ccproxy_port
console.print(
f" [green]✓ ccproxy will run on http://127.0.0.1:{ccproxy_port}[/green]"
)
def _run_ccproxy_login(provider: str, label: str) -> None:
"""Run ccproxy auth login for the given provider and show status."""
from ...ccproxy_manager import _ccproxy_exe, check_ccproxy_auth
console.print(" [dim]Opening browser for authentication...[/dim]")
try:
proc = subprocess.run(
[_ccproxy_exe() or "ccproxy", "auth", "login", provider],
capture_output=True,
text=True,
timeout=120,
)
for line in proc.stdout.splitlines():
if line.strip().startswith("https://"):
console.print(f" [dim]Visit: {line.strip()}[/dim]")
break
authed, msg = check_ccproxy_auth(provider)
if authed:
console.print(f" [green]✓ {label}: {msg}[/green]")
else:
console.print(f" [red]Authentication failed: {msg}[/red]")
except subprocess.TimeoutExpired:
console.print(" [red]Login timed out.[/red]")
except Exception as exc:
console.print(f" [red]Login error: {exc}[/red]")
def _check_npx() -> bool:
"""Check if npx is available on the system.
@@ -591,24 +412,6 @@ def validate_imessage() -> tuple[bool, str]:
return True, f"imsg{version_str} at {cli_path}"
def _install_ccproxy() -> bool:
"""Run pip install for ccproxy (evoscientist[oauth]).
Uses uv pip install when available (uv-managed envs don't ship pip).
Returns:
True if installation succeeded and ccproxy is available.
"""
from ...ccproxy_manager import is_ccproxy_available
from ...mcp.registry import install_library
ok = install_library("evoscientist[oauth]")
if not ok:
console.print(" [red]✗ Installation failed.[/red]")
return False
return is_ccproxy_available()
def _install_imsg() -> bool:
"""Run brew install for imsg CLI.
+3 -30
View File
@@ -5,31 +5,13 @@ from __future__ import annotations
from typing import Any
class GoBack(Exception):
"""Raised inside the provider sub-loop to rewind to provider selection."""
# Sentinel value the back-keybinding writes into the prompt result, and the
# value the trailing ``← Back`` menu item carries. Same string so the two
# code paths converge to a single ``GoBack`` raise.
BACK_SENTINEL = "__back__"
def install_navigation_keys(
question,
*,
with_back: bool = False,
sentinel: str = BACK_SENTINEL,
) -> None:
def install_navigation_keys(question) -> None:
"""Add keyboard shortcuts on a questionary select ``Question``.
Bindings (merged in front of questionary's defaults — Ctrl+C/Ctrl+D still
cancel the wizard):
- ``→`` — accept the option under the cursor and advance (mirrors Enter).
- ``Esc`` / ``←`` (only when ``with_back=True``) — exit with ``sentinel``
so the wizard can rewind. Used in the provider sub-loop's auth_mode
prompts.
"""
from prompt_toolkit.key_binding import KeyBindings, merge_key_bindings
@@ -49,13 +31,6 @@ def install_navigation_keys(
event.app.exit(result=pointed.value)
return
if with_back:
@kb.add("escape", eager=True)
@kb.add("left", eager=True)
def _back(event):
event.app.exit(result=sentinel)
question.application.key_bindings = merge_key_bindings(
[kb, question.application.key_bindings]
)
@@ -78,7 +53,7 @@ def select_navigation_active():
def _wrapped(*args, **kwargs):
q = original(*args, **kwargs)
try:
install_navigation_keys(q, with_back=False)
install_navigation_keys(q)
except Exception:
# Don't let a stray keybinding error block the wizard.
pass
@@ -92,7 +67,7 @@ def select_navigation_active():
class NonInteractivePrompter:
"""Container for CLI-supplied wizard answers (``--provider``, ``--model``…).
"""Container for CLI-supplied wizard answers (``--ui``, ``--tavily-key``…).
``strict=True`` makes missing presets fatal instead of falling back to
interactive. Wizard reads ``answers`` / ``skip_set`` / ``strict`` directly.
@@ -113,8 +88,6 @@ class NonInteractivePrompter:
__all__ = [
"BACK_SENTINEL",
"GoBack",
"NonInteractivePrompter",
"install_navigation_keys",
"select_navigation_active",
+2 -636
View File
@@ -1,8 +1,8 @@
"""Individual wizard step functions.
Each ``_step_*`` prompts the user for one logical decision and returns the
chosen value. Conditional steps (auth mode, base URL) are only called by
``run_onboard`` when the provider needs them.
chosen value. LLM provider/model/API-key configuration no longer happens in
the wizard — models are configured via the WebUI / model registry.
"""
from __future__ import annotations
@@ -14,24 +14,17 @@ import questionary
from prompt_toolkit.formatted_text import FormattedText
from questionary import Choice
from ...llm import get_models_for_provider
from ...llm.ollama_discovery import validate_ollama_connection
from ..settings import EvoScientistConfig
from .helpers import (
_auto_install_latexmk,
_check_latex_components,
_detect_tinytex_install_method,
_ensure_npx,
_install_ccproxy,
_install_tinytex,
_print_latex_status,
_prompt_and_validate_api_key,
_prompt_ccproxy_port,
_provider_key_info,
_run_ccproxy_login,
)
from .style import (
CONFIRM_STYLE,
QMARK,
WIZARD_STYLE,
_checkbox_ask,
@@ -228,633 +221,6 @@ def _step_webui_port(config: EvoScientistConfig) -> int:
return port
def _step_provider(
config: EvoScientistConfig,
*,
label: str | None = None,
default_value: str | None = None,
) -> str:
"""Step 1: Select LLM provider.
Args:
config: Current configuration.
label: Optional role label (e.g. "co-pilot") to clarify which model this
provider is for. When omitted, the generic main-model prompt is used.
default_value: Preselect this provider instead of ``config.provider``
(e.g. the auxiliary provider when configuring the co-pilot).
Returns:
Selected provider name.
"""
choices = [
# Direct providers
Choice(title="Anthropic (Claude models — API / OAuth)", value="anthropic"),
Choice(title="OpenAI (GPT models — API / OAuth)", value="openai"),
Choice(title="Google GenAI (Gemini models)", value="google-genai"),
Choice(
title="MiniMax (M2 — M3 models, up to 1M context, thinking)",
value="minimax",
),
Choice(title="ZhipuAI (智谱 — GLM models)", value="zhipu"),
Choice(
title="ZhipuAI CodePlan (智谱代码计划 — GLM models for coding)",
value="zhipu-code",
),
Choice(
title="Volcengine (火山引擎 — Doubao models)",
value="volcengine",
),
Choice(
title="DashScope (阿里云 — Qwen models)",
value="dashscope",
),
Choice(
title="DashScope Coding Plan (阿里云代码计划 — Qwen models)",
value="dashscope-code",
),
Choice(
title="DeepSeek (DeepSeek-R1, DeepSeek-V3)",
value="deepseek",
),
Choice(
title="Moonshot (月之暗面 — Moonshot models)",
value="moonshot",
),
Choice(
title="Kimi Coding Plan (Kimi 代码计划 — coding-focused)",
value="kimi-coding",
),
# Local
Choice(title="Ollama (local models)", value="ollama"),
# Third-party / aggregator
Choice(title="NVIDIA (third party — limited free requests)", value="nvidia"),
Choice(
title="SiliconFlow (aggregator — GLM, Kimi, MiniMax, etc.)",
value="siliconflow",
),
Choice(
title="OpenRouter (aggregator — Grok, Gemini, Qwen, etc.)",
value="openrouter",
),
Choice(
title="OpenAI-compatible (third-party OpenAI endpoint)",
value="custom-openai",
),
Choice(
title="Claude-compatible (third-party Anthropic endpoint)",
value="custom-anthropic",
),
]
# Set default based on current config (or an explicit override).
valid_providers = {c.value for c in choices}
preferred = default_value or config.provider
default = preferred if preferred in valid_providers else "anthropic"
provider = questionary.select(
f"Select {label} provider:" if label else "Select your LLM provider:",
choices=choices,
default=default,
style=WIZARD_STYLE,
qmark=QMARK,
use_indicator=True,
).ask()
if provider is None:
raise KeyboardInterrupt()
return provider
_MINIMAX_REGIONS: dict[str, str] = {
"global": "https://api.minimax.io/anthropic",
"cn": "https://api.minimaxi.com/anthropic",
}
def _step_minimax_region(config: EvoScientistConfig) -> str:
"""Step 2a (MiniMax): Select API region.
MiniMax has two regional endpoints — Global (api.minimax.io) and
Mainland China (api.minimaxi.com). API keys are region-bound.
Returns:
The selected base URL.
"""
current = config.minimax_base_url or os.environ.get("MINIMAX_BASE_URL", "")
if current == _MINIMAX_REGIONS["global"]:
default = "global"
else:
default = "cn"
region = questionary.select(
"Select MiniMax API region (must match where your key was created):",
choices=[
Choice(
title="Global (api.minimax.io — platform.minimax.io keys)",
value="global",
),
Choice(
title="Mainland China (api.minimaxi.com — platform.minimaxi.com keys)",
value="cn",
),
],
default=default,
style=WIZARD_STYLE,
qmark=QMARK,
use_indicator=True,
).ask()
if region is None:
raise KeyboardInterrupt()
return _MINIMAX_REGIONS[region]
def _step_oauth_auth_mode(
config: EvoScientistConfig,
*,
provider_label: str,
ccproxy_provider: str,
config_attr: str,
prompt_login_label: str,
oauth_choice_label: str | None = None,
status_label: str | None = None,
question_label: str | None = None,
) -> str:
"""Select API-key vs ccproxy OAuth authentication for a provider.
Args:
config: Current configuration.
provider_label: Provider display name for direct API-key access.
ccproxy_provider: ccproxy auth provider name.
config_attr: Config attribute storing this provider's auth mode.
prompt_login_label: Label used in "Log in to ..." prompts.
oauth_choice_label: Optional display label for the OAuth choice.
status_label: Optional display label for status messages.
question_label: Optional prompt label override.
Returns:
Selected auth mode: "api_key" or "oauth".
"""
from ...ccproxy_manager import check_ccproxy_auth, is_ccproxy_available
ccproxy_available = is_ccproxy_available()
from .prompter import BACK_SENTINEL, GoBack, install_navigation_keys
oauth_label = oauth_choice_label or f"{prompt_login_label} OAuth"
auth_status_label = status_label or oauth_label
auth_question_label = question_label or f"{provider_label} authentication mode"
choices = [
Choice(title=f"API Key (direct {provider_label} access)", value="api_key"),
Choice(
title=f"{oauth_label} (via ccproxy — no API key needed)"
+ (
""
if ccproxy_available
else " [requires: pip install evoscientist[oauth]]"
),
value="oauth",
),
questionary.Separator(),
Choice(title="← Back (re-pick provider)", value=BACK_SENTINEL),
]
current = getattr(config, config_attr)
if current not in ("api_key", "oauth"):
current = "api_key"
question = questionary.select(
f"{auth_question_label} [Esc/← to go back]:",
choices=choices,
default=current,
style=WIZARD_STYLE,
qmark=QMARK,
use_indicator=True,
)
install_navigation_keys(question, with_back=True)
auth_mode = question.ask()
if auth_mode is None:
raise KeyboardInterrupt()
if auth_mode == BACK_SENTINEL:
raise GoBack()
if auth_mode == "oauth" and not ccproxy_available:
console.print(" [yellow]✗ ccproxy not installed[/yellow]")
console.print()
install = questionary.confirm(
'Install ccproxy now? (pip install "evoscientist[oauth]")',
default=True,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if install is None:
raise KeyboardInterrupt()
if install:
console.print()
if _install_ccproxy():
console.print(" [green]✓ ccproxy installed successfully.[/green]")
else:
console.print(" [yellow]Falling back to API key mode.[/yellow]")
return "api_key"
else:
console.print(
' [dim]Skipped. Install manually: pip install "evoscientist[oauth]"[/dim]'
)
return "api_key"
if auth_mode == "oauth":
_prompt_ccproxy_port(config)
authed, msg = check_ccproxy_auth(ccproxy_provider)
if authed:
console.print(f" [green]✓ {auth_status_label}: {msg}[/green]")
relogin = questionary.confirm(
"Re-authenticate to refresh credentials?",
default=False,
style=CONFIRM_STYLE,
qmark=QMARK,
).ask()
if relogin is None:
raise KeyboardInterrupt()
if relogin:
_run_ccproxy_login(ccproxy_provider, auth_status_label)
else:
console.print(
f" [yellow]{auth_status_label} not authenticated: {msg}[/yellow]"
)
login = questionary.confirm(
f"Log in to {prompt_login_label} now?",
default=True,
style=CONFIRM_STYLE,
qmark=QMARK,
).ask()
if login is None:
raise KeyboardInterrupt()
if login:
_run_ccproxy_login(ccproxy_provider, auth_status_label)
return auth_mode
def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
"""Step 2a: Select Anthropic authentication mode (API key vs OAuth).
Args:
config: Current configuration.
Returns:
Selected auth mode: "api_key" or "oauth".
"""
return _step_oauth_auth_mode(
config,
provider_label="Anthropic",
ccproxy_provider="claude_api",
config_attr="anthropic_auth_mode",
prompt_login_label="Claude",
oauth_choice_label="Claude Code OAuth",
status_label="OAuth",
question_label="Authentication mode",
)
def _step_openai_auth_mode(config: EvoScientistConfig) -> str:
"""Step 2b: Select OpenAI authentication mode (API key vs Codex OAuth).
Args:
config: Current configuration.
Returns:
Selected auth mode: "api_key" or "oauth".
"""
return _step_oauth_auth_mode(
config,
provider_label="OpenAI",
ccproxy_provider="codex",
config_attr="openai_auth_mode",
prompt_login_label="Codex",
oauth_choice_label="Codex OAuth",
status_label="Codex OAuth",
question_label="OpenAI authentication mode",
)
def _step_provider_api_key(
config: EvoScientistConfig,
provider: str,
skip_validation: bool = False,
) -> str | None:
"""Step 2: Enter API key for the selected provider.
Args:
config: Current configuration.
provider: Selected provider name.
skip_validation: Skip API key validation.
Returns:
New API key or None if unchanged.
"""
key_name, current, validate_fn = _provider_key_info(config, provider)
hint = f"Current: ***{current[-4:]}" if current else "Not set"
prompt_text = f"Enter {key_name} API key ({hint}, Enter to keep):"
return _prompt_and_validate_api_key(
prompt_text,
current,
validate_fn,
skip_validation,
)
def _step_base_url(config: EvoScientistConfig, current_value: str | None = None) -> str:
"""Prompt for custom provider base URL.
Args:
config: Current configuration.
current_value: Current base URL value (if None, defaults to empty).
Returns:
Base URL string.
"""
current = current_value if current_value is not None else ""
hint = f"Current: {current}" if current else ""
default = current or ""
url = questionary.text(
f"Base URL{' (' + hint + ', Enter to keep)' if hint else ''}:",
default=default,
style=WIZARD_STYLE,
qmark=QMARK,
placeholder=FormattedText([("fg:#858585", " e.g. https://api.example.com/v1")])
if not default
else None,
).ask()
if url is None:
raise KeyboardInterrupt()
return url.strip()
def _step_ollama_base_url(config: EvoScientistConfig) -> tuple[str, list[str]]:
"""Prompt for Ollama server base URL and validate connection.
Args:
config: Current configuration.
Returns:
Tuple of (base_url, detected_model_names).
"""
current = config.ollama_base_url or os.environ.get("OLLAMA_BASE_URL", "")
default = current or "http://localhost:11434"
url = questionary.text(
f"Ollama base URL (Enter for {default}):",
default=default,
style=WIZARD_STYLE,
qmark=QMARK,
).ask()
if url is None:
raise KeyboardInterrupt()
url = url.strip()
detected_models: list[str] = []
if url:
console.print(" [dim]Checking Ollama connection...[/dim]", end="")
valid, msg, detected_models = validate_ollama_connection(url)
if valid:
console.print(f"\r [green]\u2713 {msg}[/green] ")
else:
console.print(f"\r [yellow]\u2717 {msg}[/yellow] ")
console.print(" [dim]You can start Ollama later and it will work.[/dim]")
return url, detected_models
def _step_model(
config: EvoScientistConfig,
provider: str,
*,
ollama_detected_models: list[str] | None = None,
label: str | None = None,
default_value: str | None = None,
) -> str:
"""Step 3: Select model for the provider.
Args:
config: Current configuration.
provider: Selected provider name.
ollama_detected_models: Model names detected from a live Ollama server.
label: Optional role label (e.g. "co-pilot") for the prompt. When omitted,
the generic main-model prompt is used.
default_value: Preselect this model instead of ``config.model`` (e.g. the
auxiliary model when configuring the co-pilot).
Returns:
Selected model name.
"""
model_prompt = f"Select {label} model:" if label else "Select model:"
model_default = default_value or config.model
# Ollama: show only what's actually pulled on the server
if provider == "ollama":
if ollama_detected_models:
_CUSTOM_SENTINEL = "__custom__"
choices = [
Choice(title=name, value=name) for name in ollama_detected_models
]
choices.append(Choice(title="Type a model name...", value=_CUSTOM_SENTINEL))
default = ollama_detected_models[0]
if model_default in ollama_detected_models:
default = model_default
selected = questionary.select(
model_prompt,
choices=choices,
default=default,
style=WIZARD_STYLE,
qmark=QMARK,
use_indicator=True,
).ask()
if selected is None:
raise KeyboardInterrupt()
if selected != _CUSTOM_SENTINEL:
return selected
# No detected models (server down or empty) — direct text input
if not ollama_detected_models:
console.print(
" [dim]No models detected — type the model name you plan to pull.[/dim]"
)
model = questionary.text(
"Model name:",
style=WIZARD_STYLE,
qmark=QMARK,
placeholder=FormattedText([("fg:#858585", " e.g. qwen3-coder-next")]),
).ask()
if model is None:
raise KeyboardInterrupt()
model = model.strip()
if not model:
model = "qwen3-coder-next"
console.print(f" [dim]Using default: {model}[/dim]")
return model
# Get models for the selected provider
entries = get_models_for_provider(provider)
if not entries:
# Custom / unknown provider: direct text input.
# Keep prompting until a non-empty model name is provided — saving an
# empty string here leaves the first request broken with an opaque
# "model required" error from the provider SDK.
while True:
model = questionary.text(
"Model name:",
style=WIZARD_STYLE,
qmark=QMARK,
placeholder=FormattedText([("fg:#858585", " e.g. owner/model-name")]),
default=model_default or "",
).ask()
if model is None:
raise KeyboardInterrupt()
model = model.strip()
if model:
return model
console.print(
" [yellow]Model name cannot be empty for a custom provider. "
"Press Ctrl+C to cancel.[/yellow]"
)
provider_models = [name for name, _ in entries]
# Create choices with model IDs as hints
_CUSTOM_SENTINEL = "__custom__"
choices = []
for name, model_id in entries:
choices.append(Choice(title=f"{name} ({model_id})", value=name))
choices.append(Choice(title="Type a model name...", value=_CUSTOM_SENTINEL))
# Determine default. An explicit ``default_value`` override (e.g. a saved
# co-pilot model on a re-run) that isn't a registry model is a custom name:
# preselect "Type a model name..." and prefill it. A plain ``config.model``
# that just isn't in the current provider's list (e.g. the provider was
# changed) falls back to the first model, NOT the custom entry.
custom_default = (
default_value if default_value and default_value not in provider_models else ""
)
if model_default in provider_models:
default = model_default
elif custom_default:
default = _CUSTOM_SENTINEL
else:
default = provider_models[0]
selected = questionary.select(
model_prompt,
choices=choices,
default=default,
style=WIZARD_STYLE,
qmark=QMARK,
use_indicator=True,
).ask()
if selected is None:
raise KeyboardInterrupt()
if selected != _CUSTOM_SENTINEL:
return selected
model = questionary.text(
"Model name:",
default=custom_default,
style=WIZARD_STYLE,
qmark=QMARK,
placeholder=FormattedText([("fg:#858585", " e.g. owner/model-name")]),
).ask()
if model is None:
raise KeyboardInterrupt()
model = model.strip()
if not model:
model = provider_models[0]
console.print(f" [dim]Using default: {model}[/dim]")
return model
def _step_auxiliary_enable(config: EvoScientistConfig) -> bool:
"""Step 3.25: Choose whether to assemble a co-pilot (auxiliary) model.
The co-pilot runs background/helper LLM calls — EvoMemory (memory workers)
and the main agent's tool selector — so it can be a cheaper/faster model.
Returns True when the user picks "Assemble"; the caller then runs the
provider/key/model pickers. Returns False to keep the pilot (main model)
everywhere.
"""
console.print(
" [dim]A cheaper/faster co-pilot for EvoMemory (memory workers).[/dim]"
)
choice = questionary.select(
"Co-pilot (auxiliary model):",
choices=[
Choice(
title="Skip — single pilot (main model handles everything)",
value="skip",
),
Choice(
title="Assemble a co-pilot — separate cheaper/faster model",
value="assemble",
),
],
default="assemble" if config.auxiliary_model else "skip",
style=WIZARD_STYLE,
qmark=QMARK,
use_indicator=True,
).ask()
if choice is None:
raise KeyboardInterrupt()
return choice == "assemble"
def _step_reasoning_effort(config: EvoScientistConfig) -> str:
"""Step 3.5: Configure OpenRouter reasoning effort level.
Only shown when the selected provider is OpenRouter. See:
https://openrouter.ai/docs/guides/best-practices/reasoning-tokens
Args:
config: Current configuration.
Returns:
Selected reasoning effort level, or empty string to use default.
"""
effort_choices = [
Choice(title="xhigh — ~95% of max_tokens for reasoning", value="xhigh"),
Choice(title="high — ~80% of max_tokens (recommended)", value="high"),
Choice(title="medium — ~50% of max_tokens", value="medium"),
Choice(title="low — ~20% of max_tokens", value="low"),
Choice(title="minimal — ~10% of max_tokens", value="minimal"),
Choice(title="none — disable reasoning entirely", value="none"),
]
current = config.reasoning_effort or "high"
effort = questionary.select(
"Select reasoning effort level:",
choices=effort_choices,
default=current,
style=WIZARD_STYLE,
qmark=QMARK,
use_indicator=True,
).ask()
if effort is None:
raise KeyboardInterrupt()
return effort
def _step_tavily_key(
config: EvoScientistConfig,
skip_validation: bool = False,
+1 -415
View File
@@ -1,7 +1,7 @@
"""Input validators for the onboarding wizard.
- IntegerValidator / ChoiceValidator: prompt_toolkit Validators
- validate_*_key: per-provider API key validators (live HTTP probes)
- validate_tavily_key: Tavily API key validator (live HTTP probe)
"""
from __future__ import annotations
@@ -97,415 +97,6 @@ def _classify_validation_error(error: BaseException) -> tuple[bool, str] | None:
return None
def validate_anthropic_key(api_key: str) -> tuple[bool, str]:
"""Validate an Anthropic API key by making a test request.
Args:
api_key: The API key to validate.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import anthropic
client = anthropic.Anthropic(api_key=api_key)
# Make a minimal request to validate the key
client.models.list()
return True, "Valid"
except anthropic.AuthenticationError:
return False, "Invalid API key"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_openai_key(api_key: str) -> tuple[bool, str]:
"""Validate an OpenAI API key by making a test request.
Args:
api_key: The API key to validate.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import openai
client = openai.OpenAI(api_key=api_key)
# Make a minimal request to validate the key
client.models.list()
return True, "Valid"
except openai.AuthenticationError:
return False, "Invalid API key"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_nvidia_key(api_key: str) -> tuple[bool, str]:
"""Validate an NVIDIA API key by making a test request.
Args:
api_key: The API key to validate.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
# NOTE: ``ChatNVIDIA(api_key=...)`` does NOT send a network request — it
# only stores the key in a client object. We must actually invoke the
# API (e.g. ``get_available_models()``) to verify the key is good.
try:
from langchain_nvidia_ai_endpoints import ChatNVIDIA
client = ChatNVIDIA(api_key=api_key, model="meta/llama-3.1-8b-instruct")
# Force a real authenticated request via model discovery.
client.get_available_models()
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_google_key(api_key: str) -> tuple[bool, str]:
"""Validate a Google GenAI API key by making a test request.
Args:
api_key: The API key to validate.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
from google import genai
client = genai.Client(api_key=api_key)
# Make a minimal request to validate the key
pager = client.models.list(config={"page_size": 1})
next(iter(pager)) # fetch first model only
return True, "Valid"
except StopIteration:
# Empty result but request succeeded — key is valid
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
# Google-specific 400 phrasing not in the shared hint list.
error_str = str(e).lower()
if "api_key_invalid" in error_str or "api key invalid" in error_str:
return False, "Invalid API key"
return False, f"Error: {e}"
def validate_minimax_key(
api_key: str,
base_url: str = "https://api.minimaxi.com/anthropic",
) -> tuple[bool, str]:
"""Validate a MiniMax API key without consuming tokens.
Sends a messages.create() with an empty model string. MiniMax checks
auth *before* validating request params, so a valid key returns 400
(bad model) while an invalid key returns 401.
Args:
api_key: The MiniMax API key to validate.
base_url: Anthropic-compatible endpoint (global or mainland China).
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import anthropic
client = anthropic.Anthropic(
api_key=api_key,
base_url=base_url,
)
client.messages.create(
model="",
max_tokens=1,
messages=[{"role": "user", "content": "hi"}],
)
# Unexpected success — treat as valid
return True, "Valid"
except anthropic.AuthenticationError:
return False, "Invalid API key"
except anthropic.APIStatusError:
# Any non-auth HTTP error (400 bad model, 500 insufficient balance,
# etc.) means the key itself was accepted → treat as valid.
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_siliconflow_key(api_key: str) -> tuple[bool, str]:
"""Validate a SiliconFlow API key by making a test request.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import openai
client = openai.OpenAI(
api_key=api_key, base_url="https://api.siliconflow.cn/v1"
)
client.models.list()
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_openrouter_key(api_key: str) -> tuple[bool, str]:
"""Validate an OpenRouter API key via the authenticated /auth/key endpoint.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import httpx
resp = httpx.get(
"https://openrouter.ai/api/v1/auth/key",
headers={"Authorization": f"Bearer {api_key}"},
timeout=10,
)
if resp.status_code == 200:
return True, "Valid"
# Only 401/403 mean the key is actually rejected. 429 (rate-limit)
# and 5xx (OpenRouter incident) leave the key validity unknown —
# surface the real status so the user doesn't go re-roll a good key
# during an outage.
if resp.status_code in (401, 403):
return False, "Invalid API key"
return False, f"Validation inconclusive (HTTP {resp.status_code})"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_deepseek_key(api_key: str) -> tuple[bool, str]:
"""Validate a DeepSeek API key by making a test request.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import openai
client = openai.OpenAI(api_key=api_key, base_url="https://api.deepseek.com")
client.models.list()
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_zhipu_key(api_key: str) -> tuple[bool, str]:
"""Validate a ZhipuAI API key by making a test request.
Uses the general endpoint for validation — both zhipu and zhipu-code
share the same API key, only the base_url differs at runtime.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import openai
client = openai.OpenAI(
api_key=api_key, base_url="https://open.bigmodel.cn/api/paas/v4"
)
client.models.list()
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_volcengine_key(api_key: str) -> tuple[bool, str]:
"""Validate a Volcengine API key by making a test request.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import openai
client = openai.OpenAI(
api_key=api_key,
base_url="https://ark.cn-beijing.volces.com/api/v3",
)
client.models.list()
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_dashscope_key(api_key: str) -> tuple[bool, str]:
"""Validate a DashScope API key by making a test request.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import openai
client = openai.OpenAI(
api_key=api_key,
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1",
)
client.models.list()
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_dashscope_code_key(api_key: str) -> tuple[bool, str]:
"""Validate a DashScope Coding Plan API key (sk-sp-* subscription keys).
The coding endpoint at coding.dashscope.aliyuncs.com does not expose
/models (returns 404), so validation issues a minimal chat completion
instead of the usual models.list() probe.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import openai
client = openai.OpenAI(
api_key=api_key,
base_url="https://coding.dashscope.aliyuncs.com/v1",
)
client.chat.completions.create(
model="qwen3-coder-plus",
messages=[{"role": "user", "content": "hi"}],
max_tokens=1,
)
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_moonshot_key(api_key: str) -> tuple[bool, str]:
"""Validate a Moonshot API key by making a test request.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import openai
client = openai.OpenAI(
api_key=api_key,
base_url="https://api.moonshot.cn/v1",
)
client.models.list()
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_kimi_key(api_key: str) -> tuple[bool, str]:
"""Validate a Kimi Coding Plan API key by making a test request.
Uses the Anthropic-compatible endpoint at api.kimi.com/coding/.
Returns:
Tuple of (is_valid, message).
"""
if not api_key:
return True, "Skipped (no key provided)"
try:
import anthropic
client = anthropic.Anthropic(
api_key=api_key,
base_url="https://api.kimi.com/coding/",
default_headers={"User-Agent": "claude-code/0.1.0"},
)
client.models.list()
return True, "Valid"
except Exception as e:
classified = _classify_validation_error(e)
if classified is not None:
return classified
return False, f"Error: {e}"
def validate_tavily_key(api_key: str) -> tuple[bool, str]:
"""Validate a Tavily API key by making a test request.
@@ -530,8 +121,3 @@ def validate_tavily_key(api_key: str) -> tuple[bool, str]:
if classified is not None:
return classified
return False, f"Error: {e}"
# =============================================================================
# Display Helpers
# =============================================================================
+4 -401
View File
@@ -3,7 +3,6 @@
from __future__ import annotations
import copy
import os
import questionary
from rich.panel import Panel
@@ -17,18 +16,8 @@ from ..settings import (
)
from .channels import _step_channels
from .steps import (
_step_anthropic_auth_mode,
_step_auxiliary_enable,
_step_base_url,
_step_langgraph_dev_port,
_step_mcp_servers,
_step_minimax_region,
_step_model,
_step_ollama_base_url,
_step_openai_auth_mode,
_step_provider,
_step_provider_api_key,
_step_reasoning_effort,
_step_skills,
_step_tavily_key,
_step_thinking,
@@ -41,7 +30,6 @@ from .style import (
CONFIRM_STYLE,
QMARK,
_print_header,
_print_section,
_print_step_skipped,
console,
)
@@ -49,10 +37,6 @@ from .style import (
STEPS = [
"UI",
"LangGraph Port",
"Provider",
"API Key",
"Model",
"Auxiliary Model",
"Tavily Key",
"Workspace",
"Thinking",
@@ -110,32 +94,6 @@ def render_progress(current_step: int, completed: set[int]) -> Panel:
# =============================================================================
_PROVIDER_KEY_ATTR = {
"anthropic": "anthropic_api_key",
"minimax": "minimax_api_key",
"nvidia": "nvidia_api_key",
"google-genai": "google_api_key",
"siliconflow": "siliconflow_api_key",
"openrouter": "openrouter_api_key",
"deepseek": "deepseek_api_key",
"zhipu": "zhipu_api_key",
"zhipu-code": "zhipu_api_key",
"volcengine": "volcengine_api_key",
"dashscope": "dashscope_api_key",
"dashscope-code": "dashscope_api_key",
"moonshot": "moonshot_api_key",
"kimi-coding": "kimi_api_key",
"custom-openai": "custom_openai_api_key",
"custom-anthropic": "custom_anthropic_api_key",
}
_MINIMAX_GLOBAL_BASE_URL = "https://api.minimax.io/anthropic"
_CUSTOM_PROVIDER_BASE_URL = {
"custom-openai": ("custom_openai_base_url", "CUSTOM_OPENAI_BASE_URL"),
"custom-anthropic": ("custom_anthropic_base_url", "CUSTOM_ANTHROPIC_BASE_URL"),
}
def _autosave(config: EvoScientistConfig) -> None:
"""Persist current config to disk between phases.
@@ -148,208 +106,10 @@ def _autosave(config: EvoScientistConfig) -> None:
pass
def _configure_provider_base_url(
config: EvoScientistConfig,
provider: str,
*,
strict: bool,
) -> list[str]:
"""Configure provider-specific base URL/region and return Ollama models."""
if provider in _CUSTOM_PROVIDER_BASE_URL:
attr_name, env_name = _CUSTOM_PROVIDER_BASE_URL[provider]
current_base_url = getattr(config, attr_name) or os.environ.get(env_name, "")
if strict:
if not current_base_url:
raise RuntimeError(
f"--non-interactive: {provider} provider needs a base URL. "
f"Set the {env_name} env var or run without --non-interactive."
)
setattr(config, attr_name, current_base_url)
else:
setattr(
config,
attr_name,
_step_base_url(config, current_value=current_base_url),
)
elif provider == "minimax":
if strict:
config.minimax_base_url = (
config.minimax_base_url or _MINIMAX_GLOBAL_BASE_URL
)
else:
config.minimax_base_url = _step_minimax_region(config)
elif provider == "ollama":
if strict:
config.ollama_base_url = (
config.ollama_base_url
or os.environ.get("OLLAMA_BASE_URL", "")
or "http://localhost:11434"
)
else:
ollama_url, ollama_detected_models = _step_ollama_base_url(config)
config.ollama_base_url = ollama_url
return ollama_detected_models
return []
def _configure_provider_auth_mode(
config: EvoScientistConfig,
provider: str,
*,
strict: bool,
) -> None:
"""Configure Anthropic/OpenAI auth mode for the selected provider."""
if provider == "anthropic":
if strict:
config.anthropic_auth_mode = "api_key"
else:
config.anthropic_auth_mode = _step_anthropic_auth_mode(config)
elif provider == "openai":
if strict:
config.openai_auth_mode = "api_key"
else:
config.openai_auth_mode = _step_openai_auth_mode(config)
def _active_llm_providers(config: EvoScientistConfig) -> set[str]:
"""Return providers currently selected by the main and auxiliary models."""
providers = {config.provider}
if config.auxiliary_provider:
providers.add(config.auxiliary_provider)
return providers
def _reconcile_oauth_modes(config: EvoScientistConfig) -> None:
"""Clear OAuth flags for providers no selected model uses."""
active_providers = _active_llm_providers(config)
if "anthropic" not in active_providers:
config.anthropic_auth_mode = "api_key"
if "openai" not in active_providers:
config.openai_auth_mode = "api_key"
def _provider_uses_oauth(config: EvoScientistConfig, provider: str) -> bool:
return (provider == "anthropic" and config.anthropic_auth_mode == "oauth") or (
provider == "openai" and config.openai_auth_mode == "oauth"
)
def _apply_preset_provider_api_key(
config: EvoScientistConfig,
provider: str,
preset_api_key: str,
*,
skip_validation: bool,
) -> None:
"""Validate and store a CLI-supplied provider API key."""
if not skip_validation:
from .helpers import _provider_key_info
_info = _provider_key_info(config, provider)
validate_fn = _info[2] if _info else None
if validate_fn is not None:
console.print(" [dim]Validating preset API key...[/dim]", end="")
valid, msg = validate_fn(preset_api_key)
if valid:
console.print(f"\r [green]✓ {msg}[/green] ")
else:
console.print(f"\r [red]✗ {msg}[/red] ")
raise RuntimeError(
f"--api-key rejected by {provider} validator: {msg}. "
"Pass --skip-validation to override."
)
key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key")
setattr(config, key_attr, preset_api_key)
console.print(
f" [green]✓ API key: ***{preset_api_key[-4:]}[/green] [dim](--api-key)[/dim]"
)
def _configure_provider_api_key(
config: EvoScientistConfig,
provider: str,
*,
skip_validation: bool,
preset_api_key: str | None = None,
require_api_key=None,
) -> None:
"""Configure provider API key unless the provider does not need one."""
if provider == "ollama" or _provider_uses_oauth(config, provider):
return
key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key")
if preset_api_key is not None:
_apply_preset_provider_api_key(
config,
provider,
preset_api_key,
skip_validation=skip_validation,
)
return
if require_api_key is not None:
require_api_key()
new_key = _step_provider_api_key(config, provider, skip_validation)
if new_key is not None:
setattr(config, key_attr, new_key)
elif not getattr(config, key_attr):
_print_step_skipped("API Key", "not set")
def _provider_connection_configured(config: EvoScientistConfig, provider: str) -> bool:
"""Return True when provider-level setup can be safely reused."""
if provider == "ollama":
return bool(config.ollama_base_url)
if provider == "custom-openai" and not config.custom_openai_base_url:
return False
if provider == "custom-anthropic" and not config.custom_anthropic_base_url:
return False
if provider == "minimax" and not config.minimax_base_url:
return False
if _provider_uses_oauth(config, provider):
return True
key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key")
return bool(getattr(config, key_attr))
def _configure_provider_connection(
config: EvoScientistConfig,
provider: str,
*,
strict: bool,
skip_validation: bool,
preset_api_key: str | None = None,
require_api_key=None,
) -> list[str]:
"""Configure provider base URL/region, auth mode, and API key."""
ollama_detected_models = _configure_provider_base_url(
config,
provider,
strict=strict,
)
_configure_provider_auth_mode(
config,
provider,
strict=strict,
)
_configure_provider_api_key(
config,
provider,
skip_validation=skip_validation,
preset_api_key=preset_api_key,
require_api_key=require_api_key,
)
return ollama_detected_models
# Sections offered in Keep/Modify/Reset → which step labels they enable.
_SECTION_LABELS: list[tuple[str, str]] = [
("ui", "UI backend"),
("port", "LangGraph server port"),
("provider", "LLM provider + auth + API key"),
("model", "Model + reasoning effort"),
("auxiliary_model", "Auxiliary model (optional)"),
("tavily", "Tavily search key"),
("workspace", "Workspace mode"),
("thinking", "Thinking panel"),
@@ -360,17 +120,10 @@ _SECTION_LABELS: list[tuple[str, str]] = [
]
_ALL_SECTIONS: frozenset[str] = frozenset(s for s, _ in _SECTION_LABELS)
# Each preset flag implies the section(s) it would change. ``--provider`` also
# cascades into ``model`` because the model list depends on the provider —
# silently keeping a stale model id would leave the first request broken.
# Each preset flag implies the section(s) it would change.
_FLAG_TO_SECTIONS: dict[str, frozenset[str]] = {
"ui": frozenset({"ui"}),
"port": frozenset({"port"}),
"provider": frozenset({"provider", "model"}),
# ``--api-key`` re-runs the provider section, which can change provider —
# cascade to model for the same reason ``--provider`` does.
"api_key": frozenset({"provider", "model"}),
"model": frozenset({"model"}),
"tavily_key": frozenset({"tavily"}),
"workspace_mode": frozenset({"workspace"}),
"show_thinking": frozenset({"thinking"}),
@@ -481,7 +234,7 @@ def run_onboard(
Args:
skip_validation: Skip API key validation.
prompter: Optional :class:`NonInteractivePrompter` carrying
CLI-supplied answers (``--provider``, ``--model``, …) and
CLI-supplied answers (``--ui``, ``--tavily-key``, …) and
``skip_set`` (sections to bypass). When None, all prompts
fall through to the interactive questionary form.
only_sections: If given, restrict the wizard to exactly these section
@@ -556,9 +309,9 @@ def run_onboard(
# Decide which sections this run should cover.
#
# - ``only_sections`` (programmatic, e.g. ``configure provider``):
# - ``only_sections`` (programmatic, e.g. ``configure mcp``):
# run exactly those sections, no Keep/Modify/Reset prompt.
# - Any preset flag (``--provider``/``--model``/…): treat as
# - Any preset flag (``--ui``/``--tavily-key``/…): treat as
# explicit user intent — skip Keep/Modify/Reset and run ONLY
# the sections each flag implies (see ``_FLAG_TO_SECTIONS``).
# - Strict ``--non-interactive`` with no preset flags: run all
@@ -654,156 +407,6 @@ def run_onboard(
config.langgraph_dev_port = _step_langgraph_dev_port(config)
_autosave(config)
ollama_detected_models: list[str] = []
if "provider" in sections_to_run:
from .prompter import GoBack
_print_section("EvoScientist · Pilot (Main model)")
_require("provider", "LLM provider")
# Provider sub-loop: auth_mode can raise GoBack to re-pick provider.
# We snapshot config at the top of each iteration so a GoBack can
# roll back partial writes (base_url, minimax region, ollama URL,
# provider id itself) — otherwise picking `custom-openai`, entering
# a base URL, going Back, then picking `anthropic` would leave a
# stale ``custom_openai_base_url`` in the final saved config.
while True:
loop_snapshot = copy.deepcopy(config)
preset_provider = _preset("provider")
if preset_provider is not None:
provider = preset_provider
config.provider = provider
console.print(
f" [green]✓ Provider: {provider}[/green] "
"[dim](--provider)[/dim]"
)
else:
provider = _step_provider(config)
config.provider = provider
try:
ollama_detected_models = _configure_provider_connection(
config,
provider,
strict=strict,
skip_validation=skip_validation,
preset_api_key=_preset("api_key"),
require_api_key=lambda provider=provider: _require(
"api_key", f"{provider} API key"
),
)
except GoBack:
# User picked "← Back" — restore config to its state at the
# top of this iteration (drops any base_url / region /
# provider writes), then discard ALL provider-coupled
# presets and re-prompt. Clearing only ``provider``
# leaves a stale ``--model`` / ``--api-key`` that would
# be re-applied under a different provider, producing
# an invalid pair (e.g. ``provider=openai`` +
# ``model=claude-sonnet-4-6``).
for field_name in vars(loop_snapshot):
setattr(
config, field_name, getattr(loop_snapshot, field_name)
)
if p:
for stale_key in ("provider", "model", "api_key"):
p.answers.pop(stale_key, None)
ollama_detected_models = []
console.print(" [dim]↩ Returning to provider selection.[/dim]")
continue
break # Provider setup succeeded — exit sub-loop
_reconcile_oauth_modes(config)
_autosave(config)
else:
# Provider section skipped — keep prior provider value to drive
# downstream sections that depend on it (e.g., model picker).
provider = config.provider
if "model" in sections_to_run:
_require("model", "Model")
preset_model = _preset("model")
if preset_model is not None:
config.model = preset_model
console.print(
f" [green]✓ Model: {preset_model}[/green] [dim](--model)[/dim]"
)
else:
config.model = _step_model(
config, provider, ollama_detected_models=ollama_detected_models
)
if provider == "openrouter" and _preset("model") is None:
config.reasoning_effort = _step_reasoning_effort(config)
_autosave(config)
if "auxiliary_model" in sections_to_run:
_print_section("Co-pilot (Auxiliary model)")
if strict:
# Optional; never prompt under --non-interactive. Keep
# current (default empty = use main model).
_print_step_skipped(
"Auxiliary Model",
"kept current" if config.auxiliary_model else "not set",
)
elif _step_auxiliary_enable(config):
from .prompter import GoBack
aux_ollama_detected_models: list[str] = []
while True:
loop_snapshot = copy.deepcopy(config)
aux_provider = _step_provider(
config,
label="co-pilot",
default_value=config.auxiliary_provider,
)
config.auxiliary_provider = aux_provider
if (
aux_provider == config.provider
and _provider_connection_configured(config, aux_provider)
):
if aux_provider == "ollama":
aux_ollama_detected_models = ollama_detected_models
_print_step_skipped(
"Co-pilot credentials",
"reusing main provider settings",
)
else:
try:
aux_ollama_detected_models = (
_configure_provider_connection(
config,
aux_provider,
strict=False,
skip_validation=skip_validation,
)
)
except GoBack:
for field_name in vars(loop_snapshot):
setattr(
config,
field_name,
getattr(loop_snapshot, field_name),
)
aux_ollama_detected_models = []
console.print(
" [dim]↩ Returning to co-pilot provider "
"selection.[/dim]"
)
continue
break
config.auxiliary_model = _step_model(
config,
aux_provider,
ollama_detected_models=aux_ollama_detected_models,
label="co-pilot",
default_value=config.auxiliary_model,
)
else:
# Skip: single driver — clear any prior auxiliary config.
config.auxiliary_provider = ""
config.auxiliary_model = ""
_reconcile_oauth_modes(config)
_autosave(config)
if "tavily" in sections_to_run:
preset_tavily = _preset("tavily_key")
if preset_tavily is not None:
+94 -145
View File
@@ -23,6 +23,19 @@ from dotenv import find_dotenv, load_dotenv
# (stream/display.py, channels/consumer.py) — keep aligned with the agent's
# `interrupt_on` set in EvoScientist.py.
HITL_SHELL_TOOLS = ("execute", "run_in_background")
_CONFIG_APPLIED_ENV_VALUES: dict[str, str] = {}
def _apply_config_env(name: str, value: str) -> None:
if value and not os.environ.get(name):
os.environ[name] = value
_CONFIG_APPLIED_ENV_VALUES[name] = value
def is_config_applied_env(name: str) -> bool:
"""Return whether the current env value was injected from config.yaml."""
applied = _CONFIG_APPLIED_ENV_VALUES.get(name)
return applied is not None and os.environ.get(name) == applied
class MemoryObservationTarget(StrEnum):
@@ -106,23 +119,11 @@ def _normalize_hhmm(value: Any) -> str | None:
def get_config_dir() -> Path:
"""Get the configuration directory path.
Priority:
1. EVOSCIENTIST_CONFIG_DIR
2. EVOSCIENTIST_HOME/config
3. XDG_CONFIG_HOME/evoscientist
4. ~/.config/evoscientist
Uses XDG_CONFIG_HOME if set, otherwise ~/.config/evoscientist/
"""
configured = os.environ.get("EVOSCIENTIST_CONFIG_DIR")
if configured:
return Path(configured).expanduser().resolve()
home = os.environ.get("EVOSCIENTIST_HOME")
if home:
return Path(home).expanduser().resolve() / "config"
xdg_config = os.environ.get("XDG_CONFIG_HOME")
if xdg_config:
return Path(xdg_config).expanduser() / "evoscientist"
return Path(xdg_config) / "evoscientist"
return Path.home() / ".config" / "evoscientist"
@@ -131,6 +132,11 @@ def get_config_path() -> Path:
return get_config_dir() / "config.yaml"
def get_default_workspace_dir() -> Path:
"""Return the stable workspace used when no directory is configured."""
return Path.home() / ".evoscientist" / "workspace"
# =============================================================================
# Configuration dataclass
# =============================================================================
@@ -140,55 +146,21 @@ def get_config_path() -> Path:
class EvoScientistConfig:
"""EvoScientist configuration settings.
LLM provider / model / API key configuration lives in the model
registry (model-runtime.sqlite3), not here — this dataclass only holds
platform settings (workspace, UI, channels, memory, HITL, …).
Attributes:
anthropic_api_key: Anthropic API key for Claude models.
openai_api_key: OpenAI API key for GPT models.
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.
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).
auxiliary_model: Model for memory workers + tool selector + scheduler (empty = use main model).
default_mode: Default workspace mode ('daemon' or 'run').
default_workdir: Default workspace directory (empty = use current working directory).
default_workdir: Default workspace directory (empty = use
~/.evoscientist/workspace).
show_thinking: Whether to show thinking panels in CLI.
"""
# API Keys
anthropic_api_key: str = ""
anthropic_base_url: str = ""
anthropic_auth_mode: str = "api_key" # "api_key" | "oauth"
openai_api_key: str = ""
openai_auth_mode: str = "api_key" # "api_key" | "oauth"
nvidia_api_key: str = ""
google_api_key: str = ""
minimax_api_key: str = ""
minimax_base_url: str = ""
siliconflow_api_key: str = ""
openrouter_api_key: str = ""
deepseek_api_key: str = ""
zhipu_api_key: str = ""
volcengine_api_key: str = ""
dashscope_api_key: str = ""
moonshot_api_key: str = ""
kimi_api_key: str = ""
custom_openai_api_key: str = ""
custom_openai_base_url: str = ""
custom_anthropic_api_key: str = ""
custom_anthropic_base_url: str = ""
ollama_base_url: str = ""
# API Keys (non-LLM tools only)
tavily_api_key: str = ""
# LLM Settings
provider: str = "anthropic"
model: str = "claude-sonnet-4-6"
model_fallbacks: str = "" # "model:provider,model:provider" fallback chain
# Optional auxiliary model for background/helper LLM calls (memory workers +
# tool selector). Empty = fall back to the main model/provider.
auxiliary_provider: str = "" # empty = use main provider
auxiliary_model: str = "" # empty = use main model
# Async Sub-agent Settings
# When True (default), the EvoSci CLI auto-starts a langgraph dev subprocess
# so any sub-agent flagged ``async: true`` in subagents/<name>.yaml runs
@@ -291,9 +263,6 @@ class EvoScientistConfig:
ui_backend: Literal["cli", "tui", "webui"] = "tui"
log_level: str = "warning"
reasoning_effort: str = "high"
# 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
# Channel Settings
channel_enabled: str = "" # "imessage" | "telegram" | "discord" | "slack" | "wechat" | "dingtalk" | "feishu" | "email" | "qq" | "signal" | "" (comma-separated for multiple)
@@ -410,6 +379,17 @@ class EvoScientistConfig:
# blocklist (sudo/chmod/dd/...) still applies. Implies auto_approve.
dangerous_mode: bool = False
# Conversation workspace isolation. New WebUI conversations use a registry-
# validated scope by default. ``optional`` preserves the legacy execution
# path only when a trusted caller cannot provide a scope.
workspace_isolation: str = "optional" # legacy | optional | required
scope_registry_topology: str = "single-host"
strict_executor: str = "oci" # oci
strict_executor_image: str = "" # immutable OCI digest required in required mode
strict_code_interpreter: str = "disabled" # disabled | scoped
draft_workspace_ttl_hours: int = 24
workspace_trash_retention_days: int = 7
# Agent features
enable_ask_user: bool = True # Enable ask_user tool for agent-initiated questions
@@ -438,9 +418,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
@@ -470,6 +447,51 @@ class EvoScientistConfig:
if self.dangerous_mode:
self.auto_approve = True
if self.workspace_isolation not in {"legacy", "optional", "required"}:
raise ValueError(
"workspace_isolation must be one of legacy, optional, required"
)
if self.scope_registry_topology != "single-host":
raise ValueError(
"v1 workspace isolation only supports single-host topology"
)
if self.strict_executor != "oci":
raise ValueError("strict_executor must be oci")
if self.strict_code_interpreter not in {"disabled", "scoped"}:
raise ValueError("strict_code_interpreter must be disabled or scoped")
if (
not isinstance(self.draft_workspace_ttl_hours, int)
or isinstance(self.draft_workspace_ttl_hours, bool)
or self.draft_workspace_ttl_hours <= 0
):
raise ValueError("draft_workspace_ttl_hours must be a positive integer")
if (
not isinstance(self.workspace_trash_retention_days, int)
or isinstance(self.workspace_trash_retention_days, bool)
or self.workspace_trash_retention_days <= 0
):
raise ValueError(
"workspace_trash_retention_days must be a positive integer"
)
if self.workspace_isolation == "required" and self.dangerous_mode:
raise ValueError(
"dangerous_mode is incompatible with required workspace isolation"
)
if (
self.workspace_isolation == "required"
and "@sha256:" not in self.strict_executor_image
):
raise ValueError(
"required workspace isolation needs strict_executor_image pinned by digest"
)
if (
self.workspace_isolation == "required"
and self.strict_code_interpreter != "disabled"
):
raise ValueError(
"strict_code_interpreter=scoped is not available until every PTC tool is scope-verified"
)
_normalize_str_enum_fields(self)
synthesis_time = _normalize_hhmm(self.memory_skill_synthesis_time)
@@ -739,44 +761,22 @@ def list_config() -> dict[str, Any]:
# Environment variable mappings
_ENV_MAPPINGS = {
"anthropic_api_key": "ANTHROPIC_API_KEY",
"anthropic_base_url": "ANTHROPIC_BASE_URL",
"anthropic_auth_mode": "EVOSCIENTIST_ANTHROPIC_AUTH_MODE",
"openai_api_key": "OPENAI_API_KEY",
"openai_auth_mode": "EVOSCIENTIST_OPENAI_AUTH_MODE",
"nvidia_api_key": "NVIDIA_API_KEY",
"google_api_key": "GOOGLE_API_KEY",
"minimax_api_key": "MINIMAX_API_KEY",
"minimax_base_url": "MINIMAX_BASE_URL",
"siliconflow_api_key": "SILICONFLOW_API_KEY",
"openrouter_api_key": "OPENROUTER_API_KEY",
"deepseek_api_key": "DEEPSEEK_API_KEY",
"zhipu_api_key": "ZHIPU_API_KEY",
"volcengine_api_key": "VOLCENGINE_API_KEY",
"dashscope_api_key": "DASHSCOPE_API_KEY",
"moonshot_api_key": "MOONSHOT_API_KEY",
"kimi_api_key": "KIMI_API_KEY",
"custom_openai_api_key": "CUSTOM_OPENAI_API_KEY",
"custom_openai_base_url": "CUSTOM_OPENAI_BASE_URL",
"custom_anthropic_api_key": "CUSTOM_ANTHROPIC_API_KEY",
"custom_anthropic_base_url": "CUSTOM_ANTHROPIC_BASE_URL",
"ollama_base_url": "OLLAMA_BASE_URL",
"tavily_api_key": "TAVILY_API_KEY",
"default_mode": "EVOSCIENTIST_DEFAULT_MODE",
"default_workdir": "EVOSCIENTIST_WORKSPACE_DIR",
"ui_backend": "EVOSCIENTIST_UI_BACKEND",
"log_level": "EVOSCIENTIST_LOG_LEVEL",
"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"
),
"dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE",
"workspace_isolation": "EVOSCIENTIST_WORKSPACE_ISOLATION",
"scope_registry_topology": "EVOSCIENTIST_SCOPE_REGISTRY_TOPOLOGY",
"strict_executor": "EVOSCIENTIST_STRICT_EXECUTOR",
"strict_executor_image": "EVOSCIENTIST_STRICT_EXECUTOR_IMAGE",
"strict_code_interpreter": "EVOSCIENTIST_STRICT_CODE_INTERPRETER",
"draft_workspace_ttl_hours": "EVOSCIENTIST_DRAFT_WORKSPACE_TTL_HOURS",
"workspace_trash_retention_days": "EVOSCIENTIST_WORKSPACE_TRASH_RETENTION_DAYS",
"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",
@@ -848,66 +848,19 @@ def get_effective_config(
def apply_config_to_env(config: EvoScientistConfig) -> None:
"""Apply config API keys to environment variables if not already set.
"""Apply config values to environment variables if not already set.
This allows the config file to provide API keys that downstream
libraries (like langchain-anthropic) can pick up.
LLM provider keys are no longer injected here — they live in the model
registry (model-runtime.sqlite3) and are applied by the model runtime.
Only non-LLM tool keys and platform round-trips remain.
Args:
config: Configuration to apply.
"""
if config.anthropic_api_key and not os.environ.get("ANTHROPIC_API_KEY"):
os.environ["ANTHROPIC_API_KEY"] = config.anthropic_api_key
if config.anthropic_base_url and not os.environ.get("ANTHROPIC_BASE_URL"):
os.environ["ANTHROPIC_BASE_URL"] = config.anthropic_base_url
if config.openai_api_key and not os.environ.get("OPENAI_API_KEY"):
os.environ["OPENAI_API_KEY"] = config.openai_api_key
if config.nvidia_api_key and not os.environ.get("NVIDIA_API_KEY"):
os.environ["NVIDIA_API_KEY"] = config.nvidia_api_key
if config.google_api_key and not os.environ.get("GOOGLE_API_KEY"):
os.environ["GOOGLE_API_KEY"] = config.google_api_key
if config.minimax_api_key and not os.environ.get("MINIMAX_API_KEY"):
os.environ["MINIMAX_API_KEY"] = config.minimax_api_key
if config.minimax_base_url and not os.environ.get("MINIMAX_BASE_URL"):
os.environ["MINIMAX_BASE_URL"] = config.minimax_base_url
if config.siliconflow_api_key and not os.environ.get("SILICONFLOW_API_KEY"):
os.environ["SILICONFLOW_API_KEY"] = config.siliconflow_api_key
if config.openrouter_api_key and not os.environ.get("OPENROUTER_API_KEY"):
os.environ["OPENROUTER_API_KEY"] = config.openrouter_api_key
if config.deepseek_api_key and not os.environ.get("DEEPSEEK_API_KEY"):
os.environ["DEEPSEEK_API_KEY"] = config.deepseek_api_key
if config.zhipu_api_key and not os.environ.get("ZHIPU_API_KEY"):
os.environ["ZHIPU_API_KEY"] = config.zhipu_api_key
if config.volcengine_api_key and not os.environ.get("VOLCENGINE_API_KEY"):
os.environ["VOLCENGINE_API_KEY"] = config.volcengine_api_key
if config.dashscope_api_key and not os.environ.get("DASHSCOPE_API_KEY"):
os.environ["DASHSCOPE_API_KEY"] = config.dashscope_api_key
if config.moonshot_api_key and not os.environ.get("MOONSHOT_API_KEY"):
os.environ["MOONSHOT_API_KEY"] = config.moonshot_api_key
if config.kimi_api_key and not os.environ.get("KIMI_API_KEY"):
os.environ["KIMI_API_KEY"] = config.kimi_api_key
if config.custom_openai_api_key and not os.environ.get("CUSTOM_OPENAI_API_KEY"):
os.environ["CUSTOM_OPENAI_API_KEY"] = config.custom_openai_api_key
if config.custom_openai_base_url and not os.environ.get("CUSTOM_OPENAI_BASE_URL"):
os.environ["CUSTOM_OPENAI_BASE_URL"] = config.custom_openai_base_url
if config.custom_anthropic_api_key and not os.environ.get(
"CUSTOM_ANTHROPIC_API_KEY"
):
os.environ["CUSTOM_ANTHROPIC_API_KEY"] = config.custom_anthropic_api_key
if config.custom_anthropic_base_url and not os.environ.get(
"CUSTOM_ANTHROPIC_BASE_URL"
):
os.environ["CUSTOM_ANTHROPIC_BASE_URL"] = config.custom_anthropic_base_url
if config.ollama_base_url and not os.environ.get("OLLAMA_BASE_URL"):
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 not config.openrouter_anthropic_prompt_cache and not os.environ.get(
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"
):
os.environ["EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"] = "false"
# Round-trip dangerous_mode to env so it survives a fresh get_effective_config()
# (warning banner, run_in_background) and is inherited by the langgraph dev
# subprocess — otherwise a --dangerous CLI flag (not persisted to file/env)
@@ -918,7 +871,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
+97 -11
View File
@@ -12,7 +12,7 @@ multiple clients at one hand-started server they will share the same cron store.
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from langgraph_sdk.schema import Cron, Run
@@ -47,22 +47,84 @@ def is_available() -> bool:
return bool(is_langgraph_dev_running(base_url=_scheduler_url()))
def _scope_payload(
scope: Any | None, *, owner_id: str | None = None
) -> tuple[dict[str, Any], dict[str, Any]]:
if scope is None:
return {}, {}
metadata = {
"workspace_scope_id": scope.scope_id,
"workspace_scope_owner_id": owner_id or scope.owner_id,
"workspace_deployment_id": scope.deployment_id,
}
return metadata, {
"configurable": {
**metadata,
"workspace_scope_revision": scope.revision,
}
}
def create_schedule(
*, name: str, schedule: str, prompt: str, timezone: str | None = None
*,
name: str,
schedule: str,
prompt: str,
timezone: str | None = None,
scope: Any | None = None,
) -> Cron:
"""Create a recurring scheduled task on the scheduler graph."""
# Crons are stored in the langgraph-dev process's .langgraph_api store, not
# tagged by workspace. Isolation is process-level (see module docstring).
return _client().crons.create(
assistant_id=SCHEDULER_GRAPH_ID,
schedule=schedule,
input=messages_input(prompt),
metadata={"run_kind": SCHEDULED_RUN_KIND, "name": name, "prompt": prompt},
timezone=timezone or _default_timezone(),
owner = None
if scope is not None:
from ..scope_registry import get_scope_registry
owner = get_scope_registry().register_owner(
scope.deployment_id,
scope.scope_id,
owner_type="schedule",
parent_owner_id=scope.owner_id,
state="reserved",
)
scope_metadata, config = _scope_payload(
scope, owner_id=owner.owner_id if owner else None
)
try:
cron = _client().crons.create(
assistant_id=SCHEDULER_GRAPH_ID,
schedule=schedule,
input=messages_input(prompt),
metadata={
"run_kind": SCHEDULED_RUN_KIND,
"name": name,
"prompt": prompt,
**scope_metadata,
},
**({"config": config} if config else {}),
timezone=timezone or _default_timezone(),
)
except Exception:
if owner is not None:
get_scope_registry().bind_owner(
scope.deployment_id,
scope.scope_id,
owner.owner_id,
f"failed:{owner.owner_id}",
state="terminal",
)
raise
if owner is not None:
get_scope_registry().bind_owner(
scope.deployment_id,
scope.scope_id,
owner.owner_id,
str(cron["cron_id"]),
)
return cron
def list_schedules() -> list[Cron]:
def list_schedules(scope: Any | None = None) -> list[Cron]:
"""Return only EvoScientist scheduled tasks.
Filtered server-side by ``run_kind`` metadata (the cron backend matches by
@@ -71,10 +133,17 @@ def list_schedules() -> list[Cron]:
rather than ``assistant_id`` because the stored ``assistant_id`` is a resolved
UUID, not the ``scheduler`` graph name we create with.
"""
return _client().crons.search(
rows = _client().crons.search(
metadata={"run_kind": SCHEDULED_RUN_KIND},
limit=1000,
)
if scope is None:
return rows
return [
row
for row in rows
if (row.get("metadata") or {}).get("workspace_scope_id") == scope.scope_id
]
def delete_schedule(cron_id: str) -> None:
@@ -87,13 +156,28 @@ def set_enabled(cron_id: str, enabled: bool) -> Cron:
return _client().crons.update(cron_id, enabled=enabled)
def run_now(prompt: str) -> Run:
def run_now(prompt: str, scope: Any | None = None) -> Run:
"""Fire a one-off scheduler run immediately (for ``/schedule run``).
Output goes wherever the task's prompt specifies; there is no push notification.
"""
client = _client()
thread = client.threads.create(graph_id=SCHEDULER_GRAPH_ID)
owner = None
if scope is not None:
from ..scope_registry import get_scope_registry
owner = get_scope_registry().register_owner(
scope.deployment_id,
scope.scope_id,
owner_type="scheduler_run",
resource_id=str(thread["thread_id"]),
parent_owner_id=scope.owner_id,
state="active",
)
scope_metadata, config = _scope_payload(
scope, owner_id=owner.owner_id if owner else None
)
return client.runs.create(
thread_id=str(thread["thread_id"]),
assistant_id=SCHEDULER_GRAPH_ID,
@@ -102,5 +186,7 @@ def run_now(prompt: str) -> Run:
"run_kind": SCHEDULED_RUN_KIND,
"name": "manual-run",
"prompt": prompt,
**scope_metadata,
},
**({"config": config} if config else {}),
)
+110 -9
View File
@@ -5,8 +5,8 @@ for consumption by external LangChain-compatible UIs (deep-agents-ui,
agent-chat-ui, LangSmith Studio) and SDK clients.
Differs from ``EvoSci`` / ``EvoSci serve``: no in-process CLI agent,
no session DB, no channel runtime, no TUI. The terminal only shows
startup progress, the Ready banner, and then blocks until Ctrl+C.
no session DB, no channel runtime, no TUI. The terminal shows startup
progress, the Ready banner, and the live Gateway log until Ctrl+C.
Mode dispatch happens via the ``EVOSCIENTIST_DEPLOY_MODE`` env var
injected by ``start_langgraph_dev``: ``full`` for the deploy subprocess
@@ -20,11 +20,13 @@ loads or skips MCP based on the value.
from __future__ import annotations
import atexit
import codecs
import os
import signal
import sys
import threading
from pathlib import Path
from typing import Any
from typing import Any, TextIO
import typer # type: ignore[import-untyped]
from rich.panel import Panel
@@ -32,6 +34,65 @@ from rich.text import Text
from ..cli._app import app
from ..stream.console import console
from ..usage import prepare_usage_environment
def _follow_gateway_log(
log_path: Path,
start_offset: int,
stop_event: threading.Event,
*,
output: TextIO | None = None,
poll_interval: float = 0.1,
) -> None:
"""Mirror appended Gateway log bytes to a text stream until stopped."""
sink = output if output is not None else sys.stdout
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
offset = max(0, start_offset)
while True:
try:
with log_path.open("rb") as log_file:
log_file.seek(0, os.SEEK_END)
if log_file.tell() < offset:
# Be defensive if another process replaced/truncated the log.
offset = 0
decoder.reset()
log_file.seek(offset)
chunk = log_file.read()
except OSError:
chunk = b""
if chunk:
offset += len(chunk)
text = decoder.decode(chunk)
if text:
sink.write(text)
sink.flush()
if stop_event.is_set():
break
stop_event.wait(poll_interval)
remainder = decoder.decode(b"", final=True)
if remainder:
sink.write(remainder)
sink.flush()
def _start_gateway_log_follower(
log_path: Path,
start_offset: int,
) -> tuple[threading.Event, threading.Thread]:
stop_event = threading.Event()
thread = threading.Thread(
target=_follow_gateway_log,
args=(log_path, start_offset, stop_event),
name="evoscientist-gateway-log",
daemon=True,
)
thread.start()
return stop_event, thread
@app.command()
@@ -39,7 +100,10 @@ def deploy(
workdir: str | None = typer.Option(
None,
"--workdir",
help="Workspace directory (default: config.default_workdir or cwd)",
help=(
"Workspace directory (default: config.default_workdir or "
"~/.evoscientist/workspace)"
),
),
port: int | None = typer.Option(
None,
@@ -64,11 +128,16 @@ def deploy(
Connect any LangChain-compatible UI or SDK client to the printed
endpoint. Press Ctrl+C to stop.
"""
from ..config import apply_config_to_env, get_effective_config
from ..config import (
apply_config_to_env,
get_default_workspace_dir,
get_effective_config,
)
from ..langgraph_dev.manager import (
_DEFAULT_PORT,
RUNTIME,
_is_port_occupied,
current_log_start_offset,
is_langgraph_dev_running,
read_tunnel_url,
start_langgraph_dev,
@@ -88,17 +157,21 @@ def deploy(
_configure_logging()
apply_config_to_env(config)
# 2. Resolve workspace (CLI > config.default_workdir > cwd)
# 2. Resolve workspace (CLI > config.default_workdir > stable app default)
if workdir:
ws = os.path.abspath(os.path.expanduser(workdir))
elif config.default_workdir:
ws = os.path.abspath(os.path.expanduser(config.default_workdir))
else:
ws = os.getcwd()
ws = str(get_default_workspace_dir())
# Subprocess inherits this path via EVOSCIENTIST_WORKSPACE_DIR (set inside
# start_langgraph_dev). Ensure the dir exists; do NOT mutate the parent
# process's paths module state — the deploy parent has no in-process agent.
os.makedirs(ws, exist_ok=True)
if getattr(config, "workspace_isolation", "optional") != "legacy":
from ..scope_registry import get_scope_service_token
os.environ["EVOSCIENTIST_BACKEND_SERVICE_TOKEN"] = get_scope_service_token(ws)
# 3. Resolve port (explicit None check — don't treat --port 0 as "unset"),
# then validate range so misconfigurations fail fast with a clear message
@@ -138,6 +211,18 @@ def deploy(
)
raise typer.Exit(1)
# Standalone WebUI development runs the backend and Next.js in separate
# terminals. Prepare the same durable identity, sink token, and spool used
# by the integrated launcher so the independently started WebUI can accept
# usage events through the shared data directory.
webui_port = int(getattr(config, "webui_port", 4716))
try:
usage_env = prepare_usage_environment(ws, webui_port=webui_port)
except (OSError, ValueError) as exc:
console.print(f"[yellow]Token statistics disabled:[/yellow] {exc}")
else:
os.environ.update(usage_env)
# 5. Startup banner
_auth_label = _describe_auth(config)
console.print(
@@ -169,9 +254,14 @@ def deploy(
"trust.[/bold red]"
)
# 6. ccproxy lifecycle (only if any provider uses OAuth)
# 6. ccproxy lifecycle (only if any provider uses OAuth). The legacy
# config.yaml auth-mode fields are gone with the unified model registry;
# getattr degrades to "api_key" so the OAuth branch simply stays off here.
_ccproxy_proc = None
if config.anthropic_auth_mode == "oauth" or config.openai_auth_mode == "oauth":
if (
getattr(config, "anthropic_auth_mode", "api_key") == "oauth"
or getattr(config, "openai_auth_mode", "api_key") == "oauth"
):
try:
from ..ccproxy_manager import maybe_start_ccproxy, stop_ccproxy
@@ -249,6 +339,15 @@ def deploy(
)
)
# Replay this session's startup output, then mirror new Gateway log lines
# while deploy owns the foreground terminal. The child continues writing
# directly to the file, so terminal backpressure cannot block the server.
console.print("[dim]Gateway logs (also saved to the path above):[/dim]")
log_stop_event, log_thread = _start_gateway_log_follower(
RUNTIME.log_file,
current_log_start_offset(),
)
# 10. Block on signal — mirror serve's dual-gate (threading.Event +
# explicit SIGINT/SIGTERM handlers) so SIGTERM (no default raise) also
# triggers clean shutdown.
@@ -268,6 +367,8 @@ def deploy(
except KeyboardInterrupt:
shutdown_event.set()
finally:
log_stop_event.set()
log_thread.join(timeout=2.0)
signal.signal(signal.SIGINT, _orig_sigint)
signal.signal(signal.SIGTERM, _orig_sigterm)
# stop_langgraph_dev + stop_ccproxy run via atexit during interpreter
+35 -6
View File
@@ -41,9 +41,15 @@ from ..stream.console import console
# Front-end npm package + spec. ``@latest`` → always the newest published UI.
_WEBUI_PACKAGE = "@evoscientist/webui@latest"
_WEBUI_PACKAGE_ENV = "EVOSCIENTIST_WEBUI_PACKAGE"
_DEFAULT_WEBUI_PORT = 4716
def _resolve_webui_package() -> str:
"""Return the npm package spec to launch for the WebUI front-end."""
return os.getenv(_WEBUI_PACKAGE_ENV) or _WEBUI_PACKAGE
def run_webui(config: Any, workspace_dir: str | None = None) -> None:
"""Start the deploy-style backend + the WebUI front-end, then block.
@@ -51,12 +57,12 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
config: Effective ``EvoScientistConfig`` (already env-applied upstream,
but re-applied here so this is safe to call standalone).
workspace_dir: Resolved workspace path; falls back to
``config.default_workdir`` then cwd.
``config.default_workdir`` then ``~/.evoscientist/workspace``.
Blocks until Ctrl+C / SIGTERM, or until the front-end process exits, then
tears down both subprocesses. Never returns a value.
"""
from ..config import apply_config_to_env
from ..config import apply_config_to_env, get_default_workspace_dir
from ..langgraph_dev.manager import (
_DEFAULT_PORT,
RUNTIME,
@@ -69,7 +75,7 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
apply_config_to_env(config)
# 1. Resolve workspace (CLI-resolved value > config.default_workdir > cwd),
# 1. Resolve workspace (CLI > config.default_workdir > stable app default),
# mirroring `EvoSci deploy`. The langgraph dev subprocess inherits this via
# EVOSCIENTIST_WORKSPACE_DIR (set inside start_langgraph_dev).
if workspace_dir:
@@ -77,8 +83,14 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
elif getattr(config, "default_workdir", ""):
ws = os.path.abspath(os.path.expanduser(config.default_workdir))
else:
ws = os.getcwd()
ws = str(get_default_workspace_dir())
os.makedirs(ws, exist_ok=True)
scope_service_token = ""
if getattr(config, "workspace_isolation", "optional") != "legacy":
from ..scope_registry import get_scope_service_token
scope_service_token = get_scope_service_token(ws)
os.environ["EVOSCIENTIST_BACKEND_SERVICE_TOKEN"] = scope_service_token
# 2. Resolve ports: backend = langgraph dev (browser connects here),
# webui_port = the local Next.js server the browser actually opens.
@@ -124,6 +136,18 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
)
raise typer.Exit(1)
# Token tracking is enabled only for this integrated topology. Generate the
# stable deployment/workspace identities and the dedicated transport token
# before starting either child so both processes share the exact values.
from ..usage import prepare_usage_environment
try:
usage_env = prepare_usage_environment(ws, webui_port=webui_port)
except (OSError, ValueError) as exc:
console.print(f"[yellow]Token statistics disabled:[/yellow] {exc}")
usage_env = {}
os.environ.update(usage_env)
# 4. Backend (langgraph dev): reuse an EvoSci server already on the port,
# else start a fresh deploy-mode one (full MCP + async). Refuse a foreign
# occupant — that's a configuration error, not something to silently share.
@@ -199,10 +223,15 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
# are scrubbed — the browser UI never needs LLM provider API keys.
webui_env = _scrubbed_env(
{
**usage_env,
"EVOSCIENTIST_LANGGRAPH_DEV_PORT": str(backend_port),
"EVOSCIENTIST_BACKEND_URL": f"http://127.0.0.1:{backend_port}",
"EVOSCIENTIST_BACKEND_SERVICE_TOKEN": scope_service_token,
"EVOSCIENTIST_WORKSPACE_DIR": ws,
"PORT": str(webui_port),
}
)
webui_package = _resolve_webui_package()
console.print(
Panel(
Text.from_markup(
@@ -211,7 +240,7 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
f"[bold]WebUI:[/bold] http://localhost:{webui_port} "
f"[dim](opens in your browser)[/dim]\n"
f"[bold]Logs:[/bold] {_shorten(str(RUNTIME.log_file))}\n\n"
f"[dim]Fetching {_WEBUI_PACKAGE} via npx (first run may take a "
f"[dim]Fetching {webui_package} via npx (first run may take a "
f"moment)… Press Ctrl+C to stop.[/dim]"
),
title="[bold green]✓ EvoScientist WebUI[/bold green]",
@@ -228,7 +257,7 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
popen_kwargs["creationflags"] = subprocess.CREATE_NEW_PROCESS_GROUP
try:
webui_proc = subprocess.Popen(
[npx, "--yes", _WEBUI_PACKAGE, "--port", str(webui_port)],
[npx, "--yes", webui_package, "--port", str(webui_port)],
**popen_kwargs,
)
except Exception as exc:
+3 -1
View File
@@ -476,7 +476,9 @@ async def alaunch_background_run(
thread_id,
name=request.name,
)
payload = request.run_payload(thread_id)
# Payload builders may do blocking work (source-thread metadata
# fetch, snapshot creation against SQLite); keep it off the loop.
payload = await asyncio.to_thread(request.run_payload, thread_id)
run_id = await _acreate_run(
client,
thread_id=thread_id,
+20
View File
@@ -147,14 +147,34 @@ class LocalGraphGateway:
target: GraphTarget,
request: RunRequest,
) -> AsyncIterator[GraphEvent]:
from langgraph.types import Command
from ..stream.events import stream_agent_events
# Local snapshot entry (design doc 8.1): freeze the registry
# defaults into a per-turn run snapshot and carry only its ID in
# configurable. Resumed runs (HITL Command(resume=...)) keep the
# snapshot the original turn froze — a resume must not re-freeze
# newer defaults. A bootstrap registry raises
# MODEL_REGISTRY_NOT_READY here instead of falling back to any
# implicit default model.
runtime_snapshot_id: str | None = None
if not (isinstance(request.message, Command) and request.message.resume):
from ..model_registry.runtime import get_snapshot_runtime
runtime_snapshot_id = (
get_snapshot_runtime()
.create_local_snapshot(request.thread_id)
.snapshot_id
)
inner = stream_agent_events(
local_graph,
request.message,
request.thread_id,
metadata=request.metadata,
media=request.media,
runtime_snapshot_id=runtime_snapshot_id,
)
try:
async for event in inner:
+11
View File
@@ -568,6 +568,17 @@ class LangGraphServerGateway:
request.message,
media=request.media,
)
# Local snapshot entry (design doc 8.1): the langgraph server shares
# the same model-runtime store, so the snapshot frozen here is the
# one the server-side run resolves. Resume turns returned above and
# keep the snapshot the original turn froze.
from ..model_registry.runtime import get_snapshot_runtime
config["configurable"]["runtime_snapshot_id"] = (
get_snapshot_runtime()
.create_local_snapshot(request.thread_id)
.snapshot_id
)
await stream.run.start(
input=run_input,
config=config,
+15
View File
@@ -0,0 +1,15 @@
"""Image generation: dedicated image models, adapters, and agent tools."""
from .config import (
ImageGenerationSettings,
ImageModelEntry,
is_image_generation_model,
load_image_generation_settings,
)
__all__ = [
"ImageGenerationSettings",
"ImageModelEntry",
"is_image_generation_model",
"load_image_generation_settings",
]
@@ -0,0 +1,7 @@
"""Image generation provider adapters."""
from .base import ImageGenAdapter, ImageGenError
from .gemini import GeminiImageAdapter
from .openai import OpenAIImageAdapter
__all__ = ["GeminiImageAdapter", "ImageGenAdapter", "ImageGenError", "OpenAIImageAdapter"]
+104
View File
@@ -0,0 +1,104 @@
"""Shared adapter protocol, errors, and response-parsing helpers."""
from __future__ import annotations
import base64
import binascii
import re
from collections.abc import Awaitable, Callable
from pathlib import Path
from typing import Any, Protocol
from urllib.parse import urlparse
class ImageGenError(Exception):
"""Provider or configuration failure.
The message must be safe to show an agent or a browser: no secrets, no
request headers, no URLs carrying credentials.
"""
class ImageGenAdapter(Protocol):
"""Vendor-specific image generation/editing."""
async def generate(
self, *, prompt: str, size: str, quality: str, background: str, n: int
) -> list[bytes]: ...
async def edit(
self,
*,
image: Path,
mask: Path | None,
prompt: str,
size: str,
quality: str,
) -> list[bytes]: ...
def strip_data_uri(value: str) -> str:
if value.startswith("data:image/") and "," in value:
return value.split(",", 1)[1]
return value
def looks_like_url(value: str) -> bool:
parsed = urlparse(value)
return parsed.scheme in {"http", "https"} and bool(parsed.netloc)
def looks_like_base64_image(value: str) -> bool:
if value.startswith("data:image/"):
return True
if len(value) < 16 or not re.fullmatch(r"[A-Za-z0-9+/=\s]+", value):
return False
try:
decoded = base64.b64decode(value, validate=False)
except (binascii.Error, ValueError):
return False
return decoded.startswith((b"\x89PNG", b"\xff\xd8\xff", b"RIFF", b"GIF8")) or len(
decoded
) > 256
def collect_image_values(
value: Any, *, b64_values: list[str], urls: list[str]
) -> None:
"""Recursively collect base64 payloads and URLs from a provider response."""
if isinstance(value, dict):
for key, item in value.items():
normalized = key.lower()
if isinstance(item, str):
if normalized in {"b64_json", "image_base64", "base64"}:
b64_values.append(item)
elif normalized == "result" and looks_like_base64_image(item):
b64_values.append(item)
elif normalized in {"url", "image_url"} and looks_like_url(item):
urls.append(item)
elif item.startswith("data:image/"):
b64_values.append(item)
else:
collect_image_values(item, b64_values=b64_values, urls=urls)
elif isinstance(value, list):
for item in value:
collect_image_values(item, b64_values=b64_values, urls=urls)
async def extract_image_bytes_from_values(
b64_values: list[str],
urls: list[str],
*,
download: Callable[[str], Awaitable[bytes]],
) -> list[bytes]:
images: list[bytes] = []
for value in b64_values:
try:
images.append(base64.b64decode(strip_data_uri(value), validate=False))
except (binascii.Error, ValueError):
continue
for url in urls:
images.append(await download(url))
if not images:
raise ImageGenError("Image provider returned no image data")
return images
+124
View File
@@ -0,0 +1,124 @@
"""Gemini (Imagen) image generation adapter.
Uses the Gemini API ``:predict`` endpoint. The API key travels in the
``x-goog-api-key`` header — never in the URL query — so it cannot leak
through logged URLs.
"""
from __future__ import annotations
import base64
import binascii
from pathlib import Path
from typing import Any
import httpx
from ..config import ImageModelEntry
from .base import ImageGenError
DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com"
# gpt-image-style sizes -> Imagen aspect ratios.
SIZE_TO_ASPECT = {
"1024x1024": "1:1",
"1536x1024": "16:9",
"1024x1536": "9:16",
}
class GeminiImageAdapter:
"""Calls ``{base}/v1beta/models/{model}:predict``."""
def __init__(
self,
entry: ImageModelEntry,
*,
client: httpx.AsyncClient | None = None,
timeout: float = 120.0,
) -> None:
self.entry = entry
self._client = client
self._timeout = timeout
def _api_key(self) -> str:
key = self.entry.resolved_api_key()
if not key:
raise ImageGenError(
f"Image model {self.entry.id!r} has no API key; "
"set the environment variable referenced in image_generation."
)
return key
def _endpoint_base(self) -> str:
return (self.entry.resolved_base_url() or DEFAULT_BASE_URL).rstrip("/")
def _raise_for_status(self, resp: httpx.Response) -> None:
if resp.status_code < 400:
return
detail = f"Image provider returned HTTP {resp.status_code}"
try:
payload = resp.json()
error = payload.get("error") if isinstance(payload, dict) else None
message = error.get("message") if isinstance(error, dict) else error
if isinstance(message, str) and message:
detail = f"{detail}: {message[:300]}"
except Exception:
pass
raise ImageGenError(detail)
async def generate(
self, *, prompt: str, size: str, quality: str, background: str, n: int
) -> list[bytes]:
# `quality`/`background` are not Imagen parameters and are ignored.
url = f"{self._endpoint_base()}/v1beta/models/{self.entry.id}:predict"
payload: dict[str, Any] = {
"instances": [{"prompt": prompt}],
"parameters": {
"sampleCount": n,
"aspectRatio": SIZE_TO_ASPECT.get(size, "1:1"),
**self.entry.params,
},
}
headers = {"x-goog-api-key": self._api_key()}
try:
if self._client is not None:
resp = await self._client.post(url, headers=headers, json=payload)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.post(url, headers=headers, json=payload)
except httpx.TimeoutException:
raise ImageGenError(
f"Image provider request timed out after {self._timeout:g}s"
) from None
except httpx.RequestError as exc:
raise ImageGenError(
f"Image provider request failed: {exc.__class__.__name__}"
) from None
self._raise_for_status(resp)
predictions = resp.json().get("predictions") or []
images: list[bytes] = []
for item in predictions:
value = item.get("bytesBase64Encoded") if isinstance(item, dict) else None
if not value:
continue
try:
images.append(base64.b64decode(value, validate=False))
except (binascii.Error, ValueError):
continue
if not images:
raise ImageGenError("Image provider returned no image data")
return images
async def edit(
self,
*,
image: Path,
mask: Path | None,
prompt: str,
size: str,
quality: str,
) -> list[bytes]:
raise ImageGenError(
f"Image model {self.entry.id!r} does not support edit."
)
+234
View File
@@ -0,0 +1,234 @@
"""OpenAI-compatible image generation/editing adapter."""
from __future__ import annotations
import asyncio
import mimetypes
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
import httpx
from ..config import ImageModelEntry
from .base import (
ImageGenError,
collect_image_values,
extract_image_bytes_from_values,
)
_TRANSIENT_STATUSES = {429, 500, 502, 503, 504}
_RETRY_DELAYS = (0.0, 1.0, 3.0)
_FALLBACK_MARKERS = (
"requires an image model",
"unsupported model",
"not supported",
"unknown model",
"model_not_found",
"not found",
"invalid endpoint",
)
class OpenAIImageAdapter:
"""Calls ``{base}/v1/images/*``; falls back to the Responses API."""
def __init__(
self,
entry: ImageModelEntry,
*,
client: httpx.AsyncClient | None = None,
timeout: float = 120.0,
) -> None:
self.entry = entry
self._client = client
self._timeout = timeout
# -- helpers -------------------------------------------------------------
def _api_key(self) -> str:
key = self.entry.resolved_api_key()
if not key:
raise ImageGenError(
f"Image model {self.entry.id!r} has no API key; "
"set the environment variable referenced in image_generation."
)
return key
def _endpoint_base(self) -> str:
base = (
self.entry.resolved_base_url() or "https://api.openai.com"
).rstrip("/")
return base if base.endswith("/v1") else f"{base}/v1"
def _raise_for_status(self, resp: httpx.Response) -> None:
if resp.status_code < 400:
return
detail = f"Image provider returned HTTP {resp.status_code}"
try:
payload = resp.json()
error = payload.get("error") if isinstance(payload, dict) else None
message = error.get("message") if isinstance(error, dict) else error
if isinstance(message, str) and message:
detail = f"{detail}: {message[:300]}"
except Exception:
pass
# Never include request headers, the URL, or the API key.
raise ImageGenError(detail)
async def _post_json(self, url: str, payload: dict[str, Any]) -> dict[str, Any]:
headers = {"Authorization": f"Bearer {self._api_key()}"}
last_response: httpx.Response | None = None
for delay in _RETRY_DELAYS:
if delay:
await asyncio.sleep(delay)
try:
if self._client is not None:
resp = await self._client.post(url, headers=headers, json=payload)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.post(url, headers=headers, json=payload)
except httpx.TimeoutException:
raise ImageGenError(
f"Image provider request timed out after {self._timeout:g}s"
) from None
except httpx.RequestError as exc:
raise ImageGenError(
f"Image provider request failed: {exc.__class__.__name__}"
) from None
last_response = resp
if resp.status_code not in _TRANSIENT_STATUSES:
self._raise_for_status(resp)
return resp.json()
assert last_response is not None
self._raise_for_status(last_response)
def _base_origin(self) -> str:
base = (self.entry.resolved_base_url() or "https://api.openai.com").rstrip("/")
parsed = urlparse(base)
return f"{parsed.scheme}://{parsed.netloc}"
async def _download(self, url: str) -> bytes:
# Only send the API key to the provider's own origin; provider-returned
# URLs on other hosts (e.g. pre-signed blob storage) must not leak it.
headers: dict[str, str] = {}
parsed = urlparse(url)
if f"{parsed.scheme}://{parsed.netloc}" == self._base_origin():
headers["Authorization"] = f"Bearer {self._api_key()}"
try:
if self._client is not None:
resp = await self._client.get(url, headers=headers)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.get(url, headers=headers)
except httpx.TimeoutException:
raise ImageGenError(
f"Image provider request timed out after {self._timeout:g}s"
) from None
except httpx.RequestError as exc:
raise ImageGenError(
f"Image provider request failed: {exc.__class__.__name__}"
) from None
self._raise_for_status(resp)
content_type = resp.headers.get("content-type", "")
if content_type and not content_type.startswith("image/"):
raise ImageGenError("Image provider returned a non-image URL")
return resp.content
async def _extract(self, response: dict[str, Any]) -> list[bytes]:
b64_values: list[str] = []
urls: list[str] = []
collect_image_values(response, b64_values=b64_values, urls=urls)
return await extract_image_bytes_from_values(
b64_values, urls, download=self._download
)
# -- adapter API ---------------------------------------------------------
async def generate(
self, *, prompt: str, size: str, quality: str, background: str, n: int
) -> list[bytes]:
base = self._endpoint_base()
payload: dict[str, Any] = {
**self.entry.params,
"model": self.entry.id,
"prompt": prompt,
"size": size,
"quality": quality,
"n": n,
"response_format": "b64_json",
}
if background and background != "auto":
payload["background"] = background
try:
response = await self._post_json(f"{base}/images/generations", payload)
except ImageGenError as exc:
message = str(exc).lower()
if not any(marker in message for marker in _FALLBACK_MARKERS):
raise
tool: dict[str, Any] = {"type": "image_generation", "size": size}
if quality and quality != "auto":
tool["quality"] = quality
if background and background != "auto":
tool["background"] = background
response = await self._post_json(
f"{base}/responses",
{
"model": self.entry.id,
"input": prompt,
"tools": [tool],
"tool_choice": {"type": "image_generation"},
"metadata": {"image_count": str(n)},
},
)
return await self._extract(response)
async def edit(
self,
*,
image: Path,
mask: Path | None,
prompt: str,
size: str,
quality: str,
) -> list[bytes]:
base = self._endpoint_base()
data = {
"model": self.entry.id,
"prompt": prompt,
"size": size,
"quality": quality,
"response_format": "b64_json",
}
image_mime = mimetypes.guess_type(image.name)[0] or "application/octet-stream"
# Read uploads off the event loop: the dev server aborts on blocking
# file I/O inside async tool calls.
image_bytes = await asyncio.to_thread(image.read_bytes)
files: list[tuple[str, tuple[str, Any, str]]] = [
("image", (image.name, image_bytes, image_mime))
]
if mask is not None:
mask_mime = mimetypes.guess_type(mask.name)[0] or "application/octet-stream"
mask_bytes = await asyncio.to_thread(mask.read_bytes)
files.append(("mask", (mask.name, mask_bytes, mask_mime)))
headers = {"Authorization": f"Bearer {self._api_key()}"}
try:
if self._client is not None:
resp = await self._client.post(
f"{base}/images/edits", headers=headers, data=data, files=files
)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.post(
f"{base}/images/edits", headers=headers, data=data, files=files
)
except httpx.TimeoutException:
raise ImageGenError(
f"Image provider request timed out after {self._timeout:g}s"
) from None
except httpx.RequestError as exc:
raise ImageGenError(
f"Image provider request failed: {exc.__class__.__name__}"
) from None
self._raise_for_status(resp)
return await self._extract(resp.json())
+158
View File
@@ -0,0 +1,158 @@
"""Configuration for dedicated image-generation models.
Image generation models are service models, not chat models. They live in a
separate ``image_generation`` section of ``config.yaml`` so image-only models
such as ``gpt-image-2`` are never offered in the chat model selector.
Secrets are stored as ``${ENV_VAR}`` references and resolved only at call
time, server-side. Tool signatures, skill docs, logs and error messages must
never carry a resolved key.
"""
from __future__ import annotations
import os
import re
from pathlib import Path
from typing import Any, Literal
import yaml
from pydantic import BaseModel, Field, ValidationError, field_validator
IMAGE_GENERATION_SECTION = "image_generation"
DEFAULT_IMAGE_GENERATION_TIMEOUT_SECONDS = 120.0
_IMAGE_MODEL_PREFIXES = ("gpt-image", "chatgpt-image", "dall-e", "dalle")
_IMAGE_MODEL_MARKERS = ("imagen", "wanx", "seedream")
_ENV_REF = re.compile(r"^\$\{([A-Za-z_][A-Za-z0-9_]*)\}$")
def _resolve_env_ref(value: str) -> str:
match = _ENV_REF.match(value.strip())
if match:
return os.environ.get(match.group(1), "")
return value
class ImageModelEntry(BaseModel):
"""A single dedicated image-generation model."""
id: str = Field(..., description="Model ID sent to the image API")
name: str = ""
provider: Literal["openai", "gemini"] = "openai"
api_key: str = Field("", description="API key or ${ENV_VAR} reference")
base_url: str = Field("", description="API base URL or ${ENV_VAR} reference")
enabled: bool = True
default_size: str = "1024x1024"
default_quality: str = "auto"
params: dict[str, Any] = Field(default_factory=dict)
@field_validator("id")
@classmethod
def id_not_empty(cls, value: str) -> str:
value = value.strip()
if not value:
raise ValueError("image model id must not be empty")
return value
def resolved_api_key(self) -> str:
return _resolve_env_ref(self.api_key)
def resolved_base_url(self) -> str:
return _resolve_env_ref(self.base_url)
def display_name(self) -> str:
return self.name or self.id
class ImageGenerationSettings(BaseModel):
"""The ``image_generation`` section of config.yaml."""
default_model: str = ""
timeout_seconds: float = DEFAULT_IMAGE_GENERATION_TIMEOUT_SECONDS
models: list[ImageModelEntry] = Field(default_factory=list)
@field_validator("timeout_seconds")
@classmethod
def timeout_must_be_positive(cls, value: float) -> float:
if value <= 0:
raise ValueError("timeout_seconds must be greater than 0")
return value
def find_model(self, model_id_or_name: str) -> ImageModelEntry | None:
for entry in self.models:
if entry.id == model_id_or_name or entry.name == model_id_or_name:
return entry
return None
def is_image_generation_model(model_ref: str | None) -> bool:
"""Return True when a model ID is known to be image-generation only."""
if not model_ref:
return False
value = str(model_ref).strip().lower()
if "/" in value:
value = value.rsplit("/", 1)[1]
return value.startswith(_IMAGE_MODEL_PREFIXES) or any(
marker in value for marker in _IMAGE_MODEL_MARKERS
)
def load_image_generation_settings(
*, config_path: Path | None = None
) -> ImageGenerationSettings:
"""Load the image_generation section; empty settings when absent."""
if config_path is None:
from EvoScientist.config.settings import get_config_path
config_path = get_config_path()
if not config_path.exists():
return ImageGenerationSettings()
try:
data = yaml.safe_load(config_path.read_text(encoding="utf-8")) or {}
except yaml.YAMLError:
return ImageGenerationSettings()
section = data.get(IMAGE_GENERATION_SECTION)
if not isinstance(section, dict):
return ImageGenerationSettings()
try:
return ImageGenerationSettings.model_validate(section)
except ValidationError as exc:
# Never echo input values: a mis-indented yaml can put a literal API
# key under the wrong field and pydantic's default message embeds it
# verbatim. Report only field locations and error types.
from EvoScientist.image_gen.adapters.base import ImageGenError
details = "; ".join(
f"{'.'.join(str(part) for part in error['loc'])}: {error['type']}"
for error in exc.errors(include_url=False, include_input=False)
)
raise ImageGenError(
f"invalid {IMAGE_GENERATION_SECTION} section: {details}"
) from None
def save_image_generation_settings(
settings: ImageGenerationSettings, *, config_path: Path | None = None
) -> None:
"""Replace the image_generation section, preserving other sections."""
if config_path is None:
from EvoScientist.config.settings import get_config_path
config_path = get_config_path()
data: dict[str, Any] = {}
if config_path.exists():
try:
loaded = yaml.safe_load(config_path.read_text(encoding="utf-8"))
except yaml.YAMLError:
loaded = None
if isinstance(loaded, dict):
data = loaded
data[IMAGE_GENERATION_SECTION] = settings.model_dump(mode="json")
config_path.parent.mkdir(parents=True, exist_ok=True)
config_path.write_text(
yaml.safe_dump(data, allow_unicode=True, sort_keys=False),
encoding="utf-8",
)
config_path.chmod(0o600)
+236
View File
@@ -0,0 +1,236 @@
"""Image generation service: resolve model, dispatch adapter, save safely."""
from __future__ import annotations
import asyncio
import re
import time
from datetime import datetime
from pathlib import Path, PurePosixPath
from typing import Any
from .adapters.base import ImageGenAdapter, ImageGenError
from .adapters.gemini import GeminiImageAdapter
from .adapters.openai import OpenAIImageAdapter
from .config import ImageModelEntry, load_image_generation_settings
ADAPTER_CLASSES: dict[str, type] = {
"openai": OpenAIImageAdapter,
"gemini": GeminiImageAdapter,
}
ALLOWED_INPUT_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp"}
MAX_INPUT_BYTES = 20 * 1024 * 1024
DEFAULT_MIME_TYPE = "image/png"
def _config_path() -> Path:
from EvoScientist.config.settings import get_config_path
return get_config_path()
def _available_hint(models: list[ImageModelEntry]) -> str:
names = [entry.display_name() for entry in models]
return f" Available image models: {', '.join(names)}." if names else ""
def _resolve_entry(model: str | None) -> tuple[ImageModelEntry, float]:
settings = load_image_generation_settings(config_path=_config_path())
if not settings.models:
raise ImageGenError(
"There are no image models configured; add an image_generation "
"section to config.yaml."
)
if model:
entry = settings.find_model(model)
if entry is None:
raise ImageGenError(
f"Unknown image model {model!r}."
+ _available_hint(settings.models)
)
if not entry.enabled:
raise ImageGenError(
f"Image model {entry.display_name()!r} is disabled."
)
return entry, settings.timeout_seconds
default = settings.default_model or settings.models[0].id
entry = settings.find_model(default)
if entry is not None and entry.enabled:
return entry, settings.timeout_seconds
for candidate in settings.models:
if candidate.enabled:
return candidate, settings.timeout_seconds
raise ImageGenError(
"There are no enabled image models; enable one in the "
"image_generation section of config.yaml."
)
def _adapter_for(entry: ImageModelEntry, timeout: float) -> ImageGenAdapter:
adapter_cls = ADAPTER_CLASSES.get(entry.provider)
if adapter_cls is None:
raise ImageGenError(f"Unknown image provider {entry.provider!r}.")
return adapter_cls(entry, timeout=timeout)
def _clean_logical_path(value: str) -> str:
raw = value.replace("\\", "/").strip()
if not raw:
raise ImageGenError("Path is required")
if raw.startswith("/") or re.match(r"^[A-Za-z]:/", raw):
raise ImageGenError("Absolute paths are not allowed")
path = PurePosixPath(raw)
if any(part in {"", ".", ".."} for part in path.parts):
raise ImageGenError("Path traversal is not allowed")
return str(path)
def _normalize_output_paths(
output_path: str | None, *, default_stem: str, count: int
) -> list[str]:
if output_path is None or not output_path.strip():
stamp = datetime.now().strftime("%Y%m%d_%H%M%S")
base = f"artifacts/{default_stem}_{stamp}.png"
else:
base = _clean_logical_path(output_path)
if not base.startswith("artifacts/"):
raise ImageGenError("output_path must start with artifacts/")
suffix = PurePosixPath(base).suffix.lower()
if suffix and suffix != ".png":
raise ImageGenError("Only PNG image outputs are supported")
if not suffix:
base = f"{base}.png"
if count <= 1:
return [base]
path = PurePosixPath(base)
return [f"{path.with_suffix('')}_{idx}{path.suffix}" for idx in range(1, count + 1)]
def _resolve_input_image(workspace: Path, logical_path: str) -> Path:
clean = _clean_logical_path(logical_path)
if PurePosixPath(clean).suffix.lower() not in ALLOWED_INPUT_EXTENSIONS:
raise ImageGenError("Unsupported image extension")
root = workspace.resolve()
target = (root / clean).resolve()
if not target.is_relative_to(root) or not target.is_file():
raise ImageGenError("Input image not found")
if target.stat().st_size > MAX_INPUT_BYTES:
raise ImageGenError("Input image is too large")
return target
def _save_images(
workspace: Path,
images: list[bytes],
*,
output_path: str | None,
default_stem: str,
) -> list[str]:
(workspace / "artifacts").mkdir(parents=True, exist_ok=True)
logical_paths = _normalize_output_paths(
output_path, default_stem=default_stem, count=len(images)
)
root = workspace.resolve()
saved: list[str] = []
for logical_path, data in zip(logical_paths, images, strict=True):
target = (root / logical_path).resolve()
if not target.is_relative_to(root):
raise ImageGenError("Invalid output_path")
target.parent.mkdir(parents=True, exist_ok=True)
target.write_bytes(data)
saved.append(logical_path)
return saved
def _success(paths: list[str], *, model: str, size: str) -> dict[str, Any]:
return {
"ok": True,
"path": paths[0] if paths else None,
"paths": paths,
"mime_type": DEFAULT_MIME_TYPE,
"model": model,
"size": size,
}
TEST_PROMPT = "A single red circle centered on a plain white background."
async def test_entry(entry: ImageModelEntry, *, timeout: float) -> int:
"""Generate one probe image against ``entry``; returns latency in ms.
The image bytes are discarded: this only proves the configured provider,
credentials and defaults can complete a real generation call.
"""
adapter = _adapter_for(entry, timeout)
started = time.monotonic()
await adapter.generate(
prompt=TEST_PROMPT,
size=entry.default_size,
quality=entry.default_quality,
background="auto",
n=1,
)
return int((time.monotonic() - started) * 1000)
async def generate_for_workspace(
workspace: Path,
*,
prompt: str,
model: str | None = None,
size: str = "1024x1024",
quality: str = "auto",
background: str = "auto",
output_path: str | None = None,
n: int = 1,
) -> dict[str, Any]:
# Config load and file I/O run in a thread: the langgraph dev server
# intercepts blocking calls (os.mkdir, read_text, stat) made directly on
# the event loop and aborts the tool call.
entry, timeout = await asyncio.to_thread(_resolve_entry, model)
adapter = _adapter_for(entry, timeout)
images = await adapter.generate(
prompt=prompt, size=size, quality=quality, background=background, n=n
)
saved = await asyncio.to_thread(
_save_images,
workspace,
images,
output_path=output_path,
default_stem="generated",
)
return _success(saved, model=entry.id, size=size)
async def edit_for_workspace(
workspace: Path,
*,
image_path: str,
prompt: str,
model: str | None = None,
mask_path: str | None = None,
size: str = "1024x1024",
quality: str = "auto",
output_path: str | None = None,
) -> dict[str, Any]:
entry, timeout = await asyncio.to_thread(_resolve_entry, model)
image = await asyncio.to_thread(_resolve_input_image, workspace, image_path)
mask = (
await asyncio.to_thread(_resolve_input_image, workspace, mask_path)
if mask_path
else None
)
adapter = _adapter_for(entry, timeout)
images = await adapter.edit(
image=image, mask=mask, prompt=prompt, size=size, quality=quality
)
saved = await asyncio.to_thread(
_save_images,
workspace,
images,
output_path=output_path,
default_stem="edited",
)
return _success(saved, model=entry.id, size=size)
File diff suppressed because it is too large Load Diff
+46 -10
View File
@@ -20,6 +20,7 @@ import shutil
import subprocess
import threading
import time
import uuid
from dataclasses import dataclass
from pathlib import Path
@@ -177,7 +178,9 @@ class WorkspaceMismatchError(RuntimeError):
"""
def _write_workspace_sidecar(workspace_dir: Path, pid: int) -> None:
def _write_workspace_sidecar(
workspace_dir: Path, pid: int, *, port: int = _DEFAULT_PORT
) -> None:
"""Record the workspace + pid of the langgraph dev we just started.
Atomic write via temp-file + ``os.replace``: without this, a concurrent
@@ -194,8 +197,20 @@ def _write_workspace_sidecar(workspace_dir: Path, pid: int) -> None:
try:
RUNTIME.pid_dir.mkdir(parents=True, exist_ok=True)
tmp = RUNTIME.workspace_sidecar.with_suffix(".json.tmp")
resolved_workspace = workspace_dir.resolve()
deployment_id = os.getenv("EVOSCIENTIST_DEPLOYMENT_ID") or str(
uuid.uuid5(uuid.NAMESPACE_URL, f"evoscientist:{resolved_workspace}")
)
tmp.write_text(
json.dumps({"workspace": str(workspace_dir), "pid": pid}), encoding="utf-8"
json.dumps(
{
"workspace": str(workspace_dir),
"pid": pid,
"deployment_id": deployment_id,
"api_url": _base_url(port),
}
),
encoding="utf-8",
)
os.replace(tmp, RUNTIME.workspace_sidecar)
except OSError as exc:
@@ -231,11 +246,14 @@ def _read_workspace_sidecar() -> dict | None:
return data
def _unlink_workspace_sidecar() -> None:
def _unlink_workspace_sidecar(
runtime: LanggraphRuntimePaths | None = None,
) -> None:
"""Best-effort sidecar removal — called alongside every ``RUNTIME.pid_file.unlink()``
so the workspace fingerprint never outlives the PID file it pairs with."""
runtime = runtime or RUNTIME
try:
RUNTIME.workspace_sidecar.unlink()
runtime.workspace_sidecar.unlink()
except OSError:
pass
@@ -270,6 +288,16 @@ _LOG_OFFSET_AT_START: int = 0
# langgraph dev log. Mirrors langgraph_api/tunneling/cloudflare.py.
_TUNNEL_URL_RE = re.compile(r"https://[A-Za-z0-9.-]+\.trycloudflare\.com")
def current_log_start_offset() -> int:
"""Return the byte offset where the current server session started.
Consumers that mirror the Gateway log can start here to avoid replaying
output from earlier deploy sessions that share the append-only log file.
"""
return _LOG_OFFSET_AT_START
# Whether async sub-agents are usable in this process.
#
# - CLI / serve parent process: starts False; flipped True after
@@ -743,7 +771,7 @@ def start_langgraph_dev(
except Exception:
pass
RUNTIME.pid_file.write_text(str(proc.pid), encoding="utf-8")
_write_workspace_sidecar(workspace_dir=workspace_dir, pid=proc.pid)
_write_workspace_sidecar(workspace_dir=workspace_dir, pid=proc.pid, port=port)
global _PROCESS_WORKSPACE
_PROCESS = proc
_PROCESS_WORKSPACE = workspace_dir
@@ -820,7 +848,11 @@ def read_tunnel_url(timeout: float = 35.0, poll_interval: float = 0.5) -> str |
return None
def stop_langgraph_dev(proc: subprocess.Popen | None = None) -> None:
def stop_langgraph_dev(
proc: subprocess.Popen | None = None,
*,
runtime: LanggraphRuntimePaths | None = None,
) -> None:
"""Gracefully stop a langgraph dev process.
Sends SIGTERM to the process group (langgraph dev spawns worker children),
@@ -831,6 +863,7 @@ def stop_langgraph_dev(proc: subprocess.Popen | None = None) -> None:
(which also hold ``_LOCK``) don't observe partially-cleared state.
"""
global _PROCESS, _PROCESS_WORKSPACE
runtime = runtime or RUNTIME
with _LOCK:
proc = proc if proc is not None else _PROCESS
if proc is None:
@@ -879,12 +912,12 @@ def stop_langgraph_dev(proc: subprocess.Popen | None = None) -> None:
if proc is _PROCESS:
_PROCESS = None
_PROCESS_WORKSPACE = None
if RUNTIME.pid_file.exists():
if runtime.pid_file.exists():
try:
RUNTIME.pid_file.unlink()
runtime.pid_file.unlink()
except OSError:
pass
_unlink_workspace_sidecar()
_unlink_workspace_sidecar(runtime)
# Note: ``.langgraph_api/`` is intentionally NOT removed — it holds
# langgraph dev's persisted async-task / scheduler / Store state that
@@ -1064,5 +1097,8 @@ def _ensure_langgraph_dev_locked(
return None
_ASYNC_SUBAGENTS_AVAILABLE = True
atexit.register(stop_langgraph_dev, proc)
# Capture the runtime paths alongside the process. Embedded callers and
# tests may replace the module-level RUNTIME before interpreter shutdown;
# consulting that later would unlink an unrelated live deployment's files.
atexit.register(stop_langgraph_dev, proc, runtime=RUNTIME)
return proc
+14 -16
View File
@@ -1,33 +1,31 @@
"""LLM module for EvoScientist.
Provides a unified interface for creating chat model instances
with support for multiple providers.
The static model catalog and free-string model factory were removed in the
unified model-configuration refactor: all chat-model construction now goes
through :mod:`EvoScientist.model_registry` (SnapshotRuntime +
``build_chat_model`` + ``ResolvedModelConfig``).
``models`` is attached lazily via :mod:`lazy_loader` (SPEC-1 / PEP 562) so
that importing ``EvoScientist.llm`` (or any of its submodules, like
``context_window``) does not eagerly drag in ``langchain.chat_models`` and
its transitive ``langchain_anthropic``/``langchain_openai`` stack — that's
roughly 1 s of wall time on every CLI invocation.
What remains here:
- ``context_window`` — context-window resolution helpers consumed by the
middleware layer (e.g. ``context_editing``);
- ``patches`` — LangChain provider monkey-patches, attached lazily via
:mod:`lazy_loader` (SPEC-1 / PEP 562) so that importing
``EvoScientist.llm`` does not eagerly drag in ``langchain.chat_models``
and its transitive ``langchain_anthropic``/``langchain_openai`` stack —
that's roughly 1 s of wall time on every CLI invocation.
"""
import lazy_loader as _lazy
__getattr__, __dir__, __all__ = _lazy.attach(
__name__,
submodules=["context_window", "models", "patches"],
submodules=["context_window", "patches"],
submod_attrs={
"context_window": [
"DEFAULT_CONTEXT_WINDOW_FALLBACK",
"get_context_window",
"resolve_context_window",
],
"models": [
"DEFAULT_MODEL",
"MODELS",
"get_chat_model",
"get_model_info",
"get_models_for_provider",
"list_models",
],
},
)
-725
View File
@@ -1,725 +0,0 @@
"""LLM model configuration based on LangChain init_chat_model.
This module provides a unified interface for creating chat model instances
with support for multiple providers (Anthropic, OpenAI, Google GenAI, MiniMax
(Anthropic-compatible), NVIDIA, SiliconFlow, OpenRouter, ZhipuAI, Volcengine,
DashScope, DashScope-Code, DeepSeek, Ollama, and custom OpenAI/Anthropic-compatible
endpoints) and convenient short names for common models.
"""
from __future__ import annotations
import os
import warnings
from typing import Any
from langchain.chat_models import init_chat_model
from .context_window import apply_known_context_window
from .patches import (
_is_ccproxy_codex,
_patch_ccproxy_system_to_developer,
_patch_deepseek_reasoning_passback,
_patch_openai_compat_content,
_patch_openrouter_strip_responses_reasoning,
)
_MINIMAX_ANTHROPIC_BASE_URL = "https://api.minimaxi.com/anthropic"
_SILICONFLOW_BASE_URL = "https://api.siliconflow.cn/v1"
_ZHIPU_BASE_URL = "https://open.bigmodel.cn/api/paas/v4"
_ZHIPU_CODE_BASE_URL = "https://open.bigmodel.cn/api/coding/paas/v4"
_VOLCENGINE_BASE_URL = "https://ark.cn-beijing.volces.com/api/v3"
_DASHSCOPE_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
_DASHSCOPE_CODE_BASE_URL = "https://coding.dashscope.aliyuncs.com/v1"
_DEEPSEEK_BASE_URL = "https://api.deepseek.com"
_MOONSHOT_BASE_URL = "https://api.moonshot.cn/v1"
_KIMI_CODING_BASE_URL = "https://api.kimi.com/coding/"
# 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]] = {
"deepseek": (_DEEPSEEK_BASE_URL, "DEEPSEEK_API_KEY"),
"moonshot": (_MOONSHOT_BASE_URL, "MOONSHOT_API_KEY"),
"siliconflow": (_SILICONFLOW_BASE_URL, "SILICONFLOW_API_KEY"),
"zhipu": (_ZHIPU_BASE_URL, "ZHIPU_API_KEY"),
"zhipu-code": (_ZHIPU_CODE_BASE_URL, "ZHIPU_API_KEY"),
"volcengine": (_VOLCENGINE_BASE_URL, "VOLCENGINE_API_KEY"),
"dashscope": (_DASHSCOPE_BASE_URL, "DASHSCOPE_API_KEY"),
"dashscope-code": (_DASHSCOPE_CODE_BASE_URL, "DASHSCOPE_API_KEY"),
"custom-openai": (
None,
"CUSTOM_OPENAI_API_KEY",
), # base_url from CUSTOM_OPENAI_BASE_URL env
}
# Providers routed through the Anthropic provider with a custom base_url.
# Maps provider name → (base_url or None, env var for API key).
_ANTHROPIC_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = {
"minimax": (_MINIMAX_ANTHROPIC_BASE_URL, "MINIMAX_API_KEY"),
"kimi-coding": (_KIMI_CODING_BASE_URL, "KIMI_API_KEY"),
"custom-anthropic": (None, "CUSTOM_ANTHROPIC_API_KEY"),
}
# Anthropic-routed providers that support extended thinking.
_THINKING_CAPABLE_PROVIDERS: set[str] = {"minimax"}
_TRUTHY_ENV_VALUES = {"1", "true", "yes", "on"}
_FALSEY_ENV_VALUES = {"0", "false", "no", "off"}
# Legacy/provider-specific options that are not accepted by the installed
# LangChain chat model constructors. Leaving them at the top level makes
# LangChain move them into model_kwargs and can later leak them into SDK calls.
_UNSUPPORTED_CHAT_MODEL_KWARGS = frozenset({"sanitize_openai_sdk_headers"})
# Model registry: list of (short_name, model_id, provider)
# Allows same short_name across different providers.
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
# Custom Anthropic (third-party Claude-compatible endpoints, current-gen defaults)
# Listed BEFORE native anthropic so MODELS dict defaults to native provider
("claude-sonnet-4-6", "claude-sonnet-4-6", "custom-anthropic"),
("claude-haiku-4-5", "claude-haiku-4-5", "custom-anthropic"),
# Custom OpenAI (third-party OpenAI-compatible endpoints, 3 defaults)
# Listed BEFORE native openai so MODELS dict defaults to native provider
("gpt-5.5-pro", "gpt-5.5-pro", "custom-openai"),
("gpt-5.5", "gpt-5.5", "custom-openai"),
("gpt-5.4", "gpt-5.4", "custom-openai"),
("gpt-5.3-codex", "gpt-5.3-codex", "custom-openai"),
("gpt-5-mini", "gpt-5-mini", "custom-openai"),
# Anthropic (current generation)
("claude-fable-5", "claude-fable-5", "anthropic"),
("claude-opus-4-8", "claude-opus-4-8", "anthropic"),
("claude-sonnet-5", "claude-sonnet-5", "anthropic"),
("claude-sonnet-4-6", "claude-sonnet-4-6", "anthropic"),
("claude-haiku-4-5", "claude-haiku-4-5", "anthropic"),
# OpenAI
("gpt-5.6-sol", "gpt-5.6-sol", "openai"),
("gpt-5.6-terra", "gpt-5.6-terra", "openai"),
("gpt-5.6-luna", "gpt-5.6-luna", "openai"),
("gpt-5.5-pro", "gpt-5.5-pro", "openai"),
("gpt-5.5", "gpt-5.5", "openai"),
("gpt-5.4", "gpt-5.4", "openai"),
("gpt-5.4-mini", "gpt-5.4-mini", "openai"),
("gpt-5.4-nano", "gpt-5.4-nano", "openai"),
("gpt-5.3-codex", "gpt-5.3-codex", "openai"),
("gpt-5.2-codex", "gpt-5.2-codex", "openai"),
("gpt-5.2", "gpt-5.2", "openai"),
("gpt-5.1", "gpt-5.1", "openai"),
("gpt-5", "gpt-5", "openai"),
("gpt-5-mini", "gpt-5-mini", "openai"),
("gpt-5-nano", "gpt-5-nano", "openai"),
# Google GenAI
("gemini-3.5-flash", "gemini-3.5-flash", "google-genai"),
("gemini-3.1-pro", "gemini-3.1-pro-preview", "google-genai"),
(
"gemini-3.1-pro-customtools",
"gemini-3.1-pro-preview-customtools",
"google-genai",
),
("gemini-3.1-flash-lite", "gemini-3.1-flash-lite-preview", "google-genai"),
("gemini-3-flash", "gemini-3-flash-preview", "google-genai"),
("gemini-2.5-flash", "gemini-2.5-flash", "google-genai"),
("gemini-2.5-flash-lite", "gemini-2.5-flash-lite", "google-genai"),
("gemini-2.5-pro", "gemini-2.5-pro", "google-genai"),
# MiniMax (direct API — Anthropic-compatible; default: api.minimaxi.com, global: api.minimax.io)
("minimax-m3", "MiniMax-M3", "minimax"),
("minimax-m2.7", "MiniMax-M2.7", "minimax"),
("minimax-m2.7-highspeed", "MiniMax-M2.7-highspeed", "minimax"),
("minimax-m2.5", "MiniMax-M2.5", "minimax"),
("minimax-m2.5-highspeed", "MiniMax-M2.5-highspeed", "minimax"),
# NVIDIA
("nemotron-super", "nvidia/nemotron-3-super-120b-a12b", "nvidia"),
("nemotron-nano", "nvidia/nemotron-3-nano-30b-a3b", "nvidia"),
("glm-5.2", "z-ai/glm-5.2", "nvidia"),
("glm4.7", "z-ai/glm4.7", "nvidia"),
("deepseek-v3.2", "deepseek-ai/deepseek-v3.2", "nvidia"),
("deepseek-v3.1", "deepseek-ai/deepseek-v3.1-terminus", "nvidia"),
("kimi-k2.5", "moonshotai/kimi-k2.5", "nvidia"),
("kimi-k2-thinking", "moonshotai/kimi-k2-thinking", "nvidia"),
("minimax-m2.5", "minimaxai/minimax-m2.5", "nvidia"),
("minimax-m2.1", "minimaxai/minimax-m2.1", "nvidia"),
("qwen3.5-397b", "qwen/qwen3.5-397b-a17b", "nvidia"),
("step-3.5-flash", "stepfun-ai/step-3.5-flash", "nvidia"),
# SiliconFlow
("minimax-m2.5", "Pro/MiniMaxAI/MiniMax-M2.5", "siliconflow"),
("glm-5.2", "Pro/zai-org/GLM-5.2", "siliconflow"),
("glm-5", "Pro/zai-org/GLM-5", "siliconflow"),
("kimi-k2.5", "Pro/moonshotai/Kimi-K2.5", "siliconflow"),
("glm-4.7", "Pro/zai-org/GLM-4.7", "siliconflow"),
# OpenRouter
("claude-fable-5", "anthropic/claude-fable-5", "openrouter"),
("claude-opus-4.8", "anthropic/claude-opus-4.8", "openrouter"),
("claude-opus-4.8-fast", "anthropic/claude-opus-4.8-fast", "openrouter"),
("claude-sonnet-5", "anthropic/claude-sonnet-5", "openrouter"),
("claude-sonnet-4.6", "anthropic/claude-sonnet-4.6", "openrouter"),
("gpt-5.6-sol", "openai/gpt-5.6-sol", "openrouter"),
("gpt-5.6-terra", "openai/gpt-5.6-terra", "openrouter"),
("gpt-5.6-luna", "openai/gpt-5.6-luna", "openrouter"),
("gpt-5.5-pro", "openai/gpt-5.5-pro", "openrouter"),
("gpt-5.5", "openai/gpt-5.5", "openrouter"),
("gpt-5.4", "openai/gpt-5.4", "openrouter"),
("gpt-5.3-codex", "openai/gpt-5.3-codex", "openrouter"),
("gemini-3.5-flash", "google/gemini-3.5-flash", "openrouter"),
("gemini-3.1-pro", "google/gemini-3.1-pro-preview", "openrouter"),
("gemini-3-flash", "google/gemini-3-flash-preview", "openrouter"),
("kimi-k2.6", "moonshotai/kimi-k2.6", "openrouter"),
("glm-5.2", "z-ai/glm-5.2", "openrouter"),
("glm-5v-turbo", "z-ai/glm-5v-turbo", "openrouter"),
("minimax-m3", "minimax/minimax-m3", "openrouter"),
("mimo-v2.5-pro", "xiaomi/mimo-v2.5-pro", "openrouter"),
("mimo-v2.5", "xiaomi/mimo-v2.5", "openrouter"),
("grok-build-0.1", "x-ai/grok-build-0.1", "openrouter"),
("grok-4.5", "x-ai/grok-4.5", "openrouter"),
("hy3", "tencent/hy3", "openrouter"),
("qwen3.7-max", "qwen/qwen3.7-max", "openrouter"),
("qwen3.7-plus", "qwen/qwen3.7-plus", "openrouter"),
("qwen3.6-flash", "qwen/qwen3.6-flash", "openrouter"),
("qwen3.5-122b", "qwen/qwen3.5-122b-a10b", "openrouter"),
("deepseek-v4-pro", "deepseek/deepseek-v4-pro", "openrouter"),
("deepseek-v4-flash", "deepseek/deepseek-v4-flash", "openrouter"),
# Zhipu CodePlan (智谱代码计划 — coding-only endpoint)
("glm-5.2", "glm-5.2", "zhipu-code"),
("glm-5.1", "glm-5.1", "zhipu-code"),
("glm-5", "glm-5", "zhipu-code"),
("glm-5-turbo", "glm-5-turbo", "zhipu-code"),
("glm-5v-turbo", "glm-5v-turbo", "zhipu-code"),
("glm-4.7", "glm-4.7", "zhipu-code"),
# Zhipu (智谱 — general endpoint, default for simple lookups)
("glm-5.2", "glm-5.2", "zhipu"),
("glm-5.1", "glm-5.1", "zhipu"),
("glm-5", "glm-5", "zhipu"),
("glm-5-turbo", "glm-5-turbo", "zhipu"),
("glm-5v-turbo", "glm-5v-turbo", "zhipu"),
("glm-4.7", "glm-4.7", "zhipu"),
# Volcengine (火山引擎 — Doubao models)
("doubao-seed-2.0-pro", "doubao-seed-2-0-pro-260215", "volcengine"),
("doubao-seed-2.0-lite", "doubao-seed-2-0-lite-260215", "volcengine"),
("doubao-seed-2.0-mini", "doubao-seed-2-0-mini-260215", "volcengine"),
("doubao-seed-2.0-code", "doubao-seed-2-0-code-preview-260215", "volcengine"),
("doubao-seed-1.6", "doubao-seed-1.6", "volcengine"),
("doubao-1.5-pro", "doubao-1.5-pro-256k", "volcengine"),
("doubao-1.5-thinking-pro", "doubao-1.5-thinking-pro", "volcengine"),
# DashScope Coding Plan (阿里云代码计划 — subscription sk-sp-* endpoint)
("qwen3.7-max", "qwen3.7-max", "dashscope-code"),
("qwen3.7-plus", "qwen3.7-plus", "dashscope-code"),
("qwen3.6-max", "qwen3.6-max-preview", "dashscope-code"),
("qwen3.6-plus", "qwen3.6-plus", "dashscope-code"),
("qwen3.6-flash", "qwen3.6-flash", "dashscope-code"),
("qwen3-coder", "qwen3-coder-plus", "dashscope-code"),
("qwen3-coder-next", "qwen3-coder-next", "dashscope-code"),
("qwen3-max", "qwen3-max", "dashscope-code"),
("qwen3.5-plus", "qwen3.5-plus", "dashscope-code"),
# DashScope (阿里云 — Qwen models, default for simple lookups)
("qwen3.7-max", "qwen3.7-max", "dashscope"),
("qwen3.7-plus", "qwen3.7-plus", "dashscope"),
("qwen3.6-max", "qwen3.6-max-preview", "dashscope"),
("qwen3.6-plus", "qwen3.6-plus", "dashscope"),
("qwen3.6-flash", "qwen3.6-flash", "dashscope"),
("qwen3-coder", "qwen3-coder-plus", "dashscope"),
("qwen3-235b", "qwen3-235b-a22b", "dashscope"),
("qwen-max", "qwen-max", "dashscope"),
("qwq-plus", "qwq-plus", "dashscope"),
# DeepSeek
("deepseek-v4-pro", "deepseek-v4-pro", "deepseek"),
("deepseek-v4-flash", "deepseek-v4-flash", "deepseek"),
# Legacy aliases (deprecated 2026-07-24; route to v4-flash thinking/non-thinking)
("deepseek-r1", "deepseek-reasoner", "deepseek"),
("deepseek-v3", "deepseek-chat", "deepseek"),
# Moonshot (OpenAI-compatible)
("kimi-k2.6", "kimi-k2.6", "moonshot"),
("kimi-k2.5", "kimi-k2.5", "moonshot"),
("kimi-k2-thinking", "kimi-k2-thinking", "moonshot"),
("kimi-k2-thinking-turbo", "kimi-k2-thinking-turbo", "moonshot"),
("moonshot-v1-auto", "moonshot-v1-auto", "moonshot"),
("moonshot-v1-128k", "moonshot-v1-128k", "moonshot"),
("moonshot-v1-32k", "moonshot-v1-32k", "moonshot"),
("moonshot-v1-8k", "moonshot-v1-8k", "moonshot"),
# Kimi Coding Plan (Anthropic-compatible)
("kimi-for-coding", "kimi-for-coding", "kimi-coding"),
]
# Public dict for simple lookups (last entry wins for duplicate names).
# Use get_models_for_provider() for provider-aware lookups.
MODELS: dict[str, tuple[str, str]] = {
name: (model_id, provider) for name, model_id, provider in _MODEL_ENTRIES
}
DEFAULT_MODEL = "claude-sonnet-4-6"
def get_models_for_provider(provider: str) -> list[tuple[str, str]]:
"""Get all models for a specific provider.
Args:
provider: Provider name (e.g., 'anthropic', 'openrouter').
Returns:
List of (short_name, model_id) tuples for the provider.
"""
return [(name, model_id) for name, model_id, p in _MODEL_ENTRIES if p == provider]
def _env_flag_enabled(name: str) -> bool:
return os.environ.get(name, "").strip().lower() in _TRUTHY_ENV_VALUES
def _env_flag_disabled(name: str) -> bool:
value = os.environ.get(name)
return value is not None and value.strip().lower() in _FALSEY_ENV_VALUES
def _drop_unsupported_chat_model_kwargs(kwargs: dict[str, Any]) -> None:
for key in _UNSUPPORTED_CHAT_MODEL_KWARGS:
kwargs.pop(key, None)
model_kwargs = kwargs.get("model_kwargs")
if isinstance(model_kwargs, dict):
for key in _UNSUPPORTED_CHAT_MODEL_KWARGS:
model_kwargs.pop(key, None)
def _supports_openrouter_anthropic_prompt_cache(provider: str, model_id: str) -> bool:
"""Return whether EvoScientist should declare OpenRouter Claude caching."""
return provider == "openrouter" and model_id.startswith(
("anthropic/", "~anthropic/")
)
def _has_cache_control_override(kwargs: dict[str, Any]) -> bool:
"""Return whether the caller already supplied cache-control settings."""
if "cache_control" in kwargs:
return True
model_kwargs = kwargs.get("model_kwargs")
if model_kwargs is None:
return False
if not isinstance(model_kwargs, dict):
warnings.warn(
"OpenRouter Anthropic prompt caching was not applied because "
"`model_kwargs` is not a dict; pass cache_control explicitly or use "
"a dict-shaped model_kwargs.",
UserWarning,
stacklevel=3,
)
return True
return "cache_control" in model_kwargs
def _apply_openrouter_anthropic_prompt_cache(
provider: str,
model_id: str,
kwargs: dict[str, Any],
) -> None:
"""Declare OpenRouter Claude prompt caching unless explicitly disabled.
OpenRouter already handles implicit caching for most providers, but Claude
prompt caching needs Anthropic-style cache-control declaration.
"""
if _env_flag_disabled("EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"):
return
if not _supports_openrouter_anthropic_prompt_cache(provider, model_id):
return
if _has_cache_control_override(kwargs):
return
kwargs.setdefault("model_kwargs", {})["cache_control"] = {"type": "ephemeral"}
def _apply_auto_config(
provider: str,
model_id: str,
is_third_party: bool,
kwargs: dict[str, Any],
original_provider: str | None = None,
) -> None:
"""Auto-enable provider-specific features (thinking, reasoning, etc.).
Mutates *kwargs* in place. Only sets keys that the caller hasn't already
provided, so explicit user settings are never overridden.
"""
disable_reasoning = bool(kwargs.pop("_disable_reasoning", False))
disable_thinking = bool(kwargs.pop("_disable_thinking", False))
if disable_reasoning:
kwargs.pop("reasoning", None)
kwargs.pop("include_thoughts", None)
if disable_thinking:
kwargs.pop("thinking", None)
# Anthropic: extended thinking
if provider == "anthropic" and not disable_thinking and "thinking" not in kwargs:
_supports_thinking = original_provider in _THINKING_CAPABLE_PROVIDERS
# Detect local proxy (e.g. ccproxy): thinking blocks in conversation
# history cause 422 errors because the proxy doesn't accept 'thinking'
# as a valid content block type on round-trip.
if not is_third_party:
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
_is_proxy = "127.0.0.1" in base_url or "localhost" in base_url
else:
_is_proxy = False
if _is_proxy or (is_third_party and not _supports_thinking):
pass
elif "fable" in model_id or model_id.endswith(("4-6", "4-7", "4-8")):
kwargs["thinking"] = {"type": "adaptive", "display": "summarized"}
kwargs.setdefault("effort", "max")
else:
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000}
# OpenAI (native, not third-party routed): reasoning
if (
provider == "openai"
and not is_third_party
and not disable_reasoning
and "reasoning" not in kwargs
):
if _is_ccproxy_codex(kwargs.get("base_url"), kwargs.get("api_key")):
# ccproxy uses Chat Completions which doesn't support reasoning.
pass
else:
_eff = (
"xhigh"
if ("5.4" in model_id or "5.5" in model_id or "codex" in model_id)
else "high"
)
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
# Google GenAI: surface thinking traces
if provider == "google-genai" and not disable_reasoning:
kwargs.setdefault("include_thoughts", True)
# Ollama: separate reasoning content from response for thinking models
if provider == "ollama" and not disable_reasoning and "reasoning" not in kwargs:
kwargs["reasoning"] = True
def get_chat_model(
model: str | None = None,
provider: str | None = None,
**kwargs: Any,
) -> Any:
"""Get a chat model instance.
Args:
model: Model name (short name like 'claude-sonnet-4-6' or full ID
like 'claude-sonnet-4-6-20250929'). Defaults to DEFAULT_MODEL.
provider: Override the provider (e.g., 'anthropic', 'openai').
If not specified, inferred from model name or defaults to 'anthropic'.
**kwargs: Additional arguments passed to init_chat_model (e.g., temperature).
Returns:
A LangChain chat model instance.
Examples:
>>> model = get_chat_model() # Uses default (claude-sonnet-4-6)
>>> model = get_chat_model("claude-opus-4-8") # Use short name
>>> 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
# Look up short name in registry (provider-aware)
model_id = None
if provider:
# Try exact match with provider first
for name, mid, p in _MODEL_ENTRIES:
if name == model and p == provider:
model_id = mid
break
if model_id is None and model in MODELS:
model_id, default_provider = MODELS[model]
provider = provider or default_provider
if model_id is None:
# Assume it's a full model ID
model_id = model
# Try to infer provider from model ID prefix
if provider is None:
if model_id.startswith(("claude-", "anthropic")):
provider = "anthropic"
elif model_id.startswith(("gpt-", "o1", "davinci", "text-")):
provider = "openai"
elif model_id.startswith("gemini"):
provider = "google-genai"
elif model_id.startswith("ollama:"):
provider = "ollama"
model_id = model_id.removeprefix("ollama:")
else:
provider = "anthropic" # Default fallback
# Anthropic base_url override (e.g. ccproxy at localhost:8000/api/v1)
_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
if (
runtime_resolved is not None
and provider == "openai"
and resolved_base_url
and "api.openai.com" not in resolved_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
)
if provider == "anthropic":
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
if base_url:
kwargs.setdefault("base_url", base_url)
api_key = os.environ.get("ANTHROPIC_API_KEY", "")
if api_key:
kwargs.setdefault("api_key", api_key)
# Native OpenAI base_url override (e.g. ccproxy Codex at localhost:8000/codex/v1)
elif provider == "openai":
base_url = os.environ.get("OPENAI_BASE_URL", "")
if base_url:
kwargs.setdefault("base_url", base_url)
_is_openai_proxy = _is_ccproxy_codex(
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
api_key = os.environ.get("OPENAI_API_KEY", "")
if api_key:
kwargs.setdefault("api_key", api_key)
# OpenAI-routed providers → route through OpenAI provider with base_url
elif provider in _OPENAI_ROUTED_PROVIDERS:
_original_provider = provider
base_url_default, api_key_env = _OPENAI_ROUTED_PROVIDERS[provider]
if provider == "custom-openai":
base_url = os.environ.get("CUSTOM_OPENAI_BASE_URL", "")
if not base_url:
raise ValueError(
"CUSTOM_OPENAI_BASE_URL environment variable is required when using "
"the 'custom-openai' provider. Please set it to your "
"OpenAI-compatible API endpoint URL (e.g. https://api.openai.com/v1)."
)
base_url = base_url.rstrip("/")
else:
base_url = base_url_default
if base_url:
kwargs.setdefault("base_url", base_url)
api_key = os.environ.get(api_key_env, "")
if api_key:
kwargs.setdefault("api_key", api_key)
# SiliconFlow: disable thinking — LangChain drops reasoning_content
# from history, causing error 20015 on multi-turn requests.
if provider == "siliconflow":
kwargs.setdefault("extra_body", {})["enable_thinking"] = False
# Moonshot: disable thinking for all models to prevent LangChain from dropping
# reasoning_content, which causes multi-turn conversation errors (error 20015).
# Even native thinking models like kimi-k2-thinking operate in non-thinking mode.
if provider == "moonshot":
kwargs.setdefault("extra_body", {})["thinking"] = {"type": "disabled"}
provider = "openai"
# OpenRouter → native ChatOpenRouter via init_chat_model.
elif provider == "openrouter":
_is_third_party = True
api_key = os.environ.get("OPENROUTER_API_KEY", "")
if api_key:
kwargs.setdefault("api_key", api_key)
# Reasoning via `effort` + `summary: "auto"` so a readable reasoning
# summary is returned for display. OpenAI-Responses also emits encrypted
# reasoning items (`rs_*` id) that can't be replayed on multi-turn
# 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 = os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or "high"
kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"})
_patch_openrouter_strip_responses_reasoning()
# Anthropic-routed providers → route through Anthropic provider with base_url
elif provider in _ANTHROPIC_ROUTED_PROVIDERS:
_original_provider = provider
base_url_default, api_key_env = _ANTHROPIC_ROUTED_PROVIDERS[provider]
if provider == "custom-anthropic":
base_url = os.environ.get("CUSTOM_ANTHROPIC_BASE_URL", "")
if not base_url:
raise ValueError(
"CUSTOM_ANTHROPIC_BASE_URL environment variable is required when using "
"the 'custom-anthropic' provider. Please set it to your "
"Anthropic-compatible API endpoint URL (e.g. https://api.anthropic.com)."
)
base_url = base_url.rstrip("/")
elif provider == "minimax":
base_url = os.environ.get("MINIMAX_BASE_URL", base_url_default).rstrip("/")
else:
base_url = base_url_default
if base_url:
kwargs.setdefault("base_url", base_url)
api_key = os.environ.get(api_key_env, "")
if api_key:
kwargs.setdefault("api_key", api_key)
# Kimi Coding Plan requires claude-code User-Agent header
if provider == "kimi-coding":
kwargs.setdefault("default_headers", {})["User-Agent"] = "claude-code/0.1.0"
provider = "anthropic"
elif provider == "ollama":
base_url = os.environ.get("OLLAMA_BASE_URL", "")
if base_url:
kwargs.setdefault("base_url", base_url)
_drop_unsupported_chat_model_kwargs(kwargs)
_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)
try:
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
finally:
if anthropic_auth_token is not None:
os.environ["ANTHROPIC_AUTH_TOKEN"] = anthropic_auth_token
# Flatten list content to strings for strict OpenAI-compatible providers
# (DeepSeek, SiliconFlow, OpenRouter, custom-openai, etc.) and
# native OpenAI through a proxy, to avoid "sequence expected string" errors.
# Moonshot and Kimi Coding support standard format, no patch needed.
_no_patch_providers = {"moonshot", "kimi-coding"}
if (
_is_third_party or _is_openai_proxy
) and _original_provider not in _no_patch_providers:
# 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)
# DeepSeek thinking mode requires reasoning_content passback in multi-turn
# + tool_use scenarios.
if _original_provider == "deepseek":
_patch_deepseek_reasoning_passback(chat_model)
if _is_openai_proxy:
_patch_ccproxy_system_to_developer(chat_model)
apply_known_context_window(chat_model)
return chat_model
def list_models() -> list[str]:
"""List all available model short names.
Returns:
List of unique model short names that can be passed to get_chat_model().
"""
seen = set()
result = []
for name, _, _ in _MODEL_ENTRIES:
if name not in seen:
seen.add(name)
result.append(name)
return result
def list_models_by_provider() -> list[tuple[str, str, str]]:
"""List all unique (short_name, model_id, provider) entries.
Returns:
De-duplicated list of model entries preserving registry order.
"""
seen: set[tuple[str, str]] = set()
result: list[tuple[str, str, str]] = []
for name, model_id, provider in _MODEL_ENTRIES:
key = (name, provider)
if key not in seen:
seen.add(key)
result.append((name, model_id, provider))
return result
async def list_model_picker_entries(
ollama_base_url: str | None,
*,
include_custom_ollama: bool,
) -> list[tuple[str, str, str]]:
"""Return model picker entries, optionally including local Ollama models."""
entries = list_models_by_provider()
if ollama_base_url:
from .ollama_discovery import discover_ollama_models
for detected_name in await discover_ollama_models(
ollama_base_url,
timeout=1.5,
):
entries.append((detected_name, detected_name, "ollama"))
if include_custom_ollama:
entries.append(("Custom Ollama model...", "__custom_ollama__", "ollama"))
return entries
def get_model_info(model: str) -> tuple[str, str] | None:
"""Get the (model_id, provider) tuple for a short name.
Args:
model: Short model name.
Returns:
Tuple of (model_id, provider) or None if not found.
"""
return MODELS.get(model)
+1 -1
View File
@@ -1,7 +1,7 @@
"""Ollama server probing — shared by onboard wizard and /model picker.
Ollama models are whatever the user has ``ollama pull``ed locally; they
cannot be enumerated in ``_MODEL_ENTRIES``. Both the setup wizard and the
cannot be enumerated in a static model catalog. Both the setup wizard and the
interactive model picker need to hit ``GET {base_url}/api/tags`` to see
what is actually installed.
+155 -136
View File
@@ -25,7 +25,6 @@ Utilities:
from __future__ import annotations
import hashlib
import os
from typing import Any
@@ -179,20 +178,15 @@ _patch_ccproxy_codex_compat()
# ---------------------------------------------------------------------------
# Utility: detect ccproxy's Codex adapter (as opposed to generic localhost).
# ---------------------------------------------------------------------------
def _is_ccproxy_codex(
base_url: str | None = None,
api_key: str | None = None,
) -> bool:
def _is_ccproxy_codex() -> bool:
"""Return True if the OpenAI endpoint is ccproxy's Codex adapter.
Checks for the ccproxy-specific markers set by ``setup_codex_env()``
in ``ccproxy_manager.py``: the sentinel API key and the ``/codex/v1``
path. Plain localhost endpoints (vLLM, Ollama, etc.) are not affected.
"""
if base_url is None:
base_url = os.environ.get("OPENAI_BASE_URL", "")
if api_key is None:
api_key = os.environ.get("OPENAI_API_KEY", "")
base_url = os.environ.get("OPENAI_BASE_URL", "")
api_key = os.environ.get("OPENAI_API_KEY", "")
return (
("127.0.0.1" in base_url or "localhost" in base_url)
and api_key == "ccproxy-oauth"
@@ -273,82 +267,6 @@ def _flatten_message_content(content: Any) -> str | list[Any] | Any:
return "\n\n".join(parts) if parts else ""
def _stable_tool_call_id(message: Any, message_index: int, call_index: int) -> str:
seed = ":".join(
(
str(getattr(message, "id", "") or "message"),
str(message_index),
str(call_index),
)
)
return "call_" + hashlib.sha256(seed.encode("utf-8")).hexdigest()[:24]
def _ensure_openai_tool_call_ids(messages: list[Any]) -> list[Any]:
"""Copy messages and repair missing AI/ToolMessage call identifiers."""
import copy
from collections import deque
pending_call_ids: deque[str] = deque()
normalized: list[Any] = []
for message_index, message in enumerate(messages):
message_type = getattr(message, "type", None)
if message_type == "ai":
tool_calls = list(getattr(message, "tool_calls", None) or [])
if not tool_calls:
normalized.append(message)
continue
copied = copy.copy(message)
normalized_calls: list[dict[str, Any]] = []
for call_index, original_call in enumerate(tool_calls):
call = dict(original_call)
call_id = str(call.get("id") or "") or _stable_tool_call_id(
message, message_index, call_index
)
call["id"] = call_id
normalized_calls.append(call)
pending_call_ids.append(call_id)
copied.tool_calls = normalized_calls
if isinstance(copied.content, list):
call_index = 0
blocks: list[Any] = []
for original_block in copied.content:
if not isinstance(original_block, dict):
blocks.append(original_block)
continue
block = dict(original_block)
if block.get("type") in {"tool_call", "function_call"}:
if call_index < len(normalized_calls):
block["id"] = normalized_calls[call_index]["id"]
call_index += 1
blocks.append(block)
copied.content = blocks
normalized.append(copied)
continue
if message_type == "tool":
tool_call_id = str(getattr(message, "tool_call_id", "") or "")
if tool_call_id:
try:
pending_call_ids.remove(tool_call_id)
except ValueError:
pass
normalized.append(message)
continue
if pending_call_ids:
copied = copy.copy(message)
copied.tool_call_id = pending_call_ids.popleft()
normalized.append(copied)
continue
normalized.append(message)
return normalized
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
"""Flatten list content for OpenAI-compatible APIs, preserving media.
@@ -364,7 +282,6 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
from langchain_core.messages import HumanMessage
messages = _ensure_openai_tool_call_ids(messages)
out: list[Any] = []
pending_media: list[Any] = [] # media hoisted out of a run of tool messages
@@ -956,26 +873,25 @@ def _patch_deepseek_reasoning_passback(model: Any) -> None:
# ---------------------------------------------------------------------------
# Patch: forward CLI's live (model, model_provider) into deepagents'
# start_async_task / update_async_task tool calls so the deployed graph
# (running in a separate ``langgraph dev`` subprocess) re-resolves the
# chat model per run.
# Patch: forward workspace scope and usage correlation context into
# deepagents' start_async_task / update_async_task tool calls so the deployed
# graph (running in a separate ``langgraph dev`` subprocess) inherits the
# parent run's scoping and accounting metadata.
#
# Without this, async sub-agents stay on the model their graph was compiled
# with at langgraph dev boot — `/model` switches in the CLI never reach
# them because they live in another process.
# Model configuration is deliberately NOT forwarded: the deployed graph
# resolves its chat model per run from ``configurable.runtime_snapshot_id``
# via ``ConfigurableModelMiddleware`` (design doc 8.2/8.3). Injecting
# ``model``/``model_provider`` here would be rejected with
# ``MODEL_CONFIG_OUTSIDE_SNAPSHOT``. Instead, a fresh local run snapshot
# bound to the child thread is created here (section 8.1: sub-agents are a
# local entry point sharing the same SnapshotService and snapshot table).
#
# Mechanism: wrap deepagents' ``_build_start_tool`` and ``_build_update_tool``
# factories. Each wrapped factory calls the original with a proxied client
# cache that intercepts ``runs.create(...)`` calls only and injects
# ``config={"configurable": {"model": <cfg.model>, "model_provider": <cfg.provider>}}``.
# All other client methods (``threads.create``, ``runs.get``, ``runs.cancel``,
# ``runs.join_stream``) pass through unchanged. The deployed graph picks up
# ``configurable.model`` via ``ConfigurableModelMiddleware``.
#
# Reads ``_ensure_config()`` at tool-call time (not patch time) so a
# ``/model`` switch in the CLI is reflected on the very next async tool
# call without an agent rebuild.
# cache that intercepts ``runs.create(...)`` calls only and merges the
# inherited scope/usage context into ``config``/``metadata``. All other
# client methods (``threads.create``, ``runs.get``, ``runs.cancel``,
# ``runs.join_stream``) pass through unchanged.
#
# Upstream PR opportunity: passing ``config`` through ``client.runs.create``
# is generic functionality; worth contributing back to ``langchain-ai/deepagents``
@@ -984,49 +900,152 @@ def _patch_deepseek_reasoning_passback(model: Any) -> None:
_model_passthrough_patched = False
def _read_cfg_configurable() -> dict[str, str]:
"""Read live ``(model, provider)`` from EvoScientist config.
Returns a dict suitable for inserting under
``RunnableConfig.configurable``. Empty dict on any failure (so the
patch degrades to a no-op rather than breaking async tool calls).
"""
try:
from EvoScientist.EvoScientist import _ensure_config
cfg = _ensure_config()
except Exception:
return {}
out: dict[str, str] = {}
model = getattr(cfg, "model", None)
provider = getattr(cfg, "provider", None)
if isinstance(model, str) and model:
out["model"] = model
if isinstance(provider, str) and provider:
out["model_provider"] = provider
return out
def _merge_runs_config_kwargs(kwargs: dict) -> dict:
"""Merge the live model override into ``kwargs`` for ``runs.create``.
"""Merge scope and usage correlation context into ``runs.create``.
Preserves any caller-supplied ``config.configurable`` keys. EvoScientist's
keys take precedence on conflict (callers shouldn't be passing model
overrides — the CLI is the source of truth).
Preserves any caller-supplied ``config.configurable`` keys.
"""
overrides = _read_cfg_configurable()
if not overrides:
return kwargs
usage_enabled = os.getenv("EVOSCIENTIST_USAGE_TRACKING", "").strip().lower() in {
"1",
"true",
"yes",
"on",
}
current_metadata: dict = {}
current_configurable: dict = {}
try:
from langgraph.config import get_config
# Scoped deployments need this context even when usage accounting is
# disabled: DeepAgents creates derived threads through this proxy.
current = get_config()
raw_metadata = current.get("metadata")
raw_configurable = current.get("configurable")
if isinstance(raw_metadata, dict):
current_metadata = raw_metadata
if isinstance(raw_configurable, dict):
current_configurable = raw_configurable
except (LookupError, RuntimeError):
pass
inherited_scope: dict[str, object] = {}
scope_keys = (
"workspace_scope_id",
"workspace_scope_owner_id",
"workspace_scope_revision",
"workspace_deployment_id",
)
if all(current_configurable.get(key) is not None for key in scope_keys):
inherited_scope = {key: current_configurable[key] for key in scope_keys}
target_thread_id = kwargs.get("thread_id")
source_thread_id = current_configurable.get("thread_id")
if (
isinstance(target_thread_id, str)
and target_thread_id
and target_thread_id != source_thread_id
):
try:
from EvoScientist.scope_registry import (
ScopeConflictError,
get_scope_registry,
)
registry = get_scope_registry()
try:
owner = registry.register_owner(
str(inherited_scope["workspace_deployment_id"]),
str(inherited_scope["workspace_scope_id"]),
owner_type="derived_run",
resource_id=target_thread_id,
parent_owner_id=str(
inherited_scope["workspace_scope_owner_id"]
),
state="active",
)
except ScopeConflictError:
owner = registry.get_owner_by_resource(
str(inherited_scope["workspace_deployment_id"]),
target_thread_id,
)
if owner.scope_id != inherited_scope["workspace_scope_id"]:
raise
inherited_scope["workspace_scope_owner_id"] = owner.owner_id
inherited_scope["thread_id"] = target_thread_id
except Exception:
# The required-mode backend factory rejects missing/invalid
# ownership at execution time. Optional mode keeps legacy
# integrations working when they cannot register a child.
if (
os.getenv("EVOSCIENTIST_WORKSPACE_ISOLATION", "optional").lower()
== "required"
):
raise
existing = kwargs.get("config")
if not isinstance(existing, dict):
existing = {}
existing_configurable = existing.get("configurable")
if not isinstance(existing_configurable, dict):
existing_configurable = {}
merged_configurable = {**existing_configurable, **overrides}
merged_configurable = {**existing_configurable, **inherited_scope}
# Async sub-agent snapshot entry (design doc 8.1): the deployed child
# graph runs on its own thread, so it cannot reuse the parent's
# snapshot binding. Freeze a fresh local snapshot bound to the child
# thread; the child's ConfigurableModelMiddleware resolves it per call.
# A bootstrap registry raises MODEL_REGISTRY_NOT_READY here, surfacing
# as a structured start/update_async_task tool error on the parent.
child_thread_id = kwargs.get("thread_id")
if (
isinstance(child_thread_id, str)
and child_thread_id
and "runtime_snapshot_id" not in merged_configurable
):
from EvoScientist.model_registry.runtime import get_snapshot_runtime
merged_configurable["runtime_snapshot_id"] = (
get_snapshot_runtime()
.create_local_snapshot(child_thread_id)
.snapshot_id
)
kwargs = dict(kwargs)
kwargs["config"] = {**existing, "configurable": merged_configurable}
if merged_configurable or "config" in kwargs:
kwargs["config"] = {**existing, "configurable": merged_configurable}
if not usage_enabled:
return kwargs
inherited_metadata = {
key: current_metadata[key]
for key in (
"usage_context_version",
"turn_id",
"source_session_id",
"source_agent",
"workspace_dir",
)
if current_metadata.get(key) is not None
}
inherited_metadata.setdefault(
"source_session_id",
current_metadata.get("thread_id")
or current_metadata.get("langgraph_thread_id")
or current_configurable.get("thread_id"),
)
if inherited_metadata.get("source_session_id") is None:
inherited_metadata.pop("source_session_id", None)
existing_metadata = kwargs.get("metadata")
if not isinstance(existing_metadata, dict):
existing_metadata = {}
merged_metadata = {**inherited_metadata, **existing_metadata}
run_kind = merged_metadata.get("run_kind")
if "usage_scope" not in merged_metadata:
if isinstance(run_kind, str) and run_kind.startswith("evomemory_"):
merged_metadata["usage_scope"] = "memory"
else:
merged_metadata["usage_scope"] = "async_subagent"
kwargs["metadata"] = merged_metadata
return kwargs
@@ -1100,7 +1119,7 @@ class _ClientCacheProxy:
def _patch_deepagents_model_passthrough() -> None:
"""Wrap deepagents' async-launch tool factories to inject CLI model.
"""Wrap deepagents' async-launch tool factories to inherit run context.
Idempotent: re-invocation is a no-op once the patch is active. Safe to
call from ``_maybe_swap_async_subagents`` on every CLI startup; both
-302
View File
@@ -1,302 +0,0 @@
"""Shared logging configuration helpers."""
from __future__ import annotations
import logging
import os
import sys
from datetime import UTC, datetime
from pathlib import Path
from typing import Any, TextIO
DEFAULT_LOG_RETENTION_DAYS = 30
DEFAULT_LOG_FORMAT = "%(asctime)s [%(levelname)s] %(name)s: %(message)s"
DEFAULT_LOG_DATE_FORMAT = "%Y-%m-%d %H:%M:%S"
MANAGED_HANDLER_ATTR = "_evoscientist_managed_handler"
def resolve_log_level(level: int | str | None, default: int = logging.INFO) -> int:
"""Resolve a logging level from config or environment input."""
if isinstance(level, int):
return level
raw = str(level or "").strip()
if not raw:
return default
if raw.isdigit():
return int(raw)
normalized = raw.upper()
if normalized == "WARN":
normalized = "WARNING"
resolved = logging.getLevelNamesMapping().get(normalized)
return resolved if isinstance(resolved, int) else default
def _mark_managed(handler: logging.Handler, kind: str) -> logging.Handler:
setattr(handler, MANAGED_HANDLER_ATTR, kind)
return handler
def _managed_kind(handler: logging.Handler) -> str | None:
kind = getattr(handler, MANAGED_HANDLER_ATTR, None)
return kind if isinstance(kind, str) else None
def remove_managed_handlers(
logger: logging.Logger | None = None,
*,
kinds: set[str] | None = None,
) -> None:
"""Remove handlers installed by this module without touching external ones."""
target = logger or logging.getLogger()
for handler in target.handlers[:]:
kind = _managed_kind(handler)
if kind and (kinds is None or kind in kinds):
target.removeHandler(handler)
handler.close()
def _standard_formatter() -> logging.Formatter:
return logging.Formatter(DEFAULT_LOG_FORMAT, datefmt=DEFAULT_LOG_DATE_FORMAT)
class DailyLogFileHandler(logging.FileHandler):
"""File handler that writes the active log to a date-based filename."""
def __init__(
self,
log_dir: str | Path,
*,
prefix: str = "evoscientist",
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
encoding: str = "utf-8",
utc: bool = False,
) -> None:
self.log_dir = Path(log_dir).expanduser()
self.prefix = prefix
self.retention_days = max(1, retention_days)
self.utc = utc
self.log_dir.mkdir(parents=True, exist_ok=True)
super().__init__(self._dated_log_path(), encoding=encoding, delay=True)
@property
def active_log_path(self) -> Path:
"""Return the active log path for the current date."""
return self._dated_log_path()
def _dated_log_path(self) -> Path:
now = datetime.now(UTC if self.utc else None)
return self.log_dir / f"{self.prefix}-{now:%Y-%m-%d}.log"
def emit(self, record: logging.LogRecord) -> None:
try:
expected = str(self.active_log_path)
if self.baseFilename != expected:
if self.stream:
self.stream.close()
self.stream = None
self.baseFilename = expected
self._delete_expired_logs()
super().emit(record)
except OSError:
self.handleError(record)
def getFilesToDelete(self) -> list[str]:
candidates = sorted(self.log_dir.glob(f"{self.prefix}-????-??-??.log"))
if len(candidates) <= self.retention_days:
return []
return [str(path) for path in candidates[: -self.retention_days]]
def _delete_expired_logs(self) -> None:
for path in self.getFilesToDelete():
try:
os.remove(path)
except OSError:
pass
def default_log_dir() -> Path:
"""Return the default runtime log directory."""
env_dir = os.environ.get("EVOSCIENTIST_LOG_DIR")
if env_dir:
return Path(env_dir).expanduser()
from EvoScientist.paths import DATA_DIR
return DATA_DIR / "logs"
def configure_daily_file_logging(
logger: logging.Logger | None = None,
*,
log_dir: str | Path | None = None,
prefix: str = "evoscientist",
level: int | str = logging.INFO,
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
) -> DailyLogFileHandler:
"""Attach a daily file handler, replacing older matching handlers."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
retention_days = max(1, int(retention_days))
resolved_dir = Path(log_dir).expanduser() if log_dir else default_log_dir()
for handler in target.handlers[:]:
if (
isinstance(handler, DailyLogFileHandler)
and handler.prefix == prefix
and handler.log_dir == resolved_dir
):
target.removeHandler(handler)
handler.close()
handler = DailyLogFileHandler(
resolved_dir,
prefix=prefix,
retention_days=retention_days,
)
_mark_managed(handler, "file")
handler.setLevel(resolved_level)
handler.setFormatter(_standard_formatter())
target.addHandler(handler)
if target.level == logging.NOTSET or target.level > resolved_level:
target.setLevel(resolved_level)
return handler
def configure_console_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
stream: TextIO | None = None,
replace: bool = True,
) -> logging.StreamHandler:
"""Attach a standard console handler for non-interactive entry points."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
if replace:
remove_managed_handlers(target, kinds={"console", "rich"})
handler = logging.StreamHandler(stream or sys.stderr)
_mark_managed(handler, "console")
handler.setLevel(resolved_level)
handler.setFormatter(_standard_formatter())
target.addHandler(handler)
target.setLevel(resolved_level)
return handler
def configure_rich_console_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
console: Any = None,
replace: bool = True,
dim_warnings: bool = False,
show_time: bool | None = None,
show_path: bool | None = None,
show_level: bool | None = None,
) -> logging.Handler:
"""Attach a Rich console handler for interactive CLI output."""
from rich.logging import RichHandler
from rich.markup import escape
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
verbose = resolved_level <= logging.DEBUG
if replace:
remove_managed_handlers(target, kinds={"console", "rich"})
class DimWarningHandler(RichHandler):
def emit(self, record: logging.LogRecord) -> None:
if dim_warnings and record.levelno == logging.WARNING and console is not None:
console.print(
"[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] "
f"[dim]{escape(record.getMessage())}[/dim]"
)
return
super().emit(record)
handler = DimWarningHandler(
console=console,
show_time=verbose if show_time is None else show_time,
show_path=verbose if show_path is None else show_path,
show_level=verbose if show_level is None else show_level,
)
_mark_managed(handler, "rich")
handler.setLevel(resolved_level)
target.addHandler(handler)
target.setLevel(resolved_level)
return handler
def configure_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
log_dir: str | Path | None = None,
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
prefix: str = "evoscientist",
console: bool = True,
file: bool = True,
replace_managed: bool = True,
) -> list[logging.Handler]:
"""Configure standard EvoScientist console and daily file logging."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
if replace_managed:
remove_managed_handlers(target, kinds={"console", "rich", "file"})
handlers: list[logging.Handler] = []
if console:
handlers.append(
configure_console_logging(target, level=resolved_level, replace=False)
)
if file:
handlers.append(
configure_daily_file_logging(
target,
log_dir=log_dir,
prefix=prefix,
level=resolved_level,
retention_days=retention_days,
)
)
target.setLevel(resolved_level)
return handlers
def configure_logging_from_settings(
logger: logging.Logger | None = None,
*,
default_level: int = logging.INFO,
prefix: str = "evoscientist",
console: bool = True,
file: bool = True,
) -> list[logging.Handler]:
"""Configure logging from EvoScientist settings and environment overrides."""
level: int | str | None = os.environ.get("EVOSCIENTIST_LOG_LEVEL")
log_dir: str | Path | None = os.environ.get("EVOSCIENTIST_LOG_DIR") or None
retention_days = int(
os.environ.get("EVOSCIENTIST_LOG_RETENTION_DAYS", DEFAULT_LOG_RETENTION_DAYS)
)
try:
from EvoScientist.config import get_effective_config
cfg = get_effective_config()
level = level or getattr(cfg, "log_level", None)
log_dir = log_dir or getattr(cfg, "log_dir", None) or None
retention_days = int(
getattr(cfg, "log_retention_days", DEFAULT_LOG_RETENTION_DAYS)
)
except Exception:
level = level or default_level
return configure_logging(
logger,
level=resolve_log_level(level, default=default_level),
log_dir=log_dir,
retention_days=retention_days,
prefix=prefix,
console=console,
file=file,
)
-2
View File
@@ -10,7 +10,6 @@ from .client import (
build_mcp_add_kwargs,
build_mcp_edit_fields,
edit_mcp_server,
get_mcp_server_errors,
load_mcp_config,
load_mcp_tools,
parse_mcp_add_args,
@@ -39,7 +38,6 @@ __all__ = [
"find_server_by_name",
"get_all_tags",
"get_installed_names",
"get_mcp_server_errors",
"install_mcp_server",
"install_mcp_servers",
"load_mcp_config",
+1 -16
View File
@@ -114,10 +114,6 @@ _URL_TRANSPORTS = {"http", "streamable_http", "sse", "websocket"}
# still parallelizing the common 3–7 server case to completion.
_MAX_CONCURRENT_CONNECTIONS = 8
# Last connection error per configured server. This is process-local runtime
# diagnostics for the Web/CLI status surfaces, not persisted configuration.
_MCP_SERVER_ERRORS: dict[str, str] = {}
# Env vars forwarded to stdio MCP subprocesses on top of the MCP SDK's
# minimal default set (HOME/PATH/USER/…). Without this, servers behind
# a proxy or with a custom CA bundle silently fail with long timeouts.
@@ -768,9 +764,6 @@ async def _load_tools(
if not connections:
return {}
for stale_name in set(_MCP_SERVER_ERRORS) - set(connections):
_MCP_SERVER_ERRORS.pop(stale_name, None)
client = MultiServerMCPClient(connections) # type: ignore[invalid-argument-type]
def _report(event: str, name: str, detail: str = "") -> None:
@@ -794,13 +787,10 @@ async def _load_tools(
_report("start", name)
try:
tools = await client.get_tools(server_name=name)
_MCP_SERVER_ERRORS.pop(name, None)
logger.info("MCP server %r: loaded %d tool(s)", name, len(tools))
_report("success", name, str(len(tools)))
return name, tools
except Exception as exc:
detail = str(exc) or type(exc).__name__
_MCP_SERVER_ERRORS[name] = detail
# When the caller wired up ``on_progress`` they own the
# user-facing display; downgrade the logger so we don't
# double-print.
@@ -808,7 +798,7 @@ async def _load_tools(
logger.warning("MCP server %r: failed to load tools: %s", name, exc)
else:
logger.debug("MCP server %r: failed to load tools: %s", name, exc)
_report("error", name, detail)
_report("error", name, str(exc))
return name, []
# ``return_exceptions=False`` is fine because ``_fetch`` already
@@ -817,11 +807,6 @@ async def _load_tools(
return dict(results)
def get_mcp_server_errors() -> dict[str, str]:
"""Return a snapshot of the most recent per-server connection errors."""
return dict(_MCP_SERVER_ERRORS)
async def aload_mcp_tools(
config: dict[str, Any] | None = None,
*,
+7 -3
View File
@@ -87,7 +87,8 @@ def build_memory_agent_graph(
from deepagents import create_deep_agent
from ...backends import build_memory_agent_backend
from ...EvoScientist import _ensure_auxiliary_chat_model
from ...EvoScientist import _compile_time_role_model
from ...middleware.configurable_model import ConfigurableModelMiddleware
kwargs: dict[str, Any] = {}
if response_format is not None:
@@ -101,11 +102,14 @@ def build_memory_agent_graph(
agent = create_deep_agent(
name=name,
model=_ensure_auxiliary_chat_model(),
model=_compile_time_role_model("primary"),
system_prompt=system_prompt,
tools=list(tools),
backend=backend,
middleware=list(middleware),
# The compile-time model is only a placeholder: the middleware
# re-resolves every call from the run snapshot (lazily healing
# background runs to a local snapshot of the registry default).
middleware=[ConfigurableModelMiddleware(), *middleware],
subagents=[],
skills=skills,
**kwargs,
+129 -10
View File
@@ -3,6 +3,8 @@
from __future__ import annotations
import json
import logging
import uuid
from collections.abc import Callable
from pathlib import Path
from typing import cast
@@ -17,6 +19,7 @@ from ..gateway.background_runs import (
launch_background_run,
)
from ..langgraph_dev.sdk import messages_input
from ..model_registry.schemas import ModelRef, ReasoningEffort
from .observations import build_observation_linker_index_context
from .scheduler import ObservationLinkerContext
from .source_context import MemorySourceContext, _trajectory_for_prompt
@@ -35,6 +38,8 @@ from .worker_activity import (
snapshot_observation_relations,
)
logger = logging.getLogger(__name__)
SUBAGENT_MEMORY_WORKER_GRAPH_ID = "evomemory-subagent-worker"
TURN_MEMORY_WORKER_GRAPH_ID = "evomemory-turn-worker"
OBSERVATION_LINKER_GRAPH_ID = "evomemory-observation-linker"
@@ -90,7 +95,7 @@ def _worker_workspace_dir(workspace_dir: str | Path) -> str:
def _memory_worker_metadata(context: MemorySourceContext) -> dict[str, str]:
return {
metadata = {
"run_kind": f"evomemory_{context.source_type.value}_worker",
"source_session_id": context.session_id,
"source_agent": context.source_agent,
@@ -98,6 +103,114 @@ def _memory_worker_metadata(context: MemorySourceContext) -> dict[str, str]:
"trajectory_digest": context.trajectory_digest,
"workspace_dir": _worker_workspace_dir(context.workspace_dir),
}
from ..usage.callback import usage_tracking_requested
if usage_tracking_requested():
metadata["usage_scope"] = "memory"
if usage_tracking_requested() and context.turn_id:
metadata["turn_id"] = context.turn_id
return metadata
def _source_thread_model_selection(
session_id: str,
) -> tuple[ModelRef, ReasoningEffort | None, int] | None:
"""Read the source conversation thread's explicit model selection.
Returns ``(primary, reasoning_effort, revision)`` for an explicit
ThreadModelSelection, or ``None`` when the thread is unreadable, has no
selection, or is set to ``inherit`` — in those cases the worker falls
back to the section 8.1 lazy local snapshot (registry default). A legacy
``auxiliary`` key in stored metadata is tolerated and dropped, matching
the BFF validator.
"""
from langgraph_sdk import get_sync_client
from ..langgraph_dev.sdk import (
configured_langgraph_dev_url,
langgraph_dev_headers,
)
client = get_sync_client(
url=configured_langgraph_dev_url(),
headers=langgraph_dev_headers(None),
)
thread = client.threads.get(session_id)
metadata = (
thread.get("metadata") if isinstance(thread, dict) else getattr(thread, "metadata", None)
)
if not isinstance(metadata, dict):
return None
raw = metadata.get("model_selection")
if not isinstance(raw, dict):
return None
primary_raw = raw.get("primary")
if not isinstance(primary_raw, dict):
return None
provider_id = primary_raw.get("provider_id")
model_key = primary_raw.get("model_key")
if (
not isinstance(provider_id, str)
or not provider_id
or not isinstance(model_key, str)
or not model_key
):
return None
effort_raw = raw.get("reasoning_effort")
effort: ReasoningEffort | None = (
effort_raw if effort_raw in ("low", "medium", "high") else None
)
revision_raw = metadata.get("model_selection_revision")
revision = revision_raw if isinstance(revision_raw, int) and revision_raw >= 0 else 0
return (ModelRef(provider_id=provider_id, model_key=model_key), effort, revision)
def _worker_snapshot_id(context: MemorySourceContext, worker_thread_id: str) -> str | None:
"""Freeze the source conversation's model selection for the worker run.
The worker runs on its own thread, so the conversation's snapshot cannot
be reused (bindings are per-thread); a fresh snapshot with the same
selection is created and bound to the worker thread instead. Returns
``None`` to keep the section 8.1 lazy-default behavior when the source
thread has no explicit selection or any step fails — memory work must
never fail just because selection inheritance did.
"""
try:
selection = _source_thread_model_selection(context.session_id)
except Exception:
logger.warning(
"Memory worker: could not read source thread %s model selection; "
"falling back to the registry default",
context.session_id,
exc_info=True,
)
return None
if selection is None:
return None
primary, reasoning_effort, revision = selection
from ..model_registry.runtime import get_snapshot_runtime
from ..model_registry.snapshots import SnapshotCreateRequest
try:
runtime = get_snapshot_runtime()
creation = runtime.snapshots.create( SnapshotCreateRequest(
run_request_id=uuid.uuid4().hex,
thread_id=worker_thread_id,
deployment_id=runtime.local_deployment_id,
model_selection_revision=revision,
primary=primary,
reasoning_effort=reasoning_effort,
)
)
except Exception:
logger.warning(
"Memory worker: could not freeze the source selection into a run "
"snapshot; falling back to the registry default",
exc_info=True,
)
return None
return creation.snapshot.snapshot_id
def _memory_worker_run_payload(
@@ -107,19 +220,25 @@ def _memory_worker_run_payload(
) -> BackgroundRunPayload:
"""Build the LangGraph SDK run payload for a memory worker."""
metadata = _memory_worker_metadata(context)
configurable = {
"thread_id": thread_id,
"evomemory_source_session_id": context.session_id,
"evomemory_source_agent": context.source_agent,
"evomemory_project_id": context.project_id,
"evomemory_trajectory_digest": context.trajectory_digest,
}
snapshot_id = _worker_snapshot_id(context, thread_id)
if snapshot_id is not None:
configurable["runtime_snapshot_id"] = snapshot_id
from ..usage.callback import usage_tracking_requested
if usage_tracking_requested() and context.turn_id:
configurable["evomemory_source_turn_id"] = context.turn_id
payload: BackgroundRunPayload = {
"assistant_id": _memory_worker_graph_id(context.source_type),
"input": messages_input(_memory_worker_user_prompt(context)),
"metadata": metadata,
"config": {
"configurable": {
"thread_id": thread_id,
"evomemory_source_session_id": context.session_id,
"evomemory_source_agent": context.source_agent,
"evomemory_project_id": context.project_id,
"evomemory_trajectory_digest": context.trajectory_digest,
}
},
"config": {"configurable": configurable},
}
return _runs_create_kwargs(payload)
+16
View File
@@ -40,6 +40,7 @@ class MemorySourceContext:
session_id: str
trajectory: list[CompactMessage]
trajectory_digest: str
turn_id: str | None = None
def _task_tool_call_ids(messages: list[BaseMessage]) -> set[str]:
@@ -203,6 +204,20 @@ def _runtime_thread_id(runtime: Runtime | None) -> str | None:
return None
def _active_turn_id() -> str | None:
"""Read the WebUI turn correlation value from the active run config."""
try:
from langgraph.config import get_config
metadata = get_config().get("metadata", {})
except (LookupError, RuntimeError):
return None
if not isinstance(metadata, dict):
return None
value = metadata.get("turn_id")
return value if isinstance(value, str) and value else None
def _short_hash(text: str) -> str:
"""Return the short hash fragment used in generated ids."""
return hashlib.sha256(text.encode("utf-8")).hexdigest()[:16]
@@ -240,6 +255,7 @@ def build_memory_source_context(
project_id=project_id,
source_agent=source_agent,
session_id=session_id,
turn_id=_active_turn_id(),
trajectory=trajectory,
trajectory_digest=_trajectory_digest(trajectory),
)
+13 -3
View File
@@ -27,7 +27,10 @@ from .memory_lifecycle import (
create_memory_lifecycle_middleware,
default_memory_scheduler,
)
from .model_fallback import ModelFallbackMiddleware, load_fallback_chain
from .message_budget import (
count_message_text_tokens,
create_message_budget_middleware,
)
from .runtime_context import RuntimeContextMiddleware, create_runtime_context_middleware
from .scheduler import (
SchedulerMiddleware,
@@ -37,6 +40,13 @@ from .tool_error_handler import ToolErrorHandlerMiddleware
from .tool_selector import create_tool_selector_middleware
from .utils import disable_thinking
# Patches the deepagents FilesystemMiddleware class so every read_file tool
# (main agent + sub-agents) post-processes image results. Must run before any
# create_deep_agent call; this package is imported on every agent-build path.
from .read_file_images import install_read_file_image_patch
install_read_file_image_patch()
__all__ = [
"AskUserMiddleware",
"AskUserRequest",
@@ -46,20 +56,20 @@ __all__ = [
"ContextOverflowMapperMiddleware",
"EvoMemoryLifecycleMiddleware",
"EvoMemoryMiddleware",
"ModelFallbackMiddleware",
"Question",
"RuntimeContextMiddleware",
"SchedulerMiddleware",
"ToolErrorHandlerMiddleware",
"compute_context_editing_trigger",
"count_message_text_tokens",
"create_code_interpreter_middleware",
"create_context_editing_middleware",
"create_memory_lifecycle_middleware",
"create_memory_middleware",
"create_message_budget_middleware",
"create_runtime_context_middleware",
"create_scheduler_middleware",
"create_tool_selector_middleware",
"default_memory_scheduler",
"disable_thinking",
"load_fallback_chain",
]
+96 -23
View File
@@ -29,13 +29,46 @@ from langchain.agents.middleware.types import (
)
from langchain.tools import InjectedToolCallId
from langchain_core.messages import AIMessage, SystemMessage, ToolMessage
from langchain_core.tools import tool
from langchain_core.tools import BaseTool, tool
from langgraph.types import Command, interrupt
from pydantic import BeforeValidator, Field
from typing_extensions import TypedDict
logger = logging.getLogger(__name__)
_VALID_REVIEW_MODES = frozenset({"manual", "auto", "full"})
def _review_mode() -> str:
"""Return the per-run WebUI review mode, failing closed to Manual."""
try:
from langgraph.config import get_config
config = get_config()
except Exception:
return "manual"
if not isinstance(config, dict):
return "manual"
configurable = config.get("configurable") or {}
if not isinstance(configurable, dict):
return "manual"
mode = configurable.get("review_mode")
return mode if mode in _VALID_REVIEW_MODES else "manual"
def _tool_name(tool_value: BaseTool | dict[str, Any]) -> str | None:
if isinstance(tool_value, BaseTool):
return tool_value.name or None
name = tool_value.get("name")
if isinstance(name, str) and name:
return name
function = tool_value.get("function")
if isinstance(function, dict):
nested_name = function.get("name")
if isinstance(nested_name, str) and nested_name:
return nested_name
return None
# ---------------------------------------------------------------------------
# Data types
@@ -202,6 +235,16 @@ or available tools.
- Never ask more than once per decision point — respect the user's time
- After receiving answers, summarize what you understood before proceeding"""
FULL_APPROVE_SYSTEM_PROMPT = """\
You are running in Full approve mode. Do not wait for user clarification. When
information is missing, make reasonable, conservative assumptions, continue the
task, and report material assumptions in the final response."""
FULL_APPROVE_TOOL_MESSAGE = (
"Full approve mode: continue with reasonable assumptions and do not ask "
"the user again."
)
# ---------------------------------------------------------------------------
# Validation & parsing
@@ -344,6 +387,17 @@ def _parse_answers(
# ---------------------------------------------------------------------------
def _with_system_prompt(request: ModelRequest[ContextT], prompt: str) -> SystemMessage:
if request.system_message is not None:
content = [
*request.system_message.content_blocks,
{"type": "text", "text": f"\n\n{prompt}"},
]
else:
content = [{"type": "text", "text": prompt}]
return SystemMessage(content=cast("list[str | dict[str, str]]", content))
class AskUserMiddleware(AgentMiddleware[Any, ContextT, ResponseT]):
"""Middleware that provides an ``ask_user`` tool for interactive questioning.
@@ -369,6 +423,17 @@ class AskUserMiddleware(AgentMiddleware[Any, ContextT, ResponseT]):
tool_call_id: Annotated[str, InjectedToolCallId],
) -> Command[Any]:
"""Ask the user one or more questions."""
if _review_mode() == "full":
return Command(
update={
"messages": [
ToolMessage(
FULL_APPROVE_TOOL_MESSAGE,
tool_call_id=tool_call_id,
)
]
}
)
_validate_questions(questions)
ask_request = AskUserRequest(
type="ask_user",
@@ -386,18 +451,22 @@ class AskUserMiddleware(AgentMiddleware[Any, ContextT, ResponseT]):
request: ModelRequest[ContextT],
handler: Callable[[ModelRequest[ContextT]], ModelResponse[ResponseT]],
) -> ModelResponse[ResponseT] | AIMessage:
"""Inject the ask_user system prompt."""
if request.system_message is not None:
new_system_content = [
*request.system_message.content_blocks,
{"type": "text", "text": f"\n\n{self.system_prompt}"},
]
else:
new_system_content = [{"type": "text", "text": self.system_prompt}]
new_system_message = SystemMessage(
content=cast("list[str | dict[str, str]]", new_system_content)
"""Apply the interactive or unattended prompt and tool policy."""
if _review_mode() == "full":
tools = [tool for tool in request.tools if _tool_name(tool) != "ask_user"]
return handler(
request.override(
system_message=_with_system_prompt(
request, FULL_APPROVE_SYSTEM_PROMPT
),
tools=tools,
)
)
return handler(
request.override(
system_message=_with_system_prompt(request, self.system_prompt)
)
)
return handler(request.override(system_message=new_system_message))
async def awrap_model_call(
self,
@@ -406,15 +475,19 @@ class AskUserMiddleware(AgentMiddleware[Any, ContextT, ResponseT]):
[ModelRequest[ContextT]], Awaitable[ModelResponse[ResponseT]]
],
) -> ModelResponse[ResponseT] | AIMessage:
"""Inject the ask_user system prompt (async)."""
if request.system_message is not None:
new_system_content = [
*request.system_message.content_blocks,
{"type": "text", "text": f"\n\n{self.system_prompt}"},
]
else:
new_system_content = [{"type": "text", "text": self.system_prompt}]
new_system_message = SystemMessage(
content=cast("list[str | dict[str, str]]", new_system_content)
"""Apply the interactive or unattended prompt and tool policy (async)."""
if _review_mode() == "full":
tools = [tool for tool in request.tools if _tool_name(tool) != "ask_user"]
return await handler(
request.override(
system_message=_with_system_prompt(
request, FULL_APPROVE_SYSTEM_PROMPT
),
tools=tools,
)
)
return await handler(
request.override(
system_message=_with_system_prompt(request, self.system_prompt)
)
)
return await handler(request.override(system_message=new_system_message))
+7 -3
View File
@@ -11,6 +11,7 @@ sub-agents are *tasks*, future cron is *schedules*).
from __future__ import annotations
import os
from datetime import UTC, datetime
from langchain.agents.middleware import AgentMiddleware
@@ -80,9 +81,12 @@ def run_in_background(
# apply_config_to_env round-trips at startup (and the subprocess inherits) —
# cheaper than reloading the full config from disk on every launch, and uses
# the same truthy parsing as every other bool env flag.
from ..llm.models import _env_flag_enabled
dangerous = _env_flag_enabled("EVOSCIENTIST_DANGEROUS_MODE")
dangerous = os.getenv("EVOSCIENTIST_DANGEROUS_MODE", "").strip().lower() in {
"1",
"true",
"yes",
"on",
}
# Same path-rewriting + validation as execute (shared helper) so virtual paths
# resolve to the workspace and the command can't bypass the sandbox checks.
command, error = prepare_sandbox_command(
+224 -121
View File
@@ -1,23 +1,27 @@
"""Middleware that resolves the chat model from RunnableConfig.configurable per call.
"""Middleware that resolves the per-call chat model from a run snapshot.
The deployed async sub-agents run in a separate ``langgraph dev`` subprocess
and have their model frozen into the compiled graph at subprocess boot time
(see ``EvoScientist/subagents/_factory.py``). When the user runs ``/model``
in the CLI, only the CLI process's model state changes — the subprocess
graph still uses the boot-time model.
Design doc 8.3: the middleware no longer reads ``config.yaml``, global
aliases, or ``model``/``model_provider`` overrides. The only model input a
run may carry is ``configurable["runtime_snapshot_id"]``; the middleware
loads the frozen snapshot (deployment/thread binding verified), resolves
the middleware's role through the section 6.1 role mapping, resolves the
credential against the frozen ``credential_revision``, and constructs the
model via ``build_chat_model`` with both safe HTTP clients.
This middleware closes that gap by reading ``model`` / ``model_provider``
from ``RunnableConfig.configurable`` on every model call. The CLI's patched
``start_async_task`` / ``update_async_task`` (see ``llm/patches.py``) injects
those fields into ``client.runs.create(config=...)``; the deployed graph
hits this middleware and re-resolves the chat model fresh.
``configurable`` carrying ``model``, ``model_provider``, or any other
out-of-snapshot model parameter is rejected with
``MODEL_CONFIG_OUTSIDE_SNAPSHOT`` (422 semantics, section 8.2) — run
creation must never mix snapshot and non-snapshot model configuration.
When ``configurable.model`` is absent, the middleware is a pass-through —
safe to install on the CLI's in-process agent too.
The middleware mirrors the pattern used by ``ModelFallbackMiddleware``:
``request.override(model=new_model)`` does not break tool binding, because
the downstream model-invocation node re-binds tools per request.
Local entry points (CLI/channels/scheduler/sub-agents, section 8.1) create
snapshots through ``SnapshotRuntime.create_local_snapshot`` and put the
``runtime_snapshot_id`` into ``configurable`` before the run starts. Runs
whose entry point could not inject a snapshot up front (langgraph-dev cron
fires, deployed async sub-agent graphs) are healed lazily: the middleware
creates the local snapshot itself, bound to the run's own thread. A run
carrying neither a snapshot nor a bindable thread ID is rejected, and a
bootstrap registry fails closed with ``MODEL_REGISTRY_NOT_READY`` — there
is no pass-through fallback to a compile-time model anymore.
**Reading the config**: ``Runtime`` (per its own docstring) does NOT include
``config``. The official path to reach ``RunnableConfig`` from inside any
@@ -32,8 +36,8 @@ from __future__ import annotations
import asyncio
import logging
import threading
from collections.abc import Awaitable, Callable
from typing import Any
from collections.abc import Awaitable, Callable, Mapping
from typing import Any, get_args
from langchain.agents.middleware.types import (
AgentMiddleware,
@@ -41,108 +45,202 @@ from langchain.agents.middleware.types import (
ModelResponse,
)
from ..model_registry.errors import (
MODEL_CONFIG_OUTSIDE_SNAPSHOT,
ModelRegistryError,
)
from ..model_registry.runtime import SnapshotRuntime, get_snapshot_runtime
from ..model_registry.schemas import ModelRole
from ..model_registry.snapshots import RuntimeSnapshot
logger = logging.getLogger(__name__)
_MODEL_ROLES = get_args(ModelRole)
def _read_model_override() -> tuple[str | None, str | None]:
"""Pull ``(model, model_provider)`` from the active ``RunnableConfig``.
# The only model-related keys a run's ``configurable`` may never carry: the
# snapshot is the single model configuration entry point (section 8.2).
_OUTSIDE_SNAPSHOT_MODEL_KEYS = ("model", "model_provider")
Reads via ``langgraph.config.get_config()`` (the documented entry point
for accessing the per-run ``RunnableConfig`` from inside any runnable
context — middleware, node, tool). Returns ``(None, None)`` when the
config has no ``configurable.model`` override or when called outside a
runnable context.
"""
def _current_configurable() -> Mapping[str, Any]:
"""Return the active run's ``configurable`` mapping (empty outside runs)."""
try:
from langgraph.config import get_config
cfg = get_config()
config = get_config()
except Exception:
# Outside a runnable context (most common in tests) or
# langgraph not importable — nothing to override.
return None, None
if not isinstance(cfg, dict):
return None, None
configurable = cfg.get("configurable") or {}
if not isinstance(configurable, dict):
return None, None
model = configurable.get("model")
provider = configurable.get("model_provider")
return (
model if isinstance(model, str) and model else None,
provider if isinstance(provider, str) and provider else None,
# Outside a runnable context (most common in tests) or langgraph not
# importable — treat as "no per-run configuration".
return {}
if not isinstance(config, Mapping):
return {}
configurable = config.get("configurable")
return configurable if isinstance(configurable, Mapping) else {}
def check_no_outside_snapshot_model_config(configurable: Mapping[str, Any]) -> None:
"""Reject any model configuration carried outside the run snapshot."""
for key in _OUTSIDE_SNAPSHOT_MODEL_KEYS:
if configurable.get(key) is not None:
raise ModelRegistryError(
MODEL_CONFIG_OUTSIDE_SNAPSHOT,
"Model configuration must come from the run snapshot only; "
f"configurable carries {key!r}.",
details=[{"path": key, "code": MODEL_CONFIG_OUTSIDE_SNAPSHOT}],
)
def read_snapshot_binding(
configurable: Mapping[str, Any],
) -> tuple[str, str | None] | None:
"""Extract ``(snapshot_id, thread_id)`` from configurable.
Returns ``None`` when the run carries no ``runtime_snapshot_id``. A
missing thread ID fails closed later because it can never match the
snapshot's binding. The deployment is deliberately NOT taken from
``workspace_deployment_id``: that key names the workspace-isolation
scope, not the snapshot issuer — issuer verification happens against
the platform-registered deployment set in ``get_snapshot_for_run``.
"""
snapshot_id = configurable.get("runtime_snapshot_id")
if not isinstance(snapshot_id, str) or not snapshot_id:
return None
thread_id = configurable.get("thread_id")
return (snapshot_id, thread_id if isinstance(thread_id, str) else None)
def ensure_snapshot_binding(
configurable: Mapping[str, Any],
runtime: SnapshotRuntime,
) -> tuple[str, str]:
"""Return ``(snapshot_id, thread_id)`` for the active run.
The run's explicit ``runtime_snapshot_id`` wins (section 8.2 binding).
Without one, the run is a local entry that could not inject a snapshot
up front — a langgraph-dev cron fire or a deployed sub-agent graph — so
the snapshot is created lazily through the same ``SnapshotService``
(section 8.1): ``deployment_id`` is the platform local deployment ID,
``model_selection_revision`` is ``0``, and ``primary`` inherits the
registry defaults. The lazy ``run_request_id`` is derived from the
thread, so every middleware in the run converges on one snapshot and a
fresh thread (each cron fire, each sub-agent task) freezes fresh
defaults.
Raises:
ModelRegistryError: ``MODEL_CONFIG_OUTSIDE_SNAPSHOT`` when the run
carries neither a snapshot nor a thread ID to bind one to;
``MODEL_REGISTRY_NOT_READY`` when the registry is in bootstrap.
"""
binding = read_snapshot_binding(configurable)
if binding is not None:
snapshot_id, thread_id = binding
return (snapshot_id, thread_id or "")
thread_id = configurable.get("thread_id")
if not isinstance(thread_id, str) or not thread_id:
raise ModelRegistryError(
MODEL_CONFIG_OUTSIDE_SNAPSHOT,
"Model configuration must come from a run snapshot; this run "
"carries neither 'runtime_snapshot_id' nor a 'thread_id' a "
"local snapshot could be bound to.",
details=[{"path": "runtime_snapshot_id", "code": MODEL_CONFIG_OUTSIDE_SNAPSHOT}],
)
snapshot = runtime.create_local_snapshot(
thread_id, run_request_id=f"auto:{thread_id}"
)
return (snapshot.snapshot_id, snapshot.thread_id)
class ConfigurableModelMiddleware(AgentMiddleware):
"""Re-resolve the chat model from RunnableConfig.configurable on every call.
"""Re-resolve the chat model from the run snapshot on every call.
Reads ``model`` and ``model_provider`` from the active ``RunnableConfig``
via ``langgraph.config.get_config()`` — the documented entry point for
accessing per-run config from any runnable context (middleware, node, tool).
When the override is present, calls
``EvoScientist.llm.get_chat_model(model=..., provider=...)`` and replaces
``request.model`` via ``request.override``. When absent, the middleware
passes through unchanged.
``role`` selects which frozen configuration of the snapshot feeds this
agent's model calls; every role currently maps to the snapshot's frozen
primary (section 6.1).
Note: ``Runtime`` (per its own docstring) does NOT include ``config`` as a
field — an earlier version of this middleware tried to read
``request.runtime.config`` and silently no-op'd because that attribute does
not exist. Stick with ``get_config()``.
A per-instance cache keyed by ``(model, provider)`` avoids rebuilding
identical models within a turn. The cache is a plain dict guarded by a
``threading.Lock`` because middleware instances are shared across
A per-instance cache keyed by snapshot ID avoids rebuilding the model on
every call within a run; snapshots are immutable once created, so the
cached instance stays valid for the run's lifetime. The cache is guarded
by a ``threading.Lock`` because middleware instances are shared across
concurrent requests in long-lived deployments (e.g. ``langgraph dev``
workers).
"""
name = "configurable_model"
def __init__(self) -> None:
def __init__(
self, role: ModelRole = "primary", runtime: SnapshotRuntime | None = None
) -> None:
super().__init__()
self._cache: dict[tuple[str, str | None], Any] = {}
if role not in _MODEL_ROLES:
raise ValueError(
f"Unknown model role {role!r}; expected one of {list(_MODEL_ROLES)}."
)
self._role = role
self._runtime = runtime
self._cache: dict[str, Any] = {}
self._lock = threading.Lock()
# Track the last (model, provider) pair we INFO-logged so we only
# surface a banner on transition. Without this, every LLM call in a
# long async run would emit an identical INFO line.
self._last_logged_key: tuple[str, str | None] | None = None
# Track the last snapshot we INFO-logged so we only surface a banner
# on transition. Without this, every LLM call in a long run would
# emit an identical INFO line.
self._last_logged_snapshot_id: str | None = None
def _log_override(self, model_name: str, provider: str | None) -> None:
"""INFO on transition; DEBUG on subsequent calls with same key."""
key = (model_name, provider)
def _snapshot_runtime(self) -> SnapshotRuntime:
return self._runtime if self._runtime is not None else get_snapshot_runtime()
def _log_override(self, snapshot: RuntimeSnapshot) -> None:
"""INFO on transition; DEBUG on subsequent calls of the same snapshot."""
from ..model_registry.snapshots import config_for_role
config = config_for_role(snapshot, self._role)
with self._lock:
transitioned = key != self._last_logged_key
transitioned = snapshot.snapshot_id != self._last_logged_snapshot_id
if transitioned:
self._last_logged_key = key
self._last_logged_snapshot_id = snapshot.snapshot_id
message_args = (
self._role,
config.model_ref.provider_id,
config.model_ref.model_key,
snapshot.snapshot_id,
)
if transitioned:
logger.info(
"ConfigurableModelMiddleware: overriding model to %s (%s)",
model_name,
provider,
"ConfigurableModelMiddleware: role %s bound to %s/%s (snapshot %s)",
*message_args,
)
else:
logger.debug(
"ConfigurableModelMiddleware: reusing override model=%s provider=%s",
model_name,
provider,
"ConfigurableModelMiddleware: role %s reusing %s/%s (snapshot %s)",
*message_args,
)
def _resolve(self, model: str, provider: str | None) -> Any:
"""Return a cached or freshly-built chat model for ``(model, provider)``."""
key = (model, provider)
def _load_snapshot(self) -> RuntimeSnapshot:
"""Load the run's snapshot, verifying its deployment/thread binding.
The snapshot is the only model configuration entry point: outside
model keys are rejected first, and a run without an explicit
``runtime_snapshot_id`` gets a lazily created local snapshot bound
to its own thread (never a silent compile-time fallback).
"""
configurable = _current_configurable()
check_no_outside_snapshot_model_config(configurable)
runtime = self._snapshot_runtime()
snapshot_id, thread_id = ensure_snapshot_binding(configurable, runtime)
return runtime.get_snapshot_for_run(snapshot_id, thread_id=thread_id)
def _resolve(self) -> Any:
"""Return a cached or freshly-built chat model for the run's snapshot."""
snapshot = self._load_snapshot()
with self._lock:
cached = self._cache.get(key)
cached = self._cache.get(snapshot.snapshot_id)
if cached is not None:
return cached
# Build outside the lock (network/SDK init can be slow); two
# concurrent first-time misses for the same key may build twice but
# the second result simply overwrites the first — both are equivalent.
from ..llm import get_chat_model
new_model = get_chat_model(model=model, provider=provider)
# Build outside the lock (SDK init can be slow); two concurrent
# first-time misses for the same snapshot may build twice but the
# second result simply overwrites the first — both are equivalent.
new_model = self._snapshot_runtime().build_role_model(snapshot, self._role)
with self._lock:
self._cache[key] = new_model
self._cache[snapshot.snapshot_id] = new_model
self._log_override(snapshot)
return new_model
def wrap_model_call(
@@ -150,21 +248,7 @@ class ConfigurableModelMiddleware(AgentMiddleware):
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
model_name, provider = _read_model_override()
if model_name is None:
return handler(request)
try:
new_model = self._resolve(model_name, provider)
except Exception:
logger.warning(
"ConfigurableModelMiddleware failed to resolve model=%r "
"provider=%r; falling back to compile-time model",
model_name,
provider,
exc_info=True,
)
return handler(request)
self._log_override(model_name, provider)
new_model = self._resolve()
return handler(request.override(model=new_model))
async def awrap_model_call(
@@ -172,25 +256,44 @@ class ConfigurableModelMiddleware(AgentMiddleware):
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
model_name, provider = _read_model_override()
if model_name is None:
return await handler(request)
try:
# Offload first-call SDK init off the event loop. ``_resolve`` calls
# ``get_chat_model`` on a cache miss, which can spend hundreds of ms
# building HTTP clients. Doing this synchronously inside an
# ``async def`` would block every other coroutine on the same
# langgraph dev event loop. Cache hits are still fast (a dict
# lookup); the thread-pool overhead is irrelevant once warm.
new_model = await asyncio.to_thread(self._resolve, model_name, provider)
except Exception:
logger.warning(
"ConfigurableModelMiddleware failed to resolve model=%r "
"provider=%r; falling back to compile-time model",
model_name,
provider,
exc_info=True,
)
return await handler(request)
self._log_override(model_name, provider)
# Offload first-call SDK init off the event loop: ``_resolve`` reads
# SQLite and can spend hundreds of ms building HTTP clients on a
# cache miss, which would block every other coroutine on the same
# langgraph dev event loop. Cache hits are still fast; the
# thread-pool overhead is irrelevant once warm.
new_model = await asyncio.to_thread(self._resolve)
return await handler(request.override(model=new_model))
# --- per-call resolution for in-run helper models -----------------------------
# In-run helper LLM calls (the tool selector's internal selection call) do
# not pass through an agent's model request, so ``request.override`` cannot
# reach them. They resolve the run snapshot's primary model through this
# helper instead; background runs without an explicit snapshot are healed
# through the same lazy local-snapshot path as the middleware.
_helper_model_cache: dict[str, Any] = {}
_helper_model_cache_lock = threading.Lock()
def resolve_snapshot_model(runtime: SnapshotRuntime | None = None) -> Any:
"""Resolve the active run's snapshot primary model (cached per snapshot).
Uses the same binding rules as ``ConfigurableModelMiddleware``: the
run's explicit ``runtime_snapshot_id`` wins, otherwise a local snapshot
is lazily created for the run's own thread. Must be called from inside
a runnable context (or a test that installed one).
"""
configurable = _current_configurable()
check_no_outside_snapshot_model_config(configurable)
rt = runtime if runtime is not None else get_snapshot_runtime()
snapshot_id, thread_id = ensure_snapshot_binding(configurable, rt)
snapshot = rt.get_snapshot_for_run(snapshot_id, thread_id=thread_id)
with _helper_model_cache_lock:
cached = _helper_model_cache.get(snapshot.snapshot_id)
if cached is not None:
return cached
model = rt.build_role_model(snapshot, "primary")
with _helper_model_cache_lock:
_helper_model_cache[snapshot.snapshot_id] = model
return model
+462
View File
@@ -0,0 +1,462 @@
"""Message-only context budgeting and automatic conversation compaction.
The middleware deliberately does not attempt to count the whole provider
request. System instructions, memories, tools, and attachments instead
consume the fixed reserves frozen into the run snapshot. The measured value
is only textual conversation messages and tool-result text, which is the
part compaction can actually reduce.
Snapshot mode (design doc 6.5, 8.3): every run carries — or lazily creates,
via the section 8.1 local entry convention — a run snapshot; the input
limit and the three fixed reserves come from the frozen
``ResolvedModelConfig.budget`` of the middleware's own ``snapshot_role``
(every role maps to the snapshot's frozen primary, section 6.1). There is
no 32K default and no compile-time profile fallback: a bootstrap registry
fails closed with ``MODEL_REGISTRY_NOT_READY``. The per-call message budget
is recomputed on every invocation from those frozen reserves with the
current ``has_tools``/``has_attachments`` mode; it is never carried over
from a previous call:
message_budget = resolved_input_limit
- fixed_system_reserve_tokens
- (has_tools ? fixed_tools_reserve_tokens : 0)
- (has_attachments ? fixed_attachments_reserve_tokens : 0)
``has_tools`` follows the conservative rule: it is decided once at agent
construction (the agent has tool capability and a configured toolset), not
by counting the tools bound to a single call, so it does not depend on
tool-selector middleware ordering. ``has_attachments`` is true only when
the current messages carry file/image/other attachment blocks. The frozen
base ``budget.message_budget`` (no tools, no attachments) is never reused
directly.
Because ``count_message_text_tokens`` is a conservative character estimate,
the effective hard budget is ``message_budget x 0.90`` and the soft trigger
scales proportionally (``x 0.70``), replacing the removed ``safety_reserve``
(section 6.5). Budget satisfiability against ``min_effective_input_tokens``
is guaranteed at snapshot creation (run-creation stage, all four
tool/attachment modes), so the middleware does not re-check it per call.
"""
from __future__ import annotations
import asyncio
import logging
import threading
from collections.abc import Iterable, Mapping
from contextvars import ContextVar
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, get_args
from langchain_core.messages import AnyMessage, SystemMessage
from ..model_registry.schemas import ModelRole
from .configurable_model import _current_configurable, ensure_snapshot_binding
if TYPE_CHECKING:
from ..model_registry.runtime import SnapshotRuntime
from ..model_registry.snapshots import RuntimeSnapshot
# Character-estimation budget scaling (section 6.5). The conservative char
# counter replaces the removed ``safety_reserve`` deduction.
_ESTIMATE_HARD_FRACTION = 0.90
_SOFT_FRACTION = 0.70
_KEEP_FRACTION = 0.35
_MODEL_ROLES = get_args(ModelRole)
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class MessageBudget:
"""The message budget selected for one model invocation."""
input_limit: int
hard_tokens: int
soft_tokens: int
keep_tokens: int
has_tools: bool
has_attachments: bool
# Fixed reserves (system + tools + attachments) consumed by this call;
# used_tokens + reserved_tokens approximates total context occupancy.
reserved_tokens: int = 0
_ACTIVE_BUDGET: ContextVar[MessageBudget | None] = ContextVar(
"evoscientist_message_budget", default=None
)
# The message-text total the base class computed just before deciding to
# summarize; stashed per call so a freshly created ``_summarization_event``
# can record the pre-compression estimate alongside the post-compression one.
_LAST_TOTAL_TOKENS: ContextVar[int | None] = ContextVar(
"evoscientist_message_budget_total", default=None
)
def _augment_summarization_event(
response: Any, budget: MessageBudget, tokens_after: list[int]
) -> None:
"""Attach estimate token stats to a freshly created summarization event.
``count_message_text_tokens`` is a character estimate, so the fields are
named ``estimated_*``; the provider-measured sizes keep flowing through
the usage pipeline (``input_tokens`` on the next model call reflects the
compressed prompt). Responses without a new ``_summarization_event``
update pass through untouched.
"""
command = getattr(response, "command", None)
update = getattr(command, "update", None)
if not isinstance(update, dict):
return
event = update.get("_summarization_event")
if not isinstance(event, dict):
return
before = _LAST_TOTAL_TOKENS.get()
if before is not None:
event["estimated_tokens_before"] = before
if tokens_after:
event["estimated_tokens_after"] = tokens_after[0]
event["budget"] = {
"hard_tokens": budget.hard_tokens,
"soft_tokens": budget.soft_tokens,
"keep_tokens": budget.keep_tokens,
}
def _content_text(content: Any) -> str:
"""Extract text without serializing image/file blocks or tool schemas."""
if isinstance(content, str):
return content
if not isinstance(content, list):
return ""
parts: list[str] = []
for block in content:
if isinstance(block, str):
parts.append(block)
continue
if not isinstance(block, Mapping):
continue
block_type = block.get("type")
if block_type in {"image", "image_url", "file", "document", "audio", "video"}:
continue
text = block.get("text")
if isinstance(text, str):
parts.append(text)
return "\n".join(parts)
def count_message_text_tokens(messages: Iterable[AnyMessage]) -> int:
"""Conservative, provider-neutral token estimate for compactable text only."""
characters = sum(len(_content_text(message.content)) for message in messages)
# A four-character estimate deliberately errs slightly high for Chinese and
# mixed code while remaining cheap enough to run before every model call.
return max(0, (characters + 3) // 4)
def _has_attachments(messages: Iterable[AnyMessage]) -> bool:
attachment_types = {"image", "image_url", "file", "document", "audio", "video"}
for message in messages:
if not isinstance(message.content, list):
continue
for block in message.content:
if isinstance(block, Mapping) and block.get("type") in attachment_types:
return True
return False
def _snapshot_message_budget(
snapshot: RuntimeSnapshot,
role: ModelRole,
*,
has_tools: bool,
has_attachments: bool,
) -> MessageBudget:
"""Recompute the section 6.5 per-call budget from the frozen reserves."""
from ..model_registry.snapshots import config_for_role
budget = config_for_role(snapshot, role).budget
reserves = budget.fixed_reserves
message_budget = (
budget.resolved_input_limit
- reserves.fixed_system_reserve_tokens
- (reserves.fixed_tools_reserve_tokens if has_tools else 0)
- (reserves.fixed_attachments_reserve_tokens if has_attachments else 0)
)
hard = max(1, int(message_budget * _ESTIMATE_HARD_FRACTION))
return MessageBudget(
input_limit=budget.resolved_input_limit,
hard_tokens=hard,
soft_tokens=max(1, int(hard * _SOFT_FRACTION)),
keep_tokens=max(1, int(hard * _KEEP_FRACTION)),
has_tools=has_tools,
has_attachments=has_attachments,
reserved_tokens=budget.resolved_input_limit - message_budget,
)
def _emit_context_usage(budget: MessageBudget, used_tokens: int) -> None:
"""Stream the prompt size of the model call about to happen.
Emitted on the LangGraph ``custom`` stream mode so the UI can render a
live context-occupancy indicator; ``used_tokens`` is the character
estimate of the messages actually sent (provider-measured sizes keep
flowing through the usage pipeline). Outside a runnable context (unit
tests, sync drivers) there is no stream writer — skip silently.
"""
try:
from langgraph.config import get_stream_writer
writer = get_stream_writer()
except Exception:
return
try:
writer(
{
"type": "evoscientist_context_usage",
"used_tokens": used_tokens,
"reserved_tokens": budget.reserved_tokens,
"input_limit": budget.input_limit,
"hard_tokens": budget.hard_tokens,
"soft_tokens": budget.soft_tokens,
}
)
except Exception:
logger.warning("Failed to emit context-usage stream event", exc_info=True)
class MessageBudgetMiddleware:
"""Factory namespace kept separate from DeepAgents' concrete middleware."""
@staticmethod
def create(
model: Any,
backend: Any,
*,
has_tools: bool = True,
snapshot_role: ModelRole = "primary",
runtime: SnapshotRuntime | None = None,
):
"""Create the runtime-aware DeepAgents summarization middleware.
Args:
model: Compile-time placeholder model, required by the base
summarization middleware; every per-call budget and the
summarizer model are resolved from the run snapshot.
backend: Graph backend forwarded to the summarization middleware.
has_tools: Conservative section 6.5 tool-mode flag: true when the
agent has tool capability and a configured toolset. Decided
once here, never from a single request's bound tool count.
snapshot_role: The snapshot role whose frozen limits size this
agent's budget; every role maps to the snapshot's frozen
primary (section 6.1). The summarizer uses the same frozen
primary.
runtime: Snapshot runtime override; defaults to the shared
process runtime (tests inject an isolated one).
"""
from deepagents.middleware.summarization import SummarizationMiddleware
if snapshot_role not in _MODEL_ROLES:
raise ValueError(
f"Unknown model role {snapshot_role!r}; "
f"expected one of {list(_MODEL_ROLES)}."
)
class _RuntimeMessageBudgetMiddleware(SummarizationMiddleware):
name = "message_budget"
def __init__(self) -> None:
self._snapshot_model_cache: dict[str, Any] = {}
self._snapshot_cache: dict[str, RuntimeSnapshot] = {}
self._model_cache_lock = threading.RLock()
self._has_tools = has_tools
self._snapshot_role = snapshot_role
self._runtime = runtime
# Triggering and cutoff are overridden below. The base class is
# still used for safe AI/tool-pair handling, offloading, and
# persisted summarization events.
super().__init__(
model=model,
backend=backend,
trigger=("tokens", 1_000_000_000),
keep=("messages", 8),
token_counter=count_message_text_tokens,
trim_tokens_to_summarize=4_000,
truncate_args_settings={
"trigger": ("tokens", 2_048),
"keep": ("messages", 8),
"max_length": 2_000,
"truncation_text": "...(tool arguments compacted)",
},
)
def _snapshot_runtime(self) -> SnapshotRuntime:
if self._runtime is not None:
return self._runtime
from ..model_registry.runtime import get_snapshot_runtime
return get_snapshot_runtime()
def _snapshot(self) -> RuntimeSnapshot:
"""Load the run's snapshot, verifying its binding.
A run without an explicit ``runtime_snapshot_id`` gets a
lazily created local snapshot bound to its own thread
(section 8.1); a bootstrap registry fails closed with
``MODEL_REGISTRY_NOT_READY`` instead of a 32K fallback.
"""
runtime = self._snapshot_runtime()
snapshot_id, thread_id = ensure_snapshot_binding(
_current_configurable(), runtime
)
with self._model_cache_lock:
cached = self._snapshot_cache.get(snapshot_id)
if cached is not None:
# Snapshots are immutable; the per-run read already
# verified the binding, and repeat reads must not hit
# SQLite from the event loop (e.g. the ``model``
# property during async summarization).
return cached
snapshot = runtime.get_snapshot_for_run(
snapshot_id, thread_id=thread_id
)
with self._model_cache_lock:
self._snapshot_cache[snapshot_id] = snapshot
return snapshot
@property
def model(self) -> Any: # type: ignore[override]
"""Summarizer model: the snapshot's frozen primary."""
snapshot = self._snapshot()
with self._model_cache_lock:
cached = self._snapshot_model_cache.get(snapshot.snapshot_id)
if cached is not None:
return cached
resolved = self._snapshot_runtime().build_role_model(
snapshot, "primary"
)
with self._model_cache_lock:
self._snapshot_model_cache[snapshot.snapshot_id] = resolved
return resolved
def _budget_for_request(self, request: Any) -> MessageBudget:
has_attachments = _has_attachments(getattr(request, "messages", []))
return _snapshot_message_budget(
self._snapshot(),
self._snapshot_role,
has_tools=self._has_tools,
has_attachments=has_attachments,
)
def _active_budget(self) -> MessageBudget:
active = _ACTIVE_BUDGET.get()
if active is None:
raise RuntimeError(
"message budget accessed outside a model call; "
"wrap_model_call establishes the per-call budget."
)
return active
def wrap_model_call(self, request: Any, handler: Any) -> Any:
budget = self._budget_for_request(request)
token = _ACTIVE_BUDGET.set(budget)
tokens_after: list[int] = []
def counting_handler(modified_request: Any) -> Any:
# The last handler invocation carries the exact prompt
# sent to the model (summary + preserved tail after
# compaction); earlier attempts are superseded.
tokens_after[:] = [
count_message_text_tokens(modified_request.messages)
]
return handler(modified_request)
try:
response = super().wrap_model_call(request, counting_handler)
# Emitted from the middleware body, not the handler: the
# base class may invoke the handler on an executor thread
# where the LangGraph stream-writer context is absent.
if tokens_after:
_emit_context_usage(budget, tokens_after[0])
_augment_summarization_event(response, budget, tokens_after)
return response
finally:
_ACTIVE_BUDGET.reset(token)
async def awrap_model_call(self, request: Any, handler: Any) -> Any:
# The snapshot read hits SQLite — blocking I/O that must stay
# off the event loop (langgraph dev's blockbuster rejects it).
budget = await asyncio.to_thread(self._budget_for_request, request)
token = _ACTIVE_BUDGET.set(budget)
tokens_after: list[int] = []
async def counting_handler(modified_request: Any) -> Any:
tokens_after[:] = [
count_message_text_tokens(modified_request.messages)
]
return await handler(modified_request)
try:
response = await super().awrap_model_call(
request, counting_handler
)
if tokens_after:
_emit_context_usage(budget, tokens_after[0])
_augment_summarization_event(response, budget, tokens_after)
return response
finally:
_ACTIVE_BUDGET.reset(token)
def _get_profile_limits(self) -> int | None:
return self._active_budget().hard_tokens
def _count_tokens(
self,
messages: list[AnyMessage],
system_message: SystemMessage | None,
tools: list[Any] | None,
) -> int:
del system_message, tools
return count_message_text_tokens(messages)
def _should_summarize(
self, messages: list[AnyMessage], total_tokens: int
) -> bool:
del messages
_LAST_TOTAL_TOKENS.set(total_tokens)
return total_tokens >= self._active_budget().soft_tokens
def _determine_cutoff_index(self, messages: list[AnyMessage]) -> int:
keep_tokens = self._active_budget().keep_tokens
retained = 0
cutoff = len(messages)
for index in range(len(messages) - 1, -1, -1):
message_tokens = count_message_text_tokens([messages[index]])
if retained + message_tokens > keep_tokens:
# Keep at least the newest message even when it alone
# exceeds the retention budget. The safe-cutoff helper
# below expands that boundary when it is a tool result.
cutoff = index if retained == 0 else index + 1
break
retained += message_tokens
cutoff = index
if cutoff <= 0:
return 0
return self._lc_helper._find_safe_cutoff_point(messages, cutoff)
return _RuntimeMessageBudgetMiddleware()
def create_message_budget_middleware(
model: Any,
backend: Any,
*,
has_tools: bool = True,
snapshot_role: ModelRole = "primary",
runtime: SnapshotRuntime | None = None,
):
"""Construct automatic compaction middleware for a graph backend."""
return MessageBudgetMiddleware.create(
model, backend, has_tools=has_tools, snapshot_role=snapshot_role, runtime=runtime
)
-378
View File
@@ -1,378 +0,0 @@
"""Middleware that implements model fallback on LLM call failures.
Uses LangChain's AgentMiddleware to intercept model calls. When the primary
model raises an exception, the middleware walks the configured fallback chain,
trying each alternative model in order. Every fallback attempt and its
outcome is surfaced to the user via the registered UI callback.
Errors that indicate a client-side bug (malformed request / HTTP 400) or a
context-length breach are not eligible for fallback and are re-raised
immediately so the correct handler (user or ContextOverflowMapperMiddleware)
can deal with them.
"""
from __future__ import annotations
import logging
import threading
from collections.abc import Awaitable, Callable
from langchain.agents.middleware.types import (
AgentMiddleware,
ModelRequest,
ModelResponse,
)
logger = logging.getLogger(__name__)
_ui_emit_fn: Callable[[str, str], None] | None = None
"""UI callback registered by the CLI/TUI entrypoint. ``None`` until set."""
_fallback_chain_lock = threading.Lock()
_fallback_chain: list[tuple[str, str]] = []
"""Ordered list of ``(model_name, provider)`` fallback entries."""
_CONTEXT_LIMIT_PATTERNS: list[str] = [
"context_length_exceeded",
"context length exceeded",
"too many tokens",
"maximum context length",
"output too large",
"context_window_exceeded",
"string_too_long",
"max_tokens_exceeded",
]
"""Substrings that identify a context-length error in provider messages."""
_MALFORMED_REQUEST_PATTERNS: list[str] = [
"invalid_request_error",
"invalid request",
"malformed",
]
"""Substrings that identify a malformed request (client-side bug)."""
_AUTH_ERROR_PATTERNS: list[str] = [
"invalid_api_key",
"authentication",
"permission",
]
"""Substrings that identify auth/permission errors.
These are intentionally *not* treated as non-fallbackable because a different
provider in the chain may have valid credentials."""
def set_ui_emit(fn: Callable[[str, str], None] | None) -> None:
"""Register (or clear) the UI callback for fallback status messages.
Args:
fn: Callable with signature ``fn(text, style)`` where *style* is a
Rich style string (``"yellow"``, ``"red"``, ``"green"``).
Pass ``None`` to unregister.
"""
global _ui_emit_fn
_ui_emit_fn = fn
def _emit(text: str, style: str = "yellow") -> None:
"""Surface a fallback status message to the user.
Dispatches to the registered UI callback when available (TUI mode),
otherwise falls back to the shared Rich console on stdout (CLI mode).
Args:
text: Plain-text message to display.
style: Rich style string applied to the message.
"""
if _ui_emit_fn is not None:
try:
_ui_emit_fn(text, style)
return
except Exception:
pass
from ..stream.console import console
console.print(text, style=style)
def get_fallback_chain() -> list[tuple[str, str]]:
"""Return a snapshot of the current fallback chain.
Returns:
List of ``(model_name, provider)`` tuples in priority order.
"""
with _fallback_chain_lock:
return list(_fallback_chain)
def set_fallback_chain(chain: list[tuple[str, str]]) -> None:
"""Replace the entire fallback chain.
Args:
chain: New list of ``(model_name, provider)`` tuples.
"""
global _fallback_chain
with _fallback_chain_lock:
_fallback_chain = list(chain)
def add_fallback(model: str, provider: str) -> bool:
"""Append a model to the end of the fallback chain.
Args:
model: Short model name (e.g. ``"gpt-5.5"``).
provider: Provider identifier (e.g. ``"openai"``).
Returns:
``True`` if added, ``False`` if the entry was already present.
"""
entry = (model, provider)
with _fallback_chain_lock:
if entry in _fallback_chain:
return False
_fallback_chain.append(entry)
return True
def remove_fallback(model: str) -> bool:
"""Remove all entries matching *model* regardless of provider.
Args:
model: Short model name to remove.
Returns:
``True`` if at least one entry was removed.
"""
global _fallback_chain
with _fallback_chain_lock:
before = len(_fallback_chain)
_fallback_chain = [(m, p) for m, p in _fallback_chain if m != model]
return len(_fallback_chain) < before
def remove_fallback_at(index: int) -> tuple[str, str] | None:
"""Remove the entry at a 0-based index.
Args:
index: Position in the chain (0-based).
Returns:
The removed ``(model, provider)`` tuple, or ``None`` if out of range.
"""
with _fallback_chain_lock:
if 0 <= index < len(_fallback_chain):
return _fallback_chain.pop(index)
return None
def clear_fallbacks() -> None:
"""Remove every entry from the fallback chain."""
global _fallback_chain
with _fallback_chain_lock:
_fallback_chain = []
def serialize_fallback_chain() -> str:
"""Serialize the chain to a config-friendly string.
Returns:
Comma-separated ``"model:provider,model:provider"`` string.
"""
with _fallback_chain_lock:
return ",".join(f"{m}:{p}" for m, p in _fallback_chain)
def load_fallback_chain(raw: str) -> None:
"""Populate the chain from a serialized config string.
Args:
raw: Comma-separated ``"model:provider"`` pairs. Empty or
whitespace-only segments are silently skipped.
"""
global _fallback_chain
chain: list[tuple[str, str]] = []
for part in raw.split(","):
part = part.strip()
if not part:
continue
if ":" in part:
model, provider = part.rsplit(":", 1)
chain.append((model.strip(), provider.strip()))
with _fallback_chain_lock:
_fallback_chain = chain
def _is_non_fallbackable(exc: Exception) -> str | None:
"""Determine whether an exception should bypass the fallback chain.
Args:
exc: The exception raised by a model call.
Returns:
A human-readable reason string if the error must *not* trigger
fallback, or ``None`` if fallback should proceed.
"""
from langchain_core.exceptions import ContextOverflowError
if isinstance(exc, ContextOverflowError):
return "context length exceeded"
err_msg = str(exc).lower()
is_400 = "400" in err_msg or "bad request" in err_msg
if is_400 and any(p in err_msg for p in _CONTEXT_LIMIT_PATTERNS):
return "context length exceeded"
if is_400 and any(p in err_msg for p in _MALFORMED_REQUEST_PATTERNS):
return "malformed request (client-side error)"
return None
async def _try_fallbacks(
request: ModelRequest,
invoke: Callable[[ModelRequest], Awaitable[ModelResponse]],
primary_exc: Exception,
) -> ModelResponse:
"""Walk the fallback chain, trying each model until one succeeds.
Shared implementation for both sync and async middleware entry points.
The *invoke* callable is an async function that calls the handler with
a given request — the sync path wraps the synchronous handler in a
trivial coroutine so both paths converge here.
Args:
request: The original model request.
invoke: Async callable that invokes the handler on a request.
primary_exc: The exception raised by the primary model.
Returns:
The ``ModelResponse`` from the first successful fallback.
Raises:
Exception: Re-raises the last exception if all fallbacks fail.
"""
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
)
last_exc = primary_exc
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}",
style="yellow",
)
try:
fallback_model = get_chat_model(model=model_name, provider=provider)
fb_request = request.override(model=fallback_model)
result = await invoke(fb_request)
_emit(
f" Fallback to {model_name} ({provider}) succeeded",
style="green",
)
logger.info("Fallback to %s (%s) succeeded", model_name, provider)
return result
except Exception as fb_exc:
reason = _is_non_fallbackable(fb_exc)
if reason is not None:
_emit(
f" {model_name} hit non-fallbackable error ({reason}) "
f"-- aborting fallback chain",
style="red",
)
raise
last_exc = fb_exc
_emit(
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
style="red",
)
logger.warning(
"Fallback %s (provider=%s) failed: %s: %s",
model_name,
provider,
type(fb_exc).__name__,
fb_exc,
)
_emit(" All fallbacks exhausted -- re-raising last error", style="red")
raise last_exc
def _guard_and_fallback(
primary_exc: Exception,
request: ModelRequest,
invoke: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> Awaitable[ModelResponse]:
"""Check non-fallbackable conditions, then delegate to ``_try_fallbacks``.
Args:
primary_exc: The exception raised by the primary model.
request: The original model request.
invoke: Async callable that invokes the handler on a request.
Returns:
Coroutine that resolves to the fallback ``ModelResponse``.
Raises:
Exception: Re-raises immediately for non-fallbackable errors.
"""
reason = _is_non_fallbackable(primary_exc)
if reason is not None:
_emit(
f"Model error ({reason}) -- not eligible for fallback, re-raising",
style="red",
)
raise primary_exc
return _try_fallbacks(request, invoke, primary_exc)
class ModelFallbackMiddleware(AgentMiddleware):
"""LangChain AgentMiddleware that retries failed model calls on fallbacks.
On each invocation the middleware reads the module-level
``_fallback_chain`` so that ``/model-fallback add`` takes effect
immediately without rebuilding the agent.
Attributes:
name: Middleware identifier used by the framework.
"""
name = "model_fallback"
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
if not _fallback_chain:
return handler(request)
try:
return handler(request)
except Exception as exc:
async def _sync_invoke(r: ModelRequest) -> ModelResponse:
return handler(r)
import asyncio
return asyncio.run(_guard_and_fallback(exc, request, _sync_invoke))
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
if not _fallback_chain:
return await handler(request)
try:
return await handler(request)
except Exception as exc:
return await _guard_and_fallback(exc, request, handler)
+223
View File
@@ -0,0 +1,223 @@
"""read_file image post-processing: content sniffing, downsampling, metadata.
Upstream deepagents builds multimodal ToolMessages for binary reads, but it
guesses the MIME type from the file extension and ships the raw bytes
verbatim — a multi-MB PNG goes straight to the provider, which may reject or
truncate it (the historical "read_file returns empty for images" report), and
a corrupt image silently reaches the model as undecodable base64.
This module wraps ``FilesystemMiddleware._create_read_file_tool`` (installed
at import time from ``EvoScientist.middleware``) and post-processes
successful image ToolMessages:
- content sniffing via Pillow (extension-independent);
- downsampling to a 2048px max edge, re-encoded as JPEG (or PNG when the
image carries alpha), with a metadata text block preserving the original
width/height/format and whether scaling happened;
- explicit errors ("无法读取图像数据: ...") when the payload is not a
decodable image — never a fake empty success.
Non-image results and small, correctly-typed images pass through untouched
(byte-identical to upstream).
"""
from __future__ import annotations
import asyncio
import base64
import functools
import io
from typing import Any
from langchain_core.messages import ToolMessage
from langchain_core.tools import BaseTool, StructuredTool
MAX_IMAGE_EDGE = 2048
_JPEG_QUALITY = 85
_installed = False
def _image_block(message: ToolMessage) -> dict[str, Any] | None:
"""Return the single base64 image block of an upstream read_file result."""
if message.status != "success":
return None
content = message.content
if not isinstance(content, list) or len(content) != 1:
return None
block = content[0]
if not isinstance(block, dict) or not isinstance(block.get("base64"), str):
return None
media_type = message.additional_kwargs.get("read_file_media_type") or ""
if block.get("type") == "image" or str(media_type).startswith("image/"):
return block
return None
def process_image_tool_message(message: ToolMessage) -> ToolMessage:
"""Sniff / downsample / re-encode an upstream read_file image result."""
block = _image_block(message)
if block is None:
return message
path = str(message.additional_kwargs.get("read_file_path", ""))
declared_mime = str(
message.additional_kwargs.get("read_file_media_type")
or block.get("mime_type")
or ""
)
# SVG is vector markup; Pillow cannot parse it and vision providers that
# accept it want the original bytes.
if "svg" in declared_mime or path.lower().endswith(".svg"):
return message
from PIL import Image
try:
raw = base64.standard_b64decode(block["base64"])
image = Image.open(io.BytesIO(raw))
image.load()
except Exception as exc:
return ToolMessage(
content=f"Error: 无法读取图像数据 '{path}': {exc}",
name=message.name,
tool_call_id=message.tool_call_id,
status="error",
)
orig_width, orig_height = image.size
orig_format = image.format or "unknown"
frames = int(getattr(image, "n_frames", 1))
sniffed_mime = Image.MIME.get(orig_format, declared_mime)
scaled = max(orig_width, orig_height) > MAX_IMAGE_EDGE
parts = [f"{path}: {orig_width}x{orig_height} {orig_format}"]
if frames > 1:
parts.append(f"first of {frames} frames")
# No transform needed: keep the original bytes (no re-encode quality
# loss), but still attach the metadata text block. On OpenAI-compatible
# providers the image is hoisted out of the tool message into a following
# user message, and this text is what the tool result itself carries —
# without it the model sees an empty tool result and the chat UI renders
# the tool box blank ("(empty)").
if not scaled and sniffed_mime == declared_mime:
return ToolMessage(
content=[
{"type": "text", "text": "[read_file image] " + ", ".join(parts)},
block,
],
name=message.name,
tool_call_id=message.tool_call_id,
status="success",
additional_kwargs=message.additional_kwargs,
)
has_alpha = image.mode in ("RGBA", "LA") or (
image.mode == "P" and "transparency" in image.info
)
work = image
if scaled:
work = work.copy()
work.thumbnail((MAX_IMAGE_EDGE, MAX_IMAGE_EDGE), Image.LANCZOS)
if has_alpha:
out_format, out_mime = "PNG", "image/png"
if work.mode == "P":
work = work.convert("RGBA")
else:
# JPEG drops alpha; only taken when there is none. WebP was considered
# but JPEG decodes on every vision endpoint.
out_format, out_mime = "JPEG", "image/jpeg"
if work.mode != "RGB":
work = work.convert("RGB")
try:
buffer = io.BytesIO()
work.save(buffer, out_format, quality=_JPEG_QUALITY, optimize=True)
except Exception as exc:
return ToolMessage(
content=f"Error: 图像转码失败 '{path}': {exc}",
name=message.name,
tool_call_id=message.tool_call_id,
status="error",
)
out_width, out_height = work.size
if scaled:
parts.append(
f"downsampled to {out_width}x{out_height} {out_format} "
f"(max edge {MAX_IMAGE_EDGE}px)"
)
else:
parts.append(f"re-encoded as {out_format}")
metadata = "[read_file image] " + ", ".join(parts)
return ToolMessage(
content=[
{"type": "text", "text": metadata},
{
"type": "image",
"base64": base64.standard_b64encode(buffer.getvalue()).decode(),
"mime_type": out_mime,
},
],
name=message.name,
tool_call_id=message.tool_call_id,
status="success",
additional_kwargs={
**message.additional_kwargs,
"read_file_media_type": out_mime,
},
)
def _wrap_tool(tool: BaseTool) -> BaseTool:
"""Rebuild a StructuredTool whose read_file results are post-processed."""
if not isinstance(tool, StructuredTool):
return tool
orig_func = tool.func
orig_coroutine = tool.coroutine
# functools.wraps is load-bearing, not cosmetic: StructuredTool and
# ToolNode inspect signature(fn) for ToolRuntime-annotated params to
# decide which arguments to inject. A bare *args/**kwargs wrapper hides
# `runtime`, so the framework never injects it and the original callable
# raises "missing 1 required positional argument: 'runtime'".
@functools.wraps(orig_func)
def sync_wrapper(*args: Any, **kwargs: Any) -> ToolMessage:
return process_image_tool_message(orig_func(*args, **kwargs))
@functools.wraps(orig_coroutine)
async def async_wrapper(*args: Any, **kwargs: Any) -> ToolMessage:
result = await orig_coroutine(*args, **kwargs)
# Pillow work is blocking; langgraph dev's blockbuster forbids it on
# the event loop.
return await asyncio.to_thread(process_image_tool_message, result)
return StructuredTool.from_function(
name=tool.name,
description=tool.description,
func=sync_wrapper,
coroutine=async_wrapper,
infer_schema=False,
args_schema=tool.args_schema,
)
def install_read_file_image_patch() -> None:
"""Wrap ``FilesystemMiddleware._create_read_file_tool`` (idempotent).
``create_deep_agent`` instantiates ``FilesystemMiddleware`` internally for
the main agent and every sub-agent, so patching the class method covers
all read_file tools in the process.
"""
global _installed
if _installed:
return
from deepagents.middleware.filesystem import FilesystemMiddleware
original = FilesystemMiddleware._create_read_file_tool
def _create_read_file_tool(self: Any) -> BaseTool:
return _wrap_tool(original(self))
FilesystemMiddleware._create_read_file_tool = _create_read_file_tool
_installed = True
+41 -14
View File
@@ -13,12 +13,14 @@ from __future__ import annotations
import asyncio
import time
from collections.abc import Awaitable, Callable
from types import SimpleNamespace
from langchain.agents.middleware.types import (
AgentMiddleware,
ModelRequest,
ModelResponse,
)
from langchain.tools import ToolRuntime
from langchain_core.tools import tool
from .utils import append_to_system_message
@@ -47,8 +49,20 @@ manage them.
# ---------------------------------------------------------------------------
def _scope_for_runtime(runtime: ToolRuntime | None):
from ..workspace_scope import require_scoped_runtime
return require_scoped_runtime(runtime, kind="scheduler")
@tool
def schedule_task(name: str, cron: str, prompt: str, timezone: str = "") -> str:
def schedule_task(
name: str,
cron: str,
prompt: str,
timezone: str = "",
runtime: ToolRuntime = None,
) -> str:
"""Create a recurring scheduled task that runs unattended in the background.
Translate the user's natural-language timing into a standard 5-field cron
@@ -67,7 +81,11 @@ def schedule_task(name: str, cron: str, prompt: str, timezone: str = "") -> str:
return "Scheduler unavailable: the langgraph dev backend is not running."
try:
rec = crons.create_schedule(
name=name, schedule=cron, prompt=prompt, timezone=timezone or None
name=name,
schedule=cron,
prompt=prompt,
timezone=timezone or None,
scope=_scope_for_runtime(runtime),
)
except Exception as e:
return f"Error: {e}"
@@ -78,14 +96,14 @@ def schedule_task(name: str, cron: str, prompt: str, timezone: str = "") -> str:
@tool
def list_scheduled_tasks() -> str:
def list_scheduled_tasks(runtime: ToolRuntime = None) -> str:
"""List the user's recurring scheduled tasks (id, name, schedule, enabled)."""
from ..cron import schedule as crons
if not crons.is_available():
return "Scheduler unavailable: the langgraph dev backend is not running."
try:
rows = crons.list_schedules()
rows = crons.list_schedules(_scope_for_runtime(runtime))
except Exception as e:
return f"Error: {e}"
if not rows:
@@ -101,7 +119,7 @@ def list_scheduled_tasks() -> str:
@tool
def cancel_scheduled_task(cron_id: str) -> str:
def cancel_scheduled_task(cron_id: str, runtime: ToolRuntime = None) -> str:
"""Cancel (delete) a scheduled task. Pass the id (or its prefix) shown by list_scheduled_tasks."""
from ..cron import schedule as crons
@@ -111,7 +129,7 @@ def cancel_scheduled_task(cron_id: str) -> str:
# Empty prefix would match (and delete) the only cron — refuse it.
return "Provide the id (or a prefix) of the task to cancel."
try:
rows = crons.list_schedules()
rows = crons.list_schedules(_scope_for_runtime(runtime))
# B2: collect ALL prefix matches before acting to detect ambiguity.
matches = [
r for r in rows if str(r.get("cron_id", "")).startswith(requested_id)
@@ -142,11 +160,15 @@ class SchedulerMiddleware(AgentMiddleware):
def __init__(self) -> None:
super().__init__()
self._cache: str | None = None
self._cache_at: float = 0.0
self._cache: dict[str, tuple[float, str]] = {}
self.tools = [schedule_task, list_scheduled_tasks, cancel_scheduled_task]
def _schedules_block(self) -> str:
def _runtime_scope(self):
from langgraph.config import get_config
return _scope_for_runtime(SimpleNamespace(config=get_config()))
def _schedules_block(self, scope=None) -> str:
"""Build the dynamic ``<scheduled_tasks>`` block (empty if none / down)."""
from ..cron import schedule as crons
@@ -158,7 +180,7 @@ class SchedulerMiddleware(AgentMiddleware):
try:
if not crons.is_available():
return ""
rows = crons.list_schedules()
rows = crons.list_schedules(scope)
except Exception:
return ""
if not rows:
@@ -182,10 +204,15 @@ class SchedulerMiddleware(AgentMiddleware):
def _cached_schedules_block(self) -> str:
now = time.monotonic()
if self._cache is None or (now - self._cache_at) > _CACHE_TTL_SECONDS:
self._cache = self._schedules_block()
self._cache_at = now
return self._cache
try:
scope = self._runtime_scope()
except Exception:
return ""
key = f"{scope.scope_id}:{scope.revision}" if scope is not None else "legacy"
cached = self._cache.get(key)
if cached is None or (now - cached[0]) > _CACHE_TTL_SECONDS:
self._cache[key] = (now, self._schedules_block(scope))
return self._cache[key][1]
def _injection(self, schedules_block: str) -> str:
"""Static instructions, then the dynamic list (static→dynamic, like memory)."""
+52 -14
View File
@@ -93,15 +93,12 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
self._threshold = threshold
self._always_include = always_include or frozenset()
self._track_stream_selection = track_stream_selection
# Agent tools are fixed after graph construction, so the filtered
# always-include set is stable for this middleware instance.
self._selector: AgentMiddleware | None = None
def _build_selector(self, request: ModelRequest) -> AgentMiddleware:
if self._selector is None:
names = _available_always_include(request.tools, self._always_include)
self._selector = self._selector_factory(names)
return self._selector
# Built per call: the selector's helper model is resolved from the
# run snapshot, which varies from run to run.
names = _available_always_include(request.tools, self._always_include)
return self._selector_factory(names)
def wrap_model_call(
self,
@@ -233,8 +230,10 @@ def create_tool_selector_middleware(
names for the main-agent stream UI when ``track_stream_selection`` is true
Args:
model: Chat model for tool selection. If *None*, the default
model is resolved via ``_ensure_chat_model()``.
model: Chat model for tool selection. If *None*, the helper model
is resolved from the active run's snapshot on every call
(``resolve_snapshot_model``); the registry default only applies
to snapshot-less background runs via the lazy local snapshot.
threshold: Minimum number of tools to trigger selection.
Default 26. Set to 0 to always run selection.
track_stream_selection: Whether to update process-global stream/UI
@@ -253,11 +252,50 @@ def create_tool_selector_middleware(
from .utils import disable_thinking
if model is None:
from EvoScientist.EvoScientist import _ensure_chat_model
def tag_selector_model(base: BaseChatModel) -> BaseChatModel:
safe_model = disable_thinking(base)
selector_model = safe_model
from EvoScientist.usage.callback import usage_tracking_enabled
model = _ensure_chat_model()
safe_model = disable_thinking(model)
if usage_tracking_enabled():
try:
selector_model = safe_model.model_copy(
update={
"metadata": {
**(safe_model.metadata or {}),
"usage_scope": "tool_selector",
}
}
)
if not isinstance(selector_model, BaseChatModel):
raise TypeError("selector model copy is not a BaseChatModel")
except Exception:
# The model factory preflights this path before enabling its
# callback. Keep the selector operational even for an unusual
# third-party model.
from EvoScientist.usage.spool import mark_tracking_degraded
mark_tracking_degraded("selector_model_copy_unsupported")
selector_model = safe_model
logger.exception("Could not attach Tool Selector usage scope")
return selector_model
# The resolved snapshot model is cached per snapshot upstream, so the
# tagged copy can be cached by the base model's identity.
tagged_cache: dict[int, BaseChatModel] = {}
def selector_model() -> BaseChatModel:
if model is not None:
base = model
else:
from .configurable_model import resolve_snapshot_model
base = resolve_snapshot_model()
tagged = tagged_cache.get(id(base))
if tagged is None:
tagged = tag_selector_model(base)
tagged_cache[id(base)] = tagged
return tagged
system_prompt = (
"You are selecting tools for a scientific research agent. "
@@ -270,7 +308,7 @@ def create_tool_selector_middleware(
def selector_factory(always_include: list[str]) -> AgentMiddleware:
return LLMToolSelectorMiddleware(
model=safe_model,
model=selector_model(),
system_prompt=system_prompt,
always_include=always_include,
)
+168
View File
@@ -0,0 +1,168 @@
"""Unified model registry: schema, errors, store, and network egress policy.
This subpackage is the single authority for the version 4 model registry
(design doc sections 4.2, 4.3, 5.1, 8.2, 9.5). HTTP JSON, SQLite JSON, import
tooling, and test fixtures all reuse these Pydantic models. EndpointPolicy
and SafeHttpTransport form the SSRF defense and the only network egress
for adapters.
"""
from __future__ import annotations
from .adapters import (
Adapter,
BuiltRequest,
ResolvedParameters,
adapter_specs,
compute_effective_capabilities,
find_adapter_spec,
get_adapter,
resolve_parameters,
)
from .auth import ActorContext, BffAuthenticator
from .endpoint_policy import EndpointPolicy
from .errors import ERROR_HTTP_STATUS, ErrorDetail, ErrorPayload, ModelRegistryError
from .factory import build_chat_model
from .hashing import configuration_hash
from .http_api import (
ApiServices,
CredentialReplaceRequest,
CredentialWriteResponse,
GetModelRegistryResponse,
GetSelectableModelsResponse,
ModelRegistryHttpApi,
ProviderTestRequest,
ProviderTestResponse,
PutModelRegistryRequest,
SnapshotBindRequest,
SnapshotPublicResponse,
build_openapi_document,
model_registry_routes,
validate_registry_save,
)
from .platform import (
PlatformConfigError,
PlatformSecurityConfig,
load_platform_security_config,
)
from .provider_test import ProviderTester, ProviderTestResult
from .resolver import ModelRegistryResolver
from .safe_transport import (
AsyncSafeHttpTransport,
AsyncSafeNetworkBackend,
SafeHttpTransport,
SafeNetworkBackend,
build_safe_async_http_client,
build_safe_http_client,
)
from .schemas import (
AdapterParameterSpec,
AuthConfig,
AuthRef,
AuthSpec,
Capabilities,
CredentialStatus,
CredentialWrite,
DevelopmentEndpoint,
EndpointPolicyPublic,
ModelAvailability,
ModelConfig,
ModelRef,
ModelRuntimeConfig,
ParameterRule,
ProviderConfig,
ProviderRuntimeConfig,
RegistryV4,
ResolvedModelConfig,
VerificationInfo,
)
from .snapshots import (
BOUND_RETENTION_SECONDS,
PREPARED_TTL_SECONDS,
RuntimeSnapshot,
SnapshotCreateRequest,
SnapshotCreation,
SnapshotPayload,
SnapshotService,
compute_selection_hash,
config_for_role,
public_snapshot_view,
)
from .store import ModelRuntimeStore, SharedStorageError
__all__ = [
"BOUND_RETENTION_SECONDS",
"ERROR_HTTP_STATUS",
"PREPARED_TTL_SECONDS",
"ActorContext",
"Adapter",
"AdapterParameterSpec",
"ApiServices",
"AsyncSafeHttpTransport",
"AsyncSafeNetworkBackend",
"AuthConfig",
"AuthRef",
"AuthSpec",
"BffAuthenticator",
"BuiltRequest",
"Capabilities",
"CredentialReplaceRequest",
"CredentialStatus",
"CredentialWrite",
"CredentialWriteResponse",
"DevelopmentEndpoint",
"EndpointPolicy",
"EndpointPolicyPublic",
"ErrorDetail",
"ErrorPayload",
"GetModelRegistryResponse",
"GetSelectableModelsResponse",
"ModelAvailability",
"ModelConfig",
"ModelRef",
"ModelRegistryError",
"ModelRegistryHttpApi",
"ModelRegistryResolver",
"ModelRuntimeConfig",
"ModelRuntimeStore",
"ParameterRule",
"PlatformConfigError",
"PlatformSecurityConfig",
"ProviderConfig",
"ProviderRuntimeConfig",
"ProviderTestRequest",
"ProviderTestResponse",
"ProviderTestResult",
"ProviderTester",
"PutModelRegistryRequest",
"RegistryV4",
"ResolvedModelConfig",
"ResolvedParameters",
"RuntimeSnapshot",
"SafeHttpTransport",
"SafeNetworkBackend",
"SharedStorageError",
"SnapshotBindRequest",
"SnapshotCreateRequest",
"SnapshotCreation",
"SnapshotPayload",
"SnapshotPublicResponse",
"SnapshotService",
"VerificationInfo",
"adapter_specs",
"build_chat_model",
"build_openapi_document",
"build_safe_async_http_client",
"build_safe_http_client",
"compute_effective_capabilities",
"compute_selection_hash",
"config_for_role",
"configuration_hash",
"find_adapter_spec",
"get_adapter",
"load_platform_security_config",
"model_registry_routes",
"public_snapshot_view",
"resolve_parameters",
"validate_registry_save",
]
+830
View File
@@ -0,0 +1,830 @@
"""Adapter parameter contracts and request mapping (design doc 6.1-6.3).
This module is the single authority for how unified registry parameters are
validated, normalized, and mapped onto LangChain constructor options and
provider request fields:
- ``adapter_specs`` holds the versioned built-in contracts: a generic
``model_selector: "*"`` contract for each phase-1 adapter plus the
``openai-compatible``/``glm-5.2`` model-specific contract.
- ``find_adapter_spec`` implements the matching order — exact
``upstream_model_id`` first, longest glob next, generic contract last.
- ``resolve_parameters`` applies the section 6.1 inheritance semantics and
the save-time contract checks (stable error codes from section 9.5).
- ``Adapter.build_request`` is the only entry point that turns a frozen
``ResolvedModelConfig`` into ``{client_options, request_options}``;
parameters the contract does not declare never enter the request.
A ``target_name`` of ``""`` declares a parameter the contract validates but
enforces outside the request payload (for example ollama retries, which live
in the SafeHttpTransport). Dotted ``client_option`` target names (for example
``client_kwargs.timeout``) build nested option dicts.
"""
from __future__ import annotations
import fnmatch
from typing import Any
from pydantic import BaseModel, ConfigDict
from .errors import (
ADAPTER_NOT_SUPPORTED,
AUTH_MODE_UNSUPPORTED,
CAPABILITY_UNSUPPORTED_BY_ADAPTER,
CREDENTIAL_NOT_CONFIGURED,
UNSUPPORTED_RUNTIME_PARAMETER,
ModelRegistryError,
)
from .schemas import (
AdapterParameterSpec,
AuthSpec,
Capabilities,
ConnectionSpec,
ModelConfig,
ParameterRule,
ProviderConfig,
ReasoningEffort,
ResolvedModelConfig,
SamplingOverride,
)
SUPPORTED_ADAPTER_IDS = (
"openai",
"anthropic",
"openai-compatible",
"anthropic-compatible",
"ollama",
)
# Known adapter IDs that phase 1 deliberately does not open (section 6.2).
UNOPENED_ADAPTER_IDS = ("google-genai", "grok", "openrouter", "nvidia", "antigravity")
# Unified parameter keys the mapping engine knows how to source (section 6.3).
_UNIFIED_PARAMETERS = (
"timeout_seconds",
"max_retries",
"max_output_tokens",
"temperature",
"top_p",
"reasoning_effort",
)
_CAPABILITY_NAMES = ("tools", "vision", "structured_output")
_REASONING_VALUES = ["low", "medium", "high"]
class BuiltRequest(BaseModel):
"""The output of ``Adapter.build_request`` (section 6.4)."""
model_config = ConfigDict(frozen=True)
client_options: dict[str, Any]
request_options: dict[str, Any]
class ResolvedParameters(BaseModel):
"""Save-time resolution result; the Resolver freezes it into a snapshot."""
model_config = ConfigDict(frozen=True)
timeout_seconds: int
max_retries: int
max_output_tokens: int
temperature: float | None
top_p: float | None
# "auto" means the field is omitted from the provider request.
reasoning_effort: ReasoningEffort
# --- Normalizers (named server-side functions referenced by contracts) ---
class _Omit:
"""Sentinel: the parameter is omitted from the outbound request."""
def __repr__(self) -> str: # pragma: no cover - debugging aid
return "OMIT"
OMIT = _Omit()
def _reject_parameter(name: str, message: str) -> ModelRegistryError:
return ModelRegistryError(
UNSUPPORTED_RUNTIME_PARAMETER,
message,
details=[{"path": f"runtime.{name}", "code": UNSUPPORTED_RUNTIME_PARAMETER}],
)
def _normalize_identity(name: str, rule: ParameterRule, value: Any) -> Any:
return value
def _normalize_clamp_to_model_limit(name: str, rule: ParameterRule, value: Any) -> Any:
if value is None:
return OMIT
if rule.maximum is not None and value > rule.maximum:
return rule.maximum
if rule.minimum is not None and value < rule.minimum:
raise _reject_parameter(
name,
f"{name}={value} is below the contract minimum {rule.minimum}.",
)
return value
def _normalize_omit_when_none(name: str, rule: ParameterRule, value: Any) -> Any:
return OMIT if value is None else value
def _normalize_omit_when_auto(name: str, rule: ParameterRule, value: Any) -> Any:
return OMIT if value is None or value == "auto" else value
def _normalize_reject_non_auto(name: str, rule: ParameterRule, value: Any) -> Any:
if value is None or value == "auto":
return OMIT
raise _reject_parameter(
name,
f"{name} is not supported by this adapter contract; only the "
"inheriting 'auto' value is accepted.",
)
_NORMALIZERS = {
"identity": _normalize_identity,
"clamp_to_model_limit": _normalize_clamp_to_model_limit,
"omit_when_none": _normalize_omit_when_none,
"omit_when_auto": _normalize_omit_when_auto,
"reject_non_auto": _normalize_reject_non_auto,
}
def _check_rule_value(name: str, rule: ParameterRule, value: Any) -> None:
"""Type/range/enum validation for a non-null value against the contract.
With the ``clamp_to_model_limit`` normalizer an out-of-range high value is
clamped by the normalizer instead of rejected, so the maximum check is
left to it.
"""
if rule.value_type == "integer":
if isinstance(value, bool) or not isinstance(value, int):
raise _reject_parameter(name, f"{name} must be an integer.")
elif rule.value_type == "number":
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise _reject_parameter(name, f"{name} must be a number.")
elif rule.value_type == "enum":
if rule.enum_values is not None and value not in rule.enum_values:
raise _reject_parameter(name, f"{name} must be one of {rule.enum_values}.")
if rule.value_type in ("integer", "number"):
if rule.minimum is not None and value < rule.minimum:
raise _reject_parameter(
name,
f"{name}={value} is below the contract minimum {rule.minimum}.",
)
if (
rule.normalizer != "clamp_to_model_limit"
and rule.maximum is not None
and value > rule.maximum
):
raise _reject_parameter(
name,
f"{name}={value} exceeds the contract maximum {rule.maximum}.",
)
# --- Built-in contracts (section 6.2) ---
def _connection(chat_model: str) -> ConnectionSpec:
return ConnectionSpec(
chat_model=chat_model, model_field="model", base_url_field="base_url"
)
def _timeout_rule(target_name: str) -> ParameterRule:
return ParameterRule(
supported=True,
value_type="integer",
minimum=10,
maximum=600,
nullable="forbidden",
target="client_option",
target_name=target_name,
normalizer="identity",
)
def _retries_rule(target_name: str) -> ParameterRule:
return ParameterRule(
supported=True,
value_type="integer",
minimum=0,
maximum=5,
nullable="forbidden",
target="client_option",
target_name=target_name,
normalizer="identity",
)
def _max_output_tokens_rule(
target_name: str, maximum: float | None = None
) -> ParameterRule:
return ParameterRule(
supported=True,
value_type="integer",
minimum=1,
maximum=maximum,
nullable="forbidden",
target="request_option",
target_name=target_name,
normalizer="clamp_to_model_limit",
)
def _temperature_rule(maximum: float) -> ParameterRule:
return ParameterRule(
supported=True,
value_type="number",
minimum=0,
maximum=maximum,
nullable="omit",
target="request_option",
target_name="temperature",
normalizer="omit_when_none",
)
def _top_p_rule() -> ParameterRule:
return ParameterRule(
supported=True,
value_type="number",
minimum=0.000001,
maximum=1,
nullable="omit",
target="request_option",
target_name="top_p",
normalizer="omit_when_none",
)
def _reasoning_supported_rule() -> ParameterRule:
return ParameterRule(
supported=True,
value_type="enum",
nullable="omit",
enum_values=list(_REASONING_VALUES),
target="request_option",
target_name="reasoning_effort",
normalizer="omit_when_auto",
)
def _reasoning_unsupported_rule() -> ParameterRule:
return ParameterRule(
supported=False,
value_type="enum",
nullable="omit",
enum_values=list(_REASONING_VALUES),
target="request_option",
target_name="",
normalizer="reject_non_auto",
)
def _api_key_auth() -> AuthSpec:
return AuthSpec(
credential_required=True,
credential_kind="api_key",
target="client_option",
target_name="api_key",
)
def _bearer_auth() -> AuthSpec:
return AuthSpec(
credential_required=True,
credential_kind="bearer_token",
target="request_header",
target_name="Authorization",
)
def _openai_style_parameters() -> dict[str, ParameterRule]:
return {
"timeout_seconds": _timeout_rule("timeout"),
"max_retries": _retries_rule("max_retries"),
"max_output_tokens": _max_output_tokens_rule("max_tokens"),
"temperature": _temperature_rule(2),
"top_p": _top_p_rule(),
"reasoning_effort": _reasoning_supported_rule(),
}
def _anthropic_style_parameters() -> dict[str, ParameterRule]:
return {
"timeout_seconds": _timeout_rule("timeout"),
"max_retries": _retries_rule("max_retries"),
"max_output_tokens": _max_output_tokens_rule("max_tokens"),
"temperature": _temperature_rule(1),
"top_p": _top_p_rule(),
"reasoning_effort": _reasoning_unsupported_rule(),
}
def _build_builtin_specs() -> tuple[AdapterParameterSpec, ...]:
openai_generic = AdapterParameterSpec(
adapter_id="openai",
spec_revision=1,
model_selector="*",
auth_specs={"api_key": _api_key_auth()},
parameters=_openai_style_parameters(),
protocol_capabilities=Capabilities(
tools=True, vision=True, structured_output=True
),
connection=_connection("ChatOpenAI"),
)
openai_compatible_generic = AdapterParameterSpec(
adapter_id="openai-compatible",
spec_revision=1,
model_selector="*",
auth_specs={"api_key": _api_key_auth(), "bearer": _bearer_auth()},
parameters=_openai_style_parameters(),
protocol_capabilities=Capabilities(
tools=True, vision=True, structured_output=True
),
connection=_connection("ChatOpenAI"),
)
anthropic_generic = AdapterParameterSpec(
adapter_id="anthropic",
spec_revision=1,
model_selector="*",
auth_specs={"api_key": _api_key_auth()},
parameters=_anthropic_style_parameters(),
protocol_capabilities=Capabilities(
tools=True, vision=True, structured_output=False
),
connection=_connection("ChatAnthropic"),
)
anthropic_compatible_generic = anthropic_generic.model_copy(
update={"adapter_id": "anthropic-compatible"}
)
ollama_generic = AdapterParameterSpec(
adapter_id="ollama",
spec_revision=1,
model_selector="*",
auth_specs={
"none": AuthSpec(
credential_required=False,
credential_kind="none",
target="adapter_internal",
target_name=None,
)
},
parameters={
"timeout_seconds": _timeout_rule("client_kwargs.timeout"),
# Retries are enforced by the SafeHttpTransport built in Task 2,
# not by an ollama client option; the contract still validates
# the registry value (empty target_name = validated, not sent).
"max_retries": _retries_rule(""),
"max_output_tokens": _max_output_tokens_rule("num_predict"),
"temperature": _temperature_rule(2),
"top_p": _top_p_rule(),
"reasoning_effort": _reasoning_unsupported_rule(),
},
protocol_capabilities=Capabilities(
tools=True, vision=True, structured_output=True
),
connection=_connection("ChatOllama"),
)
# The first model-specific contract (section 6.2 YAML, verbatim values).
glm_52 = AdapterParameterSpec(
adapter_id="openai-compatible",
spec_revision=1,
model_selector="glm-5.2",
auth_specs={"api_key": _api_key_auth()},
parameters={
"timeout_seconds": _timeout_rule("timeout"),
"max_retries": _retries_rule("max_retries"),
"max_output_tokens": _max_output_tokens_rule("max_tokens", maximum=32768),
"temperature": _temperature_rule(1),
"top_p": _top_p_rule(),
"reasoning_effort": _reasoning_unsupported_rule(),
},
protocol_capabilities=Capabilities(
tools=True, vision=False, structured_output=True
),
connection=_connection("ChatOpenAI"),
)
return (
openai_generic,
openai_compatible_generic,
anthropic_generic,
anthropic_compatible_generic,
ollama_generic,
glm_52,
)
_BUILTIN_SPECS: tuple[AdapterParameterSpec, ...] = _build_builtin_specs()
def adapter_specs() -> tuple[AdapterParameterSpec, ...]:
"""Return the built-in versioned adapter contracts."""
return _BUILTIN_SPECS
def find_adapter_spec(
adapter_id: str,
upstream_model_id: str,
*,
spec_revision: int | None = None,
specs: list[AdapterParameterSpec] | tuple[AdapterParameterSpec, ...] | None = None,
) -> AdapterParameterSpec | None:
"""Match a contract: exact ID first, longest glob next, ``*`` last.
Unopened or unknown adapter IDs raise ``ADAPTER_NOT_SUPPORTED``. ``None``
is returned when no contract matches (or the pinned ``spec_revision`` is
gone); such configurations may only be saved as ``configured``.
"""
if adapter_id not in SUPPORTED_ADAPTER_IDS:
note = (
"a known but not yet opened adapter"
if adapter_id in UNOPENED_ADAPTER_IDS
else "an unknown adapter"
)
raise ModelRegistryError(
ADAPTER_NOT_SUPPORTED,
f"Adapter {adapter_id!r} is {note}; phase 1 supports "
f"{list(SUPPORTED_ADAPTER_IDS)}.",
)
candidates = [
spec
for spec in (adapter_specs() if specs is None else specs)
if spec.adapter_id == adapter_id
and (spec_revision is None or spec.spec_revision == spec_revision)
]
for spec in candidates:
if spec.model_selector == upstream_model_id:
return spec
globs = [
spec
for spec in candidates
if spec.model_selector != "*"
and fnmatch.fnmatchcase(upstream_model_id, spec.model_selector)
]
if globs:
return max(globs, key=lambda spec: len(spec.model_selector))
for spec in candidates:
if spec.model_selector == "*":
return spec
return None
def get_adapter(
adapter_id: str,
upstream_model_id: str,
*,
spec_revision: int | None = None,
) -> Adapter:
"""Return the matched Adapter, failing loudly when no contract applies."""
spec = find_adapter_spec(adapter_id, upstream_model_id, spec_revision=spec_revision)
if spec is None:
revision_note = (
f" at spec_revision {spec_revision}" if spec_revision is not None else ""
)
raise ModelRegistryError(
ADAPTER_NOT_SUPPORTED,
f"No adapter contract matches {adapter_id!r}/{upstream_model_id!r}"
f"{revision_note}; the configuration may only remain 'configured'.",
)
return Adapter(spec)
def compute_effective_capabilities(
protocol_capabilities: Capabilities,
declared_capabilities: Capabilities,
verified_capabilities: Capabilities,
) -> Capabilities:
"""The single capability rule: protocol AND declared AND verified."""
return Capabilities(
**{
name: (
getattr(protocol_capabilities, name)
and getattr(declared_capabilities, name)
and getattr(verified_capabilities, name)
)
for name in _CAPABILITY_NAMES
}
)
def _resolve_nullable_parameter(
name: str,
rule: ParameterRule,
value: Any,
) -> Any:
"""Validate one already-inherited value; returns the value or ``OMIT``."""
if not rule.supported:
if value is None or value == "auto":
return OMIT
# A concrete value for an unsupported parameter is handed to the
# contract-declared normalizer: reject_non_auto raises here. Any
# normalizer that would let the value through still rejects, because
# an unsupported parameter must never enter the request.
resolved = _NORMALIZERS[rule.normalizer](name, rule, value)
if resolved is not OMIT:
raise _reject_parameter(
name, f"{name} is not supported by this adapter contract."
)
return OMIT
if value is None or (name == "reasoning_effort" and value == "auto"):
if rule.nullable == "forbidden":
raise _reject_parameter(name, f"{name} must have a concrete value.")
return OMIT
_check_rule_value(name, rule, value)
return _NORMALIZERS[rule.normalizer](name, rule, value)
def _resolve_required_parameter(
name: str,
rule: ParameterRule,
value: Any,
) -> Any:
"""Validate a registry-sourced value that must always resolve."""
if not rule.supported:
raise _reject_parameter(
name, f"{name} is not supported by this adapter contract."
)
if value is None:
raise _reject_parameter(name, f"{name} must have a concrete value.")
_check_rule_value(name, rule, value)
resolved = _NORMALIZERS[rule.normalizer](name, rule, value)
if resolved is OMIT:
raise _reject_parameter(name, f"{name} must have a concrete value.")
return resolved
def _parameter_rule(spec: AdapterParameterSpec, name: str) -> ParameterRule:
rule = spec.parameters.get(name)
if rule is None:
raise _reject_parameter(
name, f"Adapter contract {spec.adapter_id!r} does not declare {name}."
)
return rule
def resolve_parameters(
provider: ProviderConfig,
model: ModelConfig,
spec: AdapterParameterSpec,
*,
reasoning_effort_override: ReasoningEffort | None = None,
sampling_override: SamplingOverride | None = None,
) -> ResolvedParameters:
"""Resolve a model's runtime parameters against the matched contract.
Applies the section 6.1 inheritance semantics (model value overrides the
provider default; provider ``null`` means the field is omitted, never
zero; ``reasoning_effort=auto`` inherits and is omitted when still auto)
and the save-time contract checks. Raises ``ModelRegistryError`` with a
stable section 9.5 code on any violation.
The sampling override (snapshot creation only) replaces the model's
configured value for its ``kind`` and omits the other sampling
parameter entirely — registry defaults included — because temperature
and top_p must not be sent together. Overrides flow through the same
contract rules, so unsupported adapters still reject them.
"""
auth_spec = spec.auth_specs.get(provider.auth.mode)
if auth_spec is None:
raise ModelRegistryError(
AUTH_MODE_UNSUPPORTED,
f"Adapter {spec.adapter_id!r} does not support auth mode "
f"{provider.auth.mode!r}.",
details=[{"path": "auth.mode", "code": AUTH_MODE_UNSUPPORTED}],
)
if auth_spec.credential_required and provider.auth.credential_id is None:
raise ModelRegistryError(
CREDENTIAL_NOT_CONFIGURED,
f"Auth mode {provider.auth.mode!r} requires a credential reference.",
details=[{"path": "auth.credential_id", "code": CREDENTIAL_NOT_CONFIGURED}],
)
for name in _CAPABILITY_NAMES:
if getattr(model.runtime.declared_capabilities, name) and not getattr(
spec.protocol_capabilities, name
):
raise ModelRegistryError(
CAPABILITY_UNSUPPORTED_BY_ADAPTER,
f"Adapter {spec.adapter_id!r} protocol does not support the "
f"declared capability {name!r}.",
details=[
{
"path": f"runtime.declared_capabilities.{name}",
"code": CAPABILITY_UNSUPPORTED_BY_ADAPTER,
}
],
)
timeout_seconds = _resolve_required_parameter(
"timeout_seconds",
_parameter_rule(spec, "timeout_seconds"),
provider.runtime.timeout_seconds,
)
max_retries = _resolve_required_parameter(
"max_retries",
_parameter_rule(spec, "max_retries"),
provider.runtime.max_retries,
)
max_output_tokens = _resolve_required_parameter(
"max_output_tokens",
_parameter_rule(spec, "max_output_tokens"),
model.runtime.max_output_tokens,
)
if sampling_override is not None and sampling_override.kind == "temperature":
temperature = _resolve_nullable_parameter(
"temperature", _parameter_rule(spec, "temperature"), sampling_override.value
)
top_p = _resolve_nullable_parameter(
"top_p", _parameter_rule(spec, "top_p"), None
)
elif sampling_override is not None:
top_p = _resolve_nullable_parameter(
"top_p", _parameter_rule(spec, "top_p"), sampling_override.value
)
temperature = _resolve_nullable_parameter(
"temperature", _parameter_rule(spec, "temperature"), None
)
else:
temperature = model.runtime.temperature
if temperature is None:
temperature = provider.runtime.default_temperature
temperature = _resolve_nullable_parameter(
"temperature", _parameter_rule(spec, "temperature"), temperature
)
top_p = model.runtime.top_p
if top_p is None:
top_p = provider.runtime.default_top_p
top_p = _resolve_nullable_parameter("top_p", _parameter_rule(spec, "top_p"), top_p)
reasoning_effort: ReasoningEffort = (
reasoning_effort_override
if reasoning_effort_override not in (None, "auto")
else model.runtime.reasoning_effort
)
if reasoning_effort == "auto":
reasoning_effort = provider.runtime.default_reasoning_effort
reasoning_effort = _resolve_nullable_parameter(
"reasoning_effort",
_parameter_rule(spec, "reasoning_effort"),
reasoning_effort,
)
resolved = {
"timeout_seconds": timeout_seconds,
"max_retries": max_retries,
"max_output_tokens": max_output_tokens,
"temperature": None if temperature is OMIT else temperature,
"top_p": None if top_p is OMIT else top_p,
"reasoning_effort": ("auto" if reasoning_effort is OMIT else reasoning_effort),
}
_check_conflicts(spec, resolved)
return ResolvedParameters(**resolved)
def _check_conflicts(spec: AdapterParameterSpec, resolved: dict[str, Any]) -> None:
present = {
name
for name, value in resolved.items()
if value is not None and value != "auto"
}
for name in present:
rule = spec.parameters.get(name)
if rule is None:
continue
for other in rule.conflicts_with:
if other in present:
raise _reject_parameter(
name, f"{name} conflicts with {other} in this contract."
)
def _assign_option(options: dict[str, Any], dotted_name: str, value: Any) -> None:
parts = dotted_name.split(".")
target = options
for part in parts[:-1]:
existing = target.get(part)
if not isinstance(existing, dict):
existing = {}
target[part] = existing
target = existing
target[parts[-1]] = value
class Adapter:
"""A versioned parameter contract bound to one ``AdapterParameterSpec``."""
def __init__(self, spec: AdapterParameterSpec) -> None:
for name, rule in spec.parameters.items():
if name not in _UNIFIED_PARAMETERS:
raise ValueError(f"Unknown unified parameter {name!r} in contract.")
if rule.normalizer not in _NORMALIZERS:
raise ValueError(f"Unknown normalizer {rule.normalizer!r}.")
self._spec = spec
@property
def spec(self) -> AdapterParameterSpec:
return self._spec
def build_request(
self,
resolved_config: ResolvedModelConfig,
*,
credential: str | None = None,
) -> BuiltRequest:
"""Map a frozen config into ``{client_options, request_options}``.
This is the only entry point that constructs LangChain parameters and
provider request parameters. Contract validation is re-applied so a
stale snapshot fails loudly instead of silently sending values the
current contract would reject.
"""
spec = self._spec
auth_spec = spec.auth_specs.get(resolved_config.auth_ref.mode)
if auth_spec is None:
raise ModelRegistryError(
AUTH_MODE_UNSUPPORTED,
f"Adapter {spec.adapter_id!r} does not support auth mode "
f"{resolved_config.auth_ref.mode!r}.",
)
if auth_spec.credential_required and credential is None:
raise ModelRegistryError(
CREDENTIAL_NOT_CONFIGURED,
"A credential is required to build the request but none was "
"resolved for this run.",
)
client_options: dict[str, Any] = {}
if spec.connection is not None:
client_options[spec.connection.model_field] = (
resolved_config.upstream_model_id
)
client_options[spec.connection.base_url_field] = resolved_config.base_url
sources: dict[str, Any] = {
"timeout_seconds": resolved_config.client_options.timeout_seconds,
"max_retries": resolved_config.client_options.max_retries,
"max_output_tokens": resolved_config.request_options.max_output_tokens,
"temperature": resolved_config.request_options.temperature,
"top_p": resolved_config.request_options.top_p,
"reasoning_effort": resolved_config.request_options.reasoning_effort,
}
request_options: dict[str, Any] = {}
for name, rule in spec.parameters.items():
value = _resolve_nullable_parameter(name, rule, sources[name])
if value is OMIT or rule.target_name == "":
continue
if rule.target == "client_option":
_assign_option(client_options, rule.target_name, value)
elif rule.target == "request_option":
_assign_option(request_options, rule.target_name, value)
else: # extra_body_path
extra_body = request_options.get("extra_body")
if not isinstance(extra_body, dict):
extra_body = {}
request_options["extra_body"] = extra_body
_assign_option(extra_body, rule.target_name, value)
if resolved_config.auth_ref.mode != "none" and credential is not None:
self._apply_auth(auth_spec, client_options, credential)
return BuiltRequest(
client_options=client_options, request_options=request_options
)
@staticmethod
def _apply_auth(
auth_spec: AuthSpec, client_options: dict[str, Any], credential: str
) -> None:
if auth_spec.target == "client_option" and auth_spec.target_name:
client_options[auth_spec.target_name] = credential
elif auth_spec.target == "request_header" and auth_spec.target_name:
header_value = (
f"Bearer {credential}"
if auth_spec.credential_kind == "bearer_token"
else credential
)
headers = client_options.get("default_headers")
if not isinstance(headers, dict):
headers = {}
client_options["default_headers"] = headers
headers[auth_spec.target_name] = header_value
# adapter_internal credentials are handled by the adapter itself and
# never appear in client or request options.
+198
View File
@@ -0,0 +1,198 @@
"""BFF service-token and delegation-JWT authentication (design doc 7.3).
Every BFF → EvoScientist request carries::
Authorization: Bearer <BFF service token>
X-Evo-Actor: <short-lived signed delegation JWT>
The service token only proves the request comes from a trusted WebUI
deployment; it is compared in constant time (against the configured
plaintext token or the configured SHA-256 hash). The delegation JWT carries
the end-user identity: it must be signed by a registered WebUI public key,
have ``iss == "WebUI"`` and ``aud == "EvoScientist"``, live at most 60
seconds, and carry ``sub``/``scopes``/``deployment_id``/``iat``/``exp``/
``jti`` (plus ``thread_id`` for thread-level routes). Each ``jti`` is
registered atomically in the ``delegation_jtis`` table; a live duplicate is
rejected with ``401 DELEGATION_REPLAYED``.
"""
from __future__ import annotations
import hashlib
import secrets
from collections.abc import Mapping
from dataclasses import dataclass
import jwt
from .errors import (
DELEGATION_REPLAYED,
FORBIDDEN,
UNAUTHENTICATED,
ModelRegistryError,
)
from .platform import DelegationPublicKey
DELEGATION_ISSUER = "WebUI"
DELEGATION_AUDIENCE = "EvoScientist"
DELEGATION_MAX_LIFETIME_SECONDS = 60
DELEGATION_ALGORITHMS = ("ES256", "RS256")
_REQUIRED_CLAIMS = ("sub", "scopes", "deployment_id", "iat", "exp", "jti")
@dataclass(frozen=True)
class ActorContext:
"""The verified end-user identity forwarded by the BFF."""
sub: str
scopes: frozenset[str]
deployment_id: str
thread_id: str | None
jti: str
def _unauthenticated(message: str) -> ModelRegistryError:
return ModelRegistryError(UNAUTHENTICATED, message)
class BffAuthenticator:
"""Verifies the fixed BFF authentication format on every request."""
def __init__(
self,
*,
service_token: str | None,
service_token_hash: str | None,
delegation_keys: tuple[DelegationPublicKey, ...],
jti_store: object,
) -> None:
if not service_token and not service_token_hash:
raise ValueError("A BFF service token or token hash must be configured.")
if not delegation_keys:
raise ValueError("At least one WebUI delegation key must be registered.")
self._service_token = service_token
self._service_token_hash = service_token_hash
self._delegation_keys = {key.deployment_id: key for key in delegation_keys}
self._jti_store = jti_store
# --- service token ---------------------------------------------------
def _check_service_token(self, headers: Mapping[str, str]) -> None:
authorization = headers.get("authorization", "")
scheme, _, supplied = authorization.partition(" ")
if scheme.lower() != "bearer" or not supplied.strip():
raise _unauthenticated(
"The request must carry 'Authorization: Bearer <BFF service token>'."
)
supplied = supplied.strip()
if self._service_token is not None:
if secrets.compare_digest(supplied, self._service_token):
return
elif self._service_token_hash is not None:
digest = hashlib.sha256(supplied.encode("utf-8")).hexdigest()
if secrets.compare_digest(digest, self._service_token_hash.lower()):
return
raise _unauthenticated("The BFF service token is invalid.")
# --- delegation JWT ----------------------------------------------------
def _decode_delegation(self, headers: Mapping[str, str]) -> dict:
token = headers.get("x-evo-actor", "")
if not token.strip():
raise _unauthenticated(
"The request must carry an 'X-Evo-Actor' delegation JWT."
)
try:
unverified = jwt.decode(token, options={"verify_signature": False})
except jwt.PyJWTError as exc:
raise _unauthenticated("The delegation JWT is malformed.") from exc
deployment_id = unverified.get("deployment_id")
key = (
self._delegation_keys.get(deployment_id)
if isinstance(deployment_id, str)
else None
)
if key is None:
raise _unauthenticated(
"The delegation JWT references an unregistered deployment."
)
try:
return jwt.decode(
token,
key.public_key,
algorithms=list(DELEGATION_ALGORITHMS),
audience=DELEGATION_AUDIENCE,
issuer=DELEGATION_ISSUER,
options={"require": list(_REQUIRED_CLAIMS)},
)
except jwt.PyJWTError as exc:
raise _unauthenticated(
"The delegation JWT failed signature, audience, issuer, or "
"lifetime validation."
) from exc
# --- combined check ------------------------------------------------------
def authenticate(
self,
headers: Mapping[str, str],
*,
required_scope: str,
require_thread_id: bool,
) -> ActorContext:
"""Verify both credentials and return the actor, or raise 401/403.
``required_scope`` is the section 7.2 scope the route demands;
thread-level routes additionally require a ``thread_id`` claim.
"""
self._check_service_token(headers)
claims = self._decode_delegation(headers)
if not isinstance(claims["sub"], str) or not claims["sub"].strip():
raise _unauthenticated("The delegation JWT 'sub' claim is invalid.")
if not isinstance(claims["jti"], str) or not claims["jti"].strip():
raise _unauthenticated("The delegation JWT 'jti' claim is invalid.")
lifetime = int(claims["exp"]) - int(claims["iat"])
if lifetime > DELEGATION_MAX_LIFETIME_SECONDS:
raise _unauthenticated(
"The delegation JWT lifetime exceeds "
f"{DELEGATION_MAX_LIFETIME_SECONDS} seconds."
)
scopes = claims["scopes"]
if not isinstance(scopes, list) or not all(
isinstance(scope, str) for scope in scopes
):
raise _unauthenticated("The delegation JWT 'scopes' claim is invalid.")
thread_id = claims.get("thread_id")
if require_thread_id and (
not isinstance(thread_id, str) or not thread_id.strip()
):
raise _unauthenticated(
"This route requires a delegation JWT 'thread_id' claim."
)
if thread_id is not None and not isinstance(thread_id, str):
raise _unauthenticated("The delegation JWT 'thread_id' claim is invalid.")
if required_scope not in scopes:
raise ModelRegistryError(
FORBIDDEN,
f"The actor lacks the required scope {required_scope!r}.",
)
registered = self._jti_store.register_delegation_jti(
claims["jti"], int(claims["exp"])
)
if not registered:
raise ModelRegistryError(
DELEGATION_REPLAYED,
"The delegation JWT has already been used.",
)
return ActorContext(
sub=claims["sub"].strip(),
scopes=frozenset(scopes),
deployment_id=claims["deployment_id"],
thread_id=thread_id.strip() if isinstance(thread_id, str) else None,
jti=claims["jti"],
)
@@ -0,0 +1,206 @@
"""EndpointPolicy: the SSRF allow rules for provider base URLs (section 4.3).
Public endpoints require https plus a hostname with an optional port. The
deny list covers loopback, private, link-local, multicast, unspecified, and
cloud-metadata (169.254.169.254) addresses. Local addresses over http or
https are only allowed when the normalized URL exactly matches a platform
``development_endpoints`` entry — an exact full-string match, never a
prefix or wildcard match.
Normalization lowercases the scheme and host, drops the default port, and
strips trailing ``/`` characters. URLs carrying user info, a fragment, or a
non-http(s) scheme are always rejected.
"""
from __future__ import annotations
import ipaddress
import typing
from collections.abc import Iterable
from urllib.parse import SplitResult, urlsplit
from .errors import ENDPOINT_NOT_ALLOWED, ModelRegistryError
from .schemas import DevelopmentEndpoint, EndpointPolicyPublic
_DEFAULT_PORTS = {"http": 80, "https": 443}
CLOUD_METADATA_IP = ipaddress.ip_address("169.254.169.254")
def denied_network_reason(ip: str) -> str | None:
"""Return the deny-list category for ``ip``, or ``None`` if allowed."""
try:
address = ipaddress.ip_address(ip)
except ValueError:
return None
if address == CLOUD_METADATA_IP:
return "cloud_metadata"
if address.is_loopback:
return "loopback"
if address.is_link_local:
return "link_local"
if address.is_multicast:
return "multicast"
if address.is_unspecified:
return "unspecified"
if address.is_private:
return "private"
return None
def _reject(message: str) -> ModelRegistryError:
return ModelRegistryError(ENDPOINT_NOT_ALLOWED, message)
def _host_text(host: str) -> str:
return f"[{host}]" if ":" in host else host
def _normalize_parts(
scheme: str,
host: str,
port: int | None,
path: str,
query: str,
) -> str:
netloc = _host_text(host)
if port is not None and port != _DEFAULT_PORTS[scheme]:
netloc = f"{netloc}:{port}"
normalized = f"{scheme}://{netloc}{path.rstrip('/')}"
if query:
normalized = f"{normalized}?{query}"
return normalized
class EndpointPolicy:
"""Validates provider base URLs against the platform endpoint rules."""
def __init__(
self, development_endpoints: Iterable[DevelopmentEndpoint] = ()
) -> None:
self._entries = list(development_endpoints)
self._registered: dict[str, DevelopmentEndpoint] = {}
for entry in self._entries:
try:
normalized = self._parse_and_normalize(entry.url)
except ModelRegistryError as exc:
raise ValueError(
f"Invalid development endpoint {entry.id!r}: {exc.message}"
) from exc
if normalized in self._registered:
raise ValueError(
f"Duplicate development endpoint registration: {entry.url!r}."
)
self._registered[normalized] = entry
# Request origins (scheme://host:port) implied by the registered
# entries. Base URLs must match an entry exactly, but the traffic an
# allowed base URL produces necessarily spans the whole origin.
self._registered_origins = {
self._origin_of(normalized) for normalized in self._registered
}
@staticmethod
def _origin_of(normalized: str) -> str:
parts = urlsplit(normalized)
return _normalize_parts(parts.scheme, parts.hostname or "", parts.port, "", "")
@staticmethod
def _parse_and_normalize(url: str) -> str:
parts = _checked_split(url)
return _normalize_parts(
parts.scheme.lower(),
parts.hostname or "",
parts.port,
parts.path,
parts.query,
)
def validate_base_url(self, url: str) -> str:
"""Return the normalized base URL, or raise ``ENDPOINT_NOT_ALLOWED``."""
parts = _checked_split(url)
scheme = parts.scheme.lower()
host = parts.hostname or ""
normalized = _normalize_parts(scheme, host, parts.port, parts.path, parts.query)
return self._validate(normalized, scheme, host, self._registered)
def validate_request_origin(self, origin: str) -> str:
"""Validate a request origin (``scheme://host[:port]``) before any I/O.
An origin is allowed when a registered development entry implies it or
when it satisfies the public https rules on its own.
"""
parts = _checked_split(origin)
scheme = parts.scheme.lower()
host = parts.hostname or ""
normalized = _normalize_parts(scheme, host, parts.port, "", "")
return self._validate(normalized, scheme, host, self._registered_origins)
@staticmethod
def _validate(
normalized: str,
scheme: str,
host: str,
allowed: typing.Container[str],
) -> str:
if normalized in allowed:
return normalized
if scheme != "https":
raise _reject(
"Base URL must use https unless it exactly matches a registered "
"development endpoint."
)
reason = denied_network_reason(host)
if reason is not None:
raise _reject(
f"Base URL host falls in the denied {reason} range and is not a "
"registered development endpoint."
)
return normalized
def is_development_endpoint(self, url: str) -> bool:
"""Return True when ``url`` normalizes to a registered entry."""
try:
normalized = self._parse_and_normalize(url)
except ModelRegistryError:
return False
return normalized in self._registered
def allows_denied_network(self, host: str, port: int) -> bool:
"""Return True when ``host:port`` is a registered development target.
Used by the transport at connect time: registered local endpoints may
resolve to deny-listed addresses; every other target is filtered.
The scheme is unknown at connect time, so both default-port spellings
are considered.
"""
host = host.lower()
for scheme, default_port in _DEFAULT_PORTS.items():
normalized = _normalize_parts(
scheme, host, None if port == default_port else port, "", ""
)
if normalized in self._registered_origins:
return True
return False
def public_view(self) -> EndpointPolicyPublic:
"""Return the section 9.1 browser-facing view of this policy."""
return EndpointPolicyPublic(development_endpoints=list(self._entries))
def _checked_split(url: str) -> SplitResult:
"""Split ``url`` and reject structurally disallowed forms."""
try:
parts = urlsplit(url)
# Accessing .port raises ValueError for out-of-range or text ports.
_ = parts.port
except ValueError as exc:
raise _reject("Base URL is malformed.") from exc
if parts.scheme.lower() not in _DEFAULT_PORTS:
raise _reject("Base URL must use the http or https scheme.")
if parts.username is not None or parts.password is not None:
raise _reject("Base URL must not contain user info.")
if parts.fragment:
raise _reject("Base URL must not contain a fragment.")
if not parts.hostname:
raise _reject("Base URL must include a hostname.")
return parts
+153
View File
@@ -0,0 +1,153 @@
"""Stable error codes and the unified error payload (design doc section 9.5).
Every failure response uses the same structure::
{ "code": "STABLE_ERROR_CODE", "message": "safe human-readable message",
"details": [{ "path": "...", "code": "..." }], "request_id": "..." }
Messages and details must never contain secrets, full provider errors, or
unredacted request headers.
"""
from __future__ import annotations
from typing import Any
from pydantic import BaseModel, Field
# 400 — malformed request bodies.
INVALID_REQUEST = "INVALID_REQUEST"
# 409 — version or idempotency conflicts.
REGISTRY_REVISION_CONFLICT = "REGISTRY_REVISION_CONFLICT"
THREAD_MODEL_SELECTION_CONFLICT = "THREAD_MODEL_SELECTION_CONFLICT"
RUN_REQUEST_CONFLICT = "RUN_REQUEST_CONFLICT"
SNAPSHOT_ALREADY_BOUND = "SNAPSHOT_ALREADY_BOUND"
SNAPSHOT_EXPIRED = "SNAPSHOT_EXPIRED"
MODEL_CONFIGURATION_CHANGED = "MODEL_CONFIGURATION_CHANGED"
# 401 — invalid identity.
DELEGATION_REPLAYED = "DELEGATION_REPLAYED"
UNAUTHENTICATED = "UNAUTHENTICATED"
# 403 — authenticated but insufficient scope.
FORBIDDEN = "FORBIDDEN"
# 404 — missing resource.
MODEL_NOT_FOUND = "MODEL_NOT_FOUND"
SNAPSHOT_NOT_FOUND = "SNAPSHOT_NOT_FOUND"
# 422 — configuration and runtime precondition failures.
MODEL_REGISTRY_NOT_READY = "MODEL_REGISTRY_NOT_READY"
MODEL_DISABLED = "MODEL_DISABLED"
MODEL_NOT_AVAILABLE = "MODEL_NOT_AVAILABLE"
MODEL_CONFIG_OUTSIDE_SNAPSHOT = "MODEL_CONFIG_OUTSIDE_SNAPSHOT"
CREDENTIAL_NOT_CONFIGURED = "CREDENTIAL_NOT_CONFIGURED"
CREDENTIAL_REJECTED = "CREDENTIAL_REJECTED"
RUN_CREDENTIAL_REVISION_UNAVAILABLE = "RUN_CREDENTIAL_REVISION_UNAVAILABLE"
AUTH_MODE_UNSUPPORTED = "AUTH_MODE_UNSUPPORTED"
ADAPTER_NOT_SUPPORTED = "ADAPTER_NOT_SUPPORTED"
ENDPOINT_NOT_ALLOWED = "ENDPOINT_NOT_ALLOWED"
CAPABILITY_UNSUPPORTED_BY_ADAPTER = "CAPABILITY_UNSUPPORTED_BY_ADAPTER"
MODEL_CAPABILITY_UNAVAILABLE = "MODEL_CAPABILITY_UNAVAILABLE"
UNSUPPORTED_RUNTIME_PARAMETER = "UNSUPPORTED_RUNTIME_PARAMETER"
MODEL_LIMITS_UNCONFIRMED = "MODEL_LIMITS_UNCONFIRMED"
IMAGE_MODEL_NOT_CHAT_MODEL = "IMAGE_MODEL_NOT_CHAT_MODEL"
CONTEXT_BUDGET_UNSATISFIABLE = "CONTEXT_BUDGET_UNSATISFIABLE"
PROVIDER_UNREACHABLE = "PROVIDER_UNREACHABLE"
# 422 — the request payload failed Pydantic schema validation.
VALIDATION_FAILED = "VALIDATION_FAILED"
# 500 — the platform security configuration is missing or unusable.
PLATFORM_CONFIG_MISSING = "PLATFORM_CONFIG_MISSING"
ERROR_HTTP_STATUS: dict[str, int] = {
INVALID_REQUEST: 400,
REGISTRY_REVISION_CONFLICT: 409,
THREAD_MODEL_SELECTION_CONFLICT: 409,
RUN_REQUEST_CONFLICT: 409,
SNAPSHOT_ALREADY_BOUND: 409,
SNAPSHOT_EXPIRED: 409,
MODEL_CONFIGURATION_CHANGED: 409,
DELEGATION_REPLAYED: 401,
UNAUTHENTICATED: 401,
FORBIDDEN: 403,
MODEL_NOT_FOUND: 404,
SNAPSHOT_NOT_FOUND: 404,
MODEL_REGISTRY_NOT_READY: 422,
MODEL_DISABLED: 422,
MODEL_NOT_AVAILABLE: 422,
MODEL_CONFIG_OUTSIDE_SNAPSHOT: 422,
CREDENTIAL_NOT_CONFIGURED: 422,
CREDENTIAL_REJECTED: 422,
RUN_CREDENTIAL_REVISION_UNAVAILABLE: 422,
AUTH_MODE_UNSUPPORTED: 422,
ADAPTER_NOT_SUPPORTED: 422,
ENDPOINT_NOT_ALLOWED: 422,
CAPABILITY_UNSUPPORTED_BY_ADAPTER: 422,
MODEL_CAPABILITY_UNAVAILABLE: 422,
UNSUPPORTED_RUNTIME_PARAMETER: 422,
MODEL_LIMITS_UNCONFIRMED: 422,
IMAGE_MODEL_NOT_CHAT_MODEL: 422,
CONTEXT_BUDGET_UNSATISFIABLE: 422,
PROVIDER_UNREACHABLE: 422,
VALIDATION_FAILED: 422,
PLATFORM_CONFIG_MISSING: 500,
}
class ErrorDetail(BaseModel):
"""One field-level failure entry inside an error payload."""
path: str
code: str
class ErrorPayload(BaseModel):
"""The single failure response shape defined in section 9.5."""
code: str
message: str
details: list[ErrorDetail] = Field(default_factory=list)
request_id: str
class ModelRegistryError(Exception):
"""An error carrying a stable code from the section 9.5 taxonomy.
``message`` must be safe to return to a browser: no secrets, no full
provider errors, no unredacted request headers.
"""
def __init__(
self,
code: str,
message: str,
*,
details: list[dict[str, str] | ErrorDetail] | None = None,
) -> None:
if code not in ERROR_HTTP_STATUS:
raise ValueError(f"Unknown model registry error code: {code!r}")
super().__init__(message)
self.code = code
self.message = message
self.details = [
detail
if isinstance(detail, ErrorDetail)
else ErrorDetail.model_validate(detail)
for detail in (details or [])
]
@property
def http_status(self) -> int:
return ERROR_HTTP_STATUS[self.code]
def payload(self, *, request_id: str = "") -> dict[str, Any]:
"""Return the section 9.5 failure response body."""
return ErrorPayload(
code=self.code,
message=self.message,
details=self.details,
request_id=request_id,
).model_dump()
+189
View File
@@ -0,0 +1,189 @@
"""build_chat_model: the single LangChain construction entry point (6.4).
``build_chat_model`` calls ``Adapter.build_request`` for the frozen
``ResolvedModelConfig`` and constructs the LangChain chat model from the
result. There is no second construction entry point: the interface takes no
``**kwargs`` and never merges caller values into snapshot values.
Credential handling: the resolved secret arrives via the ``credential``
parameter (resolved from ``auth_ref`` by the Resolver/API layer). This layer
never reads the credential store and never falls back to provider API-key
environment variables such as ``OPENAI_API_KEY`` — a missing credential
fails with ``CREDENTIAL_NOT_CONFIGURED`` before any model is constructed.
Every builder injects the Task 2 safe HTTP clients so provider egress keeps
passing the EndpointPolicy/SSRF defenses. Callers pass the sync client
(``build_safe_http_client``), the async client
(``build_safe_async_http_client``), or both; at least one is required:
- ``ChatOpenAI`` (openai, openai-compatible) takes ``http_client`` /
``http_async_client`` directly.
- ``ChatAnthropic`` builds its SDK clients from its own ``_client_params``
plus the safe clients (LangChain exposes no constructor hook for them).
- ``ChatOllama`` routes through the safe transports via
``sync_client_kwargs`` / ``async_client_kwargs`` — never the shared
``client_kwargs``, which langchain-ollama merges into both clients and
would poison the async client with a sync transport. The ollama SDK also
reads ``OLLAMA_API_KEY`` from the environment, so the ``Authorization``
header it may inject is stripped after construction — a ``mode=none``
adapter must not read, write, or fabricate API keys.
"""
from __future__ import annotations
import copy
from collections.abc import Callable
from typing import Any
import anthropic
import httpx
from langchain_anthropic import ChatAnthropic
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_ollama import ChatOllama
from langchain_openai import ChatOpenAI
from .adapters import get_adapter
from .schemas import ResolvedModelConfig
ChatModelBuilder = Callable[
[dict[str, Any], httpx.Client | None, httpx.AsyncClient | None], BaseChatModel
]
def _build_chat_openai(
options: dict[str, Any],
http_client: httpx.Client | None,
http_async_client: httpx.AsyncClient | None,
) -> BaseChatModel:
kwargs = dict(options)
if http_client is not None:
kwargs["http_client"] = http_client
if http_async_client is not None:
kwargs["http_async_client"] = http_async_client
model = ChatOpenAI(**kwargs)
# OpenAI-compatible endpoints only accept text in tool-role messages;
# this hoists read_file image blocks into a following user message so the
# model actually receives them (otherwise the gateway strips the image
# and the model reports an empty result).
from ..llm.patches import _patch_openai_compat_content
_patch_openai_compat_content(model, hoist_tool_media=True)
return model
def _build_chat_anthropic(
options: dict[str, Any],
http_client: httpx.Client | None,
http_async_client: httpx.AsyncClient | None,
) -> BaseChatModel:
model = ChatAnthropic(**options)
# ChatAnthropic exposes no http_client constructor argument; it builds
# ``_client``/``_async_client`` (cached_property) from ``_client_params``.
# Seeding the cached properties with SDK clients wrapped around the safe
# transports keeps anthropic egress inside the EndpointPolicy defenses.
client_params = model._client_params
if http_client is not None:
model.__dict__["_client"] = anthropic.Client(
**client_params, http_client=http_client
)
if http_async_client is not None:
model.__dict__["_async_client"] = anthropic.AsyncClient(
**client_params, http_client=http_async_client
)
return model
def _transport_of(http_client: httpx.Client | httpx.AsyncClient) -> Any:
transport = getattr(http_client, "_transport", None)
if transport is None: # pragma: no cover - defensive
raise TypeError(
"http_client must be built by build_safe_http_client or "
"build_safe_async_http_client."
)
return transport
def _build_chat_ollama(
options: dict[str, Any],
http_client: httpx.Client | None,
http_async_client: httpx.AsyncClient | None,
) -> BaseChatModel:
options = copy.deepcopy(options)
# langchain-ollama merges the shared client_kwargs into BOTH clients, so
# the safe transports must go through the per-direction kwargs: a sync
# transport in client_kwargs would poison the async client (its httpx
# async calls would hit handle_async_request on a sync transport).
if http_client is not None:
options["sync_client_kwargs"] = {"transport": _transport_of(http_client)}
if http_async_client is not None:
options["async_client_kwargs"] = {"transport": _transport_of(http_async_client)}
model = ChatOllama(**options)
# ollama-python silently adds an Authorization header from OLLAMA_API_KEY.
# Phase-1 ollama contracts only allow auth mode "none", which must not
# read, write, or fabricate API keys — strip any injected header.
for ollama_client in (model._client, model._async_client):
ollama_client._client.headers.pop("authorization", None)
return model
# chat_model name (from the contract's connection spec) -> builder. Tests
# substitute fakes here to capture construction arguments per adapter.
CHAT_MODEL_BUILDERS: dict[str, ChatModelBuilder] = {
"ChatOpenAI": _build_chat_openai,
"ChatAnthropic": _build_chat_anthropic,
"ChatOllama": _build_chat_ollama,
}
def build_chat_model(
resolved_config: ResolvedModelConfig,
http_client: httpx.Client | None = None,
http_async_client: httpx.AsyncClient | None = None,
*,
credential: str | None = None,
) -> BaseChatModel:
"""Construct the LangChain chat model for a frozen run configuration.
The snapshot's ``adapter_spec_revision`` pins the contract: if that
revision no longer exists the call fails with ``ADAPTER_NOT_SUPPORTED``
instead of silently substituting a newer contract.
"""
if not isinstance(resolved_config, ResolvedModelConfig):
raise TypeError(
"resolved_config must be a ResolvedModelConfig, got "
f"{type(resolved_config).__name__}."
)
if http_client is None and http_async_client is None:
raise TypeError(
"At least one of http_client / http_async_client is required; "
"provider egress must go through the safe transports."
)
if http_client is not None and not isinstance(http_client, httpx.Client):
raise TypeError(
f"http_client must be an httpx.Client, got {type(http_client).__name__}."
)
if http_async_client is not None and not isinstance(
http_async_client, httpx.AsyncClient
):
raise TypeError(
"http_async_client must be an httpx.AsyncClient, got "
f"{type(http_async_client).__name__}."
)
adapter = get_adapter(
resolved_config.adapter_id,
resolved_config.upstream_model_id,
spec_revision=resolved_config.adapter_spec_revision,
)
built = adapter.build_request(resolved_config, credential=credential)
overlap = set(built.client_options) & set(built.request_options)
if overlap:
raise ValueError(
"Adapter contract client_options/request_options keys overlap: "
f"{sorted(overlap)}."
)
options = {**built.client_options, **built.request_options}
connection = adapter.spec.connection
if connection is None: # pragma: no cover - built-in contracts always set it
raise TypeError(f"Adapter {adapter.spec.adapter_id!r} has no connection spec.")
builder = CHAT_MODEL_BUILDERS[connection.chat_model]
return builder(options, http_client, http_async_client)
+26
View File
@@ -0,0 +1,26 @@
"""Configuration hashing for model verification records (section 4.3).
The hash covers the provider's adapter, base URL, the model's
``upstream_model_id``, all runtime parameters, declared capabilities, and
limits. Any change to these inputs invalidates prior verification records.
"""
from __future__ import annotations
import hashlib
import json
from .schemas import ModelConfig, ProviderConfig
def configuration_hash(provider: ProviderConfig, model: ModelConfig) -> str:
"""Return the normalized SHA-256 hash of one model's full configuration."""
payload = {
"adapter": provider.adapter,
"base_url": provider.base_url,
"upstream_model_id": model.upstream_model_id,
"provider_runtime": provider.runtime.model_dump(mode="json"),
"model_runtime": model.runtime.model_dump(mode="json"),
}
encoded = json.dumps(payload, sort_keys=True, separators=(",", ":"))
return hashlib.sha256(encoded.encode("utf-8")).hexdigest()
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,37 @@
"""Compile-time chat-model placeholder for a bootstrap registry.
Graph construction must tolerate a bootstrap registry — the Config API has
to serve so operators can configure the first model (design doc section
10), and only run creation is forbidden. Every graph therefore binds this
placeholder when no registry default exists; it supports the build-time
surface (``bind_tools``, profile inspection, middleware wiring) but raises
``MODEL_REGISTRY_NOT_READY`` on the first model call, so no run can ever
silently fall back to an implicit model.
"""
from typing import Any
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import BaseMessage
from langchain_core.outputs import ChatResult
from .errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError
class RegistryNotReadyChatModel(BaseChatModel):
"""Placeholder that fails every call with MODEL_REGISTRY_NOT_READY."""
detail: str
@property
def _llm_type(self) -> str:
return "evoscientist-registry-not-ready"
def _generate(
self,
messages: list[BaseMessage],
stop: list[str] | None = None,
run_manager: Any = None,
**kwargs: Any,
) -> ChatResult:
raise ModelRegistryError(MODEL_REGISTRY_NOT_READY, self.detail)
+171
View File
@@ -0,0 +1,171 @@
"""Platform security configuration from ``config.yaml`` (design doc 4.2).
The Config API cannot serve without an explicitly configured BFF service
token and at least one registered WebUI delegation public key — silently
falling back to defaults would mean accepting unauthenticated requests.
A missing or incomplete security section raises :class:`PlatformConfigError`;
the HTTP layer maps it to ``500 PLATFORM_CONFIG_MISSING`` with a safe
message.
Recognized ``config.yaml`` fields (all other platform fields stay in
``EvoScientistConfig`` and are untouched):
- ``bff_service_token``: the shared BFF → EvoScientist bearer token
(plaintext; development convenience).
- ``bff_service_token_hash``: hex SHA-256 of the token; used when the
plaintext field is absent so production deployments never persist it.
- ``webui_delegation_public_keys``: ``[{deployment_id, public_key}]`` entries
registering each WebUI's delegation JWT signing key (PEM, ES256/RS256).
- ``development_endpoints``: ``[{id, url, label}]`` EndpointPolicy entries.
- ``local_deployment_id``: deployment ID for local (non-BFF) entry points;
defaults to ``"local"``.
- ``model_runtime_db``: optional explicit path of ``model-runtime.sqlite3``.
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import yaml
from EvoScientist.config.settings import get_config_path
from .schemas import DevelopmentEndpoint
DEFAULT_LOCAL_DEPLOYMENT_ID = "local"
class PlatformConfigError(RuntimeError):
"""Raised when the platform security configuration cannot be used."""
@dataclass(frozen=True)
class DelegationPublicKey:
"""One registered WebUI delegation signing key (section 7.3)."""
deployment_id: str
public_key: str
@dataclass(frozen=True)
class PlatformSecurityConfig:
"""The parsed security-relevant platform fields of ``config.yaml``."""
bff_service_token: str | None = None
bff_service_token_hash: str | None = None
webui_delegation_public_keys: tuple[DelegationPublicKey, ...] = ()
development_endpoints: tuple[DevelopmentEndpoint, ...] = ()
local_deployment_id: str = DEFAULT_LOCAL_DEPLOYMENT_ID
model_runtime_db: Path | None = None
def _optional_string(data: dict[str, Any], key: str) -> str | None:
value = data.get(key)
if value is None:
return None
if not isinstance(value, str) or not value.strip():
raise PlatformConfigError(f"config.yaml field {key!r} must be a string.")
return value.strip()
def _parse_delegation_keys(raw: Any) -> tuple[DelegationPublicKey, ...]:
if raw is None:
return ()
if not isinstance(raw, list):
raise PlatformConfigError(
"config.yaml field 'webui_delegation_public_keys' must be a list."
)
keys: list[DelegationPublicKey] = []
seen: set[str] = set()
for index, entry in enumerate(raw):
if not isinstance(entry, dict):
raise PlatformConfigError(
f"webui_delegation_public_keys[{index}] must be an object."
)
deployment_id = entry.get("deployment_id")
public_key = entry.get("public_key")
if (
not isinstance(deployment_id, str)
or not deployment_id.strip()
or not isinstance(public_key, str)
or "PUBLIC KEY" not in public_key
):
raise PlatformConfigError(
f"webui_delegation_public_keys[{index}] requires a non-empty "
"'deployment_id' and a PEM 'public_key'."
)
deployment_id = deployment_id.strip()
if deployment_id in seen:
raise PlatformConfigError(
f"Duplicate delegation deployment_id {deployment_id!r}."
)
seen.add(deployment_id)
keys.append(
DelegationPublicKey(deployment_id=deployment_id, public_key=public_key)
)
return tuple(keys)
def _parse_development_endpoints(raw: Any) -> tuple[DevelopmentEndpoint, ...]:
if raw is None:
return ()
if not isinstance(raw, list):
raise PlatformConfigError(
"config.yaml field 'development_endpoints' must be a list."
)
try:
return tuple(DevelopmentEndpoint.model_validate(entry) for entry in raw)
except ValueError as exc:
raise PlatformConfigError(
f"Invalid development_endpoints entry: {exc}."
) from exc
def load_platform_security_config(
config_path: str | Path | None = None,
) -> PlatformSecurityConfig:
"""Load the security platform fields; raise when auth cannot be served."""
path = Path(config_path) if config_path is not None else get_config_path()
if not path.exists():
raise PlatformConfigError(
f"config.yaml not found at {path}; configure bff_service_token "
"and webui_delegation_public_keys first."
)
try:
with open(path, encoding="utf-8") as handle:
data = yaml.safe_load(handle) or {}
except yaml.YAMLError as exc:
raise PlatformConfigError(f"config.yaml is not valid YAML: {exc}.") from exc
if not isinstance(data, dict):
raise PlatformConfigError("config.yaml must contain a mapping at the top.")
token = _optional_string(data, "bff_service_token")
token_hash = _optional_string(data, "bff_service_token_hash")
if token is None and token_hash is None:
raise PlatformConfigError(
"config.yaml must set 'bff_service_token' or "
"'bff_service_token_hash'; the Config API refuses to serve "
"without an explicitly configured BFF service token."
)
keys = _parse_delegation_keys(data.get("webui_delegation_public_keys"))
if not keys:
raise PlatformConfigError(
"config.yaml must register at least one entry in "
"'webui_delegation_public_keys'."
)
raw_db = _optional_string(data, "model_runtime_db")
return PlatformSecurityConfig(
bff_service_token=token,
bff_service_token_hash=token_hash,
webui_delegation_public_keys=keys,
development_endpoints=_parse_development_endpoints(
data.get("development_endpoints")
),
local_deployment_id=(
_optional_string(data, "local_deployment_id") or DEFAULT_LOCAL_DEPLOYMENT_ID
),
model_runtime_db=None if raw_db is None else Path(raw_db).expanduser(),
)
@@ -0,0 +1,460 @@
"""Provider test execution: the only model verification entry point (9.4).
``ProviderTester.run`` implements the section 9.4 flow end to end:
1. The ``ModelRef`` must already be saved in the registry (browser free
drafts are rejected with ``MODEL_NOT_FOUND``) and its base URL must still
pass the EndpointPolicy — the save-time check is re-applied because the
test is a network operation (SSRF defense, section 4.3).
2. ``resolve_for_test`` freezes a temporary ``ResolvedModelConfig``; the
credential secret is resolved from the store against the frozen
``auth_ref`` on every test — never from a process cache (section 5.2).
3. ``build_chat_model`` constructs the LangChain model with both safe HTTP
clients, so all egress shares the SafeHttpTransport defenses; adapters
whose contract keeps retries out of the request (empty ``target_name``,
e.g. ollama) get them through the safe client builder instead.
4. One minimal real chat call runs, followed by one minimal probe per
declared capability (section 6.2). A failing probe only marks that
capability unverified; it never blocks the others.
5. The ``model_verifications`` five-tuple record is upserted inside a
transaction that re-checks ``expected_registry_revision`` and the model's
``configuration_hash``; a concurrent change aborts with
``409 MODEL_CONFIGURATION_CHANGED`` and records nothing (section 9.4).
Errors are classified into the stable section 9.5 codes; messages are
static strings and never carry secrets, stack traces, or raw provider
responses. ``effective_request_options`` reuses ``adapter.build_request``
output with the connection and credential fields stripped, proving the
configuration, the snapshot, and the actual request agree (section 6.4).
"""
from __future__ import annotations
import asyncio
import time
from collections.abc import Callable, Iterator
from dataclasses import dataclass
from typing import Any
import httpx
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import HumanMessage
from langchain_core.tools import tool
from .adapters import Adapter, BuiltRequest, get_adapter
from .endpoint_policy import EndpointPolicy
from .errors import (
CREDENTIAL_REJECTED,
MODEL_CONFIGURATION_CHANGED,
MODEL_NOT_AVAILABLE,
MODEL_NOT_FOUND,
PROVIDER_UNREACHABLE,
ModelRegistryError,
)
from .factory import build_chat_model
from .hashing import configuration_hash
from .resolver import ModelRegistryResolver
from .safe_transport import build_safe_async_http_client, build_safe_http_client
from .schemas import (
Capabilities,
ModelAvailability,
ModelRef,
ResolvedModelConfig,
)
from .store import ModelRuntimeStore
# Provider tests must fail fast; the resolved contract timeout (up to 600s)
# is capped at this configurable, deliberately short default (section 9.4).
DEFAULT_TEST_TIMEOUT_SECONDS = 30
SyncClientBuilder = Callable[[float, int], httpx.Client]
AsyncClientBuilder = Callable[[float, int], httpx.AsyncClient]
_CAPABILITY_NAMES = ("tools", "vision", "structured_output")
# A 1x1 transparent PNG for the vision capability probe (section 6.2).
_PIXEL_PNG_B64 = (
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8"
"z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg=="
)
_STRUCTURED_PROBE_SCHEMA = {
"title": "ProviderTestProbe",
"description": "Trivial structured-output capability probe.",
"type": "object",
"properties": {"answer": {"type": "string"}},
"required": ["answer"],
}
@tool
def _ping_tool(text: str) -> str:
"""Echo the input text; exists only as a trivial tools probe."""
return text
@dataclass(frozen=True)
class ProviderTestResult:
"""The successful section 9.4 test outcome; carries no secrets."""
ok: bool
registry_revision: int
model_ref: ModelRef
adapter_spec_revision: int
effective_request_options: dict[str, Any]
model_status: ModelAvailability
latency_ms: int
@dataclass(frozen=True)
class _ExecutionOutcome:
"""The network phase result: per-capability verdicts and timing."""
verified_capabilities: dict[str, bool]
latency_ms: int
def _walk_causes(exc: BaseException) -> Iterator[BaseException]:
"""Yield ``exc`` and its ``__cause__``/``__context__`` chain, cycle-safe."""
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
yield current
current = current.__cause__ or current.__context__
def _status_code(exc: BaseException) -> int | None:
"""Find an HTTP status anywhere in the exception chain.
Both the OpenAI and Anthropic SDKs expose ``status_code`` on their
status errors; plain httpx surfaces it via ``response``.
"""
for err in _walk_causes(exc):
status = getattr(err, "status_code", None)
if isinstance(status, int):
return status
response = getattr(err, "response", None)
if isinstance(response, httpx.Response):
return response.status_code
return None
def _classify_provider_error(exc: BaseException) -> ModelRegistryError:
"""Map any provider-call failure onto the stable section 9.5 codes.
SDKs wrap transport failures (including the EndpointPolicy rejection
raised by SafeHttpTransport) into their own connection errors, so the
``ModelRegistryError`` search must walk the whole cause chain.
"""
for err in _walk_causes(exc):
if isinstance(err, ModelRegistryError):
return err
status = _status_code(exc)
if status in (401, 403):
return ModelRegistryError(
CREDENTIAL_REJECTED, "The provider rejected the configured credential."
)
if status == 404:
return ModelRegistryError(
MODEL_NOT_FOUND, "The provider reports that the model does not exist."
)
for err in _walk_causes(exc):
if isinstance(err, httpx.TransportError):
return ModelRegistryError(
PROVIDER_UNREACHABLE, "The provider endpoint could not be reached."
)
if status is not None and status >= 500:
return ModelRegistryError(
PROVIDER_UNREACHABLE, "The provider endpoint could not be reached."
)
return ModelRegistryError(MODEL_NOT_AVAILABLE, "The provider test request failed.")
def _probe_tools(chat_model: BaseChatModel) -> None:
"""Bind a trivial tool; the request must carry it and not error (6.2)."""
bound = chat_model.bind_tools([_ping_tool])
messages = [HumanMessage(content="Call the ping tool with text 'probe'.")]
request_payload = getattr(bound, "_get_request_payload", None)
if callable(request_payload):
# The bound kwargs hold the formatted tools; ``_get_request_payload``
# only folds them into the payload when they are passed explicitly.
bound_kwargs = getattr(bound, "kwargs", None)
payload = request_payload(messages, **(bound_kwargs or {}))
if "tools" not in payload:
raise RuntimeError("The bound request payload does not contain tools.")
bound.invoke(messages)
def _probe_structured_output(chat_model: BaseChatModel) -> None:
"""Request a JSON-schema structured response (section 6.2)."""
try:
structured = chat_model.with_structured_output(
_STRUCTURED_PROBE_SCHEMA, method="json_schema"
)
except (TypeError, ValueError):
# Adapters without a json_schema method fall back to their native
# structured-output mechanism.
structured = chat_model.with_structured_output(_STRUCTURED_PROBE_SCHEMA)
structured.invoke([HumanMessage(content='Answer with {"answer": "ok"}.')])
def _probe_vision(chat_model: BaseChatModel) -> None:
"""Send a 1x1 test image (section 6.2)."""
chat_model.invoke(
[
HumanMessage(
content=[
{"type": "text", "text": "Describe this image in one word."},
{
"type": "image_url",
"image_url": {"url": f"data:image/png;base64,{_PIXEL_PNG_B64}"},
},
]
)
]
)
_CAPABILITY_PROBES = {
"tools": _probe_tools,
"structured_output": _probe_structured_output,
"vision": _probe_vision,
}
def _redacted_options(adapter: Adapter, built: BuiltRequest) -> dict[str, Any]:
"""The effective request parameters with connection/secret fields removed.
Reuses the exact ``adapter.build_request`` output the model was built
from, minus the connection identity (``model``/``base_url``) and every
auth target the contract declares (``api_key``/``default_headers``).
"""
options = dict(built.client_options)
connection = adapter.spec.connection
if connection is not None:
options.pop(connection.model_field, None)
options.pop(connection.base_url_field, None)
for auth_spec in adapter.spec.auth_specs.values():
if auth_spec.target == "client_option" and auth_spec.target_name:
options.pop(auth_spec.target_name, None)
elif auth_spec.target == "request_header":
options.pop("default_headers", None)
options.update(built.request_options)
return options
class ProviderTester:
"""Runs provider tests and records their verdicts (section 9.4)."""
def __init__(
self,
store: ModelRuntimeStore,
resolver: ModelRegistryResolver,
endpoint_policy: EndpointPolicy,
*,
timeout_seconds: int = DEFAULT_TEST_TIMEOUT_SECONDS,
sync_client_builder: SyncClientBuilder | None = None,
async_client_builder: AsyncClientBuilder | None = None,
) -> None:
if timeout_seconds <= 0:
raise ValueError("timeout_seconds must be positive.")
self._store = store
self._resolver = resolver
self._endpoint_policy = endpoint_policy
self._timeout_seconds = timeout_seconds
self._sync_client_builder = sync_client_builder or (
lambda timeout, retries: build_safe_http_client(
endpoint_policy, timeout=timeout, retries=retries
)
)
self._async_client_builder = async_client_builder or (
lambda timeout, retries: build_safe_async_http_client(
endpoint_policy, timeout=timeout, retries=retries
)
)
def run(
self, *, expected_registry_revision: int, model_ref: ModelRef
) -> ProviderTestResult:
"""Test one saved model and record the verdict.
Provider-call failures upsert a ``failed`` record (when the registry
is unchanged) and then raise the classified section 9.5 error;
precondition failures raise before any network I/O or record write.
"""
registry = self._store.load_registry()
provider = registry.find_provider(model_ref.provider_id)
model = provider.find_model(model_ref.model_key) if provider else None
if provider is None or model is None:
raise ModelRegistryError(
MODEL_NOT_FOUND,
"The model reference does not exist in the saved registry; "
"only saved models can be tested.",
details=[{"path": "model_ref", "code": MODEL_NOT_FOUND}],
)
# Cheap pre-check: a stale caller view fails before any network I/O.
# The transactional re-check at record time stays authoritative.
if registry.revision != expected_registry_revision:
raise ModelRegistryError(
MODEL_CONFIGURATION_CHANGED,
"The registry changed since the test was requested; reload and retry.",
details=[
{
"path": "expected_registry_revision",
"code": MODEL_CONFIGURATION_CHANGED,
}
],
)
# Re-apply the save-time URL check: the test is itself an SSRF
# channel if it trusted a stored URL blindly (section 4.3).
self._endpoint_policy.validate_base_url(provider.base_url)
resolved = self._resolver.resolve_for_test(model_ref)
credential = self._resolve_credential(resolved)
adapter = get_adapter(
resolved.adapter_id,
resolved.upstream_model_id,
spec_revision=resolved.adapter_spec_revision,
)
built = adapter.build_request(resolved, credential=credential)
effective_request_options = _redacted_options(adapter, built)
config_hash = configuration_hash(provider, model)
try:
outcome = self._execute(
resolved, adapter, credential, model.runtime.declared_capabilities
)
except ModelRegistryError as exc:
self._record(
expected_registry_revision,
resolved,
config_hash,
result="failed",
verified_capabilities=dict.fromkeys(_CAPABILITY_NAMES, False),
error_code=exc.code,
)
raise
self._record(
expected_registry_revision,
resolved,
config_hash,
result="passed",
verified_capabilities=outcome.verified_capabilities,
error_code=None,
)
return ProviderTestResult(
ok=True,
registry_revision=expected_registry_revision,
model_ref=model_ref,
adapter_spec_revision=resolved.adapter_spec_revision,
effective_request_options=effective_request_options,
model_status=self._model_status(model_ref),
latency_ms=outcome.latency_ms,
)
def _resolve_credential(self, resolved: ResolvedModelConfig) -> str | None:
"""Resolve the frozen credential version per test; never cached (5.2)."""
auth_ref = resolved.auth_ref
if auth_ref.mode == "none":
return None
if auth_ref.credential_id is None or auth_ref.credential_revision is None:
return None
return self._store.resolve_credential(
auth_ref.credential_id, auth_ref.credential_revision
)
def _execute(
self,
resolved: ResolvedModelConfig,
adapter: Adapter,
credential: str | None,
declared: Capabilities,
) -> _ExecutionOutcome:
"""Build the model and run the minimal call plus capability probes."""
# Contracts that validate max_retries without sending it (empty
# target_name, e.g. ollama) have them enforced by the safe transport.
retries = (
resolved.client_options.max_retries
if adapter.spec.parameters["max_retries"].target_name == ""
else 0
)
timeout = float(
min(resolved.client_options.timeout_seconds, self._timeout_seconds)
)
sync_client = self._sync_client_builder(timeout, retries)
async_client = self._async_client_builder(timeout, retries)
try:
chat_model = build_chat_model(
resolved, sync_client, async_client, credential=credential
)
started = time.monotonic()
self._minimal_chat(chat_model)
verified = dict.fromkeys(_CAPABILITY_NAMES, False)
for name in _CAPABILITY_NAMES:
if getattr(declared, name):
verified[name] = self._probe(name, chat_model)
latency_ms = int((time.monotonic() - started) * 1000)
return _ExecutionOutcome(
verified_capabilities=verified, latency_ms=latency_ms
)
finally:
sync_client.close()
# The tester runs in a worker thread without a running loop.
asyncio.run(async_client.aclose())
@staticmethod
def _minimal_chat(chat_model: BaseChatModel) -> None:
"""One low-cost chat request; failures classify into section 9.5 codes."""
try:
chat_model.invoke([HumanMessage(content="Reply with exactly: ok")])
except ModelRegistryError:
raise
except Exception as exc:
raise _classify_provider_error(exc) from exc
@staticmethod
def _probe(name: str, chat_model: BaseChatModel) -> bool:
"""Run one capability probe; any failure marks only that capability."""
try:
_CAPABILITY_PROBES[name](chat_model)
except Exception:
return False
return True
def _record(
self,
expected_registry_revision: int,
resolved: ResolvedModelConfig,
config_hash: str,
*,
result: str,
verified_capabilities: dict[str, bool],
error_code: str | None,
) -> None:
"""The guarded five-tuple upsert; raises 409 when the config changed."""
self._store.record_model_verification_if_current(
expected_registry_revision=expected_registry_revision,
provider_id=resolved.model_ref.provider_id,
model_key=resolved.model_ref.model_key,
configuration_hash=config_hash,
credential_revision=resolved.auth_ref.credential_revision or 0,
adapter_spec_revision=resolved.adapter_spec_revision,
result=result,
verified_capabilities=verified_capabilities,
error_code=error_code,
)
def _model_status(self, model_ref: ModelRef) -> ModelAvailability:
"""Recompute the backend ModelAvailability after the record write."""
registry = self._store.load_registry()
availability = self._resolver.compute_availability(
registry,
self._store.list_model_verifications(),
credential_revisions=self._store.list_credential_revisions(),
)
for item in availability:
if item.model_ref == model_ref:
return item
raise ModelRegistryError( # pragma: no cover - judged from same registry
MODEL_NOT_FOUND, "The model reference does not exist in the registry."
)
+471
View File
@@ -0,0 +1,471 @@
"""ModelRegistryResolver: the single registry resolution path (section 8.1).
``resolve(ModelRef, role)`` validates the provider, model, credentials,
capabilities, limits, and input budget, then freezes a complete
``ResolvedModelConfig`` (section 6.4). ``resolve_for_test`` relaxes the two
run-time visibility gates — the model need not be enabled and no passing
verification record is required (section 9.4) — because the provider test is
exactly what produces that verification; every other validation still runs.
``compute_availability`` implements the section 4.3
six-state judgement order and is the only availability computation — callers
must not derive state from ``enabled`` flags or test timestamps on their own.
The resolver never reads ``secret_value``: a credential check resolves the
pointer's ``current_revision`` and freezes it into ``auth_ref``. Secrets are
resolved per provider call by the snapshot layer against that frozen
revision (section 5.2).
"""
from __future__ import annotations
from collections.abc import Iterable, Mapping
from datetime import UTC, datetime
from typing import Any, get_args
from .adapters import (
adapter_specs,
compute_effective_capabilities,
find_adapter_spec,
resolve_parameters,
)
from .errors import (
CONTEXT_BUDGET_UNSATISFIABLE,
CREDENTIAL_NOT_CONFIGURED,
MODEL_DISABLED,
MODEL_LIMITS_UNCONFIRMED,
MODEL_NOT_AVAILABLE,
MODEL_NOT_FOUND,
ModelRegistryError,
)
from .hashing import configuration_hash
from .schemas import (
AdapterParameterSpec,
AuthRef,
Capabilities,
ClientOptions,
FixedReserves,
InputBudget,
ModelAvailability,
ModelConfig,
ModelRef,
ModelRole,
ProviderConfig,
ReasoningEffort,
RegistryV4,
RequestOptions,
ResolvedModelConfig,
SamplingOverride,
VerificationInfo,
)
from .store import ModelRuntimeStore
_CAPABILITY_NAMES = ("tools", "vision", "structured_output")
_MODEL_ROLES = get_args(ModelRole)
# Availability reason codes (section 4.3 `reason_code` strings).
REASON_PROVIDER_DISABLED = "PROVIDER_DISABLED"
REASON_NO_ADAPTER_CONTRACT = "NO_ADAPTER_CONTRACT"
REASON_MODEL_DISABLED = "MODEL_DISABLED"
REASON_VERIFICATION_FAILED = "VERIFICATION_FAILED"
REASON_VERIFICATION_STALE = "VERIFICATION_STALE"
def _rfc3339(epoch_seconds: int) -> str:
return datetime.fromtimestamp(epoch_seconds, UTC).strftime("%Y-%m-%dT%H:%M:%SZ")
def _check_role(role: str) -> None:
if role not in _MODEL_ROLES:
raise ValueError(
f"Unknown model role {role!r}; expected one of {list(_MODEL_ROLES)}."
)
class ModelRegistryResolver:
"""Validates ModelRefs and freezes per-run resolved configurations."""
def __init__(
self,
store: ModelRuntimeStore,
*,
specs: Iterable[AdapterParameterSpec] | None = None,
) -> None:
self._store = store
self._specs = tuple(adapter_specs() if specs is None else specs)
@property
def specs(self) -> tuple[AdapterParameterSpec, ...]:
"""The adapter contracts this resolver matches against."""
return self._specs
def resolve(
self,
model_ref: ModelRef,
role: ModelRole = "primary",
*,
registry: RegistryV4 | None = None,
reasoning_effort_override: ReasoningEffort | None = None,
sampling_override: SamplingOverride | None = None,
) -> ResolvedModelConfig:
"""Resolve an enabled, verified model into its frozen run config."""
return self._resolve(
model_ref,
role,
registry=registry,
require_enabled=True,
reasoning_effort_override=reasoning_effort_override,
sampling_override=sampling_override,
)
def resolve_for_test(self, model_ref: ModelRef) -> ResolvedModelConfig:
"""Resolve for a provider test (section 9.4).
The run-time visibility gates — the model must be enabled and a
passing verification must already exist — are relaxed, because the
provider test is exactly what produces that verification. Auth,
parameter, capability, limits, and budget validation all still run.
"""
return self._resolve(
model_ref,
"primary",
registry=None,
require_enabled=False,
require_verified=False,
)
def _resolve(
self,
model_ref: ModelRef,
role: ModelRole,
*,
registry: RegistryV4 | None,
require_enabled: bool,
require_verified: bool = True,
reasoning_effort_override: ReasoningEffort | None = None,
sampling_override: SamplingOverride | None = None,
) -> ResolvedModelConfig:
_check_role(role)
if registry is None:
registry = self._store.load_registry()
provider = registry.find_provider(model_ref.provider_id)
if provider is None:
raise ModelRegistryError(
MODEL_NOT_FOUND,
f"Provider {model_ref.provider_id!r} does not exist.",
details=[{"path": "provider_id", "code": MODEL_NOT_FOUND}],
)
model = provider.find_model(model_ref.model_key)
if model is None:
raise ModelRegistryError(
MODEL_NOT_FOUND,
f"Model {model_ref.model_key!r} does not exist on provider "
f"{model_ref.provider_id!r}.",
details=[{"path": "model_key", "code": MODEL_NOT_FOUND}],
)
if not provider.enabled:
raise ModelRegistryError(
MODEL_NOT_AVAILABLE,
f"Provider {provider.id!r} is disabled.",
details=[{"path": "provider.enabled", "code": MODEL_NOT_AVAILABLE}],
)
if require_enabled and not model.enabled:
raise ModelRegistryError(
MODEL_DISABLED,
f"Model {model_ref.model_key!r} is disabled.",
details=[{"path": "model.enabled", "code": MODEL_DISABLED}],
)
spec = find_adapter_spec(
provider.adapter, model.upstream_model_id, specs=self._specs
)
if spec is None:
raise ModelRegistryError(
MODEL_NOT_AVAILABLE,
"No adapter contract matches "
f"{provider.adapter!r}/{model.upstream_model_id!r}; the model "
"may only remain 'configured'.",
details=[{"path": "model.adapter", "code": MODEL_NOT_AVAILABLE}],
)
# Save-time contract checks re-applied at resolve time: auth mode,
# credential reference, declared capabilities, and parameters.
parameters = resolve_parameters(
provider,
model,
spec,
reasoning_effort_override=reasoning_effort_override,
sampling_override=sampling_override,
)
auth_spec = spec.auth_specs[provider.auth.mode]
credential_revision: int | None = None
if auth_spec.credential_required and provider.auth.credential_id is not None:
credential_revision = self._store.current_credential_revision(
provider.auth.credential_id
)
if credential_revision is None:
raise ModelRegistryError(
CREDENTIAL_NOT_CONFIGURED,
"The credential referenced by this provider is not configured.",
details=[
{
"path": "auth.credential_id",
"code": CREDENTIAL_NOT_CONFIGURED,
}
],
)
auth_ref = AuthRef(
mode=provider.auth.mode,
credential_id=provider.auth.credential_id,
credential_revision=credential_revision,
)
# Run gate (relaxed for provider tests): the current verification
# five-tuple must exist and have passed; any configuration or
# credential change invalidates it (4.3).
record = self._store.get_model_verification(
provider_id=provider.id,
model_key=model.key,
configuration_hash=configuration_hash(provider, model),
credential_revision=credential_revision or 0,
adapter_spec_revision=spec.spec_revision,
)
if require_verified and (record is None or record["result"] != "passed"):
raise ModelRegistryError(
MODEL_NOT_AVAILABLE,
"The model has no passing verification for its current configuration.",
details=[{"path": "model.verification", "code": MODEL_NOT_AVAILABLE}],
)
if model.runtime.limits_status != "confirmed":
raise ModelRegistryError(
MODEL_LIMITS_UNCONFIRMED,
"The model's context limits are not confirmed.",
details=[
{"path": "runtime.limits_status", "code": MODEL_LIMITS_UNCONFIRMED}
],
)
budget = self._resolve_budget(model, parameters.max_output_tokens)
# Without a passing record (provider tests) every capability is
# unverified until the test's own probes prove it (section 6.2).
verified = record["verified_capabilities"] if record is not None else {}
verified_capabilities = Capabilities(
**{name: bool(verified.get(name, False)) for name in _CAPABILITY_NAMES}
)
effective_capabilities = compute_effective_capabilities(
spec.protocol_capabilities,
model.runtime.declared_capabilities,
verified_capabilities,
)
return ResolvedModelConfig(
model_ref=model_ref,
role=role,
adapter_id=spec.adapter_id,
adapter_spec_revision=spec.spec_revision,
upstream_model_id=model.upstream_model_id,
base_url=provider.base_url,
auth_ref=auth_ref,
client_options=ClientOptions(
timeout_seconds=parameters.timeout_seconds,
max_retries=parameters.max_retries,
),
request_options=RequestOptions(
max_output_tokens=parameters.max_output_tokens,
temperature=parameters.temperature,
top_p=parameters.top_p,
reasoning_effort=parameters.reasoning_effort,
),
budget=budget,
effective_capabilities=effective_capabilities,
)
@staticmethod
def _resolve_budget(model: ModelConfig, max_output_tokens: int) -> InputBudget:
"""Apply the section 6.5 formulas and the four-mode invariant guard.
The frozen ``message_budget`` is the base budget (system reserve
deducted, no tools/attachments); ``MessageBudgetMiddleware``
recomputes the per-call budget from the frozen reserves. All four
tool/attachment combinations must stay above
``min_effective_input_tokens`` or the run cannot be created.
"""
runtime = model.runtime
if runtime.limit_mode == "combined":
assert runtime.context_window_tokens is not None # schema invariant
resolved_input_limit = runtime.context_window_tokens - max_output_tokens
else:
assert runtime.max_input_tokens is not None # schema invariant
resolved_input_limit = runtime.max_input_tokens
reserves = FixedReserves(
fixed_system_reserve_tokens=runtime.fixed_system_reserve_tokens,
fixed_tools_reserve_tokens=runtime.fixed_tools_reserve_tokens,
fixed_attachments_reserve_tokens=runtime.fixed_attachments_reserve_tokens,
)
base = resolved_input_limit - reserves.fixed_system_reserve_tokens
worst_case = (
base
- reserves.fixed_tools_reserve_tokens
- reserves.fixed_attachments_reserve_tokens
)
if resolved_input_limit <= 0 or worst_case < runtime.min_effective_input_tokens:
raise ModelRegistryError(
CONTEXT_BUDGET_UNSATISFIABLE,
"The input budget cannot satisfy min_effective_input_tokens "
"for every tool/attachment combination.",
details=[
{
"path": "runtime.min_effective_input_tokens",
"code": CONTEXT_BUDGET_UNSATISFIABLE,
}
],
)
return InputBudget(
resolved_input_limit=resolved_input_limit,
fixed_reserves=reserves,
message_budget=base,
)
def compute_availability(
self,
registry: RegistryV4,
verifications: Iterable[dict[str, Any]],
*,
credential_revisions: Mapping[str, int] | None = None,
) -> list[ModelAvailability]:
"""The section 4.3 six-state judgement for every configured model.
``verifications`` carries full records (five-tuple, result, verified
capabilities, timestamp, error code) as returned by
``ModelRuntimeStore.list_model_verifications``;
``credential_revisions`` maps credential IDs to their current pointer
revision for the five-tuple comparison.
"""
revisions = dict(credential_revisions or {})
records_by_model: dict[tuple[str, str], list[dict[str, Any]]] = {}
for record in verifications:
key = (str(record["provider_id"]), str(record["model_key"]))
records_by_model.setdefault(key, []).append(record)
return [
self._judge_model(
provider,
model,
records_by_model.get((provider.id, model.key), []),
revisions,
)
for provider in registry.providers
for model in provider.models
]
def _judge_model(
self,
provider: ProviderConfig,
model: ModelConfig,
records: list[dict[str, Any]],
credential_revisions: Mapping[str, int],
) -> ModelAvailability:
model_ref = ModelRef(provider_id=provider.id, model_key=model.key)
try:
spec = find_adapter_spec(
provider.adapter, model.upstream_model_id, specs=self._specs
)
except ModelRegistryError:
spec = None
current = self._current_record(
provider, model, spec, records, credential_revisions
)
verification = self._verification_info(current, records)
effective = Capabilities()
if spec is not None and current is not None and current["result"] == "passed":
verified_caps = Capabilities(
**{
name: bool(current["verified_capabilities"].get(name, False))
for name in _CAPABILITY_NAMES
}
)
effective = compute_effective_capabilities(
spec.protocol_capabilities,
model.runtime.declared_capabilities,
verified_caps,
)
# Section 4.3 judgement order; the first match wins.
if not provider.enabled or spec is None:
state = "unavailable"
reason = (
REASON_PROVIDER_DISABLED
if not provider.enabled
else REASON_NO_ADAPTER_CONTRACT
)
elif current is not None and current["result"] == "passed" and model.enabled:
state, reason = "enabled", None
elif current is not None and current["result"] == "passed":
state, reason = "verified", REASON_MODEL_DISABLED
elif current is not None:
state = "verification_failed"
reason = current["error_code"] or REASON_VERIFICATION_FAILED
elif records:
# Records exist, but none matches the current five-tuple; stale
# takes priority over configured.
state, reason = "verification_stale", REASON_VERIFICATION_STALE
else:
state, reason = "configured", None
return ModelAvailability(
model_ref=model_ref,
state=state,
selectable=state == "enabled",
reason_code=reason,
verification=verification,
effective_capabilities=effective,
)
@staticmethod
def _current_record(
provider: ProviderConfig,
model: ModelConfig,
spec: AdapterParameterSpec | None,
records: list[dict[str, Any]],
credential_revisions: Mapping[str, int],
) -> dict[str, Any] | None:
if spec is None:
return None
auth_spec = spec.auth_specs.get(provider.auth.mode)
credential_revision = 0
if (
auth_spec is not None
and auth_spec.credential_required
and provider.auth.credential_id is not None
):
credential_revision = credential_revisions.get(
provider.auth.credential_id, 0
)
config_hash = configuration_hash(provider, model)
for record in records:
if (
record["configuration_hash"] == config_hash
and record["credential_revision"] == credential_revision
and record["adapter_spec_revision"] == spec.spec_revision
):
return record
return None
@staticmethod
def _verification_info(
current: dict[str, Any] | None,
records: list[dict[str, Any]],
) -> VerificationInfo:
if current is not None:
return VerificationInfo(
status=current["result"],
verified_at=_rfc3339(int(current["verified_at"])),
adapter_spec_revision=int(current["adapter_spec_revision"]),
)
if records:
latest = max(records, key=lambda record: int(record["verified_at"]))
return VerificationInfo(
status="stale",
verified_at=_rfc3339(int(latest["verified_at"])),
adapter_spec_revision=int(latest["adapter_spec_revision"]),
)
return VerificationInfo(status="none")
+340
View File
@@ -0,0 +1,340 @@
"""Runtime-link glue for snapshot-driven model resolution (design doc 8.3).
``SnapshotRuntime`` bundles the store, resolver, and snapshot service into
the single accessor the running agent link (middleware, agent factories)
uses. It provides two resolution paths:
- ``build_role_model(snapshot, role)`` — per-call construction from a frozen
run snapshot (section 5.2 per-call credential resolution). Every role
resolves to the snapshot's frozen primary (section 6.1).
- ``build_default_role_model()`` — build-time construction from the active
registry's default primary, used only as the compile-time graph binding;
every run re-resolves its model from the run snapshot.
Every model is constructed through ``build_chat_model`` with both safe HTTP
clients (section 6.4). Ollama adapters additionally carry ``max_retries``
on the safe transports because langchain-ollama exposes no client-level
retry option. Secrets are resolved per call and never logged.
"""
from __future__ import annotations
import threading
import uuid
from pathlib import Path
import yaml
from langchain_core.language_models.chat_models import BaseChatModel
from EvoScientist.config.settings import get_config_path
from .endpoint_policy import EndpointPolicy
from .errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError
from .factory import build_chat_model
from .platform import DEFAULT_LOCAL_DEPLOYMENT_ID, PlatformConfigError
from .resolver import ModelRegistryResolver
from .safe_transport import build_safe_async_http_client, build_safe_http_client
from .schemas import DevelopmentEndpoint, ModelRef, ModelRole, ResolvedModelConfig
from .snapshots import RuntimeSnapshot, SnapshotService, config_for_role
from .store import ModelRuntimeStore
class SnapshotRuntime:
"""Store/resolver/snapshot bundle shared by the whole runtime link."""
def __init__(
self,
store: ModelRuntimeStore,
*,
endpoint_policy: EndpointPolicy | None = None,
local_deployment_id: str = DEFAULT_LOCAL_DEPLOYMENT_ID,
allowed_deployment_ids: tuple[str, ...] | None = None,
) -> None:
self._store = store
self._resolver = ModelRegistryResolver(store)
self._snapshots = SnapshotService(store, self._resolver)
self._endpoint_policy = (
endpoint_policy if endpoint_policy is not None else EndpointPolicy(())
)
self._local_deployment_id = local_deployment_id
# Deployment ids allowed to issue snapshots a run may consume: the
# local entry points plus every registered WebUI delegation
# deployment (design doc 7.3/8.2). Defaults to just the local id.
self._allowed_deployment_ids = frozenset(
allowed_deployment_ids
if allowed_deployment_ids is not None
else (local_deployment_id,)
)
@property
def store(self) -> ModelRuntimeStore:
return self._store
@property
def resolver(self) -> ModelRegistryResolver:
return self._resolver
@property
def snapshots(self) -> SnapshotService:
return self._snapshots
@property
def local_deployment_id(self) -> str:
return self._local_deployment_id
@property
def allowed_deployment_ids(self) -> frozenset[str]:
return self._allowed_deployment_ids
# --- snapshot-driven (per-call) resolution ------------------------------
def get_snapshot(
self, snapshot_id: str, *, deployment_id: str, thread_id: str
) -> RuntimeSnapshot:
"""Read a snapshot after verifying its deployment/thread binding."""
return self._snapshots.get(
snapshot_id, deployment_id=deployment_id, thread_id=thread_id
)
def get_snapshot_for_run(
self, snapshot_id: str, *, thread_id: str
) -> RuntimeSnapshot:
"""Read a run's snapshot, verifying thread binding + issuer set."""
return self._snapshots.get_for_run(
snapshot_id,
thread_id=thread_id,
allowed_deployment_ids=self._allowed_deployment_ids,
)
def build_role_model(
self, snapshot: RuntimeSnapshot, role: ModelRole
) -> BaseChatModel:
"""Build the chat model for one role of a frozen run snapshot."""
config = config_for_role(snapshot, role)
credential = self._snapshots.resolve_snapshot_credential(snapshot, role)
return self._build(config, credential)
# --- registry-default (build-time) resolution ----------------------------
def registry_default(self) -> tuple[ModelRef, int]:
"""Return ``(primary, revision)`` for an active registry."""
registry = self._store.load_registry()
if registry.state != "active" or registry.defaults.primary is None:
raise ModelRegistryError(
MODEL_REGISTRY_NOT_READY,
"The model registry is in bootstrap; configure and enable a "
"primary model first.",
)
return registry.defaults.primary, registry.revision
def resolve_default_config(self) -> ResolvedModelConfig:
"""Resolve the registry default primary (section 6.1)."""
primary, _ = self.registry_default()
return self._resolver.resolve(primary, "primary")
def build_default_role_model(self, role: ModelRole = "primary") -> BaseChatModel:
"""Build the chat model for a role from the registry default."""
config = self.resolve_default_config()
return self._build(config, self._resolve_credential(config))
# --- local entry points (design doc 8.1) ---------------------------------
def create_local_snapshot(
self, thread_id: str, *, run_request_id: str | None = None
) -> RuntimeSnapshot:
"""Create (or reuse) a run snapshot for a local entry point.
Local entries — CLI, channels, scheduled tasks, and sub-agents — do
not go through the BFF; they share the same ``SnapshotService`` and
the same snapshot table as the HTTP API (section 8.1). The binding
convention is fixed: ``deployment_id`` is the platform
``local_deployment_id``, ``thread_id`` is the local session/task
identifier, ``model_selection_revision`` is ``0``, and ``primary``
is ``None`` (inherit — the registry defaults are resolved and frozen
at creation time).
``run_request_id`` scopes idempotency: reusing the same
``{deployment_id, thread_id, run_request_id}`` triple returns the
existing snapshot instead of creating a new one. Pass a fresh ID
(the default, a generated UUID) to freeze a new snapshot per run, or
a deterministic one to share a snapshot across collaborators in the
same run.
Raises:
ModelRegistryError: ``MODEL_REGISTRY_NOT_READY`` when the
registry is still in bootstrap (no enabled primary model).
"""
from .snapshots import SnapshotCreateRequest
creation = self._snapshots.create(
SnapshotCreateRequest(
run_request_id=run_request_id or uuid.uuid4().hex,
thread_id=thread_id,
deployment_id=self._local_deployment_id,
model_selection_revision=0,
primary=None,
)
)
return creation.snapshot
# --- internals -------------------------------------------------------------
def _resolve_credential(self, config: ResolvedModelConfig) -> str:
auth_ref = config.auth_ref
if auth_ref.mode == "none":
return ""
assert auth_ref.credential_id is not None # AuthSpec validation
assert auth_ref.credential_revision is not None
return self._store.resolve_credential(
auth_ref.credential_id, auth_ref.credential_revision
)
def _build(self, config: ResolvedModelConfig, credential: str) -> BaseChatModel:
# langchain-ollama exposes no client-level retry option, so the safe
# transports carry the frozen retry budget for ollama adapters; other
# adapters receive retries through their own client options.
retries = (
config.client_options.max_retries if config.adapter_id == "ollama" else 0
)
timeout = config.client_options.timeout_seconds
model = build_chat_model(
config,
http_client=build_safe_http_client(
self._endpoint_policy, timeout=timeout, retries=retries
),
http_async_client=build_safe_async_http_client(
self._endpoint_policy, timeout=timeout, retries=retries
),
credential=credential,
)
# Usage tracking is inert unless the launcher supplied the complete
# sink environment; the adapter contract revision stands in for the
# removed per-profile runtime revision.
from ..usage import UsageModelIdentity, attach_usage_callback
return attach_usage_callback(
model,
UsageModelIdentity(
provider_profile_id=config.model_ref.provider_id,
provider_revision=str(config.adapter_spec_revision),
provider_adapter=config.adapter_id,
model_alias=config.model_ref.model_key,
upstream_model_id=config.upstream_model_id,
),
)
# --- process-shared default runtime -------------------------------------------
_default_runtime: SnapshotRuntime | None = None
_default_runtime_lock = threading.Lock()
def _read_local_platform_fields() -> tuple[
str, Path | None, tuple[DevelopmentEndpoint, ...], tuple[str, ...]
]:
"""Read the runtime-relevant platform fields without the BFF auth gates.
Unlike ``load_platform_security_config`` (which rightly refuses to serve
the Config API without a BFF token and delegation keys), the runtime
link only needs ``local_deployment_id``, ``model_runtime_db``,
``development_endpoints``, and the deployment ids of
``webui_delegation_public_keys`` — all safe to read with the documented
defaults when ``config.yaml`` is absent.
"""
path = get_config_path()
if not path.exists():
return DEFAULT_LOCAL_DEPLOYMENT_ID, None, (), ()
try:
with open(path, encoding="utf-8") as handle:
data = yaml.safe_load(handle) or {}
except yaml.YAMLError as exc:
raise PlatformConfigError(f"config.yaml is not valid YAML: {exc}.") from exc
if not isinstance(data, dict):
raise PlatformConfigError("config.yaml must contain a mapping at the top.")
deployment_id = data.get("local_deployment_id")
if deployment_id is not None and (
not isinstance(deployment_id, str) or not deployment_id.strip()
):
raise PlatformConfigError(
"config.yaml field 'local_deployment_id' must be a non-empty string."
)
database = data.get("model_runtime_db")
if database is not None and not isinstance(database, str):
raise PlatformConfigError(
"config.yaml field 'model_runtime_db' must be a string."
)
raw_endpoints = data.get("development_endpoints")
endpoints: tuple[DevelopmentEndpoint, ...] = ()
if raw_endpoints is not None:
if not isinstance(raw_endpoints, list):
raise PlatformConfigError(
"config.yaml field 'development_endpoints' must be a list."
)
try:
endpoints = tuple(
DevelopmentEndpoint.model_validate(entry) for entry in raw_endpoints
)
except ValueError as exc:
raise PlatformConfigError(
f"Invalid development_endpoints entry: {exc}."
) from exc
raw_keys = data.get("webui_delegation_public_keys")
webui_deployment_ids: tuple[str, ...] = ()
if raw_keys is not None:
if not isinstance(raw_keys, list):
raise PlatformConfigError(
"config.yaml field 'webui_delegation_public_keys' must be a list."
)
ids: list[str] = []
for index, entry in enumerate(raw_keys):
if not isinstance(entry, dict):
raise PlatformConfigError(
f"webui_delegation_public_keys[{index}] must be an object."
)
deployment = entry.get("deployment_id")
if not isinstance(deployment, str) or not deployment.strip():
raise PlatformConfigError(
f"webui_delegation_public_keys[{index}] requires a non-empty "
"'deployment_id'."
)
ids.append(deployment.strip())
webui_deployment_ids = tuple(ids)
return (
deployment_id.strip() if deployment_id else DEFAULT_LOCAL_DEPLOYMENT_ID,
Path(database).expanduser() if database else None,
endpoints,
webui_deployment_ids,
)
def _build_default_runtime() -> SnapshotRuntime:
deployment_id, database_path, endpoints, webui_ids = _read_local_platform_fields()
store = (
ModelRuntimeStore(database_path=database_path)
if database_path is not None
else ModelRuntimeStore()
)
return SnapshotRuntime(
store,
endpoint_policy=EndpointPolicy(endpoints),
local_deployment_id=deployment_id,
allowed_deployment_ids=(deployment_id, *webui_ids),
)
def get_snapshot_runtime() -> SnapshotRuntime:
"""Return the process-shared runtime, building it on first access."""
global _default_runtime
if _default_runtime is None:
with _default_runtime_lock:
if _default_runtime is None:
_default_runtime = _build_default_runtime()
return _default_runtime
def set_snapshot_runtime_for_tests(runtime: SnapshotRuntime | None) -> None:
"""Install (or clear) the shared runtime so tests stay hermetic."""
global _default_runtime
_default_runtime = runtime
@@ -0,0 +1,292 @@
"""SafeHttpTransport: the single network egress for all adapters (section 4.3).
Provider tests and real model calls must share this transport instead of an
SDK's default HTTP client. The policy is enforced in two layers, in order:
1. URL layer — every request origin passes
``EndpointPolicy.validate_request_origin`` before any I/O, so unregistered
local endpoints never open a socket.
2. IP layer — a custom ``httpcore.NetworkBackend`` performs controlled DNS
resolution inside ``connect_tcp`` on every connect (retries included),
filters the deny-listed ranges, and connects directly to the selected IP.
TLS then runs on the origin hostname, so SNI and certificate hostname
checks keep the original name, and HTTP keeps the original Host header.
Redirects are disabled on the built clients; even if a caller enables them,
every redirected origin re-passes both layers.
"""
from __future__ import annotations
import functools
import socket
import ssl
import typing
import anyio
import httpcore
import httpx
from httpcore._backends.anyio import AnyIOStream
from httpcore._backends.base import (
SOCKET_OPTION,
AsyncNetworkBackend,
AsyncNetworkStream,
NetworkBackend,
NetworkStream,
)
from httpcore._backends.sync import SyncStream
from httpcore._exceptions import (
ConnectError,
ConnectTimeout,
ExceptionMapping,
map_exceptions,
)
from httpx._config import DEFAULT_LIMITS
from .endpoint_policy import EndpointPolicy, denied_network_reason
GetAddrInfo = typing.Callable[..., list]
def _select_address(
policy: EndpointPolicy,
host: str,
port: int,
infos: list,
) -> tuple[int, tuple]:
"""Pick the first policy-compliant ``getaddrinfo`` result."""
allow_denied = policy.allows_denied_network(host, port)
for family, _type, _proto, _canonname, sockaddr in infos:
if allow_denied or denied_network_reason(sockaddr[0]) is None:
return family, sockaddr
raise ConnectError(
"DNS resolution for the endpoint yielded only denied network addresses."
)
class SafeNetworkBackend(NetworkBackend):
"""Sync ``httpcore.NetworkBackend`` with controlled DNS and IP filtering."""
def __init__(
self,
policy: EndpointPolicy,
*,
getaddrinfo: GetAddrInfo = socket.getaddrinfo,
) -> None:
self._policy = policy
self._getaddrinfo = getaddrinfo
# Test hook: the (ip, port) selected by the most recent connect.
self.last_selected: tuple[str, int] | None = None
def resolve(self, host: str, port: int) -> tuple[int, tuple]:
"""Resolve ``host`` and return the selected ``(family, sockaddr)``."""
infos = self._getaddrinfo(host, port, type=socket.SOCK_STREAM)
return _select_address(self._policy, host, port, infos)
def connect_tcp(
self,
host: str,
port: int,
timeout: float | None = None,
local_address: str | None = None,
socket_options: typing.Iterable[SOCKET_OPTION] | None = None,
) -> NetworkStream:
family, sockaddr = self.resolve(host, port)
self.last_selected = (sockaddr[0], sockaddr[1])
exc_map: ExceptionMapping = {
socket.timeout: ConnectTimeout,
OSError: ConnectError,
}
with map_exceptions(exc_map):
sock = socket.socket(family, socket.SOCK_STREAM)
try:
if local_address is not None:
sock.bind((local_address, 0))
sock.settimeout(timeout)
# Connect to the validated IP directly — no second DNS lookup.
sock.connect(sockaddr)
for option in socket_options or ():
sock.setsockopt(*option)
sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
except BaseException:
sock.close()
raise
return SyncStream(sock)
def connect_unix_socket(
self,
path: str,
timeout: float | None = None,
socket_options: typing.Iterable[SOCKET_OPTION] | None = None,
) -> NetworkStream:
raise ConnectError("Unix domain sockets are not an allowed endpoint.")
class AsyncSafeNetworkBackend(AsyncNetworkBackend):
"""Async ``httpcore.AsyncNetworkBackend`` with the same controls."""
def __init__(
self,
policy: EndpointPolicy,
*,
getaddrinfo: GetAddrInfo = socket.getaddrinfo,
) -> None:
self._policy = policy
self._getaddrinfo = getaddrinfo
self.last_selected: tuple[str, int] | None = None
async def resolve(self, host: str, port: int) -> tuple[int, tuple]:
lookup = functools.partial(
self._getaddrinfo, host, port, type=socket.SOCK_STREAM
)
infos = await anyio.to_thread.run_sync(lookup)
return _select_address(self._policy, host, port, infos)
async def connect_tcp(
self,
host: str,
port: int,
timeout: float | None = None,
local_address: str | None = None,
socket_options: typing.Iterable[SOCKET_OPTION] | None = None,
) -> AsyncNetworkStream:
_family, sockaddr = await self.resolve(host, port)
self.last_selected = (sockaddr[0], sockaddr[1])
exc_map: ExceptionMapping = {
TimeoutError: ConnectTimeout,
OSError: ConnectError,
anyio.BrokenResourceError: ConnectError,
}
with map_exceptions(exc_map):
with anyio.fail_after(timeout):
# anyio skips DNS for IP literals, so the validated IP is the
# actual connection peer.
stream: anyio.abc.ByteStream = await anyio.connect_tcp(
remote_host=sockaddr[0],
remote_port=sockaddr[1],
local_host=local_address,
)
for option in socket_options or ():
stream._raw_socket.setsockopt(*option) # type: ignore[attr-defined]
return AnyIOStream(stream)
async def connect_unix_socket(
self,
path: str,
timeout: float | None = None,
socket_options: typing.Iterable[SOCKET_OPTION] | None = None,
) -> AsyncNetworkStream:
raise ConnectError("Unix domain sockets are not an allowed endpoint.")
def _origin_string(url: httpx.URL) -> str:
host = url.host
if ":" in host:
host = f"[{host}]"
netloc = host if url.port is None else f"{host}:{url.port}"
return f"{url.scheme}://{netloc}"
class SafeHttpTransport(httpx.HTTPTransport):
"""Sync httpx transport enforcing EndpointPolicy at the URL and IP layer."""
def __init__(
self,
policy: EndpointPolicy,
*,
retries: int = 0,
limits: httpx.Limits = DEFAULT_LIMITS,
ssl_context: ssl.SSLContext | None = None,
backend: SafeNetworkBackend | None = None,
getaddrinfo: GetAddrInfo = socket.getaddrinfo,
) -> None:
ssl_context = ssl_context or ssl.create_default_context()
super().__init__(
verify=ssl_context, limits=limits, retries=retries, trust_env=False
)
self._policy = policy
# httpx.HTTPTransport has no network_backend hook, so the pool is
# rebuilt with the safe backend. The request/response mapping stays
# entirely with httpx.
self._pool = httpcore.ConnectionPool(
ssl_context=ssl_context,
max_connections=limits.max_connections,
max_keepalive_connections=limits.max_keepalive_connections,
keepalive_expiry=limits.keepalive_expiry,
retries=retries,
network_backend=backend
if backend is not None
else SafeNetworkBackend(policy, getaddrinfo=getaddrinfo),
)
def handle_request(self, request: httpx.Request) -> httpx.Response:
self._policy.validate_request_origin(_origin_string(request.url))
return super().handle_request(request)
class AsyncSafeHttpTransport(httpx.AsyncHTTPTransport):
"""Async httpx transport enforcing EndpointPolicy at the URL and IP layer."""
def __init__(
self,
policy: EndpointPolicy,
*,
retries: int = 0,
limits: httpx.Limits = DEFAULT_LIMITS,
ssl_context: ssl.SSLContext | None = None,
backend: AsyncSafeNetworkBackend | None = None,
getaddrinfo: GetAddrInfo = socket.getaddrinfo,
) -> None:
ssl_context = ssl_context or ssl.create_default_context()
super().__init__(
verify=ssl_context, limits=limits, retries=retries, trust_env=False
)
self._policy = policy
self._pool = httpcore.AsyncConnectionPool(
ssl_context=ssl_context,
max_connections=limits.max_connections,
max_keepalive_connections=limits.max_keepalive_connections,
keepalive_expiry=limits.keepalive_expiry,
retries=retries,
network_backend=backend
if backend is not None
else AsyncSafeNetworkBackend(policy, getaddrinfo=getaddrinfo),
)
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
self._policy.validate_request_origin(_origin_string(request.url))
return await super().handle_async_request(request)
def build_safe_http_client(
policy: EndpointPolicy,
timeout: float | httpx.Timeout = 120.0,
*,
retries: int = 0,
getaddrinfo: GetAddrInfo = socket.getaddrinfo,
) -> httpx.Client:
"""Build the sync client every adapter and provider test must share."""
return httpx.Client(
transport=SafeHttpTransport(policy, retries=retries, getaddrinfo=getaddrinfo),
follow_redirects=False,
trust_env=False,
timeout=timeout,
)
def build_safe_async_http_client(
policy: EndpointPolicy,
timeout: float | httpx.Timeout = 120.0,
*,
retries: int = 0,
getaddrinfo: GetAddrInfo = socket.getaddrinfo,
) -> httpx.AsyncClient:
"""Build the async client every adapter and provider test must share."""
return httpx.AsyncClient(
transport=AsyncSafeHttpTransport(
policy, retries=retries, getaddrinfo=getaddrinfo
),
follow_redirects=False,
trust_env=False,
timeout=timeout,
)
+399
View File
@@ -0,0 +1,399 @@
"""The single RegistryV4 schema (design doc sections 4.3, 6.2, 6.4, 9.1).
HTTP JSON, SQLite JSON, import tooling, and test fixtures must all reuse
these Pydantic models instead of maintaining separate shapes.
"""
from __future__ import annotations
import math
from typing import Annotated, Literal
from pydantic import (
BaseModel,
ConfigDict,
Field,
NonNegativeInt,
PositiveInt,
StringConstraints,
field_validator,
model_validator,
)
# ProviderId, ModelKey, CredentialId, and AdapterId share the lowercase ID
# rule. The anchors enforce a full match (the pattern constraint alone only
# searches for a matching substring).
_ID_PATTERN = r"\A[a-z0-9][a-z0-9._-]{0,63}\z"
ProviderId = Annotated[str, StringConstraints(pattern=_ID_PATTERN)]
ModelKey = Annotated[str, StringConstraints(pattern=_ID_PATTERN)]
CredentialId = Annotated[str, StringConstraints(pattern=_ID_PATTERN)]
AdapterId = Annotated[str, StringConstraints(pattern=_ID_PATTERN)]
NonEmptyString = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)]
# `upstream_model_id` keeps the case and characters the provider requires;
# only its length is bounded.
UpstreamModelId = Annotated[str, StringConstraints(min_length=1, max_length=300)]
# `base_url` is normalized and authorized by the backend EndpointPolicy; the
# schema only enforces a plausible http(s) endpoint string.
ValidatedEndpoint = Annotated[str, StringConstraints(min_length=1, max_length=2048)]
ReasoningEffort = Literal["auto", "low", "medium", "high"]
AuthMode = Literal["none", "api_key", "bearer", "adapter_managed"]
RegistryState = Literal["bootstrap", "active"]
LimitMode = Literal["combined", "input_only"]
LimitsStatus = Literal["confirmed", "needs_confirmation"]
LimitsSource = Literal["provider", "tested_contract", "user"]
ModelRole = Literal["primary"]
class SamplingOverride(BaseModel):
"""Mutually-exclusive per-run sampling override (2026-07-30 design).
Overriding one parameter omits the other — including its registry
default — from the request, because providers reject or misbehave when
temperature and top_p are set together.
"""
kind: Literal["temperature", "top_p"]
value: float
@model_validator(mode="after")
def _check_value_range(self) -> SamplingOverride:
if not math.isfinite(self.value):
raise ValueError("sampling override value must be finite")
if self.kind == "temperature" and not 0 <= self.value <= 2:
raise ValueError("temperature override must be within [0, 2]")
if self.kind == "top_p" and not 0 < self.value <= 1:
raise ValueError("top_p override must be within (0, 1]")
return self
class _FrozenModel(BaseModel):
model_config = ConfigDict(frozen=True)
class ModelRef(_FrozenModel):
"""The only identifier of an enabled model: ``{provider_id, model_key}``."""
provider_id: ProviderId
model_key: ModelKey
class RegistryDefaults(BaseModel):
primary: ModelRef | None = None
class Capabilities(BaseModel):
tools: bool = False
vision: bool = False
structured_output: bool = False
class AuthConfig(BaseModel):
"""The adapter-allowed authentication mode plus an optional credential ref."""
mode: AuthMode
credential_id: CredentialId | None = None
@model_validator(mode="after")
def _mode_none_has_no_credential(self) -> AuthConfig:
if self.mode == "none" and self.credential_id is not None:
raise ValueError(
"auth.credential_id must be null when auth.mode is 'none'."
)
return self
class ProviderRuntimeConfig(BaseModel):
timeout_seconds: int = Field(default=120, ge=10, le=600)
max_retries: int = Field(default=2, ge=0, le=5)
default_temperature: float | None = Field(default=None, ge=0, le=2)
default_top_p: float | None = Field(default=None, gt=0, le=1)
default_reasoning_effort: ReasoningEffort = "auto"
class ModelRuntimeConfig(BaseModel):
limit_mode: LimitMode
context_window_tokens: PositiveInt | None = None
max_input_tokens: PositiveInt | None = None
max_output_tokens: PositiveInt
min_effective_input_tokens: int = Field(default=4096, ge=1024)
fixed_system_reserve_tokens: NonNegativeInt
fixed_tools_reserve_tokens: NonNegativeInt
fixed_attachments_reserve_tokens: NonNegativeInt
limits_status: LimitsStatus
limits_source: LimitsSource
temperature: float | None = Field(default=None, ge=0, le=2)
top_p: float | None = Field(default=None, gt=0, le=1)
reasoning_effort: ReasoningEffort = "auto"
declared_capabilities: Capabilities = Field(default_factory=Capabilities)
@model_validator(mode="after")
def _limit_mode_requires_its_limit(self) -> ModelRuntimeConfig:
if self.limit_mode == "combined" and self.context_window_tokens is None:
raise ValueError(
"context_window_tokens is required when limit_mode is 'combined'."
)
if self.limit_mode == "input_only" and self.max_input_tokens is None:
raise ValueError(
"max_input_tokens is required when limit_mode is 'input_only'."
)
return self
class ModelConfig(BaseModel):
key: ModelKey
name: NonEmptyString
upstream_model_id: UpstreamModelId
enabled: bool = True
runtime: ModelRuntimeConfig
class ProviderConfig(BaseModel):
id: ProviderId
name: NonEmptyString
adapter: AdapterId
base_url: ValidatedEndpoint
auth: AuthConfig
enabled: bool = True
runtime: ProviderRuntimeConfig = Field(default_factory=ProviderRuntimeConfig)
models: list[ModelConfig] = Field(default_factory=list)
@field_validator("base_url")
@classmethod
def _base_url_uses_http(cls, value: str) -> str:
if not value.startswith(("http://", "https://")):
raise ValueError("base_url must use http:// or https://.")
return value
@model_validator(mode="after")
def _model_keys_unique(self) -> ProviderConfig:
keys = [model.key for model in self.models]
if len(keys) != len(set(keys)):
raise ValueError(f"Provider {self.id!r} contains duplicate model keys.")
return self
def find_model(self, model_key: str) -> ModelConfig | None:
for model in self.models:
if model.key == model_key:
return model
return None
class RegistryV4(BaseModel):
"""The version 4 model registry document (section 4.3)."""
version: Literal[4] = 4
revision: PositiveInt = 1
state: RegistryState = "bootstrap"
defaults: RegistryDefaults = Field(default_factory=RegistryDefaults)
providers: list[ProviderConfig] = Field(default_factory=list)
@model_validator(mode="after")
def _check_document(self) -> RegistryV4:
provider_ids = [provider.id for provider in self.providers]
if len(provider_ids) != len(set(provider_ids)):
raise ValueError("Provider IDs must be unique.")
ref = self.defaults.primary
if ref is not None:
target = self._locate(ref)
if target is None:
raise ValueError(
f"defaults.primary references an unknown ModelRef "
f"{ref.provider_id!r}/{ref.model_key!r}."
)
provider, model = target
if self.state == "active" and not (provider.enabled and model.enabled):
raise ValueError(
"defaults.primary must reference an enabled model while "
"the registry is active."
)
if self.state == "active" and self.defaults.primary is None:
raise ValueError("defaults.primary must not be null while active.")
return self
def _locate(self, ref: ModelRef) -> tuple[ProviderConfig, ModelConfig] | None:
provider = self.find_provider(ref.provider_id)
if provider is None:
return None
model = provider.find_model(ref.model_key)
if model is None:
return None
return provider, model
def find_provider(self, provider_id: str) -> ProviderConfig | None:
for provider in self.providers:
if provider.id == provider_id:
return provider
return None
# --- Adapter parameter contracts (section 6.2) ---
CredentialKind = Literal["api_key", "bearer_token", "adapter_managed", "none"]
class AuthSpec(BaseModel):
credential_required: bool
credential_kind: CredentialKind
target: Literal["client_option", "request_header", "adapter_internal"]
target_name: str | None = None
class ParameterRule(BaseModel):
supported: bool
value_type: Literal["integer", "number", "enum"]
nullable: Literal["inherit", "omit", "forbidden"]
minimum: float | None = None
maximum: float | None = None
enum_values: list[str] | None = None
target: Literal["client_option", "request_option", "extra_body_path"]
target_name: str
conflicts_with: list[str] = Field(default_factory=list)
normalizer: NonEmptyString
class ConnectionSpec(BaseModel):
chat_model: NonEmptyString
model_field: NonEmptyString
base_url_field: NonEmptyString
class AdapterParameterSpec(BaseModel):
adapter_id: AdapterId
spec_revision: PositiveInt
model_selector: NonEmptyString
auth_specs: dict[AuthMode, AuthSpec] = Field(default_factory=dict)
parameters: dict[str, ParameterRule] = Field(default_factory=dict)
protocol_capabilities: Capabilities
connection: ConnectionSpec | None = None
# --- Availability and resolved configuration (sections 4.3, 6.4, 9.1) ---
class VerificationInfo(BaseModel):
status: Literal["none", "passed", "failed", "stale"]
verified_at: str | None = None
adapter_spec_revision: PositiveInt | None = None
class ModelAvailability(BaseModel):
model_ref: ModelRef
state: Literal[
"unavailable",
"configured",
"verification_failed",
"verification_stale",
"verified",
"enabled",
]
selectable: bool
reason_code: str | None = None
verification: VerificationInfo = Field(
default_factory=lambda: VerificationInfo(status="none")
)
effective_capabilities: Capabilities = Field(default_factory=Capabilities)
class AuthRef(BaseModel):
"""Snapshot credential reference; never carries ``secret_value``."""
mode: AuthMode
credential_id: CredentialId | None = None
credential_revision: PositiveInt | None = None
@model_validator(mode="after")
def _mode_none_has_no_credential(self) -> AuthRef:
if self.mode == "none" and (
self.credential_id is not None or self.credential_revision is not None
):
raise ValueError(
"auth_ref credential fields must be null when mode is 'none'."
)
return self
class ClientOptions(BaseModel):
timeout_seconds: int = Field(ge=10, le=600)
max_retries: int = Field(ge=0, le=5)
class RequestOptions(BaseModel):
max_output_tokens: PositiveInt
temperature: float | None = Field(default=None, ge=0, le=2)
top_p: float | None = Field(default=None, gt=0, le=1)
reasoning_effort: ReasoningEffort = "auto"
class FixedReserves(BaseModel):
fixed_system_reserve_tokens: NonNegativeInt
fixed_tools_reserve_tokens: NonNegativeInt
fixed_attachments_reserve_tokens: NonNegativeInt
class InputBudget(BaseModel):
resolved_input_limit: PositiveInt
fixed_reserves: FixedReserves
message_budget: int
class ResolvedModelConfig(BaseModel):
"""The frozen per-run configuration (section 6.4); contains no secrets."""
model_ref: ModelRef
role: ModelRole
adapter_id: AdapterId
adapter_spec_revision: PositiveInt
upstream_model_id: UpstreamModelId
base_url: ValidatedEndpoint
auth_ref: AuthRef
client_options: ClientOptions
request_options: RequestOptions
budget: InputBudget
effective_capabilities: Capabilities
# --- Endpoint policy (sections 4.3, 9.1) ---
class DevelopmentEndpoint(BaseModel):
"""One local endpoint pre-registered by the deployment administrator.
Only base URLs that exactly match a registered entry (after EndpointPolicy
normalization) may point at loopback or private network addresses.
"""
id: NonEmptyString
url: NonEmptyString
label: NonEmptyString
class EndpointPolicyPublic(BaseModel):
"""The browser-facing view of the endpoint policy (section 9.1).
Exposes only the selectable local endpoints; network rules and secrets
are never part of this shape.
"""
public_https_allowed: Literal[True] = True
development_endpoints: list[DevelopmentEndpoint] = Field(default_factory=list)
# --- API helper models (sections 9.1, 9.2) ---
class CredentialStatus(BaseModel):
credential_id: CredentialId
configured: bool
hint: str | None = None
updated_at: str | None = None
class CredentialWrite(BaseModel):
credential_id: CredentialId
operation: Literal["replace"] = "replace"
secret_value: NonEmptyString
+454
View File
@@ -0,0 +1,454 @@
"""Run runtime snapshot service (design doc 5.2, 8.1, 8.2).
``SnapshotService`` is the single implementation shared by the HTTP snapshot
API (Task 5) and the local CLI/channel/scheduler entry points (Task 7): it
freezes the primary role's complete ``ResolvedModelConfig`` — adapter spec
revision, fixed budget reserves, capabilities, and credential revisions —
into ``run_runtime_snapshots.payload_json``. Payloads and logs never carry
``secret_value``.
- Idempotency: one ``{deployment_id, thread_id, run_request_id}`` triplet
maps to at most one non-terminal snapshot. A repeated request with the
same ``selection_hash`` returns the original snapshot; a different hash
raises ``RUN_REQUEST_CONFLICT``. ``selection_hash`` covers the
pre-resolution ``{primary, reasoning_effort, sampling_override}`` selection (inherit
participates as ``null``); ``model_selection_revision`` is an audit field
and never part of the hash.
- Lifecycle: ``prepared`` (TTL 15 minutes) → ``bound`` (retained 24 hours
after binding) → ``expired`` (terminal, via ``cleanup_expired``);
``prepared`` → ``aborted`` on run creation failure. Terminal snapshots
free their triplet for recreation and reject reads/binds with
``SNAPSHOT_EXPIRED``.
- Reads revalidate the frozen ``adapter_spec_revision`` against the current
contracts and verify the deployment/thread binding.
- Credentials are resolved per call from the frozen ``credential_revision``;
there is no in-process secret cache (section 5.2 rule 4).
"""
from __future__ import annotations
import hashlib
import json
import sqlite3
import time
import uuid
from typing import Any, Literal, get_args
from pydantic import BaseModel, NonNegativeInt, PositiveInt
from .adapters import find_adapter_spec
from .errors import (
ADAPTER_NOT_SUPPORTED,
MODEL_REGISTRY_NOT_READY,
RUN_REQUEST_CONFLICT,
SNAPSHOT_ALREADY_BOUND,
SNAPSHOT_EXPIRED,
SNAPSHOT_NOT_FOUND,
ModelRegistryError,
)
from .resolver import ModelRegistryResolver
from .schemas import (
ModelRef,
ModelRole,
NonEmptyString,
ReasoningEffort,
ResolvedModelConfig,
SamplingOverride,
)
from .store import ModelRuntimeStore
PREPARED_TTL_SECONDS = 15 * 60
BOUND_RETENTION_SECONDS = 24 * 60 * 60
_MODEL_ROLES = get_args(ModelRole)
_TERMINAL_STATUSES = ("expired", "aborted")
class SnapshotCreateRequest(BaseModel):
"""The section 8.2 snapshot creation request body."""
run_request_id: NonEmptyString
thread_id: NonEmptyString
deployment_id: NonEmptyString
model_selection_revision: NonNegativeInt = 0
# ``None`` means inherit: the registry default is resolved at creation.
primary: ModelRef | None = None
# Per-run reasoning-effort override (thread-level selection);
# ``None``/``auto`` keeps the registry-configured effort. Unsupported
# adapters reject it.
reasoning_effort: ReasoningEffort | None = None
# Mutually-exclusive per-run sampling override; ``None`` keeps the
# registry-configured values. Overriding one omits the other from the
# request. Out-of-contract values are rejected by the adapter rules at
# resolve time.
sampling_override: SamplingOverride | None = None
class SnapshotPayload(BaseModel):
"""The frozen private payload; never contains secrets."""
registry_revision: PositiveInt
model_selection_revision: NonNegativeInt
primary: ResolvedModelConfig
class RuntimeSnapshot(BaseModel):
"""One ``run_runtime_snapshots`` row with a typed payload."""
snapshot_id: str
deployment_id: str
thread_id: str
run_request_id: str
selection_hash: str
status: Literal["prepared", "bound", "expired", "aborted"]
langgraph_run_id: str | None
payload: SnapshotPayload
created_at: int
expires_at: int
@classmethod
def from_row(cls, row: dict[str, Any]) -> RuntimeSnapshot:
return cls(
snapshot_id=row["snapshot_id"],
deployment_id=row["deployment_id"],
thread_id=row["thread_id"],
run_request_id=row["run_request_id"],
selection_hash=row["selection_hash"],
status=row["status"],
langgraph_run_id=row["langgraph_run_id"],
payload=SnapshotPayload.model_validate(row["payload"]),
created_at=row["created_at"],
expires_at=row["expires_at"],
)
class SnapshotCreation(BaseModel):
"""``create`` result; ``created`` distinguishes 201 from 200 semantics."""
snapshot: RuntimeSnapshot
created: bool
def compute_selection_hash(
primary: ModelRef | None,
reasoning_effort: ReasoningEffort | None = None,
sampling_override: SamplingOverride | None = None,
) -> str:
"""Hash the pre-resolution selection; inherit participates as ``null``."""
encoded = json.dumps(
{
"primary": (
None
if primary is None
else {"provider_id": primary.provider_id, "model_key": primary.model_key}
),
"reasoning_effort": reasoning_effort,
"sampling_override": (
None
if sampling_override is None
else {"kind": sampling_override.kind, "value": sampling_override.value}
),
},
sort_keys=True,
separators=(",", ":"),
)
return hashlib.sha256(encoded.encode("utf-8")).hexdigest()
def config_for_role(snapshot: RuntimeSnapshot, role: ModelRole) -> ResolvedModelConfig:
"""Every role resolves to the snapshot's frozen primary (section 6.1)."""
if role not in _MODEL_ROLES:
raise ValueError(
f"Unknown model role {role!r}; expected one of {list(_MODEL_ROLES)}."
)
return snapshot.payload.primary
def public_snapshot_view(snapshot: RuntimeSnapshot) -> dict[str, Any]:
"""The section 8.2 public diagnostic subset; never contains secrets."""
def public_config(config: ResolvedModelConfig) -> dict[str, Any]:
return {
"provider_id": config.model_ref.provider_id,
"model_key": config.model_ref.model_key,
"adapter_spec_revision": config.adapter_spec_revision,
"runtime": {
"max_output_tokens": config.request_options.max_output_tokens,
"temperature": config.request_options.temperature,
"top_p": config.request_options.top_p,
"reasoning_effort": config.request_options.reasoning_effort,
"timeout_seconds": config.client_options.timeout_seconds,
"max_retries": config.client_options.max_retries,
},
}
return {
"snapshot_id": snapshot.snapshot_id,
"registry_revision": snapshot.payload.registry_revision,
"primary": public_config(snapshot.payload.primary),
}
class SnapshotService:
"""Creates, binds, reads, and cleans up run runtime snapshots."""
def __init__(
self, store: ModelRuntimeStore, resolver: ModelRegistryResolver
) -> None:
self._store = store
self._resolver = resolver
# --- creation ---------------------------------------------------------
def create(self, request: SnapshotCreateRequest) -> SnapshotCreation:
"""Freeze the requested (or inherited) selection into a snapshot."""
registry = self._store.load_registry()
if registry.state != "active":
raise ModelRegistryError(
MODEL_REGISTRY_NOT_READY,
"The model registry is in bootstrap; configure and enable a "
"primary model first.",
)
selection_hash = compute_selection_hash(
request.primary,
request.reasoning_effort,
request.sampling_override,
)
existing = self._store.find_active_run_snapshot(
deployment_id=request.deployment_id,
thread_id=request.thread_id,
run_request_id=request.run_request_id,
)
if existing is not None:
if existing["selection_hash"] == selection_hash:
return SnapshotCreation(
snapshot=RuntimeSnapshot.from_row(existing), created=False
)
raise ModelRegistryError(
RUN_REQUEST_CONFLICT,
"run_request_id already exists with a different model selection.",
details=[{"path": "run_request_id", "code": RUN_REQUEST_CONFLICT}],
)
primary_ref = (
request.primary
if request.primary is not None
else registry.defaults.primary
)
if primary_ref is None: # pragma: no cover - active registries set it
raise ModelRegistryError(
MODEL_REGISTRY_NOT_READY,
"The model registry has no default primary model.",
)
primary_config = self._resolver.resolve(
primary_ref,
"primary",
registry=registry,
reasoning_effort_override=request.reasoning_effort,
sampling_override=request.sampling_override,
)
payload = {
"registry_revision": registry.revision,
"model_selection_revision": request.model_selection_revision,
"primary": primary_config.model_dump(mode="json"),
}
snapshot_id = f"snap-{uuid.uuid4().hex}"
try:
self._store.insert_run_snapshot(
snapshot_id=snapshot_id,
deployment_id=request.deployment_id,
thread_id=request.thread_id,
run_request_id=request.run_request_id,
selection_hash=selection_hash,
payload=payload,
expires_at=int(time.time()) + PREPARED_TTL_SECONDS,
)
except sqlite3.IntegrityError:
# Lost a creation race: the winner now occupies the triplet, so
# the standard idempotency comparison applies.
winner = self._store.find_active_run_snapshot(
deployment_id=request.deployment_id,
thread_id=request.thread_id,
run_request_id=request.run_request_id,
)
if winner is not None and winner["selection_hash"] == selection_hash:
return SnapshotCreation(
snapshot=RuntimeSnapshot.from_row(winner), created=False
)
raise ModelRegistryError(
RUN_REQUEST_CONFLICT,
"run_request_id already exists with a different model selection.",
details=[{"path": "run_request_id", "code": RUN_REQUEST_CONFLICT}],
) from None
row = self._store.get_run_snapshot(snapshot_id)
assert row is not None # pragma: no cover - inserted above
return SnapshotCreation(snapshot=RuntimeSnapshot.from_row(row), created=True)
# --- lifecycle ----------------------------------------------------------
def bind(self, snapshot_id: str, langgraph_run_id: str) -> RuntimeSnapshot:
"""Transition ``prepared`` → ``bound`` once and extend the retention.
Re-binding with the same ``langgraph_run_id`` is idempotent; a
different value conflicts with ``SNAPSHOT_ALREADY_BOUND``.
"""
while True:
row = self._store.get_run_snapshot(snapshot_id)
if row is None:
raise _snapshot_not_found()
if row["status"] == "bound":
if row["langgraph_run_id"] == langgraph_run_id:
return RuntimeSnapshot.from_row(row)
raise ModelRegistryError(
SNAPSHOT_ALREADY_BOUND,
"The snapshot is already bound to a LangGraph run.",
)
if row["status"] in _TERMINAL_STATUSES:
raise _snapshot_expired()
expires_at = int(time.time()) + BOUND_RETENTION_SECONDS
if self._store.bind_run_snapshot(
snapshot_id,
langgraph_run_id=langgraph_run_id,
expires_at=expires_at,
):
bound = self._store.get_run_snapshot(snapshot_id)
assert bound is not None # pragma: no cover - just updated
return RuntimeSnapshot.from_row(bound)
# Lost a state-transition race; re-read and apply the rules.
def abort(self, snapshot_id: str) -> None:
"""Mark a ``prepared`` snapshot ``aborted`` (run creation failed).
The write is a conditional update: a bind that commits between the
read and the write wins the race, and the abort re-reads and raises
instead of overwriting the bound row.
"""
while True:
row = self._store.get_run_snapshot(snapshot_id)
if row is None:
raise _snapshot_not_found()
if row["status"] == "aborted":
return
if row["status"] == "bound":
raise ModelRegistryError(
SNAPSHOT_ALREADY_BOUND,
"A bound snapshot cannot be aborted.",
)
if row["status"] == "expired":
raise _snapshot_expired()
if self._store.abort_run_snapshot(snapshot_id):
return
# Lost a state-transition race; re-read and apply the rules.
def get(
self,
snapshot_id: str,
*,
deployment_id: str,
thread_id: str,
) -> RuntimeSnapshot:
"""Read a snapshot after verifying its deployment/thread binding.
The frozen ``adapter_spec_revision`` is revalidated against the
current contracts: a removed revision fails loudly with
``ADAPTER_NOT_SUPPORTED`` instead of being silently substituted.
"""
row = self._store.get_run_snapshot(snapshot_id)
if (
row is None
or row["deployment_id"] != deployment_id
or row["thread_id"] != thread_id
):
raise _snapshot_not_found()
if row["status"] in _TERMINAL_STATUSES:
raise _snapshot_expired()
snapshot = RuntimeSnapshot.from_row(row)
self._ensure_frozen_specs_available(snapshot)
return snapshot
def get_for_run(
self,
snapshot_id: str,
*,
thread_id: str,
allowed_deployment_ids: frozenset[str],
) -> RuntimeSnapshot:
"""Read a snapshot for an executing run (section 8.2).
A run cannot self-assert its issuing deployment: the LangGraph
``configurable`` carries ``workspace_deployment_id``, which names
the workspace-isolation scope, not the snapshot issuer. The binding
is therefore verified as ``thread_id`` equality plus membership of
the snapshot's ``deployment_id`` in the platform-registered set
(``local_deployment_id`` and every ``webui_delegation_public_keys``
entry), so a snapshot can never be reused by another thread or by
an unregistered deployment.
"""
row = self._store.get_run_snapshot(snapshot_id)
if (
row is None
or row["thread_id"] != thread_id
or row["deployment_id"] not in allowed_deployment_ids
):
raise _snapshot_not_found()
if row["status"] in _TERMINAL_STATUSES:
raise _snapshot_expired()
snapshot = RuntimeSnapshot.from_row(row)
self._ensure_frozen_specs_available(snapshot)
return snapshot
def cleanup_expired(self, now: int | None = None) -> list[str]:
"""Mark due snapshots ``expired``; returns the transitioned IDs."""
return self._store.expire_due_run_snapshots(
int(time.time()) if now is None else now
)
# --- credentials --------------------------------------------------------
def resolve_snapshot_credential(
self, snapshot: RuntimeSnapshot, role: ModelRole
) -> str:
"""Resolve the secret for the role's frozen credential revision.
Reads the credential store on every call — secrets are never cached
in process memory (section 5.2). A destroyed revision raises
``RUN_CREDENTIAL_REVISION_UNAVAILABLE``; ``mode=none`` models have
no credential and yield an empty string.
"""
config = config_for_role(snapshot, role)
auth_ref = config.auth_ref
if auth_ref.mode == "none":
return ""
assert auth_ref.credential_id is not None # AuthSpec validation
assert auth_ref.credential_revision is not None
return self._store.resolve_credential(
auth_ref.credential_id, auth_ref.credential_revision
)
# --- internals ------------------------------------------------------------
def _ensure_frozen_specs_available(self, snapshot: RuntimeSnapshot) -> None:
config = snapshot.payload.primary
spec = find_adapter_spec(
config.adapter_id,
config.upstream_model_id,
spec_revision=config.adapter_spec_revision,
specs=self._resolver.specs,
)
if spec is None:
raise ModelRegistryError(
ADAPTER_NOT_SUPPORTED,
f"The adapter contract {config.adapter_id!r} at "
f"spec_revision {config.adapter_spec_revision} frozen by "
"this snapshot no longer exists.",
)
def _snapshot_not_found() -> ModelRegistryError:
return ModelRegistryError(SNAPSHOT_NOT_FOUND, "The snapshot does not exist.")
def _snapshot_expired() -> ModelRegistryError:
return ModelRegistryError(
SNAPSHOT_EXPIRED, "The snapshot has reached a terminal state."
)
File diff suppressed because it is too large Load Diff
-63
View File
@@ -5,7 +5,6 @@ from __future__ import annotations
import logging
import os
import shutil
from collections.abc import Iterator
from datetime import datetime
from pathlib import Path
@@ -216,65 +215,3 @@ def resolve_virtual_path(virtual_path: str) -> Path:
"""Resolve a virtual workspace path (e.g. /image.png) to a real filesystem path."""
vpath = virtual_path if virtual_path.startswith("/") else "/" + virtual_path
return (_active_workspace / vpath.lstrip("/")).resolve()
def evoscientist_root() -> Path:
"""Return the application root used by Gateway-managed runtime data."""
env_root = os.environ.get("EVOSCIENTIST_HOME")
if env_root:
return Path(env_root).expanduser().resolve()
return DATA_DIR.expanduser().resolve()
_EVOSCIENTIST_DATA_ROOT: Path | None = None
def _data_root() -> Path:
"""Return the root directory for isolated Web user workspaces."""
global _EVOSCIENTIST_DATA_ROOT
if _EVOSCIENTIST_DATA_ROOT is not None:
return _EVOSCIENTIST_DATA_ROOT
env_root = os.environ.get("EVOSCIENTIST_DATA_ROOT")
if env_root:
root = Path(env_root).expanduser().resolve()
else:
root = evoscientist_root() / "data"
_EVOSCIENTIST_DATA_ROOT = root
return root
def user_data_dir(user_id: str) -> Path:
"""Return and create the isolated data directory for a Web user."""
path = _data_root() / user_id
path.mkdir(parents=True, exist_ok=True)
return path
def iter_user_data_dirs() -> Iterator[Path]:
"""Yield existing Web user directories without creating the data root."""
root = _data_root()
if not root.exists():
return
for path in root.iterdir():
if path.is_dir():
yield path
def thread_data_dir(user_id: str, thread_id: str) -> Path:
"""Return and create a user's isolated thread workspace."""
path = user_data_dir(user_id) / thread_id
path.mkdir(parents=True, exist_ok=True)
return path
def global_data_dir(user_id: str) -> Path:
"""Return and create a user's directory shared across all threads."""
path = user_data_dir(user_id) / "__global__"
path.mkdir(parents=True, exist_ok=True)
return path
def uploads_dir() -> Path:
"""Return the Gateway upload staging directory."""
return evoscientist_root() / "uploads"
+14 -3
View File
@@ -236,6 +236,15 @@ WRITING_GUIDELINES = """# Writing Guidelines
- Professional, objective tone. Be precise, technical, and concise.
"""
# =============================================================================
# Workspace file references (how replies must cite files)
# =============================================================================
FILE_REFERENCES = """# Referencing Workspace Files
When you mention a workspace file in a reply, always use its exact workspace-relative path WITH the directory prefix (e.g. `artifacts/plot.png`, `data/results.csv`). Two forms are both wrong: a bare filename (`plot.png`), and an absolute host path (`/Users/.../files/artifacts/plot.png`) — shell commands show you real host paths (e.g. via `pwd`), but you MUST convert them back to workspace-relative before referencing them in a reply. To show an image inline, embed it as markdown with that same relative path: `![description](artifacts/plot.png)`. The chat UI resolves workspace-relative paths against the conversation workspace; any other form renders as a broken link the user cannot open.
"""
# =============================================================================
# Shell execution guidelines (rules for the `execute` tool)
# =============================================================================
@@ -410,9 +419,10 @@ def get_system_prompt(
2. :data:`EXPERIMENT_WORKFLOW`
3. :data:`REPORT_TEMPLATE`
4. :data:`WRITING_GUIDELINES`
5. :data:`SHELL_GUIDELINES` (or :data:`SHELL_GUIDELINES_DANGEROUS`)
6. :data:`DELEGATION_STRATEGY`
7. :data:`ASYNC_NOTIFICATIONS`
5. :data:`FILE_REFERENCES`
6. :data:`SHELL_GUIDELINES` (or :data:`SHELL_GUIDELINES_DANGEROUS`)
7. :data:`DELEGATION_STRATEGY`
8. :data:`ASYNC_NOTIFICATIONS`
Runtime context is injected per-turn by
:class:`EvoScientist.middleware.RuntimeContextMiddleware`, so dates and
@@ -439,6 +449,7 @@ def get_system_prompt(
EXPERIMENT_WORKFLOW,
REPORT_TEMPLATE,
WRITING_GUIDELINES,
FILE_REFERENCES,
shell_guidelines,
DELEGATION_STRATEGY,
ASYNC_NOTIFICATIONS,
-116
View File
@@ -1,116 +0,0 @@
"""Optional runtime services supplied by an application embedding EvoScientist.
The CLI package must not import a concrete web gateway. Applications such as
Ai4Sci-Web can register their database, storage, metering, and media services
at process startup through this module.
"""
from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, replace
from datetime import date
from pathlib import Path
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):
"""Raised when an optional host-provided service is not configured."""
@dataclass(frozen=True)
class RuntimeIntegrations:
app_connection_provider: AsyncProvider | None = None
session_connection_provider: AsyncProvider | None = None
session_dsn_provider: Callable[[], str | None] | None = None
current_date_provider: Callable[[], date] | None = None
user_storage_root_provider: Callable[[str], Path] | None = None
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()
def configure_runtime_integrations(**services: Any) -> RuntimeIntegrations:
"""Register host-provided services and return the resulting configuration."""
global _integrations
_integrations = replace(_integrations, **services)
return _integrations
def reset_runtime_integrations() -> None:
"""Clear all host-provided services, primarily for tests."""
global _integrations
_integrations = RuntimeIntegrations()
def has_session_connection_provider() -> bool:
return _integrations.session_connection_provider is not None
def get_session_dsn() -> str | None:
provider = _integrations.session_dsn_provider
return provider() if provider is not None else None
async def get_session_connection() -> Any:
provider = _integrations.session_connection_provider
if provider is None:
raise RuntimeIntegrationUnavailable(
"No session connection provider is configured"
)
return await provider()
async def get_app_connection() -> Any:
provider = _integrations.app_connection_provider
if provider is None:
raise RuntimeIntegrationUnavailable(
"No application connection provider is configured"
)
return await provider()
def current_date() -> date:
provider = _integrations.current_date_provider
return provider() if provider is not None else date.today()
def resolve_user_storage_root(user_id: str) -> Path | None:
provider = _integrations.user_storage_root_provider
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:
await handler(path)
async def record_service_usage(service: str, action: str) -> None:
recorder = _integrations.usage_recorder
if recorder is not None:
await recorder(service, action)
def get_image_backend() -> Any:
factory = _integrations.image_backend_factory
if factory is None:
raise RuntimeIntegrationUnavailable(
"Image generation is unavailable in this runtime. Configure an image backend first."
)
return factory()
File diff suppressed because it is too large Load Diff
+62
View File
@@ -0,0 +1,62 @@
---
name: image-artist
description: Generate or edit images with the configured dedicated image models. Use when the user asks to draw, illustrate, create a cover/diagram/figure, or modify an existing workspace image.
---
# Image Artist
Workflow for producing AI-generated images. Execution happens through the
`generate_image` and `edit_image` tools — this skill tells you when and how
to call them.
## When to Use
- The user asks to create, draw, or generate an image, cover, illustration,
poster, or diagram from a description → `generate_image`.
- The user provides an existing image (in `artifacts/` or another workspace
path) and asks to modify, restyle, or adapt it → `edit_image`.
- The user asks for a chart/plot of **data** → do NOT use this skill; compute
the plot with code (matplotlib) instead.
## Prompt Construction
Write the `prompt` as a single dense paragraph covering:
1. **Subject** — what is in the image, concretely.
2. **Style** — e.g. flat vector, watercolor, photorealistic, line art.
3. **Composition** — layout, perspective, framing.
4. **Palette/lighting** — only when the user cares.
5. **Negative constraints** — what to avoid (text artifacts, extra limbs),
phrased as "no ..." at the end.
Do not pad with fluff ("masterpiece", "best quality"); describe the image.
## Parameters
- `size`: "1024x1024" (default), "1536x1024" (landscape), "1024x1536"
(portrait). Pick from the content: covers/posters often landscape or portrait.
- `model`: omit to use the configured default; the tool description lists the
available image models. Only pass a name from that list.
- `output_path`: suggest a meaningful name like `artifacts/topic_cover.png`;
must start with `artifacts/` and end with `.png`.
- `n`: number of variations (default 1). For "give me a few options" use 2-4.
## After the Call
- The tool returns JSON: `{"ok": true, "paths": [...]}` or
`{"ok": false, "error": "..."}`.
- On success, reference the saved paths in your reply.
- On failure, read the error: if it lists available models, retry with one of
them or explain the configuration gap to the user. If the provider rejected
the prompt, rewrite it more concretely and retry once.
## Variations
For "make variations", call `generate_image` once with `n>1`, or call
`edit_image` on a previously generated image with a modified prompt.
## Security Rules
- Never read, print, or quote the `image_generation` section of config.yaml
or any API key. The tools handle credentials internally.
- Never pass file paths outside the workspace as `image_path`/`mask_path`.
@@ -2,7 +2,8 @@
"""Improve a skill description based on eval results.
Takes eval results (from run_eval.py) and generates an improved description
using EvoSci's LLM layer (multi-provider support).
using ``langchain.chat_models.init_chat_model`` (multi-provider support via
environment-provided API keys, optionally seeded from the EvoSci config).
"""
import argparse
@@ -16,8 +17,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from langchain_core.messages import HumanMessage
from EvoScientist.config import apply_config_to_env, get_effective_config
from EvoScientist.llm import get_chat_model
from scripts.run_eval import _config_defaults, _init_chat_model
from scripts.utils import parse_skill_md
@@ -141,13 +141,15 @@ I'd encourage you to be creative and mix up the style in different iterations si
Please respond with only the new description text in <new_description> tags, nothing else."""
# Initialise model via EvoSci's LLM layer
config = get_effective_config()
apply_config_to_env(config)
chat_model = get_chat_model(
model=model or config.model,
provider=provider or config.provider,
)
# Initialise model (API keys come from the environment / EvoSci config)
default_model, default_provider = _config_defaults()
effective_model = model or default_model
if not effective_model:
raise RuntimeError(
"No model specified: pass --model or configure one (or set a "
"provider's default model via environment)."
)
chat_model = _init_chat_model(effective_model, provider or default_provider)
response = chat_model.invoke([HumanMessage(content=prompt)])
@@ -253,9 +255,8 @@ def main():
name, _, content = parse_skill_md(skill_path)
current_description = eval_results["description"]
# Load EvoSci config for defaults
config = get_effective_config()
apply_config_to_env(config)
# Load EvoSci config for defaults (best-effort; env vars alone suffice)
default_model, default_provider = _config_defaults()
if args.verbose:
print(f"Current: {current_description}", file=sys.stderr)
@@ -270,8 +271,8 @@ def main():
current_description=current_description,
eval_results=eval_results,
history=history,
model=args.model or config.model,
provider=args.provider or config.provider,
model=args.model or default_model,
provider=args.provider or default_provider,
)
if args.verbose:
@@ -2,8 +2,10 @@
"""Run trigger evaluation for a skill description.
Tests whether a skill's description causes an LLM to trigger (load the skill)
for a set of queries. Uses EvoSci's multi-provider LLM layer with tool calling
to simulate the agent's skill selection behavior.
for a set of queries. Uses ``langchain.chat_models.init_chat_model`` with
tool calling to simulate the agent's skill selection behavior. API keys come
from the environment (optionally seeded from the EvoSci config when
available).
"""
import argparse
@@ -18,13 +20,28 @@ sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from scripts.utils import parse_skill_md
def _init_config():
"""Initialize EvoSci config and apply env vars (once per process)."""
from EvoScientist.config import apply_config_to_env, get_effective_config
def _config_defaults() -> tuple[str | None, str | None]:
"""Best-effort (model, provider) defaults from the EvoSci config.
config = get_effective_config()
apply_config_to_env(config)
return config
Also applies configured API keys to the process environment so
``init_chat_model`` can pick them up. Returns ``(None, None)`` when the
config layer is unavailable — environment variables alone suffice.
"""
try:
from EvoScientist.config import apply_config_to_env, get_effective_config
config = get_effective_config()
apply_config_to_env(config)
return getattr(config, "model", None), getattr(config, "provider", None)
except Exception:
return None, None
def _init_chat_model(model: str, provider: str | None, **kwargs):
"""Build a chat model; API keys are read from the environment."""
from langchain.chat_models import init_chat_model
return init_chat_model(model=model, model_provider=provider, **kwargs)
def run_single_query(
@@ -43,12 +60,14 @@ def run_single_query(
from langchain_core.messages import HumanMessage, SystemMessage
from langchain_core.tools import tool
from EvoScientist.llm import get_chat_model
config = _init_config()
effective_model = model or config.model
effective_provider = provider or config.provider
default_model, default_provider = _config_defaults()
effective_model = model or default_model
effective_provider = provider or default_provider
if not effective_model:
raise RuntimeError(
"No model specified: pass --model or configure one (or set a "
"provider's default model via environment)."
)
@tool
def load_skill(name: str) -> str:
@@ -72,9 +91,9 @@ If no skill is relevant, respond directly to the user without calling any tools.
pass
try:
chat_model = get_chat_model(
model=effective_model,
provider=effective_provider,
chat_model = _init_chat_model(
effective_model,
effective_provider,
**eval_kwargs,
)
model_with_tools = chat_model.bind_tools([load_skill])
@@ -98,9 +117,9 @@ If no skill is relevant, respond directly to the user without calling any tools.
# Fallback for providers that don't support tool calling:
# Use a text-based approach
try:
chat_model = get_chat_model(
model=effective_model,
provider=effective_provider,
chat_model = _init_chat_model(
effective_model,
effective_provider,
**eval_kwargs,
)
fallback_prompt = f"""You are a helpful AI assistant with specialized skills available.
@@ -221,12 +240,14 @@ def main():
"--trigger-threshold", type=float, default=0.5, help="Trigger rate threshold"
)
parser.add_argument(
"--model", default=None, help="Model to use (default: user's configured model)"
"--model",
default=None,
help="Model to use (default: EvoSci config model; required otherwise)",
)
parser.add_argument(
"--provider",
default=None,
help="LLM provider (default: user's configured provider)",
help="LLM provider (default: EvoSci config provider, else inferred from model)",
)
parser.add_argument(
"--verbose", action="store_true", help="Print progress to stderr"
@@ -18,10 +18,9 @@ from pathlib import Path
# Ensure skill-creator root is on sys.path for `from scripts.xxx` imports
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from EvoScientist.config import apply_config_to_env, get_effective_config
from scripts.generate_report import generate_html
from scripts.improve_description import improve_description
from scripts.run_eval import run_eval
from scripts.run_eval import _config_defaults, run_eval
from scripts.utils import parse_skill_md
@@ -66,8 +65,10 @@ def run_loop(
log_dir: Path | None = None,
) -> dict:
"""Run the eval + improvement loop."""
config = get_effective_config()
apply_config_to_env(config)
# Seed API keys into the environment from the EvoSci config (best-effort);
# model/provider defaults are resolved per call by run_eval /
# improve_description when not passed explicitly.
_config_defaults()
name, original_description, content = parse_skill_md(skill_path)
current_description = description_override or original_description
+6 -8
View File
@@ -8,12 +8,10 @@ import base64
import inspect
import mimetypes
import os
import warnings
from collections.abc import AsyncGenerator, AsyncIterator, Mapping
from dataclasses import dataclass
from typing import Any, TypeAlias
from langchain_core._api import LangChainBetaWarning
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
from langgraph.graph import END
from langgraph.types import Command, Interrupt
@@ -45,12 +43,6 @@ GraphRunInput: TypeAlias = str | Command
LangGraphStreamInput: TypeAlias = dict[str, list[dict[str, object]]] | Command
_ValueMessageKey: TypeAlias = tuple[str, ...]
warnings.filterwarnings(
"ignore",
message=r"The v3 streaming protocol on Pregel is experimental\.",
category=LangChainBetaWarning,
)
@dataclass(frozen=True, slots=True)
class _AssistantValueMessage:
@@ -805,6 +797,7 @@ async def stream_agent_events(
thread_id: str,
metadata: dict[str, Any] | None = None,
media: list[str] | None = None,
runtime_snapshot_id: str | None = None,
) -> AsyncGenerator[dict[str, Any], None]:
"""Stream events from a DeepAgents/LangGraph v3 run.
@@ -820,6 +813,9 @@ async def stream_agent_events(
metadata: Optional metadata dict merged into the LangGraph config
(e.g. agent_name, updated_at for checkpoint persistence).
media: Optional list of local file paths for attachments.
runtime_snapshot_id: Optional frozen run snapshot ID (design doc
8.1/8.2) placed into ``configurable`` — the only model
configuration a run may carry.
Yields:
Event dicts: thinking, text, tool_call, tool_result,
@@ -829,6 +825,8 @@ async def stream_agent_events(
config: dict[str, Any] = {"configurable": {"thread_id": thread_id}}
if metadata:
config["metadata"] = metadata
if runtime_snapshot_id is not None:
config["configurable"]["runtime_snapshot_id"] = runtime_snapshot_id
emitter = StreamEventEmitter()
existing_summarization_event: Mapping[str, object] | None = None
try:
+36 -12
View File
@@ -7,10 +7,12 @@ construction utility, not a deployment concern. Any deployment surface
servers) can call ``build_async_subagent_graph(name)`` to materialize the
runnable graph.
Reuses the main EvoScientist agent's chat model, backend, and middleware so
the deployed sub-agent has full capability parity with its in-process
synchronous counterpart: same workspace files, same ``/skills/`` and
``/memories/`` routes, same error-handling and context-overflow middleware.
Reuses the main EvoScientist agent's backend and middleware so the deployed
sub-agent has full capability parity with its in-process synchronous
counterpart: same workspace files, same ``/skills/`` and ``/memories/``
routes, same error-handling and context-overflow middleware. The chat model
is resolved from the active model registry's default primary and re-resolved
per run from the run snapshot by ``ConfigurableModelMiddleware``.
"""
from __future__ import annotations
@@ -40,13 +42,12 @@ def build_async_subagent_graph(name: str) -> Any:
from EvoScientist.config import apply_config_to_env, get_effective_config
from EvoScientist.EvoScientist import (
SUBAGENTS_CONFIG,
_ensure_auxiliary_chat_model,
_ensure_chat_model,
_ensure_general_purpose_subagent,
_get_default_backend,
_get_default_middleware,
_inject_subagent_middleware,
)
from EvoScientist.model_registry.runtime import get_snapshot_runtime
from EvoScientist.tools import skill_manager, tavily_search, think_tool
from EvoScientist.utils import load_subagents
@@ -99,18 +100,41 @@ def build_async_subagent_graph(name: str) -> Any:
#
# Memory middleware is included so async sub-agents get the same profile
# context and `/memories/profile/...` file guidance as the main agent.
#
# The compile-time model binding is resolved from the active registry's
# default primary — never from config.yaml free strings. Per-run calls
# are re-resolved from the run snapshot by ConfigurableModelMiddleware.
#
# Bootstrap registry: no default exists yet, so no compile-time model
# can be built. The graph must still materialize — one failing factory
# must not take down the whole langgraph dev service — so we bind a
# placeholder that raises MODEL_REGISTRY_NOT_READY on the first model
# call. Every run fails with that clear structured error until a
# primary model is configured and enabled; no 32K/implicit fallback.
from EvoScientist.model_registry.errors import (
MODEL_REGISTRY_NOT_READY,
ModelRegistryError,
)
from EvoScientist.model_registry.placeholder import RegistryNotReadyChatModel
runtime = get_snapshot_runtime()
snapshot_role = "primary"
try:
model = runtime.build_default_role_model(snapshot_role)
except ModelRegistryError as exc:
if exc.code != MODEL_REGISTRY_NOT_READY:
raise
model = RegistryNotReadyChatModel(detail=str(exc))
subagents = []
_ensure_general_purpose_subagent(subagents)
_inject_subagent_middleware(subagents)
_inject_subagent_middleware(subagents, chat_model=model)
middleware = _get_default_middleware(
for_async_subagent=True,
memory_source_agent=name,
)
# Scheduler is an unattended timer task → use the cheaper auxiliary model.
model = (
_ensure_auxiliary_chat_model() if name == "scheduler" else _ensure_chat_model()
chat_model=model,
snapshot_role=snapshot_role,
)
return create_deep_agent(
+4
View File
@@ -4,12 +4,16 @@ External imports like ``from EvoScientist.tools import tavily_search`` continue
to work unchanged thanks to these re-exports.
"""
from .image import edit_image, generate_image, refresh_image_tool_descriptions
from .search import fetch_webpage_content, tavily_search
from .skill_manager import skill_manager
from .think import think_tool
__all__ = [
"edit_image",
"fetch_webpage_content",
"generate_image",
"refresh_image_tool_descriptions",
"skill_manager",
"tavily_search",
"think_tool",
+165
View File
@@ -0,0 +1,165 @@
"""Agent tools for image generation and editing.
Thin wrappers over ``image_gen.service``. They never see API keys: entries
resolve ``${ENV_VAR}`` references server-side at call time.
"""
from __future__ import annotations
import asyncio
import json
from pathlib import Path
from langchain.tools import ToolRuntime
from langchain_core.tools import tool
def _resolve_workspace(runtime: ToolRuntime | None) -> Path:
"""Resolve the directory the agent's filesystem writes to.
Scoped conversations must land in the scope files dir — that is what the
WebUI file browser and download API serve. Unscoped (legacy CLI) runs
fall back to the deployment workspace root.
"""
from EvoScientist.workspace_scope import require_scoped_runtime
context = require_scoped_runtime(runtime)
if context is not None:
return context.files_dir
from EvoScientist.paths import _active_workspace
return Path(_active_workspace).resolve()
async def _workspace(runtime: ToolRuntime | None = None) -> Path:
# Scope validation hits the sqlite Registry; run it off the event loop or
# the dev server's blocking-call detector aborts the tool.
return await asyncio.to_thread(_resolve_workspace, runtime)
def _settings_path() -> Path:
from EvoScientist.config.settings import get_config_path
return get_config_path()
# Late-bound indirections so tests can patch without importing the service's
# adapter stack.
async def _generate_for_workspace(workspace: Path, **kwargs):
from EvoScientist.image_gen import service
return await service.generate_for_workspace(workspace, **kwargs)
async def _edit_for_workspace(workspace: Path, **kwargs):
from EvoScientist.image_gen import service
return await service.edit_for_workspace(workspace, **kwargs)
def _json_result(payload: dict) -> str:
if payload.get("ok") and payload.get("paths"):
# Ready-to-use embeds so the agent references the exact
# workspace-relative paths instead of paraphrasing them.
payload["markdown"] = [
f"![{Path(p).stem}]({p})" for p in payload["paths"]
]
return json.dumps(payload, ensure_ascii=False)
def _available_models_hint() -> str:
try:
from EvoScientist.image_gen.config import load_image_generation_settings
settings = load_image_generation_settings(config_path=_settings_path())
names = [entry.display_name() for entry in settings.models]
if not names:
return ""
default = settings.default_model or settings.models[0].id
return (
f" Available image models: {', '.join(names)}. Default: {default}."
)
except Exception:
return ""
def refresh_image_tool_descriptions() -> None:
"""Refresh tool descriptions with the current image model list.
Called on every agent build so config.yaml edits take effect without a
restart. Descriptions must never include credentials.
"""
hint = _available_models_hint()
reference_rule = (
" In your reply, embed each saved file with the exact "
"workspace-relative path from the result (the ready-made snippets "
"in the result's markdown field, e.g. ![...](artifacts/example.png))"
" — never drop the directory prefix."
)
generate_image.description = (
"Generate one or more images with a dedicated image model and save "
f"them under artifacts/.{hint}{reference_rule}"
)
edit_image.description = (
"Edit an existing workspace image with a dedicated image model and "
f"save the result under artifacts/.{hint}{reference_rule}"
)
@tool
async def generate_image(
prompt: str,
model: str | None = None,
size: str = "1024x1024",
quality: str = "auto",
background: str = "auto",
output_path: str | None = None,
n: int = 1,
runtime: ToolRuntime = None,
) -> str:
"""Generate one or more images with a dedicated image model and save them under artifacts/."""
try:
result = await _generate_for_workspace(
await _workspace(runtime),
prompt=prompt,
model=model,
size=size,
quality=quality,
background=background,
output_path=output_path,
n=n,
)
return _json_result(result)
except Exception as exc:
return _json_result({"ok": False, "error": str(exc) or exc.__class__.__name__})
@tool
async def edit_image(
image_path: str,
prompt: str,
model: str | None = None,
mask_path: str | None = None,
size: str = "1024x1024",
quality: str = "auto",
output_path: str | None = None,
runtime: ToolRuntime = None,
) -> str:
"""Edit an existing workspace image with a dedicated image model and save the result under artifacts/."""
try:
result = await _edit_for_workspace(
await _workspace(runtime),
image_path=image_path,
prompt=prompt,
model=model,
mask_path=mask_path,
size=size,
quality=quality,
output_path=output_path,
)
return _json_result(result)
except Exception as exc:
return _json_result({"ok": False, "error": str(exc) or exc.__class__.__name__})
refresh_image_tool_descriptions()
+420
View File
@@ -3,13 +3,28 @@
Compares the installed version against PyPI and caches the result
(see ``CACHE_TTL``). All errors are silently swallowed so startup
is never blocked or degraded.
Also exposes a Gitea-based checker (``get_update_info``) used by the
``/internal/system/version`` HTTP route. It resolves "latest" as the
max semver across the instance's releases and tags — Gitea orders
``releases/latest`` by tag creation date, not semver, so a single
endpoint cannot be trusted.
``download_update`` stages a release artifact under the update staging
dir (``~/.evoscientist/updates/`` by default) and returns the suggested
install command; it never applies the update itself.
"""
from __future__ import annotations
import hashlib
import json
import logging
import os
import time
from pathlib import Path
from typing import Any, TypedDict
from urllib.parse import urlparse
from .config.settings import get_config_dir
@@ -91,3 +106,408 @@ def is_update_available() -> tuple[bool, str | None]:
logger.debug("Failed to compare versions", exc_info=True)
return False, None
# ---------------------------------------------------------------------------
# Gitea-based checker (powers /internal/system/version)
# ---------------------------------------------------------------------------
GITEA_CACHE_TTL = 1_200 # 20 minutes
_UPDATE_CACHE: dict[str, Any] = {"info": None, "fetched_at": 0.0}
class UpdateInfo(TypedDict):
current_version: str
latest_version: str
has_update: bool
release_url: str | None
release_notes: str | None
published_at: str | None
cached: bool
warning: str | None
breaking_db: bool
def _semver_key(v: str) -> tuple[int, int, int]:
"""Map a version string to a comparable 3-tuple; junk segments count as 0."""
parts = v.strip().lstrip("vV").split(".")[:3]
out = []
for p in parts:
try:
out.append(int(p))
except ValueError:
out.append(0)
while len(out) < 3:
out.append(0)
return tuple(out) # type: ignore[return-value]
def _gitea_base() -> str:
return os.environ.get("EVOSCIENTIST_UPDATE_BASE_URL", "https://git.foksai.com").rstrip("/")
def _gitea_repo() -> str:
return os.environ.get("EVOSCIENTIST_UPDATE_REPO", "ouyangbo/EvoScientist")
def _http_get_json(url: str, *, timeout: float = 10.0) -> Any:
import httpx
headers = {"User-Agent": "EvoScientist update-check"}
token = os.environ.get("EVOSCIENTIST_UPDATE_TOKEN")
if token:
headers["Authorization"] = f"token {token}"
resp = httpx.get(url, headers=headers, timeout=timeout, follow_redirects=True)
if resp.status_code == 404:
return None
resp.raise_for_status()
return resp.json()
def _current_version() -> str:
try:
return _installed_version()
except Exception:
return "0.0.0-dev"
def _fetch_latest() -> tuple[str, dict | None]:
"""Return (latest_version, release_dict_or_None) as max semver over releases+tags."""
base, repo = _gitea_base(), _gitea_repo()
releases = _http_get_json(f"{base}/api/v1/repos/{repo}/releases?limit=20") or []
tags = _http_get_json(f"{base}/api/v1/repos/{repo}/tags?limit=20") or []
candidates: list[tuple[tuple[int, int, int], str, dict | None]] = []
for r in releases:
tag = r.get("tag_name", "")
if tag:
candidates.append((_semver_key(tag), tag, r))
for t in tags:
name = t.get("name", "")
if name:
candidates.append((_semver_key(name), name, None))
if not candidates:
return _current_version(), None
candidates.sort(key=lambda c: c[0])
_key, tag, release = candidates[-1]
if release is None:
# the newest candidate may still have a release lower in the list;
# prefer its metadata only when it matches the winning tag
release = next((r for k, tg, r in candidates if r and tg == tag), None)
return tag.lstrip("vV"), release
def _is_breaking_release(release: dict | None) -> bool:
if not release:
return False
return bool(release.get("prerelease")) and "BREAKING-DB" in (release.get("body") or "")
def get_update_info(*, force: bool = False) -> UpdateInfo:
"""Resolve current vs latest published version, with a 20-minute cache."""
current = _current_version()
def stale(latest: str | None = None, warning: str | None = None, cached: bool = False) -> UpdateInfo:
return UpdateInfo(
current_version=current,
latest_version=latest or current,
has_update=False,
release_url=None,
release_notes=None,
published_at=None,
cached=cached,
warning=warning,
breaking_db=False,
)
if os.environ.get("EVOSCIENTIST_UPDATE_CHECK_DISABLED") == "1":
return stale()
now = time.time()
cached_info: UpdateInfo | None = _UPDATE_CACHE["info"]
if cached_info is not None and not force:
age = now - _UPDATE_CACHE["fetched_at"]
if age < GITEA_CACHE_TTL:
return UpdateInfo(
**{**cached_info, "cached": True, "current_version": current}
)
try:
latest, release = _fetch_latest()
except Exception as exc: # network/HTTP failure: serve stale cache if any
logger.debug("Gitea update check failed", exc_info=True)
if cached_info is not None:
return UpdateInfo(
**{
**cached_info,
"cached": True,
"current_version": current,
"warning": f"update check failed: {exc}",
}
)
return stale(warning=f"update check failed: {exc}")
info = UpdateInfo(
current_version=current,
latest_version=latest,
has_update=_semver_key(latest) > _semver_key(current),
release_url=release.get("html_url") if release else None,
release_notes=release.get("body") if release else None,
published_at=release.get("published_at") if release else None,
cached=False,
warning=None,
breaking_db=_is_breaking_release(release),
)
_UPDATE_CACHE["info"] = info
_UPDATE_CACHE["fetched_at"] = now
return info
# ---------------------------------------------------------------------------
# Rollback version list
# ---------------------------------------------------------------------------
class RollbackVersion(TypedDict):
version: str
published_at: str | None
release_url: str | None
def list_rollback_versions(limit: int = 3) -> list[RollbackVersion]:
"""Non-draft, non-prerelease releases strictly older than the current version.
Stops collecting at a BREAKING-DB marker release: rollback past a breaking
schema change is not offered (and the breaking release itself, being a
prerelease, is never listed).
"""
base, repo = _gitea_base(), _gitea_repo()
releases = _http_get_json(f"{base}/api/v1/repos/{repo}/releases?limit=20") or []
current = _current_version()
out: list[RollbackVersion] = []
seen: set[str] = set()
ordered = sorted(
(r for r in releases if r.get("tag_name")),
key=lambda r: _semver_key(r["tag_name"]),
reverse=True,
)
for r in ordered:
if r.get("draft"):
continue
if r.get("prerelease"):
if "BREAKING-DB" in (r.get("body") or ""):
break # breaking point: nothing older is rollback-safe
continue
v = r["tag_name"].lstrip("vV")
if v in seen or _semver_key(v) >= _semver_key(current):
continue
seen.add(v)
out.append(
RollbackVersion(
version=v,
published_at=r.get("published_at"),
release_url=r.get("html_url"),
)
)
if len(out) >= limit:
break
return out
def is_allowed_rollback(version: str) -> bool:
target = version.strip().lstrip("vV")
return any(v["version"] == target for v in list_rollback_versions())
# ---------------------------------------------------------------------------
# Update download (powers POST /internal/system/version/download)
# ---------------------------------------------------------------------------
MAX_DOWNLOAD_BYTES = 200 * 1024 * 1024
class UpdateDownloadError(RuntimeError):
"""A user-safe download failure (message may reach the browser)."""
class DownloadResult(TypedDict):
version: str
file: str
path: str
suggested_command: str
def _staging_dir() -> Path:
override = os.environ.get("EVOSCIENTIST_UPDATE_STAGING_DIR", "").strip()
if override:
return Path(override)
return get_config_dir() / "updates"
def _updates_dir(version: str) -> Path:
return _staging_dir() / f"v{version}"
def _http_download(
url: str, dest: Path, *, max_bytes: int, timeout: float = 120.0
) -> Path:
import httpx
headers = {"User-Agent": "EvoScientist update-check"}
token = os.environ.get("EVOSCIENTIST_UPDATE_TOKEN")
if token:
headers["Authorization"] = f"token {token}"
with httpx.stream(
"GET", url, headers=headers, timeout=timeout, follow_redirects=True
) as resp:
resp.raise_for_status()
total = 0
with open(dest, "wb") as fh:
for chunk in resp.iter_bytes():
total += len(chunk)
if total > max_bytes:
raise UpdateDownloadError(
f"download exceeds the {max_bytes}-byte cap"
)
fh.write(chunk)
return dest
def _check_asset_url(url: str) -> None:
"""SSRF guard: only download from the configured Gitea host (any port)."""
if urlparse(url).hostname != urlparse(_gitea_base()).hostname:
raise UpdateDownloadError(f"download URL host is not allowed: {url!r}")
def _pick_asset(release: dict) -> tuple[str, str]:
"""Choose (filename, url): wheel > sdist > repo tarball."""
assets = release.get("assets") or []
wheels = [a for a in assets if str(a.get("name", "")).endswith(".whl")]
sdists = [a for a in assets if str(a.get("name", "")).endswith(".tar.gz")]
for asset in (*wheels, *sdists):
return asset["name"], asset["browser_download_url"]
tarball = release.get("tarball_url")
if tarball:
return f"{release.get('tag_name', 'release')}.tar.gz", tarball
raise UpdateDownloadError("release has no downloadable assets")
def _verify_local_checksum(sums_path: Path, filename: str, path: Path) -> None:
"""Verify path against the checksums.txt entry for filename (if listed)."""
expected = None
for line in sums_path.read_text(encoding="utf-8").splitlines():
parts = line.split()
if len(parts) == 2 and parts[1] == filename:
expected = parts[0]
break
if expected is None:
return # artifact not listed — nothing to verify against
actual = hashlib.sha256(path.read_bytes()).hexdigest()
if actual != expected:
raise UpdateDownloadError(f"checksum mismatch for {filename}")
def _verify_checksum(release: dict, updates: Path, filename: str, path: Path) -> None:
"""If the release ships checksums.txt with an entry for filename, verify it."""
sums = next(
(a for a in (release.get("assets") or []) if a.get("name") == "checksums.txt"),
None,
)
if sums is None:
return
_check_asset_url(sums["browser_download_url"])
sums_path = updates / "checksums.txt"
_http_download(sums["browser_download_url"], sums_path, max_bytes=MAX_DOWNLOAD_BYTES)
_verify_local_checksum(sums_path, filename, path)
def _find_verified_local_artifact(updates: Path) -> tuple[str, Path] | None:
"""Return (filename, path) of a staged artifact that passes checksums.txt."""
sums_path = updates / "checksums.txt"
if not sums_path.exists():
return None
for line in sums_path.read_text(encoding="utf-8").splitlines():
parts = line.split()
if len(parts) != 2:
continue
filename = parts[1]
if not (filename.endswith(".whl") or filename.endswith(".tar.gz")):
continue
dest = updates / filename
if not dest.exists():
continue
try:
_verify_local_checksum(sums_path, filename, dest)
except UpdateDownloadError:
dest.unlink(missing_ok=True) # corrupt staging — fall through to download
continue
return filename, dest
return None
def _resolve_release(version: str | None) -> tuple[str, dict]:
base, repo = _gitea_base(), _gitea_repo()
if version is not None:
tag = version if version.startswith(("v", "V")) else f"v{version}"
release = _http_get_json(f"{base}/api/v1/repos/{repo}/releases/tags/{tag}")
if not release:
raise UpdateDownloadError(f"no release found for version {version}")
return tag.lstrip("vV"), release
latest, release = _fetch_latest()
if release is None:
# newest candidate is a bare tag — fall back to its source archive
release = {
"tag_name": f"v{latest}",
"assets": [],
"tarball_url": f"{base}/api/v1/repos/{repo}/archive/v{latest}.tar.gz",
}
return latest, release
def download_update(version: str | None = None) -> DownloadResult:
"""Stage a release artifact locally; never installs it."""
try:
resolved_version, release = _resolve_release(version)
except UpdateDownloadError:
raise
except Exception as exc:
raise UpdateDownloadError(f"could not resolve release: {exc}") from exc
updates = _updates_dir(resolved_version)
updates.mkdir(parents=True, exist_ok=True)
reused = _find_verified_local_artifact(updates)
if reused is not None:
filename, dest = reused
return DownloadResult(
version=resolved_version,
file=filename,
path=str(dest),
suggested_command=f"uv pip install {dest} # then restart the backend",
)
try:
filename, url = _pick_asset(release)
_check_asset_url(url)
except UpdateDownloadError:
raise
except Exception as exc:
raise UpdateDownloadError(f"could not resolve release: {exc}") from exc
dest = updates / filename
try:
_http_download(url, dest, max_bytes=MAX_DOWNLOAD_BYTES)
_verify_checksum(release, updates, filename, dest)
except Exception:
dest.unlink(missing_ok=True)
raise
return DownloadResult(
version=resolved_version,
file=filename,
path=str(dest),
suggested_command=f"uv pip install {dest} # then restart the backend",
)
+296
View File
@@ -0,0 +1,296 @@
"""Self-update planner and standalone updater runner.
This module is STDLIB-ONLY on purpose: the file gets copied into the
staging dir and executed after the EvoScientist package has been replaced
on disk, so importing anything from EvoScientist would load mixed code.
"""
from __future__ import annotations
import json
import os
import shutil
import subprocess
import sys
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Mapping, Sequence
DEPLOY_DOCKER = "docker"
DEPLOY_SYSTEMD = "systemd"
DEPLOY_UVTOOL = "cli-uvtool"
DEPLOY_PIPX = "cli-pipx"
DEPLOY_VENV = "cli-venv"
PID_WAIT_TIMEOUT = 300.0 # seconds waiting for the parent to exit
@dataclass(frozen=True)
class SystemProbe:
in_container: bool
invocation_id: str | None
uv_tool_pkg: bool
pipx_pkg: bool
def detect_deployment(probe: SystemProbe) -> str:
if probe.in_container:
return DEPLOY_DOCKER
if probe.invocation_id:
return DEPLOY_SYSTEMD
if probe.uv_tool_pkg:
return DEPLOY_UVTOOL
if probe.pipx_pkg:
return DEPLOY_PIPX
return DEPLOY_VENV
def _uv_tool_has_package(environ: Mapping[str, str]) -> bool:
uv = shutil.which("uv")
if uv is None:
return False
try:
proc = subprocess.run(
[uv, "tool", "dir"], capture_output=True, text=True, timeout=10
)
except (OSError, subprocess.SubprocessError):
return False
if proc.returncode != 0:
return False
tool_dir = Path(proc.stdout.strip())
return (tool_dir / "evoscientist").is_dir() or (tool_dir / "EvoScientist").is_dir()
def _pipx_has_package() -> bool:
home = Path.home()
for base in (
home / ".local" / "pipx" / "venvs",
home / ".local" / "share" / "pipx" / "venvs",
):
if (base / "evoscientist").is_dir() or (base / "EvoScientist").is_dir():
return True
return False
def probe_system(environ: Mapping[str, str] | None = None) -> SystemProbe:
env = os.environ if environ is None else environ
in_container = Path("/.dockerenv").exists() or bool(env.get("container"))
return SystemProbe(
in_container=in_container,
invocation_id=env.get("INVOCATION_ID") or None,
uv_tool_pkg=False if in_container else _uv_tool_has_package(env),
pipx_pkg=False if in_container else _pipx_has_package(),
)
def build_install_command(kind: str, artifact: Path, python: str) -> list[str]:
if kind in (DEPLOY_UVTOOL, DEPLOY_SYSTEMD):
# systemd only decides the restart strategy; the install command for a
# systemd + non-uv host is overridden by the caller (build_plan).
return ["uv", "tool", "install", "--force", str(artifact)]
if kind == DEPLOY_PIPX:
return ["pipx", "install", "--force", str(artifact)]
return [python, "-m", "pip", "install", "--force-reinstall", str(artifact)]
class UpdateInProgressError(RuntimeError):
"""Another update holds the plan lock."""
def build_plan(
*,
version: str,
artifact: Path,
kind: str,
python: str,
argv: Sequence[str],
cwd: Path,
deploy_mode: str,
staging: Path,
systemd_unit: str | None,
) -> dict:
target_dir = staging / f"v{version}"
if kind == DEPLOY_SYSTEMD:
strategy, unit = "systemd", systemd_unit or "evoscientist.service"
elif kind == DEPLOY_DOCKER:
strategy, unit = "none", None
else:
strategy, unit = "cli", None
try:
from importlib.metadata import version as _pkg_version
previous = _pkg_version("EvoScientist")
except Exception:
previous = "0.0.0-dev"
return {
"version": version,
"artifact": str(artifact),
"install_command": build_install_command(kind, artifact, python),
"restart": {"strategy": strategy, "unit": unit},
"respawn_command": list(argv),
"cwd": str(cwd),
"deploy_mode": deploy_mode,
"previous_version": previous,
"result_path": str(target_dir / "update-result.json"),
"log_path": str(target_dir / "updater.log"),
}
def _lock_path(plan_path: Path) -> Path:
return plan_path.with_suffix(".lock")
def write_plan_locked(plan: dict, plan_path: Path) -> None:
"""Write plan.json, claiming the update lock atomically.
Lock semantics: the lock FILE existing means an update is in progress.
Uses O_CREAT|O_EXCL instead of filelock so this module stays stdlib-only
for the copied runner; the lock is released by release_plan_lock (called
by the runner on completion or by the next operation after a crash).
"""
plan_path.parent.mkdir(parents=True, exist_ok=True)
lock_path = _lock_path(plan_path)
try:
fd = os.open(lock_path, os.O_CREAT | os.O_EXCL | os.O_WRONLY)
except FileExistsError:
raise UpdateInProgressError(str(plan_path)) from None
else:
os.close(fd)
plan_path.write_text(json.dumps(plan, indent=2), encoding="utf-8")
def release_plan_lock(plan_path: Path) -> None:
_lock_path(plan_path).unlink(missing_ok=True)
def copy_updater(dest_dir: Path) -> Path:
dest_dir.mkdir(parents=True, exist_ok=True)
dest = dest_dir / "updater.py"
shutil.copyfile(Path(__file__).resolve(), dest)
return dest
def spawn_updater(updater_copy: Path, plan_path: Path, parent_pid: int) -> subprocess.Popen:
return subprocess.Popen(
[
sys.executable,
str(updater_copy),
"--plan",
str(plan_path),
"--parent-pid",
str(parent_pid),
],
start_new_session=True,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
stdin=subprocess.DEVNULL,
)
# ---------------------------------------------------------------------------
# Standalone runner (executed from the copied updater.py)
# ---------------------------------------------------------------------------
def _pid_alive_default(pid: int) -> bool:
try:
os.kill(pid, 0)
except OSError:
return False
return True
def wait_for_pid_exit(pid, timeout, *, poll=0.5, sleep=time.sleep, pid_alive=None):
alive = pid_alive or _pid_alive_default
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
if not alive(pid):
return True
sleep(poll)
return not alive(pid)
def _append_log(log_path: Path, text: str) -> None:
with open(log_path, "a", encoding="utf-8") as fh:
fh.write(text)
def _write_result(plan: dict, staging: Path, status: str, stage: str | None, message: str) -> None:
payload = {
"status": status,
"stage": stage,
"version": plan["version"],
"previous_version": plan.get("previous_version", ""),
"message": message,
"log": plan["log_path"],
}
result_path = Path(plan["result_path"])
result_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
# fixed read location for the status endpoint
(staging / "last-result.json").write_text(json.dumps(payload, indent=2), encoding="utf-8")
def run_plan(plan_path, *, parent_pid, sleep=time.sleep, pid_alive=None, run=subprocess.run, popen=subprocess.Popen):
plan_path = Path(plan_path)
plan = json.loads(plan_path.read_text(encoding="utf-8"))
staging = Path(plan["result_path"]).parent.parent # <staging>/v<ver>/update-result.json
log_path = Path(plan["log_path"])
_append_log(log_path, f"updater start: pid={os.getpid()} parent={parent_pid}\n")
try:
if not wait_for_pid_exit(parent_pid, PID_WAIT_TIMEOUT, sleep=sleep, pid_alive=pid_alive):
_write_result(plan, staging, "failed", "wait", f"parent {parent_pid} still alive after {PID_WAIT_TIMEOUT}s")
return 1
proc = run(plan["install_command"], capture_output=True, text=True, cwd=plan["cwd"])
_append_log(log_path, f"install rc={proc.returncode}\n{proc.stdout}\n{proc.stderr}\n")
if proc.returncode != 0:
_write_result(plan, staging, "failed", "install", proc.stderr.strip() or "install failed")
return 1
strategy = plan["restart"]["strategy"]
if strategy == "cli":
# The respawned service is long-lived, so we cannot judge success by
# waiting on it; only a spawn failure (missing binary etc.) counts
# as a restart failure.
try:
popen(
plan["respawn_command"],
cwd=plan["cwd"],
start_new_session=True,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
stdin=subprocess.DEVNULL,
)
except OSError as exc:
_write_result(plan, staging, "installed_restart_failed", "restart", str(exc))
return 1
elif strategy == "systemd":
unit = plan["restart"]["unit"]
proc = run(["systemctl", "--user", "restart", unit], capture_output=True, text=True)
_append_log(log_path, f"systemctl rc={proc.returncode}\n{proc.stdout}\n{proc.stderr}\n")
if proc.returncode != 0:
_write_result(
plan, staging, "installed_restart_failed", "restart",
f"systemctl --user restart {unit} failed; run it manually",
)
return 1
_write_result(plan, staging, "success", None, "update applied")
return 0
finally:
release_plan_lock(plan_path)
def main(argv=None):
import argparse
parser = argparse.ArgumentParser(prog="EvoScientist.updater")
parser.add_argument("--plan", required=True)
parser.add_argument("--parent-pid", required=True, type=int)
args = parser.parse_args(argv)
return run_plan(Path(args.plan), parent_pid=args.parent_pid)
if __name__ == "__main__":
sys.exit(main())
+6
View File
@@ -0,0 +1,6 @@
"""Fail-open model usage capture for the integrated EvoScientist WebUI."""
from .callback import UsageModelIdentity, attach_usage_callback
from .identity import prepare_usage_environment
__all__ = ["UsageModelIdentity", "attach_usage_callback", "prepare_usage_environment"]
+356
View File
@@ -0,0 +1,356 @@
"""LangChain terminal callback that emits one UsageEvent per real model call."""
from __future__ import annotations
import copy
import logging
import os
import threading
from collections.abc import Mapping
from dataclasses import dataclass, field
from datetime import UTC, datetime
from typing import Any
from uuid import UUID
from langchain_core.callbacks import BaseCallbackHandler
from .schema import MAX_SAFE_TOKEN_INTEGER, UsageEventV1, UsageScope
from .spool import get_usage_spool, mark_tracking_degraded
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class UsageModelIdentity:
provider_profile_id: str
provider_revision: str | None
provider_adapter: str
model_alias: str
upstream_model_id: str
@dataclass
class _CallState:
started_at: datetime
parent_run_id: str | None
metadata: dict[str, Any] = field(default_factory=dict)
last_usage: dict[str, Any] | None = None
provider_request_id: str | None = None
def usage_tracking_requested() -> bool:
"""Return whether this process was launched in usage-tracking mode."""
enabled = os.getenv("EVOSCIENTIST_USAGE_TRACKING", "").strip().lower()
return enabled in {"1", "true", "yes", "on"}
def usage_tracking_enabled() -> bool:
"""Return whether the integrated launcher supplied a complete usage sink."""
return usage_tracking_requested() and all(
os.getenv(name, "").strip()
for name in (
"EVOSCIENTIST_USAGE_SINK_URL",
"EVOSCIENTIST_USAGE_SINK_TOKEN",
"EVOSCIENTIST_USAGE_DEPLOYMENT_ID",
"EVOSCIENTIST_WORKSPACE_ID",
"EVOSCIENTIST_USAGE_SPOOL_DIR",
)
)
def _text(value: Any, *, limit: int = 256) -> str | None:
if value is None:
return None
rendered = "".join(
character
for character in str(value)
if ord(character) > 31 and ord(character) != 127
)
return rendered[:limit] if rendered else None
def _metadata_value(metadata: Mapping[str, Any], *keys: str) -> Any:
for key in keys:
value = metadata.get(key)
if value is not None:
return value
return None
def _scope(metadata: Mapping[str, Any]) -> UsageScope:
explicit = metadata.get("usage_scope")
allowed = {
"main", "sync_subagent", "async_subagent", "tool_selector", "summarizer",
"memory", "scheduler", "autoskills", "diagnostic", "skill_eval", "unattributed",
}
if explicit in allowed:
return explicit # type: ignore[return-value]
if metadata.get("lc_source") == "summarization":
return "summarizer"
run_kind = metadata.get("run_kind")
if run_kind == "scheduled_task":
return "scheduler"
if run_kind == "evomemory_autoskills":
return "autoskills"
if isinstance(run_kind, str) and run_kind.startswith("evomemory_"):
return "memory"
# deepagents attaches this stable metadata field to every compiled agent
# graph. Remote async runs carry an explicit usage_scope and have already
# returned above; remaining non-main names are synchronous subagents.
agent_name = metadata.get("lc_agent_name")
if agent_name == "EvoScientist":
return "main"
if isinstance(agent_name, str) and agent_name:
return "sync_subagent"
if metadata.get("async_task_id") or metadata.get("source_session_id"):
return "async_subagent"
if metadata.get("thread_id") or metadata.get("langgraph_thread_id"):
return "main"
return "unattributed"
def _clean_details(value: Any) -> dict[str, int]:
if not isinstance(value, Mapping):
return {}
result: dict[str, int] = {}
for key, count in value.items():
if (
isinstance(key, str)
and 0 < len(key) <= 128
and isinstance(count, int)
and not isinstance(count, bool)
and 0 <= count <= MAX_SAFE_TOKEN_INTEGER
):
result[key] = count
return result
def _normalize_usage(value: Any) -> dict[str, Any] | None:
if not isinstance(value, Mapping):
return None
input_tokens = value.get("input_tokens", value.get("prompt_tokens"))
output_tokens = value.get("output_tokens", value.get("completion_tokens"))
if not all(
isinstance(token, int)
and not isinstance(token, bool)
and 0 <= token <= MAX_SAFE_TOKEN_INTEGER
for token in (input_tokens, output_tokens)
):
return None
if input_tokens + output_tokens > MAX_SAFE_TOKEN_INTEGER:
return None
provider_total = value.get("total_tokens")
if not (
isinstance(provider_total, int)
and not isinstance(provider_total, bool)
and 0 <= provider_total <= MAX_SAFE_TOKEN_INTEGER
):
provider_total = None
return {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"provider_total_tokens": provider_total,
"input_token_details": _clean_details(value.get("input_token_details")),
"output_token_details": _clean_details(value.get("output_token_details")),
}
def _message_from_chunk(chunk: Any) -> Any:
return getattr(chunk, "message", chunk)
def _message_usage(message: Any) -> dict[str, Any] | None:
return _normalize_usage(getattr(message, "usage_metadata", None))
def _request_id(message: Any) -> str | None:
metadata = getattr(message, "response_metadata", None)
if not isinstance(metadata, Mapping):
return None
return _text(_metadata_value(metadata, "request_id", "id", "x_request_id"))
class UsageCaptureCallback(BaseCallbackHandler):
"""Capture final provider usage without influencing the model invocation."""
def __init__(self, identity: UsageModelIdentity) -> None:
self.identity = identity
self._lock = threading.Lock()
self._calls: dict[str, _CallState] = {}
def on_chat_model_start(
self,
serialized: dict[str, Any],
messages: list[list[Any]],
*,
run_id: UUID,
parent_run_id: UUID | None = None,
metadata: dict[str, Any] | None = None,
**kwargs: Any,
) -> None:
try:
with self._lock:
self._calls[str(run_id)] = _CallState(
started_at=datetime.now(UTC),
parent_run_id=str(parent_run_id) if parent_run_id else None,
metadata=dict(metadata or {}),
)
except Exception:
logger.exception("Usage callback could not initialize run state")
def on_llm_new_token(
self, token: Any, *, run_id: UUID, chunk: Any = None, **kwargs: Any
) -> None:
try:
message = _message_from_chunk(chunk)
usage = _message_usage(message)
if usage is None:
return
with self._lock:
state = self._calls.get(str(run_id))
if state:
state.last_usage = usage
state.provider_request_id = _request_id(message) or state.provider_request_id
except Exception:
logger.exception("Usage callback could not inspect a stream chunk")
def on_llm_end(self, response: Any, *, run_id: UUID, **kwargs: Any) -> None:
usage: dict[str, Any] | None = None
request_id: str | None = None
try:
for group in getattr(response, "generations", []) or []:
for generation in group or []:
message = getattr(generation, "message", generation)
usage = _message_usage(message) or usage
request_id = _request_id(message) or request_id
if usage is None:
output = getattr(response, "llm_output", None)
if isinstance(output, Mapping):
usage = _normalize_usage(output.get("token_usage") or output.get("usage"))
except Exception:
logger.exception("Usage callback could not inspect terminal usage")
try:
self._finish(str(run_id), usage=usage, request_id=request_id)
except Exception:
logger.exception("Usage callback could not emit terminal usage")
def on_llm_error(self, error: BaseException, *, run_id: UUID, **kwargs: Any) -> None:
try:
self._finish(str(run_id), usage=None, request_id=None)
except Exception:
logger.exception("Usage callback could not emit unknown terminal usage")
def _finish(
self, run_id: str, *, usage: dict[str, Any] | None, request_id: str | None
) -> None:
with self._lock:
state = self._calls.pop(run_id, None)
if state is None:
return
usage = usage or state.last_usage
completed = datetime.now(UTC)
metadata = state.metadata
confirmed = usage is not None
event = UsageEventV1(
schema_version=1,
event_id=(
f"{os.environ['EVOSCIENTIST_USAGE_DEPLOYMENT_ID']}:{run_id}:callback_final:1"
),
event_type="usage_observed",
source="callback_final",
authority_class="observed_final",
revision=1,
deployment_id=os.environ["EVOSCIENTIST_USAGE_DEPLOYMENT_ID"],
workspace_id=os.environ["EVOSCIENTIST_WORKSPACE_ID"],
model_call_id=run_id,
parent_run_id=state.parent_run_id,
provider_request_id=request_id or state.provider_request_id,
thread_id=_text(
_metadata_value(metadata, "thread_id", "langgraph_thread_id")
),
source_session_id=_text(
_metadata_value(
metadata, "source_session_id", "evomemory_source_session_id"
)
),
turn_id=_text(
_metadata_value(metadata, "turn_id", "evomemory_source_turn_id")
),
workspace_dir=_text(
_metadata_value(metadata, "workspace_dir"), limit=4096
)
or _text(os.getenv("EVOSCIENTIST_WORKSPACE_DIR"), limit=4096),
scope=_scope(metadata),
source_agent=_text(
_metadata_value(metadata, "source_agent", "evomemory_source_agent")
),
provider_profile_id=_text(
self.identity.provider_profile_id, limit=512
)
or "unknown",
provider_revision=_text(self.identity.provider_revision),
provider_adapter=_text(self.identity.provider_adapter, limit=512)
or "unknown",
model_alias=_text(self.identity.model_alias, limit=512) or "unknown",
upstream_model_id=_text(
self.identity.upstream_model_id, limit=512
)
or "unknown",
usage_status="confirmed" if confirmed else "unknown",
input_tokens=usage["input_tokens"] if usage else None,
output_tokens=usage["output_tokens"] if usage else None,
provider_total_tokens=usage["provider_total_tokens"] if usage else None,
input_token_details=usage["input_token_details"] if usage else {},
output_token_details=usage["output_token_details"] if usage else {},
started_at=state.started_at,
observed_at=completed,
completed_at=completed,
)
get_usage_spool().enqueue(event)
def _merged_callbacks(existing: Any, callback: UsageCaptureCallback) -> Any:
if existing is None:
return [callback]
if isinstance(existing, (list, tuple)):
return [*existing, callback]
if hasattr(existing, "add_handler"):
manager = copy.copy(existing)
manager.add_handler(callback, inherit=True)
return manager
raise TypeError("unsupported callbacks value")
def attach_usage_callback(model: Any, identity: UsageModelIdentity) -> Any:
"""Return a callback-enabled model copy, or the original model on degradation."""
if not usage_tracking_enabled():
return model
try:
from langchain_core.language_models import BaseChatModel
if not isinstance(model, BaseChatModel):
raise TypeError("model is not a BaseChatModel")
callback = UsageCaptureCallback(identity)
updated = model.model_copy(
update={"callbacks": _merged_callbacks(model.callbacks, callback)}
)
if not isinstance(updated, BaseChatModel):
raise TypeError("callback model_copy did not return a BaseChatModel")
# Tool Selector depends on a second metadata-only model copy. Verify the
# contract before enabling tracking for this model so a selector call
# can never be silently counted as a main-agent call.
probe = updated.model_copy(
update={"metadata": {**(updated.metadata or {}), "usage_scope": "tool_selector"}}
)
if not isinstance(probe, BaseChatModel):
raise TypeError("selector metadata model_copy did not return a BaseChatModel")
# Start capability negotiation and heartbeat at backend startup rather
# than waiting for the first completed model call. This lets the UI
# distinguish a healthy empty database from unsupported tracking.
get_usage_spool()
return updated
except Exception:
mark_tracking_degraded("selector_model_copy_unsupported")
logger.exception("Usage callback injection failed; model behavior is unchanged")
return model
+122
View File
@@ -0,0 +1,122 @@
"""Stable deployment/workspace identity and launcher environment setup."""
from __future__ import annotations
import hashlib
import ntpath
import os
import secrets
import time
import unicodedata
import uuid
from collections.abc import Callable
from pathlib import Path
def _fsync_directory(path: Path) -> None:
try:
descriptor = os.open(path, os.O_RDONLY)
except OSError:
return
try:
os.fsync(descriptor)
except OSError:
pass
finally:
os.close(descriptor)
def _load_or_create(path: Path, factory: Callable[[], str]) -> str:
path.parent.mkdir(parents=True, exist_ok=True, mode=0o700)
value = factory()
try:
descriptor = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
except FileExistsError:
# Another launcher may have won O_EXCL but not completed its fsync yet.
for _ in range(100):
existing = path.read_text(encoding="utf-8").strip()
if existing:
return existing
time.sleep(0.01)
raise ValueError(f"Identity file is empty: {path}") from None
try:
os.write(descriptor, f"{value}\n".encode())
os.fsync(descriptor)
finally:
os.close(descriptor)
_fsync_directory(path.parent)
return value
def resolve_data_dir() -> Path:
configured = os.getenv("EVOSCIENTIST_DATA_DIR", "").strip()
if configured:
path = Path(configured)
if not path.is_absolute():
raise ValueError("EVOSCIENTIST_DATA_DIR must be absolute")
return path
return Path.home() / ".evoscientist"
def normalize_workspace_path_v1(
resolved_path: str, *, windows: bool | None = None
) -> str:
"""Normalize an already-real path using the frozen ws1 cross-platform rules."""
use_windows = os.name == "nt" if windows is None else windows
normalized = unicodedata.normalize("NFC", resolved_path)
if use_windows:
windows_path = ntpath.normcase(normalized)
_drive, tail = ntpath.splitdrive(windows_path)
normalized = windows_path.replace("\\", "/")
if tail in {"\\", "/"}:
return normalized.rstrip("/") + "/"
return normalized.rstrip("/")
if normalized == "/":
return normalized
return normalized.rstrip("/")
def workspace_identity_from_normalized(deployment_id: str, normalized_path: str) -> str:
"""Hash a normalized path; split out for cross-language contract fixtures."""
digest = hashlib.sha256(f"{deployment_id}\0{normalized_path}".encode()).hexdigest()
return f"ws1_{digest}"
def workspace_identity(deployment_id: str, workspace_dir: str | Path) -> str:
resolved = os.path.realpath(os.path.expanduser(str(workspace_dir)), strict=True)
normalized = normalize_workspace_path_v1(resolved)
return workspace_identity_from_normalized(deployment_id, normalized)
def prepare_usage_environment(
workspace_dir: str | Path, *, webui_port: int
) -> dict[str, str]:
"""Create stable secrets/identities and return the shared child environment."""
if not 1 <= webui_port <= 65535:
raise ValueError("webui_port must be an integer in [1, 65535]")
data_dir = resolve_data_dir().expanduser().resolve()
data_dir.mkdir(parents=True, exist_ok=True, mode=0o700)
deployment_id = _load_or_create(
data_dir / "deployment-id", lambda: str(uuid.uuid4())
)
# Reject corrupt identity files early instead of silently creating a second identity.
deployment_id = str(uuid.UUID(deployment_id))
sink_token = _load_or_create(
data_dir / "usage-sink-token", lambda: secrets.token_urlsafe(32)
)
workspace_id = workspace_identity(deployment_id, workspace_dir)
return {
"EVOSCIENTIST_DATA_DIR": str(data_dir),
"EVOSCIENTIST_USAGE_TRACKING": "true",
"EVOSCIENTIST_USAGE_SINK_URL": (
f"http://127.0.0.1:{webui_port}/api/usage/events"
),
"EVOSCIENTIST_USAGE_SINK_TOKEN": sink_token,
"EVOSCIENTIST_USAGE_DEPLOYMENT_ID": deployment_id,
"EVOSCIENTIST_WORKSPACE_ID": workspace_id,
"EVOSCIENTIST_USAGE_SPOOL_DIR": str(data_dir / "usage-spool"),
}

Some files were not shown because too many files have changed in this diff Show More