Compare commits
56 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 11181c14d5 | |||
| 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
|
# EvoScientist — cp .env.example .env && fill in your keys
|
||||||
|
#
|
||||||
# LLM provider (pick at least one)
|
# LLM providers, models, and API keys are managed exclusively through the
|
||||||
ANTHROPIC_API_KEY= # console.anthropic.com
|
# Model Registry (WebUI 大模型配置 / Config API); no provider credential is
|
||||||
OPENAI_API_KEY= # platform.openai.com
|
# read from environment variables. See
|
||||||
GOOGLE_API_KEY= # aistudio.google.com/api-keys
|
# docs/unified-model-configuration-architecture.md.
|
||||||
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)
|
|
||||||
|
|
||||||
# Web search (optional)
|
# Web search (optional)
|
||||||
TAVILY_API_KEY= # app.tavily.com
|
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:
|
strategy:
|
||||||
fail-fast: false
|
fail-fast: false
|
||||||
matrix:
|
matrix:
|
||||||
os: [ubuntu-latest, windows-latest]
|
os: [ubuntu-latest, macos-latest, windows-latest]
|
||||||
python-version: ["3.11", "3.12"]
|
python-version: ["3.11", "3.12"]
|
||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v5
|
- 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
|
- uses: astral-sh/setup-uv@v6
|
||||||
with:
|
with:
|
||||||
python-version: ${{ matrix.python-version }}
|
python-version: ${{ matrix.python-version }}
|
||||||
@@ -29,3 +34,23 @@ jobs:
|
|||||||
run: uv sync --dev
|
run: uv sync --dev
|
||||||
- name: Run pytest
|
- name: Run pytest
|
||||||
run: uv run pytest -v --timeout=30
|
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/
|
*meals/
|
||||||
botpy.log
|
botpy.log
|
||||||
large_tool_results/
|
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`(涉及文件)→ 全净
|
||||||
+179
-278
@@ -19,7 +19,6 @@ Usage:
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from collections.abc import Sequence
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
@@ -56,19 +55,11 @@ DEFAULT_SKILL_SOURCES = ("/skills/",)
|
|||||||
|
|
||||||
_config = None
|
_config = None
|
||||||
_chat_model = None
|
_chat_model = None
|
||||||
# Track the (model, provider) binding of _chat_model so cache invalidates
|
# Track the (provider_id, model_key, registry_revision) binding of
|
||||||
# when config.model/provider change (e.g. via /model). Without this,
|
# _chat_model so the cache invalidates when the registry default changes.
|
||||||
# _ensure_chat_model() returns the stale cached instance even after
|
# The compile-time binding is only a placeholder — per-run resolution
|
||||||
# _ensure_config(new_cfg) has overwritten the active config — causing
|
# always comes from the run snapshot via ConfigurableModelMiddleware.
|
||||||
# /model switch to lag one step (see issue #179).
|
_chat_model_key: tuple[str, str, int] | None = None
|
||||||
_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
|
|
||||||
|
|
||||||
# Cache MCP tools by the effective config signature to avoid reconnecting
|
# Cache MCP tools by the effective config signature to avoid reconnecting
|
||||||
# to MCP servers on every `/new` when config is unchanged.
|
# to MCP servers on every `/new` when config is unchanged.
|
||||||
@@ -89,10 +80,10 @@ _EvoScientist_agent = None
|
|||||||
def set_active_config(cfg) -> None:
|
def set_active_config(cfg) -> None:
|
||||||
"""Commit *cfg* as the active module config.
|
"""Commit *cfg* as the active module config.
|
||||||
|
|
||||||
Public commit path for callers (e.g. ``/model``) that built an agent on
|
Public commit path for callers that built an agent on the pure
|
||||||
the pure ``create_cli_agent(config=..., chat_model=...)`` path and now
|
``create_cli_agent(config=..., chat_model=...)`` path and now want it
|
||||||
want it to become the session-wide active config. This is the write half
|
to become the session-wide active config. This is the write half of
|
||||||
of ``_ensure_config(cfg)`` extracted so the pure path can defer the commit
|
``_ensure_config(cfg)`` extracted so the pure path can defer the commit
|
||||||
until the agent has been built successfully.
|
until the agent has been built successfully.
|
||||||
"""
|
"""
|
||||||
global _config
|
global _config
|
||||||
@@ -119,25 +110,11 @@ def _ensure_config(config=None):
|
|||||||
return _config
|
return _config
|
||||||
|
|
||||||
|
|
||||||
def _build_chat_model(cfg):
|
def _replace_chat_model(instance, key: tuple[str, str, int]) -> None:
|
||||||
"""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:
|
|
||||||
"""Install a new chat model and propagate the related invariants.
|
"""Install a new chat model and propagate the related invariants.
|
||||||
|
|
||||||
Single write point for ``_chat_model`` / ``_chat_model_key`` /
|
Single write point for ``_chat_model`` / ``_chat_model_key`` /
|
||||||
``_EvoScientist_agent``: both ``_ensure_chat_model`` (cache-miss
|
``_EvoScientist_agent``: ``_ensure_chat_model`` cache-miss rebuilds
|
||||||
rebuild) and ``set_chat_model`` (explicit switch via ``/model``)
|
|
||||||
funnel through here so the three globals can never drift.
|
funnel through here so the three globals can never drift.
|
||||||
"""
|
"""
|
||||||
global _chat_model, _chat_model_key, _EvoScientist_agent
|
global _chat_model, _chat_model_key, _EvoScientist_agent
|
||||||
@@ -149,81 +126,50 @@ def _replace_chat_model(instance, key: tuple[str | None, str | None]) -> None:
|
|||||||
|
|
||||||
|
|
||||||
def _ensure_chat_model():
|
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
|
The model is resolved from the active registry's ``defaults.primary``
|
||||||
differs from the key that built ``_chat_model``, rebuild — this makes
|
and cached under its ``(provider_id, model_key, registry_revision)``
|
||||||
``create_cli_agent(config=temp_cfg)`` bind the freshly requested model
|
key. It is only the compile-time placeholder: every run re-resolves its
|
||||||
into the new agent without requiring callers to interleave
|
model from the run snapshot via ``ConfigurableModelMiddleware``.
|
||||||
``set_chat_model()`` calls in any particular order.
|
|
||||||
|
Raises:
|
||||||
|
ModelRegistryError: ``MODEL_REGISTRY_NOT_READY`` when the registry
|
||||||
|
is still in bootstrap (no enabled primary model configured).
|
||||||
"""
|
"""
|
||||||
cfg = _ensure_config()
|
from .model_registry.runtime import get_snapshot_runtime
|
||||||
key = (cfg.model, cfg.provider)
|
|
||||||
|
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:
|
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
|
return _chat_model
|
||||||
|
|
||||||
|
|
||||||
def _ensure_auxiliary_chat_model():
|
def _compile_time_role_model(role: str = "primary"):
|
||||||
"""Return the auxiliary chat model for background/helper LLM calls.
|
"""Return the compile-time model binding for graph construction.
|
||||||
|
|
||||||
Resolves ``(cfg.auxiliary_model or cfg.model, cfg.auxiliary_provider or
|
Unlike ``_ensure_chat_model()`` this never raises on a bootstrap
|
||||||
cfg.provider)``. When the auxiliary fields are empty — or resolve to the same
|
registry: graphs must still materialize so the Config API can serve
|
||||||
``(model, provider)`` pair as the main model — returns the main
|
(design doc section 10 — only run creation is forbidden in bootstrap),
|
||||||
``_ensure_chat_model()`` instance directly, so no second client is built.
|
so a ``RegistryNotReadyChatModel`` placeholder is bound instead. It
|
||||||
Otherwise it is cached separately under its own key. Onboard sets the
|
raises ``MODEL_REGISTRY_NOT_READY`` on the first model call; per-run
|
||||||
provider alongside the model, so the ``or cfg.provider`` fallback only
|
resolution still comes from the run snapshot via
|
||||||
matters for a model set without an explicit auxiliary provider.
|
``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 .model_registry.errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError
|
||||||
from .llm import get_chat_model
|
from .model_registry.placeholder import RegistryNotReadyChatModel
|
||||||
|
from .model_registry.runtime import get_snapshot_runtime
|
||||||
|
|
||||||
cfg = _ensure_config()
|
runtime = get_snapshot_runtime()
|
||||||
aux_model = cfg.auxiliary_model or cfg.model
|
try:
|
||||||
aux_provider = cfg.auxiliary_provider or cfg.provider
|
return runtime.build_default_role_model(role)
|
||||||
if (aux_model, aux_provider) == (cfg.model, cfg.provider):
|
except ModelRegistryError as exc:
|
||||||
return _ensure_chat_model()
|
if exc.code != MODEL_REGISTRY_NOT_READY:
|
||||||
key = (aux_model, aux_provider)
|
raise
|
||||||
if _auxiliary_chat_model is None or _auxiliary_chat_model_key != key:
|
return RegistryNotReadyChatModel(detail=str(exc))
|
||||||
_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)
|
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -305,12 +251,8 @@ def _inject_subagent_middleware(
|
|||||||
path doesn't fall back to the global-writing ``_ensure_chat_model()``.
|
path doesn't fall back to the global-writing ``_ensure_chat_model()``.
|
||||||
"""
|
"""
|
||||||
from .middleware import (
|
from .middleware import (
|
||||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
|
||||||
ContextOverflowMapperMiddleware,
|
ContextOverflowMapperMiddleware,
|
||||||
ErrorNormalizationMiddleware,
|
|
||||||
RepetitiveToolCallGuardMiddleware,
|
|
||||||
ToolErrorHandlerMiddleware,
|
ToolErrorHandlerMiddleware,
|
||||||
ToolProtocolGuardMiddleware,
|
|
||||||
create_context_editing_middleware,
|
create_context_editing_middleware,
|
||||||
create_memory_lifecycle_middleware,
|
create_memory_lifecycle_middleware,
|
||||||
create_memory_middleware,
|
create_memory_middleware,
|
||||||
@@ -319,16 +261,6 @@ def _inject_subagent_middleware(
|
|||||||
)
|
)
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
cfg = cfg if cfg is not None else _ensure_config()
|
||||||
repetitive_tool_call_threshold = getattr(
|
|
||||||
cfg,
|
|
||||||
"repetitive_tool_call_threshold",
|
|
||||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
|
||||||
)
|
|
||||||
if not isinstance(repetitive_tool_call_threshold, int):
|
|
||||||
repetitive_tool_call_threshold = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD
|
|
||||||
max_consecutive_tool_errors = getattr(cfg, "max_consecutive_tool_errors", 3)
|
|
||||||
if not isinstance(max_consecutive_tool_errors, int):
|
|
||||||
max_consecutive_tool_errors = 3
|
|
||||||
memory_controls = MemoryControls.from_config(cfg)
|
memory_controls = MemoryControls.from_config(cfg)
|
||||||
memory_dir = str(_paths_mod.MEMORIES_DIR)
|
memory_dir = str(_paths_mod.MEMORIES_DIR)
|
||||||
memory_scheduler = default_memory_scheduler()
|
memory_scheduler = default_memory_scheduler()
|
||||||
@@ -348,16 +280,6 @@ def _inject_subagent_middleware(
|
|||||||
memory_scheduler=memory_scheduler,
|
memory_scheduler=memory_scheduler,
|
||||||
)
|
)
|
||||||
middleware = [
|
middleware = [
|
||||||
# Outermost — catches provider-SDK exceptions from the
|
|
||||||
# model call (including inner middlewares) and normalizes
|
|
||||||
# them into a non-dataclass envelope wrapper before
|
|
||||||
# anything downstream sees them.
|
|
||||||
ErrorNormalizationMiddleware(),
|
|
||||||
RepetitiveToolCallGuardMiddleware(
|
|
||||||
threshold=repetitive_tool_call_threshold,
|
|
||||||
max_consecutive_errors=max_consecutive_tool_errors,
|
|
||||||
),
|
|
||||||
ToolProtocolGuardMiddleware(),
|
|
||||||
# Subagents share the main agent's model: use the threaded
|
# Subagents share the main agent's model: use the threaded
|
||||||
# ``chat_model`` on the pure path, else defer to the factory's
|
# ``chat_model`` on the pure path, else defer to the factory's
|
||||||
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
|
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
|
||||||
@@ -366,9 +288,12 @@ def _inject_subagent_middleware(
|
|||||||
ToolErrorHandlerMiddleware(),
|
ToolErrorHandlerMiddleware(),
|
||||||
ContextOverflowMapperMiddleware(),
|
ContextOverflowMapperMiddleware(),
|
||||||
]
|
]
|
||||||
if memory_controls.memory_enabled:
|
if memory_controls.memory_enabled and cfg.workspace_isolation != "required":
|
||||||
middleware.append(memory_middleware)
|
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(
|
middleware.append(
|
||||||
create_memory_lifecycle_middleware(
|
create_memory_lifecycle_middleware(
|
||||||
memory_dir,
|
memory_dir,
|
||||||
@@ -488,9 +413,10 @@ def _maybe_swap_async_subagents(
|
|||||||
|
|
||||||
middleware.append(AsyncWatcherMiddleware(agent_specs))
|
middleware.append(AsyncWatcherMiddleware(agent_specs))
|
||||||
|
|
||||||
# Forward the CLI's live (model, provider) into deepagents'
|
# Wrap deepagents' start/update_async_task tool calls so workspace scope
|
||||||
# start/update_async_task tool calls so the deployed graph can
|
# and usage correlation metadata reach the deployed graph's runs. Model
|
||||||
# re-resolve its chat model per run via ConfigurableModelMiddleware.
|
# 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.
|
# Idempotent — safe to call on every CLI startup.
|
||||||
if agent_specs:
|
if agent_specs:
|
||||||
from .llm.patches import _patch_deepagents_model_passthrough
|
from .llm.patches import _patch_deepagents_model_passthrough
|
||||||
@@ -504,14 +430,25 @@ def _build_base_kwargs(
|
|||||||
base_backend, base_middleware, *, cfg=None, chat_model=None, workspace_dir=None
|
base_backend, base_middleware, *, cfg=None, chat_model=None, workspace_dir=None
|
||||||
):
|
):
|
||||||
"""Build agent kwargs *without* MCP (fast, no subprocess spawning)."""
|
"""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
|
from .utils import load_subagents
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
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}
|
tool_registry = {"think_tool": think_tool}
|
||||||
if os.environ.get("TAVILY_API_KEY"):
|
if os.environ.get("TAVILY_API_KEY"):
|
||||||
tool_registry["tavily_search"] = tavily_search
|
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(
|
subs = load_subagents(
|
||||||
SUBAGENTS_CONFIG,
|
SUBAGENTS_CONFIG,
|
||||||
@@ -519,12 +456,12 @@ def _build_base_kwargs(
|
|||||||
)
|
)
|
||||||
_ensure_general_purpose_subagent(subs)
|
_ensure_general_purpose_subagent(subs)
|
||||||
_inject_subagent_middleware(
|
_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)
|
subs = _maybe_swap_async_subagents(subs, base_middleware, cfg=cfg)
|
||||||
return {
|
return {
|
||||||
"name": "EvoScientist",
|
"name": "EvoScientist",
|
||||||
"model": chat_model if chat_model is not None else _ensure_chat_model(),
|
"model": model,
|
||||||
"tools": list(base_tools),
|
"tools": list(base_tools),
|
||||||
"backend": base_backend,
|
"backend": base_backend,
|
||||||
"subagents": subs,
|
"subagents": subs,
|
||||||
@@ -556,24 +493,35 @@ def load_mcp_and_build_kwargs(
|
|||||||
chat_model: Explicit chat model to bind instead of
|
chat_model: Explicit chat model to bind instead of
|
||||||
``_ensure_chat_model()`` (which would write module globals).
|
``_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
|
from .utils import load_subagents
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
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)
|
mcp_by_agent = _load_mcp_tools_cached(on_progress=on_mcp_progress)
|
||||||
if not mcp_by_agent:
|
if not mcp_by_agent:
|
||||||
return _build_base_kwargs(
|
return _build_base_kwargs(
|
||||||
base_backend,
|
base_backend,
|
||||||
base_middleware,
|
base_middleware,
|
||||||
cfg=cfg,
|
cfg=cfg,
|
||||||
chat_model=chat_model,
|
chat_model=model,
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
)
|
)
|
||||||
|
|
||||||
tool_registry = {"think_tool": think_tool}
|
tool_registry = {"think_tool": think_tool}
|
||||||
if os.environ.get("TAVILY_API_KEY"):
|
if os.environ.get("TAVILY_API_KEY"):
|
||||||
tool_registry["tavily_search"] = tavily_search
|
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
|
# Fresh tool registry — start from base tools + MCP tools
|
||||||
registry = dict(tool_registry)
|
registry = dict(tool_registry)
|
||||||
@@ -590,7 +538,7 @@ def load_mcp_and_build_kwargs(
|
|||||||
|
|
||||||
_ensure_general_purpose_subagent(subs)
|
_ensure_general_purpose_subagent(subs)
|
||||||
_inject_subagent_middleware(
|
_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
|
# Inject MCP tools into subagents by name
|
||||||
@@ -604,7 +552,7 @@ def load_mcp_and_build_kwargs(
|
|||||||
|
|
||||||
return {
|
return {
|
||||||
"name": "EvoScientist",
|
"name": "EvoScientist",
|
||||||
"model": chat_model if chat_model is not None else _ensure_chat_model(),
|
"model": model,
|
||||||
"tools": base_tools + mcp_main,
|
"tools": base_tools + mcp_main,
|
||||||
"backend": base_backend,
|
"backend": base_backend,
|
||||||
"subagents": subs,
|
"subagents": subs,
|
||||||
@@ -619,8 +567,8 @@ def load_mcp_and_build_kwargs(
|
|||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
def _get_default_backend():
|
def _get_legacy_backend():
|
||||||
"""Build the default composite backend from current paths."""
|
"""Build the deployment-root backend used by CLI and legacy mode only."""
|
||||||
from deepagents.backends import CompositeBackend
|
from deepagents.backends import CompositeBackend
|
||||||
|
|
||||||
from .backends import (
|
from .backends import (
|
||||||
@@ -662,18 +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(
|
def _get_default_middleware(
|
||||||
*,
|
*,
|
||||||
for_async_subagent: bool = False,
|
for_async_subagent: bool = False,
|
||||||
workspace_dir: str | Path | None = None,
|
workspace_dir: str | Path | None = None,
|
||||||
memory_dir: str | Path | None = None,
|
|
||||||
cfg=None,
|
cfg=None,
|
||||||
chat_model=None,
|
chat_model=None,
|
||||||
|
backend=None,
|
||||||
memory_source_agent: str = "EvoScientist",
|
memory_source_agent: str = "EvoScientist",
|
||||||
tool_selector_threshold: int | None = None,
|
snapshot_role: str = "primary",
|
||||||
memory_max_inline_profile_chars: int | None = None,
|
|
||||||
enable_background_execution: bool = True,
|
|
||||||
enable_legacy_model_fallback: bool = True,
|
|
||||||
):
|
):
|
||||||
"""Build the default middleware list.
|
"""Build the default middleware list.
|
||||||
|
|
||||||
@@ -693,42 +668,35 @@ def _get_default_middleware(
|
|||||||
(avoids writing module globals on the pure path).
|
(avoids writing module globals on the pure path).
|
||||||
memory_source_agent: Attribution name for profile/observation writes.
|
memory_source_agent: Attribution name for profile/observation writes.
|
||||||
Async sub-agent factories pass their deployed agent name here.
|
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 (
|
from .middleware import (
|
||||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
|
||||||
ConfigurableModelMiddleware,
|
ConfigurableModelMiddleware,
|
||||||
ContextOverflowMapperMiddleware,
|
ContextOverflowMapperMiddleware,
|
||||||
ErrorNormalizationMiddleware,
|
|
||||||
ModelFallbackMiddleware,
|
|
||||||
RepetitiveToolCallGuardMiddleware,
|
|
||||||
ToolErrorHandlerMiddleware,
|
ToolErrorHandlerMiddleware,
|
||||||
ToolProtocolGuardMiddleware,
|
|
||||||
create_code_interpreter_middleware,
|
create_code_interpreter_middleware,
|
||||||
create_context_editing_middleware,
|
create_context_editing_middleware,
|
||||||
create_memory_lifecycle_middleware,
|
create_memory_lifecycle_middleware,
|
||||||
create_memory_middleware,
|
create_memory_middleware,
|
||||||
|
create_message_budget_middleware,
|
||||||
create_runtime_context_middleware,
|
create_runtime_context_middleware,
|
||||||
create_scheduler_middleware,
|
create_scheduler_middleware,
|
||||||
create_tool_selector_middleware,
|
create_tool_selector_middleware,
|
||||||
default_memory_scheduler,
|
default_memory_scheduler,
|
||||||
load_fallback_chain,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
cfg = cfg if cfg is not None else _ensure_config()
|
||||||
repetitive_tool_call_threshold = getattr(
|
model = chat_model if chat_model is not None else _compile_time_role_model("primary")
|
||||||
cfg,
|
if backend is None:
|
||||||
"repetitive_tool_call_threshold",
|
# Preserve the factory's pure path for callers that provide an
|
||||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
# explicit model/configuration (notably tests and subagent assembly).
|
||||||
)
|
# Production graph factories always pass their real composite backend.
|
||||||
if not isinstance(repetitive_tool_call_threshold, int):
|
from deepagents.backends import StateBackend
|
||||||
repetitive_tool_call_threshold = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD
|
|
||||||
max_consecutive_tool_errors = getattr(cfg, "max_consecutive_tool_errors", 3)
|
backend = StateBackend()
|
||||||
if not isinstance(max_consecutive_tool_errors, int):
|
memory_dir = str(_paths_mod.MEMORIES_DIR)
|
||||||
max_consecutive_tool_errors = 3
|
|
||||||
if cfg.model_fallbacks:
|
|
||||||
load_fallback_chain(cfg.model_fallbacks)
|
|
||||||
model = chat_model if chat_model is not None else _ensure_chat_model()
|
|
||||||
memory_dir = str(memory_dir or _paths_mod.MEMORIES_DIR)
|
|
||||||
source_type = (
|
source_type = (
|
||||||
MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN
|
MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN
|
||||||
)
|
)
|
||||||
@@ -739,80 +707,59 @@ def _get_default_middleware(
|
|||||||
if for_async_subagent
|
if for_async_subagent
|
||||||
else MemoryObservationTarget.TURN_WORKER
|
else MemoryObservationTarget.TURN_WORKER
|
||||||
)
|
)
|
||||||
# ``ConfigurableModelMiddleware`` is placed first so it wraps
|
# ``ConfigurableModelMiddleware`` sits first so the snapshot-driven model
|
||||||
# ``ModelFallbackMiddleware``: a configurable.model override sets the
|
# override applies before any other middleware inspects the request.
|
||||||
# PRIMARY model only, leaving the fallback chain free to try its own
|
memory_middleware = create_memory_middleware(
|
||||||
# alternatives instead of re-overriding every retry to the same model.
|
memory_dir,
|
||||||
memory_kwargs = {
|
workspace_dir=workspace_dir,
|
||||||
"workspace_dir": workspace_dir,
|
source_type=source_type,
|
||||||
"source_type": source_type,
|
source_agent=memory_source_agent,
|
||||||
"source_agent": memory_source_agent,
|
enable_profile_memory=memory_controls.profile_enabled,
|
||||||
"enable_profile_memory": memory_controls.profile_enabled,
|
enable_observation_memory=memory_controls.observations_enabled,
|
||||||
"enable_observation_memory": memory_controls.observations_enabled,
|
enable_observation_tool=memory_controls.observation_tool_enabled(
|
||||||
"enable_observation_tool": memory_controls.observation_tool_enabled(
|
|
||||||
MemoryObservationTarget.AGENT
|
MemoryObservationTarget.AGENT
|
||||||
),
|
),
|
||||||
"memory_scheduler": memory_scheduler,
|
memory_scheduler=memory_scheduler,
|
||||||
}
|
)
|
||||||
if memory_max_inline_profile_chars is not None:
|
# Main-agent tool selection resolves its helper model from the run
|
||||||
memory_kwargs["max_inline_profile_chars"] = memory_max_inline_profile_chars
|
# snapshot on every call (model=None below); async sub-agents and the
|
||||||
memory_middleware = create_memory_middleware(memory_dir, **memory_kwargs)
|
# pure path (explicit model + config) keep their threaded model.
|
||||||
# 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).
|
|
||||||
# context_editing stays on the main model — its model only sizes the
|
# context_editing stays on the main model — its model only sizes the
|
||||||
# context-window trigger for the main agent's own history.
|
# context-window trigger for the main agent's own history.
|
||||||
if for_async_subagent:
|
tool_selector_model = None if (not for_async_subagent and chat_model is None) else model
|
||||||
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)
|
|
||||||
selector_middlewares = create_tool_selector_middleware(
|
|
||||||
**(
|
|
||||||
{"threshold": tool_selector_threshold}
|
|
||||||
if tool_selector_threshold is not None
|
|
||||||
else {}
|
|
||||||
),
|
|
||||||
model=tool_selector_model,
|
|
||||||
track_stream_selection=not for_async_subagent,
|
|
||||||
)
|
|
||||||
mw = [
|
mw = [
|
||||||
# Outermost — catches provider-SDK exceptions from the model
|
ConfigurableModelMiddleware(role=snapshot_role),
|
||||||
# call (including exceptions surfaced through inner
|
create_message_budget_middleware(model, backend, snapshot_role=snapshot_role),
|
||||||
# middlewares) and normalizes them into a non-dataclass
|
|
||||||
# envelope wrapper before anything downstream sees them.
|
|
||||||
ErrorNormalizationMiddleware(),
|
|
||||||
ConfigurableModelMiddleware(),
|
|
||||||
create_context_editing_middleware(model),
|
create_context_editing_middleware(model),
|
||||||
*([ModelFallbackMiddleware()] if enable_legacy_model_fallback else []),
|
|
||||||
RepetitiveToolCallGuardMiddleware(
|
|
||||||
threshold=repetitive_tool_call_threshold,
|
|
||||||
max_consecutive_errors=max_consecutive_tool_errors,
|
|
||||||
),
|
|
||||||
ContextOverflowMapperMiddleware(),
|
ContextOverflowMapperMiddleware(),
|
||||||
ToolErrorHandlerMiddleware(),
|
ToolErrorHandlerMiddleware(),
|
||||||
*selector_middlewares,
|
*create_tool_selector_middleware(
|
||||||
ToolProtocolGuardMiddleware(),
|
model=tool_selector_model,
|
||||||
|
track_stream_selection=not for_async_subagent,
|
||||||
|
),
|
||||||
# Interpreter prompt must land before runtime/memory context, so this
|
# Interpreter prompt must land before runtime/memory context, so this
|
||||||
# middleware sits ahead of runtime_context in the stack.
|
# middleware sits ahead of runtime_context in the stack.
|
||||||
|
*(
|
||||||
|
[]
|
||||||
|
if cfg.workspace_isolation == "required"
|
||||||
|
and cfg.strict_code_interpreter == "disabled"
|
||||||
|
else [
|
||||||
create_code_interpreter_middleware(
|
create_code_interpreter_middleware(
|
||||||
timeout=cfg.code_interpreter_timeout,
|
timeout=cfg.code_interpreter_timeout,
|
||||||
max_result_chars=cfg.code_interpreter_max_result_chars,
|
max_result_chars=cfg.code_interpreter_max_result_chars,
|
||||||
|
)
|
||||||
|
]
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
if cfg.enable_scheduler and not for_async_subagent:
|
if cfg.enable_scheduler and not for_async_subagent:
|
||||||
mw.append(create_scheduler_middleware())
|
mw.append(create_scheduler_middleware())
|
||||||
mw.append(create_runtime_context_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)
|
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(
|
mw.append(
|
||||||
create_memory_lifecycle_middleware(
|
create_memory_lifecycle_middleware(
|
||||||
memory_dir,
|
memory_dir,
|
||||||
@@ -832,7 +779,7 @@ def _get_default_middleware(
|
|||||||
# Background-process tools (run_in_background / check_process / stop_process /
|
# Background-process tools (run_in_background / check_process / stop_process /
|
||||||
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
|
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
|
||||||
# must not spawn local OS processes.
|
# must not spawn local OS processes.
|
||||||
if not for_async_subagent and enable_background_execution:
|
if not for_async_subagent and cfg.workspace_isolation != "required":
|
||||||
from .middleware.background import BackgroundExecutionMiddleware
|
from .middleware.background import BackgroundExecutionMiddleware
|
||||||
|
|
||||||
mw.append(BackgroundExecutionMiddleware())
|
mw.append(BackgroundExecutionMiddleware())
|
||||||
@@ -869,7 +816,7 @@ def _get_default_agent():
|
|||||||
|
|
||||||
cfg = _ensure_config()
|
cfg = _ensure_config()
|
||||||
be = _get_default_backend()
|
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,
|
# HITL on main agent only (mirrors create_cli_agent). Use middleware,
|
||||||
# not interrupt_on= kwarg — the kwarg propagates to every subagent and
|
# not interrupt_on= kwarg — the kwarg propagates to every subagent and
|
||||||
@@ -886,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(
|
kwargs = _build_base_kwargs(
|
||||||
be,
|
be,
|
||||||
mw,
|
mw,
|
||||||
@@ -930,14 +880,6 @@ def create_cli_agent(
|
|||||||
chat_model=None,
|
chat_model=None,
|
||||||
*,
|
*,
|
||||||
on_mcp_progress=None,
|
on_mcp_progress=None,
|
||||||
workspace_backend=None,
|
|
||||||
memory_dir: str | Path | None = None,
|
|
||||||
tool_selector_threshold: int | None = None,
|
|
||||||
memory_max_inline_profile_chars: int | None = None,
|
|
||||||
enable_subagents: bool = True,
|
|
||||||
enable_background_execution: bool = True,
|
|
||||||
main_agent_outer_middlewares: Sequence[AgentMiddleware] | None = None,
|
|
||||||
main_agent_route_middleware: AgentMiddleware | None = None,
|
|
||||||
) -> "CompiledStateGraph":
|
) -> "CompiledStateGraph":
|
||||||
"""Create agent with checkpointer for CLI multi-turn support.
|
"""Create agent with checkpointer for CLI multi-turn support.
|
||||||
|
|
||||||
@@ -948,10 +890,10 @@ def create_cli_agent(
|
|||||||
**Pure path:** when *both* ``config`` and ``chat_model`` are explicit, this
|
**Pure path:** when *both* ``config`` and ``chat_model`` are explicit, this
|
||||||
writes none of the cached config/model module globals (``_config``,
|
writes none of the cached config/model module globals (``_config``,
|
||||||
``_chat_model``, ``_chat_model_key``, ``_EvoScientist_agent``) — the agent
|
``_chat_model``, ``_chat_model_key``, ``_EvoScientist_agent``) — the agent
|
||||||
is built purely from the passed-in locals. The caller commits the switch
|
is built purely from the passed-in locals. The caller commits the config
|
||||||
on success via ``set_active_config`` / ``set_chat_model_instance`` (see
|
on success via ``set_active_config``. Otherwise the existing
|
||||||
``/model``). Otherwise the existing module-global path runs (langgraph
|
module-global path runs (langgraph dev, notebooks, and CLI startup, which
|
||||||
dev, notebooks, and CLI startup, which pass ``config=`` only).
|
pass ``config=`` only).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
workspace_dir: Per-session workspace directory. If ``None``,
|
workspace_dir: Per-session workspace directory. If ``None``,
|
||||||
@@ -964,22 +906,6 @@ def create_cli_agent(
|
|||||||
chat_model: Optional pre-built chat model. Only triggers the pure
|
chat_model: Optional pre-built chat model. Only triggers the pure
|
||||||
path when ``config`` is also explicit; otherwise it is ignored in
|
path when ``config`` is also explicit; otherwise it is ignored in
|
||||||
favor of the ``_ensure_chat_model()`` fallback.
|
favor of the ``_ensure_chat_model()`` fallback.
|
||||||
workspace_backend: Optional host-provided backend for the workspace
|
|
||||||
route. The default remains ``CustomSandboxBackend``.
|
|
||||||
memory_dir: Optional memory root used by both the backend route and
|
|
||||||
memory middleware.
|
|
||||||
tool_selector_threshold: Optional adaptive tool-selection threshold.
|
|
||||||
memory_max_inline_profile_chars: Optional memory profile injection cap.
|
|
||||||
enable_subagents: Whether configured subagents are available to the agent.
|
|
||||||
enable_background_execution: Whether local background-process tools are
|
|
||||||
installed. Embedding hosts should disable this when process execution
|
|
||||||
is provided by an external backend.
|
|
||||||
main_agent_outer_middlewares: Optional host-owned middleware installed
|
|
||||||
only on the top-level agent, outside EvoScientist's default chain.
|
|
||||||
main_agent_route_middleware: Optional host-owned route middleware placed
|
|
||||||
after ConfigurableModelMiddleware and before tool selection. When
|
|
||||||
provided, EvoScientist's legacy model fallback is disabled for the
|
|
||||||
top-level agent so the host is the only fallback authority.
|
|
||||||
"""
|
"""
|
||||||
import os as _os
|
import os as _os
|
||||||
|
|
||||||
@@ -1021,15 +947,13 @@ def create_cli_agent(
|
|||||||
workspace_dir = str(_paths.WORKSPACE_ROOT)
|
workspace_dir = str(_paths.WORKSPACE_ROOT)
|
||||||
|
|
||||||
# Read paths dynamically so runtime set_workspace_root() changes are picked up
|
# Read paths dynamically so runtime set_workspace_root() changes are picked up
|
||||||
_mem_dir = str(memory_dir or _paths.MEMORIES_DIR)
|
_mem_dir = str(_paths.MEMORIES_DIR)
|
||||||
_usr_skills_dir = str(_paths.USER_SKILLS_DIR)
|
_usr_skills_dir = str(_paths.USER_SKILLS_DIR)
|
||||||
_global_skills_dir = str(_paths.GLOBAL_SKILLS_DIR)
|
_global_skills_dir = str(_paths.GLOBAL_SKILLS_DIR)
|
||||||
|
|
||||||
# Always construct fresh backends from current paths (avoids stale
|
# Always construct fresh backends from current paths (avoids stale
|
||||||
# module-level backend when workspace root changed at runtime).
|
# module-level backend when workspace root changed at runtime).
|
||||||
set_active_workspace(workspace_dir)
|
set_active_workspace(workspace_dir)
|
||||||
ws_backend = workspace_backend
|
|
||||||
if ws_backend is None:
|
|
||||||
ws_backend = CustomSandboxBackend(
|
ws_backend = CustomSandboxBackend(
|
||||||
root_dir=workspace_dir,
|
root_dir=workspace_dir,
|
||||||
virtual_mode=True,
|
virtual_mode=True,
|
||||||
@@ -1057,29 +981,8 @@ def create_cli_agent(
|
|||||||
# CLI agent never drifts from the default chain. Anything CLI-specific
|
# CLI agent never drifts from the default chain. Anything CLI-specific
|
||||||
# (e.g. ``HumanInTheLoopMiddleware``) is appended below.
|
# (e.g. ``HumanInTheLoopMiddleware``) is appended below.
|
||||||
mw: list[AgentMiddleware] = _get_default_middleware(
|
mw: list[AgentMiddleware] = _get_default_middleware(
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model, backend=be
|
||||||
memory_dir=_mem_dir,
|
|
||||||
cfg=cfg,
|
|
||||||
chat_model=chat_model,
|
|
||||||
tool_selector_threshold=tool_selector_threshold,
|
|
||||||
memory_max_inline_profile_chars=memory_max_inline_profile_chars,
|
|
||||||
enable_background_execution=enable_background_execution,
|
|
||||||
enable_legacy_model_fallback=main_agent_route_middleware is None,
|
|
||||||
)
|
)
|
||||||
if main_agent_route_middleware is not None:
|
|
||||||
configurable_index = next(
|
|
||||||
(
|
|
||||||
index
|
|
||||||
for index, middleware in enumerate(mw)
|
|
||||||
if getattr(middleware, "name", "") == "configurable_model"
|
|
||||||
),
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
if configurable_index is None:
|
|
||||||
raise RuntimeError("ConfigurableModelMiddleware route slot is unavailable")
|
|
||||||
mw.insert(configurable_index + 1, main_agent_route_middleware)
|
|
||||||
if main_agent_outer_middlewares:
|
|
||||||
mw = [*main_agent_outer_middlewares, *mw]
|
|
||||||
|
|
||||||
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
|
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
|
||||||
# would propagate it to every subagent, breaking parallel execute calls
|
# would propagate it to every subagent, breaking parallel execute calls
|
||||||
@@ -1104,8 +1007,6 @@ def create_cli_agent(
|
|||||||
chat_model=chat_model,
|
chat_model=chat_model,
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
)
|
)
|
||||||
if not enable_subagents:
|
|
||||||
kwargs = {**kwargs, "subagents": []}
|
|
||||||
|
|
||||||
return create_deep_agent(
|
return create_deep_agent(
|
||||||
**kwargs,
|
**kwargs,
|
||||||
|
|||||||
@@ -9,8 +9,6 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from importlib import import_module
|
from importlib import import_module
|
||||||
|
|
||||||
__version__ = "0.2.2"
|
|
||||||
|
|
||||||
_EXPORTS: dict[str, tuple[str, str]] = {
|
_EXPORTS: dict[str, tuple[str, str]] = {
|
||||||
# Agent graph (lazy to avoid expensive initialization at import time)
|
# Agent graph (lazy to avoid expensive initialization at import time)
|
||||||
"EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"),
|
"EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"),
|
||||||
@@ -25,11 +23,6 @@ _EXPORTS: dict[str, tuple[str, str]] = {
|
|||||||
"save_config": (".config", "save_config"),
|
"save_config": (".config", "save_config"),
|
||||||
"get_effective_config": (".config", "get_effective_config"),
|
"get_effective_config": (".config", "get_effective_config"),
|
||||||
"get_config_path": (".config", "get_config_path"),
|
"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
|
# Prompts
|
||||||
"get_system_prompt": (".prompts", "get_system_prompt"),
|
"get_system_prompt": (".prompts", "get_system_prompt"),
|
||||||
# Tools
|
# Tools
|
||||||
|
|||||||
@@ -682,6 +682,27 @@ def _guard_bare_absolute(result: str | None) -> str | None:
|
|||||||
return result
|
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(
|
def _rewrite_quoted_path(
|
||||||
path: str,
|
path: str,
|
||||||
workspace_name: str | None,
|
workspace_name: str | None,
|
||||||
@@ -720,6 +741,7 @@ def _rewrite_quoted_path(
|
|||||||
def convert_virtual_paths_in_command(
|
def convert_virtual_paths_in_command(
|
||||||
command: str,
|
command: str,
|
||||||
workspace_name: str | None = None,
|
workspace_name: str | None = None,
|
||||||
|
workspace_root: str | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Convert virtual paths (starting with ``/``) in commands to relative paths.
|
"""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
|
mount (``/skills/...``, ``/memories/...``) or a workspace-prefixed
|
||||||
system path are rewritten as a single shell token — this fixes #237
|
system path are rewritten as a single shell token — this fixes #237
|
||||||
where ``python "/skills/my skill/main.py"`` was truncated at the
|
where ``python "/skills/my skill/main.py"`` was truncated at the
|
||||||
embedded space. Bare quoted ``/...`` paths (e.g. ``echo "/hi"``)
|
embedded space. When *workspace_root* is given, quoted ``/...`` paths
|
||||||
are left untouched since their semantics are ambiguous.
|
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
|
After pre-processing, the original regex handles unquoted
|
||||||
paths and workspace-name correction as before.
|
paths and workspace-name correction as before.
|
||||||
"""
|
"""
|
||||||
# Pre-process: rewrite quoted paths whose decoded content starts with /
|
# 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(
|
command = re.sub(
|
||||||
r'(["\'])((?:\\.|(?!\1).)*?)\1',
|
r'(["\'])((?:\\.|(?!\1).)*?)\1',
|
||||||
lambda m: (
|
_rewrite_quoted,
|
||||||
_rewrite_quoted_path(
|
|
||||||
re.sub(r"\\(.)", r"\1", m.group(2)),
|
|
||||||
workspace_name,
|
|
||||||
)
|
|
||||||
or m.group(0)
|
|
||||||
),
|
|
||||||
command,
|
command,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1103,6 +1140,7 @@ def prepare_sandbox_command(
|
|||||||
command = convert_virtual_paths_in_command(
|
command = convert_virtual_paths_in_command(
|
||||||
command=command,
|
command=command,
|
||||||
workspace_name=Path(cwd_str).name,
|
workspace_name=Path(cwd_str).name,
|
||||||
|
workspace_root=cwd_str,
|
||||||
)
|
)
|
||||||
# Skills/memory dirs must be allowlisted: the workspace-literal replace above runs
|
# 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.
|
# before the resolver, so any absolute path it later injects reaches validate unstripped.
|
||||||
|
|||||||
@@ -20,9 +20,6 @@ from EvoScientist.config import EvoScientistConfig
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_CCPROXY_AUTH_TIMEOUT_SECONDS = 30
|
|
||||||
_CCPROXY_HEALTH_TIMEOUT_SECONDS = 180
|
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# Availability & auth checks
|
# Availability & auth checks
|
||||||
@@ -130,11 +127,7 @@ def check_ccproxy_auth(provider: str = "claude_api") -> tuple[bool, str]:
|
|||||||
[exe, "auth", "status", provider],
|
[exe, "auth", "status", provider],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True,
|
text=True,
|
||||||
# ccproxy's CLI initializes its full plugin system on every
|
timeout=10,
|
||||||
# invocation — a cold start takes ~10s on Apple Silicon, so a
|
|
||||||
# 10s timeout made OAuth startup fail intermittently with
|
|
||||||
# "Auth check timed out".
|
|
||||||
timeout=_CCPROXY_AUTH_TIMEOUT_SECONDS,
|
|
||||||
)
|
)
|
||||||
import re as _re
|
import re as _re
|
||||||
|
|
||||||
@@ -183,33 +176,6 @@ def is_ccproxy_running(port: int) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def write_ccproxy_config() -> str:
|
|
||||||
"""Write the ccproxy config file EvoScientist passes to ``serve --config``.
|
|
||||||
|
|
||||||
Disables ccproxy's default Codex model mappings, which rewrite any
|
|
||||||
``gpt-*``/``o1-*``/``o3-*``/``claude-*`` model to ``gpt-5.3-codex``
|
|
||||||
before forwarding — silently overriding the model the user configured
|
|
||||||
(and failing outright on accounts where ``gpt-5.3-codex`` is not
|
|
||||||
served). With no mappings, the requested model reaches the Codex
|
|
||||||
backend unmodified.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Absolute path to the generated config file.
|
|
||||||
"""
|
|
||||||
from EvoScientist.config import get_config_dir
|
|
||||||
|
|
||||||
path = get_config_dir() / "ccproxy.toml"
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
path.write_text(
|
|
||||||
"# Generated by EvoScientist (ccproxy_manager) — do not edit;\n"
|
|
||||||
"# regenerated on every ccproxy start.\n"
|
|
||||||
"[plugins.codex]\n"
|
|
||||||
"model_mappings = []\n",
|
|
||||||
encoding="utf-8",
|
|
||||||
)
|
|
||||||
return str(path)
|
|
||||||
|
|
||||||
|
|
||||||
def start_ccproxy(port: int) -> subprocess.Popen:
|
def start_ccproxy(port: int) -> subprocess.Popen:
|
||||||
"""Start ccproxy serve as a background process.
|
"""Start ccproxy serve as a background process.
|
||||||
|
|
||||||
@@ -220,32 +186,18 @@ def start_ccproxy(port: int) -> subprocess.Popen:
|
|||||||
The Popen handle for the ccproxy process.
|
The Popen handle for the ccproxy process.
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
RuntimeError: If ccproxy fails to become healthy within
|
RuntimeError: If ccproxy fails to become healthy within 30 seconds.
|
||||||
``_CCPROXY_HEALTH_TIMEOUT_SECONDS``.
|
|
||||||
FileNotFoundError: If ccproxy binary is not found.
|
FileNotFoundError: If ccproxy binary is not found.
|
||||||
"""
|
"""
|
||||||
exe = _ccproxy_exe() or "ccproxy"
|
exe = _ccproxy_exe() or "ccproxy"
|
||||||
cmd = [exe, "serve", "--port", str(port)]
|
|
||||||
try:
|
|
||||||
cmd += ["--config", write_ccproxy_config()]
|
|
||||||
except (OSError, UnicodeError) as exc:
|
|
||||||
logger.warning(
|
|
||||||
"Could not write ccproxy config (%s); starting with defaults — "
|
|
||||||
"Codex model mappings will rewrite gpt-* models to gpt-5.3-codex",
|
|
||||||
exc,
|
|
||||||
)
|
|
||||||
logger.warning(
|
|
||||||
"Starting ccproxy on port %d; first startup may take up to %d seconds",
|
|
||||||
port,
|
|
||||||
_CCPROXY_HEALTH_TIMEOUT_SECONDS,
|
|
||||||
)
|
|
||||||
proc = subprocess.Popen(
|
proc = subprocess.Popen(
|
||||||
cmd,
|
[exe, "serve", "--port", str(port)],
|
||||||
stdout=subprocess.DEVNULL,
|
stdout=subprocess.DEVNULL,
|
||||||
stderr=subprocess.DEVNULL,
|
stderr=subprocess.DEVNULL,
|
||||||
)
|
)
|
||||||
|
|
||||||
deadline = time.monotonic() + _CCPROXY_HEALTH_TIMEOUT_SECONDS
|
# Wait for health (ccproxy can take up to ~11s on first start)
|
||||||
|
deadline = time.monotonic() + 30
|
||||||
while time.monotonic() < deadline:
|
while time.monotonic() < deadline:
|
||||||
if proc.poll() is not None:
|
if proc.poll() is not None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -261,10 +213,7 @@ def start_ccproxy(port: int) -> subprocess.Popen:
|
|||||||
proc.wait(timeout=3)
|
proc.wait(timeout=3)
|
||||||
except subprocess.TimeoutExpired:
|
except subprocess.TimeoutExpired:
|
||||||
proc.kill()
|
proc.kill()
|
||||||
raise RuntimeError(
|
raise RuntimeError("ccproxy did not become healthy within 30 seconds")
|
||||||
"ccproxy did not become healthy within "
|
|
||||||
f"{_CCPROXY_HEALTH_TIMEOUT_SECONDS} seconds"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def stop_ccproxy(proc: subprocess.Popen | None) -> None:
|
def stop_ccproxy(proc: subprocess.Popen | None) -> None:
|
||||||
|
|||||||
@@ -150,11 +150,10 @@ class BackgroundAgentLoader(Generic[AgentT]):
|
|||||||
def adopt(self, agent: AgentT) -> None:
|
def adopt(self, agent: AgentT) -> None:
|
||||||
"""Install an externally-built agent and supersede any in-flight load.
|
"""Install an externally-built agent and supersede any in-flight load.
|
||||||
|
|
||||||
Used by ``/model`` (and any other caller that constructs a
|
Used by any caller that constructs a replacement agent directly:
|
||||||
replacement agent directly): bumps the generation token so a
|
bumps the generation token so a late-arriving background load can't
|
||||||
late-arriving background load can't clobber ``self.agent`` via
|
clobber ``self.agent`` via the done-callback, cancels the in-flight
|
||||||
the done-callback, cancels the in-flight wrapper, and seats the
|
wrapper, and seats the new agent immediately.
|
||||||
new agent immediately.
|
|
||||||
"""
|
"""
|
||||||
prev = self._task
|
prev = self._task
|
||||||
if prev is not None and not prev.done():
|
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
|
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(
|
def _load_agent(
|
||||||
workspace_dir: str | None = None,
|
workspace_dir: str | None = None,
|
||||||
checkpointer=None,
|
checkpointer=None,
|
||||||
|
|||||||
@@ -306,7 +306,7 @@ async def dispatch_channel_slash_command(
|
|||||||
resolution, or the dispatcher's input agent when no resolver is
|
resolution, or the dispatcher's input agent when no resolver is
|
||||||
supplied. Callers can compare ``ctx.agent`` with
|
supplied. Callers can compare ``ctx.agent`` with
|
||||||
``original_agent`` to detect command-driven swaps. Used by Rich
|
``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
|
running session and (b) refresh the status snapshot for
|
||||||
commands that mutate session-level state (``/new``,
|
commands that mutate session-level state (``/new``,
|
||||||
``/compact``) — mirrors the REPL dispatch at
|
``/compact``) — mirrors the REPL dispatch at
|
||||||
|
|||||||
@@ -74,16 +74,9 @@ if TYPE_CHECKING:
|
|||||||
@app.command()
|
@app.command()
|
||||||
def onboard(
|
def onboard(
|
||||||
skip_validation: bool = typer.Option(
|
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)
|
# ---- 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(
|
tavily_key: str | None = typer.Option(
|
||||||
None, "--tavily-key", help="Pre-set Tavily API key"
|
None, "--tavily-key", help="Pre-set Tavily API key"
|
||||||
),
|
),
|
||||||
@@ -120,16 +113,16 @@ def onboard(
|
|||||||
):
|
):
|
||||||
"""Interactive setup wizard for EvoScientist.
|
"""Interactive setup wizard for EvoScientist.
|
||||||
|
|
||||||
Guides you through configuring API keys, model selection,
|
Guides you through workspace settings, channels, and agent parameters.
|
||||||
workspace settings, 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
|
Any answer can be pre-set via a flag; prompts for unset answers stay
|
||||||
--model claude-sonnet-4-6 ...``); prompts for unset answers stay
|
|
||||||
interactive unless ``--non-interactive`` is passed, in which case any
|
interactive unless ``--non-interactive`` is passed, in which case any
|
||||||
missing required answer aborts the wizard.
|
missing required answer aborts the wizard.
|
||||||
"""
|
"""
|
||||||
from ..config.onboard.constants import (
|
from ..config.onboard.constants import (
|
||||||
VALID_PROVIDERS,
|
|
||||||
VALID_UI_BACKENDS,
|
VALID_UI_BACKENDS,
|
||||||
VALID_WORKSPACE_MODES,
|
VALID_WORKSPACE_MODES,
|
||||||
)
|
)
|
||||||
@@ -149,11 +142,6 @@ def onboard(
|
|||||||
f"--workspace-mode must be one of {sorted(VALID_WORKSPACE_MODES)}",
|
f"--workspace-mode must be one of {sorted(VALID_WORKSPACE_MODES)}",
|
||||||
param_hint="--workspace-mode",
|
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
|
# Match the interactive prompt's range (1024 < port < 65536). Without
|
||||||
# this check, --port 80 or --port 99999 would land in config and break
|
# this check, --port 80 or --port 99999 would land in config and break
|
||||||
# the langgraph dev server on startup.
|
# the langgraph dev server on startup.
|
||||||
@@ -169,12 +157,6 @@ def onboard(
|
|||||||
answers["ui"] = ui
|
answers["ui"] = ui
|
||||||
if port is not None:
|
if port is not None:
|
||||||
answers["port"] = str(port)
|
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:
|
if tavily_key is not None:
|
||||||
answers["tavily_key"] = tavily_key
|
answers["tavily_key"] = tavily_key
|
||||||
if workspace_mode is not None:
|
if workspace_mode is not None:
|
||||||
@@ -212,8 +194,6 @@ def onboard(
|
|||||||
_CONFIGURE_SECTIONS = {
|
_CONFIGURE_SECTIONS = {
|
||||||
"ui": "UI backend",
|
"ui": "UI backend",
|
||||||
"port": "LangGraph server port",
|
"port": "LangGraph server port",
|
||||||
"provider": "LLM provider + auth + API key",
|
|
||||||
"model": "Model + reasoning effort",
|
|
||||||
"tavily": "Tavily search key",
|
"tavily": "Tavily search key",
|
||||||
"workspace": "Workspace mode",
|
"workspace": "Workspace mode",
|
||||||
"thinking": "Thinking panel",
|
"thinking": "Thinking panel",
|
||||||
@@ -263,29 +243,6 @@ def configure_port():
|
|||||||
_configure_section("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")
|
@configure_app.command("tavily")
|
||||||
def configure_tavily(
|
def configure_tavily(
|
||||||
skip_validation: bool = typer.Option(False, "--skip-validation"),
|
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.
|
"""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
|
swaps back into ``runtime_state`` so the outer poll loop picks up
|
||||||
the new handles on subsequent messages. Also keeps
|
the new handles on subsequent messages. Also keeps
|
||||||
``channel_runtime`` in sync so the bus sees the new values.
|
``channel_runtime`` in sync so the bus sees the new values.
|
||||||
@@ -1389,6 +1346,47 @@ def _serve_drain_notifications(
|
|||||||
_notif_loop.close()
|
_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()
|
@app.command()
|
||||||
def serve(
|
def serve(
|
||||||
no_thinking: bool = typer.Option(
|
no_thinking: bool = typer.Option(
|
||||||
@@ -1452,20 +1450,7 @@ def serve(
|
|||||||
if debug:
|
if debug:
|
||||||
_configure_logging()
|
_configure_logging()
|
||||||
|
|
||||||
# Auto-start ccproxy if any provider uses OAuth mode
|
_startup_gates()
|
||||||
_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
|
|
||||||
|
|
||||||
if not config.channel_enabled:
|
if not config.channel_enabled:
|
||||||
console.print("[red]No channels configured.[/red]")
|
console.print("[red]No channels configured.[/red]")
|
||||||
@@ -1565,6 +1550,9 @@ def serve(
|
|||||||
_orig_sigterm = signal.signal(signal.SIGTERM, _handle_shutdown)
|
_orig_sigterm = signal.signal(signal.SIGTERM, _handle_shutdown)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
from .agent import current_model_label
|
||||||
|
|
||||||
|
model_label = current_model_label()
|
||||||
while not shutdown_event.is_set():
|
while not shutdown_event.is_set():
|
||||||
try:
|
try:
|
||||||
msg = _message_queue.get(timeout=0.5)
|
msg = _message_queue.get(timeout=0.5)
|
||||||
@@ -1577,7 +1565,7 @@ def serve(
|
|||||||
_serve_process_message(
|
_serve_process_message(
|
||||||
msg,
|
msg,
|
||||||
runtime_state=runtime_state,
|
runtime_state=runtime_state,
|
||||||
model=config.model,
|
model=model_label,
|
||||||
workspace_dir=ws,
|
workspace_dir=ws,
|
||||||
show_thinking=effective_channel_thinking,
|
show_thinking=effective_channel_thinking,
|
||||||
on_cmd_completed=_serve_on_cmd_completed,
|
on_cmd_completed=_serve_on_cmd_completed,
|
||||||
@@ -1593,7 +1581,7 @@ def serve(
|
|||||||
if async_notifier.has_pending_notifications(runtime_state.thread_id):
|
if async_notifier.has_pending_notifications(runtime_state.thread_id):
|
||||||
_serve_drain_notifications(
|
_serve_drain_notifications(
|
||||||
runtime_state=runtime_state,
|
runtime_state=runtime_state,
|
||||||
model=config.model,
|
model=model_label,
|
||||||
workspace_dir=ws,
|
workspace_dir=ws,
|
||||||
show_thinking=effective_channel_thinking,
|
show_thinking=effective_channel_thinking,
|
||||||
)
|
)
|
||||||
@@ -2084,11 +2072,6 @@ def _main_callback(
|
|||||||
"--dangerous",
|
"--dangerous",
|
||||||
help="DANGEROUS: real-filesystem access (no workspace confinement); implies --auto-approve",
|
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(
|
ui: str | None = typer.Option(
|
||||||
None,
|
None,
|
||||||
"--ui",
|
"--ui",
|
||||||
@@ -2168,29 +2151,11 @@ def _main_callback(
|
|||||||
cli_overrides["enable_ask_user"] = True
|
cli_overrides["enable_ask_user"] = True
|
||||||
if dangerous:
|
if dangerous:
|
||||||
cli_overrides["dangerous_mode"] = True
|
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)
|
config = get_effective_config(cli_overrides)
|
||||||
apply_config_to_env(config)
|
apply_config_to_env(config)
|
||||||
|
|
||||||
# Auto-start ccproxy if any provider uses OAuth mode
|
_startup_gates()
|
||||||
_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
|
|
||||||
|
|
||||||
show_thinking = config.show_thinking if not no_thinking else False
|
show_thinking = config.show_thinking if not no_thinking else False
|
||||||
effective_channel_thinking = config.channel_send_thinking and (not no_thinking)
|
effective_channel_thinking = config.channel_send_thinking and (not no_thinking)
|
||||||
@@ -2356,6 +2321,9 @@ def _main_callback(
|
|||||||
config=config,
|
config=config,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
|
from .agent import current_model_label
|
||||||
|
|
||||||
|
model_label = current_model_label()
|
||||||
if effective_output_format == "stream-json":
|
if effective_output_format == "stream-json":
|
||||||
# Headless JSONL path: drive the sink through the gateway
|
# Headless JSONL path: drive the sink through the gateway
|
||||||
# directly. We are already inside the async single-shot
|
# directly. We are already inside the async single-shot
|
||||||
@@ -2364,7 +2332,7 @@ def _main_callback(
|
|||||||
request = RunRequest(
|
request = RunRequest(
|
||||||
message=prompt,
|
message=prompt,
|
||||||
thread_id=tid,
|
thread_id=tid,
|
||||||
metadata=build_metadata(workspace_dir, config.model),
|
metadata=build_metadata(workspace_dir, model_label),
|
||||||
target=GraphTarget(
|
target=GraphTarget(
|
||||||
local_graph=agent, workspace_dir=workspace_dir
|
local_graph=agent, workspace_dir=workspace_dir
|
||||||
),
|
),
|
||||||
@@ -2388,7 +2356,7 @@ def _main_callback(
|
|||||||
thread_id=tid,
|
thread_id=tid,
|
||||||
show_thinking=show_thinking,
|
show_thinking=show_thinking,
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
model=config.model,
|
model=model_label,
|
||||||
ui_backend=config.ui_backend,
|
ui_backend=config.ui_backend,
|
||||||
runtime_gateways=runtime_gateways,
|
runtime_gateways=runtime_gateways,
|
||||||
)
|
)
|
||||||
@@ -2403,6 +2371,7 @@ def _main_callback(
|
|||||||
nest_asyncio.apply()
|
nest_asyncio.apply()
|
||||||
asyncio.get_event_loop().run_until_complete(_single_shot())
|
asyncio.get_event_loop().run_until_complete(_single_shot())
|
||||||
else:
|
else:
|
||||||
|
from .agent import current_model_label
|
||||||
from .interactive import cmd_interactive
|
from .interactive import cmd_interactive
|
||||||
|
|
||||||
# Interactive mode (default) — checkpointer managed inside cmd_interactive
|
# Interactive mode (default) — checkpointer managed inside cmd_interactive
|
||||||
@@ -2412,8 +2381,7 @@ def _main_callback(
|
|||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
workspace_fixed=workspace_fixed,
|
workspace_fixed=workspace_fixed,
|
||||||
mode=effective_mode,
|
mode=effective_mode,
|
||||||
model=config.model,
|
model=current_model_label(),
|
||||||
provider=config.provider,
|
|
||||||
run_name=name,
|
run_name=name,
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
ui_backend=config.ui_backend,
|
ui_backend=config.ui_backend,
|
||||||
|
|||||||
@@ -961,7 +961,7 @@ def cmd_interactive(
|
|||||||
ctx: CommandContext, original_agent: Any, cmd: Command
|
ctx: CommandContext, original_agent: Any, cmd: Command
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Mirror the REPL adoption block at
|
"""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
|
state-mutating commands invoked via a channel actually
|
||||||
rebind the running session and keep the status bar
|
rebind the running session and keep the status bar
|
||||||
in sync."""
|
in sync."""
|
||||||
@@ -970,11 +970,10 @@ def cmd_interactive(
|
|||||||
ctx.agent is not None and ctx.agent is not original_agent
|
ctx.agent is not None and ctx.agent is not original_agent
|
||||||
)
|
)
|
||||||
if agent_swapped:
|
if agent_swapped:
|
||||||
from ..EvoScientist import _ensure_config
|
from .agent import current_model_label
|
||||||
|
|
||||||
agent_loader.adopt(ctx.agent)
|
agent_loader.adopt(ctx.agent)
|
||||||
cfg = _ensure_config()
|
model = current_model_label()
|
||||||
model = cfg.model
|
|
||||||
state["status_base_snapshot"] = make_empty_status_snapshot(
|
state["status_base_snapshot"] = make_empty_status_snapshot(
|
||||||
model
|
model
|
||||||
)
|
)
|
||||||
@@ -1326,7 +1325,7 @@ def cmd_interactive(
|
|||||||
if not state["running"]:
|
if not state["running"]:
|
||||||
break
|
break
|
||||||
|
|
||||||
# Agent swap (e.g. /model successfully built a
|
# Agent swap (a command successfully built a
|
||||||
# new agent): adopt into loader + reset status
|
# new agent): adopt into loader + reset status
|
||||||
# snapshot + sync channel runtime.
|
# snapshot + sync channel runtime.
|
||||||
agent_swapped = (
|
agent_swapped = (
|
||||||
@@ -1334,11 +1333,10 @@ def cmd_interactive(
|
|||||||
and ctx.agent is not _agent_for_ctx
|
and ctx.agent is not _agent_for_ctx
|
||||||
)
|
)
|
||||||
if agent_swapped:
|
if agent_swapped:
|
||||||
from ..EvoScientist import _ensure_config
|
from .agent import current_model_label
|
||||||
|
|
||||||
agent_loader.adopt(ctx.agent)
|
agent_loader.adopt(ctx.agent)
|
||||||
cfg = _ensure_config()
|
model = current_model_label()
|
||||||
model = cfg.model
|
|
||||||
state["status_base_snapshot"] = (
|
state["status_base_snapshot"] = (
|
||||||
make_empty_status_snapshot(model)
|
make_empty_status_snapshot(model)
|
||||||
)
|
)
|
||||||
@@ -1361,8 +1359,8 @@ def cmd_interactive(
|
|||||||
|
|
||||||
# Commands that mutate status fields need an
|
# Commands that mutate status fields need an
|
||||||
# async refresh here (/compact + /new use sync
|
# async refresh here (/compact + /new use sync
|
||||||
# callbacks; /model swaps the agent). /resume
|
# callbacks; an agent swap rebuilds the snapshot).
|
||||||
# awaits its own refresh inline inside the
|
# /resume awaits its own refresh inline inside the
|
||||||
# async callback.
|
# async callback.
|
||||||
if agent_swapped or _cmd.name in ("/compact", "/new"):
|
if agent_swapped or _cmd.name in ("/compact", "/new"):
|
||||||
await _refresh_status_snapshot(
|
await _refresh_status_snapshot(
|
||||||
|
|||||||
@@ -19,7 +19,6 @@ from collections.abc import Awaitable, Callable
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
from rich.table import Table
|
|
||||||
|
|
||||||
from ..commands.base import CommandUI
|
from ..commands.base import CommandUI
|
||||||
|
|
||||||
@@ -73,40 +72,6 @@ class RichCLICommandUI(CommandUI):
|
|||||||
# Rich console flushes synchronously; nothing to await.
|
# Rich console flushes synchronously; nothing to await.
|
||||||
return
|
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 ────────────────────────────────
|
# ── Interactive pickers ────────────────────────────────
|
||||||
|
|
||||||
async def wait_for_thread_pick(
|
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."""
|
"""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
|
agent_swapped = ctx.agent is not None and ctx.agent is not original_agent
|
||||||
if agent_swapped:
|
if agent_swapped:
|
||||||
from ..EvoScientist import _ensure_config
|
from .agent import current_model_label
|
||||||
|
|
||||||
app._agent_loader.adopt(ctx.agent)
|
app._agent_loader.adopt(ctx.agent)
|
||||||
cfg = _ensure_config()
|
|
||||||
update_model = getattr(app, "update_status_after_model_change", None)
|
update_model = getattr(app, "update_status_after_model_change", None)
|
||||||
if callable(update_model):
|
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
|
# Rebind the runtime whenever the agent OR thread_id may have moved
|
||||||
# — ``/new`` and ``/resume`` rotate ``app._conversation_tid``
|
# — ``/new`` and ``/resume`` rotate ``app._conversation_tid``
|
||||||
@@ -456,7 +455,6 @@ def run_textual_interactive(
|
|||||||
self._picker_future: asyncio.Future | None = None
|
self._picker_future: asyncio.Future | None = None
|
||||||
self._browser_future: asyncio.Future | None = None
|
self._browser_future: asyncio.Future | None = None
|
||||||
self._mcp_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_suggester = HistorySuggester(DATA_DIR / "history")
|
||||||
self._history_index: int = -1 # -1 = not browsing history
|
self._history_index: int = -1 # -1 = not browsing history
|
||||||
self._history_saved_input: str = "" # saved current input before browsing
|
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)
|
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:
|
def clear_chat(self) -> None:
|
||||||
container = self.query_one("#chat", VerticalScroll)
|
container = self.query_one("#chat", VerticalScroll)
|
||||||
welcome = self.query_one("#welcome", Static)
|
welcome = self.query_one("#welcome", Static)
|
||||||
@@ -794,12 +772,6 @@ def run_textual_interactive(
|
|||||||
yield Static("", id="status")
|
yield Static("", id="status")
|
||||||
|
|
||||||
def on_mount(self) -> None:
|
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_welcome()
|
||||||
self._render_status()
|
self._render_status()
|
||||||
self.set_interval(1.0, 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():
|
if self._mcp_browser_future and not self._mcp_browser_future.done():
|
||||||
self._mcp_browser_future.set_result(None)
|
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 ─────────────────────────────────────
|
# ── Streaming core ─────────────────────────────────────
|
||||||
|
|
||||||
async def _stream_with_widgets(
|
async def _stream_with_widgets(
|
||||||
@@ -2495,7 +2439,6 @@ def run_textual_interactive(
|
|||||||
if focused is not None:
|
if focused is not None:
|
||||||
from .widgets.approval_widget import ApprovalWidget
|
from .widgets.approval_widget import ApprovalWidget
|
||||||
from .widgets.mcp_browser import MCPBrowserWidget
|
from .widgets.mcp_browser import MCPBrowserWidget
|
||||||
from .widgets.model_picker import ModelPickerWidget
|
|
||||||
from .widgets.skill_browser import SkillBrowserWidget
|
from .widgets.skill_browser import SkillBrowserWidget
|
||||||
from .widgets.thread_selector import ThreadPickerWidget
|
from .widgets.thread_selector import ThreadPickerWidget
|
||||||
|
|
||||||
@@ -2511,33 +2454,6 @@ def run_textual_interactive(
|
|||||||
if isinstance(focused, MCPBrowserWidget):
|
if isinstance(focused, MCPBrowserWidget):
|
||||||
focused.action_cancel()
|
focused.action_cancel()
|
||||||
return
|
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:
|
if self._queued_messages:
|
||||||
self._queued_messages.pop()
|
self._queued_messages.pop()
|
||||||
self._render_queue_indicator()
|
self._render_queue_indicator()
|
||||||
@@ -2561,7 +2477,6 @@ def run_textual_interactive(
|
|||||||
from .widgets.approval_widget import ApprovalWidget
|
from .widgets.approval_widget import ApprovalWidget
|
||||||
from .widgets.ask_user_widget import AskUserWidget
|
from .widgets.ask_user_widget import AskUserWidget
|
||||||
from .widgets.mcp_browser import MCPBrowserWidget
|
from .widgets.mcp_browser import MCPBrowserWidget
|
||||||
from .widgets.model_picker import ModelPickerWidget
|
|
||||||
from .widgets.skill_browser import SkillBrowserWidget
|
from .widgets.skill_browser import SkillBrowserWidget
|
||||||
from .widgets.thread_selector import ThreadPickerWidget
|
from .widgets.thread_selector import ThreadPickerWidget
|
||||||
|
|
||||||
@@ -2580,20 +2495,6 @@ def run_textual_interactive(
|
|||||||
if isinstance(focused, MCPBrowserWidget):
|
if isinstance(focused, MCPBrowserWidget):
|
||||||
focused.action_move_up()
|
focused.action_move_up()
|
||||||
return
|
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:
|
if self._queued_messages:
|
||||||
last = self._queued_messages.pop()
|
last = self._queued_messages.pop()
|
||||||
prompt = self.query_one("#prompt", ChatTextArea)
|
prompt = self.query_one("#prompt", ChatTextArea)
|
||||||
@@ -2629,7 +2530,6 @@ def run_textual_interactive(
|
|||||||
from .widgets.approval_widget import ApprovalWidget
|
from .widgets.approval_widget import ApprovalWidget
|
||||||
from .widgets.ask_user_widget import AskUserWidget
|
from .widgets.ask_user_widget import AskUserWidget
|
||||||
from .widgets.mcp_browser import MCPBrowserWidget
|
from .widgets.mcp_browser import MCPBrowserWidget
|
||||||
from .widgets.model_picker import ModelPickerWidget
|
|
||||||
from .widgets.skill_browser import SkillBrowserWidget
|
from .widgets.skill_browser import SkillBrowserWidget
|
||||||
from .widgets.thread_selector import ThreadPickerWidget
|
from .widgets.thread_selector import ThreadPickerWidget
|
||||||
|
|
||||||
@@ -2648,18 +2548,6 @@ def run_textual_interactive(
|
|||||||
if isinstance(focused, MCPBrowserWidget):
|
if isinstance(focused, MCPBrowserWidget):
|
||||||
focused.action_move_down()
|
focused.action_move_down()
|
||||||
return
|
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)
|
# History browsing (down key)
|
||||||
if self._history_index >= 0:
|
if self._history_index >= 0:
|
||||||
@@ -2970,9 +2858,6 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
def _do_exit(self) -> None:
|
def _do_exit(self) -> None:
|
||||||
"""Clean up channels, unregister callbacks, and exit."""
|
"""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:
|
if self._channel_timer is not None:
|
||||||
self._channel_timer.stop()
|
self._channel_timer.stop()
|
||||||
self._channel_timer = None
|
self._channel_timer = None
|
||||||
@@ -3078,7 +2963,7 @@ def run_textual_interactive(
|
|||||||
def update_status_after_model_change(
|
def update_status_after_model_change(
|
||||||
self, new_model: str, new_provider: str | None = None
|
self, new_model: str, new_provider: str | None = 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
|
self._current_model = new_model
|
||||||
if new_provider is not None:
|
if new_provider is not None:
|
||||||
self._current_provider = new_provider
|
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"
|
EMPTY = "empty"
|
||||||
|
|
||||||
|
|
||||||
_CATEGORY_ORDER = ["Session", "Skills", "MCP", "Channels", "Model", "General"]
|
_CATEGORY_ORDER = ["Session", "Skills", "MCP", "Channels", "General"]
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
|
|||||||
@@ -47,12 +47,6 @@ class CommandUI(Protocol):
|
|||||||
async def wait_for_mcp_browse(
|
async def wait_for_mcp_browse(
|
||||||
self, servers: list, installed_names: set[str], pre_filter_tag: str
|
self, servers: list, installed_names: set[str], pre_filter_tag: str
|
||||||
) -> list | None: ...
|
) -> 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 clear_chat(self) -> None: ...
|
||||||
def request_quit(self) -> None: ...
|
def request_quit(self) -> None: ...
|
||||||
def force_quit(self) -> None: ...
|
def force_quit(self) -> None: ...
|
||||||
|
|||||||
@@ -5,8 +5,6 @@ from . import (
|
|||||||
channel,
|
channel,
|
||||||
general,
|
general,
|
||||||
mcp,
|
mcp,
|
||||||
model,
|
|
||||||
model_fallback,
|
|
||||||
schedule,
|
schedule,
|
||||||
session,
|
session,
|
||||||
skills,
|
skills,
|
||||||
@@ -17,8 +15,6 @@ __all__ = [
|
|||||||
"channel",
|
"channel",
|
||||||
"general",
|
"general",
|
||||||
"mcp",
|
"mcp",
|
||||||
"model",
|
|
||||||
"model_fallback",
|
|
||||||
"schedule",
|
"schedule",
|
||||||
"session",
|
"session",
|
||||||
"skills",
|
"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_dir,
|
||||||
get_config_path,
|
get_config_path,
|
||||||
get_config_value,
|
get_config_value,
|
||||||
|
get_default_workspace_dir,
|
||||||
get_effective_config,
|
get_effective_config,
|
||||||
|
is_config_applied_env,
|
||||||
list_config,
|
list_config,
|
||||||
load_config,
|
load_config,
|
||||||
reset_config,
|
reset_config,
|
||||||
@@ -39,7 +41,9 @@ __all__ = [
|
|||||||
"get_config_dir",
|
"get_config_dir",
|
||||||
"get_config_path",
|
"get_config_path",
|
||||||
"get_config_value",
|
"get_config_value",
|
||||||
|
"get_default_workspace_dir",
|
||||||
"get_effective_config",
|
"get_effective_config",
|
||||||
|
"is_config_applied_env",
|
||||||
"list_config",
|
"list_config",
|
||||||
"load_config",
|
"load_config",
|
||||||
"reset_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``
|
``STEPS``, ``render_progress``
|
||||||
- :mod:`EvoScientist.config.onboard.steps` — per-step functions
|
- :mod:`EvoScientist.config.onboard.steps` — per-step functions
|
||||||
- :mod:`EvoScientist.config.onboard.channels` — channel selection + setup
|
- :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
|
npx/node, LaTeX, iMessage helpers
|
||||||
- :mod:`EvoScientist.config.onboard.style` — Rich styles + ``_checkbox_ask``
|
- :mod:`EvoScientist.config.onboard.style` — Rich styles + ``_checkbox_ask``
|
||||||
- :mod:`EvoScientist.config.onboard.validators` — input validators
|
- :mod:`EvoScientist.config.onboard.validators` — input validators
|
||||||
- :mod:`EvoScientist.config.onboard.prompter` — ``NonInteractivePrompter``
|
- :mod:`EvoScientist.config.onboard.prompter` — ``NonInteractivePrompter``
|
||||||
(CLI-answer container) + ``select_navigation_active`` / ``GoBack`` for
|
(CLI-answer container) + ``select_navigation_active`` for keyboard nav
|
||||||
keyboard nav
|
|
||||||
- :mod:`EvoScientist.config.onboard.constants` — canonical valid-value sets
|
- :mod:`EvoScientist.config.onboard.constants` — canonical valid-value sets
|
||||||
|
|
||||||
This module used to re-export every symbol from every submodule for
|
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``
|
The interactive ``_step_*`` functions in ``steps.py`` use these for ``Choice``
|
||||||
construction (or are checked against them by tests). The CLI ``onboard``
|
construction (or are checked against them by tests). The CLI ``onboard``
|
||||||
command in ``cli/commands.py`` uses them to validate ``--provider`` /
|
command in ``cli/commands.py`` uses them to validate ``--ui`` /
|
||||||
``--ui`` / ``--workspace-mode`` flag inputs.
|
``--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
|
``Choice(value=...)`` in ``steps.py`` is required; a drift test in
|
||||||
``tests/test_onboard.py`` keeps both sides in sync.
|
``tests/test_onboard.py`` keeps both sides in sync.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
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_UI_BACKENDS: frozenset[str] = frozenset({"tui", "cli", "webui"})
|
||||||
|
|
||||||
VALID_WORKSPACE_MODES: frozenset[str] = frozenset({"daemon", "run"})
|
VALID_WORKSPACE_MODES: frozenset[str] = frozenset({"daemon", "run"})
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"VALID_PROVIDERS",
|
|
||||||
"VALID_UI_BACKENDS",
|
"VALID_UI_BACKENDS",
|
||||||
"VALID_WORKSPACE_MODES",
|
"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.
|
LaTeX detection/install, iMessage setup.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -11,126 +11,7 @@ import sys
|
|||||||
|
|
||||||
import questionary
|
import questionary
|
||||||
|
|
||||||
from ..settings import EvoScientistConfig
|
|
||||||
from .style import QMARK, WIZARD_STYLE, console
|
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(
|
def _prompt_and_validate_api_key(
|
||||||
@@ -192,66 +73,6 @@ def _prompt_and_validate_api_key(
|
|||||||
return new_key or None
|
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:
|
def _check_npx() -> bool:
|
||||||
"""Check if npx is available on the system.
|
"""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}"
|
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:
|
def _install_imsg() -> bool:
|
||||||
"""Run brew install for imsg CLI.
|
"""Run brew install for imsg CLI.
|
||||||
|
|
||||||
|
|||||||
@@ -5,31 +5,13 @@ from __future__ import annotations
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
class GoBack(Exception):
|
def install_navigation_keys(question) -> None:
|
||||||
"""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:
|
|
||||||
"""Add keyboard shortcuts on a questionary select ``Question``.
|
"""Add keyboard shortcuts on a questionary select ``Question``.
|
||||||
|
|
||||||
Bindings (merged in front of questionary's defaults — Ctrl+C/Ctrl+D still
|
Bindings (merged in front of questionary's defaults — Ctrl+C/Ctrl+D still
|
||||||
cancel the wizard):
|
cancel the wizard):
|
||||||
|
|
||||||
- ``→`` — accept the option under the cursor and advance (mirrors Enter).
|
- ``→`` — 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
|
from prompt_toolkit.key_binding import KeyBindings, merge_key_bindings
|
||||||
|
|
||||||
@@ -49,13 +31,6 @@ def install_navigation_keys(
|
|||||||
event.app.exit(result=pointed.value)
|
event.app.exit(result=pointed.value)
|
||||||
return
|
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(
|
question.application.key_bindings = merge_key_bindings(
|
||||||
[kb, question.application.key_bindings]
|
[kb, question.application.key_bindings]
|
||||||
)
|
)
|
||||||
@@ -78,7 +53,7 @@ def select_navigation_active():
|
|||||||
def _wrapped(*args, **kwargs):
|
def _wrapped(*args, **kwargs):
|
||||||
q = original(*args, **kwargs)
|
q = original(*args, **kwargs)
|
||||||
try:
|
try:
|
||||||
install_navigation_keys(q, with_back=False)
|
install_navigation_keys(q)
|
||||||
except Exception:
|
except Exception:
|
||||||
# Don't let a stray keybinding error block the wizard.
|
# Don't let a stray keybinding error block the wizard.
|
||||||
pass
|
pass
|
||||||
@@ -92,7 +67,7 @@ def select_navigation_active():
|
|||||||
|
|
||||||
|
|
||||||
class NonInteractivePrompter:
|
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
|
``strict=True`` makes missing presets fatal instead of falling back to
|
||||||
interactive. Wizard reads ``answers`` / ``skip_set`` / ``strict`` directly.
|
interactive. Wizard reads ``answers`` / ``skip_set`` / ``strict`` directly.
|
||||||
@@ -113,8 +88,6 @@ class NonInteractivePrompter:
|
|||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"BACK_SENTINEL",
|
|
||||||
"GoBack",
|
|
||||||
"NonInteractivePrompter",
|
"NonInteractivePrompter",
|
||||||
"install_navigation_keys",
|
"install_navigation_keys",
|
||||||
"select_navigation_active",
|
"select_navigation_active",
|
||||||
|
|||||||
@@ -1,8 +1,8 @@
|
|||||||
"""Individual wizard step functions.
|
"""Individual wizard step functions.
|
||||||
|
|
||||||
Each ``_step_*`` prompts the user for one logical decision and returns the
|
Each ``_step_*`` prompts the user for one logical decision and returns the
|
||||||
chosen value. Conditional steps (auth mode, base URL) are only called by
|
chosen value. LLM provider/model/API-key configuration no longer happens in
|
||||||
``run_onboard`` when the provider needs them.
|
the wizard — models are configured via the WebUI / model registry.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -14,24 +14,17 @@ import questionary
|
|||||||
from prompt_toolkit.formatted_text import FormattedText
|
from prompt_toolkit.formatted_text import FormattedText
|
||||||
from questionary import Choice
|
from questionary import Choice
|
||||||
|
|
||||||
from ...llm import get_models_for_provider
|
|
||||||
from ...llm.ollama_discovery import validate_ollama_connection
|
|
||||||
from ..settings import EvoScientistConfig
|
from ..settings import EvoScientistConfig
|
||||||
from .helpers import (
|
from .helpers import (
|
||||||
_auto_install_latexmk,
|
_auto_install_latexmk,
|
||||||
_check_latex_components,
|
_check_latex_components,
|
||||||
_detect_tinytex_install_method,
|
_detect_tinytex_install_method,
|
||||||
_ensure_npx,
|
_ensure_npx,
|
||||||
_install_ccproxy,
|
|
||||||
_install_tinytex,
|
_install_tinytex,
|
||||||
_print_latex_status,
|
_print_latex_status,
|
||||||
_prompt_and_validate_api_key,
|
_prompt_and_validate_api_key,
|
||||||
_prompt_ccproxy_port,
|
|
||||||
_provider_key_info,
|
|
||||||
_run_ccproxy_login,
|
|
||||||
)
|
)
|
||||||
from .style import (
|
from .style import (
|
||||||
CONFIRM_STYLE,
|
|
||||||
QMARK,
|
QMARK,
|
||||||
WIZARD_STYLE,
|
WIZARD_STYLE,
|
||||||
_checkbox_ask,
|
_checkbox_ask,
|
||||||
@@ -228,633 +221,6 @@ def _step_webui_port(config: EvoScientistConfig) -> int:
|
|||||||
return port
|
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(
|
def _step_tavily_key(
|
||||||
config: EvoScientistConfig,
|
config: EvoScientistConfig,
|
||||||
skip_validation: bool = False,
|
skip_validation: bool = False,
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
"""Input validators for the onboarding wizard.
|
"""Input validators for the onboarding wizard.
|
||||||
|
|
||||||
- IntegerValidator / ChoiceValidator: prompt_toolkit Validators
|
- 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
|
from __future__ import annotations
|
||||||
@@ -97,415 +97,6 @@ def _classify_validation_error(error: BaseException) -> tuple[bool, str] | None:
|
|||||||
return 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]:
|
def validate_tavily_key(api_key: str) -> tuple[bool, str]:
|
||||||
"""Validate a Tavily API key by making a test request.
|
"""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:
|
if classified is not None:
|
||||||
return classified
|
return classified
|
||||||
return False, f"Error: {e}"
|
return False, f"Error: {e}"
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
|
||||||
# Display Helpers
|
|
||||||
# =============================================================================
|
|
||||||
|
|||||||
@@ -3,7 +3,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import copy
|
import copy
|
||||||
import os
|
|
||||||
|
|
||||||
import questionary
|
import questionary
|
||||||
from rich.panel import Panel
|
from rich.panel import Panel
|
||||||
@@ -17,18 +16,8 @@ from ..settings import (
|
|||||||
)
|
)
|
||||||
from .channels import _step_channels
|
from .channels import _step_channels
|
||||||
from .steps import (
|
from .steps import (
|
||||||
_step_anthropic_auth_mode,
|
|
||||||
_step_auxiliary_enable,
|
|
||||||
_step_base_url,
|
|
||||||
_step_langgraph_dev_port,
|
_step_langgraph_dev_port,
|
||||||
_step_mcp_servers,
|
_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_skills,
|
||||||
_step_tavily_key,
|
_step_tavily_key,
|
||||||
_step_thinking,
|
_step_thinking,
|
||||||
@@ -41,7 +30,6 @@ from .style import (
|
|||||||
CONFIRM_STYLE,
|
CONFIRM_STYLE,
|
||||||
QMARK,
|
QMARK,
|
||||||
_print_header,
|
_print_header,
|
||||||
_print_section,
|
|
||||||
_print_step_skipped,
|
_print_step_skipped,
|
||||||
console,
|
console,
|
||||||
)
|
)
|
||||||
@@ -49,10 +37,6 @@ from .style import (
|
|||||||
STEPS = [
|
STEPS = [
|
||||||
"UI",
|
"UI",
|
||||||
"LangGraph Port",
|
"LangGraph Port",
|
||||||
"Provider",
|
|
||||||
"API Key",
|
|
||||||
"Model",
|
|
||||||
"Auxiliary Model",
|
|
||||||
"Tavily Key",
|
"Tavily Key",
|
||||||
"Workspace",
|
"Workspace",
|
||||||
"Thinking",
|
"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:
|
def _autosave(config: EvoScientistConfig) -> None:
|
||||||
"""Persist current config to disk between phases.
|
"""Persist current config to disk between phases.
|
||||||
|
|
||||||
@@ -148,208 +106,10 @@ def _autosave(config: EvoScientistConfig) -> None:
|
|||||||
pass
|
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.
|
# Sections offered in Keep/Modify/Reset → which step labels they enable.
|
||||||
_SECTION_LABELS: list[tuple[str, str]] = [
|
_SECTION_LABELS: list[tuple[str, str]] = [
|
||||||
("ui", "UI backend"),
|
("ui", "UI backend"),
|
||||||
("port", "LangGraph server port"),
|
("port", "LangGraph server port"),
|
||||||
("provider", "LLM provider + auth + API key"),
|
|
||||||
("model", "Model + reasoning effort"),
|
|
||||||
("auxiliary_model", "Auxiliary model (optional)"),
|
|
||||||
("tavily", "Tavily search key"),
|
("tavily", "Tavily search key"),
|
||||||
("workspace", "Workspace mode"),
|
("workspace", "Workspace mode"),
|
||||||
("thinking", "Thinking panel"),
|
("thinking", "Thinking panel"),
|
||||||
@@ -360,17 +120,10 @@ _SECTION_LABELS: list[tuple[str, str]] = [
|
|||||||
]
|
]
|
||||||
_ALL_SECTIONS: frozenset[str] = frozenset(s for s, _ in _SECTION_LABELS)
|
_ALL_SECTIONS: frozenset[str] = frozenset(s for s, _ in _SECTION_LABELS)
|
||||||
|
|
||||||
# Each preset flag implies the section(s) it would change. ``--provider`` also
|
# Each preset flag implies the section(s) it would change.
|
||||||
# cascades into ``model`` because the model list depends on the provider —
|
|
||||||
# silently keeping a stale model id would leave the first request broken.
|
|
||||||
_FLAG_TO_SECTIONS: dict[str, frozenset[str]] = {
|
_FLAG_TO_SECTIONS: dict[str, frozenset[str]] = {
|
||||||
"ui": frozenset({"ui"}),
|
"ui": frozenset({"ui"}),
|
||||||
"port": frozenset({"port"}),
|
"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"}),
|
"tavily_key": frozenset({"tavily"}),
|
||||||
"workspace_mode": frozenset({"workspace"}),
|
"workspace_mode": frozenset({"workspace"}),
|
||||||
"show_thinking": frozenset({"thinking"}),
|
"show_thinking": frozenset({"thinking"}),
|
||||||
@@ -481,7 +234,7 @@ def run_onboard(
|
|||||||
Args:
|
Args:
|
||||||
skip_validation: Skip API key validation.
|
skip_validation: Skip API key validation.
|
||||||
prompter: Optional :class:`NonInteractivePrompter` carrying
|
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
|
``skip_set`` (sections to bypass). When None, all prompts
|
||||||
fall through to the interactive questionary form.
|
fall through to the interactive questionary form.
|
||||||
only_sections: If given, restrict the wizard to exactly these section
|
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.
|
# 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.
|
# 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
|
# explicit user intent — skip Keep/Modify/Reset and run ONLY
|
||||||
# the sections each flag implies (see ``_FLAG_TO_SECTIONS``).
|
# the sections each flag implies (see ``_FLAG_TO_SECTIONS``).
|
||||||
# - Strict ``--non-interactive`` with no preset flags: run all
|
# - 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)
|
config.langgraph_dev_port = _step_langgraph_dev_port(config)
|
||||||
_autosave(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:
|
if "tavily" in sections_to_run:
|
||||||
preset_tavily = _preset("tavily_key")
|
preset_tavily = _preset("tavily_key")
|
||||||
if preset_tavily is not None:
|
if preset_tavily is not None:
|
||||||
|
|||||||
+95
-210
@@ -23,6 +23,19 @@ from dotenv import find_dotenv, load_dotenv
|
|||||||
# (stream/display.py, channels/consumer.py) — keep aligned with the agent's
|
# (stream/display.py, channels/consumer.py) — keep aligned with the agent's
|
||||||
# `interrupt_on` set in EvoScientist.py.
|
# `interrupt_on` set in EvoScientist.py.
|
||||||
HITL_SHELL_TOOLS = ("execute", "run_in_background")
|
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):
|
class MemoryObservationTarget(StrEnum):
|
||||||
@@ -106,23 +119,11 @@ def _normalize_hhmm(value: Any) -> str | None:
|
|||||||
def get_config_dir() -> Path:
|
def get_config_dir() -> Path:
|
||||||
"""Get the configuration directory path.
|
"""Get the configuration directory path.
|
||||||
|
|
||||||
Priority:
|
Uses XDG_CONFIG_HOME if set, otherwise ~/.config/evoscientist/
|
||||||
1. EVOSCIENTIST_CONFIG_DIR
|
|
||||||
2. EVOSCIENTIST_HOME/config
|
|
||||||
3. XDG_CONFIG_HOME/evoscientist
|
|
||||||
4. ~/.config/evoscientist
|
|
||||||
"""
|
"""
|
||||||
configured = os.environ.get("EVOSCIENTIST_CONFIG_DIR")
|
|
||||||
if configured:
|
|
||||||
return Path(configured).expanduser().resolve()
|
|
||||||
|
|
||||||
home = os.environ.get("EVOSCIENTIST_HOME")
|
|
||||||
if home:
|
|
||||||
return Path(home).expanduser().resolve() / "config"
|
|
||||||
|
|
||||||
xdg_config = os.environ.get("XDG_CONFIG_HOME")
|
xdg_config = os.environ.get("XDG_CONFIG_HOME")
|
||||||
if xdg_config:
|
if xdg_config:
|
||||||
return Path(xdg_config).expanduser() / "evoscientist"
|
return Path(xdg_config) / "evoscientist"
|
||||||
return Path.home() / ".config" / "evoscientist"
|
return Path.home() / ".config" / "evoscientist"
|
||||||
|
|
||||||
|
|
||||||
@@ -131,74 +132,35 @@ def get_config_path() -> Path:
|
|||||||
return get_config_dir() / "config.yaml"
|
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
|
# Configuration dataclass
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
# OpenRouter app-attribution defaults (issue #339). Single source of truth: the
|
|
||||||
# EvoScientistConfig fields below default to these, and llm/models.py imports
|
|
||||||
# them for its env-fallback, so the values never drift across the two layers.
|
|
||||||
OPENROUTER_DEFAULT_HTTP_REFERER = "https://github.com/EvoScientist/EvoScientist"
|
|
||||||
OPENROUTER_DEFAULT_APP_TITLE = "EvoScientist"
|
|
||||||
# OpenRouter honors only the first 2 categories per request (server-side limit)
|
|
||||||
# and silently ignores the rest, so keep the two most relevant ones. Chosen per
|
|
||||||
# maintainer review — creative-writing is a less competitive marketplace group.
|
|
||||||
OPENROUTER_DEFAULT_APP_CATEGORIES = "creative-writing,personal-agent"
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class EvoScientistConfig:
|
class EvoScientistConfig:
|
||||||
"""EvoScientist configuration settings.
|
"""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:
|
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.
|
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_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.
|
show_thinking: Whether to show thinking panels in CLI.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# API Keys
|
# API Keys (non-LLM tools only)
|
||||||
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 = ""
|
|
||||||
tavily_api_key: str = ""
|
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
|
# Async Sub-agent Settings
|
||||||
# When True (default), the EvoSci CLI auto-starts a langgraph dev subprocess
|
# When True (default), the EvoSci CLI auto-starts a langgraph dev subprocess
|
||||||
# so any sub-agent flagged ``async: true`` in subagents/<name>.yaml runs
|
# so any sub-agent flagged ``async: true`` in subagents/<name>.yaml runs
|
||||||
@@ -262,13 +224,6 @@ class EvoScientistConfig:
|
|||||||
# Lower (e.g., 5000) if you want a tighter safety net against runaway loops.
|
# Lower (e.g., 5000) if you want a tighter safety net against runaway loops.
|
||||||
recursion_limit: int = 1_000_000
|
recursion_limit: int = 1_000_000
|
||||||
|
|
||||||
# Number of consecutive model rounds with the same structured tool name and
|
|
||||||
# arguments that activates provider-facing loop repair. Set 0 to disable.
|
|
||||||
repetitive_tool_call_threshold: int = 2
|
|
||||||
# Number of consecutive deterministic tool errors allowed before the next
|
|
||||||
# model call is blocked. Transient provider/network errors are not counted.
|
|
||||||
max_consecutive_tool_errors: int = 3
|
|
||||||
|
|
||||||
# Memory Settings
|
# Memory Settings
|
||||||
# Profile memory injects and maintains `/memories/profile/...` files.
|
# Profile memory injects and maintains `/memories/profile/...` files.
|
||||||
memory_profile_enabled: bool = True
|
memory_profile_enabled: bool = True
|
||||||
@@ -307,21 +262,7 @@ class EvoScientistConfig:
|
|||||||
# a deploy-style langgraph server instead of the in-terminal CLI/TUI.
|
# a deploy-style langgraph server instead of the in-terminal CLI/TUI.
|
||||||
ui_backend: Literal["cli", "tui", "webui"] = "tui"
|
ui_backend: Literal["cli", "tui", "webui"] = "tui"
|
||||||
log_level: str = "warning"
|
log_level: str = "warning"
|
||||||
# Empty means use the provider/model default. A non-empty value is an
|
reasoning_effort: str = "high"
|
||||||
# explicit user override exported as EVOSCIENTIST_REASONING_EFFORT.
|
|
||||||
reasoning_effort: str = ""
|
|
||||||
# Anthropic prompt caching for OpenRouter anthropic/* models. Opt out if
|
|
||||||
# cache-write costs outweigh the benefit for a workflow.
|
|
||||||
openrouter_anthropic_prompt_cache: bool = True
|
|
||||||
# OpenRouter app attribution (issue #339). Sent only for the openrouter
|
|
||||||
# provider; identifies EvoScientist in OpenRouter's app rankings/analytics.
|
|
||||||
# Override (e.g. a private fork) via these fields or their env vars.
|
|
||||||
# Defaults live in the module constants above (also imported by llm/models.py).
|
|
||||||
openrouter_http_referer: str = OPENROUTER_DEFAULT_HTTP_REFERER
|
|
||||||
openrouter_app_title: str = OPENROUTER_DEFAULT_APP_TITLE
|
|
||||||
# Comma-separated; split into a list before being passed to
|
|
||||||
# langchain-openrouter (its app_categories kwarg expects list[str]).
|
|
||||||
openrouter_app_categories: str = OPENROUTER_DEFAULT_APP_CATEGORIES
|
|
||||||
|
|
||||||
# Channel Settings
|
# Channel Settings
|
||||||
channel_enabled: str = "" # "imessage" | "telegram" | "discord" | "slack" | "wechat" | "dingtalk" | "feishu" | "email" | "qq" | "signal" | "" (comma-separated for multiple)
|
channel_enabled: str = "" # "imessage" | "telegram" | "discord" | "slack" | "wechat" | "dingtalk" | "feishu" | "email" | "qq" | "signal" | "" (comma-separated for multiple)
|
||||||
@@ -438,6 +379,17 @@ class EvoScientistConfig:
|
|||||||
# blocklist (sudo/chmod/dd/...) still applies. Implies auto_approve.
|
# blocklist (sudo/chmod/dd/...) still applies. Implies auto_approve.
|
||||||
dangerous_mode: bool = False
|
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
|
# Agent features
|
||||||
enable_ask_user: bool = True # Enable ask_user tool for agent-initiated questions
|
enable_ask_user: bool = True # Enable ask_user tool for agent-initiated questions
|
||||||
|
|
||||||
@@ -466,9 +418,6 @@ class EvoScientistConfig:
|
|||||||
# DM access control policy
|
# DM access control policy
|
||||||
dm_policy: str = "allowlist"
|
dm_policy: str = "allowlist"
|
||||||
|
|
||||||
# OpenAI API mode - "" = auto, "true" = force Responses, "false" = force Completions
|
|
||||||
use_responses_api: str = ""
|
|
||||||
|
|
||||||
# ccproxy
|
# ccproxy
|
||||||
ccproxy_port: int = 8000
|
ccproxy_port: int = 8000
|
||||||
|
|
||||||
@@ -480,14 +429,6 @@ class EvoScientistConfig:
|
|||||||
stt_compute_type: str = "int8" # "int8" | "float16" | "float32"
|
stt_compute_type: str = "int8" # "int8" | "float16" | "float32"
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
for field_name in (
|
|
||||||
"repetitive_tool_call_threshold",
|
|
||||||
"max_consecutive_tool_errors",
|
|
||||||
):
|
|
||||||
value = getattr(self, field_name)
|
|
||||||
if not isinstance(value, int) or isinstance(value, bool) or value < 0:
|
|
||||||
raise ValueError(f"{field_name} must be a non-negative integer")
|
|
||||||
|
|
||||||
# A non-positive or non-int sandbox_execute_timeout (e.g. a hand-edited
|
# A non-positive or non-int sandbox_execute_timeout (e.g. a hand-edited
|
||||||
# config file value — load_config does not coerce file values — or a
|
# config file value — load_config does not coerce file values — or a
|
||||||
# 0/negative env value) would raise inside CustomSandboxBackend.__init__
|
# 0/negative env value) would raise inside CustomSandboxBackend.__init__
|
||||||
@@ -506,6 +447,51 @@ class EvoScientistConfig:
|
|||||||
if self.dangerous_mode:
|
if self.dangerous_mode:
|
||||||
self.auto_approve = True
|
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)
|
_normalize_str_enum_fields(self)
|
||||||
|
|
||||||
synthesis_time = _normalize_hhmm(self.memory_skill_synthesis_time)
|
synthesis_time = _normalize_hhmm(self.memory_skill_synthesis_time)
|
||||||
@@ -750,11 +736,6 @@ def set_config_value(key: str, value: Any) -> bool:
|
|||||||
|
|
||||||
if key == "sandbox_execute_timeout" and value <= 0:
|
if key == "sandbox_execute_timeout" and value <= 0:
|
||||||
return False
|
return False
|
||||||
if key in {
|
|
||||||
"repetitive_tool_call_threshold",
|
|
||||||
"max_consecutive_tool_errors",
|
|
||||||
} and (isinstance(value, bool) or value < 0):
|
|
||||||
return False
|
|
||||||
if key == "memory_skill_synthesis_time":
|
if key == "memory_skill_synthesis_time":
|
||||||
value = _normalize_hhmm(value)
|
value = _normalize_hhmm(value)
|
||||||
if value is None:
|
if value is None:
|
||||||
@@ -780,47 +761,22 @@ def list_config() -> dict[str, Any]:
|
|||||||
|
|
||||||
# Environment variable mappings
|
# Environment variable mappings
|
||||||
_ENV_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",
|
"tavily_api_key": "TAVILY_API_KEY",
|
||||||
"default_mode": "EVOSCIENTIST_DEFAULT_MODE",
|
"default_mode": "EVOSCIENTIST_DEFAULT_MODE",
|
||||||
"default_workdir": "EVOSCIENTIST_WORKSPACE_DIR",
|
"default_workdir": "EVOSCIENTIST_WORKSPACE_DIR",
|
||||||
"ui_backend": "EVOSCIENTIST_UI_BACKEND",
|
"ui_backend": "EVOSCIENTIST_UI_BACKEND",
|
||||||
"log_level": "EVOSCIENTIST_LOG_LEVEL",
|
"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",
|
"reasoning_effort": "EVOSCIENTIST_REASONING_EFFORT",
|
||||||
"openrouter_anthropic_prompt_cache": (
|
|
||||||
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"
|
|
||||||
),
|
|
||||||
"openrouter_http_referer": "EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
|
|
||||||
"openrouter_app_title": "EVOSCIENTIST_OPENROUTER_APP_TITLE",
|
|
||||||
"openrouter_app_categories": "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
|
|
||||||
"dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE",
|
"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",
|
"channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING",
|
||||||
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
|
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
|
||||||
"use_responses_api": "EVOSCIENTIST_USE_RESPONSES_API",
|
|
||||||
"checkpoint_keep_per_thread": "EVOSCIENTIST_CHECKPOINT_KEEP_PER_THREAD",
|
"checkpoint_keep_per_thread": "EVOSCIENTIST_CHECKPOINT_KEEP_PER_THREAD",
|
||||||
"enable_async_subagents": "EVOSCIENTIST_ENABLE_ASYNC_SUBAGENTS",
|
"enable_async_subagents": "EVOSCIENTIST_ENABLE_ASYNC_SUBAGENTS",
|
||||||
"langgraph_dev_port": "EVOSCIENTIST_LANGGRAPH_DEV_PORT",
|
"langgraph_dev_port": "EVOSCIENTIST_LANGGRAPH_DEV_PORT",
|
||||||
@@ -833,10 +789,6 @@ _ENV_MAPPINGS = {
|
|||||||
"langgraph_dev_file_persistence": "EVOSCIENTIST_LANGGRAPH_DEV_FILE_PERSISTENCE",
|
"langgraph_dev_file_persistence": "EVOSCIENTIST_LANGGRAPH_DEV_FILE_PERSISTENCE",
|
||||||
"langgraph_dev_jobs_per_worker": "EVOSCIENTIST_LANGGRAPH_DEV_JOBS_PER_WORKER",
|
"langgraph_dev_jobs_per_worker": "EVOSCIENTIST_LANGGRAPH_DEV_JOBS_PER_WORKER",
|
||||||
"recursion_limit": "EVOSCIENTIST_RECURSION_LIMIT",
|
"recursion_limit": "EVOSCIENTIST_RECURSION_LIMIT",
|
||||||
"repetitive_tool_call_threshold": (
|
|
||||||
"EVOSCIENTIST_REPETITIVE_TOOL_CALL_THRESHOLD"
|
|
||||||
),
|
|
||||||
"max_consecutive_tool_errors": "EVOSCIENTIST_MAX_CONSECUTIVE_TOOL_ERRORS",
|
|
||||||
"memory_profile_enabled": "EVOSCIENTIST_MEMORY_PROFILE_ENABLED",
|
"memory_profile_enabled": "EVOSCIENTIST_MEMORY_PROFILE_ENABLED",
|
||||||
"memory_observations_enabled": "EVOSCIENTIST_MEMORY_OBSERVATIONS_ENABLED",
|
"memory_observations_enabled": "EVOSCIENTIST_MEMORY_OBSERVATIONS_ENABLED",
|
||||||
"memory_observation_writer": "EVOSCIENTIST_MEMORY_OBSERVATION_WRITER",
|
"memory_observation_writer": "EVOSCIENTIST_MEMORY_OBSERVATION_WRITER",
|
||||||
@@ -896,82 +848,19 @@ def get_effective_config(
|
|||||||
|
|
||||||
|
|
||||||
def apply_config_to_env(config: EvoScientistConfig) -> None:
|
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
|
LLM provider keys are no longer injected here — they live in the model
|
||||||
libraries (like langchain-anthropic) can pick up.
|
registry (model-runtime.sqlite3) and are applied by the model runtime.
|
||||||
|
Only non-LLM tool keys and platform round-trips remain.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Configuration to apply.
|
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"):
|
if config.tavily_api_key and not os.environ.get("TAVILY_API_KEY"):
|
||||||
os.environ["TAVILY_API_KEY"] = config.tavily_api_key
|
os.environ["TAVILY_API_KEY"] = config.tavily_api_key
|
||||||
if config.reasoning_effort and not os.environ.get("EVOSCIENTIST_REASONING_EFFORT"):
|
if config.reasoning_effort and not os.environ.get("EVOSCIENTIST_REASONING_EFFORT"):
|
||||||
os.environ["EVOSCIENTIST_REASONING_EFFORT"] = config.reasoning_effort
|
os.environ["EVOSCIENTIST_REASONING_EFFORT"] = config.reasoning_effort
|
||||||
if config.openrouter_http_referer and not os.environ.get(
|
|
||||||
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER"
|
|
||||||
):
|
|
||||||
os.environ["EVOSCIENTIST_OPENROUTER_HTTP_REFERER"] = (
|
|
||||||
config.openrouter_http_referer
|
|
||||||
)
|
|
||||||
if config.openrouter_app_title and not os.environ.get(
|
|
||||||
"EVOSCIENTIST_OPENROUTER_APP_TITLE"
|
|
||||||
):
|
|
||||||
os.environ["EVOSCIENTIST_OPENROUTER_APP_TITLE"] = config.openrouter_app_title
|
|
||||||
if config.openrouter_app_categories and not os.environ.get(
|
|
||||||
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES"
|
|
||||||
):
|
|
||||||
os.environ["EVOSCIENTIST_OPENROUTER_APP_CATEGORIES"] = (
|
|
||||||
config.openrouter_app_categories
|
|
||||||
)
|
|
||||||
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()
|
# 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
|
# (warning banner, run_in_background) and is inherited by the langgraph dev
|
||||||
# subprocess — otherwise a --dangerous CLI flag (not persisted to file/env)
|
# subprocess — otherwise a --dangerous CLI flag (not persisted to file/env)
|
||||||
@@ -982,7 +871,3 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
|
|||||||
os.environ["EVOSCIENTIST_DANGEROUS_MODE"] = "true"
|
os.environ["EVOSCIENTIST_DANGEROUS_MODE"] = "true"
|
||||||
else:
|
else:
|
||||||
os.environ.pop("EVOSCIENTIST_DANGEROUS_MODE", None)
|
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 __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph_sdk.schema import Cron, Run
|
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()))
|
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(
|
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:
|
) -> Cron:
|
||||||
"""Create a recurring scheduled task on the scheduler graph."""
|
"""Create a recurring scheduled task on the scheduler graph."""
|
||||||
# Crons are stored in the langgraph-dev process's .langgraph_api store, not
|
# Crons are stored in the langgraph-dev process's .langgraph_api store, not
|
||||||
# tagged by workspace. Isolation is process-level (see module docstring).
|
# tagged by workspace. Isolation is process-level (see module docstring).
|
||||||
return _client().crons.create(
|
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,
|
assistant_id=SCHEDULER_GRAPH_ID,
|
||||||
schedule=schedule,
|
schedule=schedule,
|
||||||
input=messages_input(prompt),
|
input=messages_input(prompt),
|
||||||
metadata={"run_kind": SCHEDULED_RUN_KIND, "name": name, "prompt": prompt},
|
metadata={
|
||||||
|
"run_kind": SCHEDULED_RUN_KIND,
|
||||||
|
"name": name,
|
||||||
|
"prompt": prompt,
|
||||||
|
**scope_metadata,
|
||||||
|
},
|
||||||
|
**({"config": config} if config else {}),
|
||||||
timezone=timezone or _default_timezone(),
|
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.
|
"""Return only EvoScientist scheduled tasks.
|
||||||
|
|
||||||
Filtered server-side by ``run_kind`` metadata (the cron backend matches by
|
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
|
rather than ``assistant_id`` because the stored ``assistant_id`` is a resolved
|
||||||
UUID, not the ``scheduler`` graph name we create with.
|
UUID, not the ``scheduler`` graph name we create with.
|
||||||
"""
|
"""
|
||||||
return _client().crons.search(
|
rows = _client().crons.search(
|
||||||
metadata={"run_kind": SCHEDULED_RUN_KIND},
|
metadata={"run_kind": SCHEDULED_RUN_KIND},
|
||||||
limit=1000,
|
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:
|
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)
|
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``).
|
"""Fire a one-off scheduler run immediately (for ``/schedule run``).
|
||||||
|
|
||||||
Output goes wherever the task's prompt specifies; there is no push notification.
|
Output goes wherever the task's prompt specifies; there is no push notification.
|
||||||
"""
|
"""
|
||||||
client = _client()
|
client = _client()
|
||||||
thread = client.threads.create(graph_id=SCHEDULER_GRAPH_ID)
|
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(
|
return client.runs.create(
|
||||||
thread_id=str(thread["thread_id"]),
|
thread_id=str(thread["thread_id"]),
|
||||||
assistant_id=SCHEDULER_GRAPH_ID,
|
assistant_id=SCHEDULER_GRAPH_ID,
|
||||||
@@ -102,5 +186,7 @@ def run_now(prompt: str) -> Run:
|
|||||||
"run_kind": SCHEDULED_RUN_KIND,
|
"run_kind": SCHEDULED_RUN_KIND,
|
||||||
"name": "manual-run",
|
"name": "manual-run",
|
||||||
"prompt": prompt,
|
"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.
|
agent-chat-ui, LangSmith Studio) and SDK clients.
|
||||||
|
|
||||||
Differs from ``EvoSci`` / ``EvoSci serve``: no in-process CLI agent,
|
Differs from ``EvoSci`` / ``EvoSci serve``: no in-process CLI agent,
|
||||||
no session DB, no channel runtime, no TUI. The terminal only shows
|
no session DB, no channel runtime, no TUI. The terminal shows startup
|
||||||
startup progress, the Ready banner, and then blocks until Ctrl+C.
|
progress, the Ready banner, and the live Gateway log until Ctrl+C.
|
||||||
|
|
||||||
Mode dispatch happens via the ``EVOSCIENTIST_DEPLOY_MODE`` env var
|
Mode dispatch happens via the ``EVOSCIENTIST_DEPLOY_MODE`` env var
|
||||||
injected by ``start_langgraph_dev``: ``full`` for the deploy subprocess
|
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
|
from __future__ import annotations
|
||||||
|
|
||||||
import atexit
|
import atexit
|
||||||
|
import codecs
|
||||||
import os
|
import os
|
||||||
import signal
|
import signal
|
||||||
|
import sys
|
||||||
import threading
|
import threading
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any, TextIO
|
||||||
|
|
||||||
import typer # type: ignore[import-untyped]
|
import typer # type: ignore[import-untyped]
|
||||||
from rich.panel import Panel
|
from rich.panel import Panel
|
||||||
@@ -32,6 +34,65 @@ from rich.text import Text
|
|||||||
|
|
||||||
from ..cli._app import app
|
from ..cli._app import app
|
||||||
from ..stream.console import console
|
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()
|
@app.command()
|
||||||
@@ -39,7 +100,10 @@ def deploy(
|
|||||||
workdir: str | None = typer.Option(
|
workdir: str | None = typer.Option(
|
||||||
None,
|
None,
|
||||||
"--workdir",
|
"--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(
|
port: int | None = typer.Option(
|
||||||
None,
|
None,
|
||||||
@@ -64,11 +128,16 @@ def deploy(
|
|||||||
Connect any LangChain-compatible UI or SDK client to the printed
|
Connect any LangChain-compatible UI or SDK client to the printed
|
||||||
endpoint. Press Ctrl+C to stop.
|
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 (
|
from ..langgraph_dev.manager import (
|
||||||
_DEFAULT_PORT,
|
_DEFAULT_PORT,
|
||||||
RUNTIME,
|
RUNTIME,
|
||||||
_is_port_occupied,
|
_is_port_occupied,
|
||||||
|
current_log_start_offset,
|
||||||
is_langgraph_dev_running,
|
is_langgraph_dev_running,
|
||||||
read_tunnel_url,
|
read_tunnel_url,
|
||||||
start_langgraph_dev,
|
start_langgraph_dev,
|
||||||
@@ -88,17 +157,21 @@ def deploy(
|
|||||||
_configure_logging()
|
_configure_logging()
|
||||||
apply_config_to_env(config)
|
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:
|
if workdir:
|
||||||
ws = os.path.abspath(os.path.expanduser(workdir))
|
ws = os.path.abspath(os.path.expanduser(workdir))
|
||||||
elif config.default_workdir:
|
elif config.default_workdir:
|
||||||
ws = os.path.abspath(os.path.expanduser(config.default_workdir))
|
ws = os.path.abspath(os.path.expanduser(config.default_workdir))
|
||||||
else:
|
else:
|
||||||
ws = os.getcwd()
|
ws = str(get_default_workspace_dir())
|
||||||
# Subprocess inherits this path via EVOSCIENTIST_WORKSPACE_DIR (set inside
|
# Subprocess inherits this path via EVOSCIENTIST_WORKSPACE_DIR (set inside
|
||||||
# start_langgraph_dev). Ensure the dir exists; do NOT mutate the parent
|
# 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.
|
# process's paths module state — the deploy parent has no in-process agent.
|
||||||
os.makedirs(ws, exist_ok=True)
|
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"),
|
# 3. Resolve port (explicit None check — don't treat --port 0 as "unset"),
|
||||||
# then validate range so misconfigurations fail fast with a clear message
|
# then validate range so misconfigurations fail fast with a clear message
|
||||||
@@ -138,6 +211,18 @@ def deploy(
|
|||||||
)
|
)
|
||||||
raise typer.Exit(1)
|
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
|
# 5. Startup banner
|
||||||
_auth_label = _describe_auth(config)
|
_auth_label = _describe_auth(config)
|
||||||
console.print(
|
console.print(
|
||||||
@@ -169,9 +254,14 @@ def deploy(
|
|||||||
"trust.[/bold red]"
|
"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
|
_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:
|
try:
|
||||||
from ..ccproxy_manager import maybe_start_ccproxy, stop_ccproxy
|
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 +
|
# 10. Block on signal — mirror serve's dual-gate (threading.Event +
|
||||||
# explicit SIGINT/SIGTERM handlers) so SIGTERM (no default raise) also
|
# explicit SIGINT/SIGTERM handlers) so SIGTERM (no default raise) also
|
||||||
# triggers clean shutdown.
|
# triggers clean shutdown.
|
||||||
@@ -268,6 +367,8 @@ def deploy(
|
|||||||
except KeyboardInterrupt:
|
except KeyboardInterrupt:
|
||||||
shutdown_event.set()
|
shutdown_event.set()
|
||||||
finally:
|
finally:
|
||||||
|
log_stop_event.set()
|
||||||
|
log_thread.join(timeout=2.0)
|
||||||
signal.signal(signal.SIGINT, _orig_sigint)
|
signal.signal(signal.SIGINT, _orig_sigint)
|
||||||
signal.signal(signal.SIGTERM, _orig_sigterm)
|
signal.signal(signal.SIGTERM, _orig_sigterm)
|
||||||
# stop_langgraph_dev + stop_ccproxy run via atexit during interpreter
|
# 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.
|
# Front-end npm package + spec. ``@latest`` → always the newest published UI.
|
||||||
_WEBUI_PACKAGE = "@evoscientist/webui@latest"
|
_WEBUI_PACKAGE = "@evoscientist/webui@latest"
|
||||||
|
_WEBUI_PACKAGE_ENV = "EVOSCIENTIST_WEBUI_PACKAGE"
|
||||||
_DEFAULT_WEBUI_PORT = 4716
|
_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:
|
def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
||||||
"""Start the deploy-style backend + the WebUI front-end, then block.
|
"""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,
|
config: Effective ``EvoScientistConfig`` (already env-applied upstream,
|
||||||
but re-applied here so this is safe to call standalone).
|
but re-applied here so this is safe to call standalone).
|
||||||
workspace_dir: Resolved workspace path; falls back to
|
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
|
Blocks until Ctrl+C / SIGTERM, or until the front-end process exits, then
|
||||||
tears down both subprocesses. Never returns a value.
|
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 (
|
from ..langgraph_dev.manager import (
|
||||||
_DEFAULT_PORT,
|
_DEFAULT_PORT,
|
||||||
RUNTIME,
|
RUNTIME,
|
||||||
@@ -69,7 +75,7 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
|
|
||||||
apply_config_to_env(config)
|
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
|
# mirroring `EvoSci deploy`. The langgraph dev subprocess inherits this via
|
||||||
# EVOSCIENTIST_WORKSPACE_DIR (set inside start_langgraph_dev).
|
# EVOSCIENTIST_WORKSPACE_DIR (set inside start_langgraph_dev).
|
||||||
if workspace_dir:
|
if workspace_dir:
|
||||||
@@ -77,8 +83,14 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
elif getattr(config, "default_workdir", ""):
|
elif getattr(config, "default_workdir", ""):
|
||||||
ws = os.path.abspath(os.path.expanduser(config.default_workdir))
|
ws = os.path.abspath(os.path.expanduser(config.default_workdir))
|
||||||
else:
|
else:
|
||||||
ws = os.getcwd()
|
ws = str(get_default_workspace_dir())
|
||||||
os.makedirs(ws, exist_ok=True)
|
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),
|
# 2. Resolve ports: backend = langgraph dev (browser connects here),
|
||||||
# webui_port = the local Next.js server the browser actually opens.
|
# 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)
|
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,
|
# 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
|
# else start a fresh deploy-mode one (full MCP + async). Refuse a foreign
|
||||||
# occupant — that's a configuration error, not something to silently share.
|
# 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.
|
# are scrubbed — the browser UI never needs LLM provider API keys.
|
||||||
webui_env = _scrubbed_env(
|
webui_env = _scrubbed_env(
|
||||||
{
|
{
|
||||||
|
**usage_env,
|
||||||
"EVOSCIENTIST_LANGGRAPH_DEV_PORT": str(backend_port),
|
"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),
|
"PORT": str(webui_port),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
webui_package = _resolve_webui_package()
|
||||||
console.print(
|
console.print(
|
||||||
Panel(
|
Panel(
|
||||||
Text.from_markup(
|
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"[bold]WebUI:[/bold] http://localhost:{webui_port} "
|
||||||
f"[dim](opens in your browser)[/dim]\n"
|
f"[dim](opens in your browser)[/dim]\n"
|
||||||
f"[bold]Logs:[/bold] {_shorten(str(RUNTIME.log_file))}\n\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]"
|
f"moment)… Press Ctrl+C to stop.[/dim]"
|
||||||
),
|
),
|
||||||
title="[bold green]✓ EvoScientist WebUI[/bold green]",
|
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
|
popen_kwargs["creationflags"] = subprocess.CREATE_NEW_PROCESS_GROUP
|
||||||
try:
|
try:
|
||||||
webui_proc = subprocess.Popen(
|
webui_proc = subprocess.Popen(
|
||||||
[npx, "--yes", _WEBUI_PACKAGE, "--port", str(webui_port)],
|
[npx, "--yes", webui_package, "--port", str(webui_port)],
|
||||||
**popen_kwargs,
|
**popen_kwargs,
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
|
|||||||
@@ -476,7 +476,9 @@ async def alaunch_background_run(
|
|||||||
thread_id,
|
thread_id,
|
||||||
name=request.name,
|
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(
|
run_id = await _acreate_run(
|
||||||
client,
|
client,
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
|
|||||||
@@ -147,14 +147,34 @@ class LocalGraphGateway:
|
|||||||
target: GraphTarget,
|
target: GraphTarget,
|
||||||
request: RunRequest,
|
request: RunRequest,
|
||||||
) -> AsyncIterator[GraphEvent]:
|
) -> AsyncIterator[GraphEvent]:
|
||||||
|
from langgraph.types import Command
|
||||||
|
|
||||||
from ..stream.events import stream_agent_events
|
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(
|
inner = stream_agent_events(
|
||||||
local_graph,
|
local_graph,
|
||||||
request.message,
|
request.message,
|
||||||
request.thread_id,
|
request.thread_id,
|
||||||
metadata=request.metadata,
|
metadata=request.metadata,
|
||||||
media=request.media,
|
media=request.media,
|
||||||
|
runtime_snapshot_id=runtime_snapshot_id,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
async for event in inner:
|
async for event in inner:
|
||||||
|
|||||||
@@ -568,6 +568,17 @@ class LangGraphServerGateway:
|
|||||||
request.message,
|
request.message,
|
||||||
media=request.media,
|
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(
|
await stream.run.start(
|
||||||
input=run_input,
|
input=run_input,
|
||||||
config=config,
|
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)
|
||||||
+1177
-44
File diff suppressed because it is too large
Load Diff
@@ -5,244 +5,8 @@ in ``EvoScientist/EvoScientist.py`` so it doesn't construct on plain
|
|||||||
``import EvoScientist``. ``langgraph dev`` 's symbol resolver inspects
|
``import EvoScientist``. ``langgraph dev`` 's symbol resolver inspects
|
||||||
module attributes directly and doesn't trigger ``__getattr__``, so we
|
module attributes directly and doesn't trigger ``__getattr__``, so we
|
||||||
re-export here to make it visible.
|
re-export here to make it visible.
|
||||||
|
|
||||||
Before re-export we upgrade the compiled graph's class in place to
|
|
||||||
``_EvoFilteredGraph``, which strips ``PrivateStateAttr``-marked fields
|
|
||||||
(currently just ``_quickjs_snapshot_payload``) from ``get_state`` /
|
|
||||||
``get_state_history`` responses. Upstream ``langchain_quickjs`` annotates
|
|
||||||
the field ``PrivateStateAttr = OmitFromSchema(input=True, output=True)``,
|
|
||||||
but LangGraph's ``_prepare_state_snapshot`` doesn't honor that on
|
|
||||||
checkpoint reads — every ``getState`` materializes the delta chain back
|
|
||||||
into a full ~1.4 MB blob, which the WebUI then downloads. The subclass
|
|
||||||
closes the gap without touching the middleware's write path, preserving
|
|
||||||
cross-turn REPL persistence as ``langchain-ai/deepagents#3064`` shipped it.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from langgraph.graph.state import CompiledStateGraph
|
from EvoScientist.EvoScientist import EvoScientist_agent
|
||||||
from langgraph.types import PregelTask, StateSnapshot
|
|
||||||
|
|
||||||
from EvoScientist.EvoScientist import EvoScientist_agent as _agent
|
|
||||||
|
|
||||||
_PRIVATE_STATE_FIELDS = frozenset({"_quickjs_snapshot_payload"})
|
|
||||||
|
|
||||||
# Sanity check on the LangGraph internals ``_strip_private`` scrubs. If any
|
|
||||||
# of these attributes disappear or get renamed in a future upstream bump,
|
|
||||||
# the assertion fires at import time — the deployment refuses to start,
|
|
||||||
# instead of silently degrading (the filter would ``.get()`` its way to a
|
|
||||||
# no-op and the private-field payload would come back on the wire without
|
|
||||||
# anyone noticing until a user reports slow thread switches again).
|
|
||||||
#
|
|
||||||
# Doesn't cover every internal we depend on — ``metadata["writes"]`` /
|
|
||||||
# ``metadata["counters_since_delta_snapshot"]`` dict keys aren't a canary
|
|
||||||
# target because ``dict.get`` already tolerates their absence. What we
|
|
||||||
# canary here is the ``NamedTuple`` field set: renames there would be the
|
|
||||||
# highest-impact silent regression.
|
|
||||||
_EXPECTED_SNAPSHOT_FIELDS = frozenset({"values", "metadata", "tasks"})
|
|
||||||
_EXPECTED_TASK_FIELDS = frozenset({"result", "state"})
|
|
||||||
|
|
||||||
_missing_snap = _EXPECTED_SNAPSHOT_FIELDS - set(StateSnapshot._fields)
|
|
||||||
_missing_task = _EXPECTED_TASK_FIELDS - set(PregelTask._fields)
|
|
||||||
if _missing_snap or _missing_task:
|
|
||||||
raise RuntimeError(
|
|
||||||
"LangGraph state shape drifted from the version _strip_private was "
|
|
||||||
f"written against. Missing StateSnapshot fields: {_missing_snap or set()}. "
|
|
||||||
f"Missing PregelTask fields: {_missing_task or set()}. Review "
|
|
||||||
"_strip_private and re-verify against the current upstream shape "
|
|
||||||
"before removing this assertion."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _strip_private(snap):
|
|
||||||
"""Strip ``PrivateStateAttr``-marked fields from a ``StateSnapshot``.
|
|
||||||
|
|
||||||
Empirically verified against a live history response for a thread with
|
|
||||||
a single touched turn: the private field leaks on four surfaces — three
|
|
||||||
trivial, one heavy:
|
|
||||||
|
|
||||||
* ``snap.values`` — the materialized channel state exposed as the main
|
|
||||||
payload. For DeltaChannels this is the delta chain replayed into full
|
|
||||||
bytes (~1.4 MB for the quickjs snapshot). ``get_state`` and every
|
|
||||||
history entry.
|
|
||||||
* ``snap.metadata['writes']`` — ``{node_name: {channel: value}}`` map of
|
|
||||||
the raw writes that produced each checkpoint. On the ``after_agent``
|
|
||||||
step that first snapshots the REPL, ``value`` is the encoded write
|
|
||||||
record ``("snap", full_bytes)`` ≈ 1.4 MB.
|
|
||||||
* ``snap.tasks[*].result`` — the return dict of each completed
|
|
||||||
``PregelTask``. ``after_agent`` returns
|
|
||||||
``{"_quickjs_snapshot_payload": ("snap", bytes)}``; this dict becomes
|
|
||||||
the task's ``result`` field, which the API surfaces verbatim under
|
|
||||||
``tasks[*].result`` (``langgraph_api.state:106``). This is the
|
|
||||||
dominant leak: 1.7 MB in the last history entry of any thread whose
|
|
||||||
most-recent-in-window checkpoint had a snapshot anchor.
|
|
||||||
* ``snap.metadata['counters_since_delta_snapshot']`` — DeltaChannel's
|
|
||||||
snapshot cadence bookkeeping ``{channel: [count, superstep]}``. Tiny
|
|
||||||
(~20 B) but exposes the private field name; strip for cleanliness.
|
|
||||||
* ``snap.tasks[*].state`` (nested ``StateSnapshot``) — populated when the
|
|
||||||
caller passes ``subgraphs=True``. Repeats all of the above surfaces
|
|
||||||
for each subgraph task, so recurse into it. Not exercised by the
|
|
||||||
current WebUI (which doesn't pass ``subgraphs=True`` on REST reads),
|
|
||||||
but SDK / curl / gRPC callers can.
|
|
||||||
"""
|
|
||||||
if snap is None:
|
|
||||||
return snap
|
|
||||||
values = {k: v for k, v in snap.values.items() if k not in _PRIVATE_STATE_FIELDS}
|
|
||||||
metadata = snap.metadata
|
|
||||||
if metadata:
|
|
||||||
new_metadata = metadata
|
|
||||||
if new_metadata.get("writes"):
|
|
||||||
scrubbed_writes = {
|
|
||||||
node: {
|
|
||||||
k: v for k, v in ch_writes.items() if k not in _PRIVATE_STATE_FIELDS
|
|
||||||
}
|
|
||||||
for node, ch_writes in new_metadata["writes"].items()
|
|
||||||
}
|
|
||||||
new_metadata = {**new_metadata, "writes": scrubbed_writes}
|
|
||||||
if new_metadata.get("counters_since_delta_snapshot"):
|
|
||||||
scrubbed_counters = {
|
|
||||||
k: v
|
|
||||||
for k, v in new_metadata["counters_since_delta_snapshot"].items()
|
|
||||||
if k not in _PRIVATE_STATE_FIELDS
|
|
||||||
}
|
|
||||||
new_metadata = {
|
|
||||||
**new_metadata,
|
|
||||||
"counters_since_delta_snapshot": scrubbed_counters,
|
|
||||||
}
|
|
||||||
metadata = new_metadata
|
|
||||||
tasks = snap.tasks
|
|
||||||
if tasks:
|
|
||||||
new_tasks = []
|
|
||||||
changed = False
|
|
||||||
for t in tasks:
|
|
||||||
replace_kwargs: dict = {}
|
|
||||||
result = getattr(t, "result", None)
|
|
||||||
if isinstance(result, dict) and any(
|
|
||||||
k in result for k in _PRIVATE_STATE_FIELDS
|
|
||||||
):
|
|
||||||
replace_kwargs["result"] = {
|
|
||||||
k: v for k, v in result.items() if k not in _PRIVATE_STATE_FIELDS
|
|
||||||
}
|
|
||||||
# ``t.state`` is a ``RunnableConfig | StateSnapshot | None`` per
|
|
||||||
# ``PregelTask``'s typing. When ``subgraphs=True`` on the caller,
|
|
||||||
# this holds the subgraph's fully-materialized ``StateSnapshot`` —
|
|
||||||
# which repeats the same four leak surfaces (``values``,
|
|
||||||
# ``metadata.writes``, ``metadata.counters_since_delta_snapshot``,
|
|
||||||
# ``tasks[*].result/state``). Recurse so the whole tree is clean.
|
|
||||||
nested_state = getattr(t, "state", None)
|
|
||||||
if isinstance(nested_state, StateSnapshot):
|
|
||||||
scrubbed_state = _strip_private(nested_state)
|
|
||||||
if scrubbed_state is not nested_state:
|
|
||||||
replace_kwargs["state"] = scrubbed_state
|
|
||||||
if replace_kwargs:
|
|
||||||
new_tasks.append(t._replace(**replace_kwargs))
|
|
||||||
changed = True
|
|
||||||
else:
|
|
||||||
new_tasks.append(t)
|
|
||||||
if changed:
|
|
||||||
tasks = tuple(new_tasks)
|
|
||||||
return snap._replace(values=values, metadata=metadata, tasks=tasks)
|
|
||||||
|
|
||||||
|
|
||||||
class _EvoFilteredGraph(CompiledStateGraph):
|
|
||||||
"""Filters ``PrivateStateAttr``-marked state fields from checkpoint reads.
|
|
||||||
|
|
||||||
``Pregel.copy`` uses ``self.__class__(**attrs)`` so this subclass
|
|
||||||
survives the ``graph_obj.copy(update=...)`` call in
|
|
||||||
``langgraph_api.graph.get_graph`` that binds the checkpointer / store
|
|
||||||
before yielding to endpoint handlers.
|
|
||||||
|
|
||||||
**Known gap — streaming paths.** The overrides only cover ``get_state``
|
|
||||||
/ ``get_state_history``. On this compiled graph,
|
|
||||||
``self.output_channels`` correctly excludes ``_quickjs_snapshot_payload``
|
|
||||||
(respects ``OmitFromSchema(output=True)``), but
|
|
||||||
``self.stream_channels_asis`` includes it alongside other private
|
|
||||||
fields (``jump_to``, ``_summarization_event``) — the two lists are
|
|
||||||
built by ``langgraph.graph.state``'s graph builder and only the first
|
|
||||||
checks the output schema. So a client streaming with
|
|
||||||
``stream_mode="values"`` or ``stream_mode="events"`` (which fall back
|
|
||||||
to ``stream_channels_asis`` when ``output_keys`` is ``None``) can pull
|
|
||||||
the anchor blob in per-run event data. Empirically the WebUI's
|
|
||||||
``stream_mode=["updates"]`` path is clean, so this is transient per-run
|
|
||||||
rather than the persistent per-getState download this PR targets.
|
|
||||||
Filter here first; extend into the stream layer if a client relying on
|
|
||||||
``values`` / ``events`` reports it.
|
|
||||||
"""
|
|
||||||
|
|
||||||
async def aget_state(self, config, *, subgraphs=False):
|
|
||||||
return _strip_private(await super().aget_state(config, subgraphs=subgraphs))
|
|
||||||
|
|
||||||
def get_state(self, config, *, subgraphs=False):
|
|
||||||
return _strip_private(super().get_state(config, subgraphs=subgraphs))
|
|
||||||
|
|
||||||
async def aget_state_history(self, config, **kw):
|
|
||||||
async for snap in super().aget_state_history(config, **kw):
|
|
||||||
yield _strip_private(snap)
|
|
||||||
|
|
||||||
def get_state_history(self, config, **kw):
|
|
||||||
for snap in super().get_state_history(config, **kw):
|
|
||||||
yield _strip_private(snap)
|
|
||||||
|
|
||||||
|
|
||||||
# In-place ``__class__`` swap: the subclass adds only methods (no new
|
|
||||||
# instance attributes) so the memory layout is identical and the swap is
|
|
||||||
# safe. Constructing a fresh ``_EvoFilteredGraph`` via ``.copy()`` would
|
|
||||||
# require reproducing the deep-agent build pipeline; the swap avoids that.
|
|
||||||
_agent.__class__ = _EvoFilteredGraph
|
|
||||||
EvoScientist_agent = _agent
|
|
||||||
|
|
||||||
|
|
||||||
def _apply_filter_to_all_registered_graphs() -> None:
|
|
||||||
"""Extend the class swap to every graph registered in ``langgraph.json``.
|
|
||||||
|
|
||||||
``EvoScientist.py:_build_middleware_stack`` installs
|
|
||||||
``create_code_interpreter_middleware`` unconditionally — it's not gated
|
|
||||||
on the ``for_async_subagent`` flag — so every subagent (sync ``task``
|
|
||||||
dispatch and async ``start_async_task``) carries the QuickJS REPL and
|
|
||||||
can produce ``_quickjs_snapshot_payload`` writes on its own checkpoint
|
|
||||||
namespace.
|
|
||||||
|
|
||||||
Async subagents get their own ``thread_id`` and their ``/threads/{id}/state``
|
|
||||||
endpoint is served by their own compiled graph. Without swapping the
|
|
||||||
class on those graphs, the filter we applied to ``EvoScientist_agent``
|
|
||||||
doesn't reach that endpoint and any real code_interpreter touch inside
|
|
||||||
a subagent leaks the anchor snapshot verbatim.
|
|
||||||
|
|
||||||
Reads the graph registry straight from ``langgraph.json`` so a new
|
|
||||||
subagent added to the config picks up the swap automatically — no
|
|
||||||
hardcoded list to keep in sync.
|
|
||||||
|
|
||||||
Idempotent (skips graphs already swapped) and safe on graphs that don't
|
|
||||||
use the middleware — ``_strip_private`` returns snapshots unchanged when
|
|
||||||
the private field is absent. Best-effort: if the config is unreadable
|
|
||||||
or an entry can't be resolved, the deployment still starts — only the
|
|
||||||
unresolvable subagents remain unfiltered.
|
|
||||||
"""
|
|
||||||
import json
|
|
||||||
from importlib import import_module
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
config_path = Path(__file__).parent / "langgraph.json"
|
|
||||||
try:
|
|
||||||
config = json.loads(config_path.read_text())
|
|
||||||
except (OSError, json.JSONDecodeError):
|
|
||||||
return
|
|
||||||
|
|
||||||
for path in config.get("graphs", {}).values():
|
|
||||||
# Format: "module.dotted.path:attr_name"
|
|
||||||
if ":" not in path:
|
|
||||||
continue
|
|
||||||
module_path, attr = path.rsplit(":", 1)
|
|
||||||
try:
|
|
||||||
module = import_module(module_path)
|
|
||||||
except ImportError:
|
|
||||||
continue
|
|
||||||
graph = getattr(module, attr, None)
|
|
||||||
if isinstance(graph, CompiledStateGraph) and not isinstance(
|
|
||||||
graph, _EvoFilteredGraph
|
|
||||||
):
|
|
||||||
graph.__class__ = _EvoFilteredGraph
|
|
||||||
|
|
||||||
|
|
||||||
_apply_filter_to_all_registered_graphs()
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["EvoScientist_agent"]
|
__all__ = ["EvoScientist_agent"]
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ import shutil
|
|||||||
import subprocess
|
import subprocess
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
|
import uuid
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
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.
|
"""Record the workspace + pid of the langgraph dev we just started.
|
||||||
|
|
||||||
Atomic write via temp-file + ``os.replace``: without this, a concurrent
|
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:
|
try:
|
||||||
RUNTIME.pid_dir.mkdir(parents=True, exist_ok=True)
|
RUNTIME.pid_dir.mkdir(parents=True, exist_ok=True)
|
||||||
tmp = RUNTIME.workspace_sidecar.with_suffix(".json.tmp")
|
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(
|
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)
|
os.replace(tmp, RUNTIME.workspace_sidecar)
|
||||||
except OSError as exc:
|
except OSError as exc:
|
||||||
@@ -231,11 +246,14 @@ def _read_workspace_sidecar() -> dict | None:
|
|||||||
return data
|
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()``
|
"""Best-effort sidecar removal — called alongside every ``RUNTIME.pid_file.unlink()``
|
||||||
so the workspace fingerprint never outlives the PID file it pairs with."""
|
so the workspace fingerprint never outlives the PID file it pairs with."""
|
||||||
|
runtime = runtime or RUNTIME
|
||||||
try:
|
try:
|
||||||
RUNTIME.workspace_sidecar.unlink()
|
runtime.workspace_sidecar.unlink()
|
||||||
except OSError:
|
except OSError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@@ -270,6 +288,16 @@ _LOG_OFFSET_AT_START: int = 0
|
|||||||
# langgraph dev log. Mirrors langgraph_api/tunneling/cloudflare.py.
|
# langgraph dev log. Mirrors langgraph_api/tunneling/cloudflare.py.
|
||||||
_TUNNEL_URL_RE = re.compile(r"https://[A-Za-z0-9.-]+\.trycloudflare\.com")
|
_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.
|
# Whether async sub-agents are usable in this process.
|
||||||
#
|
#
|
||||||
# - CLI / serve parent process: starts False; flipped True after
|
# - CLI / serve parent process: starts False; flipped True after
|
||||||
@@ -743,7 +771,7 @@ def start_langgraph_dev(
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
RUNTIME.pid_file.write_text(str(proc.pid), encoding="utf-8")
|
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
|
global _PROCESS_WORKSPACE
|
||||||
_PROCESS = proc
|
_PROCESS = proc
|
||||||
_PROCESS_WORKSPACE = workspace_dir
|
_PROCESS_WORKSPACE = workspace_dir
|
||||||
@@ -820,7 +848,11 @@ def read_tunnel_url(timeout: float = 35.0, poll_interval: float = 0.5) -> str |
|
|||||||
return None
|
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.
|
"""Gracefully stop a langgraph dev process.
|
||||||
|
|
||||||
Sends SIGTERM to the process group (langgraph dev spawns worker children),
|
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.
|
(which also hold ``_LOCK``) don't observe partially-cleared state.
|
||||||
"""
|
"""
|
||||||
global _PROCESS, _PROCESS_WORKSPACE
|
global _PROCESS, _PROCESS_WORKSPACE
|
||||||
|
runtime = runtime or RUNTIME
|
||||||
with _LOCK:
|
with _LOCK:
|
||||||
proc = proc if proc is not None else _PROCESS
|
proc = proc if proc is not None else _PROCESS
|
||||||
if proc is None:
|
if proc is None:
|
||||||
@@ -879,12 +912,12 @@ def stop_langgraph_dev(proc: subprocess.Popen | None = None) -> None:
|
|||||||
if proc is _PROCESS:
|
if proc is _PROCESS:
|
||||||
_PROCESS = None
|
_PROCESS = None
|
||||||
_PROCESS_WORKSPACE = None
|
_PROCESS_WORKSPACE = None
|
||||||
if RUNTIME.pid_file.exists():
|
if runtime.pid_file.exists():
|
||||||
try:
|
try:
|
||||||
RUNTIME.pid_file.unlink()
|
runtime.pid_file.unlink()
|
||||||
except OSError:
|
except OSError:
|
||||||
pass
|
pass
|
||||||
_unlink_workspace_sidecar()
|
_unlink_workspace_sidecar(runtime)
|
||||||
|
|
||||||
# Note: ``.langgraph_api/`` is intentionally NOT removed — it holds
|
# Note: ``.langgraph_api/`` is intentionally NOT removed — it holds
|
||||||
# langgraph dev's persisted async-task / scheduler / Store state that
|
# langgraph dev's persisted async-task / scheduler / Store state that
|
||||||
@@ -1064,5 +1097,8 @@ def _ensure_langgraph_dev_locked(
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
_ASYNC_SUBAGENTS_AVAILABLE = True
|
_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
|
return proc
|
||||||
|
|||||||
@@ -1,33 +1,31 @@
|
|||||||
"""LLM module for EvoScientist.
|
"""LLM module for EvoScientist.
|
||||||
|
|
||||||
Provides a unified interface for creating chat model instances
|
The static model catalog and free-string model factory were removed in the
|
||||||
with support for multiple providers.
|
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
|
What remains here:
|
||||||
that importing ``EvoScientist.llm`` (or any of its submodules, like
|
|
||||||
``context_window``) does not eagerly drag in ``langchain.chat_models`` and
|
- ``context_window`` — context-window resolution helpers consumed by the
|
||||||
its transitive ``langchain_anthropic``/``langchain_openai`` stack — that's
|
middleware layer (e.g. ``context_editing``);
|
||||||
roughly 1 s of wall time on every CLI invocation.
|
- ``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
|
import lazy_loader as _lazy
|
||||||
|
|
||||||
__getattr__, __dir__, __all__ = _lazy.attach(
|
__getattr__, __dir__, __all__ = _lazy.attach(
|
||||||
__name__,
|
__name__,
|
||||||
submodules=["context_window", "models", "patches"],
|
submodules=["context_window", "patches"],
|
||||||
submod_attrs={
|
submod_attrs={
|
||||||
"context_window": [
|
"context_window": [
|
||||||
"DEFAULT_CONTEXT_WINDOW_FALLBACK",
|
"DEFAULT_CONTEXT_WINDOW_FALLBACK",
|
||||||
"get_context_window",
|
"get_context_window",
|
||||||
"resolve_context_window",
|
"resolve_context_window",
|
||||||
],
|
],
|
||||||
"models": [
|
|
||||||
"DEFAULT_MODEL",
|
|
||||||
"MODELS",
|
|
||||||
"get_chat_model",
|
|
||||||
"get_model_info",
|
|
||||||
"get_models_for_provider",
|
|
||||||
"list_models",
|
|
||||||
],
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,386 +0,0 @@
|
|||||||
"""Provider-error surface for langgraph SSE frames.
|
|
||||||
|
|
||||||
Provides :class:`ProviderStreamError` — a normalized, non-dataclass
|
|
||||||
exception raised by ``ErrorNormalizationMiddleware`` in place of the
|
|
||||||
provider SDK exception that a chat model call raised. Non-dataclass on
|
|
||||||
purpose: since orjson 3.0, dataclass instances are serialized natively
|
|
||||||
via their field enumeration, skipping the ``default=`` hook that
|
|
||||||
would otherwise build our SSE envelope. Some provider SDKs (openrouter
|
|
||||||
today) decorate their exceptions with ``@dataclass``, so their errors
|
|
||||||
emerge on the wire as raw dataclass fields — no envelope, no way for
|
|
||||||
the WebUI to distinguish quota / auth / rate-limit. Wrapping them in
|
|
||||||
a plain ``Exception`` subclass here keeps orjson on the ``default=``
|
|
||||||
path, which then calls :meth:`ProviderStreamError.model_dump`
|
|
||||||
(upstream ``langgraph_api.serde.default`` checks that hook before its
|
|
||||||
``BaseException`` branch) — no serde monkey-patch needed.
|
|
||||||
|
|
||||||
Also lives here: the pure-function helpers the middleware uses to
|
|
||||||
build the envelope (provider tag from ``ModelRequest.model``, SDK
|
|
||||||
field extractors, env-driven API-key redaction). They stay next to
|
|
||||||
:class:`ProviderStreamError` because the middleware is their only
|
|
||||||
consumer.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import os
|
|
||||||
import re
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# ProviderStreamError
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
class AgentControlError(Exception):
|
|
||||||
"""Host-defined terminal control error that must bypass model fallback."""
|
|
||||||
|
|
||||||
non_fallbackable = True
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
code: str,
|
|
||||||
message: str,
|
|
||||||
*,
|
|
||||||
status_code: int = 403,
|
|
||||||
retryable: bool = False,
|
|
||||||
) -> None:
|
|
||||||
super().__init__(message)
|
|
||||||
self.code = code
|
|
||||||
self.message = message
|
|
||||||
self.status_code = status_code
|
|
||||||
self.retryable = retryable
|
|
||||||
|
|
||||||
def model_dump(self) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"error": type(self).__name__,
|
|
||||||
"code": self.code,
|
|
||||||
"message": self.message,
|
|
||||||
"status_code": self.status_code,
|
|
||||||
"retryable": self.retryable,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class ModelToolProtocolError(AgentControlError):
|
|
||||||
"""A completed model response contained an invalid tool-call protocol."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
reason: str,
|
|
||||||
*,
|
|
||||||
provider: str | None = None,
|
|
||||||
model: str | None = None,
|
|
||||||
route_key: str | None = None,
|
|
||||||
config_generation: int | None = None,
|
|
||||||
api_mode: str | None = None,
|
|
||||||
endpoint: str | None = None,
|
|
||||||
tool_call_transport: str | None = None,
|
|
||||||
call_id: str | None = None,
|
|
||||||
call_diagnostic: dict[str, Any] | None = None,
|
|
||||||
) -> None:
|
|
||||||
super().__init__(
|
|
||||||
"MODEL_TOOL_PROTOCOL_INVALID",
|
|
||||||
"The model returned an invalid structured tool call.",
|
|
||||||
status_code=502,
|
|
||||||
retryable=False,
|
|
||||||
)
|
|
||||||
self.reason = reason
|
|
||||||
self.provider = provider
|
|
||||||
self.model = model
|
|
||||||
self.route_key = route_key
|
|
||||||
self.config_generation = config_generation
|
|
||||||
self.api_mode = api_mode
|
|
||||||
self.endpoint = endpoint
|
|
||||||
self.tool_call_transport = tool_call_transport
|
|
||||||
self.call_id = call_id
|
|
||||||
# Internal-only, redacted structure for server logs. Deliberately omitted
|
|
||||||
# from model_dump() so it never becomes part of the public SSE contract.
|
|
||||||
self.call_diagnostic = dict(call_diagnostic or {})
|
|
||||||
self.fallbackable = True
|
|
||||||
self.recoverable = True
|
|
||||||
|
|
||||||
def model_dump(self) -> dict[str, Any]:
|
|
||||||
payload = super().model_dump()
|
|
||||||
payload.update(
|
|
||||||
{
|
|
||||||
"reason": self.reason,
|
|
||||||
"fallbackable": self.fallbackable,
|
|
||||||
"recoverable": self.recoverable,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
for key in (
|
|
||||||
"provider",
|
|
||||||
"model",
|
|
||||||
"route_key",
|
|
||||||
"config_generation",
|
|
||||||
"api_mode",
|
|
||||||
"endpoint",
|
|
||||||
"tool_call_transport",
|
|
||||||
"call_id",
|
|
||||||
):
|
|
||||||
value = getattr(self, key)
|
|
||||||
if value is not None:
|
|
||||||
payload[key] = value
|
|
||||||
return payload
|
|
||||||
|
|
||||||
|
|
||||||
class ProviderStreamError(Exception):
|
|
||||||
"""Envelope-shaped wrapper for a provider SDK exception raised
|
|
||||||
inside a chat model call.
|
|
||||||
|
|
||||||
Attributes mirror the SSE envelope one-for-one:
|
|
||||||
|
|
||||||
- ``provider`` — concrete provider tag (``openai`` / ``anthropic``
|
|
||||||
/ ``deepseek`` / ``openrouter`` / ``openai_compat`` / …)
|
|
||||||
- ``class_qualname`` — fully qualified name of the underlying
|
|
||||||
exception's class (e.g. ``openrouter.errors.…``)
|
|
||||||
- ``message`` — API-key-redacted ``str(exc)``
|
|
||||||
- ``status_code`` — HTTP status if the SDK exposed one
|
|
||||||
- ``code`` — provider error code (``insufficient_quota``, …)
|
|
||||||
- ``err_type`` — provider error type label (openai's ``.type``)
|
|
||||||
- ``request_id`` — SDK-provided correlation id
|
|
||||||
|
|
||||||
The underlying exception is available via ``__cause__`` (set by
|
|
||||||
``raise ProviderStreamError(...) from exc`` in the middleware).
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
provider: str,
|
|
||||||
class_qualname: str,
|
|
||||||
message: str,
|
|
||||||
*,
|
|
||||||
status_code: int | None = None,
|
|
||||||
code: str | None = None,
|
|
||||||
err_type: str | None = None,
|
|
||||||
request_id: str | None = None,
|
|
||||||
) -> None:
|
|
||||||
super().__init__(message)
|
|
||||||
self.provider = provider
|
|
||||||
self.class_qualname = class_qualname
|
|
||||||
self.message = message
|
|
||||||
self.status_code = status_code
|
|
||||||
self.code = code
|
|
||||||
self.err_type = err_type
|
|
||||||
self.request_id = request_id
|
|
||||||
|
|
||||||
def as_envelope(self) -> dict[str, Any]:
|
|
||||||
"""Return the SSE envelope dict — the shape the WebUI consumes."""
|
|
||||||
payload: dict[str, Any] = {
|
|
||||||
"error": self.class_qualname.rsplit(".", 1)[-1],
|
|
||||||
"class": self.class_qualname,
|
|
||||||
"message": self.message,
|
|
||||||
"provider": self.provider,
|
|
||||||
}
|
|
||||||
if self.status_code is not None:
|
|
||||||
payload["status_code"] = self.status_code
|
|
||||||
if self.code is not None:
|
|
||||||
payload["code"] = self.code
|
|
||||||
if self.err_type is not None:
|
|
||||||
payload["type"] = self.err_type
|
|
||||||
if self.request_id:
|
|
||||||
payload["request_id"] = self.request_id
|
|
||||||
return payload
|
|
||||||
|
|
||||||
def model_dump(self) -> dict[str, Any]:
|
|
||||||
"""Serialization hook consumed by ``langgraph_api.serde.default``.
|
|
||||||
|
|
||||||
Upstream's dispatch checks ``hasattr(obj, 'model_dump')`` BEFORE
|
|
||||||
the ``isinstance(obj, BaseException)`` branch, so exposing this
|
|
||||||
method lets upstream emit our envelope with no monkey-patch on
|
|
||||||
its ``default`` callable. The name matches Pydantic's
|
|
||||||
convention deliberately — it's the hook upstream is looking
|
|
||||||
for.
|
|
||||||
"""
|
|
||||||
return self.as_envelope()
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# API-key redaction — env-driven, prefix-only
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
#
|
|
||||||
# Redaction is built from credentials actually deployed via env vars,
|
|
||||||
# not from generic key shapes. Rationale: (a) zero false positives —
|
|
||||||
# we only scrub strings we know are secrets, (b) defense-in-depth —
|
|
||||||
# the compiled regex holds only the first 8 chars of each key, so a
|
|
||||||
# leak of the regex object itself (traceback locals, process dump)
|
|
||||||
# can't expose the secret. Suffix-greedy match consumes the rest of
|
|
||||||
# the key shape at runtime. The table is rebuilt on every
|
|
||||||
# ``_redact_api_keys`` call so credentials loaded after import
|
|
||||||
# (typically ``load_dotenv`` in a main entry point) still get
|
|
||||||
# scrubbed. ``re.compile`` caches by source string internally, so an
|
|
||||||
# unchanged env costs a dict lookup.
|
|
||||||
|
|
||||||
_API_KEY_ENV_SUFFIXES = ("_API_KEY", "_TOKEN", "_SECRET")
|
|
||||||
_API_KEY_MIN_LEN = 12
|
|
||||||
_API_KEY_PREFIX_LEN = 8
|
|
||||||
|
|
||||||
|
|
||||||
def _build_env_key_redaction_re() -> re.Pattern[str] | None:
|
|
||||||
prefixes: list[str] = []
|
|
||||||
for k, v in os.environ.items():
|
|
||||||
if not k.endswith(_API_KEY_ENV_SUFFIXES):
|
|
||||||
continue
|
|
||||||
if not isinstance(v, str) or len(v) < _API_KEY_MIN_LEN:
|
|
||||||
continue
|
|
||||||
prefixes.append(re.escape(v[:_API_KEY_PREFIX_LEN]))
|
|
||||||
if not prefixes:
|
|
||||||
return None
|
|
||||||
alternation = "|".join(f"{p}[A-Za-z0-9_+/=.-]*" for p in prefixes)
|
|
||||||
return re.compile(alternation)
|
|
||||||
|
|
||||||
|
|
||||||
def _redact_api_keys(message: str) -> str:
|
|
||||||
"""Replace any deployed key prefix in *message* with ``<redacted>``.
|
|
||||||
|
|
||||||
Defensive; provider error messages occasionally echo the
|
|
||||||
authorization header back. Rebuilt per call so credentials loaded
|
|
||||||
after import (typical ``load_dotenv`` pattern) are still redacted.
|
|
||||||
"""
|
|
||||||
pattern = _build_env_key_redaction_re()
|
|
||||||
if pattern is None:
|
|
||||||
return message
|
|
||||||
return pattern.sub("<redacted>", message)
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Provider inference from ModelRequest.model
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
#
|
|
||||||
# Host → concrete provider. Hand-maintained snapshot mirroring the
|
|
||||||
# routed-provider tables in ``llm/models.py``
|
|
||||||
# (``_OPENAI_ROUTED_PROVIDERS`` + ``_ANTHROPIC_ROUTED_PROVIDERS``).
|
|
||||||
# Kept here rather than imported from ``models.py`` to keep the
|
|
||||||
# import surface of ``errors.py`` minimal — importing ``models.py``
|
|
||||||
# would pull in every langchain chat-model client at first
|
|
||||||
# middleware access. Consumed by ``_lookup_host_or_compat``; unknown
|
|
||||||
# hosts fall back to ``<module>_compat`` so the WebUI knows
|
|
||||||
# "openai/anthropic SDK, but not native" instead of getting a
|
|
||||||
# misleading concrete tag. Update when a new routed provider is
|
|
||||||
# added to ``models.py``.
|
|
||||||
#
|
|
||||||
# Related sibling: ``_PROVIDER_EXC_MODULE_PREFIXES`` in
|
|
||||||
# ``middleware/error_normalization.py`` — the exception-side
|
|
||||||
# provider allow-list. Adding a whole new provider SDK (not just a
|
|
||||||
# new base_url routed through an existing one) means updating that
|
|
||||||
# list too.
|
|
||||||
|
|
||||||
_HOST_TO_PROVIDER: dict[str, str] = {
|
|
||||||
"api.openai.com": "openai",
|
|
||||||
"api.anthropic.com": "anthropic",
|
|
||||||
"api.deepseek.com": "deepseek",
|
|
||||||
"api.moonshot.cn": "moonshot",
|
|
||||||
"api.siliconflow.cn": "siliconflow",
|
|
||||||
"open.bigmodel.cn": "zhipu", # zhipu + zhipu-code share this host
|
|
||||||
"ark.cn-beijing.volces.com": "volcengine",
|
|
||||||
"dashscope.aliyuncs.com": "dashscope",
|
|
||||||
"coding.dashscope.aliyuncs.com": "dashscope",
|
|
||||||
"api.minimaxi.com": "minimax",
|
|
||||||
"api.kimi.com": "kimi", # kimi-coding shares this host
|
|
||||||
"openrouter.ai": "openrouter",
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _provider_from_model(model: Any) -> str | None:
|
|
||||||
"""Derive the concrete provider tag from a chat model instance.
|
|
||||||
|
|
||||||
Class-based dispatch for unambiguous providers (``ChatOpenRouter``,
|
|
||||||
``ChatGoogleGenerativeAI``); ``openai_api_base`` /
|
|
||||||
``anthropic_api_url`` looked up in ``_HOST_TO_PROVIDER`` for
|
|
||||||
openai/anthropic-shape clients (native + routed). Returns ``None``
|
|
||||||
when the model isn't from a recognized provider SDK — the caller
|
|
||||||
(``ErrorNormalizationMiddleware``) then passes the exception
|
|
||||||
through unchanged.
|
|
||||||
"""
|
|
||||||
cls_module = type(model).__module__ or ""
|
|
||||||
if cls_module.startswith("langchain_openrouter"):
|
|
||||||
return "openrouter"
|
|
||||||
if cls_module.startswith("langchain_google_genai"):
|
|
||||||
return "google_genai"
|
|
||||||
if cls_module.startswith("langchain_openai"):
|
|
||||||
return _lookup_host_or_compat(
|
|
||||||
getattr(model, "openai_api_base", None), module_tag="openai"
|
|
||||||
)
|
|
||||||
if cls_module.startswith("langchain_anthropic"):
|
|
||||||
return _lookup_host_or_compat(
|
|
||||||
getattr(model, "anthropic_api_url", None), module_tag="anthropic"
|
|
||||||
)
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _lookup_host_or_compat(base_url: str | None, module_tag: str) -> str:
|
|
||||||
"""Extract host from *base_url* and look up in ``_HOST_TO_PROVIDER``.
|
|
||||||
|
|
||||||
Falls back to *module_tag* when no ``base_url`` is set (native SDK
|
|
||||||
default endpoint) or ``<module_tag>_compat`` for an unrecognized
|
|
||||||
host — the honest "openai SDK shape but unknown upstream" tag.
|
|
||||||
"""
|
|
||||||
if not base_url:
|
|
||||||
return module_tag
|
|
||||||
try:
|
|
||||||
from urllib.parse import urlparse
|
|
||||||
|
|
||||||
host = urlparse(base_url).hostname
|
|
||||||
except Exception:
|
|
||||||
host = None
|
|
||||||
if not host:
|
|
||||||
return module_tag
|
|
||||||
return _HOST_TO_PROVIDER.get(host.lower(), f"{module_tag}_compat")
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# SDK-field extractors — populate the envelope's optional fields
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
def _extract_status_code(exc: BaseException) -> int | None:
|
|
||||||
"""Best-effort HTTP status code from a provider SDK exception.
|
|
||||||
|
|
||||||
Order matters: openai/anthropic store it on ``.status_code``;
|
|
||||||
httpx-wrappers expose it via ``.response.status_code``;
|
|
||||||
``google.genai.errors.APIError`` (unusually) stores it as an
|
|
||||||
integer ``.code`` — type-disambiguated from openai/anthropic's
|
|
||||||
string ``.code`` (provider error code, surfaced separately).
|
|
||||||
"""
|
|
||||||
status_code = getattr(exc, "status_code", None)
|
|
||||||
if isinstance(status_code, int):
|
|
||||||
return status_code
|
|
||||||
response = getattr(exc, "response", None)
|
|
||||||
if response is not None:
|
|
||||||
rsc = getattr(response, "status_code", None)
|
|
||||||
if isinstance(rsc, int):
|
|
||||||
return rsc
|
|
||||||
code = getattr(exc, "code", None)
|
|
||||||
if isinstance(code, int):
|
|
||||||
return code
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _extract_provider_code(exc: BaseException) -> str | None:
|
|
||||||
"""Provider error code (e.g. ``insufficient_quota``,
|
|
||||||
``invalid_api_key``). Distinct from HTTP status; higher signal for
|
|
||||||
a WebUI toast than the integer alone.
|
|
||||||
"""
|
|
||||||
code = getattr(exc, "code", None)
|
|
||||||
if isinstance(code, str) and code:
|
|
||||||
return code
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _extract_error_type(exc: BaseException) -> str | None:
|
|
||||||
"""Provider error type label.
|
|
||||||
|
|
||||||
- openai exposes this as ``.type`` (``rate_limit_error`` etc.)
|
|
||||||
- ``google.genai.errors.APIError`` stores a string label at
|
|
||||||
``.status`` (``"NOT_FOUND"``, ``"RESOURCE_EXHAUSTED"``, …) — a
|
|
||||||
good fit for the same field.
|
|
||||||
|
|
||||||
``.type`` takes precedence when both are set.
|
|
||||||
"""
|
|
||||||
err_type = getattr(exc, "type", None)
|
|
||||||
if isinstance(err_type, str) and err_type:
|
|
||||||
return err_type
|
|
||||||
status = getattr(exc, "status", None)
|
|
||||||
if isinstance(status, str) and status:
|
|
||||||
return status
|
|
||||||
return None
|
|
||||||
@@ -1,847 +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 re
|
|
||||||
import subprocess
|
|
||||||
import warnings
|
|
||||||
from functools import lru_cache
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from langchain.chat_models import init_chat_model
|
|
||||||
|
|
||||||
from ..config.settings import (
|
|
||||||
OPENROUTER_DEFAULT_APP_CATEGORIES,
|
|
||||||
OPENROUTER_DEFAULT_APP_TITLE,
|
|
||||||
OPENROUTER_DEFAULT_HTTP_REFERER,
|
|
||||||
)
|
|
||||||
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/"
|
|
||||||
|
|
||||||
# Minimum Codex CLI version advertised when no explicit override is set. Newer
|
|
||||||
# installed versions are advertised automatically.
|
|
||||||
_CODEX_CLIENT_VERSION_FALLBACK = "0.144.1"
|
|
||||||
|
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
|
||||||
def _installed_codex_client_version() -> str:
|
|
||||||
"""Return the installed Codex CLI version, or an empty string."""
|
|
||||||
try:
|
|
||||||
result = subprocess.run(
|
|
||||||
["codex", "--version"],
|
|
||||||
capture_output=True,
|
|
||||||
text=True,
|
|
||||||
timeout=2,
|
|
||||||
check=False,
|
|
||||||
)
|
|
||||||
except (OSError, subprocess.TimeoutExpired):
|
|
||||||
return ""
|
|
||||||
|
|
||||||
if result.returncode != 0:
|
|
||||||
return ""
|
|
||||||
match = re.search(r"\b(\d+\.\d+\.\d+)\b", result.stdout + result.stderr)
|
|
||||||
return match.group(1) if match else ""
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve_codex_client_version() -> str:
|
|
||||||
"""Resolve an explicit override or the newer of installed and minimum versions."""
|
|
||||||
override = os.environ.get("EVOSCIENTIST_CODEX_CLIENT_VERSION", "").strip()
|
|
||||||
if override:
|
|
||||||
return override
|
|
||||||
|
|
||||||
installed = _installed_codex_client_version()
|
|
||||||
if installed and tuple(map(int, installed.split("."))) >= tuple(
|
|
||||||
map(int, _CODEX_CLIENT_VERSION_FALLBACK.split("."))
|
|
||||||
):
|
|
||||||
return installed
|
|
||||||
return _CODEX_CLIENT_VERSION_FALLBACK
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve_reasoning_effort(default: str) -> str:
|
|
||||||
"""Return the configured reasoning effort or a provider-specific default."""
|
|
||||||
return os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or default
|
|
||||||
|
|
||||||
|
|
||||||
# Providers routed through the OpenAI provider with a custom base_url.
|
|
||||||
# Maps provider name → (base_url or None, env var for API key).
|
|
||||||
_OPENAI_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = {
|
|
||||||
"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"}
|
|
||||||
|
|
||||||
# OpenRouter app attribution (issue #339). Default values are the single source
|
|
||||||
# of truth in config/settings.py (imported above); langchain-openrouter maps
|
|
||||||
# app_url → HTTP-Referer, app_title → X-Title, app_categories →
|
|
||||||
# X-OpenRouter-Categories. OpenRouter honors at most this many categories per
|
|
||||||
# request (server-side limit) and silently ignores the rest, so the sent list is
|
|
||||||
# capped to this many below. https://openrouter.ai/docs/app-attribution
|
|
||||||
_OPENROUTER_MAX_CATEGORIES_PER_REQUEST = 2
|
|
||||||
|
|
||||||
# Legacy/provider-specific options that are not accepted by the installed
|
|
||||||
# LangChain chat model constructors. Leaving them at the top level makes
|
|
||||||
# LangChain move them into model_kwargs and can later leak them into SDK calls.
|
|
||||||
_UNSUPPORTED_CHAT_MODEL_KWARGS = frozenset({"sanitize_openai_sdk_headers"})
|
|
||||||
|
|
||||||
# Model registry: list of (short_name, model_id, provider)
|
|
||||||
# Allows same short_name across different providers.
|
|
||||||
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
|
||||||
# Custom Anthropic (third-party Claude-compatible endpoints, current-gen defaults)
|
|
||||||
# Listed BEFORE native anthropic so MODELS dict defaults to native provider
|
|
||||||
("claude-sonnet-4-6", "claude-sonnet-4-6", "custom-anthropic"),
|
|
||||||
("claude-haiku-4-5", "claude-haiku-4-5", "custom-anthropic"),
|
|
||||||
# Custom OpenAI (third-party OpenAI-compatible endpoints, 3 defaults)
|
|
||||||
# Listed BEFORE native openai so MODELS dict defaults to native provider
|
|
||||||
("gpt-5.5-pro", "gpt-5.5-pro", "custom-openai"),
|
|
||||||
("gpt-5.5", "gpt-5.5", "custom-openai"),
|
|
||||||
("gpt-5.4", "gpt-5.4", "custom-openai"),
|
|
||||||
("gpt-5.3-codex", "gpt-5.3-codex", "custom-openai"),
|
|
||||||
("gpt-5-mini", "gpt-5-mini", "custom-openai"),
|
|
||||||
# Anthropic (current generation)
|
|
||||||
("claude-fable-5", "claude-fable-5", "anthropic"),
|
|
||||||
("claude-opus-4-8", "claude-opus-4-8", "anthropic"),
|
|
||||||
("claude-sonnet-5", "claude-sonnet-5", "anthropic"),
|
|
||||||
("claude-sonnet-4-6", "claude-sonnet-4-6", "anthropic"),
|
|
||||||
("claude-haiku-4-5", "claude-haiku-4-5", "anthropic"),
|
|
||||||
# OpenAI
|
|
||||||
("gpt-5.6-sol", "gpt-5.6-sol", "openai"),
|
|
||||||
("gpt-5.6-terra", "gpt-5.6-terra", "openai"),
|
|
||||||
("gpt-5.6-luna", "gpt-5.6-luna", "openai"),
|
|
||||||
("gpt-5.5-pro", "gpt-5.5-pro", "openai"),
|
|
||||||
("gpt-5.5", "gpt-5.5", "openai"),
|
|
||||||
("gpt-5.4", "gpt-5.4", "openai"),
|
|
||||||
("gpt-5.4-mini", "gpt-5.4-mini", "openai"),
|
|
||||||
("gpt-5.4-nano", "gpt-5.4-nano", "openai"),
|
|
||||||
("gpt-5.3-codex", "gpt-5.3-codex", "openai"),
|
|
||||||
("gpt-5.2-codex", "gpt-5.2-codex", "openai"),
|
|
||||||
("gpt-5.2", "gpt-5.2", "openai"),
|
|
||||||
("gpt-5.1", "gpt-5.1", "openai"),
|
|
||||||
("gpt-5", "gpt-5", "openai"),
|
|
||||||
("gpt-5-mini", "gpt-5-mini", "openai"),
|
|
||||||
("gpt-5-nano", "gpt-5-nano", "openai"),
|
|
||||||
# Google GenAI
|
|
||||||
("gemini-3.5-flash", "gemini-3.5-flash", "google-genai"),
|
|
||||||
("gemini-3.1-pro", "gemini-3.1-pro-preview", "google-genai"),
|
|
||||||
(
|
|
||||||
"gemini-3.1-pro-customtools",
|
|
||||||
"gemini-3.1-pro-preview-customtools",
|
|
||||||
"google-genai",
|
|
||||||
),
|
|
||||||
("gemini-3.1-flash-lite", "gemini-3.1-flash-lite-preview", "google-genai"),
|
|
||||||
("gemini-3-flash", "gemini-3-flash-preview", "google-genai"),
|
|
||||||
("gemini-2.5-flash", "gemini-2.5-flash", "google-genai"),
|
|
||||||
("gemini-2.5-flash-lite", "gemini-2.5-flash-lite", "google-genai"),
|
|
||||||
("gemini-2.5-pro", "gemini-2.5-pro", "google-genai"),
|
|
||||||
# MiniMax (direct API — Anthropic-compatible; default: api.minimaxi.com, global: api.minimax.io)
|
|
||||||
("minimax-m3", "MiniMax-M3", "minimax"),
|
|
||||||
("minimax-m2.7", "MiniMax-M2.7", "minimax"),
|
|
||||||
("minimax-m2.7-highspeed", "MiniMax-M2.7-highspeed", "minimax"),
|
|
||||||
("minimax-m2.5", "MiniMax-M2.5", "minimax"),
|
|
||||||
("minimax-m2.5-highspeed", "MiniMax-M2.5-highspeed", "minimax"),
|
|
||||||
# NVIDIA
|
|
||||||
("nemotron-super", "nvidia/nemotron-3-super-120b-a12b", "nvidia"),
|
|
||||||
("nemotron-nano", "nvidia/nemotron-3-nano-30b-a3b", "nvidia"),
|
|
||||||
("glm-5.2", "z-ai/glm-5.2", "nvidia"),
|
|
||||||
("glm4.7", "z-ai/glm4.7", "nvidia"),
|
|
||||||
("deepseek-v3.2", "deepseek-ai/deepseek-v3.2", "nvidia"),
|
|
||||||
("deepseek-v3.1", "deepseek-ai/deepseek-v3.1-terminus", "nvidia"),
|
|
||||||
("kimi-k2.5", "moonshotai/kimi-k2.5", "nvidia"),
|
|
||||||
("kimi-k2-thinking", "moonshotai/kimi-k2-thinking", "nvidia"),
|
|
||||||
("minimax-m2.5", "minimaxai/minimax-m2.5", "nvidia"),
|
|
||||||
("minimax-m2.1", "minimaxai/minimax-m2.1", "nvidia"),
|
|
||||||
("qwen3.5-397b", "qwen/qwen3.5-397b-a17b", "nvidia"),
|
|
||||||
("step-3.5-flash", "stepfun-ai/step-3.5-flash", "nvidia"),
|
|
||||||
# SiliconFlow
|
|
||||||
("minimax-m2.5", "Pro/MiniMaxAI/MiniMax-M2.5", "siliconflow"),
|
|
||||||
("glm-5.2", "Pro/zai-org/GLM-5.2", "siliconflow"),
|
|
||||||
("glm-5", "Pro/zai-org/GLM-5", "siliconflow"),
|
|
||||||
("kimi-k2.5", "Pro/moonshotai/Kimi-K2.5", "siliconflow"),
|
|
||||||
("glm-4.7", "Pro/zai-org/GLM-4.7", "siliconflow"),
|
|
||||||
# OpenRouter
|
|
||||||
("claude-fable-5", "anthropic/claude-fable-5", "openrouter"),
|
|
||||||
("claude-opus-4.8", "anthropic/claude-opus-4.8", "openrouter"),
|
|
||||||
("claude-opus-4.8-fast", "anthropic/claude-opus-4.8-fast", "openrouter"),
|
|
||||||
("claude-sonnet-5", "anthropic/claude-sonnet-5", "openrouter"),
|
|
||||||
("claude-sonnet-4.6", "anthropic/claude-sonnet-4.6", "openrouter"),
|
|
||||||
("gpt-5.6-sol", "openai/gpt-5.6-sol", "openrouter"),
|
|
||||||
("gpt-5.6-terra", "openai/gpt-5.6-terra", "openrouter"),
|
|
||||||
("gpt-5.6-luna", "openai/gpt-5.6-luna", "openrouter"),
|
|
||||||
("gpt-5.5-pro", "openai/gpt-5.5-pro", "openrouter"),
|
|
||||||
("gpt-5.5", "openai/gpt-5.5", "openrouter"),
|
|
||||||
("gpt-5.4", "openai/gpt-5.4", "openrouter"),
|
|
||||||
("gpt-5.3-codex", "openai/gpt-5.3-codex", "openrouter"),
|
|
||||||
("gemini-3.5-flash", "google/gemini-3.5-flash", "openrouter"),
|
|
||||||
("gemini-3.1-pro", "google/gemini-3.1-pro-preview", "openrouter"),
|
|
||||||
("gemini-3-flash", "google/gemini-3-flash-preview", "openrouter"),
|
|
||||||
("kimi-k2.6", "moonshotai/kimi-k2.6", "openrouter"),
|
|
||||||
("glm-5.2", "z-ai/glm-5.2", "openrouter"),
|
|
||||||
("glm-5v-turbo", "z-ai/glm-5v-turbo", "openrouter"),
|
|
||||||
("minimax-m3", "minimax/minimax-m3", "openrouter"),
|
|
||||||
("mimo-v2.5-pro", "xiaomi/mimo-v2.5-pro", "openrouter"),
|
|
||||||
("mimo-v2.5", "xiaomi/mimo-v2.5", "openrouter"),
|
|
||||||
("grok-build-0.1", "x-ai/grok-build-0.1", "openrouter"),
|
|
||||||
("grok-4.5", "x-ai/grok-4.5", "openrouter"),
|
|
||||||
("hy3", "tencent/hy3", "openrouter"),
|
|
||||||
("qwen3.7-max", "qwen/qwen3.7-max", "openrouter"),
|
|
||||||
("qwen3.7-plus", "qwen/qwen3.7-plus", "openrouter"),
|
|
||||||
("qwen3.6-flash", "qwen/qwen3.6-flash", "openrouter"),
|
|
||||||
("qwen3.5-122b", "qwen/qwen3.5-122b-a10b", "openrouter"),
|
|
||||||
("deepseek-v4-pro", "deepseek/deepseek-v4-pro", "openrouter"),
|
|
||||||
("deepseek-v4-flash", "deepseek/deepseek-v4-flash", "openrouter"),
|
|
||||||
# Zhipu CodePlan (智谱代码计划 — coding-only endpoint)
|
|
||||||
("glm-5.2", "glm-5.2", "zhipu-code"),
|
|
||||||
("glm-5.1", "glm-5.1", "zhipu-code"),
|
|
||||||
("glm-5", "glm-5", "zhipu-code"),
|
|
||||||
("glm-5-turbo", "glm-5-turbo", "zhipu-code"),
|
|
||||||
("glm-5v-turbo", "glm-5v-turbo", "zhipu-code"),
|
|
||||||
("glm-4.7", "glm-4.7", "zhipu-code"),
|
|
||||||
# Zhipu (智谱 — general endpoint, default for simple lookups)
|
|
||||||
("glm-5.2", "glm-5.2", "zhipu"),
|
|
||||||
("glm-5.1", "glm-5.1", "zhipu"),
|
|
||||||
("glm-5", "glm-5", "zhipu"),
|
|
||||||
("glm-5-turbo", "glm-5-turbo", "zhipu"),
|
|
||||||
("glm-5v-turbo", "glm-5v-turbo", "zhipu"),
|
|
||||||
("glm-4.7", "glm-4.7", "zhipu"),
|
|
||||||
# Volcengine (火山引擎 — Doubao models)
|
|
||||||
("doubao-seed-2.0-pro", "doubao-seed-2-0-pro-260215", "volcengine"),
|
|
||||||
("doubao-seed-2.0-lite", "doubao-seed-2-0-lite-260215", "volcengine"),
|
|
||||||
("doubao-seed-2.0-mini", "doubao-seed-2-0-mini-260215", "volcengine"),
|
|
||||||
("doubao-seed-2.0-code", "doubao-seed-2-0-code-preview-260215", "volcengine"),
|
|
||||||
("doubao-seed-1.6", "doubao-seed-1.6", "volcengine"),
|
|
||||||
("doubao-1.5-pro", "doubao-1.5-pro-256k", "volcengine"),
|
|
||||||
("doubao-1.5-thinking-pro", "doubao-1.5-thinking-pro", "volcengine"),
|
|
||||||
# DashScope Coding Plan (阿里云代码计划 — subscription sk-sp-* endpoint)
|
|
||||||
("qwen3.7-max", "qwen3.7-max", "dashscope-code"),
|
|
||||||
("qwen3.7-plus", "qwen3.7-plus", "dashscope-code"),
|
|
||||||
("qwen3.6-max", "qwen3.6-max-preview", "dashscope-code"),
|
|
||||||
("qwen3.6-plus", "qwen3.6-plus", "dashscope-code"),
|
|
||||||
("qwen3.6-flash", "qwen3.6-flash", "dashscope-code"),
|
|
||||||
("qwen3-coder", "qwen3-coder-plus", "dashscope-code"),
|
|
||||||
("qwen3-coder-next", "qwen3-coder-next", "dashscope-code"),
|
|
||||||
("qwen3-max", "qwen3-max", "dashscope-code"),
|
|
||||||
("qwen3.5-plus", "qwen3.5-plus", "dashscope-code"),
|
|
||||||
# DashScope (阿里云 — Qwen models, default for simple lookups)
|
|
||||||
("qwen3.7-max", "qwen3.7-max", "dashscope"),
|
|
||||||
("qwen3.7-plus", "qwen3.7-plus", "dashscope"),
|
|
||||||
("qwen3.6-max", "qwen3.6-max-preview", "dashscope"),
|
|
||||||
("qwen3.6-plus", "qwen3.6-plus", "dashscope"),
|
|
||||||
("qwen3.6-flash", "qwen3.6-flash", "dashscope"),
|
|
||||||
("qwen3-coder", "qwen3-coder-plus", "dashscope"),
|
|
||||||
("qwen3-235b", "qwen3-235b-a22b", "dashscope"),
|
|
||||||
("qwen-max", "qwen-max", "dashscope"),
|
|
||||||
("qwq-plus", "qwq-plus", "dashscope"),
|
|
||||||
# DeepSeek
|
|
||||||
("deepseek-v4-pro", "deepseek-v4-pro", "deepseek"),
|
|
||||||
("deepseek-v4-flash", "deepseek-v4-flash", "deepseek"),
|
|
||||||
# Legacy aliases (deprecated 2026-07-24; route to v4-flash thinking/non-thinking)
|
|
||||||
("deepseek-r1", "deepseek-reasoner", "deepseek"),
|
|
||||||
("deepseek-v3", "deepseek-chat", "deepseek"),
|
|
||||||
# Moonshot (OpenAI-compatible)
|
|
||||||
("kimi-k2.6", "kimi-k2.6", "moonshot"),
|
|
||||||
("kimi-k2.5", "kimi-k2.5", "moonshot"),
|
|
||||||
("kimi-k2-thinking", "kimi-k2-thinking", "moonshot"),
|
|
||||||
("kimi-k2-thinking-turbo", "kimi-k2-thinking-turbo", "moonshot"),
|
|
||||||
("moonshot-v1-auto", "moonshot-v1-auto", "moonshot"),
|
|
||||||
("moonshot-v1-128k", "moonshot-v1-128k", "moonshot"),
|
|
||||||
("moonshot-v1-32k", "moonshot-v1-32k", "moonshot"),
|
|
||||||
("moonshot-v1-8k", "moonshot-v1-8k", "moonshot"),
|
|
||||||
# Kimi Coding Plan (Anthropic-compatible)
|
|
||||||
("kimi-for-coding", "kimi-for-coding", "kimi-coding"),
|
|
||||||
]
|
|
||||||
|
|
||||||
# Public dict for simple lookups (last entry wins for duplicate names).
|
|
||||||
# Use get_models_for_provider() for provider-aware lookups.
|
|
||||||
MODELS: dict[str, tuple[str, str]] = {
|
|
||||||
name: (model_id, provider) for name, model_id, provider in _MODEL_ENTRIES
|
|
||||||
}
|
|
||||||
|
|
||||||
DEFAULT_MODEL = "claude-sonnet-4-6"
|
|
||||||
|
|
||||||
|
|
||||||
def get_models_for_provider(provider: str) -> list[tuple[str, str]]:
|
|
||||||
"""Get all models for a specific provider.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
provider: Provider name (e.g., 'anthropic', 'openrouter').
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of (short_name, model_id) tuples for the provider.
|
|
||||||
"""
|
|
||||||
return [(name, model_id) for name, model_id, p in _MODEL_ENTRIES if p == provider]
|
|
||||||
|
|
||||||
|
|
||||||
def _env_flag_enabled(name: str) -> bool:
|
|
||||||
return os.environ.get(name, "").strip().lower() in _TRUTHY_ENV_VALUES
|
|
||||||
|
|
||||||
|
|
||||||
def _env_flag_disabled(name: str) -> bool:
|
|
||||||
value = os.environ.get(name)
|
|
||||||
return value is not None and value.strip().lower() in _FALSEY_ENV_VALUES
|
|
||||||
|
|
||||||
|
|
||||||
def _drop_unsupported_chat_model_kwargs(kwargs: dict[str, Any]) -> None:
|
|
||||||
for key in _UNSUPPORTED_CHAT_MODEL_KWARGS:
|
|
||||||
kwargs.pop(key, None)
|
|
||||||
model_kwargs = kwargs.get("model_kwargs")
|
|
||||||
if isinstance(model_kwargs, dict):
|
|
||||||
for key in _UNSUPPORTED_CHAT_MODEL_KWARGS:
|
|
||||||
model_kwargs.pop(key, None)
|
|
||||||
|
|
||||||
|
|
||||||
def _supports_openrouter_anthropic_prompt_cache(provider: str, model_id: str) -> bool:
|
|
||||||
"""Return whether EvoScientist should declare OpenRouter Claude caching."""
|
|
||||||
return provider == "openrouter" and model_id.startswith(
|
|
||||||
("anthropic/", "~anthropic/")
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _has_cache_control_override(kwargs: dict[str, Any]) -> bool:
|
|
||||||
"""Return whether the caller already supplied cache-control settings."""
|
|
||||||
if "cache_control" in kwargs:
|
|
||||||
return True
|
|
||||||
model_kwargs = kwargs.get("model_kwargs")
|
|
||||||
if model_kwargs is None:
|
|
||||||
return False
|
|
||||||
if not isinstance(model_kwargs, dict):
|
|
||||||
warnings.warn(
|
|
||||||
"OpenRouter Anthropic prompt caching was not applied because "
|
|
||||||
"`model_kwargs` is not a dict; pass cache_control explicitly or use "
|
|
||||||
"a dict-shaped model_kwargs.",
|
|
||||||
UserWarning,
|
|
||||||
stacklevel=3,
|
|
||||||
)
|
|
||||||
return True
|
|
||||||
return "cache_control" in model_kwargs
|
|
||||||
|
|
||||||
|
|
||||||
def _apply_openrouter_anthropic_prompt_cache(
|
|
||||||
provider: str,
|
|
||||||
model_id: str,
|
|
||||||
kwargs: dict[str, Any],
|
|
||||||
) -> None:
|
|
||||||
"""Declare OpenRouter Claude prompt caching unless explicitly disabled.
|
|
||||||
|
|
||||||
OpenRouter already handles implicit caching for most providers, but Claude
|
|
||||||
prompt caching needs Anthropic-style cache-control declaration.
|
|
||||||
"""
|
|
||||||
if _env_flag_disabled("EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"):
|
|
||||||
return
|
|
||||||
if not _supports_openrouter_anthropic_prompt_cache(provider, model_id):
|
|
||||||
return
|
|
||||||
if _has_cache_control_override(kwargs):
|
|
||||||
return
|
|
||||||
kwargs.setdefault("model_kwargs", {})["cache_control"] = {"type": "ephemeral"}
|
|
||||||
|
|
||||||
|
|
||||||
def _apply_auto_config(
|
|
||||||
provider: str,
|
|
||||||
model_id: str,
|
|
||||||
is_third_party: bool,
|
|
||||||
kwargs: dict[str, Any],
|
|
||||||
original_provider: str | None = None,
|
|
||||||
) -> None:
|
|
||||||
"""Auto-enable provider-specific features (thinking, reasoning, etc.).
|
|
||||||
|
|
||||||
Mutates *kwargs* in place. Only sets keys that the caller hasn't already
|
|
||||||
provided, so explicit user settings are never overridden.
|
|
||||||
"""
|
|
||||||
disable_reasoning = bool(kwargs.pop("_disable_reasoning", False))
|
|
||||||
disable_thinking = bool(kwargs.pop("_disable_thinking", False))
|
|
||||||
if disable_reasoning:
|
|
||||||
kwargs.pop("reasoning", None)
|
|
||||||
kwargs.pop("include_thoughts", None)
|
|
||||||
if disable_thinking:
|
|
||||||
kwargs.pop("thinking", None)
|
|
||||||
|
|
||||||
# Anthropic: extended thinking
|
|
||||||
if provider == "anthropic" and not disable_thinking and "thinking" not in kwargs:
|
|
||||||
_supports_thinking = original_provider in _THINKING_CAPABLE_PROVIDERS
|
|
||||||
# Detect local proxy (e.g. ccproxy): thinking blocks in conversation
|
|
||||||
# history cause 422 errors because the proxy doesn't accept 'thinking'
|
|
||||||
# as a valid content block type on round-trip.
|
|
||||||
if not is_third_party:
|
|
||||||
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
|
|
||||||
_is_proxy = "127.0.0.1" in base_url or "localhost" in base_url
|
|
||||||
else:
|
|
||||||
_is_proxy = False
|
|
||||||
if _is_proxy or (is_third_party and not _supports_thinking):
|
|
||||||
pass
|
|
||||||
elif "fable" in model_id or model_id.endswith(("4-6", "4-7", "4-8")):
|
|
||||||
kwargs["thinking"] = {"type": "adaptive", "display": "summarized"}
|
|
||||||
kwargs.setdefault("effort", "max")
|
|
||||||
else:
|
|
||||||
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000}
|
|
||||||
|
|
||||||
# OpenAI (native, not third-party routed): reasoning
|
|
||||||
if (
|
|
||||||
provider == "openai"
|
|
||||||
and not is_third_party
|
|
||||||
and not disable_reasoning
|
|
||||||
and "reasoning" not in kwargs
|
|
||||||
):
|
|
||||||
_default_effort = (
|
|
||||||
"xhigh"
|
|
||||||
if (
|
|
||||||
"5.4" in model_id
|
|
||||||
or "5.5" in model_id
|
|
||||||
or "5.6" in model_id
|
|
||||||
or "codex" in model_id
|
|
||||||
)
|
|
||||||
else "high"
|
|
||||||
)
|
|
||||||
_eff = _resolve_reasoning_effort(_default_effort)
|
|
||||||
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
|
|
||||||
|
|
||||||
# Google GenAI: surface thinking traces
|
|
||||||
if provider == "google-genai" and not disable_reasoning:
|
|
||||||
kwargs.setdefault("include_thoughts", True)
|
|
||||||
|
|
||||||
# Ollama: separate reasoning content from response for thinking models
|
|
||||||
if provider == "ollama" and not disable_reasoning and "reasoning" not in kwargs:
|
|
||||||
kwargs["reasoning"] = True
|
|
||||||
|
|
||||||
|
|
||||||
def get_chat_model(
|
|
||||||
model: str | None = None,
|
|
||||||
provider: str | None = None,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> Any:
|
|
||||||
"""Get a chat model instance.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
model: Model name (short name like 'claude-sonnet-4-6' or full ID
|
|
||||||
like 'claude-sonnet-4-6-20250929'). Defaults to DEFAULT_MODEL.
|
|
||||||
provider: Override the provider (e.g., 'anthropic', 'openai').
|
|
||||||
If not specified, inferred from model name or defaults to 'anthropic'.
|
|
||||||
**kwargs: Additional arguments passed to init_chat_model (e.g., temperature).
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
A LangChain chat model instance.
|
|
||||||
|
|
||||||
Examples:
|
|
||||||
>>> model = get_chat_model() # Uses default (claude-sonnet-4-6)
|
|
||||||
>>> model = get_chat_model("claude-opus-4-8") # Use short name
|
|
||||||
>>> model = get_chat_model("gpt-4o") # OpenAI model
|
|
||||||
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
|
|
||||||
"""
|
|
||||||
skip_runtime_resolver = bool(kwargs.pop("_skip_runtime_model_resolver", False))
|
|
||||||
runtime_provider_name: str | None = None
|
|
||||||
runtime_supports_reasoning: bool | None = None
|
|
||||||
runtime_resolved = None
|
|
||||||
if not skip_runtime_resolver:
|
|
||||||
from EvoScientist.runtime_integrations import resolve_runtime_model
|
|
||||||
|
|
||||||
runtime_resolved = resolve_runtime_model(model, provider)
|
|
||||||
|
|
||||||
if runtime_resolved is not None:
|
|
||||||
resolved_params = dict(getattr(runtime_resolved, "params", {}) or {})
|
|
||||||
extra_body = resolved_params.pop("_extra_body", None)
|
|
||||||
default_headers = resolved_params.pop("_default_headers", None)
|
|
||||||
if extra_body:
|
|
||||||
resolved_params["extra_body"] = extra_body
|
|
||||||
if default_headers:
|
|
||||||
resolved_params["default_headers"] = default_headers
|
|
||||||
resolved_params.update(kwargs)
|
|
||||||
kwargs = resolved_params
|
|
||||||
|
|
||||||
resolved_api_key = str(getattr(runtime_resolved, "api_key", "") or "")
|
|
||||||
resolved_base_url = str(getattr(runtime_resolved, "base_url", "") or "")
|
|
||||||
if resolved_api_key:
|
|
||||||
kwargs.setdefault("api_key", resolved_api_key)
|
|
||||||
if resolved_base_url:
|
|
||||||
kwargs.setdefault("base_url", resolved_base_url.rstrip("/"))
|
|
||||||
|
|
||||||
runtime_provider_name = str(
|
|
||||||
getattr(runtime_resolved, "provider_name", "") or ""
|
|
||||||
)
|
|
||||||
runtime_supports_reasoning = bool(
|
|
||||||
getattr(runtime_resolved, "supports_reasoning", False)
|
|
||||||
)
|
|
||||||
if not runtime_supports_reasoning:
|
|
||||||
kwargs.setdefault("_disable_reasoning", True)
|
|
||||||
kwargs.setdefault("_disable_thinking", True)
|
|
||||||
model = str(runtime_resolved.model_id)
|
|
||||||
provider = str(runtime_resolved.protocol)
|
|
||||||
else:
|
|
||||||
model = model or DEFAULT_MODEL
|
|
||||||
|
|
||||||
# Look up short name in registry (provider-aware)
|
|
||||||
model_id = None
|
|
||||||
if provider:
|
|
||||||
# Try exact match with provider first
|
|
||||||
for name, mid, p in _MODEL_ENTRIES:
|
|
||||||
if name == model and p == provider:
|
|
||||||
model_id = mid
|
|
||||||
break
|
|
||||||
if model_id is None and model in MODELS:
|
|
||||||
model_id, default_provider = MODELS[model]
|
|
||||||
provider = provider or default_provider
|
|
||||||
|
|
||||||
if model_id is None:
|
|
||||||
# Assume it's a full model ID
|
|
||||||
model_id = model
|
|
||||||
# Try to infer provider from model ID prefix
|
|
||||||
if provider is None:
|
|
||||||
if model_id.startswith(("claude-", "anthropic")):
|
|
||||||
provider = "anthropic"
|
|
||||||
elif model_id.startswith(("gpt-", "o1", "davinci", "text-")):
|
|
||||||
provider = "openai"
|
|
||||||
elif model_id.startswith("gemini"):
|
|
||||||
provider = "google-genai"
|
|
||||||
elif model_id.startswith("ollama:"):
|
|
||||||
provider = "ollama"
|
|
||||||
model_id = model_id.removeprefix("ollama:")
|
|
||||||
else:
|
|
||||||
provider = "anthropic" # Default fallback
|
|
||||||
|
|
||||||
# Anthropic base_url override (e.g. ccproxy at localhost:8000/api/v1)
|
|
||||||
_is_third_party = (
|
|
||||||
provider in _OPENAI_ROUTED_PROVIDERS or provider in _ANTHROPIC_ROUTED_PROVIDERS
|
|
||||||
)
|
|
||||||
if runtime_provider_name and runtime_provider_name != provider:
|
|
||||||
_is_third_party = True
|
|
||||||
if (
|
|
||||||
runtime_resolved is not None
|
|
||||||
and provider == "openai"
|
|
||||||
and resolved_base_url
|
|
||||||
and "api.openai.com" not in resolved_base_url.lower()
|
|
||||||
):
|
|
||||||
_is_third_party = True
|
|
||||||
_is_openai_proxy = False
|
|
||||||
_original_provider: str | None = (
|
|
||||||
runtime_provider_name if runtime_provider_name != provider else None
|
|
||||||
)
|
|
||||||
if provider == "anthropic":
|
|
||||||
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
|
|
||||||
if base_url:
|
|
||||||
kwargs.setdefault("base_url", base_url)
|
|
||||||
api_key = os.environ.get("ANTHROPIC_API_KEY", "")
|
|
||||||
if api_key:
|
|
||||||
kwargs.setdefault("api_key", api_key)
|
|
||||||
|
|
||||||
# Native OpenAI base_url override (e.g. ccproxy Codex at localhost:8000/codex/v1)
|
|
||||||
elif provider == "openai":
|
|
||||||
base_url = os.environ.get("OPENAI_BASE_URL", "")
|
|
||||||
if base_url:
|
|
||||||
kwargs.setdefault("base_url", base_url)
|
|
||||||
_is_openai_proxy = _is_ccproxy_codex(
|
|
||||||
kwargs.get("base_url"), kwargs.get("api_key")
|
|
||||||
)
|
|
||||||
if _is_openai_proxy:
|
|
||||||
# Use Responses API for ccproxy: bypasses the format chain
|
|
||||||
# converter (Chat→Responses→Chat) which returns 502 on
|
|
||||||
# complex responses. System messages are converted to
|
|
||||||
# developer role by _patch_ccproxy_system_to_developer().
|
|
||||||
kwargs.setdefault("use_responses_api", True)
|
|
||||||
# Streaming must stay ON for Responses API: ccproxy's
|
|
||||||
# StreamingBufferService loses output when assembling
|
|
||||||
# non-streaming responses. (The old streaming=False was
|
|
||||||
# for Chat Completions tool_call duplication — not an issue
|
|
||||||
# with the Responses API SSE format.)
|
|
||||||
kwargs.pop("streaming", None) # remove if set elsewhere
|
|
||||||
# ccproxy forwards client headers upstream and only
|
|
||||||
# gap-fills its own, so the Codex backend sees this
|
|
||||||
# client's identity. Without Codex-CLI-shaped headers it
|
|
||||||
# rejects current models ("The '<model>' model requires
|
|
||||||
# a newer version of Codex").
|
|
||||||
_codex_ver = _resolve_codex_client_version()
|
|
||||||
_headers = kwargs.get("default_headers") or {}
|
|
||||||
kwargs["default_headers"] = _headers
|
|
||||||
_headers.setdefault("originator", "codex_cli_rs")
|
|
||||||
_headers.setdefault("version", _codex_ver)
|
|
||||||
_headers.setdefault(
|
|
||||||
"User-Agent",
|
|
||||||
f"codex_cli_rs/{_headers['version']} (EvoScientist)",
|
|
||||||
)
|
|
||||||
api_key = os.environ.get("OPENAI_API_KEY", "")
|
|
||||||
if api_key:
|
|
||||||
kwargs.setdefault("api_key", api_key)
|
|
||||||
|
|
||||||
# OpenAI-routed providers → route through OpenAI provider with base_url
|
|
||||||
elif provider in _OPENAI_ROUTED_PROVIDERS:
|
|
||||||
_original_provider = provider
|
|
||||||
base_url_default, api_key_env = _OPENAI_ROUTED_PROVIDERS[provider]
|
|
||||||
if provider == "custom-openai":
|
|
||||||
base_url = os.environ.get("CUSTOM_OPENAI_BASE_URL", "")
|
|
||||||
if not base_url:
|
|
||||||
raise ValueError(
|
|
||||||
"CUSTOM_OPENAI_BASE_URL environment variable is required when using "
|
|
||||||
"the 'custom-openai' provider. Please set it to your "
|
|
||||||
"OpenAI-compatible API endpoint URL (e.g. https://api.openai.com/v1)."
|
|
||||||
)
|
|
||||||
base_url = base_url.rstrip("/")
|
|
||||||
else:
|
|
||||||
base_url = base_url_default
|
|
||||||
if base_url:
|
|
||||||
kwargs.setdefault("base_url", base_url)
|
|
||||||
api_key = os.environ.get(api_key_env, "")
|
|
||||||
if api_key:
|
|
||||||
kwargs.setdefault("api_key", api_key)
|
|
||||||
# SiliconFlow: disable thinking — LangChain drops reasoning_content
|
|
||||||
# from history, causing error 20015 on multi-turn requests.
|
|
||||||
if provider == "siliconflow":
|
|
||||||
kwargs.setdefault("extra_body", {})["enable_thinking"] = False
|
|
||||||
# Moonshot: disable thinking for all models to prevent LangChain from dropping
|
|
||||||
# reasoning_content, which causes multi-turn conversation errors (error 20015).
|
|
||||||
# Even native thinking models like kimi-k2-thinking operate in non-thinking mode.
|
|
||||||
if provider == "moonshot":
|
|
||||||
kwargs.setdefault("extra_body", {})["thinking"] = {"type": "disabled"}
|
|
||||||
provider = "openai"
|
|
||||||
|
|
||||||
# OpenRouter → native ChatOpenRouter via init_chat_model.
|
|
||||||
elif provider == "openrouter":
|
|
||||||
_is_third_party = True
|
|
||||||
api_key = os.environ.get("OPENROUTER_API_KEY", "")
|
|
||||||
if api_key:
|
|
||||||
kwargs.setdefault("api_key", api_key)
|
|
||||||
# Reasoning via `effort` + `summary: "auto"` so a readable reasoning
|
|
||||||
# summary is returned for display. OpenAI-Responses also emits encrypted
|
|
||||||
# reasoning items (`rs_*` id) that can't be replayed on multi-turn
|
|
||||||
# passback (OpenRouter's `/responses` beta is stateless, store=false —
|
|
||||||
# "Item with id 'rs_...' not found"); the patch strips them on passback,
|
|
||||||
# so enabling `summary` is safe. See langchain-ai/langchain#37777.
|
|
||||||
effort = _resolve_reasoning_effort("high")
|
|
||||||
kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"})
|
|
||||||
# App attribution (issue #339): identify EvoScientist to OpenRouter so
|
|
||||||
# usage is credited to the project (app rankings, model app tabs,
|
|
||||||
# analytics) rather than langchain-openrouter's LangChain-branded
|
|
||||||
# defaults. setdefault so an explicit caller kwarg wins; values are
|
|
||||||
# configurable via EVOSCIENTIST_OPENROUTER_* env (fed from the config
|
|
||||||
# file by apply_config_to_env). Applied only here, so no other provider
|
|
||||||
# ever receives these kwargs.
|
|
||||||
kwargs.setdefault(
|
|
||||||
"app_url",
|
|
||||||
os.environ.get("EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "").strip()
|
|
||||||
or OPENROUTER_DEFAULT_HTTP_REFERER,
|
|
||||||
)
|
|
||||||
kwargs.setdefault(
|
|
||||||
"app_title",
|
|
||||||
os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE", "").strip()
|
|
||||||
or OPENROUTER_DEFAULT_APP_TITLE,
|
|
||||||
)
|
|
||||||
# app_categories must be a list[str] (langchain-openrouter joins it into
|
|
||||||
# the X-OpenRouter-Categories header); split the comma-separated config
|
|
||||||
# value and drop blanks so a stray comma/space can't emit an empty one.
|
|
||||||
_app_categories_raw = (
|
|
||||||
os.environ.get("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "").strip()
|
|
||||||
or OPENROUTER_DEFAULT_APP_CATEGORIES
|
|
||||||
)
|
|
||||||
_app_categories = [
|
|
||||||
c.strip() for c in _app_categories_raw.split(",") if c.strip()
|
|
||||||
]
|
|
||||||
# Cap to the per-request limit and warn, so a misconfigured extra is
|
|
||||||
# dropped predictably here (and surfaced to the user) rather than being
|
|
||||||
# silently truncated server-side.
|
|
||||||
_limit = _OPENROUTER_MAX_CATEGORIES_PER_REQUEST
|
|
||||||
if len(_app_categories) > _limit:
|
|
||||||
warnings.warn(
|
|
||||||
f"OpenRouter accepts at most {_limit} app categories per "
|
|
||||||
f"request, so only the first {_limit} are sent: "
|
|
||||||
f"{_app_categories[:_limit]}. Ignoring the rest: "
|
|
||||||
f"{_app_categories[_limit:]}. Set "
|
|
||||||
f"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES (or the "
|
|
||||||
f"openrouter_app_categories config) to at most {_limit} "
|
|
||||||
f"categories to silence this warning.",
|
|
||||||
UserWarning,
|
|
||||||
stacklevel=2,
|
|
||||||
)
|
|
||||||
_app_categories = _app_categories[:_limit]
|
|
||||||
if _app_categories:
|
|
||||||
kwargs.setdefault("app_categories", _app_categories)
|
|
||||||
_patch_openrouter_strip_responses_reasoning()
|
|
||||||
|
|
||||||
# Anthropic-routed providers → route through Anthropic provider with base_url
|
|
||||||
elif provider in _ANTHROPIC_ROUTED_PROVIDERS:
|
|
||||||
_original_provider = provider
|
|
||||||
base_url_default, api_key_env = _ANTHROPIC_ROUTED_PROVIDERS[provider]
|
|
||||||
if provider == "custom-anthropic":
|
|
||||||
base_url = os.environ.get("CUSTOM_ANTHROPIC_BASE_URL", "")
|
|
||||||
if not base_url:
|
|
||||||
raise ValueError(
|
|
||||||
"CUSTOM_ANTHROPIC_BASE_URL environment variable is required when using "
|
|
||||||
"the 'custom-anthropic' provider. Please set it to your "
|
|
||||||
"Anthropic-compatible API endpoint URL (e.g. https://api.anthropic.com)."
|
|
||||||
)
|
|
||||||
base_url = base_url.rstrip("/")
|
|
||||||
elif provider == "minimax":
|
|
||||||
base_url = os.environ.get("MINIMAX_BASE_URL", base_url_default).rstrip("/")
|
|
||||||
else:
|
|
||||||
base_url = base_url_default
|
|
||||||
if base_url:
|
|
||||||
kwargs.setdefault("base_url", base_url)
|
|
||||||
api_key = os.environ.get(api_key_env, "")
|
|
||||||
if api_key:
|
|
||||||
kwargs.setdefault("api_key", api_key)
|
|
||||||
# Kimi Coding Plan requires claude-code User-Agent header
|
|
||||||
if provider == "kimi-coding":
|
|
||||||
kwargs.setdefault("default_headers", {})["User-Agent"] = "claude-code/0.1.0"
|
|
||||||
provider = "anthropic"
|
|
||||||
|
|
||||||
elif provider == "ollama":
|
|
||||||
base_url = os.environ.get("OLLAMA_BASE_URL", "")
|
|
||||||
if base_url:
|
|
||||||
kwargs.setdefault("base_url", base_url)
|
|
||||||
|
|
||||||
_drop_unsupported_chat_model_kwargs(kwargs)
|
|
||||||
_apply_auto_config(provider, model_id, _is_third_party, kwargs, _original_provider)
|
|
||||||
_apply_openrouter_anthropic_prompt_cache(provider, model_id, kwargs)
|
|
||||||
|
|
||||||
# User-level override for the OpenAI Responses API vs Chat Completions.
|
|
||||||
# When "false", force Chat Completions and drop reasoning (which triggers
|
|
||||||
# the Responses API path in langchain-openai). Only applies to OpenAI.
|
|
||||||
if provider == "openai":
|
|
||||||
_responses_api_setting = (
|
|
||||||
os.environ.get("EVOSCIENTIST_USE_RESPONSES_API", "").strip().lower()
|
|
||||||
)
|
|
||||||
if _responses_api_setting == "false":
|
|
||||||
kwargs["use_responses_api"] = False
|
|
||||||
kwargs.pop("reasoning", None)
|
|
||||||
elif _responses_api_setting == "true":
|
|
||||||
kwargs["use_responses_api"] = True
|
|
||||||
|
|
||||||
anthropic_auth_token = None
|
|
||||||
if provider == "anthropic" and kwargs.get("api_key"):
|
|
||||||
anthropic_auth_token = os.environ.pop("ANTHROPIC_AUTH_TOKEN", None)
|
|
||||||
try:
|
|
||||||
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
|
|
||||||
finally:
|
|
||||||
if anthropic_auth_token is not None:
|
|
||||||
os.environ["ANTHROPIC_AUTH_TOKEN"] = anthropic_auth_token
|
|
||||||
|
|
||||||
# Flatten list content to strings for strict OpenAI-compatible providers
|
|
||||||
# (DeepSeek, SiliconFlow, OpenRouter, custom-openai, etc.) and
|
|
||||||
# native OpenAI through a proxy, to avoid "sequence expected string" errors.
|
|
||||||
# Moonshot and Kimi Coding support standard format, no patch needed.
|
|
||||||
_no_patch_providers = {"moonshot", "kimi-coding"}
|
|
||||||
if (
|
|
||||||
_is_third_party or _is_openai_proxy
|
|
||||||
) and _original_provider not in _no_patch_providers:
|
|
||||||
# Anthropic-routed providers accept media in tool results natively;
|
|
||||||
# only OpenAI-compatible providers need tool-media hoisting.
|
|
||||||
_hoist = _original_provider not in _ANTHROPIC_ROUTED_PROVIDERS
|
|
||||||
_patch_openai_compat_content(chat_model, hoist_tool_media=_hoist)
|
|
||||||
|
|
||||||
# DeepSeek thinking mode requires reasoning_content passback in multi-turn
|
|
||||||
# + tool_use scenarios.
|
|
||||||
if _original_provider == "deepseek":
|
|
||||||
_patch_deepseek_reasoning_passback(chat_model)
|
|
||||||
|
|
||||||
if _is_openai_proxy:
|
|
||||||
_patch_ccproxy_system_to_developer(chat_model)
|
|
||||||
|
|
||||||
apply_known_context_window(chat_model)
|
|
||||||
|
|
||||||
return chat_model
|
|
||||||
|
|
||||||
|
|
||||||
def list_models() -> list[str]:
|
|
||||||
"""List all available model short names.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of unique model short names that can be passed to get_chat_model().
|
|
||||||
"""
|
|
||||||
seen = set()
|
|
||||||
result = []
|
|
||||||
for name, _, _ in _MODEL_ENTRIES:
|
|
||||||
if name not in seen:
|
|
||||||
seen.add(name)
|
|
||||||
result.append(name)
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
def list_models_by_provider() -> list[tuple[str, str, str]]:
|
|
||||||
"""List all unique (short_name, model_id, provider) entries.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
De-duplicated list of model entries preserving registry order.
|
|
||||||
"""
|
|
||||||
seen: set[tuple[str, str]] = set()
|
|
||||||
result: list[tuple[str, str, str]] = []
|
|
||||||
for name, model_id, provider in _MODEL_ENTRIES:
|
|
||||||
key = (name, provider)
|
|
||||||
if key not in seen:
|
|
||||||
seen.add(key)
|
|
||||||
result.append((name, model_id, provider))
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
async def list_model_picker_entries(
|
|
||||||
ollama_base_url: str | None,
|
|
||||||
*,
|
|
||||||
include_custom_ollama: bool,
|
|
||||||
) -> list[tuple[str, str, str]]:
|
|
||||||
"""Return model picker entries, optionally including local Ollama models."""
|
|
||||||
entries = list_models_by_provider()
|
|
||||||
if ollama_base_url:
|
|
||||||
from .ollama_discovery import discover_ollama_models
|
|
||||||
|
|
||||||
for detected_name in await discover_ollama_models(
|
|
||||||
ollama_base_url,
|
|
||||||
timeout=1.5,
|
|
||||||
):
|
|
||||||
entries.append((detected_name, detected_name, "ollama"))
|
|
||||||
if include_custom_ollama:
|
|
||||||
entries.append(("Custom Ollama model...", "__custom_ollama__", "ollama"))
|
|
||||||
return entries
|
|
||||||
|
|
||||||
|
|
||||||
def get_model_info(model: str) -> tuple[str, str] | None:
|
|
||||||
"""Get the (model_id, provider) tuple for a short name.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
model: Short model name.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (model_id, provider) or None if not found.
|
|
||||||
"""
|
|
||||||
return MODELS.get(model)
|
|
||||||
@@ -1,7 +1,7 @@
|
|||||||
"""Ollama server probing — shared by onboard wizard and /model picker.
|
"""Ollama server probing — shared by onboard wizard and /model picker.
|
||||||
|
|
||||||
Ollama models are whatever the user has ``ollama pull``ed locally; they
|
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
|
interactive model picker need to hit ``GET {base_url}/api/tags`` to see
|
||||||
what is actually installed.
|
what is actually installed.
|
||||||
|
|
||||||
|
|||||||
+152
-414
@@ -25,7 +25,6 @@ Utilities:
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import hashlib
|
|
||||||
import os
|
import os
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -179,19 +178,14 @@ _patch_ccproxy_codex_compat()
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Utility: detect ccproxy's Codex adapter (as opposed to generic localhost).
|
# Utility: detect ccproxy's Codex adapter (as opposed to generic localhost).
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
def _is_ccproxy_codex(
|
def _is_ccproxy_codex() -> bool:
|
||||||
base_url: str | None = None,
|
|
||||||
api_key: str | None = None,
|
|
||||||
) -> bool:
|
|
||||||
"""Return True if the OpenAI endpoint is ccproxy's Codex adapter.
|
"""Return True if the OpenAI endpoint is ccproxy's Codex adapter.
|
||||||
|
|
||||||
Checks for the ccproxy-specific markers set by ``setup_codex_env()``
|
Checks for the ccproxy-specific markers set by ``setup_codex_env()``
|
||||||
in ``ccproxy_manager.py``: the sentinel API key and the ``/codex/v1``
|
in ``ccproxy_manager.py``: the sentinel API key and the ``/codex/v1``
|
||||||
path. Plain localhost endpoints (vLLM, Ollama, etc.) are not affected.
|
path. Plain localhost endpoints (vLLM, Ollama, etc.) are not affected.
|
||||||
"""
|
"""
|
||||||
if base_url is None:
|
|
||||||
base_url = os.environ.get("OPENAI_BASE_URL", "")
|
base_url = os.environ.get("OPENAI_BASE_URL", "")
|
||||||
if api_key is None:
|
|
||||||
api_key = os.environ.get("OPENAI_API_KEY", "")
|
api_key = os.environ.get("OPENAI_API_KEY", "")
|
||||||
return (
|
return (
|
||||||
("127.0.0.1" in base_url or "localhost" in base_url)
|
("127.0.0.1" in base_url or "localhost" in base_url)
|
||||||
@@ -273,299 +267,6 @@ def _flatten_message_content(content: Any) -> str | list[Any] | Any:
|
|||||||
return "\n\n".join(parts) if parts else ""
|
return "\n\n".join(parts) if parts else ""
|
||||||
|
|
||||||
|
|
||||||
def _stable_tool_call_id(message: Any, message_index: int, call_index: int) -> str:
|
|
||||||
seed = ":".join(
|
|
||||||
(
|
|
||||||
str(getattr(message, "id", "") or "message"),
|
|
||||||
str(message_index),
|
|
||||||
str(call_index),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
return "call_" + hashlib.sha256(seed.encode("utf-8")).hexdigest()[:24]
|
|
||||||
|
|
||||||
|
|
||||||
def _tool_message_match_index(
|
|
||||||
tool_messages: list[Any],
|
|
||||||
used_indexes: set[int],
|
|
||||||
*,
|
|
||||||
call_id: str,
|
|
||||||
call_name: str,
|
|
||||||
) -> int | None:
|
|
||||||
"""Find the best unused result for one assistant tool call."""
|
|
||||||
|
|
||||||
def _matches(index: int, *, require_id: bool, require_name: bool) -> bool:
|
|
||||||
if index in used_indexes:
|
|
||||||
return False
|
|
||||||
message = tool_messages[index]
|
|
||||||
result_id = str(getattr(message, "tool_call_id", "") or "")
|
|
||||||
result_name = str(getattr(message, "name", "") or "")
|
|
||||||
if require_id and result_id != call_id:
|
|
||||||
return False
|
|
||||||
if not require_id and result_id:
|
|
||||||
return False
|
|
||||||
return not require_name or not result_name or result_name == call_name
|
|
||||||
|
|
||||||
if call_id:
|
|
||||||
for require_name in (True, False):
|
|
||||||
for index in range(len(tool_messages)):
|
|
||||||
if _matches(index, require_id=True, require_name=require_name):
|
|
||||||
return index
|
|
||||||
for require_name in (True, False):
|
|
||||||
for index in range(len(tool_messages)):
|
|
||||||
if _matches(index, require_id=False, require_name=require_name):
|
|
||||||
return index
|
|
||||||
return None
|
|
||||||
|
|
||||||
# A result-side identifier is more authoritative than a generated fallback.
|
|
||||||
for require_name in (True, False):
|
|
||||||
for index, message in enumerate(tool_messages):
|
|
||||||
if index in used_indexes:
|
|
||||||
continue
|
|
||||||
result_id = str(getattr(message, "tool_call_id", "") or "")
|
|
||||||
result_name = str(getattr(message, "name", "") or "")
|
|
||||||
if result_id and (
|
|
||||||
not require_name or not result_name or result_name == call_name
|
|
||||||
):
|
|
||||||
return index
|
|
||||||
for require_name in (True, False):
|
|
||||||
for index in range(len(tool_messages)):
|
|
||||||
if _matches(index, require_id=False, require_name=require_name):
|
|
||||||
return index
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _copy_ai_message_with_tool_pairs(
|
|
||||||
message: Any,
|
|
||||||
message_index: int,
|
|
||||||
tool_messages: list[Any],
|
|
||||||
) -> tuple[Any | None, list[Any]]:
|
|
||||||
"""Return a replay-safe assistant message and its matched tool results."""
|
|
||||||
import copy
|
|
||||||
|
|
||||||
copied = copy.copy(message)
|
|
||||||
additional_kwargs = dict(getattr(message, "additional_kwargs", None) or {})
|
|
||||||
# Parsed tool_calls are canonical. Raw copies can otherwise re-introduce an
|
|
||||||
# invalid call after invalid_tool_calls has been cleared.
|
|
||||||
additional_kwargs.pop("tool_calls", None)
|
|
||||||
copied.additional_kwargs = additional_kwargs
|
|
||||||
copied.invalid_tool_calls = []
|
|
||||||
|
|
||||||
original_calls = list(getattr(message, "tool_calls", None) or [])
|
|
||||||
used_results: set[int] = set()
|
|
||||||
matched_calls: list[dict[str, Any]] = []
|
|
||||||
matched_result_indexes: list[int] = []
|
|
||||||
original_to_matched_call: dict[int, tuple[str, str]] = {}
|
|
||||||
|
|
||||||
for call_index, original_call in enumerate(original_calls):
|
|
||||||
call = dict(original_call)
|
|
||||||
call_id = str(call.get("id") or "")
|
|
||||||
call_name = str(call.get("name") or "").strip()
|
|
||||||
# A missing name is structurally unreplayable. Never infer it from
|
|
||||||
# arguments or retain its paired ToolMessage in provider history.
|
|
||||||
if not call_name:
|
|
||||||
continue
|
|
||||||
call["name"] = call_name
|
|
||||||
result_index = _tool_message_match_index(
|
|
||||||
tool_messages,
|
|
||||||
used_results,
|
|
||||||
call_id=call_id,
|
|
||||||
call_name=call_name,
|
|
||||||
)
|
|
||||||
# A historical client-side function call is only replayable together
|
|
||||||
# with its result. Incomplete calls are discarded instead of asking the
|
|
||||||
# provider to continue a broken tool turn.
|
|
||||||
if result_index is None:
|
|
||||||
continue
|
|
||||||
if not call_id:
|
|
||||||
result_id = str(
|
|
||||||
getattr(tool_messages[result_index], "tool_call_id", "") or ""
|
|
||||||
)
|
|
||||||
call_id = result_id or _stable_tool_call_id(
|
|
||||||
message, message_index, call_index
|
|
||||||
)
|
|
||||||
call["id"] = call_id
|
|
||||||
matched_calls.append(call)
|
|
||||||
matched_result_indexes.append(result_index)
|
|
||||||
original_to_matched_call[call_index] = (call_id, call_name)
|
|
||||||
used_results.add(result_index)
|
|
||||||
|
|
||||||
copied.tool_calls = matched_calls
|
|
||||||
if isinstance(copied.content, list):
|
|
||||||
original_call_index = 0
|
|
||||||
blocks: list[Any] = []
|
|
||||||
for original_block in copied.content:
|
|
||||||
if not isinstance(original_block, dict):
|
|
||||||
blocks.append(original_block)
|
|
||||||
continue
|
|
||||||
block = dict(original_block)
|
|
||||||
if block.get("type") in {"tool_call", "function_call"}:
|
|
||||||
matched_call = original_to_matched_call.get(original_call_index)
|
|
||||||
original_call_index += 1
|
|
||||||
if matched_call is None:
|
|
||||||
continue
|
|
||||||
call_id, call_name = matched_call
|
|
||||||
# LangChain content blocks use id; the Responses converter later
|
|
||||||
# maps it to call_id.
|
|
||||||
block["id"] = call_id
|
|
||||||
block["name"] = call_name
|
|
||||||
if isinstance(block.get("function"), dict):
|
|
||||||
block["function"] = {**block["function"], "name": call_name}
|
|
||||||
blocks.append(block)
|
|
||||||
copied.content = blocks
|
|
||||||
|
|
||||||
matched_results: list[Any] = []
|
|
||||||
result_to_call_id = {
|
|
||||||
result_index: matched_calls[index]["id"]
|
|
||||||
for index, result_index in enumerate(matched_result_indexes)
|
|
||||||
}
|
|
||||||
for result_index, result in enumerate(tool_messages):
|
|
||||||
call_id = result_to_call_id.get(result_index)
|
|
||||||
if call_id is None:
|
|
||||||
continue
|
|
||||||
copied_result = copy.copy(result)
|
|
||||||
copied_result.tool_call_id = call_id
|
|
||||||
matched_results.append(copied_result)
|
|
||||||
|
|
||||||
had_tool_protocol = bool(original_calls) or bool(
|
|
||||||
getattr(message, "invalid_tool_calls", None)
|
|
||||||
)
|
|
||||||
if not matched_calls and had_tool_protocol:
|
|
||||||
replayable_content = _flatten_message_content(copied.content)
|
|
||||||
if not replayable_content:
|
|
||||||
return None, matched_results
|
|
||||||
|
|
||||||
return copied, matched_results
|
|
||||||
|
|
||||||
|
|
||||||
def _sanitize_openai_tool_history(messages: list[Any]) -> list[Any]:
|
|
||||||
"""Copy history while retaining only complete, replayable tool turns."""
|
|
||||||
|
|
||||||
normalized: list[Any] = []
|
|
||||||
index = 0
|
|
||||||
while index < len(messages):
|
|
||||||
message = messages[index]
|
|
||||||
message_type = getattr(message, "type", None)
|
|
||||||
if message_type == "tool":
|
|
||||||
# A tool result without its immediately preceding assistant call is
|
|
||||||
# invalid for both Chat Completions and Responses APIs.
|
|
||||||
index += 1
|
|
||||||
continue
|
|
||||||
if message_type != "ai":
|
|
||||||
normalized.append(message)
|
|
||||||
index += 1
|
|
||||||
continue
|
|
||||||
|
|
||||||
next_index = index + 1
|
|
||||||
tool_messages: list[Any] = []
|
|
||||||
while (
|
|
||||||
next_index < len(messages)
|
|
||||||
and getattr(messages[next_index], "type", None) == "tool"
|
|
||||||
):
|
|
||||||
tool_messages.append(messages[next_index])
|
|
||||||
next_index += 1
|
|
||||||
copied, matched_results = _copy_ai_message_with_tool_pairs(
|
|
||||||
message,
|
|
||||||
index,
|
|
||||||
tool_messages,
|
|
||||||
)
|
|
||||||
if copied is not None:
|
|
||||||
normalized.append(copied)
|
|
||||||
normalized.extend(matched_results)
|
|
||||||
index = next_index
|
|
||||||
|
|
||||||
return normalized
|
|
||||||
|
|
||||||
|
|
||||||
def _ensure_openai_tool_call_ids(messages: list[Any]) -> list[Any]:
|
|
||||||
"""Backward-compatible alias for replay-safe tool history normalization."""
|
|
||||||
|
|
||||||
return _sanitize_openai_tool_history(messages)
|
|
||||||
|
|
||||||
|
|
||||||
def _has_assistant_tool_protocol(messages: list[Any]) -> bool:
|
|
||||||
"""Return whether history contains assistant-side tool protocol state."""
|
|
||||||
|
|
||||||
for message in messages:
|
|
||||||
if getattr(message, "type", None) != "ai":
|
|
||||||
continue
|
|
||||||
if getattr(message, "tool_calls", None) or getattr(
|
|
||||||
message, "invalid_tool_calls", None
|
|
||||||
):
|
|
||||||
return True
|
|
||||||
additional_kwargs = getattr(message, "additional_kwargs", None) or {}
|
|
||||||
if additional_kwargs.get("tool_calls"):
|
|
||||||
return True
|
|
||||||
content = getattr(message, "content", None)
|
|
||||||
if isinstance(content, list) and any(
|
|
||||||
isinstance(block, dict)
|
|
||||||
and block.get("type") in {"tool_call", "function_call"}
|
|
||||||
for block in content
|
|
||||||
):
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_openai_tool_history(messages: list[Any]) -> None:
|
|
||||||
"""Raise when sanitized history still contains an invalid tool protocol."""
|
|
||||||
|
|
||||||
available_call_ids: set[str] = set()
|
|
||||||
for message in messages:
|
|
||||||
message_type = getattr(message, "type", None)
|
|
||||||
if message_type == "ai":
|
|
||||||
if getattr(message, "invalid_tool_calls", None):
|
|
||||||
raise ValueError("invalid_tool_calls must not be replayed")
|
|
||||||
response_call_ids: set[str] = set()
|
|
||||||
response_calls: dict[str, str] = {}
|
|
||||||
for call in getattr(message, "tool_calls", None) or []:
|
|
||||||
call_name = str(call.get("name") or "").strip()
|
|
||||||
if not call_name:
|
|
||||||
raise ValueError("assistant tool call is missing a name")
|
|
||||||
call_id = str(call.get("id") or "").strip()
|
|
||||||
if not call_id:
|
|
||||||
raise ValueError("assistant tool call is missing an id")
|
|
||||||
if call_id in response_call_ids or call_id in available_call_ids:
|
|
||||||
raise ValueError(
|
|
||||||
"assistant tool call id is duplicated while outstanding"
|
|
||||||
)
|
|
||||||
response_call_ids.add(call_id)
|
|
||||||
available_call_ids.add(call_id)
|
|
||||||
response_calls[call_id] = call_name
|
|
||||||
content = getattr(message, "content", None)
|
|
||||||
content_call_ids: set[str] = set()
|
|
||||||
if isinstance(content, list):
|
|
||||||
for block in content:
|
|
||||||
if not isinstance(block, dict) or block.get("type") not in {
|
|
||||||
"tool_call",
|
|
||||||
"function_call",
|
|
||||||
}:
|
|
||||||
continue
|
|
||||||
block_id = str(
|
|
||||||
block.get("id") or block.get("call_id") or ""
|
|
||||||
).strip()
|
|
||||||
block_name = block.get("name") or block.get("tool_name")
|
|
||||||
function = block.get("function")
|
|
||||||
if not block_name and isinstance(function, dict):
|
|
||||||
block_name = function.get("name")
|
|
||||||
block_name = str(block_name or "").strip()
|
|
||||||
if (
|
|
||||||
not block_id
|
|
||||||
or block_id in content_call_ids
|
|
||||||
or response_calls.get(block_id) != block_name
|
|
||||||
):
|
|
||||||
raise ValueError(
|
|
||||||
"assistant content block does not match parsed tool call"
|
|
||||||
)
|
|
||||||
content_call_ids.add(block_id)
|
|
||||||
elif message_type == "tool":
|
|
||||||
call_id = str(getattr(message, "tool_call_id", "") or "")
|
|
||||||
if not call_id or call_id not in available_call_ids:
|
|
||||||
raise ValueError("tool result does not match a prior tool call")
|
|
||||||
available_call_ids.remove(call_id)
|
|
||||||
|
|
||||||
if available_call_ids:
|
|
||||||
raise ValueError("assistant tool call is missing its tool result")
|
|
||||||
|
|
||||||
|
|
||||||
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
|
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
|
||||||
"""Flatten list content for OpenAI-compatible APIs, preserving media.
|
"""Flatten list content for OpenAI-compatible APIs, preserving media.
|
||||||
|
|
||||||
@@ -581,9 +282,6 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
|
|||||||
|
|
||||||
from langchain_core.messages import HumanMessage
|
from langchain_core.messages import HumanMessage
|
||||||
|
|
||||||
sanitize_tool_history = _has_assistant_tool_protocol(messages)
|
|
||||||
if sanitize_tool_history:
|
|
||||||
messages = _sanitize_openai_tool_history(messages)
|
|
||||||
out: list[Any] = []
|
out: list[Any] = []
|
||||||
pending_media: list[Any] = [] # media hoisted out of a run of tool messages
|
pending_media: list[Any] = [] # media hoisted out of a run of tool messages
|
||||||
|
|
||||||
@@ -622,8 +320,6 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
|
|||||||
msg.content = flat
|
msg.content = flat
|
||||||
out.append(msg)
|
out.append(msg)
|
||||||
_flush() # conversation may end with tool messages
|
_flush() # conversation may end with tool messages
|
||||||
if sanitize_tool_history:
|
|
||||||
_validate_openai_tool_history(out)
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
@@ -1033,66 +729,6 @@ def _patch_openai_capture_reasoning_content() -> None:
|
|||||||
_patch_openai_capture_reasoning_content()
|
_patch_openai_capture_reasoning_content()
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Patch (module-level): silence langgraph_api's OpenAPI schema-generation
|
|
||||||
# warnings for endpoints whose docstrings aren't valid YAML.
|
|
||||||
#
|
|
||||||
# Upstream ``langgraph_api.utils.SchemaGenerator.get_schema`` calls
|
|
||||||
# ``parse_docstring`` (inherited from Starlette's ``BaseSchemaGenerator``)
|
|
||||||
# on every registered endpoint. When the docstring is prose with stray
|
|
||||||
# ``:`` characters, ``yaml.safe_load`` raises and upstream logs the
|
|
||||||
# failure + full traceback at WARNING level. It then falls back to
|
|
||||||
# ``{"description": docstring}`` — the endpoint still ends up in the
|
|
||||||
# schema with its prose as the description, just without structured
|
|
||||||
# ``parameters``/``responses``/``tags`` fields.
|
|
||||||
#
|
|
||||||
# The fallback path is fine; the warning + traceback is just noise. And
|
|
||||||
# it's only triggered for our deploy because mounting any custom Starlette
|
|
||||||
# app (``EvoScientist/langgraph_dev/http.py``) makes upstream call
|
|
||||||
# ``update_openapi_spec`` at startup — which iterates EVERY route,
|
|
||||||
# including upstream's own endpoints whose prose docstrings predate the
|
|
||||||
# YAML convention.
|
|
||||||
#
|
|
||||||
# Fix: wrap ``parse_docstring`` itself and absorb ``yaml.YAMLError`` by
|
|
||||||
# returning the same fallback shape upstream's except branch produces.
|
|
||||||
# Non-YAML exceptions are deliberately left to propagate — upstream's
|
|
||||||
# ``get_schema`` already catches them and logs WARNING + traceback, so
|
|
||||||
# unexpected failures remain debuggable. Patching ``parse_docstring`` (a
|
|
||||||
# small, stable method) instead of ``get_schema`` (the larger loop body)
|
|
||||||
# minimizes our exposure to upstream churn.
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
_langgraph_schema_silenced_patched = False
|
|
||||||
|
|
||||||
|
|
||||||
def _patch_langgraph_schema_generator_silence_warnings() -> None:
|
|
||||||
global _langgraph_schema_silenced_patched
|
|
||||||
if _langgraph_schema_silenced_patched:
|
|
||||||
return
|
|
||||||
try:
|
|
||||||
import langgraph_api.utils as _lgapi_utils
|
|
||||||
import yaml
|
|
||||||
|
|
||||||
_SchemaGenerator = _lgapi_utils.SchemaGenerator
|
|
||||||
_orig_parse_docstring = _SchemaGenerator.parse_docstring
|
|
||||||
|
|
||||||
def _patched_parse_docstring(self: Any, func: Any) -> dict[str, Any]:
|
|
||||||
try:
|
|
||||||
return _orig_parse_docstring(self, func)
|
|
||||||
except yaml.YAMLError:
|
|
||||||
return {"description": getattr(func, "__doc__", None) or ""}
|
|
||||||
|
|
||||||
_SchemaGenerator.parse_docstring = _patched_parse_docstring
|
|
||||||
_langgraph_schema_silenced_patched = True
|
|
||||||
except Exception:
|
|
||||||
# Patches are loader-safe: never crash the import. Silent failure
|
|
||||||
# here just leaves the upstream warnings visible in deploy logs,
|
|
||||||
# which is a benign fallback.
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
_patch_langgraph_schema_generator_silence_warnings()
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Patch (lazy, OpenRouter only): strip OpenAI-Responses encrypted reasoning
|
# Patch (lazy, OpenRouter only): strip OpenAI-Responses encrypted reasoning
|
||||||
# items from outgoing assistant messages.
|
# items from outgoing assistant messages.
|
||||||
@@ -1237,26 +873,25 @@ def _patch_deepseek_reasoning_passback(model: Any) -> None:
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Patch: forward CLI's live (model, model_provider) into deepagents'
|
# Patch: forward workspace scope and usage correlation context into
|
||||||
# start_async_task / update_async_task tool calls so the deployed graph
|
# deepagents' start_async_task / update_async_task tool calls so the deployed
|
||||||
# (running in a separate ``langgraph dev`` subprocess) re-resolves the
|
# graph (running in a separate ``langgraph dev`` subprocess) inherits the
|
||||||
# chat model per run.
|
# parent run's scoping and accounting metadata.
|
||||||
#
|
#
|
||||||
# Without this, async sub-agents stay on the model their graph was compiled
|
# Model configuration is deliberately NOT forwarded: the deployed graph
|
||||||
# with at langgraph dev boot — `/model` switches in the CLI never reach
|
# resolves its chat model per run from ``configurable.runtime_snapshot_id``
|
||||||
# them because they live in another process.
|
# 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``
|
# Mechanism: wrap deepagents' ``_build_start_tool`` and ``_build_update_tool``
|
||||||
# factories. Each wrapped factory calls the original with a proxied client
|
# factories. Each wrapped factory calls the original with a proxied client
|
||||||
# cache that intercepts ``runs.create(...)`` calls only and injects
|
# cache that intercepts ``runs.create(...)`` calls only and merges the
|
||||||
# ``config={"configurable": {"model": <cfg.model>, "model_provider": <cfg.provider>}}``.
|
# inherited scope/usage context into ``config``/``metadata``. All other
|
||||||
# All other client methods (``threads.create``, ``runs.get``, ``runs.cancel``,
|
# client methods (``threads.create``, ``runs.get``, ``runs.cancel``,
|
||||||
# ``runs.join_stream``) pass through unchanged. The deployed graph picks up
|
# ``runs.join_stream``) pass through unchanged.
|
||||||
# ``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.
|
|
||||||
#
|
#
|
||||||
# Upstream PR opportunity: passing ``config`` through ``client.runs.create``
|
# Upstream PR opportunity: passing ``config`` through ``client.runs.create``
|
||||||
# is generic functionality; worth contributing back to ``langchain-ai/deepagents``
|
# is generic functionality; worth contributing back to ``langchain-ai/deepagents``
|
||||||
@@ -1265,49 +900,152 @@ def _patch_deepseek_reasoning_passback(model: Any) -> None:
|
|||||||
_model_passthrough_patched = False
|
_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:
|
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
|
Preserves any caller-supplied ``config.configurable`` keys.
|
||||||
keys take precedence on conflict (callers shouldn't be passing model
|
|
||||||
overrides — the CLI is the source of truth).
|
|
||||||
"""
|
"""
|
||||||
overrides = _read_cfg_configurable()
|
usage_enabled = os.getenv("EVOSCIENTIST_USAGE_TRACKING", "").strip().lower() in {
|
||||||
if not overrides:
|
"1",
|
||||||
return kwargs
|
"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")
|
existing = kwargs.get("config")
|
||||||
if not isinstance(existing, dict):
|
if not isinstance(existing, dict):
|
||||||
existing = {}
|
existing = {}
|
||||||
existing_configurable = existing.get("configurable")
|
existing_configurable = existing.get("configurable")
|
||||||
if not isinstance(existing_configurable, dict):
|
if not isinstance(existing_configurable, dict):
|
||||||
existing_configurable = {}
|
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 = dict(kwargs)
|
||||||
|
if merged_configurable or "config" in kwargs:
|
||||||
kwargs["config"] = {**existing, "configurable": merged_configurable}
|
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
|
return kwargs
|
||||||
|
|
||||||
|
|
||||||
@@ -1381,7 +1119,7 @@ class _ClientCacheProxy:
|
|||||||
|
|
||||||
|
|
||||||
def _patch_deepagents_model_passthrough() -> None:
|
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
|
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
|
call from ``_maybe_swap_async_subagents`` on every CLI startup; both
|
||||||
|
|||||||
@@ -1,302 +0,0 @@
|
|||||||
"""Shared logging configuration helpers."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
import os
|
|
||||||
import sys
|
|
||||||
from datetime import UTC, datetime
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, TextIO
|
|
||||||
|
|
||||||
DEFAULT_LOG_RETENTION_DAYS = 30
|
|
||||||
DEFAULT_LOG_FORMAT = "%(asctime)s [%(levelname)s] %(name)s: %(message)s"
|
|
||||||
DEFAULT_LOG_DATE_FORMAT = "%Y-%m-%d %H:%M:%S"
|
|
||||||
MANAGED_HANDLER_ATTR = "_evoscientist_managed_handler"
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_log_level(level: int | str | None, default: int = logging.INFO) -> int:
|
|
||||||
"""Resolve a logging level from config or environment input."""
|
|
||||||
if isinstance(level, int):
|
|
||||||
return level
|
|
||||||
raw = str(level or "").strip()
|
|
||||||
if not raw:
|
|
||||||
return default
|
|
||||||
if raw.isdigit():
|
|
||||||
return int(raw)
|
|
||||||
normalized = raw.upper()
|
|
||||||
if normalized == "WARN":
|
|
||||||
normalized = "WARNING"
|
|
||||||
resolved = logging.getLevelNamesMapping().get(normalized)
|
|
||||||
return resolved if isinstance(resolved, int) else default
|
|
||||||
|
|
||||||
|
|
||||||
def _mark_managed(handler: logging.Handler, kind: str) -> logging.Handler:
|
|
||||||
setattr(handler, MANAGED_HANDLER_ATTR, kind)
|
|
||||||
return handler
|
|
||||||
|
|
||||||
|
|
||||||
def _managed_kind(handler: logging.Handler) -> str | None:
|
|
||||||
kind = getattr(handler, MANAGED_HANDLER_ATTR, None)
|
|
||||||
return kind if isinstance(kind, str) else None
|
|
||||||
|
|
||||||
|
|
||||||
def remove_managed_handlers(
|
|
||||||
logger: logging.Logger | None = None,
|
|
||||||
*,
|
|
||||||
kinds: set[str] | None = None,
|
|
||||||
) -> None:
|
|
||||||
"""Remove handlers installed by this module without touching external ones."""
|
|
||||||
target = logger or logging.getLogger()
|
|
||||||
for handler in target.handlers[:]:
|
|
||||||
kind = _managed_kind(handler)
|
|
||||||
if kind and (kinds is None or kind in kinds):
|
|
||||||
target.removeHandler(handler)
|
|
||||||
handler.close()
|
|
||||||
|
|
||||||
|
|
||||||
def _standard_formatter() -> logging.Formatter:
|
|
||||||
return logging.Formatter(DEFAULT_LOG_FORMAT, datefmt=DEFAULT_LOG_DATE_FORMAT)
|
|
||||||
|
|
||||||
|
|
||||||
class DailyLogFileHandler(logging.FileHandler):
|
|
||||||
"""File handler that writes the active log to a date-based filename."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
log_dir: str | Path,
|
|
||||||
*,
|
|
||||||
prefix: str = "evoscientist",
|
|
||||||
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
|
|
||||||
encoding: str = "utf-8",
|
|
||||||
utc: bool = False,
|
|
||||||
) -> None:
|
|
||||||
self.log_dir = Path(log_dir).expanduser()
|
|
||||||
self.prefix = prefix
|
|
||||||
self.retention_days = max(1, retention_days)
|
|
||||||
self.utc = utc
|
|
||||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
super().__init__(self._dated_log_path(), encoding=encoding, delay=True)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def active_log_path(self) -> Path:
|
|
||||||
"""Return the active log path for the current date."""
|
|
||||||
return self._dated_log_path()
|
|
||||||
|
|
||||||
def _dated_log_path(self) -> Path:
|
|
||||||
now = datetime.now(UTC if self.utc else None)
|
|
||||||
return self.log_dir / f"{self.prefix}-{now:%Y-%m-%d}.log"
|
|
||||||
|
|
||||||
def emit(self, record: logging.LogRecord) -> None:
|
|
||||||
try:
|
|
||||||
expected = str(self.active_log_path)
|
|
||||||
if self.baseFilename != expected:
|
|
||||||
if self.stream:
|
|
||||||
self.stream.close()
|
|
||||||
self.stream = None
|
|
||||||
self.baseFilename = expected
|
|
||||||
self._delete_expired_logs()
|
|
||||||
super().emit(record)
|
|
||||||
except OSError:
|
|
||||||
self.handleError(record)
|
|
||||||
|
|
||||||
def getFilesToDelete(self) -> list[str]:
|
|
||||||
candidates = sorted(self.log_dir.glob(f"{self.prefix}-????-??-??.log"))
|
|
||||||
if len(candidates) <= self.retention_days:
|
|
||||||
return []
|
|
||||||
return [str(path) for path in candidates[: -self.retention_days]]
|
|
||||||
|
|
||||||
def _delete_expired_logs(self) -> None:
|
|
||||||
for path in self.getFilesToDelete():
|
|
||||||
try:
|
|
||||||
os.remove(path)
|
|
||||||
except OSError:
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
def default_log_dir() -> Path:
|
|
||||||
"""Return the default runtime log directory."""
|
|
||||||
env_dir = os.environ.get("EVOSCIENTIST_LOG_DIR")
|
|
||||||
if env_dir:
|
|
||||||
return Path(env_dir).expanduser()
|
|
||||||
|
|
||||||
from EvoScientist.paths import DATA_DIR
|
|
||||||
|
|
||||||
return DATA_DIR / "logs"
|
|
||||||
|
|
||||||
|
|
||||||
def configure_daily_file_logging(
|
|
||||||
logger: logging.Logger | None = None,
|
|
||||||
*,
|
|
||||||
log_dir: str | Path | None = None,
|
|
||||||
prefix: str = "evoscientist",
|
|
||||||
level: int | str = logging.INFO,
|
|
||||||
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
|
|
||||||
) -> DailyLogFileHandler:
|
|
||||||
"""Attach a daily file handler, replacing older matching handlers."""
|
|
||||||
target = logger or logging.getLogger()
|
|
||||||
resolved_level = resolve_log_level(level, default=logging.INFO)
|
|
||||||
retention_days = max(1, int(retention_days))
|
|
||||||
resolved_dir = Path(log_dir).expanduser() if log_dir else default_log_dir()
|
|
||||||
|
|
||||||
for handler in target.handlers[:]:
|
|
||||||
if (
|
|
||||||
isinstance(handler, DailyLogFileHandler)
|
|
||||||
and handler.prefix == prefix
|
|
||||||
and handler.log_dir == resolved_dir
|
|
||||||
):
|
|
||||||
target.removeHandler(handler)
|
|
||||||
handler.close()
|
|
||||||
|
|
||||||
handler = DailyLogFileHandler(
|
|
||||||
resolved_dir,
|
|
||||||
prefix=prefix,
|
|
||||||
retention_days=retention_days,
|
|
||||||
)
|
|
||||||
_mark_managed(handler, "file")
|
|
||||||
handler.setLevel(resolved_level)
|
|
||||||
handler.setFormatter(_standard_formatter())
|
|
||||||
target.addHandler(handler)
|
|
||||||
if target.level == logging.NOTSET or target.level > resolved_level:
|
|
||||||
target.setLevel(resolved_level)
|
|
||||||
return handler
|
|
||||||
|
|
||||||
|
|
||||||
def configure_console_logging(
|
|
||||||
logger: logging.Logger | None = None,
|
|
||||||
*,
|
|
||||||
level: int | str | None = logging.INFO,
|
|
||||||
stream: TextIO | None = None,
|
|
||||||
replace: bool = True,
|
|
||||||
) -> logging.StreamHandler:
|
|
||||||
"""Attach a standard console handler for non-interactive entry points."""
|
|
||||||
target = logger or logging.getLogger()
|
|
||||||
resolved_level = resolve_log_level(level, default=logging.INFO)
|
|
||||||
if replace:
|
|
||||||
remove_managed_handlers(target, kinds={"console", "rich"})
|
|
||||||
|
|
||||||
handler = logging.StreamHandler(stream or sys.stderr)
|
|
||||||
_mark_managed(handler, "console")
|
|
||||||
handler.setLevel(resolved_level)
|
|
||||||
handler.setFormatter(_standard_formatter())
|
|
||||||
target.addHandler(handler)
|
|
||||||
target.setLevel(resolved_level)
|
|
||||||
return handler
|
|
||||||
|
|
||||||
|
|
||||||
def configure_rich_console_logging(
|
|
||||||
logger: logging.Logger | None = None,
|
|
||||||
*,
|
|
||||||
level: int | str | None = logging.INFO,
|
|
||||||
console: Any = None,
|
|
||||||
replace: bool = True,
|
|
||||||
dim_warnings: bool = False,
|
|
||||||
show_time: bool | None = None,
|
|
||||||
show_path: bool | None = None,
|
|
||||||
show_level: bool | None = None,
|
|
||||||
) -> logging.Handler:
|
|
||||||
"""Attach a Rich console handler for interactive CLI output."""
|
|
||||||
from rich.logging import RichHandler
|
|
||||||
from rich.markup import escape
|
|
||||||
|
|
||||||
target = logger or logging.getLogger()
|
|
||||||
resolved_level = resolve_log_level(level, default=logging.INFO)
|
|
||||||
verbose = resolved_level <= logging.DEBUG
|
|
||||||
if replace:
|
|
||||||
remove_managed_handlers(target, kinds={"console", "rich"})
|
|
||||||
|
|
||||||
class DimWarningHandler(RichHandler):
|
|
||||||
def emit(self, record: logging.LogRecord) -> None:
|
|
||||||
if dim_warnings and record.levelno == logging.WARNING and console is not None:
|
|
||||||
console.print(
|
|
||||||
"[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] "
|
|
||||||
f"[dim]{escape(record.getMessage())}[/dim]"
|
|
||||||
)
|
|
||||||
return
|
|
||||||
super().emit(record)
|
|
||||||
|
|
||||||
handler = DimWarningHandler(
|
|
||||||
console=console,
|
|
||||||
show_time=verbose if show_time is None else show_time,
|
|
||||||
show_path=verbose if show_path is None else show_path,
|
|
||||||
show_level=verbose if show_level is None else show_level,
|
|
||||||
)
|
|
||||||
_mark_managed(handler, "rich")
|
|
||||||
handler.setLevel(resolved_level)
|
|
||||||
target.addHandler(handler)
|
|
||||||
target.setLevel(resolved_level)
|
|
||||||
return handler
|
|
||||||
|
|
||||||
|
|
||||||
def configure_logging(
|
|
||||||
logger: logging.Logger | None = None,
|
|
||||||
*,
|
|
||||||
level: int | str | None = logging.INFO,
|
|
||||||
log_dir: str | Path | None = None,
|
|
||||||
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
|
|
||||||
prefix: str = "evoscientist",
|
|
||||||
console: bool = True,
|
|
||||||
file: bool = True,
|
|
||||||
replace_managed: bool = True,
|
|
||||||
) -> list[logging.Handler]:
|
|
||||||
"""Configure standard EvoScientist console and daily file logging."""
|
|
||||||
target = logger or logging.getLogger()
|
|
||||||
resolved_level = resolve_log_level(level, default=logging.INFO)
|
|
||||||
if replace_managed:
|
|
||||||
remove_managed_handlers(target, kinds={"console", "rich", "file"})
|
|
||||||
|
|
||||||
handlers: list[logging.Handler] = []
|
|
||||||
if console:
|
|
||||||
handlers.append(
|
|
||||||
configure_console_logging(target, level=resolved_level, replace=False)
|
|
||||||
)
|
|
||||||
if file:
|
|
||||||
handlers.append(
|
|
||||||
configure_daily_file_logging(
|
|
||||||
target,
|
|
||||||
log_dir=log_dir,
|
|
||||||
prefix=prefix,
|
|
||||||
level=resolved_level,
|
|
||||||
retention_days=retention_days,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
target.setLevel(resolved_level)
|
|
||||||
return handlers
|
|
||||||
|
|
||||||
|
|
||||||
def configure_logging_from_settings(
|
|
||||||
logger: logging.Logger | None = None,
|
|
||||||
*,
|
|
||||||
default_level: int = logging.INFO,
|
|
||||||
prefix: str = "evoscientist",
|
|
||||||
console: bool = True,
|
|
||||||
file: bool = True,
|
|
||||||
) -> list[logging.Handler]:
|
|
||||||
"""Configure logging from EvoScientist settings and environment overrides."""
|
|
||||||
level: int | str | None = os.environ.get("EVOSCIENTIST_LOG_LEVEL")
|
|
||||||
log_dir: str | Path | None = os.environ.get("EVOSCIENTIST_LOG_DIR") or None
|
|
||||||
retention_days = int(
|
|
||||||
os.environ.get("EVOSCIENTIST_LOG_RETENTION_DAYS", DEFAULT_LOG_RETENTION_DAYS)
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
from EvoScientist.config import get_effective_config
|
|
||||||
|
|
||||||
cfg = get_effective_config()
|
|
||||||
level = level or getattr(cfg, "log_level", None)
|
|
||||||
log_dir = log_dir or getattr(cfg, "log_dir", None) or None
|
|
||||||
retention_days = int(
|
|
||||||
getattr(cfg, "log_retention_days", DEFAULT_LOG_RETENTION_DAYS)
|
|
||||||
)
|
|
||||||
except Exception:
|
|
||||||
level = level or default_level
|
|
||||||
|
|
||||||
return configure_logging(
|
|
||||||
logger,
|
|
||||||
level=resolve_log_level(level, default=default_level),
|
|
||||||
log_dir=log_dir,
|
|
||||||
retention_days=retention_days,
|
|
||||||
prefix=prefix,
|
|
||||||
console=console,
|
|
||||||
file=file,
|
|
||||||
)
|
|
||||||
@@ -10,7 +10,6 @@ from .client import (
|
|||||||
build_mcp_add_kwargs,
|
build_mcp_add_kwargs,
|
||||||
build_mcp_edit_fields,
|
build_mcp_edit_fields,
|
||||||
edit_mcp_server,
|
edit_mcp_server,
|
||||||
get_mcp_server_errors,
|
|
||||||
load_mcp_config,
|
load_mcp_config,
|
||||||
load_mcp_tools,
|
load_mcp_tools,
|
||||||
parse_mcp_add_args,
|
parse_mcp_add_args,
|
||||||
@@ -39,7 +38,6 @@ __all__ = [
|
|||||||
"find_server_by_name",
|
"find_server_by_name",
|
||||||
"get_all_tags",
|
"get_all_tags",
|
||||||
"get_installed_names",
|
"get_installed_names",
|
||||||
"get_mcp_server_errors",
|
|
||||||
"install_mcp_server",
|
"install_mcp_server",
|
||||||
"install_mcp_servers",
|
"install_mcp_servers",
|
||||||
"load_mcp_config",
|
"load_mcp_config",
|
||||||
|
|||||||
@@ -114,10 +114,6 @@ _URL_TRANSPORTS = {"http", "streamable_http", "sse", "websocket"}
|
|||||||
# still parallelizing the common 3–7 server case to completion.
|
# still parallelizing the common 3–7 server case to completion.
|
||||||
_MAX_CONCURRENT_CONNECTIONS = 8
|
_MAX_CONCURRENT_CONNECTIONS = 8
|
||||||
|
|
||||||
# Last connection error per configured server. This is process-local runtime
|
|
||||||
# diagnostics for the Web/CLI status surfaces, not persisted configuration.
|
|
||||||
_MCP_SERVER_ERRORS: dict[str, str] = {}
|
|
||||||
|
|
||||||
# Env vars forwarded to stdio MCP subprocesses on top of the MCP SDK's
|
# Env vars forwarded to stdio MCP subprocesses on top of the MCP SDK's
|
||||||
# minimal default set (HOME/PATH/USER/…). Without this, servers behind
|
# minimal default set (HOME/PATH/USER/…). Without this, servers behind
|
||||||
# a proxy or with a custom CA bundle silently fail with long timeouts.
|
# a proxy or with a custom CA bundle silently fail with long timeouts.
|
||||||
@@ -768,9 +764,6 @@ async def _load_tools(
|
|||||||
if not connections:
|
if not connections:
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
for stale_name in set(_MCP_SERVER_ERRORS) - set(connections):
|
|
||||||
_MCP_SERVER_ERRORS.pop(stale_name, None)
|
|
||||||
|
|
||||||
client = MultiServerMCPClient(connections) # type: ignore[invalid-argument-type]
|
client = MultiServerMCPClient(connections) # type: ignore[invalid-argument-type]
|
||||||
|
|
||||||
def _report(event: str, name: str, detail: str = "") -> None:
|
def _report(event: str, name: str, detail: str = "") -> None:
|
||||||
@@ -794,13 +787,10 @@ async def _load_tools(
|
|||||||
_report("start", name)
|
_report("start", name)
|
||||||
try:
|
try:
|
||||||
tools = await client.get_tools(server_name=name)
|
tools = await client.get_tools(server_name=name)
|
||||||
_MCP_SERVER_ERRORS.pop(name, None)
|
|
||||||
logger.info("MCP server %r: loaded %d tool(s)", name, len(tools))
|
logger.info("MCP server %r: loaded %d tool(s)", name, len(tools))
|
||||||
_report("success", name, str(len(tools)))
|
_report("success", name, str(len(tools)))
|
||||||
return name, tools
|
return name, tools
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
detail = str(exc) or type(exc).__name__
|
|
||||||
_MCP_SERVER_ERRORS[name] = detail
|
|
||||||
# When the caller wired up ``on_progress`` they own the
|
# When the caller wired up ``on_progress`` they own the
|
||||||
# user-facing display; downgrade the logger so we don't
|
# user-facing display; downgrade the logger so we don't
|
||||||
# double-print.
|
# double-print.
|
||||||
@@ -808,7 +798,7 @@ async def _load_tools(
|
|||||||
logger.warning("MCP server %r: failed to load tools: %s", name, exc)
|
logger.warning("MCP server %r: failed to load tools: %s", name, exc)
|
||||||
else:
|
else:
|
||||||
logger.debug("MCP server %r: failed to load tools: %s", name, exc)
|
logger.debug("MCP server %r: failed to load tools: %s", name, exc)
|
||||||
_report("error", name, detail)
|
_report("error", name, str(exc))
|
||||||
return name, []
|
return name, []
|
||||||
|
|
||||||
# ``return_exceptions=False`` is fine because ``_fetch`` already
|
# ``return_exceptions=False`` is fine because ``_fetch`` already
|
||||||
@@ -817,11 +807,6 @@ async def _load_tools(
|
|||||||
return dict(results)
|
return dict(results)
|
||||||
|
|
||||||
|
|
||||||
def get_mcp_server_errors() -> dict[str, str]:
|
|
||||||
"""Return a snapshot of the most recent per-server connection errors."""
|
|
||||||
return dict(_MCP_SERVER_ERRORS)
|
|
||||||
|
|
||||||
|
|
||||||
async def aload_mcp_tools(
|
async def aload_mcp_tools(
|
||||||
config: dict[str, Any] | None = None,
|
config: dict[str, Any] | None = None,
|
||||||
*,
|
*,
|
||||||
|
|||||||
@@ -87,7 +87,8 @@ def build_memory_agent_graph(
|
|||||||
from deepagents import create_deep_agent
|
from deepagents import create_deep_agent
|
||||||
|
|
||||||
from ...backends import build_memory_agent_backend
|
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] = {}
|
kwargs: dict[str, Any] = {}
|
||||||
if response_format is not None:
|
if response_format is not None:
|
||||||
@@ -101,11 +102,14 @@ def build_memory_agent_graph(
|
|||||||
|
|
||||||
agent = create_deep_agent(
|
agent = create_deep_agent(
|
||||||
name=name,
|
name=name,
|
||||||
model=_ensure_auxiliary_chat_model(),
|
model=_compile_time_role_model("primary"),
|
||||||
system_prompt=system_prompt,
|
system_prompt=system_prompt,
|
||||||
tools=list(tools),
|
tools=list(tools),
|
||||||
backend=backend,
|
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=[],
|
subagents=[],
|
||||||
skills=skills,
|
skills=skills,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
|
|||||||
@@ -428,7 +428,6 @@ def _memory_worker_middleware(
|
|||||||
enable_observation_memory: bool = True,
|
enable_observation_memory: bool = True,
|
||||||
):
|
):
|
||||||
"""Build middleware for memory workers, excluding task execution tools."""
|
"""Build middleware for memory workers, excluding task execution tools."""
|
||||||
from ...middleware.error_normalization import ErrorNormalizationMiddleware
|
|
||||||
from ...middleware.memory import create_memory_middleware
|
from ...middleware.memory import create_memory_middleware
|
||||||
|
|
||||||
memory_controls = MemoryControls(
|
memory_controls = MemoryControls(
|
||||||
@@ -440,11 +439,7 @@ def _memory_worker_middleware(
|
|||||||
enable_observation_tool = memory_controls.observation_tool_enabled(
|
enable_observation_tool = memory_controls.observation_tool_enabled(
|
||||||
_memory_worker_observation_target(source_type)
|
_memory_worker_observation_target(source_type)
|
||||||
)
|
)
|
||||||
return [
|
return memory_agent_middleware(
|
||||||
# Outermost — normalize provider-SDK exceptions from the
|
|
||||||
# auxiliary model call before any inner middleware sees them.
|
|
||||||
ErrorNormalizationMiddleware(),
|
|
||||||
*memory_agent_middleware(
|
|
||||||
create_memory_middleware(
|
create_memory_middleware(
|
||||||
str(memory_dir),
|
str(memory_dir),
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
@@ -455,8 +450,7 @@ def _memory_worker_middleware(
|
|||||||
enable_observation_tool=enable_observation_tool,
|
enable_observation_tool=enable_observation_tool,
|
||||||
),
|
),
|
||||||
excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS,
|
excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS,
|
||||||
),
|
)
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
def _build_memory_worker_agent(
|
def _build_memory_worker_agent(
|
||||||
|
|||||||
@@ -71,8 +71,6 @@ def build_observation_linker_graph(
|
|||||||
workspace_dir: str | Path | None = None,
|
workspace_dir: str | Path | None = None,
|
||||||
) -> CompiledStateGraph:
|
) -> CompiledStateGraph:
|
||||||
"""Build the registered LangGraph observation linker."""
|
"""Build the registered LangGraph observation linker."""
|
||||||
from ...middleware.error_normalization import ErrorNormalizationMiddleware
|
|
||||||
|
|
||||||
agent_paths = resolve_memory_agent_paths(
|
agent_paths = resolve_memory_agent_paths(
|
||||||
memory_dir=memory_dir,
|
memory_dir=memory_dir,
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
@@ -87,7 +85,5 @@ def build_observation_linker_graph(
|
|||||||
tools=tools,
|
tools=tools,
|
||||||
memory_dir=agent_paths.memory_dir,
|
memory_dir=agent_paths.memory_dir,
|
||||||
workspace_dir=agent_paths.workspace_dir,
|
workspace_dir=agent_paths.workspace_dir,
|
||||||
# Outermost — normalize provider-SDK exceptions from the
|
middleware=memory_agent_middleware(),
|
||||||
# auxiliary model call before any inner middleware sees them.
|
|
||||||
middleware=[ErrorNormalizationMiddleware(), *memory_agent_middleware()],
|
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -3,6 +3,8 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import cast
|
from typing import cast
|
||||||
@@ -17,6 +19,7 @@ from ..gateway.background_runs import (
|
|||||||
launch_background_run,
|
launch_background_run,
|
||||||
)
|
)
|
||||||
from ..langgraph_dev.sdk import messages_input
|
from ..langgraph_dev.sdk import messages_input
|
||||||
|
from ..model_registry.schemas import ModelRef, ReasoningEffort
|
||||||
from .observations import build_observation_linker_index_context
|
from .observations import build_observation_linker_index_context
|
||||||
from .scheduler import ObservationLinkerContext
|
from .scheduler import ObservationLinkerContext
|
||||||
from .source_context import MemorySourceContext, _trajectory_for_prompt
|
from .source_context import MemorySourceContext, _trajectory_for_prompt
|
||||||
@@ -35,6 +38,8 @@ from .worker_activity import (
|
|||||||
snapshot_observation_relations,
|
snapshot_observation_relations,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
SUBAGENT_MEMORY_WORKER_GRAPH_ID = "evomemory-subagent-worker"
|
SUBAGENT_MEMORY_WORKER_GRAPH_ID = "evomemory-subagent-worker"
|
||||||
TURN_MEMORY_WORKER_GRAPH_ID = "evomemory-turn-worker"
|
TURN_MEMORY_WORKER_GRAPH_ID = "evomemory-turn-worker"
|
||||||
OBSERVATION_LINKER_GRAPH_ID = "evomemory-observation-linker"
|
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]:
|
def _memory_worker_metadata(context: MemorySourceContext) -> dict[str, str]:
|
||||||
return {
|
metadata = {
|
||||||
"run_kind": f"evomemory_{context.source_type.value}_worker",
|
"run_kind": f"evomemory_{context.source_type.value}_worker",
|
||||||
"source_session_id": context.session_id,
|
"source_session_id": context.session_id,
|
||||||
"source_agent": context.source_agent,
|
"source_agent": context.source_agent,
|
||||||
@@ -98,6 +103,114 @@ def _memory_worker_metadata(context: MemorySourceContext) -> dict[str, str]:
|
|||||||
"trajectory_digest": context.trajectory_digest,
|
"trajectory_digest": context.trajectory_digest,
|
||||||
"workspace_dir": _worker_workspace_dir(context.workspace_dir),
|
"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(
|
def _memory_worker_run_payload(
|
||||||
@@ -107,19 +220,25 @@ def _memory_worker_run_payload(
|
|||||||
) -> BackgroundRunPayload:
|
) -> BackgroundRunPayload:
|
||||||
"""Build the LangGraph SDK run payload for a memory worker."""
|
"""Build the LangGraph SDK run payload for a memory worker."""
|
||||||
metadata = _memory_worker_metadata(context)
|
metadata = _memory_worker_metadata(context)
|
||||||
payload: BackgroundRunPayload = {
|
configurable = {
|
||||||
"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,
|
"thread_id": thread_id,
|
||||||
"evomemory_source_session_id": context.session_id,
|
"evomemory_source_session_id": context.session_id,
|
||||||
"evomemory_source_agent": context.source_agent,
|
"evomemory_source_agent": context.source_agent,
|
||||||
"evomemory_project_id": context.project_id,
|
"evomemory_project_id": context.project_id,
|
||||||
"evomemory_trajectory_digest": context.trajectory_digest,
|
"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": configurable},
|
||||||
}
|
}
|
||||||
return _runs_create_kwargs(payload)
|
return _runs_create_kwargs(payload)
|
||||||
|
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ class MemorySourceContext:
|
|||||||
session_id: str
|
session_id: str
|
||||||
trajectory: list[CompactMessage]
|
trajectory: list[CompactMessage]
|
||||||
trajectory_digest: str
|
trajectory_digest: str
|
||||||
|
turn_id: str | None = None
|
||||||
|
|
||||||
|
|
||||||
def _task_tool_call_ids(messages: list[BaseMessage]) -> set[str]:
|
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
|
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:
|
def _short_hash(text: str) -> str:
|
||||||
"""Return the short hash fragment used in generated ids."""
|
"""Return the short hash fragment used in generated ids."""
|
||||||
return hashlib.sha256(text.encode("utf-8")).hexdigest()[:16]
|
return hashlib.sha256(text.encode("utf-8")).hexdigest()[:16]
|
||||||
@@ -240,6 +255,7 @@ def build_memory_source_context(
|
|||||||
project_id=project_id,
|
project_id=project_id,
|
||||||
source_agent=source_agent,
|
source_agent=source_agent,
|
||||||
session_id=session_id,
|
session_id=session_id,
|
||||||
|
turn_id=_active_turn_id(),
|
||||||
trajectory=trajectory,
|
trajectory=trajectory,
|
||||||
trajectory_digest=_trajectory_digest(trajectory),
|
trajectory_digest=_trajectory_digest(trajectory),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -18,7 +18,6 @@ from .context_editing import (
|
|||||||
create_context_editing_middleware,
|
create_context_editing_middleware,
|
||||||
)
|
)
|
||||||
from .context_overflow import ContextOverflowMapperMiddleware
|
from .context_overflow import ContextOverflowMapperMiddleware
|
||||||
from .error_normalization import ErrorNormalizationMiddleware
|
|
||||||
from .memory import (
|
from .memory import (
|
||||||
EvoMemoryMiddleware,
|
EvoMemoryMiddleware,
|
||||||
create_memory_middleware,
|
create_memory_middleware,
|
||||||
@@ -28,12 +27,9 @@ from .memory_lifecycle import (
|
|||||||
create_memory_lifecycle_middleware,
|
create_memory_lifecycle_middleware,
|
||||||
default_memory_scheduler,
|
default_memory_scheduler,
|
||||||
)
|
)
|
||||||
from .model_fallback import ModelFallbackMiddleware, load_fallback_chain
|
from .message_budget import (
|
||||||
from .repetitive_tool_guard import (
|
count_message_text_tokens,
|
||||||
DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS,
|
create_message_budget_middleware,
|
||||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
|
||||||
RepetitiveToolCallGuardMiddleware,
|
|
||||||
collapse_repetitive_tool_rounds,
|
|
||||||
)
|
)
|
||||||
from .runtime_context import RuntimeContextMiddleware, create_runtime_context_middleware
|
from .runtime_context import RuntimeContextMiddleware, create_runtime_context_middleware
|
||||||
from .scheduler import (
|
from .scheduler import (
|
||||||
@@ -41,39 +37,39 @@ from .scheduler import (
|
|||||||
create_scheduler_middleware,
|
create_scheduler_middleware,
|
||||||
)
|
)
|
||||||
from .tool_error_handler import ToolErrorHandlerMiddleware
|
from .tool_error_handler import ToolErrorHandlerMiddleware
|
||||||
from .tool_protocol_guard import ToolProtocolGuardMiddleware
|
|
||||||
from .tool_selector import create_tool_selector_middleware
|
from .tool_selector import create_tool_selector_middleware
|
||||||
from .utils import disable_thinking
|
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__ = [
|
__all__ = [
|
||||||
"DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS",
|
|
||||||
"DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD",
|
|
||||||
"AskUserMiddleware",
|
"AskUserMiddleware",
|
||||||
"AskUserRequest",
|
"AskUserRequest",
|
||||||
"AskUserWidgetResult",
|
"AskUserWidgetResult",
|
||||||
"Choice",
|
"Choice",
|
||||||
"ConfigurableModelMiddleware",
|
"ConfigurableModelMiddleware",
|
||||||
"ContextOverflowMapperMiddleware",
|
"ContextOverflowMapperMiddleware",
|
||||||
"ErrorNormalizationMiddleware",
|
|
||||||
"EvoMemoryLifecycleMiddleware",
|
"EvoMemoryLifecycleMiddleware",
|
||||||
"EvoMemoryMiddleware",
|
"EvoMemoryMiddleware",
|
||||||
"ModelFallbackMiddleware",
|
|
||||||
"Question",
|
"Question",
|
||||||
"RepetitiveToolCallGuardMiddleware",
|
|
||||||
"RuntimeContextMiddleware",
|
"RuntimeContextMiddleware",
|
||||||
"SchedulerMiddleware",
|
"SchedulerMiddleware",
|
||||||
"ToolErrorHandlerMiddleware",
|
"ToolErrorHandlerMiddleware",
|
||||||
"ToolProtocolGuardMiddleware",
|
|
||||||
"collapse_repetitive_tool_rounds",
|
|
||||||
"compute_context_editing_trigger",
|
"compute_context_editing_trigger",
|
||||||
|
"count_message_text_tokens",
|
||||||
"create_code_interpreter_middleware",
|
"create_code_interpreter_middleware",
|
||||||
"create_context_editing_middleware",
|
"create_context_editing_middleware",
|
||||||
"create_memory_lifecycle_middleware",
|
"create_memory_lifecycle_middleware",
|
||||||
"create_memory_middleware",
|
"create_memory_middleware",
|
||||||
|
"create_message_budget_middleware",
|
||||||
"create_runtime_context_middleware",
|
"create_runtime_context_middleware",
|
||||||
"create_scheduler_middleware",
|
"create_scheduler_middleware",
|
||||||
"create_tool_selector_middleware",
|
"create_tool_selector_middleware",
|
||||||
"default_memory_scheduler",
|
"default_memory_scheduler",
|
||||||
"disable_thinking",
|
"disable_thinking",
|
||||||
"load_fallback_chain",
|
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -29,13 +29,46 @@ from langchain.agents.middleware.types import (
|
|||||||
)
|
)
|
||||||
from langchain.tools import InjectedToolCallId
|
from langchain.tools import InjectedToolCallId
|
||||||
from langchain_core.messages import AIMessage, SystemMessage, ToolMessage
|
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 langgraph.types import Command, interrupt
|
||||||
from pydantic import BeforeValidator, Field
|
from pydantic import BeforeValidator, Field
|
||||||
from typing_extensions import TypedDict
|
from typing_extensions import TypedDict
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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
|
# Data types
|
||||||
@@ -202,6 +235,16 @@ or available tools.
|
|||||||
- Never ask more than once per decision point — respect the user's time
|
- Never ask more than once per decision point — respect the user's time
|
||||||
- After receiving answers, summarize what you understood before proceeding"""
|
- 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
|
# 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]):
|
class AskUserMiddleware(AgentMiddleware[Any, ContextT, ResponseT]):
|
||||||
"""Middleware that provides an ``ask_user`` tool for interactive questioning.
|
"""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],
|
tool_call_id: Annotated[str, InjectedToolCallId],
|
||||||
) -> Command[Any]:
|
) -> Command[Any]:
|
||||||
"""Ask the user one or more questions."""
|
"""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)
|
_validate_questions(questions)
|
||||||
ask_request = AskUserRequest(
|
ask_request = AskUserRequest(
|
||||||
type="ask_user",
|
type="ask_user",
|
||||||
@@ -386,18 +451,22 @@ class AskUserMiddleware(AgentMiddleware[Any, ContextT, ResponseT]):
|
|||||||
request: ModelRequest[ContextT],
|
request: ModelRequest[ContextT],
|
||||||
handler: Callable[[ModelRequest[ContextT]], ModelResponse[ResponseT]],
|
handler: Callable[[ModelRequest[ContextT]], ModelResponse[ResponseT]],
|
||||||
) -> ModelResponse[ResponseT] | AIMessage:
|
) -> ModelResponse[ResponseT] | AIMessage:
|
||||||
"""Inject the ask_user system prompt."""
|
"""Apply the interactive or unattended prompt and tool policy."""
|
||||||
if request.system_message is not None:
|
if _review_mode() == "full":
|
||||||
new_system_content = [
|
tools = [tool for tool in request.tools if _tool_name(tool) != "ask_user"]
|
||||||
*request.system_message.content_blocks,
|
return handler(
|
||||||
{"type": "text", "text": f"\n\n{self.system_prompt}"},
|
request.override(
|
||||||
]
|
system_message=_with_system_prompt(
|
||||||
else:
|
request, FULL_APPROVE_SYSTEM_PROMPT
|
||||||
new_system_content = [{"type": "text", "text": self.system_prompt}]
|
),
|
||||||
new_system_message = SystemMessage(
|
tools=tools,
|
||||||
content=cast("list[str | dict[str, str]]", new_system_content)
|
)
|
||||||
|
)
|
||||||
|
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(
|
async def awrap_model_call(
|
||||||
self,
|
self,
|
||||||
@@ -406,15 +475,19 @@ class AskUserMiddleware(AgentMiddleware[Any, ContextT, ResponseT]):
|
|||||||
[ModelRequest[ContextT]], Awaitable[ModelResponse[ResponseT]]
|
[ModelRequest[ContextT]], Awaitable[ModelResponse[ResponseT]]
|
||||||
],
|
],
|
||||||
) -> ModelResponse[ResponseT] | AIMessage:
|
) -> ModelResponse[ResponseT] | AIMessage:
|
||||||
"""Inject the ask_user system prompt (async)."""
|
"""Apply the interactive or unattended prompt and tool policy (async)."""
|
||||||
if request.system_message is not None:
|
if _review_mode() == "full":
|
||||||
new_system_content = [
|
tools = [tool for tool in request.tools if _tool_name(tool) != "ask_user"]
|
||||||
*request.system_message.content_blocks,
|
return await handler(
|
||||||
{"type": "text", "text": f"\n\n{self.system_prompt}"},
|
request.override(
|
||||||
]
|
system_message=_with_system_prompt(
|
||||||
else:
|
request, FULL_APPROVE_SYSTEM_PROMPT
|
||||||
new_system_content = [{"type": "text", "text": self.system_prompt}]
|
),
|
||||||
new_system_message = SystemMessage(
|
tools=tools,
|
||||||
content=cast("list[str | dict[str, str]]", new_system_content)
|
)
|
||||||
|
)
|
||||||
|
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
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
from langchain.agents.middleware import AgentMiddleware
|
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) —
|
# 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
|
# cheaper than reloading the full config from disk on every launch, and uses
|
||||||
# the same truthy parsing as every other bool env flag.
|
# the same truthy parsing as every other bool env flag.
|
||||||
from ..llm.models import _env_flag_enabled
|
dangerous = os.getenv("EVOSCIENTIST_DANGEROUS_MODE", "").strip().lower() in {
|
||||||
|
"1",
|
||||||
dangerous = _env_flag_enabled("EVOSCIENTIST_DANGEROUS_MODE")
|
"true",
|
||||||
|
"yes",
|
||||||
|
"on",
|
||||||
|
}
|
||||||
# Same path-rewriting + validation as execute (shared helper) so virtual paths
|
# Same path-rewriting + validation as execute (shared helper) so virtual paths
|
||||||
# resolve to the workspace and the command can't bypass the sandbox checks.
|
# resolve to the workspace and the command can't bypass the sandbox checks.
|
||||||
command, error = prepare_sandbox_command(
|
command, error = prepare_sandbox_command(
|
||||||
|
|||||||
@@ -45,21 +45,7 @@ _MEMORY_FIRST_INTERPRETER_PROMPT = (
|
|||||||
|
|
||||||
|
|
||||||
class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware):
|
class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware):
|
||||||
"""Code interpreter middleware with EvoScientist's memory preflight hint.
|
"""Code interpreter middleware with EvoScientist's memory preflight hint."""
|
||||||
|
|
||||||
``after_agent`` / ``aafter_agent`` are intentionally NOT overridden. An
|
|
||||||
earlier "conditional snapshot" gate that skipped ``after_agent`` on turns
|
|
||||||
where ``code_interpreter`` wasn't called saved ~50 ms/turn of
|
|
||||||
``create_snapshot()`` work, but also skipped the slot eviction upstream
|
|
||||||
performs in the same hook (``finally: self._registry.evict(thread_id)``
|
|
||||||
in ``langchain_quickjs.middleware.CodeInterpreterMiddleware.after_agent``).
|
|
||||||
``before_agent`` restores the REPL on every turn that follows a touched
|
|
||||||
one via ``self._registry.get(thread_id)`` (get-or-create), so skipping
|
|
||||||
eviction leaked one ``ThreadWorker`` + QuickJS Runtime per persistent
|
|
||||||
``thread_id`` that ever went touched → quiet. The regression test
|
|
||||||
``test_after_agent_evicts_slot_on_untouched_turn`` guards against
|
|
||||||
reintroducing the gate.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def _prepare_for_call(self, request: ModelRequest) -> str:
|
def _prepare_for_call(self, request: ModelRequest) -> str:
|
||||||
return super()._prepare_for_call(request) + _MEMORY_FIRST_INTERPRETER_PROMPT
|
return super()._prepare_for_call(request) + _MEMORY_FIRST_INTERPRETER_PROMPT
|
||||||
|
|||||||
@@ -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
|
Design doc 8.3: the middleware no longer reads ``config.yaml``, global
|
||||||
and have their model frozen into the compiled graph at subprocess boot time
|
aliases, or ``model``/``model_provider`` overrides. The only model input a
|
||||||
(see ``EvoScientist/subagents/_factory.py``). When the user runs ``/model``
|
run may carry is ``configurable["runtime_snapshot_id"]``; the middleware
|
||||||
in the CLI, only the CLI process's model state changes — the subprocess
|
loads the frozen snapshot (deployment/thread binding verified), resolves
|
||||||
graph still uses the boot-time model.
|
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``
|
``configurable`` carrying ``model``, ``model_provider``, or any other
|
||||||
from ``RunnableConfig.configurable`` on every model call. The CLI's patched
|
out-of-snapshot model parameter is rejected with
|
||||||
``start_async_task`` / ``update_async_task`` (see ``llm/patches.py``) injects
|
``MODEL_CONFIG_OUTSIDE_SNAPSHOT`` (422 semantics, section 8.2) — run
|
||||||
those fields into ``client.runs.create(config=...)``; the deployed graph
|
creation must never mix snapshot and non-snapshot model configuration.
|
||||||
hits this middleware and re-resolves the chat model fresh.
|
|
||||||
|
|
||||||
When ``configurable.model`` is absent, the middleware is a pass-through —
|
Local entry points (CLI/channels/scheduler/sub-agents, section 8.1) create
|
||||||
safe to install on the CLI's in-process agent too.
|
snapshots through ``SnapshotRuntime.create_local_snapshot`` and put the
|
||||||
|
``runtime_snapshot_id`` into ``configurable`` before the run starts. Runs
|
||||||
The middleware mirrors the pattern used by ``ModelFallbackMiddleware``:
|
whose entry point could not inject a snapshot up front (langgraph-dev cron
|
||||||
``request.override(model=new_model)`` does not break tool binding, because
|
fires, deployed async sub-agent graphs) are healed lazily: the middleware
|
||||||
the downstream model-invocation node re-binds tools per request.
|
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
|
**Reading the config**: ``Runtime`` (per its own docstring) does NOT include
|
||||||
``config``. The official path to reach ``RunnableConfig`` from inside any
|
``config``. The official path to reach ``RunnableConfig`` from inside any
|
||||||
@@ -32,8 +36,8 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable, Mapping
|
||||||
from typing import Any
|
from typing import Any, get_args
|
||||||
|
|
||||||
from langchain.agents.middleware.types import (
|
from langchain.agents.middleware.types import (
|
||||||
AgentMiddleware,
|
AgentMiddleware,
|
||||||
@@ -41,108 +45,202 @@ from langchain.agents.middleware.types import (
|
|||||||
ModelResponse,
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_MODEL_ROLES = get_args(ModelRole)
|
||||||
|
|
||||||
def _read_model_override() -> tuple[str | None, str | None]:
|
# The only model-related keys a run's ``configurable`` may never carry: the
|
||||||
"""Pull ``(model, model_provider)`` from the active ``RunnableConfig``.
|
# 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
|
def _current_configurable() -> Mapping[str, Any]:
|
||||||
context — middleware, node, tool). Returns ``(None, None)`` when the
|
"""Return the active run's ``configurable`` mapping (empty outside runs)."""
|
||||||
config has no ``configurable.model`` override or when called outside a
|
|
||||||
runnable context.
|
|
||||||
"""
|
|
||||||
try:
|
try:
|
||||||
from langgraph.config import get_config
|
from langgraph.config import get_config
|
||||||
|
|
||||||
cfg = get_config()
|
config = get_config()
|
||||||
except Exception:
|
except Exception:
|
||||||
# Outside a runnable context (most common in tests) or
|
# Outside a runnable context (most common in tests) or langgraph not
|
||||||
# langgraph not importable — nothing to override.
|
# importable — treat as "no per-run configuration".
|
||||||
return None, None
|
return {}
|
||||||
if not isinstance(cfg, dict):
|
if not isinstance(config, Mapping):
|
||||||
return None, None
|
return {}
|
||||||
configurable = cfg.get("configurable") or {}
|
configurable = config.get("configurable")
|
||||||
if not isinstance(configurable, dict):
|
return configurable if isinstance(configurable, Mapping) else {}
|
||||||
return None, None
|
|
||||||
model = configurable.get("model")
|
|
||||||
provider = configurable.get("model_provider")
|
def check_no_outside_snapshot_model_config(configurable: Mapping[str, Any]) -> None:
|
||||||
return (
|
"""Reject any model configuration carried outside the run snapshot."""
|
||||||
model if isinstance(model, str) and model else None,
|
for key in _OUTSIDE_SNAPSHOT_MODEL_KEYS:
|
||||||
provider if isinstance(provider, str) and provider else None,
|
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):
|
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``
|
``role`` selects which frozen configuration of the snapshot feeds this
|
||||||
via ``langgraph.config.get_config()`` — the documented entry point for
|
agent's model calls; every role currently maps to the snapshot's frozen
|
||||||
accessing per-run config from any runnable context (middleware, node, tool).
|
primary (section 6.1).
|
||||||
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.
|
|
||||||
|
|
||||||
Note: ``Runtime`` (per its own docstring) does NOT include ``config`` as a
|
A per-instance cache keyed by snapshot ID avoids rebuilding the model on
|
||||||
field — an earlier version of this middleware tried to read
|
every call within a run; snapshots are immutable once created, so the
|
||||||
``request.runtime.config`` and silently no-op'd because that attribute does
|
cached instance stays valid for the run's lifetime. The cache is guarded
|
||||||
not exist. Stick with ``get_config()``.
|
by a ``threading.Lock`` because middleware instances are shared across
|
||||||
|
|
||||||
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
|
|
||||||
concurrent requests in long-lived deployments (e.g. ``langgraph dev``
|
concurrent requests in long-lived deployments (e.g. ``langgraph dev``
|
||||||
workers).
|
workers).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
name = "configurable_model"
|
name = "configurable_model"
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(
|
||||||
|
self, role: ModelRole = "primary", runtime: SnapshotRuntime | None = None
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
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()
|
self._lock = threading.Lock()
|
||||||
# Track the last (model, provider) pair we INFO-logged so we only
|
# Track the last snapshot we INFO-logged so we only surface a banner
|
||||||
# surface a banner on transition. Without this, every LLM call in a
|
# on transition. Without this, every LLM call in a long run would
|
||||||
# long async run would emit an identical INFO line.
|
# emit an identical INFO line.
|
||||||
self._last_logged_key: tuple[str, str | None] | None = None
|
self._last_logged_snapshot_id: str | None = None
|
||||||
|
|
||||||
def _log_override(self, model_name: str, provider: str | None) -> None:
|
def _snapshot_runtime(self) -> SnapshotRuntime:
|
||||||
"""INFO on transition; DEBUG on subsequent calls with same key."""
|
return self._runtime if self._runtime is not None else get_snapshot_runtime()
|
||||||
key = (model_name, provider)
|
|
||||||
|
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:
|
with self._lock:
|
||||||
transitioned = key != self._last_logged_key
|
transitioned = snapshot.snapshot_id != self._last_logged_snapshot_id
|
||||||
if transitioned:
|
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:
|
if transitioned:
|
||||||
logger.info(
|
logger.info(
|
||||||
"ConfigurableModelMiddleware: overriding model to %s (%s)",
|
"ConfigurableModelMiddleware: role %s bound to %s/%s (snapshot %s)",
|
||||||
model_name,
|
*message_args,
|
||||||
provider,
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"ConfigurableModelMiddleware: reusing override model=%s provider=%s",
|
"ConfigurableModelMiddleware: role %s reusing %s/%s (snapshot %s)",
|
||||||
model_name,
|
*message_args,
|
||||||
provider,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _resolve(self, model: str, provider: str | None) -> Any:
|
def _load_snapshot(self) -> RuntimeSnapshot:
|
||||||
"""Return a cached or freshly-built chat model for ``(model, provider)``."""
|
"""Load the run's snapshot, verifying its deployment/thread binding.
|
||||||
key = (model, provider)
|
|
||||||
|
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:
|
with self._lock:
|
||||||
cached = self._cache.get(key)
|
cached = self._cache.get(snapshot.snapshot_id)
|
||||||
if cached is not None:
|
if cached is not None:
|
||||||
return cached
|
return cached
|
||||||
# Build outside the lock (network/SDK init can be slow); two
|
# Build outside the lock (SDK init can be slow); two concurrent
|
||||||
# concurrent first-time misses for the same key may build twice but
|
# first-time misses for the same snapshot may build twice but the
|
||||||
# the second result simply overwrites the first — both are equivalent.
|
# second result simply overwrites the first — both are equivalent.
|
||||||
from ..llm import get_chat_model
|
new_model = self._snapshot_runtime().build_role_model(snapshot, self._role)
|
||||||
|
|
||||||
new_model = get_chat_model(model=model, provider=provider)
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._cache[key] = new_model
|
self._cache[snapshot.snapshot_id] = new_model
|
||||||
|
self._log_override(snapshot)
|
||||||
return new_model
|
return new_model
|
||||||
|
|
||||||
def wrap_model_call(
|
def wrap_model_call(
|
||||||
@@ -150,21 +248,7 @@ class ConfigurableModelMiddleware(AgentMiddleware):
|
|||||||
request: ModelRequest,
|
request: ModelRequest,
|
||||||
handler: Callable[[ModelRequest], ModelResponse],
|
handler: Callable[[ModelRequest], ModelResponse],
|
||||||
) -> ModelResponse:
|
) -> ModelResponse:
|
||||||
model_name, provider = _read_model_override()
|
new_model = self._resolve()
|
||||||
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)
|
|
||||||
return handler(request.override(model=new_model))
|
return handler(request.override(model=new_model))
|
||||||
|
|
||||||
async def awrap_model_call(
|
async def awrap_model_call(
|
||||||
@@ -172,25 +256,44 @@ class ConfigurableModelMiddleware(AgentMiddleware):
|
|||||||
request: ModelRequest,
|
request: ModelRequest,
|
||||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||||
) -> ModelResponse:
|
) -> ModelResponse:
|
||||||
model_name, provider = _read_model_override()
|
# Offload first-call SDK init off the event loop: ``_resolve`` reads
|
||||||
if model_name is None:
|
# SQLite and can spend hundreds of ms building HTTP clients on a
|
||||||
return await handler(request)
|
# cache miss, which would block every other coroutine on the same
|
||||||
try:
|
# langgraph dev event loop. Cache hits are still fast; the
|
||||||
# Offload first-call SDK init off the event loop. ``_resolve`` calls
|
# thread-pool overhead is irrelevant once warm.
|
||||||
# ``get_chat_model`` on a cache miss, which can spend hundreds of ms
|
new_model = await asyncio.to_thread(self._resolve)
|
||||||
# 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)
|
|
||||||
return await handler(request.override(model=new_model))
|
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
|
||||||
|
|||||||
@@ -1,240 +0,0 @@
|
|||||||
"""ErrorNormalizationMiddleware — catch provider-SDK exceptions at the
|
|
||||||
model boundary and re-raise as a normalized non-dataclass wrapper.
|
|
||||||
|
|
||||||
Some provider SDKs (openrouter.errors.* today) decorate their exception
|
|
||||||
classes with ``@dataclass``. When langgraph_api emits an SSE error
|
|
||||||
frame via ``json_dumpb`` → ``orjson.dumps(obj, default=default,
|
|
||||||
option=OPT_SERIALIZE_DATACLASS)``, orjson's dataclass fast-path
|
|
||||||
enumerates the fields directly and skips the ``default=`` hook that
|
|
||||||
builds our envelope. The wire payload comes out as
|
|
||||||
``{"message": …, "status_code": …, "body": …, "headers": null,
|
|
||||||
"raw_response": null, "data": {…}}`` with no ``error`` / ``class`` /
|
|
||||||
``provider`` envelope and no way for the WebUI to distinguish quota /
|
|
||||||
auth / rate-limit / model-not-found.
|
|
||||||
|
|
||||||
This middleware sits at the model-call boundary. It catches
|
|
||||||
``BaseException`` from ``handler()``, and if ``request.model`` is a
|
|
||||||
recognized provider SDK client, wraps the exception in a
|
|
||||||
:class:`~EvoScientist.llm.errors.ProviderStreamError` (a plain
|
|
||||||
``Exception`` subclass, not a dataclass). The wrapper carries the SSE
|
|
||||||
envelope pre-baked on its instance attributes.
|
|
||||||
|
|
||||||
Contract: the wrap decision is based on the **model**, not the
|
|
||||||
exception, after platform and graph control signals have been excluded.
|
|
||||||
Provider SDK exceptions, httpx errors, langchain-wrapper failures, and
|
|
||||||
even builtins like ``RuntimeError`` get wrapped for a recognized model.
|
|
||||||
At the middleware boundary we can tell which provider was in use, but
|
|
||||||
not the exception's precise origin; a uniform envelope is more useful
|
|
||||||
to the WebUI than gambling on the exception class. If the model isn't
|
|
||||||
from a recognized provider, or the request carries no ``.model``, the
|
|
||||||
exception re-raises unchanged and upstream's whitelist / catch-all
|
|
||||||
behavior takes over.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from collections.abc import Awaitable, Callable
|
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
|
|
||||||
from langchain.agents.middleware.types import (
|
|
||||||
AgentMiddleware,
|
|
||||||
ModelRequest,
|
|
||||||
ModelResponse,
|
|
||||||
)
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from ..llm.errors import ProviderStreamError
|
|
||||||
|
|
||||||
|
|
||||||
def _should_pass_through(exc: BaseException) -> bool:
|
|
||||||
"""True if *exc* is a LangGraph-level signal that must propagate
|
|
||||||
untouched — either a control-flow signal or a structural error
|
|
||||||
that isn't a provider failure.
|
|
||||||
|
|
||||||
Covers everything in ``langgraph.errors.*``:
|
|
||||||
|
|
||||||
- **Control flow** (breaking these would corrupt the interrupt /
|
|
||||||
resume protocol): ``GraphBubbleUp`` and its subclasses
|
|
||||||
``GraphInterrupt``, ``NodeInterrupt``, ``ParentCommand``,
|
|
||||||
``GraphDrained``.
|
|
||||||
- **Structural** (wrapping would mis-attribute a graph-level
|
|
||||||
issue as a provider failure): ``InvalidUpdateError``,
|
|
||||||
``EmptyInputError``, ``EmptyChannelError``, ``TaskNotFound``,
|
|
||||||
``GraphRecursionError``, ``NodeCancelledError``,
|
|
||||||
``NodeTimeoutError``.
|
|
||||||
|
|
||||||
Symmetric with upstream ``langgraph_api.serde.default``'s
|
|
||||||
whitelist, which also exposes these classes' ``str(exc)`` untouched
|
|
||||||
rather than swallowing them behind a provider envelope.
|
|
||||||
|
|
||||||
``KeyboardInterrupt``, ``SystemExit``, and ``asyncio.CancelledError``
|
|
||||||
are handled implicitly by catching ``Exception`` — they inherit
|
|
||||||
from ``BaseException``.
|
|
||||||
"""
|
|
||||||
return (type(exc).__module__ or "").startswith("langgraph.errors")
|
|
||||||
|
|
||||||
|
|
||||||
# Module prefixes for provider SDK exceptions. Consumed by
|
|
||||||
# ``_is_provider_error`` to decide whether an exception raised inside
|
|
||||||
# a model call should surface as a provider incident or gracefully
|
|
||||||
# degrade (used by ``_ConditionalToolSelectorMiddleware``).
|
|
||||||
#
|
|
||||||
# Related sibling: ``_HOST_TO_PROVIDER`` in ``llm/errors.py`` — the
|
|
||||||
# host-side allow-list. Adding a whole new provider SDK means updating
|
|
||||||
# both; adding a new routed provider (new base_url through an existing
|
|
||||||
# SDK) only touches ``_HOST_TO_PROVIDER``.
|
|
||||||
_PROVIDER_EXC_MODULE_PREFIXES: tuple[str, ...] = (
|
|
||||||
"openai",
|
|
||||||
"anthropic",
|
|
||||||
"google.genai",
|
|
||||||
"google.api_core",
|
|
||||||
"openrouter",
|
|
||||||
"langchain_openai",
|
|
||||||
"langchain_anthropic",
|
|
||||||
"langchain_google_genai",
|
|
||||||
"langchain_openrouter",
|
|
||||||
"httpx",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _is_provider_error(exc: BaseException) -> bool:
|
|
||||||
"""True if *exc* looks like it originated inside a provider SDK
|
|
||||||
(openai, anthropic, google.genai, openrouter, httpx, or their
|
|
||||||
langchain wrappers), as opposed to a shape / config error (structured
|
|
||||||
output not supported, malformed schema, missing tool, …).
|
|
||||||
|
|
||||||
Used by callers that need to decide whether an exception from the
|
|
||||||
model call is worth surfacing to the user (provider errors) or
|
|
||||||
can be silently degraded around (shape errors). Cheap alternative
|
|
||||||
to inspecting ``status_code`` / ``request`` because some provider
|
|
||||||
errors — connection errors, timeouts — don't carry those attributes.
|
|
||||||
"""
|
|
||||||
module = type(exc).__module__ or ""
|
|
||||||
return any(module.startswith(p) for p in _PROVIDER_EXC_MODULE_PREFIXES)
|
|
||||||
|
|
||||||
|
|
||||||
def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError | None:
|
|
||||||
"""Return a :class:`ProviderStreamError` wrapping *exc* if the model
|
|
||||||
on *request* comes from a recognized provider SDK, or ``None`` if
|
|
||||||
the caller should re-raise *exc* unchanged.
|
|
||||||
|
|
||||||
Provider is read from ``request.model`` — the definitive config
|
|
||||||
the exception was raised under, not inferred from the exception
|
|
||||||
class / URL. Status / code / redaction still come from the raised
|
|
||||||
exception because those fields are populated by the SDK at raise
|
|
||||||
time.
|
|
||||||
|
|
||||||
Returns ``None`` (caller re-raises unchanged) for:
|
|
||||||
|
|
||||||
- Already-normalized wrappers (would double-attribute).
|
|
||||||
- LangGraph control-flow / structural errors — see
|
|
||||||
``_should_pass_through``. This gate lives here so every caller
|
|
||||||
of ``_normalize`` (not just the wrap sites of this middleware)
|
|
||||||
gets the protection automatically. Notably
|
|
||||||
``ModelFallbackMiddleware`` also calls ``_normalize`` at the
|
|
||||||
raise point of its fallback chain.
|
|
||||||
- ``ContextOverflowError`` — a cross-layer control signal that
|
|
||||||
deepagents' ``SummarizationMiddleware`` catches by type from
|
|
||||||
**outside** the user middleware stack to compress history and
|
|
||||||
retry. Wrapping it here would change the type and break that
|
|
||||||
self-healing fallback.
|
|
||||||
- ``AgentControlError`` — a platform-owned typed decision. Gateway route
|
|
||||||
fallback and canonical error mapping depend on its concrete type and
|
|
||||||
structured fields, so it must never become a provider incident.
|
|
||||||
- Models we don't recognize as a provider SDK.
|
|
||||||
"""
|
|
||||||
from langchain_core.exceptions import ContextOverflowError
|
|
||||||
|
|
||||||
from ..llm.errors import (
|
|
||||||
AgentControlError,
|
|
||||||
ProviderStreamError,
|
|
||||||
_extract_error_type,
|
|
||||||
_extract_provider_code,
|
|
||||||
_extract_status_code,
|
|
||||||
_provider_from_model,
|
|
||||||
_redact_api_keys,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Already normalized (e.g. by ModelFallbackMiddleware wrapping against
|
|
||||||
# the actual failing model rather than the original request's model).
|
|
||||||
# Pass through — re-wrapping would double-attribute.
|
|
||||||
if isinstance(exc, ProviderStreamError):
|
|
||||||
return None
|
|
||||||
|
|
||||||
# Platform control errors are raised by inner middleware after the provider
|
|
||||||
# response has already been interpreted. Wrapping them would erase routing,
|
|
||||||
# retry and recovery semantics such as ModelToolProtocolError.fallbackable.
|
|
||||||
if isinstance(exc, AgentControlError):
|
|
||||||
return None
|
|
||||||
|
|
||||||
# LangGraph control-flow / structural signals must propagate
|
|
||||||
# untouched, regardless of which caller invoked us.
|
|
||||||
if _should_pass_through(exc):
|
|
||||||
return None
|
|
||||||
|
|
||||||
# SummarizationMiddleware sits outside our stack and catches this
|
|
||||||
# by exact type to trigger reactive history compression + retry.
|
|
||||||
if isinstance(exc, ContextOverflowError):
|
|
||||||
return None
|
|
||||||
|
|
||||||
provider = _provider_from_model(getattr(request, "model", None))
|
|
||||||
if provider is None:
|
|
||||||
return None
|
|
||||||
cls = type(exc)
|
|
||||||
mod = cls.__module__ or ""
|
|
||||||
class_qualname = f"{mod}.{cls.__qualname__}" if mod else cls.__qualname__
|
|
||||||
|
|
||||||
request_id_attr = getattr(exc, "request_id", None)
|
|
||||||
request_id = (
|
|
||||||
request_id_attr
|
|
||||||
if isinstance(request_id_attr, str) and request_id_attr
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
|
|
||||||
return ProviderStreamError(
|
|
||||||
provider=provider,
|
|
||||||
class_qualname=class_qualname,
|
|
||||||
message=_redact_api_keys(str(exc)),
|
|
||||||
status_code=_extract_status_code(exc),
|
|
||||||
code=_extract_provider_code(exc),
|
|
||||||
err_type=_extract_error_type(exc),
|
|
||||||
request_id=request_id,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class ErrorNormalizationMiddleware(AgentMiddleware):
|
|
||||||
"""Wrap the model call in try/except and normalize provider SDK
|
|
||||||
exceptions into a non-dataclass envelope wrapper.
|
|
||||||
|
|
||||||
Place this middleware **outermost** in the chain (first in the
|
|
||||||
middleware list) so it catches exceptions raised by inner
|
|
||||||
middlewares as well as the model handler itself.
|
|
||||||
"""
|
|
||||||
|
|
||||||
name = "error_normalization"
|
|
||||||
|
|
||||||
def wrap_model_call(
|
|
||||||
self,
|
|
||||||
request: ModelRequest,
|
|
||||||
handler: Callable[[ModelRequest], ModelResponse],
|
|
||||||
) -> ModelResponse:
|
|
||||||
try:
|
|
||||||
return handler(request)
|
|
||||||
except Exception as exc:
|
|
||||||
normalized = _normalize(request, exc)
|
|
||||||
if normalized is None:
|
|
||||||
raise
|
|
||||||
raise normalized from exc
|
|
||||||
|
|
||||||
async def awrap_model_call(
|
|
||||||
self,
|
|
||||||
request: ModelRequest,
|
|
||||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
|
||||||
) -> ModelResponse:
|
|
||||||
try:
|
|
||||||
return await handler(request)
|
|
||||||
except Exception as exc:
|
|
||||||
normalized = _normalize(request, exc)
|
|
||||||
if normalized is None:
|
|
||||||
raise
|
|
||||||
raise normalized from exc
|
|
||||||
@@ -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,409 +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",
|
|
||||||
"repetitive tool calls",
|
|
||||||
"identical name and arguments",
|
|
||||||
]
|
|
||||||
"""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 getattr(exc, "non_fallbackable", False):
|
|
||||||
return f"platform control error: {getattr(exc, 'code', type(exc).__name__)}"
|
|
||||||
|
|
||||||
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
|
|
||||||
)
|
|
||||||
|
|
||||||
# Track the request whose model actually raised ``last_exc`` so we
|
|
||||||
# can attribute the exception to the failing model, not the
|
|
||||||
# original ``request.model``. Without this, a fallback chain
|
|
||||||
# ``deepseek → moonshot`` where moonshot exhausts its quota would
|
|
||||||
# surface as ``provider: deepseek`` — the model the user never
|
|
||||||
# actually saw fail.
|
|
||||||
last_exc = primary_exc
|
|
||||||
last_failing_request = request
|
|
||||||
|
|
||||||
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_normalized(fb_request, fb_exc)
|
|
||||||
last_exc = fb_exc
|
|
||||||
last_failing_request = fb_request
|
|
||||||
_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_normalized(last_failing_request, last_exc)
|
|
||||||
|
|
||||||
|
|
||||||
def _raise_normalized(request: ModelRequest, exc: Exception) -> None:
|
|
||||||
"""Wrap *exc* in a ``ProviderStreamError`` attributed to
|
|
||||||
``request.model`` and raise, so the outer chain sees the failure
|
|
||||||
tagged with the model that actually raised.
|
|
||||||
|
|
||||||
Falls back to a plain ``raise`` when the model isn't from a
|
|
||||||
recognized provider (``_normalize`` returns None) — nothing useful
|
|
||||||
to add.
|
|
||||||
"""
|
|
||||||
from .error_normalization import _normalize
|
|
||||||
|
|
||||||
normalized = _normalize(request, exc)
|
|
||||||
if normalized is not None:
|
|
||||||
raise normalized from exc
|
|
||||||
raise 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_normalized(request, 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
|
||||||
@@ -1,350 +0,0 @@
|
|||||||
"""Detect deterministic tool loops and compact only provider-facing history."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
|
||||||
import logging
|
|
||||||
import re
|
|
||||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from langchain.agents.middleware.types import (
|
|
||||||
AgentMiddleware,
|
|
||||||
ModelRequest,
|
|
||||||
ModelResponse,
|
|
||||||
)
|
|
||||||
|
|
||||||
from ..llm.errors import AgentControlError
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD = 2
|
|
||||||
DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS = 3
|
|
||||||
|
|
||||||
_TRANSIENT_PATTERNS = (
|
|
||||||
"timeout",
|
|
||||||
"timed out",
|
|
||||||
"cancelled",
|
|
||||||
"canceled",
|
|
||||||
"connection",
|
|
||||||
"rate limit",
|
|
||||||
"too many requests",
|
|
||||||
"temporarily unavailable",
|
|
||||||
"service unavailable",
|
|
||||||
"overloaded",
|
|
||||||
"bad gateway",
|
|
||||||
"gateway timeout",
|
|
||||||
"http 500",
|
|
||||||
"http 502",
|
|
||||||
"http 503",
|
|
||||||
"http 504",
|
|
||||||
)
|
|
||||||
_DETERMINISTIC_PATTERNS: tuple[tuple[str, tuple[str, ...]], ...] = (
|
|
||||||
(
|
|
||||||
"INVALID_ARGUMENTS",
|
|
||||||
("invalid argument", "validation error", "schema", "bad input"),
|
|
||||||
),
|
|
||||||
("UNKNOWN_TOOL", ("not a valid tool", "unknown tool", "tool not found")),
|
|
||||||
("UNSUPPORTED", ("not supported", "unsupported", "not implemented")),
|
|
||||||
(
|
|
||||||
"POLICY_DENIED",
|
|
||||||
("permission denied", "forbidden", "policy denied", "not allowed"),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
_SAFE_CODE_RE = re.compile(r"^[A-Za-z][A-Za-z0-9_.:-]{0,95}$")
|
|
||||||
_DETERMINISTIC_CODE_MARKERS = (
|
|
||||||
"INVALID",
|
|
||||||
"VALIDATION",
|
|
||||||
"SCHEMA",
|
|
||||||
"UNKNOWN_TOOL",
|
|
||||||
"NOT_FOUND",
|
|
||||||
"UNSUPPORTED",
|
|
||||||
"NOT_IMPLEMENTED",
|
|
||||||
"POLICY",
|
|
||||||
"PERMISSION",
|
|
||||||
"FORBIDDEN",
|
|
||||||
"DENIED",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
|
||||||
class RepetitiveToolHistoryRepair:
|
|
||||||
messages: list[Any]
|
|
||||||
blocked_tool_names: frozenset[str]
|
|
||||||
removed_rounds: int
|
|
||||||
tail_repetitions: int = 0
|
|
||||||
tail_consecutive_errors: int = 0
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
|
||||||
class _ToolRound:
|
|
||||||
messages: tuple[Any, ...]
|
|
||||||
signature: tuple[tuple[str, str, str], ...]
|
|
||||||
tool_names: frozenset[str]
|
|
||||||
deterministic_error: bool
|
|
||||||
|
|
||||||
|
|
||||||
def _canonical_tool_args(value: Any) -> str:
|
|
||||||
if isinstance(value, str):
|
|
||||||
try:
|
|
||||||
value = json.loads(value)
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
return value.strip()
|
|
||||||
try:
|
|
||||||
return json.dumps(
|
|
||||||
value,
|
|
||||||
ensure_ascii=False,
|
|
||||||
sort_keys=True,
|
|
||||||
separators=(",", ":"),
|
|
||||||
default=str,
|
|
||||||
)
|
|
||||||
except (TypeError, ValueError):
|
|
||||||
return repr(value)
|
|
||||||
|
|
||||||
|
|
||||||
def _deterministic_result_code(message: Any) -> str | None:
|
|
||||||
additional = getattr(message, "additional_kwargs", None)
|
|
||||||
additional = additional if isinstance(additional, Mapping) else {}
|
|
||||||
raw_code = additional.get("error_code") or additional.get("code")
|
|
||||||
status = str(getattr(message, "status", "") or "").lower()
|
|
||||||
content = str(getattr(message, "content", "") or "")
|
|
||||||
lowered = content.lower()
|
|
||||||
|
|
||||||
if any(pattern in lowered for pattern in _TRANSIENT_PATTERNS):
|
|
||||||
return None
|
|
||||||
if isinstance(raw_code, str) and _SAFE_CODE_RE.fullmatch(raw_code.strip()):
|
|
||||||
normalized = raw_code.strip().upper()
|
|
||||||
if any(
|
|
||||||
pattern.replace(" ", "_") in normalized for pattern in _TRANSIENT_PATTERNS
|
|
||||||
):
|
|
||||||
return None
|
|
||||||
if any(marker in normalized for marker in _DETERMINISTIC_CODE_MARKERS):
|
|
||||||
return normalized
|
|
||||||
return None
|
|
||||||
is_error = status == "error" or lowered.startswith("error:")
|
|
||||||
if not is_error:
|
|
||||||
return None
|
|
||||||
for code, patterns in _DETERMINISTIC_PATTERNS:
|
|
||||||
if any(pattern in lowered for pattern in patterns):
|
|
||||||
return code
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _parse_tool_round(
|
|
||||||
messages: Sequence[Any], start: int
|
|
||||||
) -> tuple[_ToolRound, int] | None:
|
|
||||||
assistant = messages[start]
|
|
||||||
if getattr(assistant, "type", None) != "ai":
|
|
||||||
return None
|
|
||||||
raw_calls = list(getattr(assistant, "tool_calls", None) or [])
|
|
||||||
calls = [call for call in raw_calls if isinstance(call, Mapping)]
|
|
||||||
if not calls or len(calls) != len(raw_calls):
|
|
||||||
return None
|
|
||||||
|
|
||||||
end = start + 1
|
|
||||||
results: list[Any] = []
|
|
||||||
while end < len(messages) and getattr(messages[end], "type", None) == "tool":
|
|
||||||
results.append(messages[end])
|
|
||||||
end += 1
|
|
||||||
if not results:
|
|
||||||
return None
|
|
||||||
results_by_id = {
|
|
||||||
str(getattr(result, "tool_call_id", "") or "").strip(): result
|
|
||||||
for result in results
|
|
||||||
if str(getattr(result, "tool_call_id", "") or "").strip()
|
|
||||||
}
|
|
||||||
|
|
||||||
signature: list[tuple[str, str, str]] = []
|
|
||||||
tool_names: set[str] = set()
|
|
||||||
for index, call in enumerate(calls):
|
|
||||||
name = str(call.get("name") or "").strip()
|
|
||||||
call_id = str(call.get("id") or "").strip()
|
|
||||||
if not name or not call_id:
|
|
||||||
return None
|
|
||||||
result = results_by_id.get(call_id)
|
|
||||||
if result is None and index < len(results):
|
|
||||||
candidate = results[index]
|
|
||||||
if not str(getattr(candidate, "tool_call_id", "") or "").strip():
|
|
||||||
result = candidate
|
|
||||||
if result is None:
|
|
||||||
return None
|
|
||||||
result_code = _deterministic_result_code(result)
|
|
||||||
if result_code is None:
|
|
||||||
return _ToolRound(
|
|
||||||
messages=(assistant, *results),
|
|
||||||
signature=(),
|
|
||||||
tool_names=frozenset(),
|
|
||||||
deterministic_error=False,
|
|
||||||
), end
|
|
||||||
signature.append((name, _canonical_tool_args(call.get("args")), result_code))
|
|
||||||
tool_names.add(name)
|
|
||||||
|
|
||||||
return (
|
|
||||||
_ToolRound(
|
|
||||||
messages=(assistant, *results),
|
|
||||||
signature=tuple(signature),
|
|
||||||
tool_names=frozenset(tool_names),
|
|
||||||
deterministic_error=True,
|
|
||||||
),
|
|
||||||
end,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def collapse_repetitive_tool_rounds(
|
|
||||||
messages: Sequence[Any],
|
|
||||||
*,
|
|
||||||
threshold: int = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
|
||||||
) -> RepetitiveToolHistoryRepair:
|
|
||||||
"""Build a provider-only projection while preserving audit history.
|
|
||||||
|
|
||||||
Only the middle rounds of three-or-more identical deterministic error
|
|
||||||
groups are omitted. The first and last observations remain, and callers
|
|
||||||
must never persist this projection back to a checkpoint.
|
|
||||||
"""
|
|
||||||
if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0:
|
|
||||||
raise ValueError("repetitive tool call threshold must be non-negative")
|
|
||||||
original = list(messages)
|
|
||||||
segments: list[Any | _ToolRound] = []
|
|
||||||
index = 0
|
|
||||||
while index < len(original):
|
|
||||||
parsed = _parse_tool_round(original, index)
|
|
||||||
if parsed is None:
|
|
||||||
segments.append(original[index])
|
|
||||||
index += 1
|
|
||||||
continue
|
|
||||||
tool_round, index = parsed
|
|
||||||
segments.append(tool_round)
|
|
||||||
|
|
||||||
tail_repetitions = 0
|
|
||||||
tail_consecutive_errors = 0
|
|
||||||
if segments and isinstance(segments[-1], _ToolRound):
|
|
||||||
tail = segments[-1]
|
|
||||||
if tail.deterministic_error:
|
|
||||||
cursor = len(segments) - 1
|
|
||||||
while cursor >= 0 and isinstance(segments[cursor], _ToolRound):
|
|
||||||
current = segments[cursor]
|
|
||||||
if not current.deterministic_error:
|
|
||||||
break
|
|
||||||
tail_consecutive_errors += 1
|
|
||||||
cursor -= 1
|
|
||||||
cursor = len(segments) - 1
|
|
||||||
while cursor >= 0 and isinstance(segments[cursor], _ToolRound):
|
|
||||||
current = segments[cursor]
|
|
||||||
if (
|
|
||||||
not current.deterministic_error
|
|
||||||
or current.signature != tail.signature
|
|
||||||
):
|
|
||||||
break
|
|
||||||
tail_repetitions += 1
|
|
||||||
cursor -= 1
|
|
||||||
|
|
||||||
projected: list[Any] = []
|
|
||||||
removed_rounds = 0
|
|
||||||
index = 0
|
|
||||||
while index < len(segments):
|
|
||||||
segment = segments[index]
|
|
||||||
if not isinstance(segment, _ToolRound) or not segment.deterministic_error:
|
|
||||||
if isinstance(segment, _ToolRound):
|
|
||||||
projected.extend(segment.messages)
|
|
||||||
else:
|
|
||||||
projected.append(segment)
|
|
||||||
index += 1
|
|
||||||
continue
|
|
||||||
end = index + 1
|
|
||||||
while (
|
|
||||||
end < len(segments)
|
|
||||||
and isinstance(segments[end], _ToolRound)
|
|
||||||
and segments[end].deterministic_error
|
|
||||||
and segments[end].signature == segment.signature
|
|
||||||
):
|
|
||||||
end += 1
|
|
||||||
group = segments[index:end]
|
|
||||||
should_compact = threshold > 0 and len(group) >= threshold and len(group) > 2
|
|
||||||
if should_compact:
|
|
||||||
projected.extend(group[0].messages)
|
|
||||||
projected.extend(group[-1].messages)
|
|
||||||
removed_rounds += len(group) - 2
|
|
||||||
else:
|
|
||||||
for item in group:
|
|
||||||
projected.extend(item.messages)
|
|
||||||
index = end
|
|
||||||
|
|
||||||
blocked = (
|
|
||||||
segments[-1].tool_names
|
|
||||||
if tail_repetitions and isinstance(segments[-1], _ToolRound)
|
|
||||||
else frozenset()
|
|
||||||
)
|
|
||||||
return RepetitiveToolHistoryRepair(
|
|
||||||
messages=projected,
|
|
||||||
blocked_tool_names=blocked,
|
|
||||||
removed_rounds=removed_rounds,
|
|
||||||
tail_repetitions=tail_repetitions,
|
|
||||||
tail_consecutive_errors=tail_consecutive_errors,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class RepetitiveToolCallGuardMiddleware(AgentMiddleware):
|
|
||||||
"""Stop deterministic loops before another model request is made."""
|
|
||||||
|
|
||||||
name = "repetitive_tool_call_guard"
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
threshold: int = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
|
||||||
max_consecutive_errors: int = DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS,
|
|
||||||
) -> None:
|
|
||||||
super().__init__()
|
|
||||||
for name, value in {
|
|
||||||
"threshold": threshold,
|
|
||||||
"max_consecutive_errors": max_consecutive_errors,
|
|
||||||
}.items():
|
|
||||||
if not isinstance(value, int) or isinstance(value, bool) or value < 0:
|
|
||||||
raise ValueError(f"{name} must be a non-negative integer")
|
|
||||||
self.threshold = threshold
|
|
||||||
self.max_consecutive_errors = max_consecutive_errors
|
|
||||||
|
|
||||||
def _prepare_request(self, request: ModelRequest) -> ModelRequest:
|
|
||||||
repair = collapse_repetitive_tool_rounds(
|
|
||||||
request.messages,
|
|
||||||
threshold=self.threshold,
|
|
||||||
)
|
|
||||||
if self.threshold and repair.tail_repetitions >= self.threshold:
|
|
||||||
raise AgentControlError(
|
|
||||||
"MODEL_TOOL_LOOP_DETECTED",
|
|
||||||
"A deterministic repeated tool-call loop was stopped.",
|
|
||||||
status_code=422,
|
|
||||||
retryable=False,
|
|
||||||
)
|
|
||||||
if (
|
|
||||||
self.max_consecutive_errors
|
|
||||||
and repair.tail_consecutive_errors >= self.max_consecutive_errors
|
|
||||||
):
|
|
||||||
raise AgentControlError(
|
|
||||||
"MODEL_TOOL_ERROR_LIMIT",
|
|
||||||
"Too many consecutive deterministic tool errors were stopped.",
|
|
||||||
status_code=422,
|
|
||||||
retryable=False,
|
|
||||||
)
|
|
||||||
if repair.removed_rounds:
|
|
||||||
logger.info(
|
|
||||||
"Compacted deterministic tool errors for provider projection: removed_rounds=%d",
|
|
||||||
repair.removed_rounds,
|
|
||||||
)
|
|
||||||
return request.override(messages=repair.messages)
|
|
||||||
return request
|
|
||||||
|
|
||||||
def wrap_model_call(
|
|
||||||
self,
|
|
||||||
request: ModelRequest,
|
|
||||||
handler: Callable[[ModelRequest], ModelResponse],
|
|
||||||
) -> ModelResponse:
|
|
||||||
return handler(self._prepare_request(request))
|
|
||||||
|
|
||||||
async def awrap_model_call(
|
|
||||||
self,
|
|
||||||
request: ModelRequest,
|
|
||||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
|
||||||
) -> ModelResponse:
|
|
||||||
return await handler(self._prepare_request(request))
|
|
||||||
@@ -13,12 +13,14 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
from langchain.agents.middleware.types import (
|
from langchain.agents.middleware.types import (
|
||||||
AgentMiddleware,
|
AgentMiddleware,
|
||||||
ModelRequest,
|
ModelRequest,
|
||||||
ModelResponse,
|
ModelResponse,
|
||||||
)
|
)
|
||||||
|
from langchain.tools import ToolRuntime
|
||||||
from langchain_core.tools import tool
|
from langchain_core.tools import tool
|
||||||
|
|
||||||
from .utils import append_to_system_message
|
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
|
@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.
|
"""Create a recurring scheduled task that runs unattended in the background.
|
||||||
|
|
||||||
Translate the user's natural-language timing into a standard 5-field cron
|
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."
|
return "Scheduler unavailable: the langgraph dev backend is not running."
|
||||||
try:
|
try:
|
||||||
rec = crons.create_schedule(
|
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:
|
except Exception as e:
|
||||||
return f"Error: {e}"
|
return f"Error: {e}"
|
||||||
@@ -78,14 +96,14 @@ def schedule_task(name: str, cron: str, prompt: str, timezone: str = "") -> str:
|
|||||||
|
|
||||||
|
|
||||||
@tool
|
@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)."""
|
"""List the user's recurring scheduled tasks (id, name, schedule, enabled)."""
|
||||||
from ..cron import schedule as crons
|
from ..cron import schedule as crons
|
||||||
|
|
||||||
if not crons.is_available():
|
if not crons.is_available():
|
||||||
return "Scheduler unavailable: the langgraph dev backend is not running."
|
return "Scheduler unavailable: the langgraph dev backend is not running."
|
||||||
try:
|
try:
|
||||||
rows = crons.list_schedules()
|
rows = crons.list_schedules(_scope_for_runtime(runtime))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
return f"Error: {e}"
|
return f"Error: {e}"
|
||||||
if not rows:
|
if not rows:
|
||||||
@@ -101,7 +119,7 @@ def list_scheduled_tasks() -> str:
|
|||||||
|
|
||||||
|
|
||||||
@tool
|
@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."""
|
"""Cancel (delete) a scheduled task. Pass the id (or its prefix) shown by list_scheduled_tasks."""
|
||||||
from ..cron import schedule as crons
|
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.
|
# Empty prefix would match (and delete) the only cron — refuse it.
|
||||||
return "Provide the id (or a prefix) of the task to cancel."
|
return "Provide the id (or a prefix) of the task to cancel."
|
||||||
try:
|
try:
|
||||||
rows = crons.list_schedules()
|
rows = crons.list_schedules(_scope_for_runtime(runtime))
|
||||||
# B2: collect ALL prefix matches before acting to detect ambiguity.
|
# B2: collect ALL prefix matches before acting to detect ambiguity.
|
||||||
matches = [
|
matches = [
|
||||||
r for r in rows if str(r.get("cron_id", "")).startswith(requested_id)
|
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:
|
def __init__(self) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self._cache: str | None = None
|
self._cache: dict[str, tuple[float, str]] = {}
|
||||||
self._cache_at: float = 0.0
|
|
||||||
self.tools = [schedule_task, list_scheduled_tasks, cancel_scheduled_task]
|
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)."""
|
"""Build the dynamic ``<scheduled_tasks>`` block (empty if none / down)."""
|
||||||
from ..cron import schedule as crons
|
from ..cron import schedule as crons
|
||||||
|
|
||||||
@@ -158,7 +180,7 @@ class SchedulerMiddleware(AgentMiddleware):
|
|||||||
try:
|
try:
|
||||||
if not crons.is_available():
|
if not crons.is_available():
|
||||||
return ""
|
return ""
|
||||||
rows = crons.list_schedules()
|
rows = crons.list_schedules(scope)
|
||||||
except Exception:
|
except Exception:
|
||||||
return ""
|
return ""
|
||||||
if not rows:
|
if not rows:
|
||||||
@@ -182,10 +204,15 @@ class SchedulerMiddleware(AgentMiddleware):
|
|||||||
|
|
||||||
def _cached_schedules_block(self) -> str:
|
def _cached_schedules_block(self) -> str:
|
||||||
now = time.monotonic()
|
now = time.monotonic()
|
||||||
if self._cache is None or (now - self._cache_at) > _CACHE_TTL_SECONDS:
|
try:
|
||||||
self._cache = self._schedules_block()
|
scope = self._runtime_scope()
|
||||||
self._cache_at = now
|
except Exception:
|
||||||
return self._cache
|
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:
|
def _injection(self, schedules_block: str) -> str:
|
||||||
"""Static instructions, then the dynamic list (static→dynamic, like memory)."""
|
"""Static instructions, then the dynamic list (static→dynamic, like memory)."""
|
||||||
|
|||||||
@@ -1,361 +0,0 @@
|
|||||||
"""Validate completed model tool calls before they can reach ToolNode."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import hashlib
|
|
||||||
import json
|
|
||||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from langchain.agents.middleware.types import (
|
|
||||||
AgentMiddleware,
|
|
||||||
ExtendedModelResponse,
|
|
||||||
ModelRequest,
|
|
||||||
ModelResponse,
|
|
||||||
)
|
|
||||||
from langchain_core.messages import AIMessage
|
|
||||||
from langchain_core.tools import BaseTool
|
|
||||||
|
|
||||||
from ..llm.errors import ModelToolProtocolError, _provider_from_model
|
|
||||||
|
|
||||||
_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call"})
|
|
||||||
_MAX_DIAGNOSTIC_KEYS = 16
|
|
||||||
_MAX_DIAGNOSTIC_KEY_CHARS = 64
|
|
||||||
|
|
||||||
|
|
||||||
def _tool_name(tool: BaseTool | Mapping[str, Any] | Any) -> str | None:
|
|
||||||
if isinstance(tool, BaseTool):
|
|
||||||
return tool.name.strip() or None
|
|
||||||
if isinstance(tool, Mapping):
|
|
||||||
value = tool.get("name")
|
|
||||||
if not value and isinstance(tool.get("function"), Mapping):
|
|
||||||
value = tool["function"].get("name")
|
|
||||||
if isinstance(value, str) and value.strip():
|
|
||||||
return value.strip()
|
|
||||||
return None
|
|
||||||
value = getattr(tool, "name", None)
|
|
||||||
return value.strip() if isinstance(value, str) and value.strip() else None
|
|
||||||
|
|
||||||
|
|
||||||
def _ai_messages(response: Any) -> list[AIMessage]:
|
|
||||||
"""Extract final AI messages from every LangChain middleware response shape."""
|
|
||||||
if isinstance(response, AIMessage):
|
|
||||||
return [response]
|
|
||||||
if isinstance(response, ExtendedModelResponse):
|
|
||||||
response = response.model_response
|
|
||||||
elif not isinstance(response, ModelResponse):
|
|
||||||
nested = getattr(response, "model_response", None)
|
|
||||||
if nested is not None:
|
|
||||||
response = nested
|
|
||||||
result = getattr(response, "result", None)
|
|
||||||
if not isinstance(result, Sequence) or isinstance(result, str | bytes):
|
|
||||||
return []
|
|
||||||
return [message for message in result if isinstance(message, AIMessage)]
|
|
||||||
|
|
||||||
|
|
||||||
def _block_identity(block: Mapping[str, Any]) -> tuple[str, str]:
|
|
||||||
call_id = str(block.get("id") or block.get("call_id") or "").strip()
|
|
||||||
name = block.get("name") or block.get("tool_name")
|
|
||||||
function = block.get("function")
|
|
||||||
if not name and isinstance(function, Mapping):
|
|
||||||
name = function.get("name")
|
|
||||||
return call_id, str(name or "").strip()
|
|
||||||
|
|
||||||
|
|
||||||
def _value_digest(value: Any) -> str:
|
|
||||||
try:
|
|
||||||
encoded = json.dumps(
|
|
||||||
value,
|
|
||||||
ensure_ascii=False,
|
|
||||||
sort_keys=True,
|
|
||||||
separators=(",", ":"),
|
|
||||||
default=lambda item: f"<{type(item).__name__}>",
|
|
||||||
).encode("utf-8")
|
|
||||||
except (TypeError, ValueError):
|
|
||||||
encoded = f"<{type(value).__name__}:unserializable>".encode()
|
|
||||||
return "sha256:" + hashlib.sha256(encoded).hexdigest()[:16]
|
|
||||||
|
|
||||||
|
|
||||||
def _argument_diagnostic(value: Any, *, present: bool) -> dict[str, Any]:
|
|
||||||
if not present:
|
|
||||||
return {"args_present": False, "args_type": "missing"}
|
|
||||||
if isinstance(value, Mapping):
|
|
||||||
keys = sorted(str(key)[:_MAX_DIAGNOSTIC_KEY_CHARS] for key in value)
|
|
||||||
return {
|
|
||||||
"args_present": True,
|
|
||||||
"args_type": "object",
|
|
||||||
"args_key_count": len(keys),
|
|
||||||
"args_keys": keys[:_MAX_DIAGNOSTIC_KEYS],
|
|
||||||
"args_keys_truncated": len(keys) > _MAX_DIAGNOSTIC_KEYS,
|
|
||||||
"args_digest": _value_digest(value),
|
|
||||||
}
|
|
||||||
if isinstance(value, Sequence) and not isinstance(value, str | bytes):
|
|
||||||
value_type = "array"
|
|
||||||
elif isinstance(value, str):
|
|
||||||
value_type = "string"
|
|
||||||
elif value is None:
|
|
||||||
value_type = "null"
|
|
||||||
else:
|
|
||||||
value_type = type(value).__name__
|
|
||||||
return {
|
|
||||||
"args_present": True,
|
|
||||||
"args_type": value_type,
|
|
||||||
"args_digest": _value_digest(value),
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _summarize_call(call: Any) -> dict[str, Any]:
|
|
||||||
if not isinstance(call, Mapping):
|
|
||||||
return {"call_type": type(call).__name__}
|
|
||||||
function = call.get("function")
|
|
||||||
function = function if isinstance(function, Mapping) else {}
|
|
||||||
call_id = str(call.get("id") or call.get("call_id") or "").strip()
|
|
||||||
name = call.get("name") or call.get("tool_name") or function.get("name")
|
|
||||||
name = str(name or "").strip()
|
|
||||||
if "args" in call:
|
|
||||||
args = call.get("args")
|
|
||||||
args_present = True
|
|
||||||
elif "arguments" in call:
|
|
||||||
args = call.get("arguments")
|
|
||||||
args_present = True
|
|
||||||
elif "arguments" in function:
|
|
||||||
args = function.get("arguments")
|
|
||||||
args_present = True
|
|
||||||
else:
|
|
||||||
args = None
|
|
||||||
args_present = False
|
|
||||||
summary = {
|
|
||||||
"call_type": "object",
|
|
||||||
"name": name or "<missing>",
|
|
||||||
"id_present": bool(call_id),
|
|
||||||
**_argument_diagnostic(args, present=args_present),
|
|
||||||
}
|
|
||||||
if call_id:
|
|
||||||
summary["id_fingerprint"] = _value_digest(call_id)
|
|
||||||
return summary
|
|
||||||
|
|
||||||
|
|
||||||
def _raw_openai_call(message: AIMessage, call_index: int) -> Any | None:
|
|
||||||
additional = getattr(message, "additional_kwargs", None)
|
|
||||||
additional = additional if isinstance(additional, Mapping) else {}
|
|
||||||
raw_calls = additional.get("tool_calls")
|
|
||||||
if (
|
|
||||||
isinstance(raw_calls, Sequence)
|
|
||||||
and not isinstance(raw_calls, str | bytes)
|
|
||||||
and call_index < len(raw_calls)
|
|
||||||
):
|
|
||||||
return raw_calls[call_index]
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _call_diagnostic(
|
|
||||||
message: AIMessage,
|
|
||||||
call: Any,
|
|
||||||
*,
|
|
||||||
source: str,
|
|
||||||
call_index: int,
|
|
||||||
call_count: int,
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
diagnostic = {
|
|
||||||
"source": source,
|
|
||||||
"call_index": call_index,
|
|
||||||
"call_count": call_count,
|
|
||||||
**_summarize_call(call),
|
|
||||||
}
|
|
||||||
raw_call = _raw_openai_call(message, call_index)
|
|
||||||
diagnostic["raw_openai_call_available"] = raw_call is not None
|
|
||||||
if raw_call is not None:
|
|
||||||
diagnostic["raw_openai_call"] = _summarize_call(raw_call)
|
|
||||||
return diagnostic
|
|
||||||
|
|
||||||
|
|
||||||
def _route_metadata(request: ModelRequest) -> dict[str, Any]:
|
|
||||||
model = request.model
|
|
||||||
metadata = getattr(model, "metadata", None)
|
|
||||||
metadata = metadata if isinstance(metadata, Mapping) else {}
|
|
||||||
provider = metadata.get("route_provider") or _provider_from_model(model)
|
|
||||||
model_id = metadata.get("route_model")
|
|
||||||
if not model_id:
|
|
||||||
model_id = (
|
|
||||||
getattr(model, "model_name", None)
|
|
||||||
or getattr(model, "model", None)
|
|
||||||
or getattr(model, "model_id", None)
|
|
||||||
)
|
|
||||||
generation = metadata.get("route_config_generation")
|
|
||||||
try:
|
|
||||||
config_generation = int(generation) if generation is not None else None
|
|
||||||
except (TypeError, ValueError):
|
|
||||||
config_generation = None
|
|
||||||
return {
|
|
||||||
"provider": str(provider) if provider else None,
|
|
||||||
"model": str(model_id) if model_id else None,
|
|
||||||
"route_key": str(metadata.get("route_key"))
|
|
||||||
if metadata.get("route_key")
|
|
||||||
else None,
|
|
||||||
"config_generation": config_generation,
|
|
||||||
"api_mode": str(metadata.get("route_api_mode"))
|
|
||||||
if metadata.get("route_api_mode")
|
|
||||||
else None,
|
|
||||||
"endpoint": str(metadata.get("route_endpoint"))
|
|
||||||
if metadata.get("route_endpoint")
|
|
||||||
else None,
|
|
||||||
"tool_call_transport": str(metadata.get("route_tool_call_transport"))
|
|
||||||
if metadata.get("route_tool_call_transport")
|
|
||||||
else None,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _raise_protocol_error(
|
|
||||||
request: ModelRequest,
|
|
||||||
reason: str,
|
|
||||||
*,
|
|
||||||
call_id: str | None = None,
|
|
||||||
call_diagnostic: dict[str, Any] | None = None,
|
|
||||||
) -> None:
|
|
||||||
raise ModelToolProtocolError(
|
|
||||||
reason,
|
|
||||||
call_id=call_id or None,
|
|
||||||
call_diagnostic=call_diagnostic,
|
|
||||||
**_route_metadata(request),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_message(
|
|
||||||
message: AIMessage,
|
|
||||||
request: ModelRequest,
|
|
||||||
allowed_names: frozenset[str],
|
|
||||||
) -> None:
|
|
||||||
invalid_calls = list(getattr(message, "invalid_tool_calls", None) or [])
|
|
||||||
if invalid_calls:
|
|
||||||
invalid = invalid_calls[0]
|
|
||||||
call_id = str(invalid.get("id") or "") if isinstance(invalid, Mapping) else ""
|
|
||||||
_raise_protocol_error(
|
|
||||||
request,
|
|
||||||
"invalid_final_call",
|
|
||||||
call_id=call_id,
|
|
||||||
call_diagnostic=_call_diagnostic(
|
|
||||||
message,
|
|
||||||
invalid,
|
|
||||||
source="invalid_tool_calls",
|
|
||||||
call_index=0,
|
|
||||||
call_count=len(invalid_calls),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
parsed_by_id: dict[str, str] = {}
|
|
||||||
parsed_calls = list(getattr(message, "tool_calls", None) or [])
|
|
||||||
for call_index, raw_call in enumerate(parsed_calls):
|
|
||||||
diagnostic = _call_diagnostic(
|
|
||||||
message,
|
|
||||||
raw_call,
|
|
||||||
source="parsed_tool_calls",
|
|
||||||
call_index=call_index,
|
|
||||||
call_count=len(parsed_calls),
|
|
||||||
)
|
|
||||||
if not isinstance(raw_call, Mapping):
|
|
||||||
_raise_protocol_error(
|
|
||||||
request, "invalid_final_call", call_diagnostic=diagnostic
|
|
||||||
)
|
|
||||||
call_id = str(raw_call.get("id") or raw_call.get("call_id") or "").strip()
|
|
||||||
name = str(raw_call.get("name") or "").strip()
|
|
||||||
if not name:
|
|
||||||
_raise_protocol_error(
|
|
||||||
request,
|
|
||||||
"missing_name",
|
|
||||||
call_id=call_id,
|
|
||||||
call_diagnostic=diagnostic,
|
|
||||||
)
|
|
||||||
if name not in allowed_names:
|
|
||||||
_raise_protocol_error(
|
|
||||||
request,
|
|
||||||
"unknown_name",
|
|
||||||
call_id=call_id,
|
|
||||||
call_diagnostic=diagnostic,
|
|
||||||
)
|
|
||||||
if not call_id:
|
|
||||||
_raise_protocol_error(request, "missing_id", call_diagnostic=diagnostic)
|
|
||||||
if call_id in parsed_by_id:
|
|
||||||
_raise_protocol_error(
|
|
||||||
request,
|
|
||||||
"duplicate_id",
|
|
||||||
call_id=call_id,
|
|
||||||
call_diagnostic=diagnostic,
|
|
||||||
)
|
|
||||||
args = raw_call.get("args")
|
|
||||||
if not isinstance(args, Mapping):
|
|
||||||
_raise_protocol_error(
|
|
||||||
request,
|
|
||||||
"invalid_args",
|
|
||||||
call_id=call_id,
|
|
||||||
call_diagnostic=diagnostic,
|
|
||||||
)
|
|
||||||
parsed_by_id[call_id] = name
|
|
||||||
|
|
||||||
content = getattr(message, "content", None)
|
|
||||||
if not isinstance(content, list):
|
|
||||||
return
|
|
||||||
seen_block_ids: set[str] = set()
|
|
||||||
tool_blocks = [
|
|
||||||
block
|
|
||||||
for block in content
|
|
||||||
if isinstance(block, Mapping) and block.get("type") in _TOOL_BLOCK_TYPES
|
|
||||||
]
|
|
||||||
for block_index, block in enumerate(tool_blocks):
|
|
||||||
diagnostic = _call_diagnostic(
|
|
||||||
message,
|
|
||||||
block,
|
|
||||||
source="content_blocks",
|
|
||||||
call_index=block_index,
|
|
||||||
call_count=len(tool_blocks),
|
|
||||||
)
|
|
||||||
call_id, name = _block_identity(block)
|
|
||||||
if not call_id:
|
|
||||||
_raise_protocol_error(request, "missing_id", call_diagnostic=diagnostic)
|
|
||||||
if call_id in seen_block_ids:
|
|
||||||
_raise_protocol_error(
|
|
||||||
request,
|
|
||||||
"duplicate_id",
|
|
||||||
call_id=call_id,
|
|
||||||
call_diagnostic=diagnostic,
|
|
||||||
)
|
|
||||||
seen_block_ids.add(call_id)
|
|
||||||
parsed_name = parsed_by_id.get(call_id)
|
|
||||||
if parsed_name is None or (name and name != parsed_name):
|
|
||||||
_raise_protocol_error(
|
|
||||||
request,
|
|
||||||
"inconsistent_block",
|
|
||||||
call_id=call_id,
|
|
||||||
call_diagnostic=diagnostic,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class ToolProtocolGuardMiddleware(AgentMiddleware):
|
|
||||||
"""Fail closed on malformed final tool calls using the actual request tools."""
|
|
||||||
|
|
||||||
name = "tool_protocol_guard"
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _validate(response: Any, request: ModelRequest) -> None:
|
|
||||||
allowed_names = frozenset(
|
|
||||||
name for tool in request.tools if (name := _tool_name(tool)) is not None
|
|
||||||
)
|
|
||||||
for message in _ai_messages(response):
|
|
||||||
_validate_message(message, request, allowed_names)
|
|
||||||
|
|
||||||
def wrap_model_call(
|
|
||||||
self,
|
|
||||||
request: ModelRequest,
|
|
||||||
handler: Callable[[ModelRequest], ModelResponse],
|
|
||||||
) -> ModelResponse:
|
|
||||||
response = handler(request)
|
|
||||||
self._validate(response, request)
|
|
||||||
return response
|
|
||||||
|
|
||||||
async def awrap_model_call(
|
|
||||||
self,
|
|
||||||
request: ModelRequest,
|
|
||||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
|
||||||
) -> ModelResponse:
|
|
||||||
response = await handler(request)
|
|
||||||
self._validate(response, request)
|
|
||||||
return response
|
|
||||||
@@ -48,7 +48,6 @@ DEFAULT_ALWAYS_INCLUDE_TOOLS: frozenset[str] = frozenset(
|
|||||||
"read_memory",
|
"read_memory",
|
||||||
"record_observation",
|
"record_observation",
|
||||||
"search_observations",
|
"search_observations",
|
||||||
"write_todos",
|
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -94,15 +93,12 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
|
|||||||
self._threshold = threshold
|
self._threshold = threshold
|
||||||
self._always_include = always_include or frozenset()
|
self._always_include = always_include or frozenset()
|
||||||
self._track_stream_selection = track_stream_selection
|
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:
|
def _build_selector(self, request: ModelRequest) -> AgentMiddleware:
|
||||||
if self._selector is None:
|
# 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)
|
names = _available_always_include(request.tools, self._always_include)
|
||||||
self._selector = self._selector_factory(names)
|
return self._selector_factory(names)
|
||||||
return self._selector
|
|
||||||
|
|
||||||
def wrap_model_call(
|
def wrap_model_call(
|
||||||
self,
|
self,
|
||||||
@@ -133,21 +129,10 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
|
|||||||
return self._build_selector(request).wrap_model_call(
|
return self._build_selector(request).wrap_model_call(
|
||||||
request, _handler_after_selection
|
request, _handler_after_selection
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception:
|
||||||
if _handler_called:
|
if _handler_called:
|
||||||
raise # Error from downstream model — don't retry
|
raise # Error from downstream model — don't retry
|
||||||
from ..llm.errors import ProviderStreamError
|
# Selector itself failed (e.g., structured output not supported).
|
||||||
from .error_normalization import _is_provider_error
|
|
||||||
|
|
||||||
if isinstance(exc, ProviderStreamError) or _is_provider_error(exc):
|
|
||||||
# Auth / quota / connection failures on the selector's
|
|
||||||
# own model. Falling back to "use all tools" would hit
|
|
||||||
# the same provider anyway (same client, likely same
|
|
||||||
# credentials). Surface it instead so the user sees
|
|
||||||
# the real cause.
|
|
||||||
raise
|
|
||||||
# Structured-output shape / config failure — gracefully
|
|
||||||
# degrade to using all tools.
|
|
||||||
logger.debug("Tool selector failed, using all tools", exc_info=True)
|
logger.debug("Tool selector failed, using all tools", exc_info=True)
|
||||||
if self._track_stream_selection:
|
if self._track_stream_selection:
|
||||||
_selector_active = False
|
_selector_active = False
|
||||||
@@ -183,16 +168,9 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
|
|||||||
return await self._build_selector(request).awrap_model_call(
|
return await self._build_selector(request).awrap_model_call(
|
||||||
request, _handler_after_selection
|
request, _handler_after_selection
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception:
|
||||||
if _handler_called:
|
if _handler_called:
|
||||||
raise
|
raise
|
||||||
from ..llm.errors import ProviderStreamError
|
|
||||||
from .error_normalization import _is_provider_error
|
|
||||||
|
|
||||||
if isinstance(exc, ProviderStreamError) or _is_provider_error(exc):
|
|
||||||
# See sync path — surface provider errors, degrade only
|
|
||||||
# on shape / config failures.
|
|
||||||
raise
|
|
||||||
logger.debug("Tool selector failed, using all tools", exc_info=True)
|
logger.debug("Tool selector failed, using all tools", exc_info=True)
|
||||||
if self._track_stream_selection:
|
if self._track_stream_selection:
|
||||||
_selector_active = False
|
_selector_active = False
|
||||||
@@ -252,8 +230,10 @@ def create_tool_selector_middleware(
|
|||||||
names for the main-agent stream UI when ``track_stream_selection`` is true
|
names for the main-agent stream UI when ``track_stream_selection`` is true
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model: Chat model for tool selection. If *None*, the default
|
model: Chat model for tool selection. If *None*, the helper model
|
||||||
model is resolved via ``_ensure_chat_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.
|
threshold: Minimum number of tools to trigger selection.
|
||||||
Default 26. Set to 0 to always run selection.
|
Default 26. Set to 0 to always run selection.
|
||||||
track_stream_selection: Whether to update process-global stream/UI
|
track_stream_selection: Whether to update process-global stream/UI
|
||||||
@@ -272,20 +252,50 @@ def create_tool_selector_middleware(
|
|||||||
|
|
||||||
from .utils import disable_thinking
|
from .utils import disable_thinking
|
||||||
|
|
||||||
if model is None:
|
def tag_selector_model(base: BaseChatModel) -> BaseChatModel:
|
||||||
from EvoScientist.EvoScientist import _ensure_chat_model
|
safe_model = disable_thinking(base)
|
||||||
|
selector_model = safe_model
|
||||||
|
from EvoScientist.usage.callback import usage_tracking_enabled
|
||||||
|
|
||||||
model = _ensure_chat_model()
|
if usage_tracking_enabled():
|
||||||
safe_model = disable_thinking(model)
|
try:
|
||||||
safe_model = safe_model.model_copy(
|
selector_model = safe_model.model_copy(
|
||||||
update={
|
update={
|
||||||
"tags": [*(safe_model.tags or []), "metering:tool_selector"],
|
|
||||||
"metadata": {
|
"metadata": {
|
||||||
**(safe_model.metadata or {}),
|
**(safe_model.metadata or {}),
|
||||||
"metering_scope": "tool_selector",
|
"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 = (
|
system_prompt = (
|
||||||
"You are selecting tools for a scientific research agent. "
|
"You are selecting tools for a scientific research agent. "
|
||||||
@@ -298,7 +308,7 @@ def create_tool_selector_middleware(
|
|||||||
|
|
||||||
def selector_factory(always_include: list[str]) -> AgentMiddleware:
|
def selector_factory(always_include: list[str]) -> AgentMiddleware:
|
||||||
return LLMToolSelectorMiddleware(
|
return LLMToolSelectorMiddleware(
|
||||||
model=safe_model,
|
model=selector_model(),
|
||||||
system_prompt=system_prompt,
|
system_prompt=system_prompt,
|
||||||
always_include=always_include,
|
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
@@ -5,7 +5,6 @@ from __future__ import annotations
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import shutil
|
import shutil
|
||||||
from collections.abc import Iterator
|
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
@@ -216,65 +215,3 @@ def resolve_virtual_path(virtual_path: str) -> Path:
|
|||||||
"""Resolve a virtual workspace path (e.g. /image.png) to a real filesystem path."""
|
"""Resolve a virtual workspace path (e.g. /image.png) to a real filesystem path."""
|
||||||
vpath = virtual_path if virtual_path.startswith("/") else "/" + virtual_path
|
vpath = virtual_path if virtual_path.startswith("/") else "/" + virtual_path
|
||||||
return (_active_workspace / vpath.lstrip("/")).resolve()
|
return (_active_workspace / vpath.lstrip("/")).resolve()
|
||||||
|
|
||||||
|
|
||||||
def evoscientist_root() -> Path:
|
|
||||||
"""Return the application root used by Gateway-managed runtime data."""
|
|
||||||
env_root = os.environ.get("EVOSCIENTIST_HOME")
|
|
||||||
if env_root:
|
|
||||||
return Path(env_root).expanduser().resolve()
|
|
||||||
return DATA_DIR.expanduser().resolve()
|
|
||||||
|
|
||||||
|
|
||||||
_EVOSCIENTIST_DATA_ROOT: Path | None = None
|
|
||||||
|
|
||||||
|
|
||||||
def _data_root() -> Path:
|
|
||||||
"""Return the root directory for isolated Web user workspaces."""
|
|
||||||
global _EVOSCIENTIST_DATA_ROOT
|
|
||||||
if _EVOSCIENTIST_DATA_ROOT is not None:
|
|
||||||
return _EVOSCIENTIST_DATA_ROOT
|
|
||||||
|
|
||||||
env_root = os.environ.get("EVOSCIENTIST_DATA_ROOT")
|
|
||||||
if env_root:
|
|
||||||
root = Path(env_root).expanduser().resolve()
|
|
||||||
else:
|
|
||||||
root = evoscientist_root() / "data"
|
|
||||||
_EVOSCIENTIST_DATA_ROOT = root
|
|
||||||
return root
|
|
||||||
|
|
||||||
|
|
||||||
def user_data_dir(user_id: str) -> Path:
|
|
||||||
"""Return and create the isolated data directory for a Web user."""
|
|
||||||
path = _data_root() / user_id
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
return path
|
|
||||||
|
|
||||||
|
|
||||||
def iter_user_data_dirs() -> Iterator[Path]:
|
|
||||||
"""Yield existing Web user directories without creating the data root."""
|
|
||||||
root = _data_root()
|
|
||||||
if not root.exists():
|
|
||||||
return
|
|
||||||
for path in root.iterdir():
|
|
||||||
if path.is_dir():
|
|
||||||
yield path
|
|
||||||
|
|
||||||
|
|
||||||
def thread_data_dir(user_id: str, thread_id: str) -> Path:
|
|
||||||
"""Return and create a user's isolated thread workspace."""
|
|
||||||
path = user_data_dir(user_id) / thread_id
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
return path
|
|
||||||
|
|
||||||
|
|
||||||
def global_data_dir(user_id: str) -> Path:
|
|
||||||
"""Return and create a user's directory shared across all threads."""
|
|
||||||
path = user_data_dir(user_id) / "__global__"
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
return path
|
|
||||||
|
|
||||||
|
|
||||||
def uploads_dir() -> Path:
|
|
||||||
"""Return the Gateway upload staging directory."""
|
|
||||||
return evoscientist_root() / "uploads"
|
|
||||||
|
|||||||
+14
-3
@@ -236,6 +236,15 @@ WRITING_GUIDELINES = """# Writing Guidelines
|
|||||||
- Professional, objective tone. Be precise, technical, and concise.
|
- 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)
|
# Shell execution guidelines (rules for the `execute` tool)
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -410,9 +419,10 @@ def get_system_prompt(
|
|||||||
2. :data:`EXPERIMENT_WORKFLOW`
|
2. :data:`EXPERIMENT_WORKFLOW`
|
||||||
3. :data:`REPORT_TEMPLATE`
|
3. :data:`REPORT_TEMPLATE`
|
||||||
4. :data:`WRITING_GUIDELINES`
|
4. :data:`WRITING_GUIDELINES`
|
||||||
5. :data:`SHELL_GUIDELINES` (or :data:`SHELL_GUIDELINES_DANGEROUS`)
|
5. :data:`FILE_REFERENCES`
|
||||||
6. :data:`DELEGATION_STRATEGY`
|
6. :data:`SHELL_GUIDELINES` (or :data:`SHELL_GUIDELINES_DANGEROUS`)
|
||||||
7. :data:`ASYNC_NOTIFICATIONS`
|
7. :data:`DELEGATION_STRATEGY`
|
||||||
|
8. :data:`ASYNC_NOTIFICATIONS`
|
||||||
|
|
||||||
Runtime context is injected per-turn by
|
Runtime context is injected per-turn by
|
||||||
:class:`EvoScientist.middleware.RuntimeContextMiddleware`, so dates and
|
:class:`EvoScientist.middleware.RuntimeContextMiddleware`, so dates and
|
||||||
@@ -439,6 +449,7 @@ def get_system_prompt(
|
|||||||
EXPERIMENT_WORKFLOW,
|
EXPERIMENT_WORKFLOW,
|
||||||
REPORT_TEMPLATE,
|
REPORT_TEMPLATE,
|
||||||
WRITING_GUIDELINES,
|
WRITING_GUIDELINES,
|
||||||
|
FILE_REFERENCES,
|
||||||
shell_guidelines,
|
shell_guidelines,
|
||||||
DELEGATION_STRATEGY,
|
DELEGATION_STRATEGY,
|
||||||
ASYNC_NOTIFICATIONS,
|
ASYNC_NOTIFICATIONS,
|
||||||
|
|||||||
@@ -1,116 +0,0 @@
|
|||||||
"""Optional runtime services supplied by an application embedding EvoScientist.
|
|
||||||
|
|
||||||
The CLI package must not import a concrete web gateway. Applications such as
|
|
||||||
Ai4Sci-Web can register their database, storage, metering, and media services
|
|
||||||
at process startup through this module.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from collections.abc import Awaitable, Callable
|
|
||||||
from dataclasses import dataclass, replace
|
|
||||||
from datetime import date
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
AsyncProvider = Callable[[], Awaitable[Any]]
|
|
||||||
AsyncFileHandler = Callable[[Path], Awaitable[Any]]
|
|
||||||
AsyncUsageRecorder = Callable[[str, str], Awaitable[Any]]
|
|
||||||
ModelResolver = Callable[[str | None, str | None], Any | None]
|
|
||||||
|
|
||||||
|
|
||||||
class RuntimeIntegrationUnavailable(RuntimeError):
|
|
||||||
"""Raised when an optional host-provided service is not configured."""
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
|
||||||
class RuntimeIntegrations:
|
|
||||||
app_connection_provider: AsyncProvider | None = None
|
|
||||||
session_connection_provider: AsyncProvider | None = None
|
|
||||||
session_dsn_provider: Callable[[], str | None] | None = None
|
|
||||||
current_date_provider: Callable[[], date] | None = None
|
|
||||||
user_storage_root_provider: Callable[[str], Path] | None = None
|
|
||||||
knowledge_file_handler: AsyncFileHandler | None = None
|
|
||||||
usage_recorder: AsyncUsageRecorder | None = None
|
|
||||||
image_backend_factory: Callable[[], Any] | None = None
|
|
||||||
model_resolver: ModelResolver | None = None
|
|
||||||
|
|
||||||
|
|
||||||
_integrations = RuntimeIntegrations()
|
|
||||||
|
|
||||||
|
|
||||||
def configure_runtime_integrations(**services: Any) -> RuntimeIntegrations:
|
|
||||||
"""Register host-provided services and return the resulting configuration."""
|
|
||||||
global _integrations
|
|
||||||
_integrations = replace(_integrations, **services)
|
|
||||||
return _integrations
|
|
||||||
|
|
||||||
|
|
||||||
def reset_runtime_integrations() -> None:
|
|
||||||
"""Clear all host-provided services, primarily for tests."""
|
|
||||||
global _integrations
|
|
||||||
_integrations = RuntimeIntegrations()
|
|
||||||
|
|
||||||
|
|
||||||
def has_session_connection_provider() -> bool:
|
|
||||||
return _integrations.session_connection_provider is not None
|
|
||||||
|
|
||||||
|
|
||||||
def get_session_dsn() -> str | None:
|
|
||||||
provider = _integrations.session_dsn_provider
|
|
||||||
return provider() if provider is not None else None
|
|
||||||
|
|
||||||
|
|
||||||
async def get_session_connection() -> Any:
|
|
||||||
provider = _integrations.session_connection_provider
|
|
||||||
if provider is None:
|
|
||||||
raise RuntimeIntegrationUnavailable(
|
|
||||||
"No session connection provider is configured"
|
|
||||||
)
|
|
||||||
return await provider()
|
|
||||||
|
|
||||||
|
|
||||||
async def get_app_connection() -> Any:
|
|
||||||
provider = _integrations.app_connection_provider
|
|
||||||
if provider is None:
|
|
||||||
raise RuntimeIntegrationUnavailable(
|
|
||||||
"No application connection provider is configured"
|
|
||||||
)
|
|
||||||
return await provider()
|
|
||||||
|
|
||||||
|
|
||||||
def current_date() -> date:
|
|
||||||
provider = _integrations.current_date_provider
|
|
||||||
return provider() if provider is not None else date.today()
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_user_storage_root(user_id: str) -> Path | None:
|
|
||||||
provider = _integrations.user_storage_root_provider
|
|
||||||
return provider(user_id) if provider is not None else None
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_runtime_model(model: str | None, provider: str | None = None) -> Any | None:
|
|
||||||
"""Resolve a host-managed model configuration when one is registered."""
|
|
||||||
resolver = _integrations.model_resolver
|
|
||||||
return resolver(model, provider) if resolver is not None else None
|
|
||||||
|
|
||||||
|
|
||||||
async def handle_knowledge_file(path: Path) -> None:
|
|
||||||
handler = _integrations.knowledge_file_handler
|
|
||||||
if handler is not None:
|
|
||||||
await handler(path)
|
|
||||||
|
|
||||||
|
|
||||||
async def record_service_usage(service: str, action: str) -> None:
|
|
||||||
recorder = _integrations.usage_recorder
|
|
||||||
if recorder is not None:
|
|
||||||
await recorder(service, action)
|
|
||||||
|
|
||||||
|
|
||||||
def get_image_backend() -> Any:
|
|
||||||
factory = _integrations.image_backend_factory
|
|
||||||
if factory is None:
|
|
||||||
raise RuntimeIntegrationUnavailable(
|
|
||||||
"Image generation is unavailable in this runtime. Configure an image backend first."
|
|
||||||
)
|
|
||||||
return factory()
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -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.
|
"""Improve a skill description based on eval results.
|
||||||
|
|
||||||
Takes eval results (from run_eval.py) and generates an improved description
|
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
|
import argparse
|
||||||
@@ -16,8 +17,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
|||||||
|
|
||||||
from langchain_core.messages import HumanMessage
|
from langchain_core.messages import HumanMessage
|
||||||
|
|
||||||
from EvoScientist.config import apply_config_to_env, get_effective_config
|
from scripts.run_eval import _config_defaults, _init_chat_model
|
||||||
from EvoScientist.llm import get_chat_model
|
|
||||||
from scripts.utils import parse_skill_md
|
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."""
|
Please respond with only the new description text in <new_description> tags, nothing else."""
|
||||||
|
|
||||||
# Initialise model via EvoSci's LLM layer
|
# Initialise model (API keys come from the environment / EvoSci config)
|
||||||
config = get_effective_config()
|
default_model, default_provider = _config_defaults()
|
||||||
apply_config_to_env(config)
|
effective_model = model or default_model
|
||||||
chat_model = get_chat_model(
|
if not effective_model:
|
||||||
model=model or config.model,
|
raise RuntimeError(
|
||||||
provider=provider or config.provider,
|
"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)])
|
response = chat_model.invoke([HumanMessage(content=prompt)])
|
||||||
|
|
||||||
@@ -253,9 +255,8 @@ def main():
|
|||||||
name, _, content = parse_skill_md(skill_path)
|
name, _, content = parse_skill_md(skill_path)
|
||||||
current_description = eval_results["description"]
|
current_description = eval_results["description"]
|
||||||
|
|
||||||
# Load EvoSci config for defaults
|
# Load EvoSci config for defaults (best-effort; env vars alone suffice)
|
||||||
config = get_effective_config()
|
default_model, default_provider = _config_defaults()
|
||||||
apply_config_to_env(config)
|
|
||||||
|
|
||||||
if args.verbose:
|
if args.verbose:
|
||||||
print(f"Current: {current_description}", file=sys.stderr)
|
print(f"Current: {current_description}", file=sys.stderr)
|
||||||
@@ -270,8 +271,8 @@ def main():
|
|||||||
current_description=current_description,
|
current_description=current_description,
|
||||||
eval_results=eval_results,
|
eval_results=eval_results,
|
||||||
history=history,
|
history=history,
|
||||||
model=args.model or config.model,
|
model=args.model or default_model,
|
||||||
provider=args.provider or config.provider,
|
provider=args.provider or default_provider,
|
||||||
)
|
)
|
||||||
|
|
||||||
if args.verbose:
|
if args.verbose:
|
||||||
|
|||||||
@@ -2,8 +2,10 @@
|
|||||||
"""Run trigger evaluation for a skill description.
|
"""Run trigger evaluation for a skill description.
|
||||||
|
|
||||||
Tests whether a skill's description causes an LLM to trigger (load the skill)
|
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
|
for a set of queries. Uses ``langchain.chat_models.init_chat_model`` with
|
||||||
to simulate the agent's skill selection behavior.
|
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
|
import argparse
|
||||||
@@ -18,13 +20,28 @@ sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
|||||||
from scripts.utils import parse_skill_md
|
from scripts.utils import parse_skill_md
|
||||||
|
|
||||||
|
|
||||||
def _init_config():
|
def _config_defaults() -> tuple[str | None, str | None]:
|
||||||
"""Initialize EvoSci config and apply env vars (once per process)."""
|
"""Best-effort (model, provider) defaults from the EvoSci 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
|
from EvoScientist.config import apply_config_to_env, get_effective_config
|
||||||
|
|
||||||
config = get_effective_config()
|
config = get_effective_config()
|
||||||
apply_config_to_env(config)
|
apply_config_to_env(config)
|
||||||
return 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(
|
def run_single_query(
|
||||||
@@ -43,12 +60,14 @@ def run_single_query(
|
|||||||
from langchain_core.messages import HumanMessage, SystemMessage
|
from langchain_core.messages import HumanMessage, SystemMessage
|
||||||
from langchain_core.tools import tool
|
from langchain_core.tools import tool
|
||||||
|
|
||||||
from EvoScientist.llm import get_chat_model
|
default_model, default_provider = _config_defaults()
|
||||||
|
effective_model = model or default_model
|
||||||
config = _init_config()
|
effective_provider = provider or default_provider
|
||||||
|
if not effective_model:
|
||||||
effective_model = model or config.model
|
raise RuntimeError(
|
||||||
effective_provider = provider or config.provider
|
"No model specified: pass --model or configure one (or set a "
|
||||||
|
"provider's default model via environment)."
|
||||||
|
)
|
||||||
|
|
||||||
@tool
|
@tool
|
||||||
def load_skill(name: str) -> str:
|
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
|
pass
|
||||||
|
|
||||||
try:
|
try:
|
||||||
chat_model = get_chat_model(
|
chat_model = _init_chat_model(
|
||||||
model=effective_model,
|
effective_model,
|
||||||
provider=effective_provider,
|
effective_provider,
|
||||||
**eval_kwargs,
|
**eval_kwargs,
|
||||||
)
|
)
|
||||||
model_with_tools = chat_model.bind_tools([load_skill])
|
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:
|
# Fallback for providers that don't support tool calling:
|
||||||
# Use a text-based approach
|
# Use a text-based approach
|
||||||
try:
|
try:
|
||||||
chat_model = get_chat_model(
|
chat_model = _init_chat_model(
|
||||||
model=effective_model,
|
effective_model,
|
||||||
provider=effective_provider,
|
effective_provider,
|
||||||
**eval_kwargs,
|
**eval_kwargs,
|
||||||
)
|
)
|
||||||
fallback_prompt = f"""You are a helpful AI assistant with specialized skills available.
|
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"
|
"--trigger-threshold", type=float, default=0.5, help="Trigger rate threshold"
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
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(
|
parser.add_argument(
|
||||||
"--provider",
|
"--provider",
|
||||||
default=None,
|
default=None,
|
||||||
help="LLM provider (default: user's configured provider)",
|
help="LLM provider (default: EvoSci config provider, else inferred from model)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--verbose", action="store_true", help="Print progress to stderr"
|
"--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
|
# Ensure skill-creator root is on sys.path for `from scripts.xxx` imports
|
||||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
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.generate_report import generate_html
|
||||||
from scripts.improve_description import improve_description
|
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
|
from scripts.utils import parse_skill_md
|
||||||
|
|
||||||
|
|
||||||
@@ -66,8 +65,10 @@ def run_loop(
|
|||||||
log_dir: Path | None = None,
|
log_dir: Path | None = None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""Run the eval + improvement loop."""
|
"""Run the eval + improvement loop."""
|
||||||
config = get_effective_config()
|
# Seed API keys into the environment from the EvoSci config (best-effort);
|
||||||
apply_config_to_env(config)
|
# 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)
|
name, original_description, content = parse_skill_md(skill_path)
|
||||||
current_description = description_override or original_description
|
current_description = description_override or original_description
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user