Compare commits
55 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 96c380caa3 | |||
| 59c77c65b1 | |||
| ae461d1c2d | |||
| 2a7cccd598 | |||
| aae8d0a379 | |||
| f3ca381ab3 | |||
| 8f6a568646 | |||
| 91e2a87be5 | |||
| 03771df508 | |||
| 90e773b30b | |||
| 40b896bcb5 | |||
| 174f03b92d | |||
| 194402fc88 | |||
| c6efdaa13f | |||
| 08aa0d0e05 | |||
| 8eff551bb6 | |||
| 862c1e9743 | |||
| d7711484ef | |||
| ba7d908276 | |||
| e57ecd4588 | |||
| 261845830d | |||
| 8efd4ad0ab | |||
| 01e674e1bc | |||
| 7f26ecc19a | |||
| 395baab7d5 | |||
| ccf4173990 | |||
| a57c52c676 | |||
| 384bc13a5b | |||
| 5622b40cf3 | |||
| 802f71bd46 | |||
| cc9dfb1cc9 | |||
| 12e4f34005 | |||
| b0f9a8d785 | |||
| 15cc389b3d | |||
| 3e67e64067 | |||
| bb9bed82e1 | |||
| 1a01fb5d74 | |||
| 087781556b | |||
| a1bfbd92ca | |||
| 421a664336 | |||
| 57176b359a | |||
| b2660fc38c | |||
| 940db565b3 | |||
| dbb6b7abde | |||
| c8c46eab16 | |||
| 0cc995eb80 | |||
| b1233d42dc | |||
| b2e28249fd | |||
| af4ae1aef5 | |||
| c46ae17084 | |||
| c21fc0a272 | |||
| 8a0ab17936 | |||
| 38668c4ce5 | |||
| 7a3fcc7c8e | |||
| e0acc6155e |
+37
-28
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -48,3 +48,8 @@ conversation_history/
|
||||
*meals/
|
||||
botpy.log
|
||||
large_tool_results/
|
||||
runs/
|
||||
|
||||
# local runtime artifacts (scope tokens, control DBs)
|
||||
.evoscientist/
|
||||
.release-state.json
|
||||
|
||||
@@ -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 内同名分支仅作防御。
|
||||
@@ -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`(涉及文件)→ 全净
|
||||
+166
-154
@@ -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,13 +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,
|
||||
cfg=None,
|
||||
chat_model=None,
|
||||
backend=None,
|
||||
memory_source_agent: str = "EvoScientist",
|
||||
snapshot_role: str = "primary",
|
||||
):
|
||||
"""Build the default middleware list.
|
||||
|
||||
@@ -663,27 +668,34 @@ 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()
|
||||
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
|
||||
@@ -695,10 +707,8 @@ 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.
|
||||
# ``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,
|
||||
@@ -711,27 +721,16 @@ def _get_default_middleware(
|
||||
),
|
||||
memory_scheduler=memory_scheduler,
|
||||
)
|
||||
# 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).
|
||||
# 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(
|
||||
@@ -740,17 +739,27 @@ def _get_default_middleware(
|
||||
),
|
||||
# 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,
|
||||
@@ -770,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:
|
||||
if not for_async_subagent and cfg.workspace_isolation != "required":
|
||||
from .middleware.background import BackgroundExecutionMiddleware
|
||||
|
||||
mw.append(BackgroundExecutionMiddleware())
|
||||
@@ -807,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
|
||||
@@ -824,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,
|
||||
@@ -878,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``,
|
||||
@@ -969,7 +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, cfg=cfg, chat_model=chat_model
|
||||
workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model, backend=be
|
||||
)
|
||||
|
||||
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
|
||||
|
||||
@@ -23,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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"))
|
||||
@@ -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)
|
||||
|
||||
@@ -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())
|
||||
@@ -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",
|
||||
|
||||
@@ -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}")
|
||||
@@ -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
|
||||
|
||||
@@ -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,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.
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,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
|
||||
# =============================================================================
|
||||
|
||||
@@ -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:
|
||||
|
||||
+92
-131
@@ -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):
|
||||
@@ -119,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
|
||||
# =============================================================================
|
||||
@@ -128,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
|
||||
@@ -279,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)
|
||||
@@ -398,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
|
||||
|
||||
@@ -426,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
|
||||
|
||||
@@ -458,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)
|
||||
@@ -727,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",
|
||||
@@ -836,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)
|
||||
@@ -906,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
|
||||
|
||||
@@ -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 {}),
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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."
|
||||
)
|
||||
@@ -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())
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
+1173
-40
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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",
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
@@ -1,638 +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"}
|
||||
|
||||
# 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 _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.
|
||||
"""
|
||||
# Anthropic: extended thinking
|
||||
if provider == "anthropic" 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 "reasoning" not in kwargs:
|
||||
if _is_ccproxy_codex():
|
||||
# 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":
|
||||
kwargs.setdefault("include_thoughts", True)
|
||||
|
||||
# Ollama: separate reasoning content from response for thinking models
|
||||
if provider == "ollama" 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
|
||||
"""
|
||||
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
|
||||
)
|
||||
_is_openai_proxy = False
|
||||
_original_provider: str | None = None
|
||||
if provider == "anthropic":
|
||||
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
|
||||
if base_url:
|
||||
kwargs["base_url"] = base_url
|
||||
api_key = os.environ.get("ANTHROPIC_API_KEY", "")
|
||||
if api_key:
|
||||
kwargs["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["base_url"] = base_url
|
||||
_is_openai_proxy = _is_ccproxy_codex()
|
||||
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["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["base_url"] = base_url
|
||||
api_key = os.environ.get(api_key_env, "")
|
||||
if api_key:
|
||||
kwargs["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["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["base_url"] = base_url
|
||||
api_key = os.environ.get(api_key_env, "")
|
||||
if api_key:
|
||||
kwargs["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["base_url"] = base_url
|
||||
|
||||
_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
|
||||
|
||||
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
|
||||
|
||||
# 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,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.
|
||||
|
||||
|
||||
+152
-50
@@ -873,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``
|
||||
@@ -901,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
|
||||
|
||||
|
||||
@@ -1017,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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)."""
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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."
|
||||
)
|
||||
@@ -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")
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
+14
-3
@@ -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: ``. 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,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -797,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.
|
||||
|
||||
@@ -812,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,
|
||||
@@ -821,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:
|
||||
|
||||
@@ -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,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",
|
||||
|
||||
@@ -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"" 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. )"
|
||||
" — 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()
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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())
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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"),
|
||||
}
|
||||
@@ -0,0 +1,170 @@
|
||||
"""UsageEvent v1 model shared by the callback and durable spool."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
MAX_SAFE_TOKEN_INTEGER = 9_007_199_254_740_991
|
||||
_UTC_TIMESTAMP = re.compile(r"^\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d{1,9})?Z$")
|
||||
UsageScope = Literal[
|
||||
"main",
|
||||
"sync_subagent",
|
||||
"async_subagent",
|
||||
"tool_selector",
|
||||
"summarizer",
|
||||
"memory",
|
||||
"scheduler",
|
||||
"autoskills",
|
||||
"diagnostic",
|
||||
"skill_eval",
|
||||
"unattributed",
|
||||
]
|
||||
|
||||
|
||||
class UsageEventV1(BaseModel):
|
||||
"""Validated terminal observation for exactly one LangChain model run."""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
schema_version: Literal[1]
|
||||
event_id: str = Field(min_length=1, max_length=1024)
|
||||
event_type: Literal["usage_observed"]
|
||||
source: Literal["callback_final"]
|
||||
authority_class: Literal["observed_final"]
|
||||
revision: Literal[1]
|
||||
|
||||
deployment_id: str = Field(min_length=1, max_length=256)
|
||||
workspace_id: str = Field(min_length=1, max_length=256)
|
||||
model_call_id: str = Field(min_length=1, max_length=256)
|
||||
parent_run_id: str | None = Field(min_length=1, max_length=256)
|
||||
provider_request_id: str | None = Field(min_length=1, max_length=256)
|
||||
thread_id: str | None = Field(min_length=1, max_length=256)
|
||||
source_session_id: str | None = Field(min_length=1, max_length=256)
|
||||
turn_id: str | None = Field(min_length=1, max_length=256)
|
||||
workspace_dir: str | None = Field(min_length=1, max_length=4096)
|
||||
scope: UsageScope
|
||||
source_agent: str | None = Field(min_length=1, max_length=256)
|
||||
|
||||
provider_profile_id: str = Field(min_length=1, max_length=512)
|
||||
provider_revision: str | None = Field(min_length=1, max_length=256)
|
||||
provider_adapter: str = Field(min_length=1, max_length=512)
|
||||
model_alias: str = Field(min_length=1, max_length=512)
|
||||
upstream_model_id: str = Field(min_length=1, max_length=512)
|
||||
|
||||
usage_status: Literal["confirmed", "unknown"]
|
||||
input_tokens: int | None = Field(ge=0, le=MAX_SAFE_TOKEN_INTEGER)
|
||||
output_tokens: int | None = Field(ge=0, le=MAX_SAFE_TOKEN_INTEGER)
|
||||
provider_total_tokens: int | None = Field(ge=0, le=MAX_SAFE_TOKEN_INTEGER)
|
||||
input_token_details: dict[str, int]
|
||||
output_token_details: dict[str, int]
|
||||
|
||||
started_at: datetime | None
|
||||
observed_at: datetime
|
||||
completed_at: datetime
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def validate_json_types(cls, value: object) -> object:
|
||||
if not isinstance(value, dict):
|
||||
return value
|
||||
integer_fields = (
|
||||
"schema_version",
|
||||
"revision",
|
||||
"input_tokens",
|
||||
"output_tokens",
|
||||
"provider_total_tokens",
|
||||
)
|
||||
for name in integer_fields:
|
||||
item = value.get(name)
|
||||
if item is not None and (
|
||||
not isinstance(item, int) or isinstance(item, bool)
|
||||
):
|
||||
raise ValueError(f"{name} must be an integer or null")
|
||||
for name in ("input_token_details", "output_token_details"):
|
||||
details = value.get(name)
|
||||
if isinstance(details, dict) and any(
|
||||
not isinstance(item, int) or isinstance(item, bool)
|
||||
for item in details.values()
|
||||
):
|
||||
raise ValueError(f"{name} values must be integers")
|
||||
for name in ("started_at", "observed_at", "completed_at"):
|
||||
timestamp = value.get(name)
|
||||
if timestamp is not None and not isinstance(timestamp, (str, datetime)):
|
||||
raise ValueError(f"{name} must be an RFC 3339 UTC timestamp")
|
||||
if isinstance(timestamp, str) and not _UTC_TIMESTAMP.fullmatch(timestamp):
|
||||
raise ValueError(f"{name} must be an RFC 3339 UTC timestamp")
|
||||
return value
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_contract(self) -> UsageEventV1:
|
||||
expected = f"{self.deployment_id}:{self.model_call_id}:callback_final:1"
|
||||
if self.event_id != expected:
|
||||
raise ValueError("event_id does not match the v1 identity")
|
||||
if self.usage_status == "unknown":
|
||||
if any(
|
||||
value is not None
|
||||
for value in (
|
||||
self.input_tokens,
|
||||
self.output_tokens,
|
||||
self.provider_total_tokens,
|
||||
)
|
||||
):
|
||||
raise ValueError("unknown usage must have null token fields")
|
||||
else:
|
||||
if self.input_tokens is None or self.output_tokens is None:
|
||||
raise ValueError("confirmed usage requires input and output tokens")
|
||||
if self.input_tokens + self.output_tokens > MAX_SAFE_TOKEN_INTEGER:
|
||||
raise ValueError("input plus output tokens exceeds safe integer")
|
||||
for details in (self.input_token_details, self.output_token_details):
|
||||
if any(
|
||||
not key
|
||||
or len(key) > 128
|
||||
or any(
|
||||
ord(character) <= 31 or ord(character) == 127 for character in key
|
||||
)
|
||||
or value < 0
|
||||
or value > MAX_SAFE_TOKEN_INTEGER
|
||||
for key, value in details.items()
|
||||
):
|
||||
raise ValueError("invalid token details")
|
||||
for name in (
|
||||
"event_id",
|
||||
"deployment_id",
|
||||
"workspace_id",
|
||||
"model_call_id",
|
||||
"parent_run_id",
|
||||
"provider_request_id",
|
||||
"thread_id",
|
||||
"source_session_id",
|
||||
"turn_id",
|
||||
"workspace_dir",
|
||||
"source_agent",
|
||||
"provider_profile_id",
|
||||
"provider_revision",
|
||||
"provider_adapter",
|
||||
"model_alias",
|
||||
"upstream_model_id",
|
||||
):
|
||||
text = getattr(self, name)
|
||||
if text is not None and any(
|
||||
ord(character) <= 31 or ord(character) == 127 for character in text
|
||||
):
|
||||
raise ValueError(f"{name} contains a control character")
|
||||
for timestamp in (self.started_at, self.observed_at, self.completed_at):
|
||||
if timestamp is not None:
|
||||
offset = timestamp.utcoffset()
|
||||
if (
|
||||
timestamp.tzinfo is None
|
||||
or offset is None
|
||||
or offset.total_seconds() != 0
|
||||
):
|
||||
raise ValueError("usage timestamps must be UTC")
|
||||
if self.started_at and (
|
||||
self.started_at > self.observed_at or self.started_at > self.completed_at
|
||||
):
|
||||
raise ValueError("started_at cannot follow terminal timestamps")
|
||||
return self
|
||||
@@ -0,0 +1,540 @@
|
||||
"""Local durable outbox and background sender for UsageEvent v1."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import atexit
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
import httpx
|
||||
|
||||
from .schema import UsageEventV1
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _now() -> str:
|
||||
return datetime.now(UTC).isoformat().replace("+00:00", "Z")
|
||||
|
||||
|
||||
def _int_env(name: str, default: int, minimum: int = 1) -> int:
|
||||
try:
|
||||
return max(minimum, int(os.getenv(name, str(default))))
|
||||
except ValueError:
|
||||
return default
|
||||
|
||||
|
||||
def _float_env(name: str, default: float) -> float:
|
||||
try:
|
||||
return max(0.05, float(os.getenv(name, str(default))))
|
||||
except ValueError:
|
||||
return default
|
||||
|
||||
|
||||
def _fsync_dir(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 _endpoint(sink_url: str, suffix: str) -> str:
|
||||
parts = urlsplit(sink_url)
|
||||
base = parts.path.removesuffix("/api/usage/events")
|
||||
return urlunsplit((parts.scheme, parts.netloc, f"{base}{suffix}", "", ""))
|
||||
|
||||
|
||||
class UsageSpool:
|
||||
"""Synchronous durable enqueue with an asynchronous at-least-once sender."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.sink_url = os.environ["EVOSCIENTIST_USAGE_SINK_URL"]
|
||||
self.token = os.environ["EVOSCIENTIST_USAGE_SINK_TOKEN"]
|
||||
self.deployment_id = os.environ["EVOSCIENTIST_USAGE_DEPLOYMENT_ID"]
|
||||
self.workspace_id = os.environ["EVOSCIENTIST_WORKSPACE_ID"]
|
||||
self.root = Path(os.environ["EVOSCIENTIST_USAGE_SPOOL_DIR"])
|
||||
if not self.root.is_absolute():
|
||||
raise ValueError("EVOSCIENTIST_USAGE_SPOOL_DIR must be absolute")
|
||||
self.tmp = self.root / "tmp"
|
||||
self.pending = self.root / "pending"
|
||||
self.inflight = self.root / "inflight"
|
||||
self.quarantine = self.root / "quarantine"
|
||||
for directory in (self.tmp, self.pending, self.inflight, self.quarantine):
|
||||
directory.mkdir(parents=True, exist_ok=True, mode=0o700)
|
||||
self.status_path = self.root / "status.json"
|
||||
self.max_files = _int_env("EVOSCIENTIST_USAGE_SPOOL_MAX_FILES", 100_000)
|
||||
self.max_bytes = _int_env("EVOSCIENTIST_USAGE_SPOOL_MAX_BYTES", 1_073_741_824)
|
||||
self.max_event_bytes = _int_env("EVOSCIENTIST_USAGE_MAX_EVENT_BYTES", 262_144)
|
||||
self.lease_seconds = _int_env("EVOSCIENTIST_USAGE_INFLIGHT_LEASE_SECONDS", 120)
|
||||
self.heartbeat_interval = _float_env(
|
||||
"EVOSCIENTIST_USAGE_HEARTBEAT_INTERVAL_SECONDS", 15
|
||||
)
|
||||
self.retry_initial = _float_env("EVOSCIENTIST_USAGE_RETRY_INITIAL_SECONDS", 1)
|
||||
self.retry_max = _float_env("EVOSCIENTIST_USAGE_RETRY_MAX_SECONDS", 60)
|
||||
self.connect_timeout = _float_env(
|
||||
"EVOSCIENTIST_USAGE_HTTP_CONNECT_TIMEOUT_SECONDS", 1
|
||||
)
|
||||
self.request_timeout = _float_env("EVOSCIENTIST_USAGE_HTTP_TIMEOUT_SECONDS", 3)
|
||||
self.unsupported_reprobe = _float_env(
|
||||
"EVOSCIENTIST_USAGE_UNSUPPORTED_REPROBE_SECONDS", 300
|
||||
)
|
||||
self.schema_reprobe = _float_env(
|
||||
"EVOSCIENTIST_USAGE_SCHEMA_REPROBE_SECONDS", 60
|
||||
)
|
||||
self.quarantine_retention_days = _int_env(
|
||||
"EVOSCIENTIST_USAGE_QUARANTINE_RETENTION_DAYS", 90
|
||||
)
|
||||
self._state_lock = threading.Lock()
|
||||
self._size_lock = threading.Lock()
|
||||
self._cached_files = 0
|
||||
self._cached_bytes = 0
|
||||
self._size_cache_at = 0.0
|
||||
self.first_loss_at: str | None = None
|
||||
self.degraded_reason: str | None = None
|
||||
self.last_error_code: str | None = None
|
||||
self._load_status()
|
||||
self._clean_expired_quarantine()
|
||||
self._clean_stale_local_artifacts()
|
||||
self._last_cleanup_at = time.monotonic()
|
||||
self._stop = threading.Event()
|
||||
self._wake = threading.Event()
|
||||
self._thread = threading.Thread(
|
||||
target=self._run, name="evoscientist-usage-sender", daemon=True
|
||||
)
|
||||
self._thread.start()
|
||||
|
||||
def _load_status(self) -> None:
|
||||
try:
|
||||
data = json.loads(self.status_path.read_text(encoding="utf-8"))
|
||||
self.first_loss_at = data.get("first_loss_at")
|
||||
self.degraded_reason = data.get("tracking_degraded_reason")
|
||||
except (OSError, ValueError, TypeError):
|
||||
pass
|
||||
|
||||
def _persist_status(self) -> None:
|
||||
temporary = self.root / f"status.{os.getpid()}.{uuid.uuid4().hex}.tmp"
|
||||
payload = json.dumps(
|
||||
{
|
||||
"first_loss_at": self.first_loss_at,
|
||||
"tracking_degraded_reason": self.degraded_reason,
|
||||
},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
).encode()
|
||||
try:
|
||||
descriptor = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
|
||||
try:
|
||||
os.write(descriptor, payload)
|
||||
os.fsync(descriptor)
|
||||
finally:
|
||||
os.close(descriptor)
|
||||
os.replace(temporary, self.status_path)
|
||||
_fsync_dir(self.root)
|
||||
except OSError:
|
||||
logger.exception("Could not persist usage tracking degraded status")
|
||||
try:
|
||||
temporary.unlink(missing_ok=True)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
def mark_degraded(self, reason: str) -> None:
|
||||
with self._state_lock:
|
||||
if self.first_loss_at is None:
|
||||
self.first_loss_at = _now()
|
||||
if self.degraded_reason is None:
|
||||
self.degraded_reason = reason[:512]
|
||||
self._persist_status()
|
||||
self._wake.set()
|
||||
|
||||
def _spool_size(self) -> tuple[int, int]:
|
||||
with self._size_lock:
|
||||
if time.monotonic() - self._size_cache_at < 1:
|
||||
return self._cached_files, self._cached_bytes
|
||||
count = 0
|
||||
size = 0
|
||||
for directory in (self.pending, self.inflight, self.quarantine):
|
||||
try:
|
||||
for item in os.scandir(directory):
|
||||
if item.is_file() and item.name.endswith(".json"):
|
||||
count += 1
|
||||
try:
|
||||
size += item.stat().st_size
|
||||
except OSError:
|
||||
pass
|
||||
except OSError:
|
||||
continue
|
||||
with self._size_lock:
|
||||
self._cached_files = count
|
||||
self._cached_bytes = size
|
||||
self._size_cache_at = time.monotonic()
|
||||
return count, size
|
||||
|
||||
def _adjust_spool_size(self, count: int, size: int) -> None:
|
||||
with self._size_lock:
|
||||
self._cached_files = max(0, self._cached_files + count)
|
||||
self._cached_bytes = max(0, self._cached_bytes + size)
|
||||
|
||||
def enqueue(self, event: UsageEventV1) -> None:
|
||||
try:
|
||||
payload = json.dumps(
|
||||
event.model_dump(mode="json"),
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
if len(payload) > self.max_event_bytes:
|
||||
self.mark_degraded("event_too_large")
|
||||
logger.error("Usage event %s exceeds spool event limit", event.event_id)
|
||||
return
|
||||
count, size = self._spool_size()
|
||||
if count >= self.max_files or size + len(payload) > self.max_bytes:
|
||||
self.mark_degraded("spool_soft_limit_reached")
|
||||
logger.error("Usage spool soft limit reached; event was not persisted")
|
||||
return
|
||||
key = hashlib.sha256(event.event_id.encode()).hexdigest()
|
||||
temporary = self.tmp / f"{key}.{os.getpid()}.{uuid.uuid4().hex}.tmp"
|
||||
descriptor = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
|
||||
try:
|
||||
os.write(descriptor, payload)
|
||||
os.fsync(descriptor)
|
||||
finally:
|
||||
os.close(descriptor)
|
||||
lock_path = self.root / f"{key}.lock"
|
||||
deadline = time.monotonic() + 0.05
|
||||
while True:
|
||||
try:
|
||||
lock_fd = os.open(
|
||||
lock_path, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600
|
||||
)
|
||||
break
|
||||
except FileExistsError:
|
||||
if time.monotonic() < deadline:
|
||||
time.sleep(0.001)
|
||||
continue
|
||||
os.replace(
|
||||
temporary,
|
||||
self.quarantine
|
||||
/ f"{key}.lock-contention.{uuid.uuid4().hex}.json",
|
||||
)
|
||||
self._adjust_spool_size(1, len(payload))
|
||||
_fsync_dir(self.quarantine)
|
||||
logger.error(
|
||||
"Usage event %s quarantined after spool lock contention",
|
||||
event.event_id,
|
||||
)
|
||||
self.mark_degraded("spool_lock_contention")
|
||||
self._wake.set()
|
||||
return
|
||||
try:
|
||||
target = self.pending / f"{key}.json"
|
||||
existing = target if target.exists() else self.inflight / target.name
|
||||
if existing.exists():
|
||||
if existing.read_bytes() != payload:
|
||||
os.replace(
|
||||
temporary,
|
||||
self.quarantine / f"{key}.conflict.{uuid.uuid4().hex}.json",
|
||||
)
|
||||
self._adjust_spool_size(1, len(payload))
|
||||
_fsync_dir(self.quarantine)
|
||||
self.mark_degraded("local_event_payload_conflict")
|
||||
else:
|
||||
temporary.unlink(missing_ok=True)
|
||||
else:
|
||||
os.replace(temporary, target)
|
||||
self._adjust_spool_size(1, len(payload))
|
||||
_fsync_dir(self.pending)
|
||||
finally:
|
||||
os.close(lock_fd)
|
||||
lock_path.unlink(missing_ok=True)
|
||||
except Exception:
|
||||
self.mark_degraded("spool_write_failed")
|
||||
logger.exception("Usage capture could not persist a terminal event")
|
||||
finally:
|
||||
self._wake.set()
|
||||
|
||||
def _recover_stale_inflight(self) -> None:
|
||||
cutoff = time.time() - self.lease_seconds
|
||||
try:
|
||||
items = list(self.inflight.glob("*.json"))
|
||||
except OSError:
|
||||
return
|
||||
for item in items:
|
||||
try:
|
||||
if item.stat().st_mtime < cutoff:
|
||||
os.replace(item, self.pending / item.name)
|
||||
_fsync_dir(self.inflight)
|
||||
_fsync_dir(self.pending)
|
||||
except (FileNotFoundError, OSError):
|
||||
continue
|
||||
|
||||
def _clean_expired_quarantine(self) -> None:
|
||||
cutoff = time.time() - self.quarantine_retention_days * 86_400
|
||||
removed = 0
|
||||
for item in self.quarantine.glob("*.json"):
|
||||
try:
|
||||
if item.stat().st_mtime < cutoff:
|
||||
item.unlink()
|
||||
removed += 1
|
||||
except (FileNotFoundError, OSError):
|
||||
continue
|
||||
if removed:
|
||||
logger.warning("Removed %d expired usage quarantine events", removed)
|
||||
_fsync_dir(self.quarantine)
|
||||
|
||||
def _clean_stale_local_artifacts(self) -> None:
|
||||
cutoff = time.time() - self.lease_seconds
|
||||
candidates = [*self.tmp.glob("*.tmp"), *self.root.glob("*.lock")]
|
||||
for item in candidates:
|
||||
try:
|
||||
if item.stat().st_mtime < cutoff:
|
||||
item.unlink()
|
||||
except (FileNotFoundError, OSError):
|
||||
continue
|
||||
|
||||
def _counts(self) -> tuple[int, int, int, int]:
|
||||
def files(directory: Path) -> list[Path]:
|
||||
try:
|
||||
return list(directory.glob("*.json"))
|
||||
except OSError:
|
||||
return []
|
||||
|
||||
pending = files(self.pending)
|
||||
inflight = files(self.inflight)
|
||||
quarantine = files(self.quarantine)
|
||||
size = 0
|
||||
for item in pending + inflight + quarantine:
|
||||
try:
|
||||
size += item.stat().st_size
|
||||
except OSError:
|
||||
pass
|
||||
return len(pending), len(inflight), len(quarantine), size
|
||||
|
||||
def _heartbeat(self, client: httpx.Client) -> None:
|
||||
pending, inflight, quarantine, size = self._counts()
|
||||
with self._state_lock:
|
||||
body = {
|
||||
"deployment_id": self.deployment_id,
|
||||
"workspace_id": self.workspace_id,
|
||||
"emitter_version": "2.0",
|
||||
"schema_version": 1,
|
||||
"sender_status": (
|
||||
"degraded"
|
||||
if self.first_loss_at or self.degraded_reason
|
||||
else "healthy"
|
||||
),
|
||||
"spool_pending": pending,
|
||||
"spool_inflight": inflight,
|
||||
"spool_quarantined": quarantine,
|
||||
"spool_bytes": size,
|
||||
"first_loss_at": self.first_loss_at,
|
||||
"tracking_degraded_reason": self.degraded_reason,
|
||||
"last_error_code": self.last_error_code,
|
||||
"sent_at": _now(),
|
||||
}
|
||||
try:
|
||||
response = client.post(
|
||||
_endpoint(self.sink_url, "/api/usage/sources/heartbeat"), json=body
|
||||
)
|
||||
if response.status_code >= 400:
|
||||
self.last_error_code = f"heartbeat_http_{response.status_code}"
|
||||
except httpx.HTTPError:
|
||||
self.last_error_code = "heartbeat_unreachable"
|
||||
|
||||
def _capable(self, client: httpx.Client) -> bool:
|
||||
try:
|
||||
response = client.get(_endpoint(self.sink_url, "/api/usage/capabilities"))
|
||||
if response.status_code != 200:
|
||||
self.last_error_code = f"capabilities_http_{response.status_code}"
|
||||
return False
|
||||
versions = response.json().get("supported_schema_versions", [])
|
||||
if 1 not in versions:
|
||||
self.last_error_code = "schema_incompatible"
|
||||
return False
|
||||
self.last_error_code = None
|
||||
return True
|
||||
except (httpx.HTTPError, ValueError, TypeError):
|
||||
self.last_error_code = "collector_unreachable"
|
||||
return False
|
||||
|
||||
def _send_one(self, client: httpx.Client) -> bool:
|
||||
pending: Path | None = None
|
||||
lock_fd: int | None = None
|
||||
lock_path: Path | None = None
|
||||
try:
|
||||
candidates = self.pending.glob("*.json")
|
||||
for candidate in candidates:
|
||||
candidate_lock = self.root / f"{candidate.stem}.lock"
|
||||
try:
|
||||
descriptor = os.open(
|
||||
candidate_lock,
|
||||
os.O_WRONLY | os.O_CREAT | os.O_EXCL,
|
||||
0o600,
|
||||
)
|
||||
except FileExistsError:
|
||||
continue
|
||||
pending = candidate
|
||||
lock_fd = descriptor
|
||||
lock_path = candidate_lock
|
||||
break
|
||||
except OSError:
|
||||
return False
|
||||
if pending is None:
|
||||
return False
|
||||
inflight = self.inflight / pending.name
|
||||
try:
|
||||
os.replace(pending, inflight)
|
||||
_fsync_dir(self.pending)
|
||||
_fsync_dir(self.inflight)
|
||||
except (FileNotFoundError, OSError):
|
||||
return True
|
||||
finally:
|
||||
if lock_fd is not None:
|
||||
os.close(lock_fd)
|
||||
if lock_path is not None:
|
||||
lock_path.unlink(missing_ok=True)
|
||||
try:
|
||||
response = client.post(
|
||||
self.sink_url,
|
||||
content=inflight.read_bytes(),
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
if response.status_code == 200:
|
||||
status = response.json().get("status")
|
||||
if status in {"accepted", "duplicate"}:
|
||||
event_size = inflight.stat().st_size
|
||||
inflight.unlink(missing_ok=True)
|
||||
self._adjust_spool_size(-1, -event_size)
|
||||
_fsync_dir(self.inflight)
|
||||
self.last_error_code = None
|
||||
return True
|
||||
if response.status_code in {400, 409, 413, 422}:
|
||||
os.replace(inflight, self.quarantine / inflight.name)
|
||||
_fsync_dir(self.inflight)
|
||||
_fsync_dir(self.quarantine)
|
||||
self.mark_degraded(f"collector_rejected_event_{response.status_code}")
|
||||
self.last_error_code = f"event_http_{response.status_code}"
|
||||
return True
|
||||
self.last_error_code = f"event_http_{response.status_code}"
|
||||
except (httpx.HTTPError, OSError, ValueError, TypeError):
|
||||
self.last_error_code = "event_send_failed"
|
||||
try:
|
||||
os.replace(inflight, self.pending / inflight.name)
|
||||
_fsync_dir(self.inflight)
|
||||
_fsync_dir(self.pending)
|
||||
except (FileNotFoundError, OSError):
|
||||
pass
|
||||
return False
|
||||
|
||||
def _run(self) -> None:
|
||||
headers = {"Authorization": f"Bearer {self.token}"}
|
||||
timeout = httpx.Timeout(self.request_timeout, connect=self.connect_timeout)
|
||||
retry = self.retry_initial
|
||||
last_heartbeat = 0.0
|
||||
collector_capable = False
|
||||
next_probe = 0.0
|
||||
with httpx.Client(headers=headers, timeout=timeout) as client:
|
||||
while not self._stop.is_set():
|
||||
self._recover_stale_inflight()
|
||||
now = time.monotonic()
|
||||
if now - self._last_cleanup_at >= 3_600:
|
||||
self._clean_expired_quarantine()
|
||||
self._clean_stale_local_artifacts()
|
||||
self._last_cleanup_at = now
|
||||
if now - last_heartbeat >= self.heartbeat_interval:
|
||||
self._heartbeat(client)
|
||||
last_heartbeat = now
|
||||
if not collector_capable and now >= next_probe:
|
||||
collector_capable = self._capable(client)
|
||||
if not collector_capable:
|
||||
if self.last_error_code == "capabilities_http_404":
|
||||
probe_delay = self.unsupported_reprobe
|
||||
elif self.last_error_code == "schema_incompatible":
|
||||
probe_delay = self.schema_reprobe
|
||||
else:
|
||||
probe_delay = min(self.retry_max, retry)
|
||||
next_probe = time.monotonic() + probe_delay
|
||||
if not collector_capable:
|
||||
wait_for = max(0.05, next_probe - time.monotonic())
|
||||
self._wake.wait(wait_for)
|
||||
self._wake.clear()
|
||||
retry = min(self.retry_max, retry * 2)
|
||||
continue
|
||||
progressed = self._send_one(client)
|
||||
if progressed:
|
||||
retry = self.retry_initial
|
||||
continue
|
||||
if self.last_error_code in {
|
||||
"event_http_401",
|
||||
"event_http_403",
|
||||
"event_http_404",
|
||||
"event_http_426",
|
||||
}:
|
||||
collector_capable = False
|
||||
next_probe = time.monotonic() + (
|
||||
self.unsupported_reprobe
|
||||
if self.last_error_code == "event_http_404"
|
||||
else self.schema_reprobe
|
||||
)
|
||||
self._wake.wait(
|
||||
self.heartbeat_interval
|
||||
if not any(self.pending.glob("*.json"))
|
||||
else retry * random.uniform(0.8, 1.2)
|
||||
)
|
||||
self._wake.clear()
|
||||
retry = min(self.retry_max, retry * 2)
|
||||
|
||||
def close(self) -> None:
|
||||
self._stop.set()
|
||||
self._wake.set()
|
||||
if self._thread.is_alive():
|
||||
self._thread.join(timeout=1)
|
||||
|
||||
|
||||
_singleton_lock = threading.Lock()
|
||||
_singleton: tuple[int, UsageSpool] | None = None
|
||||
|
||||
|
||||
def get_usage_spool() -> UsageSpool:
|
||||
global _singleton
|
||||
pid = os.getpid()
|
||||
with _singleton_lock:
|
||||
if _singleton is None or _singleton[0] != pid:
|
||||
_singleton = (pid, UsageSpool())
|
||||
return _singleton[1]
|
||||
|
||||
|
||||
def mark_tracking_degraded(reason: str) -> None:
|
||||
if os.getenv("EVOSCIENTIST_USAGE_TRACKING", "").strip().lower() not in {
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
"on",
|
||||
}:
|
||||
return
|
||||
try:
|
||||
get_usage_spool().mark_degraded(reason)
|
||||
except Exception:
|
||||
logger.exception("Could not mark usage tracking as degraded: %s", reason)
|
||||
|
||||
|
||||
def _close_singleton() -> None:
|
||||
if _singleton is not None:
|
||||
_singleton[1].close()
|
||||
|
||||
|
||||
atexit.register(_close_singleton)
|
||||
@@ -0,0 +1,417 @@
|
||||
"""Administrative migration from a shared workspace to conversation scopes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .scope_registry import (
|
||||
ScopeRegistry,
|
||||
deployment_id_for_workspace,
|
||||
get_scope_registry,
|
||||
)
|
||||
from .workspace_scope import provision_conversation_scope, workspace_metadata
|
||||
|
||||
_TERMINAL_RUN_STATUSES = frozenset(
|
||||
{"success", "error", "timeout", "cancelled", "interrupted"}
|
||||
)
|
||||
_VALID_SCOPE_STATES = frozenset({"draft", "active"})
|
||||
|
||||
|
||||
def _now() -> str:
|
||||
return datetime.now(UTC).isoformat()
|
||||
|
||||
|
||||
def _write_report(workspace_root: Path, report: dict[str, Any]) -> Path:
|
||||
reports = workspace_root / ".evoscientist" / "control" / "cutover-reports"
|
||||
reports.mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
body = dict(report)
|
||||
digest = hashlib.sha256(
|
||||
json.dumps(body, sort_keys=True, separators=(",", ":")).encode("utf-8")
|
||||
).hexdigest()
|
||||
report["sha256"] = digest
|
||||
encoded = json.dumps(report, sort_keys=True, indent=2).encode("utf-8") + b"\n"
|
||||
target = reports / f"{report['operation_id']}.json"
|
||||
for destination in (target, reports / "latest.json"):
|
||||
temporary = destination.with_suffix(".json.tmp")
|
||||
fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
|
||||
with os.fdopen(fd, "wb") as file:
|
||||
file.write(encoded)
|
||||
os.replace(temporary, destination)
|
||||
return target
|
||||
|
||||
|
||||
def _active_run_ids(client: Any, thread_id: str) -> list[str]:
|
||||
return [
|
||||
str(run.get("run_id"))
|
||||
for run in client.runs.list(thread_id=thread_id, limit=1000)
|
||||
if str(run.get("status")) not in _TERMINAL_RUN_STATUSES
|
||||
]
|
||||
|
||||
|
||||
def _scope_owner_error(
|
||||
registry: ScopeRegistry,
|
||||
*,
|
||||
deployment_id: str,
|
||||
resource_id: str,
|
||||
metadata: dict[str, Any],
|
||||
allowed_owner_types: frozenset[str] | None = None,
|
||||
) -> str | None:
|
||||
"""Return a reason unless metadata names a live Registry-owned resource."""
|
||||
|
||||
scope_id = metadata.get("workspace_scope_id")
|
||||
if not isinstance(scope_id, str) or not scope_id:
|
||||
return "missing workspace_scope_id"
|
||||
try:
|
||||
scope = registry.get(deployment_id, scope_id)
|
||||
owner = registry.get_owner_by_resource(deployment_id, resource_id)
|
||||
except Exception as exc:
|
||||
return str(exc)
|
||||
if scope.state not in _VALID_SCOPE_STATES:
|
||||
return f"workspace scope is {scope.state}"
|
||||
if owner.scope_id != scope.scope_id:
|
||||
return "resource owner belongs to another workspace scope"
|
||||
if owner.state != "active":
|
||||
return f"resource owner is {owner.state}"
|
||||
if owner.owner_type in {"primary_thread", "primary_run"}:
|
||||
return f"unexpected owner type {owner.owner_type}"
|
||||
if allowed_owner_types is not None and owner.owner_type not in allowed_owner_types:
|
||||
return f"unexpected owner type {owner.owner_type}"
|
||||
metadata_deployment = metadata.get("workspace_deployment_id")
|
||||
if (
|
||||
metadata_deployment is not None
|
||||
and metadata_deployment != ""
|
||||
and metadata_deployment != deployment_id
|
||||
):
|
||||
return "metadata deployment does not match the active deployment"
|
||||
metadata_owner = metadata.get("workspace_scope_owner_id")
|
||||
if (
|
||||
metadata_owner is not None
|
||||
and metadata_owner != ""
|
||||
and metadata_owner != owner.owner_id
|
||||
):
|
||||
return "metadata owner does not match the resource owner"
|
||||
return None
|
||||
|
||||
|
||||
def _quarantine_derived_thread(
|
||||
client: Any,
|
||||
*,
|
||||
thread_id: str,
|
||||
metadata: dict[str, Any],
|
||||
operation_id: str,
|
||||
reason: str,
|
||||
report: dict[str, Any],
|
||||
) -> None:
|
||||
"""Interrupt a non-primary thread and remove any untrusted scope claims."""
|
||||
|
||||
try:
|
||||
active_runs = _active_run_ids(client, thread_id)
|
||||
for run_id in active_runs:
|
||||
client.runs.cancel(thread_id, run_id, wait=True, action="interrupt")
|
||||
remaining_active = _active_run_ids(client, thread_id)
|
||||
if remaining_active:
|
||||
report["active_runs"].append(
|
||||
{"thread_id": thread_id, "run_ids": remaining_active}
|
||||
)
|
||||
report["unmanaged_derived_threads"].append(thread_id)
|
||||
return
|
||||
cleaned_metadata = {
|
||||
key: value
|
||||
for key, value in metadata.items()
|
||||
if not key.startswith("workspace_")
|
||||
}
|
||||
client.threads.update(
|
||||
thread_id,
|
||||
metadata={
|
||||
**cleaned_metadata,
|
||||
"workspace_quarantine": {
|
||||
"operation_id": operation_id,
|
||||
"reason": reason,
|
||||
"quarantined_at": _now(),
|
||||
},
|
||||
},
|
||||
)
|
||||
report["quarantined_derived_threads"].append(
|
||||
{
|
||||
"thread_id": thread_id,
|
||||
"cancelled_run_ids": active_runs,
|
||||
"reason": reason,
|
||||
}
|
||||
)
|
||||
except Exception as exc:
|
||||
report["unmanaged_derived_threads"].append(thread_id)
|
||||
report["quarantine_failures"].append(
|
||||
{"thread_id": thread_id, "error": str(exc)}
|
||||
)
|
||||
|
||||
|
||||
def verify_required_cutover(workspace_root: Path) -> None:
|
||||
"""Reject strict mode unless the current deployment has a passing report."""
|
||||
|
||||
latest = (
|
||||
workspace_root / ".evoscientist" / "control" / "cutover-reports" / "latest.json"
|
||||
)
|
||||
try:
|
||||
report = json.loads(latest.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
raise RuntimeError(
|
||||
"required workspace isolation needs a passing workspace-cutover report"
|
||||
) from exc
|
||||
digest = report.pop("sha256", None)
|
||||
expected = hashlib.sha256(
|
||||
json.dumps(report, sort_keys=True, separators=(",", ":")).encode("utf-8")
|
||||
).hexdigest()
|
||||
if (
|
||||
digest != expected
|
||||
or report.get("status") != "passed"
|
||||
or report.get("browser_sdk_gate") is not True
|
||||
or report.get("deployment_id") != deployment_id_for_workspace(workspace_root)
|
||||
):
|
||||
raise RuntimeError("workspace-cutover report is missing, stale, or failed")
|
||||
operation = get_scope_registry(workspace_root).get_operation(
|
||||
str(report["deployment_id"]), str(report["operation_id"])
|
||||
)
|
||||
if (
|
||||
operation.kind != "workspace-cutover"
|
||||
or operation.state != "completed"
|
||||
or operation.result_sha256 != digest
|
||||
):
|
||||
raise RuntimeError("workspace-cutover registry operation does not match report")
|
||||
|
||||
|
||||
def run_workspace_cutover(
|
||||
*,
|
||||
workspace_root: Path,
|
||||
client: Any,
|
||||
assistant_id: str = "EvoScientist",
|
||||
browser_sdk_gate: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""Provision primary threads and quarantine legacy unscoped work."""
|
||||
|
||||
workspace_root = workspace_root.expanduser().resolve()
|
||||
deployment_id = deployment_id_for_workspace(workspace_root)
|
||||
registry: ScopeRegistry = get_scope_registry(workspace_root)
|
||||
operation_id = str(uuid.uuid4())
|
||||
registry.acquire_lock(
|
||||
deployment_id, "workspace-lifecycle", operation_id, lease_seconds=300
|
||||
)
|
||||
try:
|
||||
registry.acquire_lock(
|
||||
deployment_id, "workspace-cutover", operation_id, lease_seconds=300
|
||||
)
|
||||
except Exception:
|
||||
registry.release_lock(deployment_id, "workspace-lifecycle", operation_id)
|
||||
raise
|
||||
registry.begin_operation(deployment_id, operation_id, kind="workspace-cutover")
|
||||
report: dict[str, Any] = {
|
||||
"operation_id": operation_id,
|
||||
"deployment_id": deployment_id,
|
||||
"started_at": _now(),
|
||||
"primary_threads": 0,
|
||||
"scopes": 0,
|
||||
"quarantined_crons": 0,
|
||||
"quarantined_cron_records": [],
|
||||
"quarantined_derived_threads": [],
|
||||
"validated_derived_threads": [],
|
||||
"validated_crons": [],
|
||||
"invalid_scoped_derived_threads": [],
|
||||
"invalid_scoped_crons": [],
|
||||
"active_runs": [],
|
||||
"unmanaged_derived_threads": [],
|
||||
"quarantine_failures": [],
|
||||
"metadata_mismatches": [],
|
||||
"browser_sdk_gate": browser_sdk_gate,
|
||||
"errors": [],
|
||||
"status": "failed",
|
||||
}
|
||||
try:
|
||||
threads: list[dict[str, Any]] = []
|
||||
offset = 0
|
||||
while True:
|
||||
registry.renew_lock(
|
||||
deployment_id, "workspace-cutover", operation_id, lease_seconds=300
|
||||
)
|
||||
registry.renew_lock(
|
||||
deployment_id, "workspace-lifecycle", operation_id, lease_seconds=300
|
||||
)
|
||||
page = list(client.threads.search(limit=100, offset=offset))
|
||||
threads.extend(page)
|
||||
if len(page) < 100:
|
||||
break
|
||||
offset += len(page)
|
||||
for thread in threads:
|
||||
registry.renew_lock(
|
||||
deployment_id, "workspace-cutover", operation_id, lease_seconds=300
|
||||
)
|
||||
registry.renew_lock(
|
||||
deployment_id, "workspace-lifecycle", operation_id, lease_seconds=300
|
||||
)
|
||||
metadata = dict(thread.get("metadata") or {})
|
||||
thread_id = str(thread.get("thread_id") or "")
|
||||
is_primary = (
|
||||
metadata.get("graph_id") == assistant_id
|
||||
or metadata.get("agent_name") == assistant_id
|
||||
)
|
||||
if is_primary:
|
||||
if not thread_id:
|
||||
continue
|
||||
active_runs = _active_run_ids(client, thread_id)
|
||||
if active_runs:
|
||||
report["active_runs"].append(
|
||||
{"thread_id": thread_id, "run_ids": active_runs}
|
||||
)
|
||||
continue
|
||||
record = provision_conversation_scope(
|
||||
thread_id,
|
||||
deployment_id=deployment_id,
|
||||
workspace_root=workspace_root,
|
||||
lock_operation_id=operation_id,
|
||||
)
|
||||
updated = client.threads.update(
|
||||
thread_id, metadata={**metadata, **workspace_metadata(record)}
|
||||
)
|
||||
if isinstance(updated, dict):
|
||||
updated_metadata = dict(updated.get("metadata") or {})
|
||||
if updated_metadata.get("workspace_scope_id") != record.scope_id:
|
||||
report["metadata_mismatches"].append(thread_id)
|
||||
report["primary_threads"] += 1
|
||||
report["scopes"] += 1
|
||||
elif thread_id:
|
||||
scope_id = metadata.get("workspace_scope_id")
|
||||
if scope_id is None:
|
||||
_quarantine_derived_thread(
|
||||
client,
|
||||
thread_id=thread_id,
|
||||
metadata=metadata,
|
||||
operation_id=operation_id,
|
||||
reason="unscoped-derived-thread",
|
||||
report=report,
|
||||
)
|
||||
continue
|
||||
ownership_error = _scope_owner_error(
|
||||
registry,
|
||||
deployment_id=deployment_id,
|
||||
resource_id=thread_id,
|
||||
metadata=metadata,
|
||||
)
|
||||
if ownership_error is not None:
|
||||
report["invalid_scoped_derived_threads"].append(
|
||||
{"thread_id": thread_id, "error": ownership_error}
|
||||
)
|
||||
_quarantine_derived_thread(
|
||||
client,
|
||||
thread_id=thread_id,
|
||||
metadata=metadata,
|
||||
operation_id=operation_id,
|
||||
reason="invalid-scoped-derived-thread",
|
||||
report=report,
|
||||
)
|
||||
continue
|
||||
report["validated_derived_threads"].append(thread_id)
|
||||
active_runs = _active_run_ids(client, thread_id)
|
||||
if active_runs:
|
||||
report["active_runs"].append(
|
||||
{"thread_id": thread_id, "run_ids": active_runs}
|
||||
)
|
||||
for cron in client.crons.search(limit=1000):
|
||||
registry.renew_lock(
|
||||
deployment_id, "workspace-cutover", operation_id, lease_seconds=300
|
||||
)
|
||||
registry.renew_lock(
|
||||
deployment_id, "workspace-lifecycle", operation_id, lease_seconds=300
|
||||
)
|
||||
metadata = dict(cron.get("metadata") or {})
|
||||
cron_id = str(cron.get("cron_id") or "")
|
||||
scope_id = metadata.get("workspace_scope_id")
|
||||
if scope_id is not None:
|
||||
ownership_error = _scope_owner_error(
|
||||
registry,
|
||||
deployment_id=deployment_id,
|
||||
resource_id=cron_id,
|
||||
metadata=metadata,
|
||||
allowed_owner_types=frozenset({"schedule"}),
|
||||
)
|
||||
if ownership_error is None:
|
||||
report["validated_crons"].append(cron_id)
|
||||
continue
|
||||
client.crons.update(cron_id, enabled=False)
|
||||
report["quarantined_crons"] += 1
|
||||
report["invalid_scoped_crons"].append(
|
||||
{"cron_id": cron_id, "error": ownership_error}
|
||||
)
|
||||
report["quarantined_cron_records"].append(
|
||||
{"cron_id": cron_id, "reason": "invalid-scoped-cron"}
|
||||
)
|
||||
continue
|
||||
if metadata.get("run_kind") != "scheduled_task":
|
||||
continue
|
||||
client.crons.update(cron_id, enabled=False)
|
||||
report["quarantined_crons"] += 1
|
||||
report["quarantined_cron_records"].append(
|
||||
{"cron_id": cron_id, "reason": "unscoped-scheduled-task"}
|
||||
)
|
||||
report["status"] = (
|
||||
"passed"
|
||||
if (
|
||||
report["browser_sdk_gate"]
|
||||
and not report["unmanaged_derived_threads"]
|
||||
and not report["quarantine_failures"]
|
||||
and not report["active_runs"]
|
||||
and not report["metadata_mismatches"]
|
||||
)
|
||||
else "failed"
|
||||
)
|
||||
report["completed_at"] = _now()
|
||||
report["report_path"] = str(
|
||||
workspace_root
|
||||
/ ".evoscientist"
|
||||
/ "control"
|
||||
/ "cutover-reports"
|
||||
/ f"{operation_id}.json"
|
||||
)
|
||||
_write_report(workspace_root, report)
|
||||
registry.finish_operation(
|
||||
deployment_id,
|
||||
operation_id,
|
||||
state="completed" if report["status"] == "passed" else "failed",
|
||||
result_sha256=str(report["sha256"]),
|
||||
last_error_code=None
|
||||
if report["status"] == "passed"
|
||||
else "cutover-gates-failed",
|
||||
)
|
||||
return report
|
||||
except Exception as exc:
|
||||
report["errors"].append(str(exc))
|
||||
report["completed_at"] = _now()
|
||||
report["report_path"] = str(
|
||||
workspace_root
|
||||
/ ".evoscientist"
|
||||
/ "control"
|
||||
/ "cutover-reports"
|
||||
/ f"{operation_id}.json"
|
||||
)
|
||||
_write_report(workspace_root, report)
|
||||
registry.finish_operation(
|
||||
deployment_id,
|
||||
operation_id,
|
||||
state="failed",
|
||||
result_sha256=str(report["sha256"]),
|
||||
last_error_code="cutover-exception",
|
||||
)
|
||||
return report
|
||||
finally:
|
||||
try:
|
||||
registry.release_lock(deployment_id, "workspace-cutover", operation_id)
|
||||
except Exception:
|
||||
# A lost/expired lease is already recorded as a failed report.
|
||||
pass
|
||||
try:
|
||||
registry.release_lock(deployment_id, "workspace-lifecycle", operation_id)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -0,0 +1,168 @@
|
||||
"""Idempotent lifecycle maintenance for conversation workspaces."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .scope_registry import (
|
||||
ScopeRecord,
|
||||
deployment_id_for_workspace,
|
||||
get_scope_registry,
|
||||
)
|
||||
|
||||
_TERMINAL_RUN_STATES = frozenset(
|
||||
{"success", "error", "timeout", "cancelled", "interrupted"}
|
||||
)
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
def _parse_timestamp(value: str) -> datetime:
|
||||
parsed = datetime.fromisoformat(value)
|
||||
return parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=UTC)
|
||||
|
||||
|
||||
def _trash_root(workspace_root: Path) -> Path:
|
||||
return workspace_root / ".evoscientist" / "trash"
|
||||
|
||||
|
||||
def _conversation_root(workspace_root: Path, scope_id: str) -> Path:
|
||||
return workspace_root / ".evoscientist" / "conversations" / scope_id
|
||||
|
||||
|
||||
def _move_to_trash(workspace_root: Path, scope_id: str, now: datetime) -> None:
|
||||
source = _conversation_root(workspace_root, scope_id)
|
||||
target_root = _trash_root(workspace_root)
|
||||
target_root.mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
target = target_root / f"{scope_id}-{int(now.timestamp() * 1000)}"
|
||||
try:
|
||||
os.replace(source, target)
|
||||
os.utime(target, (now.timestamp(), now.timestamp()))
|
||||
except FileNotFoundError:
|
||||
return
|
||||
|
||||
|
||||
def _purge_trash(workspace_root: Path, cutoff: datetime) -> tuple[int, list[str]]:
|
||||
root = _trash_root(workspace_root)
|
||||
if not root.exists():
|
||||
return 0, []
|
||||
removed = 0
|
||||
errors: list[str] = []
|
||||
for candidate in root.iterdir():
|
||||
try:
|
||||
if candidate.is_symlink() or not candidate.is_dir():
|
||||
continue
|
||||
modified = datetime.fromtimestamp(candidate.stat().st_mtime, tz=UTC)
|
||||
if modified >= cutoff:
|
||||
continue
|
||||
shutil.rmtree(candidate)
|
||||
removed += 1
|
||||
except OSError as exc:
|
||||
errors.append(f"trash:{candidate.name}:{exc}")
|
||||
return removed, errors
|
||||
|
||||
|
||||
def _draft_can_be_deleted(scope: ScopeRecord, registry: Any, client: Any) -> bool:
|
||||
owners = registry.owners(scope.deployment_id, scope.scope_id)
|
||||
if any(owner.owner_type != "primary_thread" for owner in owners):
|
||||
return False
|
||||
try:
|
||||
runs = list(client.runs.list(thread_id=scope.primary_thread_id, limit=1000))
|
||||
if any(str(run.get("status")) not in _TERMINAL_RUN_STATES for run in runs):
|
||||
return False
|
||||
state = client.threads.get_state(scope.primary_thread_id)
|
||||
values = state.get("values") if isinstance(state, dict) else None
|
||||
messages = values.get("messages") if isinstance(values, dict) else None
|
||||
return not messages
|
||||
except Exception:
|
||||
# Maintenance must not delete a draft whose content cannot be proven empty.
|
||||
return False
|
||||
|
||||
|
||||
def run_workspace_maintenance(
|
||||
*,
|
||||
workspace_root: Path,
|
||||
client: Any,
|
||||
draft_ttl_hours: int = 24,
|
||||
trash_retention_days: int = 7,
|
||||
now: datetime | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Delete safely-drained stale drafts and purge aged trash directories."""
|
||||
|
||||
if draft_ttl_hours <= 0 or trash_retention_days <= 0:
|
||||
raise ValueError("workspace maintenance retention values must be positive")
|
||||
workspace_root = workspace_root.expanduser().resolve()
|
||||
current = now or _now()
|
||||
deployment_id = deployment_id_for_workspace(workspace_root)
|
||||
registry = get_scope_registry(workspace_root)
|
||||
if registry.active_lock(deployment_id, "workspace-cutover") is not None:
|
||||
raise RuntimeError("workspace maintenance is blocked by workspace cutover")
|
||||
operation_id = str(uuid.uuid4())
|
||||
registry.acquire_lock(
|
||||
deployment_id, "workspace-lifecycle", operation_id, lease_seconds=300
|
||||
)
|
||||
registry.begin_operation(deployment_id, operation_id, kind="workspace-maintenance")
|
||||
report: dict[str, Any] = {
|
||||
"operation_id": operation_id,
|
||||
"deployment_id": deployment_id,
|
||||
"started_at": current.isoformat(),
|
||||
"deleted_drafts": [],
|
||||
"skipped_drafts": [],
|
||||
"purged_trash": 0,
|
||||
"errors": [],
|
||||
}
|
||||
try:
|
||||
draft_cutoff = current - timedelta(hours=draft_ttl_hours)
|
||||
for scope in registry.list_scopes(deployment_id, state="draft"):
|
||||
registry.renew_lock(
|
||||
deployment_id, "workspace-lifecycle", operation_id, lease_seconds=300
|
||||
)
|
||||
if _parse_timestamp(scope.created_at) >= draft_cutoff:
|
||||
continue
|
||||
if not _draft_can_be_deleted(scope, registry, client):
|
||||
report["skipped_drafts"].append(scope.primary_thread_id)
|
||||
continue
|
||||
deleting = registry.transition_scope(
|
||||
deployment_id,
|
||||
scope.scope_id,
|
||||
expected_revision=scope.revision,
|
||||
state="deleting",
|
||||
)
|
||||
_move_to_trash(workspace_root, deleting.scope_id, current)
|
||||
try:
|
||||
client.threads.delete(deleting.primary_thread_id)
|
||||
except Exception as exc:
|
||||
status = getattr(exc, "status", None)
|
||||
if status != 404 and "not found" not in str(exc).lower():
|
||||
raise
|
||||
registry.transition_scope(
|
||||
deployment_id,
|
||||
deleting.scope_id,
|
||||
expected_revision=deleting.revision,
|
||||
state="deleted",
|
||||
)
|
||||
report["deleted_drafts"].append(scope.primary_thread_id)
|
||||
purged, errors = _purge_trash(
|
||||
workspace_root, current - timedelta(days=trash_retention_days)
|
||||
)
|
||||
report["purged_trash"] = purged
|
||||
report["errors"].extend(errors)
|
||||
except Exception as exc:
|
||||
report["errors"].append(str(exc))
|
||||
finally:
|
||||
report["completed_at"] = _now().isoformat()
|
||||
registry.finish_operation(
|
||||
deployment_id,
|
||||
operation_id,
|
||||
state="completed" if not report["errors"] else "failed",
|
||||
last_error_code=None if not report["errors"] else "maintenance-failed",
|
||||
)
|
||||
registry.release_lock(deployment_id, "workspace-lifecycle", operation_id)
|
||||
return report
|
||||
@@ -0,0 +1,541 @@
|
||||
"""Conversation-scoped workspace resolution and DeepAgents backend factory."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from deepagents.backends.protocol import (
|
||||
EditResult,
|
||||
ExecuteResponse,
|
||||
FileDownloadResponse,
|
||||
FileUploadResponse,
|
||||
GlobResult,
|
||||
GrepResult,
|
||||
LsResult,
|
||||
ReadResult,
|
||||
SandboxBackendProtocol,
|
||||
WriteResult,
|
||||
)
|
||||
from langchain.tools import ToolRuntime
|
||||
|
||||
from . import paths
|
||||
from .scope_registry import (
|
||||
ScopeAccessError,
|
||||
ScopeRecord,
|
||||
deployment_id_for_workspace,
|
||||
get_scope_registry,
|
||||
)
|
||||
|
||||
IsolationMode = str
|
||||
|
||||
|
||||
def workspace_isolation_mode() -> IsolationMode:
|
||||
value = os.getenv("EVOSCIENTIST_WORKSPACE_ISOLATION", "optional").strip().lower()
|
||||
if value not in {"legacy", "optional", "required"}:
|
||||
raise RuntimeError("EVOSCIENTIST_WORKSPACE_ISOLATION must be legacy, optional or required")
|
||||
return value
|
||||
|
||||
|
||||
def is_required() -> bool:
|
||||
return workspace_isolation_mode() == "required"
|
||||
|
||||
|
||||
def verify_required_executor() -> None:
|
||||
"""Fail startup unless the pinned scope executor is locally usable."""
|
||||
|
||||
docker = shutil.which("docker")
|
||||
if not docker:
|
||||
raise RuntimeError("required workspace isolation needs the docker OCI runtime")
|
||||
image = os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_IMAGE", "").strip()
|
||||
if "@sha256:" not in image:
|
||||
raise RuntimeError("required workspace isolation needs an OCI image pinned by digest")
|
||||
try:
|
||||
probe = subprocess.run(
|
||||
[docker, "image", "inspect", image],
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
except (OSError, subprocess.TimeoutExpired) as exc:
|
||||
raise RuntimeError("required workspace isolation cannot verify the OCI executor") from exc
|
||||
if probe.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"required workspace isolation needs local OCI image {image!r}"
|
||||
)
|
||||
|
||||
|
||||
def current_deployment_id() -> str:
|
||||
return deployment_id_for_workspace(paths.WORKSPACE_ROOT)
|
||||
|
||||
|
||||
def conversation_root(scope_id: str, workspace_root: Path | None = None) -> Path:
|
||||
scope = str(uuid.UUID(scope_id))
|
||||
# The deploy process supplies an absolute workspace root. This helper is
|
||||
# called from synchronous DeepAgents backend factories on the ASGI loop.
|
||||
root = (workspace_root or paths.WORKSPACE_ROOT).expanduser()
|
||||
return root / ".evoscientist" / "conversations" / scope
|
||||
|
||||
|
||||
def conversation_files_dir(scope_id: str, workspace_root: Path | None = None) -> Path:
|
||||
return conversation_root(scope_id, workspace_root) / "files"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScopeContext:
|
||||
deployment_id: str
|
||||
scope_id: str
|
||||
owner_id: str
|
||||
thread_id: str
|
||||
revision: int
|
||||
files_dir: Path
|
||||
runtime_dir: Path
|
||||
primary_thread_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RuntimeScopeConfig:
|
||||
"""Untrusted runtime identifiers parsed without filesystem or Registry I/O."""
|
||||
|
||||
scope_id: str
|
||||
owner_id: str
|
||||
thread_id: str
|
||||
deployment_id: str | None
|
||||
|
||||
|
||||
class ScopedContainerBackend:
|
||||
"""Filesystem backend whose shell commands execute in a scope-only OCI container."""
|
||||
|
||||
def __init__(self, root_dir: Path, *, timeout: int) -> None:
|
||||
from .backends import CustomSandboxBackend
|
||||
|
||||
# Reuse the hardened filesystem operations; only ``execute`` is
|
||||
# replaced so no agent shell runs in the host process.
|
||||
self._filesystem = CustomSandboxBackend(
|
||||
root_dir=str(root_dir), virtual_mode=True, timeout=timeout, dangerous=False
|
||||
)
|
||||
self._root_dir = root_dir
|
||||
self._timeout = timeout
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._filesystem, name)
|
||||
|
||||
def execute(self, command: str, *, timeout: int | None = None) -> Any:
|
||||
from .backends import ExecuteResponse, prepare_sandbox_command
|
||||
|
||||
command, error = prepare_sandbox_command(
|
||||
command, self._filesystem.cwd, virtual_mode=True, dangerous=False
|
||||
)
|
||||
if error:
|
||||
return ExecuteResponse(output=error, exit_code=1, truncated=False)
|
||||
image = os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_IMAGE", "").strip()
|
||||
if "@sha256:" not in image:
|
||||
return ExecuteResponse(
|
||||
output="Required workspace isolation needs an OCI image pinned by digest.",
|
||||
exit_code=125,
|
||||
truncated=False,
|
||||
)
|
||||
effective_timeout = max(1, min(timeout or self._timeout, 3600))
|
||||
invocation = [
|
||||
"docker",
|
||||
"run",
|
||||
"--rm",
|
||||
"--network",
|
||||
"none",
|
||||
"--read-only",
|
||||
"--tmpfs",
|
||||
"/tmp:rw,noexec,nosuid,size=64m",
|
||||
"--cap-drop",
|
||||
"ALL",
|
||||
"--security-opt",
|
||||
"no-new-privileges",
|
||||
"--pids-limit",
|
||||
os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_PIDS", "128"),
|
||||
"--memory",
|
||||
os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_MEMORY", "1g"),
|
||||
"--cpus",
|
||||
os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_CPUS", "1"),
|
||||
"--mount",
|
||||
f"type=bind,src={self._root_dir},dst=/workspace",
|
||||
"--workdir",
|
||||
"/workspace",
|
||||
image,
|
||||
"sh",
|
||||
"-lc",
|
||||
command,
|
||||
]
|
||||
try:
|
||||
completed = subprocess.run(
|
||||
invocation,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=effective_timeout,
|
||||
)
|
||||
except FileNotFoundError:
|
||||
return ExecuteResponse(
|
||||
output="Required workspace isolation needs an OCI runtime (docker was not found).",
|
||||
exit_code=127,
|
||||
truncated=False,
|
||||
)
|
||||
except subprocess.TimeoutExpired as exc:
|
||||
output = (exc.stdout or "") + (exc.stderr or "")
|
||||
return ExecuteResponse(output=output, exit_code=124, truncated=False)
|
||||
output = completed.stdout + completed.stderr
|
||||
return ExecuteResponse(output=output, exit_code=completed.returncode, truncated=False)
|
||||
|
||||
|
||||
def _configurable(runtime: ToolRuntime[Any, Any] | Any | None) -> dict[str, Any]:
|
||||
"""Return the active runnable config, with a non-graph fallback.
|
||||
|
||||
``ToolRuntime`` deliberately does not expose ``RunnableConfig`` during a
|
||||
graph execution. LangGraph keeps it in a context variable instead. The
|
||||
fallback preserves direct callers and unit tests that supply a lightweight
|
||||
runtime object outside a runnable context.
|
||||
"""
|
||||
|
||||
config: Any = None
|
||||
try:
|
||||
from langgraph.config import get_config
|
||||
|
||||
config = get_config()
|
||||
except (ImportError, LookupError, RuntimeError):
|
||||
pass
|
||||
if not isinstance(config, dict) and runtime is not None:
|
||||
config = getattr(runtime, "config", None) or {}
|
||||
if not isinstance(config, dict):
|
||||
return {}
|
||||
configurable = config.get("configurable") or {}
|
||||
return dict(configurable) if isinstance(configurable, dict) else {}
|
||||
|
||||
|
||||
def _required_string(configurable: dict[str, Any], key: str) -> str:
|
||||
value = configurable.get(key)
|
||||
if not isinstance(value, str) or not value:
|
||||
raise ScopeAccessError(f"missing {key}")
|
||||
return value
|
||||
|
||||
|
||||
def _runtime_scope_config(
|
||||
runtime: ToolRuntime[Any, Any] | Any | None,
|
||||
*,
|
||||
kind: str,
|
||||
) -> _RuntimeScopeConfig | None:
|
||||
"""Parse scope identifiers without treating config as an authorization grant."""
|
||||
|
||||
configurable = _configurable(runtime)
|
||||
scope_id = configurable.get("workspace_scope_id")
|
||||
owner_id = configurable.get("workspace_scope_owner_id")
|
||||
thread_id = configurable.get("thread_id")
|
||||
|
||||
if scope_id is None and owner_id is None:
|
||||
if workspace_isolation_mode() == "required":
|
||||
raise ScopeAccessError(f"{kind} requires a workspace scope")
|
||||
return None
|
||||
if (
|
||||
not isinstance(scope_id, str)
|
||||
or not isinstance(owner_id, str)
|
||||
or not isinstance(thread_id, str)
|
||||
):
|
||||
raise ScopeAccessError(f"{kind} has an incomplete workspace scope")
|
||||
try:
|
||||
canonical_scope_id = str(uuid.UUID(scope_id))
|
||||
canonical_owner_id = str(uuid.UUID(owner_id))
|
||||
except ValueError as exc:
|
||||
raise ScopeAccessError(f"{kind} has an invalid workspace scope") from exc
|
||||
deployment_id = configurable.get("workspace_deployment_id")
|
||||
if deployment_id is not None and (
|
||||
not isinstance(deployment_id, str) or not deployment_id
|
||||
):
|
||||
raise ScopeAccessError(f"{kind} has an invalid workspace deployment")
|
||||
return _RuntimeScopeConfig(
|
||||
scope_id=canonical_scope_id,
|
||||
owner_id=canonical_owner_id,
|
||||
thread_id=thread_id,
|
||||
deployment_id=deployment_id,
|
||||
)
|
||||
|
||||
|
||||
def _validated_scope_directories(scope_id: str) -> tuple[Path, Path]:
|
||||
"""Return canonical private directories after preventing symlink escape.
|
||||
|
||||
This function intentionally resolves paths and must run only from a
|
||||
filesystem-operation worker, never from the runtime backend factory.
|
||||
"""
|
||||
|
||||
conversations_dir = (
|
||||
paths.WORKSPACE_ROOT.expanduser() / ".evoscientist" / "conversations"
|
||||
).resolve(strict=True)
|
||||
scope_root = (conversations_dir / scope_id).resolve(strict=True)
|
||||
files_dir = (scope_root / "files").resolve(strict=True)
|
||||
runtime_dir = (scope_root / "runtime").resolve(strict=True)
|
||||
if (
|
||||
scope_root.parent != conversations_dir
|
||||
or files_dir.parent != scope_root
|
||||
or runtime_dir.parent != scope_root
|
||||
):
|
||||
raise ScopeAccessError("workspace directory escapes its scope")
|
||||
if not files_dir.is_dir() or not runtime_dir.is_dir():
|
||||
raise ScopeAccessError("workspace directory is missing")
|
||||
return files_dir, runtime_dir
|
||||
|
||||
|
||||
def _resolve_scope_context(config: _RuntimeScopeConfig | None) -> ScopeContext | None:
|
||||
"""Validate parsed scope identifiers against the active registry."""
|
||||
|
||||
if config is None:
|
||||
return None
|
||||
deployment_id = config.deployment_id or current_deployment_id()
|
||||
registry = get_scope_registry(paths.WORKSPACE_ROOT)
|
||||
if registry.active_lock(deployment_id, "workspace-cutover") is not None:
|
||||
raise ScopeAccessError("workspace cutover is in progress")
|
||||
record = registry.assert_runtime(
|
||||
deployment_id, config.scope_id, config.thread_id, config.owner_id
|
||||
)
|
||||
files_dir, runtime_dir = _validated_scope_directories(record.scope_id)
|
||||
return ScopeContext(
|
||||
deployment_id=deployment_id,
|
||||
scope_id=record.scope_id,
|
||||
owner_id=config.owner_id,
|
||||
thread_id=config.thread_id,
|
||||
revision=record.revision,
|
||||
files_dir=files_dir,
|
||||
runtime_dir=runtime_dir,
|
||||
primary_thread_id=record.primary_thread_id,
|
||||
)
|
||||
|
||||
|
||||
def require_scoped_runtime(
|
||||
runtime: ToolRuntime[Any, Any] | Any | None,
|
||||
*,
|
||||
kind: str = "tool",
|
||||
) -> ScopeContext | None:
|
||||
"""Resolve and validate a runtime scope.
|
||||
|
||||
``optional`` retains legacy CLI compatibility when no scope has been
|
||||
injected. ``required`` never falls back to ``WORKSPACE_ROOT``.
|
||||
"""
|
||||
|
||||
config = _runtime_scope_config(runtime, kind=kind)
|
||||
return _resolve_scope_context(config)
|
||||
|
||||
|
||||
def provision_conversation_scope(
|
||||
thread_id: str,
|
||||
*,
|
||||
deployment_id: str | None = None,
|
||||
scope_id: str | None = None,
|
||||
workspace_root: Path | None = None,
|
||||
lock_operation_id: str | None = None,
|
||||
) -> ScopeRecord:
|
||||
"""Create the registry mapping and private directory for a primary thread."""
|
||||
|
||||
root = (workspace_root or paths.WORKSPACE_ROOT).expanduser()
|
||||
deployment_id = deployment_id or deployment_id_for_workspace(root)
|
||||
registry = get_scope_registry(root)
|
||||
record = registry.provision(
|
||||
deployment_id,
|
||||
thread_id,
|
||||
scope_id=scope_id,
|
||||
lock_operation_id=lock_operation_id,
|
||||
)
|
||||
root = conversation_root(record.scope_id, root)
|
||||
try:
|
||||
(root / "files").mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
(root / "runtime").mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
for directory in (root, root / "files", root / "runtime"):
|
||||
try:
|
||||
directory.chmod(0o700)
|
||||
except OSError:
|
||||
pass
|
||||
except OSError:
|
||||
# Keep the durable reservation for the recovery job; it is safer than
|
||||
# silently falling back to the shared deployment root.
|
||||
raise
|
||||
return record
|
||||
|
||||
|
||||
def _build_backend(root_dir: Path, *, dangerous: bool) -> Any:
|
||||
from deepagents.backends import CompositeBackend
|
||||
|
||||
from .backends import (
|
||||
CustomSandboxBackend,
|
||||
MemoryFilesystemBackend,
|
||||
MergedSkillsBackend,
|
||||
)
|
||||
from .EvoScientist import SKILLS_DIR
|
||||
|
||||
cfg_timeout = int(os.getenv("EVOSCIENTIST_SANDBOX_EXECUTE_TIMEOUT", "300"))
|
||||
ws_backend: Any
|
||||
if is_required():
|
||||
ws_backend = ScopedContainerBackend(root_dir, timeout=cfg_timeout)
|
||||
else:
|
||||
ws_backend = CustomSandboxBackend(
|
||||
root_dir=str(root_dir),
|
||||
virtual_mode=True,
|
||||
timeout=cfg_timeout,
|
||||
dangerous=dangerous,
|
||||
)
|
||||
return CompositeBackend(
|
||||
default=ws_backend,
|
||||
routes={
|
||||
"/skills/": MergedSkillsBackend(
|
||||
primary_dir=str(paths.USER_SKILLS_DIR),
|
||||
global_dir=str(paths.GLOBAL_SKILLS_DIR),
|
||||
secondary_dir=SKILLS_DIR,
|
||||
),
|
||||
"/memories/": MemoryFilesystemBackend(
|
||||
root_dir=str(paths.MEMORIES_DIR), virtual_mode=True
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class DeferredScopedBackend(SandboxBackendProtocol):
|
||||
"""Resolve the scoped filesystem backend only from a worker thread.
|
||||
|
||||
DeepAgents invokes its deprecated backend factory from async middleware.
|
||||
Its concrete filesystem backends synchronously call ``Path.resolve()`` in
|
||||
their constructors, so doing that work in the factory makes every run fail
|
||||
under LangGraph's blocking-call detector. This proxy itself is I/O-free;
|
||||
the inherited async methods dispatch the synchronous operations to a
|
||||
thread, where Registry validation and concrete backend construction occur.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: _RuntimeScopeConfig,
|
||||
*,
|
||||
dangerous: bool,
|
||||
) -> None:
|
||||
self._config = config
|
||||
self._dangerous = dangerous
|
||||
self._backend: Any | None = None
|
||||
self._backend_key: tuple[str, str, str, int] | None = None
|
||||
self._lock = threading.RLock()
|
||||
|
||||
@property
|
||||
def id(self) -> str:
|
||||
# This is queried while composing the model request; do not initialize
|
||||
# the real backend or touch the Registry here.
|
||||
return f"scope-{self._config.scope_id[:8]}-{self._config.owner_id[:8]}"
|
||||
|
||||
def _delegate(self) -> Any:
|
||||
"""Validate the current scope and return a concrete backend.
|
||||
|
||||
Every operation enters here, so a deleted scope or stale owner cannot
|
||||
keep using a backend constructed before the lifecycle transition.
|
||||
"""
|
||||
|
||||
# Async backend methods run this code in a worker thread. LangGraph's
|
||||
# RunnableConfig context variable is not available there, so validate
|
||||
# the immutable scope parsed by the factory on the graph thread.
|
||||
context = _resolve_scope_context(self._config)
|
||||
if context is None:
|
||||
raise ScopeAccessError("scoped backend lost its workspace scope")
|
||||
if is_required() and self._dangerous:
|
||||
raise ScopeAccessError(
|
||||
"dangerous_mode is incompatible with required isolation"
|
||||
)
|
||||
key = (context.scope_id, context.owner_id, context.thread_id, context.revision)
|
||||
with self._lock:
|
||||
if self._backend is None or self._backend_key != key:
|
||||
self._backend = _build_backend(
|
||||
context.files_dir, dangerous=self._dangerous
|
||||
)
|
||||
self._backend_key = key
|
||||
return self._backend
|
||||
|
||||
def ls(self, path: str) -> LsResult:
|
||||
return self._delegate().ls(path)
|
||||
|
||||
def read(self, file_path: str, offset: int = 0, limit: int = 2000) -> ReadResult:
|
||||
return self._delegate().read(file_path, offset, limit)
|
||||
|
||||
def grep(
|
||||
self, pattern: str, path: str | None = None, glob: str | None = None
|
||||
) -> GrepResult:
|
||||
return self._delegate().grep(pattern, path, glob)
|
||||
|
||||
def glob(self, pattern: str, path: str | None = None) -> GlobResult:
|
||||
return self._delegate().glob(pattern, path)
|
||||
|
||||
def write(self, file_path: str, content: str) -> WriteResult:
|
||||
return self._delegate().write(file_path, content)
|
||||
|
||||
def edit(
|
||||
self,
|
||||
file_path: str,
|
||||
old_string: str,
|
||||
new_string: str,
|
||||
replace_all: bool = False,
|
||||
) -> EditResult:
|
||||
return self._delegate().edit(file_path, old_string, new_string, replace_all)
|
||||
|
||||
def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]:
|
||||
return self._delegate().upload_files(files)
|
||||
|
||||
def download_files(self, paths: list[str]) -> list[FileDownloadResponse]:
|
||||
return self._delegate().download_files(paths)
|
||||
|
||||
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
||||
return self._delegate().execute(command, timeout=timeout)
|
||||
|
||||
|
||||
def create_workspace_backend(
|
||||
runtime: ToolRuntime[Any, Any],
|
||||
*,
|
||||
legacy_backend: Callable[[], Any],
|
||||
dangerous: bool = False,
|
||||
allow_unscoped_legacy: bool = True,
|
||||
) -> Any:
|
||||
"""Return a backend handle without blocking the Agent event loop."""
|
||||
|
||||
config = _runtime_scope_config(runtime, kind="filesystem backend")
|
||||
if config is None:
|
||||
if not allow_unscoped_legacy:
|
||||
raise ScopeAccessError(
|
||||
"deployed graph runs require a workspace scope"
|
||||
)
|
||||
return legacy_backend()
|
||||
if is_required() and dangerous:
|
||||
raise ScopeAccessError("dangerous_mode is incompatible with required isolation")
|
||||
return DeferredScopedBackend(config, dangerous=dangerous)
|
||||
|
||||
|
||||
def create_workspace_backend_factory(
|
||||
legacy_backend: Callable[[], Any],
|
||||
*,
|
||||
dangerous: bool = False,
|
||||
allow_unscoped_legacy: bool = True,
|
||||
) -> Callable[[ToolRuntime[Any, Any]], Any]:
|
||||
def factory(runtime: ToolRuntime[Any, Any]) -> Any:
|
||||
return create_workspace_backend(
|
||||
runtime,
|
||||
legacy_backend=legacy_backend,
|
||||
dangerous=dangerous,
|
||||
allow_unscoped_legacy=allow_unscoped_legacy,
|
||||
)
|
||||
|
||||
return factory
|
||||
|
||||
|
||||
def workspace_metadata(record: ScopeRecord) -> dict[str, str | int]:
|
||||
"""Metadata mirrored onto the LangGraph primary thread by trusted callers."""
|
||||
|
||||
return {
|
||||
"workspace_schema_version": 1,
|
||||
"workspace_scope_id": record.scope_id,
|
||||
"workspace_status": record.state,
|
||||
"workspace_scope_owner_id": record.primary_owner_id,
|
||||
"workspace_scope_revision": record.revision,
|
||||
"workspace_deployment_id": record.deployment_id,
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user