mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-12 23:29:16 +03:00
Compare commits
62
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0c684c5a99 | ||
|
|
d3382d7e57 | ||
|
|
76f629e925 | ||
|
|
5f916bbd3a | ||
|
|
31a71d6cd5 | ||
|
|
498070d036 | ||
|
|
b7f0ae95a9 | ||
|
|
b571d3b9ff | ||
|
|
3741ecda0b | ||
|
|
edec29e997 | ||
|
|
01c7323d74 | ||
|
|
001a7492c2 | ||
|
|
6fc0807fbf | ||
|
|
cd7480945b | ||
|
|
45245b5e55 | ||
|
|
d2cbe6536e | ||
|
|
b34f1bd0e8 | ||
|
|
edaef4e4f5 | ||
|
|
5fc8303f9e | ||
|
|
e455a2b7fa | ||
|
|
19997d20bb | ||
|
|
686dd0603e | ||
|
|
4b5319b760 | ||
|
|
1656664a47 | ||
|
|
a6193932a0 | ||
|
|
bcf5d8a6ed | ||
|
|
d64b84604c | ||
|
|
abfcdd481a | ||
|
|
a0e60116a3 | ||
|
|
ec3dfb21ba | ||
|
|
72d3ce6b23 | ||
|
|
b14ac4c401 | ||
|
|
f5cf4dcd2c | ||
|
|
99e07e138e | ||
|
|
057e8f7af6 | ||
|
|
d45c893f68 | ||
|
|
1edfd268db | ||
|
|
c0e8b8afff | ||
|
|
95287f7435 | ||
|
|
43ca12960b | ||
|
|
7703cd22eb | ||
|
|
247c474e64 | ||
|
|
86c7508607 | ||
|
|
a2979c3a4b | ||
|
|
d5e0df6963 | ||
|
|
57d81bc1cd | ||
|
|
3778e7e628 | ||
|
|
cac39477ba | ||
|
|
eab017766b | ||
|
|
52e0a6a1e3 | ||
|
|
8e77f3f8a4 | ||
|
|
b3b0517611 | ||
|
|
c281e090d0 | ||
|
|
85a452e5c7 | ||
|
|
05d73803e7 | ||
|
|
5d733b1c7c | ||
|
|
71a99b0780 | ||
|
|
43511decc9 | ||
|
|
8dd2059be3 | ||
|
|
7b1646f58c | ||
|
|
e620944150 | ||
|
|
66316f21da |
@@ -189,6 +189,44 @@ jobs:
|
|||||||
- name: Build image with default channel dependencies
|
- name: Build image with default channel dependencies
|
||||||
run: docker build -t nanobot:test .
|
run: docker build -t nanobot:test .
|
||||||
|
|
||||||
|
- name: Verify Docker Compose startup and privilege boundary
|
||||||
|
env:
|
||||||
|
HOME: ${{ runner.temp }}
|
||||||
|
run: |
|
||||||
|
docker compose run --rm --no-deps --build -T nanobot-cli status
|
||||||
|
docker compose run --rm --no-deps -T --entrypoint sh nanobot-cli -s <<'OUTER'
|
||||||
|
set -eu
|
||||||
|
field() {
|
||||||
|
awk -v key="$1:" '$1 == key { print $2 }' /proc/self/status
|
||||||
|
}
|
||||||
|
test "$(id -u)" = "0"
|
||||||
|
test "$(field NoNewPrivs)" = "1"
|
||||||
|
setpriv --reuid=nanobot --regid=nanobot --init-groups sh -s <<'INNER'
|
||||||
|
set -eu
|
||||||
|
field() {
|
||||||
|
awk -v key="$1:" '$1 == key { print $2 }' /proc/self/status
|
||||||
|
}
|
||||||
|
test "$(id -u)" = "1000"
|
||||||
|
test "$(field NoNewPrivs)" = "1"
|
||||||
|
for capability_set in CapInh CapPrm CapEff CapAmb; do
|
||||||
|
test "$(field "$capability_set")" = "0000000000000000"
|
||||||
|
done
|
||||||
|
INNER
|
||||||
|
OUTER
|
||||||
|
docker compose -f docker-compose.yml -f docker-compose.bwrap.yml --profile cli \
|
||||||
|
config --format json > "${RUNNER_TEMP}/bwrap-compose.json"
|
||||||
|
python - <<'PY'
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
config = json.loads(Path(os.environ["RUNNER_TEMP"], "bwrap-compose.json").read_text())
|
||||||
|
for service_name in ("nanobot-gateway", "nanobot-api", "nanobot-cli"):
|
||||||
|
service = config["services"][service_name]
|
||||||
|
assert {"CHOWN", "SETGID", "SETUID", "SYS_ADMIN"} <= set(service["cap_add"])
|
||||||
|
assert "no-new-privileges:true" in service["security_opt"]
|
||||||
|
PY
|
||||||
|
|
||||||
- name: Verify default WhatsApp dependencies
|
- name: Verify default WhatsApp dependencies
|
||||||
run: docker run --rm --entrypoint python nanobot:test -c "import neonize, segno"
|
run: docker run --rm --entrypoint python nanobot:test -c "import neonize, segno"
|
||||||
|
|
||||||
|
|||||||
@@ -241,7 +241,7 @@ Prefer your own infrastructure? Follow the [deployment guide](./docs/deployment.
|
|||||||
|
|
||||||
## 🌐 WebUI
|
## 🌐 WebUI
|
||||||
|
|
||||||
The WebUI ships **inside the published wheel** with no separate frontend build. It is the browser workbench for persistent topics, visible agent activity, workspace controls, Apps, Skills, Automations, and settings.
|
The WebUI ships **inside the published wheel** with no separate frontend build. It is the browser workbench for persistent topics, temporary chats, visible agent activity, workspace controls, Apps, Skills, Automations, and settings.
|
||||||
|
|
||||||
<p align="center">
|
<p align="center">
|
||||||
<img src="images/nanobot_webui.png" alt="nanobot webui preview" width="900">
|
<img src="images/nanobot_webui.png" alt="nanobot webui preview" width="900">
|
||||||
@@ -250,9 +250,10 @@ The WebUI ships **inside the published wheel** with no separate frontend build.
|
|||||||
Use it to:
|
Use it to:
|
||||||
|
|
||||||
- keep separate topics for different tasks and projects;
|
- keep separate topics for different tasks and projects;
|
||||||
|
- use temporary chats when a conversation should not be saved to history or memory;
|
||||||
- inspect reasoning, tool calls, file edits, diffs, command output, and generated artifacts;
|
- inspect reasoning, tool calls, file edits, diffs, command output, and generated artifacts;
|
||||||
- switch models and workspaces without leaving the conversation;
|
- switch models and workspaces without leaving the conversation;
|
||||||
- configure providers, chat channels, Apps, Skills, and Automations from one place.
|
- configure providers and chat channels, connect Apps, discover Skills, and manage Automations from one place.
|
||||||
|
|
||||||
See the [WebUI guide](./docs/webui.md) for LAN access, background operation, workspace controls, and the full feature tour. Working on the frontend itself? Use [`webui/README.md`](./webui/README.md).
|
See the [WebUI guide](./docs/webui.md) for LAN access, background operation, workspace controls, and the full feature tour. Working on the frontend itself? Use [`webui/README.md`](./webui/README.md).
|
||||||
|
|
||||||
|
|||||||
+27
@@ -6,6 +6,7 @@ import os
|
|||||||
import ssl
|
import ssl
|
||||||
import sys
|
import sys
|
||||||
from collections.abc import Iterator
|
from collections.abc import Iterator
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
import certifi
|
import certifi
|
||||||
import pytest
|
import pytest
|
||||||
@@ -22,6 +23,32 @@ def _isolate_nanobot_log_activation() -> Iterator[None]:
|
|||||||
logger.enable("nanobot")
|
logger.enable("nanobot")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _isolate_sessions_root(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
|
||||||
|
"""Redirect session storage away from the real active config data directory.
|
||||||
|
|
||||||
|
Session storage lives under the active runtime data root (outside the workspace,
|
||||||
|
per ADR-0001), so without redirection tests would write into the real home.
|
||||||
|
"""
|
||||||
|
runtime_root = tmp_path.parent / f"{tmp_path.name}-runtime-root"
|
||||||
|
legacy_root = tmp_path.parent / f"{tmp_path.name}-legacy-sessions-root"
|
||||||
|
|
||||||
|
def runtime_subdir(name: str) -> Path:
|
||||||
|
path = runtime_root / name
|
||||||
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
|
return path
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.session.manager.get_runtime_subdir",
|
||||||
|
runtime_subdir,
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.session.manager.get_legacy_sessions_dir",
|
||||||
|
lambda: legacy_root,
|
||||||
|
)
|
||||||
|
yield
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session", autouse=True)
|
@pytest.fixture(scope="session", autouse=True)
|
||||||
def _use_windows_system_ca_for_default_http_clients() -> Iterator[None]:
|
def _use_windows_system_ca_for_default_http_clients() -> Iterator[None]:
|
||||||
"""Avoid reparsing certifi's CA bundle for every offline HTTP client.
|
"""Avoid reparsing certifi's CA bundle for every offline HTTP client.
|
||||||
|
|||||||
@@ -8,6 +8,15 @@ x-common-config: &common-config
|
|||||||
- ~/.nanobot:/home/nanobot/.nanobot
|
- ~/.nanobot:/home/nanobot/.nanobot
|
||||||
cap_drop:
|
cap_drop:
|
||||||
- ALL
|
- ALL
|
||||||
|
# Entrypoint uses these to fix bind-mount ownership and drop to the nanobot user.
|
||||||
|
cap_add:
|
||||||
|
- CHOWN
|
||||||
|
- SETGID
|
||||||
|
- SETUID
|
||||||
|
# Prevent the non-root process from regaining capabilities through setuid
|
||||||
|
# binaries or file capabilities left inside the container image.
|
||||||
|
security_opt:
|
||||||
|
- no-new-privileges:true
|
||||||
|
|
||||||
services:
|
services:
|
||||||
nanobot-gateway:
|
nanobot-gateway:
|
||||||
|
|||||||
+2
-1
@@ -19,7 +19,7 @@ The recommended first-run path is:
|
|||||||
3. Configure a provider and model in **Settings → Models**.
|
3. Configure a provider and model in **Settings → Models**.
|
||||||
4. Send `Hello!` before configuring anything else.
|
4. Send `Hello!` before configuring anything else.
|
||||||
|
|
||||||
Most people do not need to edit JSON for the first run. The WebUI handles the initial provider, model, and local browser settings. SSH, headless, existing-config, and older-release installs retain `nanobot onboard --wizard` as a terminal fallback. After the WebUI opens, use **Settings** for models and built-in capabilities, **Settings → Channels** for chat apps, and **Apps** for CLI App or MCP integrations.
|
Most people do not need to edit JSON for the first run. The WebUI handles the initial provider, model, and local browser settings. SSH, headless, existing-config, and older-release installs retain `nanobot onboard --wizard` as a terminal fallback. After the WebUI opens, use **Settings** for models and built-in capabilities, **Settings → Channels** for chat apps, and **Apps** for Agent Plugins, CLI Apps, and MCP integrations.
|
||||||
|
|
||||||
## Add One Capability
|
## Add One Capability
|
||||||
|
|
||||||
@@ -32,6 +32,7 @@ Pick the row that matches what you want to accomplish next:
|
|||||||
| Choose a hosted, OAuth, company, or local model | [Provider Cookbook](./provider-cookbook.md) |
|
| Choose a hosted, OAuth, company, or local model | [Provider Cookbook](./provider-cookbook.md) |
|
||||||
| Add model fallbacks | [Configure Model Fallback](./guides/configure-model-fallback.md) |
|
| Add model fallbacks | [Configure Model Fallback](./guides/configure-model-fallback.md) |
|
||||||
| Enable web search | [Configure Web Search](./guides/configure-web-search.md) |
|
| Enable web search | [Configure Web Search](./guides/configure-web-search.md) |
|
||||||
|
| Manage Agent Plugins, CLI Apps, or MCP integrations | [WebUI Apps](./webui.md#apps) |
|
||||||
| Add an MCP tool server | [Configure MCP Tools](./guides/configure-mcp-tools.md) |
|
| Add an MCP tool server | [Configure MCP Tools](./guides/configure-mcp-tools.md) |
|
||||||
| Generate images | [Image Generation](./image-generation.md) |
|
| Generate images | [Image Generation](./image-generation.md) |
|
||||||
| Schedule work or create a local trigger | [Automations](./automations.md) |
|
| Schedule work or create a local trigger | [Automations](./automations.md) |
|
||||||
|
|||||||
+14
-5
@@ -51,6 +51,13 @@ Main files:
|
|||||||
- feeds tool results back into the model;
|
- feeds tool results back into the model;
|
||||||
- stops when a final answer is produced or runtime limits are hit.
|
- stops when a final answer is produced or runtime limits are hit.
|
||||||
|
|
||||||
|
MCP connections are application-owned infrastructure. Composition roots create
|
||||||
|
an `MCPProvider`, share its `ToolRegistry` with `AgentLoop`, await `connect()`
|
||||||
|
before use, and guarantee `aclose()` during shutdown; the loop does not manage
|
||||||
|
that lifecycle. `AgentLoop.from_config()` therefore requires a caller-owned
|
||||||
|
`ToolRegistry`; callers using MCP share it with their application-owned
|
||||||
|
`MCPProvider`.
|
||||||
|
|
||||||
Keep this split in mind when debugging. If a problem is about channel routing, session keys, workspace selection, or outbound delivery, start in `agent/loop.py`. If it is about provider calls, tool calls, streaming, or iteration limits, start in `agent/runner.py`.
|
Keep this split in mind when debugging. If a problem is about channel routing, session keys, workspace selection, or outbound delivery, start in `agent/loop.py`. If it is about provider calls, tool calls, streaming, or iteration limits, start in `agent/runner.py`.
|
||||||
|
|
||||||
## Providers
|
## Providers
|
||||||
@@ -142,7 +149,7 @@ Defaults:
|
|||||||
|---|---|
|
|---|---|
|
||||||
| Config | `~/.nanobot/config.json` |
|
| Config | `~/.nanobot/config.json` |
|
||||||
| Workspace | `~/.nanobot/workspace/` |
|
| Workspace | `~/.nanobot/workspace/` |
|
||||||
| Sessions | `<workspace>/sessions/*.jsonl` |
|
| Sessions | `<config-dir>/sessions/<workspace-id>/*.jsonl` (default: `~/.nanobot/sessions/...`) |
|
||||||
| Memory | `<workspace>/memory/` |
|
| Memory | `<workspace>/memory/` |
|
||||||
| Cron store | `<workspace>/cron/jobs.json` |
|
| Cron store | `<workspace>/cron/jobs.json` |
|
||||||
| WebUI/media/log runtime data | config directory subdirectories such as `webui/`, `media/`, and `logs/` |
|
| WebUI/media/log runtime data | config directory subdirectories such as `webui/`, `media/`, and `logs/` |
|
||||||
@@ -157,7 +164,7 @@ a WebUI chat may select a separate project:
|
|||||||
|
|
||||||
| Concern | Path owner |
|
| Concern | Path owner |
|
||||||
|---|---|
|
|---|---|
|
||||||
| Sessions, `SOUL.md`, `USER.md`, memory, and custom skills | Configured agent workspace |
|
| Session namespace, `SOUL.md`, `USER.md`, memory, and custom skills | Configured agent workspace |
|
||||||
| Project `AGENTS.md`, relative tool paths, and shell working directory | Effective project workspace |
|
| Project `AGENTS.md`, relative tool paths, and shell working directory | Effective project workspace |
|
||||||
| Workspace access mode and project metadata | Session workspace scope |
|
| Workspace access mode and project metadata | Session workspace scope |
|
||||||
|
|
||||||
@@ -173,7 +180,7 @@ Session history is the near-term conversation replay. Memory is the longer-term
|
|||||||
|
|
||||||
| Store | File area |
|
| Store | File area |
|
||||||
|---|---|
|
|---|---|
|
||||||
| Session JSONL files | `<workspace>/sessions/` |
|
| Session JSONL files | `<config-dir>/sessions/<workspace-id>/` |
|
||||||
| Long-term memory | `<workspace>/memory/MEMORY.md` |
|
| Long-term memory | `<workspace>/memory/MEMORY.md` |
|
||||||
| Consolidation source history | `<workspace>/memory/history.jsonl` |
|
| Consolidation source history | `<workspace>/memory/history.jsonl` |
|
||||||
| Bootstrap identity files | `<workspace>/SOUL.md`, `<workspace>/USER.md`, templates under `nanobot/templates/` |
|
| Bootstrap identity files | `<workspace>/SOUL.md`, `<workspace>/USER.md`, templates under `nanobot/templates/` |
|
||||||
@@ -201,8 +208,10 @@ When changing tools, channels, file access, WebUI workspace behavior, or network
|
|||||||
| Provider | Add `ProviderSpec` in `providers/registry.py`, add schema field in `config/schema.py`, implement provider only if the generic backend is not enough |
|
| Provider | Add `ProviderSpec` in `providers/registry.py`, add schema field in `config/schema.py`, implement provider only if the generic backend is not enough |
|
||||||
| Channel | Export a `ChannelPlugin` descriptor, keep its runtime and optional setup surfaces in one package, and follow [`channel-package-guide.md`](./channel-package-guide.md) |
|
| Channel | Export a `ChannelPlugin` descriptor, keep its runtime and optional setup surfaces in one package, and follow [`channel-package-guide.md`](./channel-package-guide.md) |
|
||||||
| Tool | Implement a tool under `agent/tools/` or expose a plugin entry point |
|
| Tool | Implement a tool under `agent/tools/` or expose a plugin entry point |
|
||||||
| MCP | Add `tools.mcpServers` config |
|
| Agent Plugin | Add a v1 package under `<workspace>/plugins/` and enable it from Apps |
|
||||||
| Skill | Add workspace skill files under `<workspace>/skills/` or built-in skills under `nanobot/skills/` |
|
| MCP | Add `tools.mcpServers` config or bundle the server in an Agent Plugin |
|
||||||
|
| Skill | Add workspace skills under `<workspace>/skills/`, bundle them in an Agent Plugin, or add built-in skills under `nanobot/skills/` |
|
||||||
|
| CLI App | Add it to the CLI Apps catalog; the installer owns its executable lifecycle and writes a skills-only Agent Plugin |
|
||||||
|
|
||||||
Prefer existing registry/discovery patterns over ad hoc wiring.
|
Prefer existing registry/discovery patterns over ad hoc wiring.
|
||||||
|
|
||||||
|
|||||||
@@ -94,6 +94,24 @@ follow the printed WebUI **Settings → Models** or `nanobot onboard --wizard` r
|
|||||||
| `nanobot agent --no-markdown` | Print plain text instead of Rich-rendered Markdown |
|
| `nanobot agent --no-markdown` | Print plain text instead of Rich-rendered Markdown |
|
||||||
| `nanobot agent --logs` | Show runtime logs while chatting |
|
| `nanobot agent --logs` | Show runtime logs while chatting |
|
||||||
|
|
||||||
|
## Session Storage and Rollback
|
||||||
|
|
||||||
|
Session JSONL files live under `<config-dir>/sessions/<workspace-id>/`, outside the
|
||||||
|
agent-readable workspace. On the first upgraded start, nanobot safely migrates existing
|
||||||
|
`<workspace>/sessions/*.jsonl` files after verifying an atomic copy. Stop every old nanobot
|
||||||
|
process that uses the workspace before upgrading; old and new binaries must not write the
|
||||||
|
same session concurrently.
|
||||||
|
|
||||||
|
To prepare a downgrade, stop nanobot and copy the current sessions back to the path understood
|
||||||
|
by older releases:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
nanobot sessions restore-workspace --config ./bot-a/config.json --workspace ./bot-a/workspace
|
||||||
|
```
|
||||||
|
|
||||||
|
The command never deletes the external store and refuses to overwrite a different existing
|
||||||
|
workspace file. Back up both the config directory and workspace before changing versions.
|
||||||
|
|
||||||
In interactive mode, `Enter` sends the current message. Press `Alt+Enter` to add a newline before sending.
|
In interactive mode, `Enter` sends the current message. Press `Alt+Enter` to add a newline before sending.
|
||||||
|
|
||||||
Interactive mode exits with `exit`, `quit`, `/exit`, `/quit`, `:q`, or `Ctrl+D`.
|
Interactive mode exits with `exit`, `quit`, `/exit`, `/quit`, `:q`, or `Ctrl+D`.
|
||||||
|
|||||||
+29
-2
@@ -26,7 +26,8 @@ The default instance lives under `~/.nanobot/`:
|
|||||||
| Path | Meaning |
|
| Path | Meaning |
|
||||||
|---|---|
|
|---|---|
|
||||||
| `~/.nanobot/config.json` | Instance configuration: providers, model defaults, channels, tools, gateway, API, and runtime options |
|
| `~/.nanobot/config.json` | Instance configuration: providers, model defaults, channels, tools, gateway, API, and runtime options |
|
||||||
| `~/.nanobot/workspace/` | Agent workspace: memory, sessions, heartbeat tasks, cron jobs, skills, and generated artifacts |
|
| `~/.nanobot/workspace/` | Agent workspace: memory, heartbeat tasks, cron jobs, skills, and generated artifacts |
|
||||||
|
| `~/.nanobot/sessions/<workspace-id>/` | Session history stored outside the agent-accessible workspace; the opaque ID follows workspace moves |
|
||||||
|
|
||||||
You can override both with command flags:
|
You can override both with command flags:
|
||||||
|
|
||||||
@@ -125,13 +126,39 @@ nanobot uses two related stores:
|
|||||||
|
|
||||||
| Store | Location | Purpose |
|
| Store | Location | Purpose |
|
||||||
|---|---|---|
|
|---|---|---|
|
||||||
| Sessions | `<workspace>/sessions/*.jsonl` | Recent conversation turns replayed into context |
|
| Sessions | `<config-dir>/sessions/<workspace-id>/*.jsonl` | Recent conversation turns replayed into context |
|
||||||
| Memory | `<workspace>/memory/MEMORY.md` and `<workspace>/memory/history.jsonl` | Long-term facts and consolidated history |
|
| Memory | `<workspace>/memory/MEMORY.md` and `<workspace>/memory/history.jsonl` | Long-term facts and consolidated history |
|
||||||
|
|
||||||
Dream is a periodic consolidation job. It reads accumulated history and updates workspace memory so useful context can survive beyond short session replay.
|
Dream is a periodic consolidation job. It reads accumulated history and updates workspace memory so useful context can survive beyond short session replay.
|
||||||
|
|
||||||
|
The configured workspace contains a `.nanobot/workspace-id` file. It contains only an
|
||||||
|
opaque random identifier—never conversation content or credentials. Keep it with workspace
|
||||||
|
backups: it lets nanobot find the same external session namespace after the workspace is
|
||||||
|
renamed, moved, or restored. A live copy opened alongside the original receives a new ID so
|
||||||
|
the two workspaces do not share conversations accidentally.
|
||||||
|
|
||||||
See [`memory.md`](./memory.md) for the detailed design.
|
See [`memory.md`](./memory.md) for the detailed design.
|
||||||
|
|
||||||
|
## Apps and Agent Plugins
|
||||||
|
|
||||||
|
Agent Plugins are nanobot's common package and activation boundary for
|
||||||
|
installable capabilities. They organize existing extension types instead of
|
||||||
|
replacing them:
|
||||||
|
|
||||||
|
| Part | Role |
|
||||||
|
|---|---|
|
||||||
|
| Agent Plugin | Installable package that can bundle skills, MCP servers, or both |
|
||||||
|
| Skill | Workflow guidance loaded progressively or invoked with `$skill-name` |
|
||||||
|
| MCP server | Runtime tools exposed to the agent |
|
||||||
|
| CLI App | Locally managed executable whose adapter is packaged and activated like a plugin |
|
||||||
|
| Apps | WebUI surface for reviewing and managing these capabilities |
|
||||||
|
|
||||||
|
Native providers, channels, built-in tools, standalone workspace skills, and
|
||||||
|
directly configured MCP servers keep their existing extension paths. See
|
||||||
|
[`webui.md#apps`](./webui.md#apps) for the user-facing flow and
|
||||||
|
[`configuration.md#agent-plugins-v1`](./configuration.md#agent-plugins-v1) for
|
||||||
|
the package contract.
|
||||||
|
|
||||||
## Tools and Safety
|
## Tools and Safety
|
||||||
|
|
||||||
Tools are discovered automatically from built-in modules and plugin entry points. Common tool groups include:
|
Tools are discovered automatically from built-in modules and plugin entry points. Common tool groups include:
|
||||||
|
|||||||
+74
-11
@@ -330,7 +330,11 @@ By default, OpenAI uses `apiType: "auto"`: nanobot calls Chat Completions normal
|
|||||||
|
|
||||||
Valid `apiType` values are exactly `auto`, `chat_completions`, and `responses`.
|
Valid `apiType` values are exactly `auto`, `chat_completions`, and `responses`.
|
||||||
|
|
||||||
`extraBody` follows the selected OpenAI API surface. With Chat Completions, nanobot passes it through as the SDK `extra_body` value. With Responses, configure it in Responses API body shape; nanobot merges ordinary top-level fields into the Responses request body, appends `extraBody.tools` after generated function tools, and merges `extraBody.include` without duplicates:
|
`extraBody` follows the selected OpenAI API surface. With Chat Completions, nanobot passes
|
||||||
|
ordinary fields through as the SDK `extra_body` value; list-valued `extraBody.tools` is handled
|
||||||
|
specially and appended after generated function tools. With Responses, configure it in Responses
|
||||||
|
API body shape; nanobot merges ordinary top-level fields into the Responses request body, appends
|
||||||
|
`extraBody.tools` after generated function tools, and merges `extraBody.include` without duplicates:
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
@@ -356,7 +360,7 @@ request, while other tools such as `web_fetch` remain available.
|
|||||||
<details>
|
<details>
|
||||||
<summary><b>DeepSeek native web search</b></summary>
|
<summary><b>DeepSeek native web search</b></summary>
|
||||||
|
|
||||||
DeepSeek V4 Flash uses DeepSeek's native Responses API. Its provider-hosted web search is
|
DeepSeek V4 Flash and Pro use DeepSeek's native Responses API. Their provider-hosted web search is
|
||||||
enabled by default because it does not require a separate paid add-on. Turn it off from the
|
enabled by default because it does not require a separate paid add-on. Turn it off from the
|
||||||
WebUI provider settings, or with:
|
WebUI provider settings, or with:
|
||||||
|
|
||||||
@@ -373,9 +377,9 @@ WebUI provider settings, or with:
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
The switch applies to `deepseek-v4-flash`; DeepSeek models that remain on Chat Completions
|
The switch applies to `deepseek-v4-flash` and `deepseek-v4-pro`; DeepSeek models that remain on
|
||||||
cannot use this Responses tool. Native search calls appear in the WebUI activity stream, and
|
Chat Completions cannot use this Responses tool. Native search calls appear in the WebUI activity
|
||||||
their opaque output items are preserved for multi-turn Responses state replay.
|
stream, and their opaque output items are preserved for multi-turn Responses state replay.
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
@@ -387,7 +391,7 @@ Providers that use the Responses API can keep reasoning context across a
|
|||||||
conversation, which helps with multi-step tasks. Supported providers can also
|
conversation, which helps with multi-step tasks. Supported providers can also
|
||||||
compact long conversations automatically.
|
compact long conversations automatically.
|
||||||
|
|
||||||
nanobot preserves Responses conversation state automatically for OpenAI Responses, OpenAI Codex, Azure OpenAI, DeepSeek V4 Flash, and compatible GitHub Copilot models.
|
nanobot preserves Responses conversation state automatically for OpenAI Responses, OpenAI Codex, Azure OpenAI, DeepSeek V4, and compatible GitHub Copilot models.
|
||||||
Native compaction is also automatic when the provider supports it. The
|
Native compaction is also automatic when the provider supports it. The
|
||||||
threshold is derived from the active model's context window and reserved output
|
threshold is derived from the active model's context window and reserved output
|
||||||
headroom; no provider configuration is required.
|
headroom; no provider configuration is required.
|
||||||
@@ -1917,6 +1921,14 @@ Create a key at [serper.dev](https://serper.dev). You can also set `SERPER_API_K
|
|||||||
|
|
||||||
nanobot by default uses [Jina Reader](https://jina.ai/reader/), a third-party API, to convert arbitrary pages into Markdown format for easy digestion by the LLM, with a local fallback based on [readability-lxml](https://github.com/buriy/python-readability) if the former fails.
|
nanobot by default uses [Jina Reader](https://jina.ai/reader/), a third-party API, to convert arbitrary pages into Markdown format for easy digestion by the LLM, with a local fallback based on [readability-lxml](https://github.com/buriy/python-readability) if the former fails.
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> Using the remote reader means the fetched URL itself is disclosed to the
|
||||||
|
> third-party service. URLs that visibly carry credentials (userinfo, signed-URL
|
||||||
|
> or token-style query parameters) are detected and fetched locally instead, but
|
||||||
|
> secrets embedded in a URL's *path* (for example bot-token or webhook-style
|
||||||
|
> URLs) cannot be reliably detected. Set `useJinaReader: false` if fetched URLs
|
||||||
|
> must never leave the machine.
|
||||||
|
|
||||||
If you want to always use the local conversion, you can force it using:
|
If you want to always use the local conversion, you can force it using:
|
||||||
|
|
||||||
```json
|
```json
|
||||||
@@ -1971,15 +1983,52 @@ Add MCP servers to your `config.json`:
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
Two transport modes are supported:
|
MCP servers can run locally over stdio or connect remotely over HTTP:
|
||||||
|
|
||||||
| Mode | Config | Example |
|
| Connection | Config | Example |
|
||||||
|------|--------|---------|
|
|------|--------|---------|
|
||||||
| **Stdio** | `command` + `args` | Local process via `npx` / `uvx` |
|
| **Stdio** | `command` + `args` | Local process via `npx` / `uvx` |
|
||||||
| **HTTP** | `url` + `headers` (optional) | Remote endpoint (`https://mcp.example.com/sse`) |
|
| **Streamable HTTP / SSE** | `url` + `headers` (optional) | Remote endpoint (`https://mcp.example.com/mcp`) |
|
||||||
|
|
||||||
|
Remote HTTP servers may use browser OAuth instead of static headers. In the
|
||||||
|
WebUI, open **Apps → MCP → Add MCP server**, choose **Custom**, select HTTP or
|
||||||
|
SSE, and choose **OAuth** under **Authentication**. Save the server, then choose
|
||||||
|
**Connect**. For manual configuration, add `auth: "oauth"` and open
|
||||||
|
**Apps → MCP** to connect. Known presets such as Xmind, Notion, and Linear add
|
||||||
|
the config automatically on first click.
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"mcpServers": {
|
||||||
|
"notion": {
|
||||||
|
"type": "streamableHttp",
|
||||||
|
"url": "https://mcp.notion.com/mcp",
|
||||||
|
"auth": "oauth"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
nanobot opens the server's authorization page and handles the callback through
|
||||||
|
the gateway. The tools become available immediately when hot reload succeeds;
|
||||||
|
otherwise the WebUI asks for a restart. OAuth tokens and dynamic client
|
||||||
|
registration data are stored in the nanobot data directory under
|
||||||
|
`auth/mcp.json`; they are not written to `config.json`. Removing the MCP server
|
||||||
|
from Apps also removes its saved OAuth credentials. Normal gateway startup never
|
||||||
|
opens a browser or registers a new OAuth client when credentials are
|
||||||
|
missing—interactive authorization starts only after a user clicks **Connect**.
|
||||||
|
|
||||||
|
For a remotely accessed WebUI, HTTPS is recommended. Configure
|
||||||
|
`channels.websocket.publicWsUrl` with the browser-facing `wss://` endpoint so
|
||||||
|
nanobot can register the matching HTTPS callback and finish automatically. A
|
||||||
|
loopback WebUI may use HTTP. When a remote WebUI is served over plain HTTP,
|
||||||
|
nanobot instead registers a localhost callback and asks you to paste the complete
|
||||||
|
callback URL from the browser address bar after authorization.
|
||||||
|
|
||||||
> [!IMPORTANT]
|
> [!IMPORTANT]
|
||||||
> HTTP/SSE MCP URLs are validated before probing or connecting, and every outgoing MCP HTTP request is validated again before redirects are followed. `localhost`, `127.0.0.1`, RFC1918/private IPs, CGNAT/Tailscale ranges, link-local addresses, and cloud metadata endpoints are blocked by default. This can break previously working local or private HTTP MCP configs until the endpoint is explicitly allowed with `tools.ssrfWhitelist`, preferably with a single-host CIDR such as `127.0.0.1/32`, `::1/128`, or `192.168.1.50/32`. Stdio MCP servers are not affected.
|
> HTTP/SSE MCP URLs are validated before probing or connecting, and every outgoing MCP HTTP request—including OAuth metadata, client registration, token exchange, and redirects—is validated again. `localhost`, `127.0.0.1`, RFC1918/private IPs, CGNAT/Tailscale ranges, link-local addresses, and cloud metadata endpoints are blocked by default. This can break previously working local or private HTTP MCP configs until the endpoint is explicitly allowed with `tools.ssrfWhitelist`, preferably with a single-host CIDR such as `127.0.0.1/32`, `::1/128`, or `192.168.1.50/32`. Stdio MCP servers are not affected.
|
||||||
|
|
||||||
Use `toolTimeout` to override the default 30s per-call timeout for slow servers:
|
Use `toolTimeout` to override the default 30s per-call timeout for slow servers:
|
||||||
|
|
||||||
@@ -2054,7 +2103,7 @@ For API keys, tokens, and other secrets, see [Environment Variables for Secrets]
|
|||||||
| `tools.ssrfWhitelist` | `[]` | CIDR ranges exempted from the shared SSRF guard used by web fetches and HTTP/SSE MCP connections. Prefer exact host CIDRs such as `192.168.1.50/32`; broad ranges increase SSRF exposure. |
|
| `tools.ssrfWhitelist` | `[]` | CIDR ranges exempted from the shared SSRF guard used by web fetches and HTTP/SSE MCP connections. Prefer exact host CIDRs such as `192.168.1.50/32`; broad ranges increase SSRF exposure. |
|
||||||
| `channels.*.allowFrom` | omitted | Access control per channel. Omit to use pairing-only mode; set `["*"]` to allow everyone; or list specific user IDs. See [Pairing](#pairing) for details. |
|
| `channels.*.allowFrom` | omitted | Access control per channel. Omit to use pairing-only mode; set `["*"]` to allow everyone; or list specific user IDs. See [Pairing](#pairing) for details. |
|
||||||
|
|
||||||
**Docker security**: The official Docker image runs as a non-root user (`nanobot`, UID 1000) with bubblewrap pre-installed. The default `docker-compose.yml` drops all Linux capabilities and keeps Docker's default AppArmor/seccomp profiles enabled. If you enable `"tools.exec.sandbox": "bwrap"` inside Docker, start Compose with `docker-compose.bwrap.yml` as an additional override so bubblewrap can create nested namespaces.
|
**Docker security**: The official Docker image runs as a non-root user (`nanobot`, UID 1000) with bubblewrap pre-installed. The default `docker-compose.yml` drops all Linux capabilities except the `CHOWN`, `SETGID`, and `SETUID` capabilities required by the root entrypoint to initialize bind-mount ownership and become UID 1000. It enables `no-new-privileges` so the final non-root process cannot regain those bootstrap capabilities, and keeps Docker's default AppArmor/seccomp profiles enabled. If you enable `"tools.exec.sandbox": "bwrap"` inside Docker, start Compose with `docker-compose.bwrap.yml` as an additional override so bubblewrap can create nested namespaces. The host must also allow unprivileged user namespaces; the override cannot bypass a host-level namespace restriction.
|
||||||
|
|
||||||
|
|
||||||
## Pairing
|
## Pairing
|
||||||
@@ -2306,6 +2355,20 @@ Disabled skills are excluded from the main agent's skill summary, from always-on
|
|||||||
|--------|---------|-------------|
|
|--------|---------|-------------|
|
||||||
| `agents.defaults.disabledSkills` | `[]` | List of skill directory names to exclude from loading. Applies to both built-in skills and workspace skills. |
|
| `agents.defaults.disabledSkills` | `[]` | List of skill directory names to exclude from loading. Applies to both built-in skills and workspace skills. |
|
||||||
|
|
||||||
|
### Agent Plugins v1
|
||||||
|
|
||||||
|
nanobot discovers [Agent Plugins](https://agent-plugins.org/) under `<workspace>/plugins/`; a v1 package has `plugin.json` and may add `mcp.json`, `skills/<name>/SKILL.md`, or both. Agent Plugins are the common package and activation boundary for installable capabilities; they do not replace native providers, channels, tools, standalone workspace skills, or directly configured MCP servers.
|
||||||
|
|
||||||
|
Directory presence means installed; activation is explicit in **Apps**. Skills use progressive loading and `$skill-name` invocation, with workspace > plugin > built-in precedence.
|
||||||
|
Enabled `stdio` servers receive contained `PLUGIN_ROOT` and isolated `PLUGIN_DATA` paths; explicit
|
||||||
|
`tools.mcpServers` entries win collisions. Invalid or escaping components are ignored.
|
||||||
|
An enabled package is treated as immutable: changing any packaged file disables it until the user
|
||||||
|
reviews and enables it again. Runtime state belongs under `PLUGIN_DATA`, not the package root.
|
||||||
|
|
||||||
|
Enabled plugins run as the nanobot user; permissions are descriptive, not an OS sandbox. The optional `extensions.dev.nanobot.logo` accepts a contained PNG, JPEG, or WebP up to 256 KiB.
|
||||||
|
|
||||||
|
CLI Apps use the same skills-only package layout while their installer manages executables, updates, and removal. Future catalogs can place packages before using this activation path.
|
||||||
|
|
||||||
## Tool Hint Max Length
|
## Tool Hint Max Length
|
||||||
|
|
||||||
Tool hints are the short progress messages shown when the agent calls tools (e.g. `$ cd …/project && npm test`). By default, these are truncated at 40 characters, which can make long commands hard to read.
|
Tool hints are the short progress messages shown when the agent calls tools (e.g. `$ cd …/project && npm test`). By default, these are truncated at 40 characters, which can make long commands hard to read.
|
||||||
|
|||||||
+15
-6
@@ -11,7 +11,7 @@ Check these once before Render, Docker, systemd, or LaunchAgent:
|
|||||||
| `nanobot status` shows the expected config and workspace | Confirms the process will read the instance you meant to run |
|
| `nanobot status` shows the expected config and workspace | Confirms the process will read the instance you meant to run |
|
||||||
| `nanobot agent -m "Hello!"` works | Proves install, config, provider, model, and workspace writes before adding a service layer |
|
| `nanobot agent -m "Hello!"` works | Proves install, config, provider, model, and workspace writes before adding a service layer |
|
||||||
| Secrets are in environment variables or protected config files | API keys, bot tokens, OAuth state, and chat credentials should not be world-readable |
|
| Secrets are in environment variables or protected config files | API keys, bot tokens, OAuth state, and chat credentials should not be world-readable |
|
||||||
| `~/.nanobot/` or your custom config/workspace path is persistent | Sessions, memory, channel login state, generated artifacts, and cron jobs live there |
|
| The active config directory (including `sessions/`) and workspace are persistent | Sessions follow `--config`; memory, generated artifacts, and the workspace identity marker follow the workspace |
|
||||||
| Channel access control is intentional | Use `allowFrom`, pairing, WebSocket `token`/`tokenIssueSecret`, or private test channels before exposing the bot |
|
| Channel access control is intentional | Use `allowFrom`, pairing, WebSocket `token`/`tokenIssueSecret`, or private test channels before exposing the bot |
|
||||||
| Ports are planned | Gateway health defaults to local-only `127.0.0.1:18790`; WebUI/WebSocket defaults to `8765`; `nanobot serve` defaults to `8900` |
|
| Ports are planned | Gateway health defaults to local-only `127.0.0.1:18790`; WebUI/WebSocket defaults to `8765`; `nanobot serve` defaults to `8900` |
|
||||||
| Logs are easy to reach | Use `docker compose logs`, `journalctl`, LaunchAgent log files, or `nanobot gateway --verbose` while diagnosing startup |
|
| Logs are easy to reach | Use `docker compose logs`, `journalctl`, LaunchAgent log files, or `nanobot gateway --verbose` while diagnosing startup |
|
||||||
@@ -160,8 +160,11 @@ docker compose logs -f nanobot-gateway # view logs
|
|||||||
docker compose down # stop
|
docker compose down # stop
|
||||||
```
|
```
|
||||||
|
|
||||||
The default Compose file drops all Linux capabilities and keeps Docker's default
|
The default Compose file drops all Linux capabilities except `CHOWN`, `SETUID`, and
|
||||||
AppArmor/seccomp profiles enabled. If you explicitly set
|
`SETGID`, which the root entrypoint needs to fix bind-mount ownership and become UID
|
||||||
|
1000. It also enables `no-new-privileges`, so the non-root process cannot regain those
|
||||||
|
bootstrap capabilities through setuid binaries or file capabilities. Docker's default
|
||||||
|
AppArmor/seccomp profiles remain enabled. If you explicitly set
|
||||||
`"tools.exec.sandbox": "bwrap"` in `~/.nanobot/config.json`, add the bwrap
|
`"tools.exec.sandbox": "bwrap"` in `~/.nanobot/config.json`, add the bwrap
|
||||||
override file when starting containers:
|
override file when starting containers:
|
||||||
|
|
||||||
@@ -170,8 +173,10 @@ docker compose -f docker-compose.yml -f docker-compose.bwrap.yml up -d nanobot-g
|
|||||||
docker compose -f docker-compose.yml -f docker-compose.bwrap.yml run --rm nanobot-cli agent -m "Hello!"
|
docker compose -f docker-compose.yml -f docker-compose.bwrap.yml run --rm nanobot-cli agent -m "Hello!"
|
||||||
```
|
```
|
||||||
|
|
||||||
The override grants `CAP_SYS_ADMIN` and disables AppArmor/seccomp confinement for
|
The override adds `CAP_SYS_ADMIN` and disables AppArmor/seccomp confinement for the
|
||||||
the container so bubblewrap can create its nested namespaces. Use it only when the
|
container so bubblewrap can create its nested namespaces. It preserves
|
||||||
|
`no-new-privileges`. The host must also allow unprivileged user namespaces; the
|
||||||
|
override cannot bypass a host-level namespace restriction. Use it only when the
|
||||||
bwrap sandbox is enabled.
|
bwrap sandbox is enabled.
|
||||||
|
|
||||||
### Docker
|
### Docker
|
||||||
@@ -197,6 +202,8 @@ vim ~/.nanobot/config.json
|
|||||||
# health endpoint on 18790.
|
# health endpoint on 18790.
|
||||||
docker run \
|
docker run \
|
||||||
--cap-drop ALL \
|
--cap-drop ALL \
|
||||||
|
--cap-add CHOWN --cap-add SETGID --cap-add SETUID \
|
||||||
|
--security-opt no-new-privileges:true \
|
||||||
-v ~/.nanobot:/home/nanobot/.nanobot \
|
-v ~/.nanobot:/home/nanobot/.nanobot \
|
||||||
-p 18790:18790 -p 8765:8765 \
|
-p 18790:18790 -p 8765:8765 \
|
||||||
nanobot gateway
|
nanobot gateway
|
||||||
@@ -205,7 +212,9 @@ docker run \
|
|||||||
# bubblewrap needs for nested namespaces. Without them, `bwrap` may exit with
|
# bubblewrap needs for nested namespaces. Without them, `bwrap` may exit with
|
||||||
# `clone3: Operation not permitted`.
|
# `clone3: Operation not permitted`.
|
||||||
docker run \
|
docker run \
|
||||||
--cap-drop ALL --cap-add SYS_ADMIN \
|
--cap-drop ALL \
|
||||||
|
--cap-add CHOWN --cap-add SETGID --cap-add SETUID --cap-add SYS_ADMIN \
|
||||||
|
--security-opt no-new-privileges:true \
|
||||||
--security-opt apparmor=unconfined \
|
--security-opt apparmor=unconfined \
|
||||||
--security-opt seccomp=unconfined \
|
--security-opt seccomp=unconfined \
|
||||||
-v ~/.nanobot:/home/nanobot/.nanobot \
|
-v ~/.nanobot:/home/nanobot/.nanobot \
|
||||||
|
|||||||
@@ -30,10 +30,15 @@ remote HTTP endpoint.
|
|||||||
For local interactive setup:
|
For local interactive setup:
|
||||||
|
|
||||||
1. Run `nanobot webui` and open **Apps**.
|
1. Run `nanobot webui` and open **Apps**.
|
||||||
2. Choose a known integration preset, or add a custom stdio, HTTP, or SSE server.
|
2. Choose a known MCP server preset, or add a custom stdio, HTTP, or SSE server.
|
||||||
|
For a custom OAuth server, choose **OAuth** under **Authentication**, save it,
|
||||||
|
and click **Connect**. Presets such as Xmind, Notion, and Linear go straight to
|
||||||
|
**Connect**. Approve access in the browser window. HTTPS and localhost WebUIs
|
||||||
|
return automatically. From a remote plain-HTTP WebUI, copy the complete
|
||||||
|
localhost callback URL from the browser address bar and paste it into nanobot.
|
||||||
3. Limit the enabled tools when the server exposes more than the task needs.
|
3. Limit the enabled tools when the server exposes more than the task needs.
|
||||||
4. Save and restart when prompted.
|
4. Save and restart when prompted.
|
||||||
5. Mention the integration with `@` in the next message and ask for a small test action.
|
5. Mention the connected MCP server with `@` in the next message and ask for a small test action.
|
||||||
|
|
||||||
For manual or deployment-managed config, add this to `~/.nanobot/config.json`:
|
For manual or deployment-managed config, add this to `~/.nanobot/config.json`:
|
||||||
|
|
||||||
@@ -58,12 +63,16 @@ Restart nanobot and ask a question that requires the MCP tool.
|
|||||||
- Prefer `enabledTools` over exposing every tool by default.
|
- Prefer `enabledTools` over exposing every tool by default.
|
||||||
- Use `toolTimeout` for slow MCP operations.
|
- Use `toolTimeout` for slow MCP operations.
|
||||||
- Use HTTP MCP only for endpoints you trust.
|
- Use HTTP MCP only for endpoints you trust.
|
||||||
|
- For deployment-managed OAuth servers, set `auth` to `oauth` and complete the
|
||||||
|
browser connection from **Apps → MCP**.
|
||||||
- Keep MCP server commands stable and versioned in deployment docs or scripts.
|
- Keep MCP server commands stable and versioned in deployment docs or scripts.
|
||||||
|
|
||||||
## Security notes
|
## Security notes
|
||||||
|
|
||||||
- Stdio MCP starts a local process; review the command before enabling it.
|
- Stdio MCP starts a local process; review the command before enabling it.
|
||||||
- HTTP/SSE MCP uses nanobot's SSRF guard.
|
- HTTP/SSE MCP uses nanobot's SSRF guard, including OAuth discovery, registration,
|
||||||
|
token exchange, and redirects.
|
||||||
|
- OAuth credentials live in the nanobot data directory, not in `config.json`.
|
||||||
- Allow private HTTP MCP hosts only with narrow `tools.ssrfWhitelist` CIDRs.
|
- Allow private HTTP MCP hosts only with narrow `tools.ssrfWhitelist` CIDRs.
|
||||||
- Do not place secrets in command arguments when environment variables or
|
- Do not place secrets in command arguments when environment variables or
|
||||||
headers can be used.
|
headers can be used.
|
||||||
|
|||||||
@@ -81,6 +81,10 @@ in the WebUI or logs.
|
|||||||
- Web fetch and HTTP MCP share an SSRF guard.
|
- Web fetch and HTTP MCP share an SSRF guard.
|
||||||
- Private, loopback, link-local, and cloud metadata addresses are blocked by
|
- Private, loopback, link-local, and cloud metadata addresses are blocked by
|
||||||
default.
|
default.
|
||||||
|
- With `useJinaReader` enabled (the default), fetched URLs are disclosed to the
|
||||||
|
remote reader service. Credential-bearing URLs (userinfo or token/signature
|
||||||
|
query parameters) are fetched locally instead; path-embedded secrets cannot
|
||||||
|
be detected, so disable the remote reader when URLs must stay local.
|
||||||
- Add `tools.ssrfWhitelist` only for narrow trusted CIDRs.
|
- Add `tools.ssrfWhitelist` only for narrow trusted CIDRs.
|
||||||
- Do not give public chat users unrestricted web and shell access without
|
- Do not give public chat users unrestricted web and shell access without
|
||||||
review.
|
review.
|
||||||
|
|||||||
@@ -47,8 +47,8 @@ nanobot gateway logs
|
|||||||
- Docker Compose is the most repeatable Linux container path.
|
- Docker Compose is the most repeatable Linux container path.
|
||||||
- systemd user services are useful for Linux user-level gateway deployments.
|
- systemd user services are useful for Linux user-level gateway deployments.
|
||||||
- macOS LaunchAgent keeps the gateway alive after login.
|
- macOS LaunchAgent keeps the gateway alive after login.
|
||||||
- Persist config, workspace, sessions, memory files, channel login state, and
|
- Persist the active config directory's `sessions/` folder together with the workspace
|
||||||
generated artifacts.
|
(including `.nanobot/workspace-id`), memory files, channel login state, and generated artifacts.
|
||||||
- Restart the gateway after editing `config.json`.
|
- Restart the gateway after editing `config.json`.
|
||||||
|
|
||||||
## Security notes
|
## Security notes
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ nanobot agent -c ~/.nanobot-telegram/config.json -w /tmp/nanobot-telegram-test
|
|||||||
|-----------|---------------|---------|
|
|-----------|---------------|---------|
|
||||||
| **Config** | `--config` path | `~/.nanobot-A/config.json` |
|
| **Config** | `--config` path | `~/.nanobot-A/config.json` |
|
||||||
| **Workspace** | `--workspace` or config | `~/.nanobot-A/workspace/` |
|
| **Workspace** | `--workspace` or config | `~/.nanobot-A/workspace/` |
|
||||||
|
| **Sessions** | config directory + workspace ID | `~/.nanobot-A/sessions/<workspace-id>/` |
|
||||||
| **Cron Jobs** | workspace directory | `~/.nanobot-A/workspace/cron/` |
|
| **Cron Jobs** | workspace directory | `~/.nanobot-A/workspace/cron/` |
|
||||||
| **Media / runtime state** | config directory | `~/.nanobot-A/media/` |
|
| **Media / runtime state** | config directory | `~/.nanobot-A/media/` |
|
||||||
|
|
||||||
@@ -126,6 +127,6 @@ nanobot gateway --config ~/.nanobot-telegram/config.json --workspace /tmp/nanobo
|
|||||||
## Notes
|
## Notes
|
||||||
|
|
||||||
- Each instance must use a different port if they run at the same time
|
- Each instance must use a different port if they run at the same time
|
||||||
- Use a different workspace per instance if you want isolated memory, sessions, and skills
|
- Session data follows the active config directory; use a different workspace per instance to isolate memory, skills, and the stable session namespace ID
|
||||||
- `--workspace` overrides the workspace defined in the config file
|
- `--workspace` overrides the workspace defined in the config file
|
||||||
- Cron jobs are stored in the active workspace; runtime media/state is derived from the config directory
|
- Cron jobs are stored in the active workspace; runtime media/state is derived from the config directory
|
||||||
|
|||||||
+24
-1
@@ -100,6 +100,29 @@ Gateway-style setup for model IDs served through OpenRouter.
|
|||||||
|
|
||||||
Use the model ID exactly as OpenRouter lists it.
|
Use the model ID exactly as OpenRouter lists it.
|
||||||
|
|
||||||
|
To opt into OpenRouter server-managed search and fetch, add:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"openrouter": {
|
||||||
|
"extraBody": {
|
||||||
|
"tools": [
|
||||||
|
{ "type": "openrouter:web_search" },
|
||||||
|
{ "type": "openrouter:web_fetch" }
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Chat Completions-compatible OpenRouter
|
||||||
|
[server tools](https://openrouter.ai/docs/guides/features/server-tools), such as those above, are
|
||||||
|
appended to nanobot's generated functions. This keeps unrelated local tools such as `write_file`
|
||||||
|
available in the same request. Responses-only server tools require an API surface that the
|
||||||
|
OpenRouter provider does not currently enable.
|
||||||
|
|
||||||
### Eden AI Gateway
|
### Eden AI Gateway
|
||||||
|
|
||||||
Eden AI exposes an OpenAI-compatible chat-completions endpoint at
|
Eden AI exposes an OpenAI-compatible chat-completions endpoint at
|
||||||
@@ -264,7 +287,7 @@ Arbitrary custom provider names are OpenAI-compatible only; they do not use the
|
|||||||
|
|
||||||
`providers.openai.apiType` may be set when you need to force a specific OpenAI API surface. Other providers reject `apiType`; leave it unset outside `providers.openai`. Replace the model with a model ID available to your OpenAI account. Direct OpenAI Responses, OpenAI Codex, Azure OpenAI Responses, and eligible GitHub Copilot models share [opaque Responses state retention](./configuration.md#responses-state-and-compaction); native compaction is enabled only where the backend supports it. The WebUI exposes provider-native switches for OpenAI web search, Codex Fast mode, DeepSeek web search, and Grok X Search. These switches write the corresponding raw provider request fields under `extraBody`.
|
`providers.openai.apiType` may be set when you need to force a specific OpenAI API surface. Other providers reject `apiType`; leave it unset outside `providers.openai`. Replace the model with a model ID available to your OpenAI account. Direct OpenAI Responses, OpenAI Codex, Azure OpenAI Responses, and eligible GitHub Copilot models share [opaque Responses state retention](./configuration.md#responses-state-and-compaction); native compaction is enabled only where the backend supports it. The WebUI exposes provider-native switches for OpenAI web search, Codex Fast mode, DeepSeek web search, and Grok X Search. These switches write the corresponding raw provider request fields under `extraBody`.
|
||||||
|
|
||||||
DeepSeek is the model-level exception in the OpenAI-compatible provider: `deepseek-v4-flash` automatically uses DeepSeek's native Responses API, while `deepseek-v4-pro` remains on Chat Completions. Its native `web_search` tool is enabled by default and shows its lifecycle in WebUI chat activity; set `providers.deepseek.extraBody.tools` to `[]` to disable it.
|
DeepSeek is the model-level exception in the OpenAI-compatible provider: `deepseek-v4-flash` and `deepseek-v4-pro` automatically use DeepSeek's native Responses API. Its native `web_search` tool is enabled by default and shows its lifecycle in WebUI chat activity; set `providers.deepseek.extraBody.tools` to `[]` to disable it.
|
||||||
|
|
||||||
### Custom OpenAI-Compatible Endpoint
|
### Custom OpenAI-Compatible Endpoint
|
||||||
|
|
||||||
|
|||||||
+2
-1
@@ -48,7 +48,8 @@ The WebUI launcher creates or updates:
|
|||||||
| Path | Purpose |
|
| Path | Purpose |
|
||||||
|---|---|
|
|---|---|
|
||||||
| `~/.nanobot/config.json` | Provider, model, WebUI, channel, tool, and runtime settings |
|
| `~/.nanobot/config.json` | Provider, model, WebUI, channel, tool, and runtime settings |
|
||||||
| `~/.nanobot/workspace/` | Sessions, memory, skills, automations, and generated files |
|
| `~/.nanobot/workspace/` | Memory, skills, automations, and generated files |
|
||||||
|
| `~/.nanobot/sessions/<workspace-id>/` | Recent session history stored outside the workspace; the ID remains stable across workspace moves |
|
||||||
|
|
||||||
If the installer did not open the browser, run:
|
If the installer did not open the browser, run:
|
||||||
|
|
||||||
|
|||||||
@@ -270,6 +270,12 @@ http://127.0.0.1:8765
|
|||||||
|
|
||||||
If accessing from another device, bind the WebSocket channel to `0.0.0.0` and set `token` or `tokenIssueSecret`. The WebSocket channel refuses public binds without a token or token issue secret.
|
If accessing from another device, bind the WebSocket channel to `0.0.0.0` and set `token` or `tokenIssueSecret`. The WebSocket channel refuses public binds without a token or token issue secret.
|
||||||
|
|
||||||
|
| Symptom | Check |
|
||||||
|
|---|---|
|
||||||
|
| A temporary chat disappeared after a reload or reconnect | This is expected. Temporary chats exist only for the current WebUI connection and are not saved to history or memory. Use a regular topic for anything you need to retain. |
|
||||||
|
| A skills.sh install says that `npx` is required | Install Node.js with `npx` on the gateway machine, or choose a SkillHub skill that does not require `npx`. |
|
||||||
|
| A remote browser says skill installation is disabled | Install from a same-machine WebUI. For a private deployment where every authenticated user is trusted to install third-party skill instructions or scripts, explicitly enable `tools.webuiAllowRemotePackageInstall`. |
|
||||||
|
|
||||||
See [`webui.md#lan-access`](./webui.md#lan-access) for LAN setup and [`../webui/README.md`](../webui/README.md) for frontend development.
|
See [`webui.md#lan-access`](./webui.md#lan-access) for LAN setup and [`../webui/README.md`](../webui/README.md) for frontend development.
|
||||||
|
|
||||||
## Chat App Problems
|
## Chat App Problems
|
||||||
@@ -313,7 +319,8 @@ See [`chat-apps.md`](./chat-apps.md) for channel-specific setup.
|
|||||||
|---|---|
|
|---|---|
|
||||||
| Conversation context seems wrong | Confirm the active workspace and session. WebUI chats and chat app threads may use different sessions. |
|
| Conversation context seems wrong | Confirm the active workspace and session. WebUI chats and chat app threads may use different sessions. |
|
||||||
| Memory does not update immediately | Dream consolidation is periodic; recent turns still live in session history. |
|
| Memory does not update immediately | Dream consolidation is periodic; recent turns still live in session history. |
|
||||||
| Old sessions appear after moving config | Session files are stored under `<workspace>/sessions/`; verify the workspace path. |
|
| Sessions disappear after changing `--config` | Sessions follow the config directory at `<config-dir>/sessions/<workspace-id>/`; use the original config path or copy that `sessions/` directory into the new config directory while nanobot is stopped. |
|
||||||
|
| Sessions disappear after moving a workspace | Keep the workspace's `.nanobot/workspace-id` file with the move or backup. If it was lost, restore that marker from backup before starting nanobot. |
|
||||||
| You want one shared session across devices | Set `agents.defaults.unifiedSession` intentionally; otherwise keep separate sessions. |
|
| You want one shared session across devices | Set `agents.defaults.unifiedSession` intentionally; otherwise keep separate sessions. |
|
||||||
|
|
||||||
## Collect Useful Evidence
|
## Collect Useful Evidence
|
||||||
|
|||||||
+77
-28
@@ -1,10 +1,10 @@
|
|||||||
# Nanobot WebUI: Browser Workbench for Self-Hosted AI Agents
|
# Nanobot WebUI: Browser Workbench for Self-Hosted AI Agents
|
||||||
|
|
||||||
<!-- Meta description: Run nanobot from a browser WebUI with persistent topics, visible tool activity, workspace controls, Apps, MCP presets, Skills, settings, and Automations. -->
|
<!-- Meta description: Run nanobot from a browser WebUI with persistent and temporary chats, visible tool activity, workspace controls, Apps, skill discovery, settings, and Automations. -->
|
||||||
|
|
||||||
The WebUI is nanobot's browser workbench for persistent topics, visible
|
The WebUI is nanobot's browser workbench for persistent topics, temporary
|
||||||
agent activity, workspace controls, Apps, Skills, settings, and Automations in
|
chats, visible agent activity, workspace controls, Apps, skill discovery,
|
||||||
one place.
|
settings, and Automations in one place.
|
||||||
|
|
||||||
The published `nanobot-ai` wheel already includes the WebUI bundle. You only need
|
The published `nanobot-ai` wheel already includes the WebUI bundle. You only need
|
||||||
the `webui/` source directory when you are changing the frontend itself.
|
the `webui/` source directory when you are changing the frontend itself.
|
||||||
@@ -72,14 +72,14 @@ This path avoids hand-editing `config.json` for normal setup. Use the reference
|
|||||||
|
|
||||||
| Area | Use it for |
|
| Area | Use it for |
|
||||||
|---|---|
|
|---|---|
|
||||||
| Topics | Start, switch, search, fork, and delete browser topics |
|
| Topics | Start persistent topics or temporary chats; switch, search, reorder, fork, or delete persistent topics |
|
||||||
| Agent activity | See thinking, tool calls, file edits with diffs, command output, and generated artifacts in context |
|
| Agent activity | See thinking, tool calls, file edits with diffs, command output, and generated artifacts in context |
|
||||||
| Workspace | Pick the project workspace before asking for file or shell work |
|
| Workspace | Pick the project workspace before asking for file or shell work |
|
||||||
| Access | Choose the access mode for local capabilities allowed by your gateway configuration |
|
| Access | Choose the access mode for local capabilities allowed by your gateway configuration |
|
||||||
| Composer | Send text, images, voice input, slash commands, and `@` mentions for topics, Apps, or MCP presets |
|
| Composer | Send text, images, voice input, slash commands, and `@` mentions for topics, Apps, or MCP presets |
|
||||||
| Channels | Connect and validate chat platforms, install their optional support, and manage saved channel setup |
|
| Channels | Connect and validate chat platforms, install their optional support, and manage saved channel setup |
|
||||||
| Apps | Install, test, update, and use local CLI App adapters and MCP presets |
|
| Apps | Install, test, update, and use local CLI App adapters and MCP presets |
|
||||||
| Skills | Inspect available built-in and workspace skills before relying on them |
|
| Skills | Inspect and manage installed skills, or discover skills from supported marketplaces |
|
||||||
| Automations | Review, search, run, pause, edit, and delete scheduled and local-trigger agent turns |
|
| Automations | Review, search, run, pause, edit, and delete scheduled and local-trigger agent turns |
|
||||||
| Settings | Adjust models, providers, image generation, voice, web tools, runtime, and safety options |
|
| Settings | Adjust models, providers, image generation, voice, web tools, runtime, and safety options |
|
||||||
|
|
||||||
@@ -90,6 +90,10 @@ workspace selection, and linked automations. Use a new topic when you want a
|
|||||||
separate context; use fork when you want to continue from an existing point
|
separate context; use fork when you want to continue from an existing point
|
||||||
without changing the original thread.
|
without changing the original thread.
|
||||||
|
|
||||||
|
Drag a topic within its current sidebar group to keep frequently used work in
|
||||||
|
your preferred order. Drag a topic from the sidebar into the composer when you
|
||||||
|
want to reference it in the next message instead of switching to it.
|
||||||
|
|
||||||
The message timeline shows both user-visible replies and agent activity. Long
|
The message timeline shows both user-visible replies and agent activity. Long
|
||||||
tool or reasoning sections can be expanded when you need the details.
|
tool or reasoning sections can be expanded when you need the details.
|
||||||
|
|
||||||
@@ -103,6 +107,28 @@ File previews follow the active session access mode. Restricted workspace access
|
|||||||
previews only files under the selected workspace. Full Access can preview files
|
previews only files under the selected workspace. Full Access can preview files
|
||||||
outside the workspace when that access mode is allowed by the gateway.
|
outside the workspace when that access mode is allowed by the gateway.
|
||||||
|
|
||||||
|
## Temporary Chats
|
||||||
|
|
||||||
|
Use a temporary chat for a conversation that should not be added to nanobot's
|
||||||
|
topic history or long-term memory:
|
||||||
|
|
||||||
|
1. Select **New topic**.
|
||||||
|
2. Select the **Temporary chat** control in the page header.
|
||||||
|
3. Send the first message.
|
||||||
|
|
||||||
|
You can keep more than one temporary chat open and switch between them under
|
||||||
|
**Temporary chats** in the sidebar while the current WebUI connection remains
|
||||||
|
open. Reloading or closing the page, restarting the gateway, or losing the
|
||||||
|
WebSocket connection ends all of them. They cannot be recovered afterward.
|
||||||
|
|
||||||
|
Temporary does not mean consequence-free. Requests still go to the configured
|
||||||
|
model provider, and tools can still change files, run commands, or affect
|
||||||
|
external services. Temporary chats always use the default workspace in
|
||||||
|
Restricted mode; the project picker and Full Access are unavailable. Commands
|
||||||
|
and tools that create durable goals, automations, or subagent work are also
|
||||||
|
unavailable. Use a regular topic when you need reusable context, scheduled work,
|
||||||
|
or a result you must retain.
|
||||||
|
|
||||||
## Workspace and Access
|
## Workspace and Access
|
||||||
|
|
||||||
Use the workspace picker before starting project-specific work. This gives the
|
Use the workspace picker before starting project-specific work. This gives the
|
||||||
@@ -145,7 +171,8 @@ clients.
|
|||||||
The composer supports plain messages, image attachments, voice input when
|
The composer supports plain messages, image attachments, voice input when
|
||||||
transcription is configured, slash commands, and `@` mentions for installed Apps
|
transcription is configured, slash commands, and `@` mentions for installed Apps
|
||||||
or MCP presets. Select another topic from the `@` menu to attach a stable
|
or MCP presets. Select another topic from the `@` menu to attach a stable
|
||||||
reference; plain text that happens to start with `@` does not attach history.
|
reference, or drag that topic from the sidebar into the composer. Plain text
|
||||||
|
that happens to start with `@` does not attach history.
|
||||||
Restricted chats offer topics from the same project, while Full Access chats can
|
Restricted chats offer topics from the same project, while Full Access chats can
|
||||||
reference any WebUI topic. Nanobot reads a referenced topic only when its history
|
reference any WebUI topic. Nanobot reads a referenced topic only when its history
|
||||||
is relevant and can link it in the response. The model badge shows the current
|
is relevant and can link it in the response. The model badge shows the current
|
||||||
@@ -171,14 +198,23 @@ Test a new channel with a private DM. When a supported channel sends a pairing c
|
|||||||
|
|
||||||
## Apps
|
## Apps
|
||||||
|
|
||||||
Open Apps from the sidebar to manage tools that nanobot can attach to a chat
|
Open Apps from the sidebar to review and manage installable capabilities. The
|
||||||
turn. The default **Ready** view shows only tools that can be used immediately:
|
default **Ready** view shows only capabilities that can be used immediately:
|
||||||
|
|
||||||
- **Apps** are local command-line adapters that nanobot runs on your machine.
|
- **Agent Plugins** are local packages that can bundle skills, MCP servers, or
|
||||||
Installing an adapter does not modify the native desktop or web app it
|
both. A package under `<workspace>/plugins/` is installed but remains inactive
|
||||||
connects to.
|
until you enable it in Apps.
|
||||||
- **Integrations** are MCP servers. Presets provide known configurations, and
|
- **CLI Apps** are local command-line adapters that nanobot runs on your
|
||||||
the custom integration panel accepts stdio, HTTP, and SSE servers.
|
machine. Their installer manages the executable and exposes its adapter
|
||||||
|
through the same plugin activation model. Installing an adapter does not
|
||||||
|
modify the native desktop or web app it connects to.
|
||||||
|
- **MCP** lists Model Context Protocol servers. Presets provide known
|
||||||
|
configurations, and the **Add MCP server** panel accepts stdio, HTTP, and SSE
|
||||||
|
servers. Custom HTTP/SSE servers can use no authentication, OAuth, or request
|
||||||
|
headers. After saving an OAuth server, choose **Connect** to open its sign-in
|
||||||
|
page. Presets such as Xmind, Notion, and Linear already use OAuth. HTTPS and
|
||||||
|
localhost WebUIs return automatically; a remote plain-HTTP WebUI shows one
|
||||||
|
field for pasting the complete localhost callback URL.
|
||||||
|
|
||||||
Apps intentionally does not list nanobot runtime support packages such as
|
Apps intentionally does not list nanobot runtime support packages such as
|
||||||
`api` or `bedrock`. Those packages enable providers, servers, or channels; they
|
`api` or `bedrock`. Those packages enable providers, servers, or channels; they
|
||||||
@@ -187,6 +223,7 @@ are not tools that can be attached to a turn with `@`. Manage them from
|
|||||||
included in nanobot and activate automatically when a file is attached. The
|
included in nanobot and activate automatically when a file is attached. The
|
||||||
equivalent CLI for optional integrations remains `nanobot plugins`. See
|
equivalent CLI for optional integrations remains `nanobot plugins`. See
|
||||||
[`cli-reference.md`](./cli-reference.md#optional-features).
|
[`cli-reference.md`](./cli-reference.md#optional-features).
|
||||||
|
That command manages nanobot runtime extras, not Agent Plugin packages.
|
||||||
|
|
||||||
Some MCP presets connect to hosted keyless endpoints. For example, the Firecrawl
|
Some MCP presets connect to hosted keyless endpoints. For example, the Firecrawl
|
||||||
preset uses Firecrawl's hosted MCP endpoint for search, scrape, crawl, and
|
preset uses Firecrawl's hosted MCP endpoint for search, scrape, crawl, and
|
||||||
@@ -199,15 +236,26 @@ endpoint and exposes `web_search` and `web_fetch` without requiring an API key.
|
|||||||
It is an optional integration and does not replace nanobot's built-in web search
|
It is an optional integration and does not replace nanobot's built-in web search
|
||||||
provider; mention `@parallel-search` when a turn should use it.
|
provider; mention `@parallel-search` when a turn should use it.
|
||||||
|
|
||||||
After an App or integration is available, mention it from the composer with
|
After a CLI App or MCP server is available, mention it from the composer with
|
||||||
`@` to attach that tool to the next message.
|
`@` to attach that tool to the next message. Plugin-provided skills participate
|
||||||
|
in normal skill discovery and can be invoked with `$skill-name`.
|
||||||
|
|
||||||
## Skills
|
## Skills
|
||||||
|
|
||||||
The Skills view shows the skill instructions available to the agent, including
|
Open **Skills → Installed** to review built-in and workspace-provided skills.
|
||||||
built-in skills and workspace-provided skills. Check this view when you want to
|
You can search and filter them, inspect their instructions and setup
|
||||||
know whether nanobot already has a focused workflow for a task before you ask it
|
requirements, enable or disable them, and delete workspace skills you no longer
|
||||||
to perform that task.
|
want.
|
||||||
|
|
||||||
|
Open **Skills → Discover** to browse or search skills from skills.sh and
|
||||||
|
SkillHub. A marketplace skill is copied into the active agent workspace after
|
||||||
|
you confirm the installation. skills.sh installation requires Node.js with
|
||||||
|
`npx`; SkillHub installation does not.
|
||||||
|
|
||||||
|
Marketplace skills are third-party instructions and may include executable
|
||||||
|
scripts. Review the source and instructions before installing one, and enable
|
||||||
|
only skills you trust with the same files, tools, and credentials available to
|
||||||
|
your agent.
|
||||||
|
|
||||||
## Automations
|
## Automations
|
||||||
|
|
||||||
@@ -295,10 +343,10 @@ trusts. Configure [`sslCertfile` and `sslKeyfile`](./websocket.md#tlsssl) on the
|
|||||||
WebSocket channel and open `https://<your-host>:8765`, or terminate HTTPS at a
|
WebSocket channel and open `https://<your-host>:8765`, or terminate HTTPS at a
|
||||||
reverse proxy and use that proxy's HTTPS URL.
|
reverse proxy and use that proxy's HTTPS URL.
|
||||||
|
|
||||||
Remote WebUI clients with a valid token can view and use Apps. Actions that
|
Remote WebUI clients with a valid token can view and use Apps and installed
|
||||||
install missing nanobot support packages, such as adding a channel dependency,
|
skills. Actions that install missing nanobot support packages or third-party
|
||||||
are blocked by default. To let trusted remote administrators change the Python
|
marketplace skills are blocked by default. To let trusted remote administrators
|
||||||
environment through the WebUI, opt in explicitly:
|
perform those installations through the WebUI, opt in explicitly:
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
@@ -309,12 +357,13 @@ environment through the WebUI, opt in explicitly:
|
|||||||
```
|
```
|
||||||
|
|
||||||
Use this only for a private deployment where every authenticated WebUI user is
|
Use this only for a private deployment where every authenticated WebUI user is
|
||||||
trusted to change the Python environment that nanobot runs in. If you publish
|
trusted to change nanobot's Python environment and install workspace skill
|
||||||
the WebUI through Nginx, Caddy, Cloudflare Tunnel, or a similar service, treat it
|
instructions or scripts. If you publish the WebUI through Nginx, Caddy,
|
||||||
as remote access and leave package installs disabled unless that is intentional.
|
Cloudflare Tunnel, or a similar service, treat it as remote access and leave
|
||||||
|
package and skill installs disabled unless that is intentional.
|
||||||
|
|
||||||
Optional feature installs use pip's configured package index, including
|
Optional feature installs use pip's configured package index, including
|
||||||
`PIP_INDEX_URL`.
|
`PIP_INDEX_URL`. skills.sh marketplace installs use `npx` instead.
|
||||||
|
|
||||||
Leave remote package installs disabled when the WebUI is exposed beyond a
|
Leave remote package installs disabled when the WebUI is exposed beyond a
|
||||||
private, trusted network.
|
private, trusted network.
|
||||||
|
|||||||
@@ -42,25 +42,11 @@ def session_extra(metadata: Mapping[str, Any] | None) -> dict[str, Any]:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
async def connect_mcp(state: Any, tools: ToolRegistry) -> None:
|
|
||||||
await mcp_tools.connect_missing_servers(state, tools)
|
|
||||||
|
|
||||||
|
|
||||||
async def close_mcp(state: Any) -> None:
|
|
||||||
await mcp_tools.close_mcp_servers(state)
|
|
||||||
|
|
||||||
|
|
||||||
async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool:
|
async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool:
|
||||||
if msg.metadata.get(INBOUND_META_RUNTIME_CONTROL) == RUNTIME_CONTROL_SESSION_DISCARD:
|
if msg.metadata.get(INBOUND_META_RUNTIME_CONTROL) == RUNTIME_CONTROL_SESSION_DISCARD:
|
||||||
await state.discard_session(msg.session_key)
|
await state.discard_session(msg.session_key)
|
||||||
return True
|
return True
|
||||||
for handler in (
|
return await image_generation_tools.handle_runtime_control(state, msg, tools)
|
||||||
image_generation_tools.handle_runtime_control,
|
|
||||||
mcp_tools.handle_runtime_control,
|
|
||||||
):
|
|
||||||
if await handler(state, msg, tools):
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
class ContextBuilder:
|
class ContextBuilder:
|
||||||
|
|||||||
+35
-29
@@ -36,6 +36,7 @@ from nanobot.agent.tools.exec_session import ExecSessionManager
|
|||||||
from nanobot.agent.tools.file_state import FileStateStore, bind_file_states, reset_file_states
|
from nanobot.agent.tools.file_state import FileStateStore, bind_file_states, reset_file_states
|
||||||
from nanobot.agent.tools.message import MessageTool
|
from nanobot.agent.tools.message import MessageTool
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
|
from nanobot.agent.tools.runtime_control import AgentRuntimeControl
|
||||||
from nanobot.agent.tools.self import MyTool
|
from nanobot.agent.tools.self import MyTool
|
||||||
from nanobot.agent.turn_delivery import (
|
from nanobot.agent.turn_delivery import (
|
||||||
TurnDelivery,
|
TurnDelivery,
|
||||||
@@ -94,11 +95,9 @@ from nanobot.utils.runtime import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nanobot.agent.tools.mcp import MCPConnection
|
|
||||||
from nanobot.config.schema import (
|
from nanobot.config.schema import (
|
||||||
ChannelsConfig,
|
ChannelsConfig,
|
||||||
Config,
|
Config,
|
||||||
MCPServerConfig,
|
|
||||||
ProviderConfig,
|
ProviderConfig,
|
||||||
ToolsConfig,
|
ToolsConfig,
|
||||||
)
|
)
|
||||||
@@ -197,6 +196,11 @@ class AgentLoop:
|
|||||||
def tool_names(self) -> list[str]:
|
def tool_names(self) -> list[str]:
|
||||||
return self.tools.tool_names
|
return self.tools.tool_names
|
||||||
|
|
||||||
|
@property
|
||||||
|
def last_usage(self) -> Mapping[str, int]:
|
||||||
|
"""Latest aggregate usage exposed through the runtime-control snapshot."""
|
||||||
|
return self._last_usage
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def provider(self) -> LLMProvider:
|
def provider(self) -> LLMProvider:
|
||||||
"""Provider selected for future turn admissions."""
|
"""Provider selected for future turn admissions."""
|
||||||
@@ -265,7 +269,7 @@ class AgentLoop:
|
|||||||
cron_service: CronService | None = None,
|
cron_service: CronService | None = None,
|
||||||
restrict_to_workspace: bool = False,
|
restrict_to_workspace: bool = False,
|
||||||
session_manager: SessionManager | None = None,
|
session_manager: SessionManager | None = None,
|
||||||
mcp_servers: dict[str, MCPServerConfig] | None = None,
|
tool_registry: ToolRegistry | None = None,
|
||||||
channels_config: ChannelsConfig | None = None,
|
channels_config: ChannelsConfig | None = None,
|
||||||
timezone: str | None = None,
|
timezone: str | None = None,
|
||||||
session_ttl_minutes: int = 0,
|
session_ttl_minutes: int = 0,
|
||||||
@@ -373,7 +377,7 @@ class AgentLoop:
|
|||||||
self.context = ContextBuilder(workspace, timezone=timezone, disabled_skills=disabled_skills)
|
self.context = ContextBuilder(workspace, timezone=timezone, disabled_skills=disabled_skills)
|
||||||
self.sessions = session_manager or SessionManager(workspace)
|
self.sessions = session_manager or SessionManager(workspace)
|
||||||
self.sessions.set_file_cap_archiver(self.context.memory.raw_archive)
|
self.sessions.set_file_cap_archiver(self.context.memory.raw_archive)
|
||||||
self.tools = ToolRegistry()
|
self.tools = tool_registry if tool_registry is not None else ToolRegistry()
|
||||||
# One file-read/write tracker per logical session. The tool registry is
|
# One file-read/write tracker per logical session. The tool registry is
|
||||||
# shared by this loop, so tools resolve the active state via contextvars.
|
# shared by this loop, so tools resolve the active state via contextvars.
|
||||||
self._file_state_store = FileStateStore()
|
self._file_state_store = FileStateStore()
|
||||||
@@ -393,14 +397,11 @@ class AgentLoop:
|
|||||||
)
|
)
|
||||||
self._unified_session = unified_session
|
self._unified_session = unified_session
|
||||||
self._running = False
|
self._running = False
|
||||||
self._mcp_servers = mcp_servers or {}
|
|
||||||
self._mcp_stacks: dict[str, MCPConnection] = {}
|
|
||||||
self._mcp_connecting = False
|
|
||||||
self._runtime_context_providers: list[RuntimeContextProvider] = []
|
self._runtime_context_providers: list[RuntimeContextProvider] = []
|
||||||
self._active_tasks: dict[str, set[asyncio.Task[Any]]] = {}
|
self._active_tasks: dict[str, set[asyncio.Task[Any]]] = {}
|
||||||
self._discarding_sessions: set[str] = set()
|
self._discarding_sessions: set[str] = set()
|
||||||
self._background_tasks: set[asyncio.Task[Any]] = set()
|
self._background_tasks: set[asyncio.Task[Any]] = set()
|
||||||
self._close_mcp_lock = asyncio.Lock()
|
self._close_lock = asyncio.Lock()
|
||||||
self._session_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
|
self._session_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
|
||||||
weakref.WeakValueDictionary()
|
weakref.WeakValueDictionary()
|
||||||
)
|
)
|
||||||
@@ -448,7 +449,6 @@ class AgentLoop:
|
|||||||
if model_preset:
|
if model_preset:
|
||||||
self.set_model_preset(model_preset, publish_update=False)
|
self.set_model_preset(model_preset, publish_update=False)
|
||||||
self._register_default_tools(provider_snapshot_loader=provider_snapshot_loader)
|
self._register_default_tools(provider_snapshot_loader=provider_snapshot_loader)
|
||||||
self._runtime_vars: dict[str, Any] = {}
|
|
||||||
self._current_iteration: int = 0
|
self._current_iteration: int = 0
|
||||||
self.commands = CommandRouter()
|
self.commands = CommandRouter()
|
||||||
register_builtin_commands(self.commands)
|
register_builtin_commands(self.commands)
|
||||||
@@ -458,10 +458,15 @@ class AgentLoop:
|
|||||||
cls,
|
cls,
|
||||||
config: Config,
|
config: Config,
|
||||||
bus: MessageBus | None = None,
|
bus: MessageBus | None = None,
|
||||||
|
*,
|
||||||
|
tool_registry: ToolRegistry,
|
||||||
**extra: Any,
|
**extra: Any,
|
||||||
) -> AgentLoop:
|
) -> AgentLoop:
|
||||||
"""Create an AgentLoop from config with the common parameter set.
|
"""Create an AgentLoop from config with the common parameter set.
|
||||||
|
|
||||||
|
The tool registry is caller-owned so application composition can share
|
||||||
|
it with infrastructure such as an ``MCPProvider``.
|
||||||
|
|
||||||
Extra keyword arguments are forwarded to ``AgentLoop.__init__``,
|
Extra keyword arguments are forwarded to ``AgentLoop.__init__``,
|
||||||
allowing callers to override or extend the standard config-derived
|
allowing callers to override or extend the standard config-derived
|
||||||
parameters (e.g. ``cron_service``, ``session_manager``).
|
parameters (e.g. ``cron_service``, ``session_manager``).
|
||||||
@@ -471,6 +476,12 @@ class AgentLoop:
|
|||||||
if bus is None:
|
if bus is None:
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
defaults = config.agents.defaults
|
defaults = config.agents.defaults
|
||||||
|
if "session_manager" not in extra:
|
||||||
|
data_dir = config.runtime_data_dir
|
||||||
|
extra["session_manager"] = SessionManager(
|
||||||
|
config.workspace_path,
|
||||||
|
sessions_root=data_dir / "sessions" if data_dir is not None else None,
|
||||||
|
)
|
||||||
provider = extra.pop("provider", None) or make_provider(config)
|
provider = extra.pop("provider", None) or make_provider(config)
|
||||||
resolved = config.resolve_preset()
|
resolved = config.resolve_preset()
|
||||||
model = extra.pop("model", None) or resolved.model
|
model = extra.pop("model", None) or resolved.model
|
||||||
@@ -494,7 +505,6 @@ class AgentLoop:
|
|||||||
provider_retry_mode=defaults.provider_retry_mode,
|
provider_retry_mode=defaults.provider_retry_mode,
|
||||||
tool_hint_max_length=defaults.tool_hint_max_length,
|
tool_hint_max_length=defaults.tool_hint_max_length,
|
||||||
restrict_to_workspace=config.tools.restrict_to_workspace,
|
restrict_to_workspace=config.tools.restrict_to_workspace,
|
||||||
mcp_servers=config.tools.mcp_servers,
|
|
||||||
channels_config=config.channels,
|
channels_config=config.channels,
|
||||||
timezone=defaults.timezone,
|
timezone=defaults.timezone,
|
||||||
unified_session=defaults.unified_session,
|
unified_session=defaults.unified_session,
|
||||||
@@ -509,6 +519,7 @@ class AgentLoop:
|
|||||||
restart_mode=config.gateway.restart_mode,
|
restart_mode=config.gateway.restart_mode,
|
||||||
provider_snapshot_loader=provider_snapshot_loader,
|
provider_snapshot_loader=provider_snapshot_loader,
|
||||||
preset_snapshot_loader=preset_snapshot_loader,
|
preset_snapshot_loader=preset_snapshot_loader,
|
||||||
|
tool_registry=tool_registry,
|
||||||
**extra,
|
**extra,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -623,19 +634,18 @@ class AgentLoop:
|
|||||||
loader = ToolLoader()
|
loader = ToolLoader()
|
||||||
registered = loader.load(ctx, self.tools)
|
registered = loader.load(ctx, self.tools)
|
||||||
|
|
||||||
# MyTool needs runtime state reference — manual registration
|
# MyTool receives only the explicit runtime-control capability.
|
||||||
if self.tools_config.my.enable:
|
if self.tools_config.my.enable:
|
||||||
self.tools.register(
|
self.tools.register(
|
||||||
MyTool(runtime_state=self, modify_allowed=self.tools_config.my.allow_set)
|
MyTool(
|
||||||
|
runtime_control=AgentRuntimeControl(self),
|
||||||
|
modify_allowed=self.tools_config.my.allow_set,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
registered.append("my")
|
registered.append("my")
|
||||||
|
|
||||||
logger.info("Registered {} tools: {}", len(registered), registered)
|
logger.info("Registered {} tools: {}", len(registered), registered)
|
||||||
|
|
||||||
async def _connect_mcp(self) -> None:
|
|
||||||
"""Connect configured MCP servers."""
|
|
||||||
await agent_context.connect_mcp(self, self.tools)
|
|
||||||
|
|
||||||
def register_runtime_context_provider(
|
def register_runtime_context_provider(
|
||||||
self,
|
self,
|
||||||
provider: RuntimeContextProvider,
|
provider: RuntimeContextProvider,
|
||||||
@@ -1147,7 +1157,6 @@ class AgentLoop:
|
|||||||
"""Run the agent loop, dispatching messages as tasks to stay responsive to /stop."""
|
"""Run the agent loop, dispatching messages as tasks to stay responsive to /stop."""
|
||||||
self._running = True
|
self._running = True
|
||||||
try:
|
try:
|
||||||
await self._connect_mcp()
|
|
||||||
logger.info("Agent loop started")
|
logger.info("Agent loop started")
|
||||||
|
|
||||||
while self._running:
|
while self._running:
|
||||||
@@ -1238,8 +1247,7 @@ class AgentLoop:
|
|||||||
active_tasks.add(task)
|
active_tasks.add(task)
|
||||||
task.add_done_callback(active_tasks.discard)
|
task.add_done_callback(active_tasks.discard)
|
||||||
finally:
|
finally:
|
||||||
# MCP stdio transports use AnyIO cancel scopes; close them from the task that opened them.
|
await self.aclose()
|
||||||
await self.close_mcp()
|
|
||||||
|
|
||||||
async def _dispatch(self, msg: InboundMessage) -> None:
|
async def _dispatch(self, msg: InboundMessage) -> None:
|
||||||
"""Process a message: per-session serial, cross-session concurrent."""
|
"""Process a message: per-session serial, cross-session concurrent."""
|
||||||
@@ -1357,24 +1365,24 @@ class AgentLoop:
|
|||||||
await delivery.idle()
|
await delivery.idle()
|
||||||
await self._publish_next_deferred_automation_turn(session_key)
|
await self._publish_next_deferred_automation_turn(session_key)
|
||||||
|
|
||||||
async def close_mcp(self) -> None:
|
async def aclose(self) -> None:
|
||||||
"""Stop active work, then close exec, subagent, and MCP resources.
|
"""Stop active work, then close resources owned by the agent loop.
|
||||||
|
|
||||||
Resource teardown must still run if cancellation interrupts task draining.
|
Resource teardown must still run if cancellation interrupts task draining.
|
||||||
Gateway shutdown deliberately bounds this coroutine, so keeping the cleanup
|
Gateway shutdown deliberately bounds this coroutine, so keeping the cleanup
|
||||||
phase in ``finally`` prevents a timed-out background task from leaving
|
phase in ``finally`` prevents a timed-out background task from leaving
|
||||||
subprocess transports alive after the event loop closes.
|
subprocess transports alive after the event loop closes.
|
||||||
"""
|
"""
|
||||||
# The agent loop closes itself from ``run()`` while gateway shutdown also
|
# The loop closes itself from ``run()`` while application shutdown also
|
||||||
# performs a guaranteed final close. Serialize those owners so they cannot
|
# performs a guaranteed final close. Serialize those owners so they cannot
|
||||||
# tear down the same subprocess transports concurrently.
|
# tear down the same resources concurrently.
|
||||||
close_lock = getattr(self, "_close_mcp_lock", None)
|
close_lock = getattr(self, "_close_lock", None)
|
||||||
if close_lock is None:
|
if close_lock is None:
|
||||||
close_lock = self._close_mcp_lock = asyncio.Lock()
|
close_lock = self._close_lock = asyncio.Lock()
|
||||||
async with close_lock:
|
async with close_lock:
|
||||||
await self._close_mcp_unlocked()
|
await self._aclose_unlocked()
|
||||||
|
|
||||||
async def _close_mcp_unlocked(self) -> None:
|
async def _aclose_unlocked(self) -> None:
|
||||||
errors: list[BaseException] = []
|
errors: list[BaseException] = []
|
||||||
active_task_groups = getattr(self, "_active_tasks", {})
|
active_task_groups = getattr(self, "_active_tasks", {})
|
||||||
active_tasks = tuple({task for tasks in active_task_groups.values() for task in tasks})
|
active_tasks = tuple({task for tasks in active_task_groups.values() for task in tasks})
|
||||||
@@ -1397,7 +1405,6 @@ class AgentLoop:
|
|||||||
cleanup_steps = (
|
cleanup_steps = (
|
||||||
self.subagents.close,
|
self.subagents.close,
|
||||||
self._exec_session_manager.close_all,
|
self._exec_session_manager.close_all,
|
||||||
lambda: agent_context.close_mcp(self),
|
|
||||||
)
|
)
|
||||||
for cleanup in cleanup_steps:
|
for cleanup in cleanup_steps:
|
||||||
try:
|
try:
|
||||||
@@ -2286,7 +2293,6 @@ class AgentLoop:
|
|||||||
"""Process an external message directly and return the outbound payload."""
|
"""Process an external message directly and return the outbound payload."""
|
||||||
if channel == "system":
|
if channel == "system":
|
||||||
raise ValueError("channel 'system' is reserved for internal messages")
|
raise ValueError("channel 'system' is reserved for internal messages")
|
||||||
await self._connect_mcp()
|
|
||||||
metadata: dict[str, Any] = {}
|
metadata: dict[str, Any] = {}
|
||||||
if not persist_user_message:
|
if not persist_user_message:
|
||||||
metadata[turn_continuation.SKIP_USER_PERSIST_META] = True
|
metadata[turn_continuation.SKIP_USER_PERSIST_META] = True
|
||||||
|
|||||||
@@ -0,0 +1,435 @@
|
|||||||
|
"""Load and activate locally installed Agent Plugin packages."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import base64
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
from dataclasses import dataclass, replace
|
||||||
|
from hashlib import sha256
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import cast
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
from nanobot.agent.skills import parse_skill_metadata, valid_skill_metadata
|
||||||
|
from nanobot.config.loader import get_config_path
|
||||||
|
from nanobot.config.schema import MCPServerConfig
|
||||||
|
|
||||||
|
AGENT_PLUGIN_SCHEMA = "https://agent-plugins.org/schemas/1.0.0/plugin.schema.json"
|
||||||
|
AGENT_PLUGIN_MCP_SCHEMA = "https://agent-plugins.org/schemas/1.0.0/mcp.schema.json"
|
||||||
|
|
||||||
|
_PLUGIN_NAME = re.compile(r"^(?!.*(?:--|\.\.))[a-z0-9](?:[a-z0-9.-]*[a-z0-9])?$")
|
||||||
|
_MCP_SERVER_FIELDS = {"type", "command", "args", "env", "cwd"}
|
||||||
|
_MAX_LOGO_BYTES = 256 * 1024
|
||||||
|
_SKILL_CACHE: dict[tuple[Path, Path], tuple[tuple[str, Path], ...]] = {}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class AgentPlugin:
|
||||||
|
"""A validated, locally installed Agent Plugins v1 package."""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
root: Path
|
||||||
|
description: str
|
||||||
|
repository: str
|
||||||
|
display_name: str
|
||||||
|
category: str
|
||||||
|
accent_color: str | None
|
||||||
|
logo: str | None
|
||||||
|
permissions: tuple[str, ...]
|
||||||
|
mcp_servers: tuple[str, ...] = ()
|
||||||
|
enabled: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
def _installed_plugins(workspace: Path) -> list[AgentPlugin]:
|
||||||
|
"""Return installed packages found under ``<workspace>/plugins/*``."""
|
||||||
|
workspace = workspace.expanduser().resolve()
|
||||||
|
root = _contained(workspace / "plugins", workspace, directory=True)
|
||||||
|
if root is None:
|
||||||
|
return []
|
||||||
|
plugins: dict[str, AgentPlugin | None] = {}
|
||||||
|
for candidate in _children(root, "Agent Plugins directory"):
|
||||||
|
plugin_root = _contained(candidate, root, directory=True)
|
||||||
|
if plugin_root is None:
|
||||||
|
continue
|
||||||
|
plugin = _load_manifest(plugin_root)
|
||||||
|
if plugin is not None:
|
||||||
|
if plugin.name in plugins:
|
||||||
|
logger.warning("Ignoring duplicate Agent Plugin identity '{}'", plugin.name)
|
||||||
|
plugins[plugin.name] = None
|
||||||
|
else:
|
||||||
|
plugins[plugin.name] = plugin
|
||||||
|
return [plugin for plugin in plugins.values() if plugin is not None]
|
||||||
|
|
||||||
|
|
||||||
|
def enabled_agent_plugin_skills(workspace: Path) -> list[tuple[str, Path]]:
|
||||||
|
"""Verify and return skills from plugins the user has explicitly enabled."""
|
||||||
|
skills = [
|
||||||
|
skill
|
||||||
|
for plugin in _installed_plugins(workspace)
|
||||||
|
if _enabled(workspace, plugin)
|
||||||
|
for skill in _discover_plugin_skills(plugin.name, plugin.root)
|
||||||
|
]
|
||||||
|
_SKILL_CACHE[_skill_cache_key(workspace)] = tuple(skills)
|
||||||
|
return skills
|
||||||
|
|
||||||
|
|
||||||
|
def enabled_agent_plugin_skill_dirs(workspace: Path) -> tuple[Path, ...]:
|
||||||
|
"""Return the last verified skill roots, verifying once on a cache miss."""
|
||||||
|
key = _skill_cache_key(workspace)
|
||||||
|
skills = _SKILL_CACHE.get(key)
|
||||||
|
if skills is None:
|
||||||
|
skills = tuple(enabled_agent_plugin_skills(workspace))
|
||||||
|
return tuple(path.parent for _name, path in skills)
|
||||||
|
|
||||||
|
|
||||||
|
def _skill_cache_key(workspace: Path) -> tuple[Path, Path]:
|
||||||
|
return (
|
||||||
|
workspace.expanduser().resolve(),
|
||||||
|
get_config_path().expanduser().resolve(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _invalidate_skill_cache(workspace: Path) -> None:
|
||||||
|
_SKILL_CACHE.pop(_skill_cache_key(workspace), None)
|
||||||
|
|
||||||
|
|
||||||
|
def _load_manifest(plugin_root: Path) -> AgentPlugin | None:
|
||||||
|
payload = _read_object(plugin_root / "plugin.json", plugin_root)
|
||||||
|
if payload is None:
|
||||||
|
return None
|
||||||
|
if payload.get("$schema") != AGENT_PLUGIN_SCHEMA:
|
||||||
|
return None
|
||||||
|
name = payload.get("name")
|
||||||
|
if (
|
||||||
|
not isinstance(name, str)
|
||||||
|
or len(name) > 64
|
||||||
|
or _PLUGIN_NAME.fullmatch(name) is None
|
||||||
|
):
|
||||||
|
logger.warning("Ignoring Agent Plugin manifest in '{}': invalid name", plugin_root)
|
||||||
|
return None
|
||||||
|
extension = payload.get("extensions")
|
||||||
|
extension_payload = cast(dict[str, object], extension) if isinstance(extension, dict) else {}
|
||||||
|
nanobot_value = extension_payload.get("dev.nanobot")
|
||||||
|
nanobot = cast(dict[str, object], nanobot_value) if isinstance(nanobot_value, dict) else {}
|
||||||
|
return AgentPlugin(
|
||||||
|
name=name,
|
||||||
|
root=plugin_root,
|
||||||
|
description=_string(payload.get("description")),
|
||||||
|
repository=_string(payload.get("repository")),
|
||||||
|
display_name=_string(nanobot.get("displayName")) or name,
|
||||||
|
category=_string(nanobot.get("category")) or "Plugin",
|
||||||
|
accent_color=_accent_color(nanobot.get("accentColor")),
|
||||||
|
logo=_plugin_logo(nanobot.get("logo"), plugin_root),
|
||||||
|
permissions=_string_tuple(nanobot.get("permissions")),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def agent_plugin_mcp_servers(
|
||||||
|
workspace: Path,
|
||||||
|
configured: dict[str, MCPServerConfig] | None = None,
|
||||||
|
) -> dict[str, MCPServerConfig]:
|
||||||
|
"""Merge explicitly enabled plugin MCP servers with user configuration.
|
||||||
|
|
||||||
|
User configuration wins on the unlikely event of a namespaced collision.
|
||||||
|
"""
|
||||||
|
servers: dict[str, MCPServerConfig] = {}
|
||||||
|
for plugin in _installed_plugins(workspace):
|
||||||
|
if not _enabled(workspace, plugin):
|
||||||
|
continue
|
||||||
|
plugin_servers = _plugin_mcp_servers(workspace, plugin)
|
||||||
|
for name, server in plugin_servers.items():
|
||||||
|
# ``--`` cannot occur in a valid plugin identity, so multi-server
|
||||||
|
# namespaces cannot collide with a single-server plugin name.
|
||||||
|
host_name = plugin.name if len(plugin_servers) == 1 else f"{plugin.name}--{name}"
|
||||||
|
servers[host_name] = server
|
||||||
|
configured = configured or {}
|
||||||
|
if collisions := servers.keys() & configured.keys():
|
||||||
|
logger.warning("Configured MCP servers override Agent Plugins: {}", ", ".join(sorted(collisions)))
|
||||||
|
return servers | configured
|
||||||
|
|
||||||
|
|
||||||
|
def discover_agent_plugins(workspace: Path) -> list[AgentPlugin]:
|
||||||
|
"""Return component and lifecycle state for discovered plugins."""
|
||||||
|
return [
|
||||||
|
replace(
|
||||||
|
plugin,
|
||||||
|
mcp_servers=tuple(sorted(_plugin_mcp_servers(workspace, plugin))),
|
||||||
|
enabled=_enabled(workspace, plugin),
|
||||||
|
)
|
||||||
|
for plugin in _installed_plugins(workspace)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def set_agent_plugin_enabled(workspace: Path, name: str, enabled: bool) -> None:
|
||||||
|
"""Enable or disable one installed plugin."""
|
||||||
|
plugin = next((item for item in _installed_plugins(workspace) if item.name == name), None)
|
||||||
|
if plugin is None:
|
||||||
|
raise ValueError(f"unknown Agent Plugin '{name}'")
|
||||||
|
data = _plugin_data_dir(workspace, plugin.name, create=True)
|
||||||
|
marker = data / "enabled"
|
||||||
|
if enabled:
|
||||||
|
activation = _activation_marker(plugin)
|
||||||
|
if activation is None:
|
||||||
|
raise RuntimeError(f"Agent Plugin '{name}' changed while it was being enabled")
|
||||||
|
marker.write_text(activation, encoding="utf-8")
|
||||||
|
marker.chmod(0o600)
|
||||||
|
else:
|
||||||
|
marker.unlink(missing_ok=True)
|
||||||
|
_invalidate_skill_cache(workspace)
|
||||||
|
|
||||||
|
|
||||||
|
def _string(value: object) -> str:
|
||||||
|
return value.strip() if isinstance(value, str) else ""
|
||||||
|
|
||||||
|
|
||||||
|
def _string_tuple(value: object) -> tuple[str, ...]:
|
||||||
|
items = cast(list[object], value) if isinstance(value, list) else []
|
||||||
|
return tuple(item.strip() for item in items if isinstance(item, str) and item.strip())
|
||||||
|
|
||||||
|
|
||||||
|
def _accent_color(value: object) -> str | None:
|
||||||
|
return value if isinstance(value, str) and re.fullmatch(r"#[0-9a-fA-F]{6}", value) else None
|
||||||
|
|
||||||
|
|
||||||
|
def _plugin_logo(value: object, plugin_root: Path) -> str | None:
|
||||||
|
"""Resolve nanobot's optional packaged logo extension."""
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
if not isinstance(value, str) or not value.startswith("./"):
|
||||||
|
logger.warning("Ignoring invalid Agent Plugin logo in '{}'", plugin_root)
|
||||||
|
return None
|
||||||
|
logo = _contained(plugin_root / value[2:], plugin_root)
|
||||||
|
try:
|
||||||
|
data = logo.read_bytes() if logo is not None else b""
|
||||||
|
suffix = logo.suffix.lower() if logo is not None else ""
|
||||||
|
if len(data) <= _MAX_LOGO_BYTES and (
|
||||||
|
suffix == ".png" and data.startswith(b"\x89PNG\r\n\x1a\n")
|
||||||
|
or suffix in {".jpg", ".jpeg"} and data.startswith(b"\xff\xd8\xff")
|
||||||
|
or suffix == ".webp" and data.startswith(b"RIFF") and data[8:12] == b"WEBP"
|
||||||
|
):
|
||||||
|
mime = "jpeg" if suffix in {".jpg", ".jpeg"} else suffix[1:]
|
||||||
|
return f"data:image/{mime};base64,{base64.b64encode(data).decode('ascii')}"
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
logger.warning("Ignoring invalid Agent Plugin logo in '{}'", plugin_root)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _plugin_mcp_servers(workspace: Path, plugin: AgentPlugin) -> dict[str, MCPServerConfig]:
|
||||||
|
payload = _read_object(plugin.root / "mcp.json", plugin.root)
|
||||||
|
if payload is None:
|
||||||
|
return {}
|
||||||
|
raw_servers = payload.get("mcpServers")
|
||||||
|
if (
|
||||||
|
payload.keys() != {"$schema", "mcpServers"}
|
||||||
|
or payload.get("$schema") != AGENT_PLUGIN_MCP_SCHEMA
|
||||||
|
or not isinstance(raw_servers, dict)
|
||||||
|
):
|
||||||
|
logger.warning("Ignoring invalid MCP component for Agent Plugin '{}'", plugin.name)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
data = _plugin_data_dir(workspace, plugin.name, create=True)
|
||||||
|
servers: dict[str, MCPServerConfig] = {}
|
||||||
|
for name, raw in cast(dict[str, object], raw_servers).items():
|
||||||
|
if not name or len(name) > 128 or any(ord(char) < 32 for char in name):
|
||||||
|
logger.warning("Ignoring invalid MCP server name in Agent Plugin '{}'", plugin.name)
|
||||||
|
continue
|
||||||
|
server = _plugin_mcp_server(raw, plugin.root, data)
|
||||||
|
if server is None:
|
||||||
|
logger.warning("Ignoring invalid MCP server '{}' in Agent Plugin '{}'", name, plugin.name)
|
||||||
|
continue
|
||||||
|
servers[name] = server
|
||||||
|
return servers
|
||||||
|
|
||||||
|
|
||||||
|
def _plugin_mcp_server(raw: object, root: Path, data: Path) -> MCPServerConfig | None:
|
||||||
|
if not isinstance(raw, dict):
|
||||||
|
return None
|
||||||
|
payload = cast(dict[str, object], raw)
|
||||||
|
if payload.keys() - _MCP_SERVER_FIELDS:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
server = MCPServerConfig.model_validate(payload)
|
||||||
|
except ValidationError:
|
||||||
|
return None
|
||||||
|
command = _stdio_command(server.command, root)
|
||||||
|
cwd = _stdio_cwd(payload.get("cwd"), root, data)
|
||||||
|
if server.type != "stdio" or command is None or cwd is None:
|
||||||
|
return None
|
||||||
|
if {"PLUGIN_ROOT", "PLUGIN_DATA"} & server.env.keys():
|
||||||
|
return None
|
||||||
|
return server.model_copy(
|
||||||
|
update={
|
||||||
|
"command": command,
|
||||||
|
"args": [_expand(item, root, data) for item in server.args],
|
||||||
|
"env": {
|
||||||
|
**{key: _expand(value, root, data) for key, value in server.env.items()},
|
||||||
|
"PYTHONDONTWRITEBYTECODE": "1",
|
||||||
|
"PLUGIN_ROOT": str(root),
|
||||||
|
"PLUGIN_DATA": str(data),
|
||||||
|
},
|
||||||
|
"cwd": str(cwd),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _stdio_command(value: object, root: Path) -> str | None:
|
||||||
|
if not isinstance(value, str) or not value:
|
||||||
|
return None
|
||||||
|
if value.startswith("./"):
|
||||||
|
executable = _contained(root / value[2:], root)
|
||||||
|
return str(executable) if executable is not None else None
|
||||||
|
if any(char.isspace() for char in value) or "/" in value or "\\" in value:
|
||||||
|
return None
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _stdio_cwd(value: object, root: Path, data: Path) -> Path | None:
|
||||||
|
if value is None:
|
||||||
|
return root
|
||||||
|
if not isinstance(value, str):
|
||||||
|
return None
|
||||||
|
if value.startswith("./"):
|
||||||
|
return _contained(root / value[2:], root, directory=True)
|
||||||
|
for placeholder, base in (("${PLUGIN_ROOT}", root), ("${PLUGIN_DATA}", data)):
|
||||||
|
if value == placeholder or value.startswith(f"{placeholder}/"):
|
||||||
|
relative = value[len(placeholder):].lstrip("/")
|
||||||
|
candidate = (base / relative).resolve()
|
||||||
|
if not candidate.is_relative_to(base):
|
||||||
|
return None
|
||||||
|
if base == data:
|
||||||
|
candidate.mkdir(parents=True, exist_ok=True)
|
||||||
|
candidate.chmod(0o700)
|
||||||
|
return candidate if candidate.is_dir() else None
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _expand(value: str, root: Path, data: Path) -> str:
|
||||||
|
return value.replace("${PLUGIN_ROOT}", str(root)).replace("${PLUGIN_DATA}", str(data))
|
||||||
|
|
||||||
|
|
||||||
|
def _plugin_data_dir(workspace: Path, name: str, *, create: bool) -> Path:
|
||||||
|
workspace_id = sha256(str(workspace.expanduser().resolve()).encode()).hexdigest()[:12]
|
||||||
|
current = get_config_path().expanduser().resolve().parent
|
||||||
|
for segment in ("plugin-data", workspace_id, name):
|
||||||
|
path = current / segment
|
||||||
|
if create:
|
||||||
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
|
try:
|
||||||
|
resolved = path.resolve(strict=create)
|
||||||
|
except OSError as exc:
|
||||||
|
raise RuntimeError("Agent Plugin data directory is unavailable") from exc
|
||||||
|
if not resolved.is_relative_to(current):
|
||||||
|
raise RuntimeError("Agent Plugin data directory escapes its parent")
|
||||||
|
if create:
|
||||||
|
resolved.chmod(0o700)
|
||||||
|
current = resolved
|
||||||
|
return current
|
||||||
|
|
||||||
|
|
||||||
|
def _enabled(workspace: Path, plugin: AgentPlugin) -> bool:
|
||||||
|
marker = _plugin_data_dir(workspace, plugin.name, create=False) / "enabled"
|
||||||
|
try:
|
||||||
|
if not marker.is_file():
|
||||||
|
return False
|
||||||
|
current = marker.read_text(encoding="utf-8")
|
||||||
|
activation = _activation_marker(plugin)
|
||||||
|
if activation is None:
|
||||||
|
marker.unlink(missing_ok=True)
|
||||||
|
_invalidate_skill_cache(workspace)
|
||||||
|
return False
|
||||||
|
if current == activation:
|
||||||
|
return True
|
||||||
|
if current == str(plugin.root):
|
||||||
|
marker.write_text(activation, encoding="utf-8")
|
||||||
|
marker.chmod(0o600)
|
||||||
|
return True
|
||||||
|
marker.unlink(missing_ok=True)
|
||||||
|
_invalidate_skill_cache(workspace)
|
||||||
|
return False
|
||||||
|
except OSError:
|
||||||
|
_invalidate_skill_cache(workspace)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _activation_marker(plugin: AgentPlugin) -> str | None:
|
||||||
|
"""Bind activation to one immutable package snapshot."""
|
||||||
|
digest = sha256()
|
||||||
|
try:
|
||||||
|
for candidate in sorted(plugin.root.rglob("*")):
|
||||||
|
relative = candidate.relative_to(plugin.root).as_posix()
|
||||||
|
digest.update(relative.encode())
|
||||||
|
if candidate.is_symlink():
|
||||||
|
digest.update(b"\0link\0")
|
||||||
|
digest.update(candidate.readlink().as_posix().encode())
|
||||||
|
elif candidate.is_file():
|
||||||
|
digest.update(b"\0file\0")
|
||||||
|
digest.update(candidate.read_bytes())
|
||||||
|
elif candidate.is_dir():
|
||||||
|
digest.update(b"\0dir\0")
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
digest.update(b"\0")
|
||||||
|
except OSError:
|
||||||
|
return None
|
||||||
|
return json.dumps(
|
||||||
|
{"fingerprint": digest.hexdigest(), "root": str(plugin.root)},
|
||||||
|
separators=(",", ":"),
|
||||||
|
sort_keys=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _discover_plugin_skills(plugin_name: str, plugin_root: Path) -> list[tuple[str, Path]]:
|
||||||
|
skills_root = _contained(plugin_root / "skills", plugin_root, directory=True)
|
||||||
|
if skills_root is None:
|
||||||
|
return []
|
||||||
|
|
||||||
|
skills: list[tuple[str, Path]] = []
|
||||||
|
for candidate in _children(skills_root, f"Agent Plugin '{plugin_name}' skills"):
|
||||||
|
skill_root = _contained(candidate, skills_root, directory=True)
|
||||||
|
if skill_root is None:
|
||||||
|
continue
|
||||||
|
skill_file = _contained(skill_root / "SKILL.md", plugin_root)
|
||||||
|
if skill_file is None:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
metadata = parse_skill_metadata(skill_file.read_text(encoding="utf-8"))
|
||||||
|
except (OSError, UnicodeError):
|
||||||
|
metadata = None
|
||||||
|
if metadata is None or not valid_skill_metadata(metadata, candidate.name):
|
||||||
|
logger.warning("Ignoring Agent Plugin '{}' skill '{}': invalid metadata", plugin_name, candidate.name)
|
||||||
|
continue
|
||||||
|
skills.append((candidate.name, skill_file))
|
||||||
|
return skills
|
||||||
|
|
||||||
|
|
||||||
|
def _children(root: Path, label: str) -> list[Path]:
|
||||||
|
try:
|
||||||
|
return sorted(root.iterdir(), key=lambda path: path.name)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.warning("Could not inspect {}: {}", label, exc)
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def _contained(path: Path, root: Path, *, directory: bool = False) -> Path | None:
|
||||||
|
try:
|
||||||
|
resolved = path.resolve(strict=True)
|
||||||
|
except OSError:
|
||||||
|
return None
|
||||||
|
expected_kind = resolved.is_dir() if directory else resolved.is_file()
|
||||||
|
return resolved if expected_kind and resolved.is_relative_to(root) else None
|
||||||
|
|
||||||
|
|
||||||
|
def _read_object(path: Path, root: Path) -> dict[str, object] | None:
|
||||||
|
contained = _contained(path, root)
|
||||||
|
if contained is None:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
value = cast(object, json.loads(contained.read_text(encoding="utf-8")))
|
||||||
|
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||||||
|
logger.warning("Ignoring invalid Agent Plugin component '{}': {}", contained, exc)
|
||||||
|
return None
|
||||||
|
return cast(dict[str, object], value) if isinstance(value, dict) else None
|
||||||
+66
-30
@@ -17,9 +17,35 @@ _STRIP_SKILL_FRONTMATTER = re.compile(
|
|||||||
r"^---\s*\r?\n(.*?)\r?\n---\s*\r?\n?",
|
r"^---\s*\r?\n(.*?)\r?\n---\s*\r?\n?",
|
||||||
re.DOTALL,
|
re.DOTALL,
|
||||||
)
|
)
|
||||||
|
_SKILL_NAME = re.compile(r"^(?!.*--)[a-z0-9](?:[a-z0-9-]*[a-z0-9])?$")
|
||||||
_SKILL_REFERENCE = re.compile(r"(?<![\w$])\$([A-Za-z0-9_-]+)")
|
_SKILL_REFERENCE = re.compile(r"(?<![\w$])\$([A-Za-z0-9_-]+)")
|
||||||
|
|
||||||
|
|
||||||
|
def parse_skill_metadata(content: str) -> dict[str, object] | None:
|
||||||
|
"""Parse a skill document's YAML frontmatter."""
|
||||||
|
if not (match := _STRIP_SKILL_FRONTMATTER.match(content)):
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
parsed = yaml.safe_load(match.group(1))
|
||||||
|
except yaml.YAMLError:
|
||||||
|
return None
|
||||||
|
if not isinstance(parsed, dict):
|
||||||
|
return None
|
||||||
|
return {str(key): value for key, value in cast(dict[object, object], parsed).items()}
|
||||||
|
|
||||||
|
|
||||||
|
def valid_skill_metadata(metadata: dict[str, object], name: str) -> bool:
|
||||||
|
"""Return whether metadata satisfies the Agent Skills identity contract."""
|
||||||
|
description = metadata.get("description")
|
||||||
|
return (
|
||||||
|
metadata.get("name") == name
|
||||||
|
and len(name) <= 64
|
||||||
|
and _SKILL_NAME.fullmatch(name) is not None
|
||||||
|
and isinstance(description, str)
|
||||||
|
and 1 <= len(description.strip()) <= 1024
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class SkillsLoader:
|
class SkillsLoader:
|
||||||
"""
|
"""
|
||||||
Loader for agent skills.
|
Loader for agent skills.
|
||||||
@@ -34,6 +60,15 @@ class SkillsLoader:
|
|||||||
self.builtin_skills = builtin_skills_dir or BUILTIN_SKILLS_DIR
|
self.builtin_skills = builtin_skills_dir or BUILTIN_SKILLS_DIR
|
||||||
self.disabled_skills = disabled_skills or set()
|
self.disabled_skills = disabled_skills or set()
|
||||||
|
|
||||||
|
def _skill_aliases(self) -> dict[str, str]:
|
||||||
|
"""Return compatibility aliases owned by installed CLI Apps."""
|
||||||
|
from nanobot.apps.cli import CliAppManager
|
||||||
|
|
||||||
|
try:
|
||||||
|
return CliAppManager(workspace=self.workspace).installed_skill_aliases()
|
||||||
|
except OSError:
|
||||||
|
return {}
|
||||||
|
|
||||||
def _skill_entries_from_dir(self, base: Path, source: str, *, skip_names: set[str] | None = None) -> list[dict[str, str]]:
|
def _skill_entries_from_dir(self, base: Path, source: str, *, skip_names: set[str] | None = None) -> list[dict[str, str]]:
|
||||||
if not base.exists():
|
if not base.exists():
|
||||||
return []
|
return []
|
||||||
@@ -60,15 +95,33 @@ class SkillsLoader:
|
|||||||
Returns:
|
Returns:
|
||||||
List of skill info dicts with 'name', 'path', 'source'.
|
List of skill info dicts with 'name', 'path', 'source'.
|
||||||
"""
|
"""
|
||||||
|
from nanobot.agent.plugins import enabled_agent_plugin_skills
|
||||||
|
|
||||||
|
plugin_skills = enabled_agent_plugin_skills(self.workspace)
|
||||||
skills = self._skill_entries_from_dir(self.workspace_skills, "workspace")
|
skills = self._skill_entries_from_dir(self.workspace_skills, "workspace")
|
||||||
workspace_names = {entry["name"] for entry in skills}
|
seen_names = {entry["name"] for entry in skills}
|
||||||
|
for name, path in plugin_skills:
|
||||||
|
if name in seen_names:
|
||||||
|
continue
|
||||||
|
skills.append(
|
||||||
|
{
|
||||||
|
"name": name,
|
||||||
|
"path": str(path),
|
||||||
|
"source": "plugin",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
seen_names.add(name)
|
||||||
if self.builtin_skills and self.builtin_skills.exists():
|
if self.builtin_skills and self.builtin_skills.exists():
|
||||||
skills.extend(
|
skills.extend(
|
||||||
self._skill_entries_from_dir(self.builtin_skills, "builtin", skip_names=workspace_names)
|
self._skill_entries_from_dir(self.builtin_skills, "builtin", skip_names=seen_names)
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.disabled_skills:
|
if self.disabled_skills:
|
||||||
skills = [s for s in skills if s["name"] not in self.disabled_skills]
|
disabled = set(self.disabled_skills)
|
||||||
|
for legacy, canonical in self._skill_aliases().items():
|
||||||
|
if legacy in disabled or canonical in disabled:
|
||||||
|
disabled.update((legacy, canonical))
|
||||||
|
skills = [s for s in skills if s["name"] not in disabled]
|
||||||
|
|
||||||
if filter_unavailable:
|
if filter_unavailable:
|
||||||
return [skill for skill in skills if self._check_requirements(self._get_skill_meta(skill["name"]))]
|
return [skill for skill in skills if self._check_requirements(self._get_skill_meta(skill["name"]))]
|
||||||
@@ -84,14 +137,11 @@ class SkillsLoader:
|
|||||||
Returns:
|
Returns:
|
||||||
Skill content or None if not found.
|
Skill content or None if not found.
|
||||||
"""
|
"""
|
||||||
roots = [self.workspace_skills]
|
skills = self.list_skills(filter_unavailable=False)
|
||||||
if self.builtin_skills:
|
available = {skill["name"] for skill in skills}
|
||||||
roots.append(self.builtin_skills)
|
resolved = name if name in available else self._skill_aliases().get(name, name)
|
||||||
for root in roots:
|
entry = next((skill for skill in skills if skill["name"] == resolved), None)
|
||||||
path = root / name / "SKILL.md"
|
return Path(entry["path"]).read_text(encoding="utf-8") if entry else None
|
||||||
if path.exists():
|
|
||||||
return path.read_text(encoding="utf-8")
|
|
||||||
return None
|
|
||||||
|
|
||||||
def load_skills_for_context(self, skill_names: list[str]) -> str:
|
def load_skills_for_context(self, skill_names: list[str]) -> str:
|
||||||
"""
|
"""
|
||||||
@@ -118,9 +168,11 @@ class SkillsLoader:
|
|||||||
entry["name"]
|
entry["name"]
|
||||||
for entry in self.list_skills(filter_unavailable=True)
|
for entry in self.list_skills(filter_unavailable=True)
|
||||||
}
|
}
|
||||||
|
aliases = self._skill_aliases()
|
||||||
invoked: list[str] = []
|
invoked: list[str] = []
|
||||||
for match in _SKILL_REFERENCE.finditer(text):
|
for match in _SKILL_REFERENCE.finditer(text):
|
||||||
name = match.group(1)
|
requested = match.group(1)
|
||||||
|
name = requested if requested in available else aliases.get(requested, requested)
|
||||||
if name in available and name not in invoked:
|
if name in available and name not in invoked:
|
||||||
invoked.append(name)
|
invoked.append(name)
|
||||||
return invoked
|
return invoked
|
||||||
@@ -145,6 +197,7 @@ class SkillsLoader:
|
|||||||
sections: list[str] = []
|
sections: list[str] = []
|
||||||
groups = (
|
groups = (
|
||||||
("Workspace skills", "workspace", self.workspace_skills),
|
("Workspace skills", "workspace", self.workspace_skills),
|
||||||
|
("Agent Plugin skills", "plugin", self.workspace / "plugins"),
|
||||||
("Built-in skills", "builtin", self.builtin_skills),
|
("Built-in skills", "builtin", self.builtin_skills),
|
||||||
)
|
)
|
||||||
for label, source, root in groups:
|
for label, source, root in groups:
|
||||||
@@ -278,21 +331,4 @@ class SkillsLoader:
|
|||||||
Returns:
|
Returns:
|
||||||
Metadata dict or None.
|
Metadata dict or None.
|
||||||
"""
|
"""
|
||||||
content = self.load_skill(name)
|
return parse_skill_metadata(self.load_skill(name) or "")
|
||||||
if not content or not content.startswith("---"):
|
|
||||||
return None
|
|
||||||
match = _STRIP_SKILL_FRONTMATTER.match(content)
|
|
||||||
if not match:
|
|
||||||
return None
|
|
||||||
try:
|
|
||||||
parsed = yaml.safe_load(match.group(1))
|
|
||||||
except yaml.YAMLError:
|
|
||||||
return None
|
|
||||||
if not isinstance(parsed, dict):
|
|
||||||
return None
|
|
||||||
# yaml.safe_load returns native types (int, bool, list, etc.);
|
|
||||||
# keep values as-is so downstream consumers get correct types.
|
|
||||||
metadata: dict[str, object] = {}
|
|
||||||
for key, value in cast(dict[object, object], parsed).items():
|
|
||||||
metadata[str(key)] = value
|
|
||||||
return metadata
|
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import json
|
|||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
import warnings
|
import warnings
|
||||||
|
from collections.abc import Mapping
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Callable, TypedDict
|
from typing import Any, Callable, TypedDict
|
||||||
@@ -157,6 +158,10 @@ class SubagentManager:
|
|||||||
self._task_statuses: dict[str, SubagentStatus] = {}
|
self._task_statuses: dict[str, SubagentStatus] = {}
|
||||||
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
|
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
|
||||||
|
|
||||||
|
def runtime_statuses(self) -> Mapping[str, SubagentStatus]:
|
||||||
|
"""Return the observable task statuses used by runtime-control snapshots."""
|
||||||
|
return self._task_statuses
|
||||||
|
|
||||||
def set_provider(self, provider: LLMProvider, model: str) -> None:
|
def set_provider(self, provider: LLMProvider, model: str) -> None:
|
||||||
"""Update the deprecated runtime source used by legacy ``spawn`` calls."""
|
"""Update the deprecated runtime source used by legacy ``spawn`` calls."""
|
||||||
warnings.warn(
|
warnings.warn(
|
||||||
|
|||||||
@@ -0,0 +1,218 @@
|
|||||||
|
"""Windows Job Object ownership for subprocess trees."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import ctypes
|
||||||
|
from ctypes import wintypes
|
||||||
|
|
||||||
|
_CREATE_SUSPENDED = 0x00000004
|
||||||
|
_PROCESS_SET_QUOTA = 0x0100
|
||||||
|
_PROCESS_TERMINATE = 0x0001
|
||||||
|
_TH32CS_SNAPTHREAD = 0x00000004
|
||||||
|
_THREAD_SUSPEND_RESUME = 0x0002
|
||||||
|
_JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE = 0x00002000
|
||||||
|
_JOB_OBJECT_EXTENDED_LIMIT_INFORMATION = 9
|
||||||
|
_INVALID_HANDLE_VALUE = ctypes.c_void_p(-1).value
|
||||||
|
|
||||||
|
|
||||||
|
class _IoCounters(ctypes.Structure):
|
||||||
|
_fields_ = [
|
||||||
|
("ReadOperationCount", ctypes.c_ulonglong),
|
||||||
|
("WriteOperationCount", ctypes.c_ulonglong),
|
||||||
|
("OtherOperationCount", ctypes.c_ulonglong),
|
||||||
|
("ReadTransferCount", ctypes.c_ulonglong),
|
||||||
|
("WriteTransferCount", ctypes.c_ulonglong),
|
||||||
|
("OtherTransferCount", ctypes.c_ulonglong),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class _BasicLimitInformation(ctypes.Structure):
|
||||||
|
_fields_ = [
|
||||||
|
("PerProcessUserTimeLimit", ctypes.c_longlong),
|
||||||
|
("PerJobUserTimeLimit", ctypes.c_longlong),
|
||||||
|
("LimitFlags", wintypes.DWORD),
|
||||||
|
("MinimumWorkingSetSize", ctypes.c_size_t),
|
||||||
|
("MaximumWorkingSetSize", ctypes.c_size_t),
|
||||||
|
("ActiveProcessLimit", wintypes.DWORD),
|
||||||
|
("Affinity", ctypes.c_size_t),
|
||||||
|
("PriorityClass", wintypes.DWORD),
|
||||||
|
("SchedulingClass", wintypes.DWORD),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class _ExtendedLimitInformation(ctypes.Structure):
|
||||||
|
_fields_ = [
|
||||||
|
("BasicLimitInformation", _BasicLimitInformation),
|
||||||
|
("IoInfo", _IoCounters),
|
||||||
|
("ProcessMemoryLimit", ctypes.c_size_t),
|
||||||
|
("JobMemoryLimit", ctypes.c_size_t),
|
||||||
|
("PeakProcessMemoryUsed", ctypes.c_size_t),
|
||||||
|
("PeakJobMemoryUsed", ctypes.c_size_t),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class _ThreadEntry32(ctypes.Structure):
|
||||||
|
_fields_ = [
|
||||||
|
("dwSize", wintypes.DWORD),
|
||||||
|
("cntUsage", wintypes.DWORD),
|
||||||
|
("th32ThreadID", wintypes.DWORD),
|
||||||
|
("th32OwnerProcessID", wintypes.DWORD),
|
||||||
|
("tpBasePri", wintypes.LONG),
|
||||||
|
("tpDeltaPri", wintypes.LONG),
|
||||||
|
("dwFlags", wintypes.DWORD),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
_kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
|
||||||
|
_kernel32.CreateJobObjectW.argtypes = [ctypes.c_void_p, wintypes.LPCWSTR]
|
||||||
|
_kernel32.CreateJobObjectW.restype = wintypes.HANDLE
|
||||||
|
_kernel32.SetInformationJobObject.argtypes = [
|
||||||
|
wintypes.HANDLE,
|
||||||
|
ctypes.c_int,
|
||||||
|
ctypes.c_void_p,
|
||||||
|
wintypes.DWORD,
|
||||||
|
]
|
||||||
|
_kernel32.SetInformationJobObject.restype = wintypes.BOOL
|
||||||
|
_kernel32.OpenProcess.argtypes = [wintypes.DWORD, wintypes.BOOL, wintypes.DWORD]
|
||||||
|
_kernel32.OpenProcess.restype = wintypes.HANDLE
|
||||||
|
_kernel32.AssignProcessToJobObject.argtypes = [wintypes.HANDLE, wintypes.HANDLE]
|
||||||
|
_kernel32.AssignProcessToJobObject.restype = wintypes.BOOL
|
||||||
|
_kernel32.TerminateProcess.argtypes = [wintypes.HANDLE, wintypes.UINT]
|
||||||
|
_kernel32.TerminateProcess.restype = wintypes.BOOL
|
||||||
|
_kernel32.TerminateJobObject.argtypes = [wintypes.HANDLE, wintypes.UINT]
|
||||||
|
_kernel32.TerminateJobObject.restype = wintypes.BOOL
|
||||||
|
_kernel32.CreateToolhelp32Snapshot.argtypes = [wintypes.DWORD, wintypes.DWORD]
|
||||||
|
_kernel32.CreateToolhelp32Snapshot.restype = wintypes.HANDLE
|
||||||
|
_kernel32.Thread32First.argtypes = [wintypes.HANDLE, ctypes.POINTER(_ThreadEntry32)]
|
||||||
|
_kernel32.Thread32First.restype = wintypes.BOOL
|
||||||
|
_kernel32.Thread32Next.argtypes = [wintypes.HANDLE, ctypes.POINTER(_ThreadEntry32)]
|
||||||
|
_kernel32.Thread32Next.restype = wintypes.BOOL
|
||||||
|
_kernel32.OpenThread.argtypes = [wintypes.DWORD, wintypes.BOOL, wintypes.DWORD]
|
||||||
|
_kernel32.OpenThread.restype = wintypes.HANDLE
|
||||||
|
_kernel32.ResumeThread.argtypes = [wintypes.HANDLE]
|
||||||
|
_kernel32.ResumeThread.restype = wintypes.DWORD
|
||||||
|
_kernel32.CloseHandle.argtypes = [wintypes.HANDLE]
|
||||||
|
_kernel32.CloseHandle.restype = wintypes.BOOL
|
||||||
|
|
||||||
|
|
||||||
|
def _win_error(operation: str) -> OSError:
|
||||||
|
code = ctypes.get_last_error()
|
||||||
|
return OSError(code, f"{operation} failed (Windows error {code})")
|
||||||
|
|
||||||
|
|
||||||
|
def _close_handle(handle: int | None) -> None:
|
||||||
|
if handle:
|
||||||
|
_kernel32.CloseHandle(handle)
|
||||||
|
|
||||||
|
|
||||||
|
def _set_kill_on_close(handle: int, enabled: bool) -> None:
|
||||||
|
info = _ExtendedLimitInformation()
|
||||||
|
if enabled:
|
||||||
|
info.BasicLimitInformation.LimitFlags = _JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE
|
||||||
|
if not _kernel32.SetInformationJobObject(
|
||||||
|
handle,
|
||||||
|
_JOB_OBJECT_EXTENDED_LIMIT_INFORMATION,
|
||||||
|
ctypes.byref(info),
|
||||||
|
ctypes.sizeof(info),
|
||||||
|
):
|
||||||
|
raise _win_error("SetInformationJobObject")
|
||||||
|
|
||||||
|
|
||||||
|
def _resume_primary_thread(pid: int) -> None:
|
||||||
|
snapshot = _kernel32.CreateToolhelp32Snapshot(_TH32CS_SNAPTHREAD, 0)
|
||||||
|
if snapshot == _INVALID_HANDLE_VALUE:
|
||||||
|
raise _win_error("CreateToolhelp32Snapshot")
|
||||||
|
try:
|
||||||
|
entry = _ThreadEntry32()
|
||||||
|
entry.dwSize = ctypes.sizeof(entry)
|
||||||
|
found = _kernel32.Thread32First(snapshot, ctypes.byref(entry))
|
||||||
|
while found:
|
||||||
|
if entry.th32OwnerProcessID == pid:
|
||||||
|
thread = _kernel32.OpenThread(
|
||||||
|
_THREAD_SUSPEND_RESUME,
|
||||||
|
False,
|
||||||
|
entry.th32ThreadID,
|
||||||
|
)
|
||||||
|
if not thread:
|
||||||
|
raise _win_error("OpenThread")
|
||||||
|
try:
|
||||||
|
if _kernel32.ResumeThread(thread) == 0xFFFFFFFF:
|
||||||
|
raise _win_error("ResumeThread")
|
||||||
|
return
|
||||||
|
finally:
|
||||||
|
_close_handle(thread)
|
||||||
|
found = _kernel32.Thread32Next(snapshot, ctypes.byref(entry))
|
||||||
|
raise RuntimeError(f"suspended process {pid} has no resumable thread")
|
||||||
|
finally:
|
||||||
|
_close_handle(snapshot)
|
||||||
|
|
||||||
|
|
||||||
|
class WindowsJob:
|
||||||
|
"""Own a process tree even after its root process exits."""
|
||||||
|
|
||||||
|
creation_flags = _CREATE_SUSPENDED
|
||||||
|
|
||||||
|
def __init__(self, handle: int) -> None:
|
||||||
|
self._handle: int | None = handle
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls) -> WindowsJob:
|
||||||
|
handle = _kernel32.CreateJobObjectW(None, None)
|
||||||
|
if not handle:
|
||||||
|
raise _win_error("CreateJobObjectW")
|
||||||
|
try:
|
||||||
|
_set_kill_on_close(handle, True)
|
||||||
|
except Exception:
|
||||||
|
_close_handle(handle)
|
||||||
|
raise
|
||||||
|
return cls(handle)
|
||||||
|
|
||||||
|
def assign_and_resume(self, pid: int) -> None:
|
||||||
|
"""Atomically establish tree ownership before the root can spawn."""
|
||||||
|
if self._handle is None:
|
||||||
|
raise RuntimeError("Windows job is already closed")
|
||||||
|
process = _kernel32.OpenProcess(
|
||||||
|
_PROCESS_SET_QUOTA | _PROCESS_TERMINATE,
|
||||||
|
False,
|
||||||
|
pid,
|
||||||
|
)
|
||||||
|
if not process:
|
||||||
|
error = _win_error("OpenProcess")
|
||||||
|
self.close()
|
||||||
|
raise error
|
||||||
|
|
||||||
|
if not _kernel32.AssignProcessToJobObject(self._handle, process):
|
||||||
|
error = _win_error("AssignProcessToJobObject")
|
||||||
|
_kernel32.TerminateProcess(process, 1)
|
||||||
|
_close_handle(process)
|
||||||
|
self.close()
|
||||||
|
raise error
|
||||||
|
|
||||||
|
try:
|
||||||
|
_resume_primary_thread(pid)
|
||||||
|
except Exception:
|
||||||
|
self.terminate()
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
_close_handle(process)
|
||||||
|
|
||||||
|
def release(self) -> None:
|
||||||
|
"""Release ownership after successful output collection."""
|
||||||
|
if self._handle is None:
|
||||||
|
return
|
||||||
|
_set_kill_on_close(self._handle, False)
|
||||||
|
self.close()
|
||||||
|
|
||||||
|
def terminate(self) -> None:
|
||||||
|
"""Terminate every process in the job and close its handle."""
|
||||||
|
if self._handle is None:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
_kernel32.TerminateJobObject(self._handle, 1)
|
||||||
|
finally:
|
||||||
|
self.close()
|
||||||
|
|
||||||
|
def close(self) -> None:
|
||||||
|
handle = self._handle
|
||||||
|
self._handle = None
|
||||||
|
_close_handle(handle)
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Base class for agent tools."""
|
"""Base class for agent tools."""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import math
|
||||||
import typing
|
import typing
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
@@ -67,6 +68,8 @@ class Schema(ABC):
|
|||||||
return [f"{label} should be number"]
|
return [f"{label} should be number"]
|
||||||
if t in _JSON_TYPE_MAP and t not in ("integer", "number") and not isinstance(val, _JSON_TYPE_MAP[t]):
|
if t in _JSON_TYPE_MAP and t not in ("integer", "number") and not isinstance(val, _JSON_TYPE_MAP[t]):
|
||||||
return [f"{label} should be {t}"]
|
return [f"{label} should be {t}"]
|
||||||
|
if t == "number" and isinstance(val, float) and not math.isfinite(val):
|
||||||
|
return [f"{label} must be finite"]
|
||||||
|
|
||||||
errors: list[str] = []
|
errors: list[str] = []
|
||||||
if "enum" in schema and val not in schema["enum"]:
|
if "enum" in schema and val not in schema["enum"]:
|
||||||
|
|||||||
@@ -209,7 +209,11 @@ class _ExecSession:
|
|||||||
timeout=2.0,
|
timeout=2.0,
|
||||||
)
|
)
|
||||||
# Safety-net reap after normal exit.
|
# Safety-net reap after normal exit.
|
||||||
from nanobot.agent.tools.shell import _reap_pid # pyright: ignore[reportPrivateUsage]
|
from nanobot.agent.tools.shell import ( # pyright: ignore[reportPrivateUsage]
|
||||||
|
ExecTool,
|
||||||
|
_reap_pid, # pyright: ignore[reportPrivateUsage]
|
||||||
|
)
|
||||||
|
ExecTool._release_process_tree(self.process) # pyright: ignore[reportPrivateUsage]
|
||||||
_reap_pid(self.process.pid) # pyright: ignore[reportPrivateUsage]
|
_reap_pid(self.process.pid) # pyright: ignore[reportPrivateUsage]
|
||||||
elif yield_time_ms > 0:
|
elif yield_time_ms > 0:
|
||||||
await self._wait_for_buffered_output()
|
await self._wait_for_buffered_output()
|
||||||
|
|||||||
@@ -148,9 +148,19 @@ class _FsTool(Tool):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _resolve_read(self, path: str) -> Path:
|
def _resolve_read(self, path: str) -> Path:
|
||||||
|
plugin_skill_dirs: list[Path] = []
|
||||||
|
if self._workspace is not None:
|
||||||
|
from nanobot.agent.plugins import enabled_agent_plugin_skill_dirs
|
||||||
|
|
||||||
|
try:
|
||||||
|
plugin_skill_dirs = list(
|
||||||
|
enabled_agent_plugin_skill_dirs(Path(self._workspace))
|
||||||
|
)
|
||||||
|
except (OSError, RuntimeError):
|
||||||
|
pass
|
||||||
return self._resolve_with_extra(
|
return self._resolve_with_extra(
|
||||||
path,
|
path,
|
||||||
self._extra_read_allowed_dirs,
|
[*self._extra_read_allowed_dirs, *plugin_skill_dirs],
|
||||||
self._extra_read_allowed_files,
|
self._extra_read_allowed_files,
|
||||||
include_media_dir=True,
|
include_media_dir=True,
|
||||||
extra_files_require_allowed_root=True,
|
extra_files_require_allowed_root=True,
|
||||||
@@ -827,7 +837,8 @@ class EditFileTool(_FsTool):
|
|||||||
def description(self) -> str:
|
def description(self) -> str:
|
||||||
return (
|
return (
|
||||||
"Perform a small, exact replacement in one file by replacing "
|
"Perform a small, exact replacement in one file by replacing "
|
||||||
"old_text with new_text. Use this for narrow text substitutions "
|
"old_text with new_text. When replacing text in an existing file, "
|
||||||
|
"old_text and new_text must be different. Use this for narrow text substitutions "
|
||||||
"with old_text copied from read_file. For multi-file, structural, "
|
"with old_text copied from read_file. For multi-file, structural, "
|
||||||
"or generated code edits, prefer apply_patch. If old_text matches "
|
"or generated code edits, prefer apply_patch. If old_text matches "
|
||||||
"multiple times, provide more context or set occurrence, line_hint, "
|
"multiple times, provide more context or set occurrence, line_hint, "
|
||||||
@@ -862,9 +873,12 @@ class EditFileTool(_FsTool):
|
|||||||
return ToolResult.error("Error: expected_replacements must be >= 1.")
|
return ToolResult.error("Error: expected_replacements must be >= 1.")
|
||||||
|
|
||||||
fp = self._resolve_write(path)
|
fp = self._resolve_write(path)
|
||||||
|
file_exists = fp.exists()
|
||||||
|
if file_exists and old_text == new_text:
|
||||||
|
return ToolResult.error("Error: new_text must be different from old_text.")
|
||||||
|
|
||||||
# Create-file semantics: old_text='' + file doesn't exist → create
|
# Create-file semantics: old_text='' + file doesn't exist → create
|
||||||
if not fp.exists():
|
if not file_exists:
|
||||||
if old_text == "":
|
if old_text == "":
|
||||||
fp.parent.mkdir(parents=True, exist_ok=True)
|
fp.parent.mkdir(parents=True, exist_ok=True)
|
||||||
fp.write_text(new_text, encoding="utf-8")
|
fp.write_text(new_text, encoding="utf-8")
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
_SKIP_MODULES = frozenset({
|
_SKIP_MODULES = frozenset({
|
||||||
"base", "schema", "registry", "context", "loader", "config",
|
"base", "schema", "registry", "context", "loader", "config",
|
||||||
"file_state", "sandbox", "mcp", "__init__", "runtime_state",
|
"file_state", "sandbox", "mcp", "__init__", "runtime_control",
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+374
-240
@@ -1,4 +1,6 @@
|
|||||||
"""MCP client: connects to MCP servers and wraps their tools as native nanobot tools."""
|
"""MCP client and dynamic tool-provider lifecycle."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import hashlib
|
import hashlib
|
||||||
@@ -7,23 +9,15 @@ import os
|
|||||||
import re
|
import re
|
||||||
import shutil
|
import shutil
|
||||||
import urllib.parse
|
import urllib.parse
|
||||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
from collections.abc import AsyncIterator, Awaitable, Callable, Iterable, Mapping
|
||||||
from contextlib import AsyncExitStack, suppress
|
from contextlib import AsyncExitStack, suppress
|
||||||
from typing import TYPE_CHECKING, Any, Mapping, Protocol, cast
|
from typing import TYPE_CHECKING, Any, Literal, Protocol, cast
|
||||||
from weakref import WeakKeyDictionary
|
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.agent.tools.base import Tool, ToolResult
|
from nanobot.agent.tools.base import Tool, ToolResult
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.bus.events import (
|
|
||||||
INBOUND_META_RUNTIME_CONTROL,
|
|
||||||
RUNTIME_CONTROL_ACK,
|
|
||||||
RUNTIME_CONTROL_MCP_RELOAD,
|
|
||||||
InboundMessage,
|
|
||||||
)
|
|
||||||
from nanobot.bus.queue import MessageBus
|
|
||||||
from nanobot.security.network import (
|
from nanobot.security.network import (
|
||||||
PinnedDNSAsyncTransport,
|
PinnedDNSAsyncTransport,
|
||||||
env_proxy_applies_to_url,
|
env_proxy_applies_to_url,
|
||||||
@@ -38,7 +32,8 @@ if TYPE_CHECKING:
|
|||||||
from mcp.types import Prompt, Resource
|
from mcp.types import Prompt, Resource
|
||||||
from mcp.types import Tool as MCPToolDefinition
|
from mcp.types import Tool as MCPToolDefinition
|
||||||
|
|
||||||
from nanobot.config.schema import MCPServerConfig
|
from nanobot.agent.tools.mcp_oauth import MCPOAuthHandlers
|
||||||
|
from nanobot.config.schema import Config, MCPServerConfig
|
||||||
|
|
||||||
# Transient connection errors that warrant a single retry.
|
# Transient connection errors that warrant a single retry.
|
||||||
# These typically happen when an MCP server restarts or a network
|
# These typically happen when an MCP server restarts or a network
|
||||||
@@ -59,14 +54,37 @@ _WINDOWS_SHELL_LAUNCHERS: frozenset[str] = frozenset(("npx", "npm", "pnpm", "yar
|
|||||||
# Characters allowed in tool names by model providers (Anthropic, OpenAI, etc.).
|
# Characters allowed in tool names by model providers (Anthropic, OpenAI, etc.).
|
||||||
# Replace anything outside [a-zA-Z0-9_-] with underscore and collapse runs.
|
# Replace anything outside [a-zA-Z0-9_-] with underscore and collapse runs.
|
||||||
_SANITIZE_RE = re.compile(r"_+")
|
_SANITIZE_RE = re.compile(r"_+")
|
||||||
_RELOAD_LOCKS: WeakKeyDictionary[Any, asyncio.Lock] = WeakKeyDictionary()
|
|
||||||
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]
|
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]
|
||||||
|
MCPServerLoader = Callable[[], Mapping[str, "MCPServerConfig"]]
|
||||||
|
MCPRuntimeStatus = Literal["connecting", "connected", "failed"]
|
||||||
|
|
||||||
|
|
||||||
class MCPConnection(Protocol):
|
class MCPConnection(Protocol):
|
||||||
async def aclose(self) -> None: ...
|
async def aclose(self) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
async def _close_mcp_connection(name: str, connection: MCPConnection) -> None:
|
||||||
|
try:
|
||||||
|
await connection.aclose()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
if task_is_cancelling():
|
||||||
|
raise
|
||||||
|
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
|
||||||
|
except (RuntimeError, BaseExceptionGroup):
|
||||||
|
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
|
||||||
|
|
||||||
|
|
||||||
|
async def _close_mcp_connections(connections: Mapping[str, MCPConnection]) -> None:
|
||||||
|
cancellation: asyncio.CancelledError | None = None
|
||||||
|
for name, connection in connections.items():
|
||||||
|
try:
|
||||||
|
await _close_mcp_connection(name, connection)
|
||||||
|
except asyncio.CancelledError as exc:
|
||||||
|
cancellation = cancellation or exc
|
||||||
|
if cancellation is not None:
|
||||||
|
raise cancellation
|
||||||
|
|
||||||
|
|
||||||
class _OwnedMCPConnection:
|
class _OwnedMCPConnection:
|
||||||
"""Close an MCP transport from the task that originally opened it."""
|
"""Close an MCP transport from the task that originally opened it."""
|
||||||
|
|
||||||
@@ -184,6 +202,25 @@ def _is_transient(exc: BaseException) -> bool:
|
|||||||
return type(exc).__name__ in _TRANSIENT_EXC_NAMES
|
return type(exc).__name__ in _TRANSIENT_EXC_NAMES
|
||||||
|
|
||||||
|
|
||||||
|
def _is_transient_connection_failure(exc: BaseException) -> bool:
|
||||||
|
if isinstance(exc, BaseExceptionGroup):
|
||||||
|
group = cast(BaseExceptionGroup[BaseException], exc)
|
||||||
|
return bool(group.exceptions) and all(
|
||||||
|
_is_transient_connection_failure(nested) for nested in group.exceptions
|
||||||
|
)
|
||||||
|
return isinstance(exc, (httpx.ConnectError, httpx.ConnectTimeout)) or _is_transient(exc)
|
||||||
|
|
||||||
|
|
||||||
|
def _log_mcp_connection_failure(name: str, exc: BaseException, hint: str = "") -> None:
|
||||||
|
if _is_transient_connection_failure(exc):
|
||||||
|
logger.warning("MCP server '{}': transient connection failure", name)
|
||||||
|
logger.opt(exception=exc).debug(
|
||||||
|
"MCP server '{}' transient connection failure details", name
|
||||||
|
)
|
||||||
|
return
|
||||||
|
logger.opt(exception=exc).error("MCP server '{}': failed to connect: {}", name, hint)
|
||||||
|
|
||||||
|
|
||||||
def _is_session_terminated(exc: BaseException) -> bool:
|
def _is_session_terminated(exc: BaseException) -> bool:
|
||||||
"""Return True when the MCP SDK reports a dead client session."""
|
"""Return True when the MCP SDK reports a dead client session."""
|
||||||
if _is_transient(exc):
|
if _is_transient(exc):
|
||||||
@@ -468,11 +505,11 @@ class _MCPWrapperBase(Tool):
|
|||||||
"""Common reconnect handling for wrappers bound to one MCP server session."""
|
"""Common reconnect handling for wrappers bound to one MCP server session."""
|
||||||
|
|
||||||
_plugin_discoverable = False
|
_plugin_discoverable = False
|
||||||
_session: "ClientSession"
|
_session: ClientSession
|
||||||
_server_name: str
|
_server_name: str
|
||||||
_name: str
|
_name: str
|
||||||
|
|
||||||
def _set_mcp_connection(self, session: "ClientSession", server_name: str) -> None:
|
def _set_mcp_connection(self, session: ClientSession, server_name: str) -> None:
|
||||||
self._session = session
|
self._session = session
|
||||||
self._server_name = server_name
|
self._server_name = server_name
|
||||||
self._reconnect: _ReconnectCallback | None = None
|
self._reconnect: _ReconnectCallback | None = None
|
||||||
@@ -562,9 +599,9 @@ class MCPToolWrapper(_MCPWrapperBase):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
session: "ClientSession",
|
session: ClientSession,
|
||||||
server_name: str,
|
server_name: str,
|
||||||
tool_def: "MCPToolDefinition",
|
tool_def: MCPToolDefinition,
|
||||||
tool_timeout: int = 30,
|
tool_timeout: int = 30,
|
||||||
):
|
):
|
||||||
self._set_mcp_connection(session, server_name)
|
self._set_mcp_connection(session, server_name)
|
||||||
@@ -724,9 +761,9 @@ class MCPResourceWrapper(_MCPWrapperBase):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
session: "ClientSession",
|
session: ClientSession,
|
||||||
server_name: str,
|
server_name: str,
|
||||||
resource_def: "Resource",
|
resource_def: Resource,
|
||||||
resource_timeout: int = 30,
|
resource_timeout: int = 30,
|
||||||
):
|
):
|
||||||
self._set_mcp_connection(session, server_name)
|
self._set_mcp_connection(session, server_name)
|
||||||
@@ -828,9 +865,9 @@ class MCPPromptWrapper(_MCPWrapperBase):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
session: "ClientSession",
|
session: ClientSession,
|
||||||
server_name: str,
|
server_name: str,
|
||||||
prompt_def: "Prompt",
|
prompt_def: Prompt,
|
||||||
prompt_timeout: int = 30,
|
prompt_timeout: int = 30,
|
||||||
):
|
):
|
||||||
self._set_mcp_connection(session, server_name)
|
self._set_mcp_connection(session, server_name)
|
||||||
@@ -961,7 +998,10 @@ class MCPPromptWrapper(_MCPWrapperBase):
|
|||||||
|
|
||||||
|
|
||||||
async def connect_mcp_servers(
|
async def connect_mcp_servers(
|
||||||
mcp_servers: "dict[str, MCPServerConfig]", registry: ToolRegistry
|
mcp_servers: dict[str, MCPServerConfig],
|
||||||
|
registry: ToolRegistry,
|
||||||
|
*,
|
||||||
|
oauth_handlers: Mapping[str, MCPOAuthHandlers] | None = None,
|
||||||
) -> dict[str, MCPConnection]:
|
) -> dict[str, MCPConnection]:
|
||||||
"""Connect to configured MCP servers and register their tools, resources, prompts.
|
"""Connect to configured MCP servers and register their tools, resources, prompts.
|
||||||
|
|
||||||
@@ -975,11 +1015,8 @@ async def connect_mcp_servers(
|
|||||||
from mcp.client.streamable_http import streamable_http_client
|
from mcp.client.streamable_http import streamable_http_client
|
||||||
|
|
||||||
async def open_single_server(
|
async def open_single_server(
|
||||||
name: str, cfg: "MCPServerConfig"
|
name: str, cfg: MCPServerConfig, server_stack: AsyncExitStack
|
||||||
) -> tuple[str, AsyncExitStack | None]:
|
) -> bool:
|
||||||
server_stack = AsyncExitStack()
|
|
||||||
await server_stack.__aenter__()
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
transport_type = cfg.type
|
transport_type = cfg.type
|
||||||
if not transport_type:
|
if not transport_type:
|
||||||
@@ -991,8 +1028,7 @@ async def connect_mcp_servers(
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.warning("MCP server '{}': no command or url configured, skipping", name)
|
logger.warning("MCP server '{}': no command or url configured, skipping", name)
|
||||||
await server_stack.aclose()
|
return False
|
||||||
return name, None
|
|
||||||
|
|
||||||
if transport_type in {"sse", "streamableHttp"}:
|
if transport_type in {"sse", "streamableHttp"}:
|
||||||
ok, error = validate_url_target(cfg.url)
|
ok, error = validate_url_target(cfg.url)
|
||||||
@@ -1003,8 +1039,30 @@ async def connect_mcp_servers(
|
|||||||
_redact_url(cfg.url),
|
_redact_url(cfg.url),
|
||||||
error,
|
error,
|
||||||
)
|
)
|
||||||
await server_stack.aclose()
|
return False
|
||||||
return name, None
|
|
||||||
|
oauth_auth: httpx.Auth | None = None
|
||||||
|
if cfg.auth == "oauth":
|
||||||
|
if transport_type not in {"sse", "streamableHttp"}:
|
||||||
|
logger.warning(
|
||||||
|
"MCP server '{}': OAuth requires an SSE or Streamable HTTP transport",
|
||||||
|
name,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
from nanobot.agent.tools.mcp_oauth import (
|
||||||
|
MCPAuthorizationRequiredError,
|
||||||
|
create_mcp_oauth_auth,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
oauth_auth = await create_mcp_oauth_auth(
|
||||||
|
name,
|
||||||
|
cfg.url,
|
||||||
|
(oauth_handlers or {}).get(name),
|
||||||
|
)
|
||||||
|
except MCPAuthorizationRequiredError:
|
||||||
|
logger.info("MCP server '{}': waiting for browser authorization", name)
|
||||||
|
return False
|
||||||
|
|
||||||
if transport_type == "stdio":
|
if transport_type == "stdio":
|
||||||
command, args, env = _normalize_windows_stdio_command(
|
command, args, env = _normalize_windows_stdio_command(
|
||||||
@@ -1022,8 +1080,7 @@ async def connect_mcp_servers(
|
|||||||
elif transport_type == "sse":
|
elif transport_type == "sse":
|
||||||
if not await _probe_http_url(cfg.url):
|
if not await _probe_http_url(cfg.url):
|
||||||
logger.warning("MCP server '{}': {} unreachable, skipping", name, _redact_url(cfg.url))
|
logger.warning("MCP server '{}': {} unreachable, skipping", name, _redact_url(cfg.url))
|
||||||
await server_stack.aclose()
|
return False
|
||||||
return name, None
|
|
||||||
|
|
||||||
def httpx_client_factory(
|
def httpx_client_factory(
|
||||||
headers: dict[str, str] | None = None,
|
headers: dict[str, str] | None = None,
|
||||||
@@ -1044,31 +1101,37 @@ async def connect_mcp_servers(
|
|||||||
**_pinned_transport_kwargs(),
|
**_pinned_transport_kwargs(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
sse_kwargs: dict[str, Any] = {
|
||||||
|
"httpx_client_factory": httpx_client_factory,
|
||||||
|
}
|
||||||
|
if oauth_auth is not None:
|
||||||
|
sse_kwargs["auth"] = oauth_auth
|
||||||
read, write = await server_stack.enter_async_context(
|
read, write = await server_stack.enter_async_context(
|
||||||
sse_client(cfg.url, httpx_client_factory=httpx_client_factory)
|
sse_client(cfg.url, **sse_kwargs)
|
||||||
)
|
)
|
||||||
elif transport_type == "streamableHttp":
|
elif transport_type == "streamableHttp":
|
||||||
if not await _probe_http_url(cfg.url):
|
if not await _probe_http_url(cfg.url):
|
||||||
logger.warning("MCP server '{}': {} unreachable, skipping", name, _redact_url(cfg.url))
|
logger.warning("MCP server '{}': {} unreachable, skipping", name, _redact_url(cfg.url))
|
||||||
await server_stack.aclose()
|
return False
|
||||||
return name, None
|
|
||||||
|
|
||||||
http_client = await server_stack.enter_async_context(
|
http_client_kwargs: dict[str, Any] = {
|
||||||
httpx.AsyncClient(
|
"headers": cfg.headers or None,
|
||||||
headers=cfg.headers or None,
|
"event_hooks": {"request": [_validate_mcp_request_url]},
|
||||||
event_hooks={"request": [_validate_mcp_request_url]},
|
"follow_redirects": True,
|
||||||
follow_redirects=True,
|
"timeout": httpx.Timeout(30.0, connect=10.0),
|
||||||
timeout=httpx.Timeout(30.0, connect=10.0),
|
|
||||||
**_pinned_transport_kwargs(),
|
**_pinned_transport_kwargs(),
|
||||||
)
|
}
|
||||||
|
if oauth_auth is not None:
|
||||||
|
http_client_kwargs["auth"] = oauth_auth
|
||||||
|
http_client = await server_stack.enter_async_context(
|
||||||
|
httpx.AsyncClient(**http_client_kwargs)
|
||||||
)
|
)
|
||||||
read, write, _ = await server_stack.enter_async_context(
|
read, write, _ = await server_stack.enter_async_context(
|
||||||
streamable_http_client(cfg.url, http_client=http_client)
|
streamable_http_client(cfg.url, http_client=http_client)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.warning("MCP server '{}': unknown transport type '{}'", name, transport_type)
|
logger.warning("MCP server '{}': unknown transport type '{}'", name, transport_type)
|
||||||
await server_stack.aclose()
|
return False
|
||||||
return name, None
|
|
||||||
|
|
||||||
read = _filter_malformed_mcp_progress_notifications(read, name)
|
read = _filter_malformed_mcp_progress_notifications(read, name)
|
||||||
session = await server_stack.enter_async_context(ClientSession(read, write))
|
session = await server_stack.enter_async_context(ClientSession(read, write))
|
||||||
@@ -1171,7 +1234,7 @@ async def connect_mcp_servers(
|
|||||||
logger.info(
|
logger.info(
|
||||||
"MCP server '{}': connected, {} capabilities registered", name, registered_count
|
"MCP server '{}': connected, {} capabilities registered", name, registered_count
|
||||||
)
|
)
|
||||||
return name, server_stack
|
return True
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
hint = ""
|
hint = ""
|
||||||
@@ -1190,43 +1253,41 @@ async def connect_mcp_servers(
|
|||||||
" Hint: this looks like stdio protocol pollution. Make sure the MCP server writes "
|
" Hint: this looks like stdio protocol pollution. Make sure the MCP server writes "
|
||||||
"only JSON-RPC to stdout and sends logs/debug output to stderr instead."
|
"only JSON-RPC to stdout and sends logs/debug output to stderr instead."
|
||||||
)
|
)
|
||||||
logger.exception("MCP server '{}': failed to connect: {}", name, hint)
|
_log_mcp_connection_failure(name, e, hint)
|
||||||
with suppress(Exception):
|
return False
|
||||||
await server_stack.aclose()
|
|
||||||
return name, None
|
|
||||||
|
|
||||||
async def connect_single_server(
|
async def connect_single_server(
|
||||||
name: str, cfg: "MCPServerConfig"
|
name: str, cfg: MCPServerConfig
|
||||||
) -> tuple[str, MCPConnection | None]:
|
) -> tuple[str, MCPConnection | None]:
|
||||||
loop = asyncio.get_running_loop()
|
loop = asyncio.get_running_loop()
|
||||||
ready: asyncio.Future[bool] = loop.create_future()
|
ready: asyncio.Future[bool] = loop.create_future()
|
||||||
close_requested = asyncio.Event()
|
close_requested = asyncio.Event()
|
||||||
|
|
||||||
async def own_connection() -> None:
|
async def own_connection() -> None:
|
||||||
stack: AsyncExitStack | None = None
|
|
||||||
try:
|
try:
|
||||||
_, stack = await open_single_server(name, cfg)
|
async with AsyncExitStack() as stack:
|
||||||
|
connected = await open_single_server(name, cfg, stack)
|
||||||
if not ready.done():
|
if not ready.done():
|
||||||
ready.set_result(stack is not None)
|
ready.set_result(connected)
|
||||||
if stack is not None:
|
if connected:
|
||||||
await close_requested.wait()
|
await close_requested.wait()
|
||||||
except BaseException as exc:
|
except BaseException as exc:
|
||||||
if not ready.done():
|
if not ready.done():
|
||||||
ready.set_exception(exc)
|
ready.set_exception(exc)
|
||||||
raise
|
raise
|
||||||
finally:
|
|
||||||
if stack is not None:
|
|
||||||
await stack.aclose()
|
|
||||||
|
|
||||||
owner = asyncio.create_task(own_connection(), name=f"mcp:{name}")
|
owner = asyncio.create_task(own_connection(), name=f"mcp:{name}")
|
||||||
connection = _OwnedMCPConnection(owner, close_requested)
|
connection = _OwnedMCPConnection(owner, close_requested)
|
||||||
try:
|
try:
|
||||||
connected = await ready
|
connected = await ready
|
||||||
except BaseException:
|
except BaseException as exc:
|
||||||
close_requested.set()
|
close_requested.set()
|
||||||
owner.cancel()
|
owner.cancel()
|
||||||
with suppress(BaseException):
|
with suppress(BaseException):
|
||||||
await asyncio.shield(owner)
|
await asyncio.shield(owner)
|
||||||
|
if isinstance(exc, asyncio.CancelledError) and not task_is_cancelling():
|
||||||
|
logger.warning("MCP server '{}': connection cancelled by server/SDK", name)
|
||||||
|
return name, None
|
||||||
raise
|
raise
|
||||||
if not connected:
|
if not connected:
|
||||||
await connection.aclose()
|
await connection.aclose()
|
||||||
@@ -1234,15 +1295,29 @@ async def connect_mcp_servers(
|
|||||||
return name, connection
|
return name, connection
|
||||||
|
|
||||||
server_stacks: dict[str, MCPConnection] = {}
|
server_stacks: dict[str, MCPConnection] = {}
|
||||||
|
attempted_names: list[str] = []
|
||||||
|
|
||||||
|
try:
|
||||||
for name, cfg in mcp_servers.items():
|
for name, cfg in mcp_servers.items():
|
||||||
|
attempted_names.append(name)
|
||||||
try:
|
try:
|
||||||
result = await connect_single_server(name, cfg)
|
result = await connect_single_server(name, cfg)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception("MCP server '{}' connection failed: {}", name, e)
|
_log_mcp_connection_failure(name, e)
|
||||||
continue
|
continue
|
||||||
if result[1] is not None:
|
if result[1] is not None:
|
||||||
server_stacks[result[0]] = result[1]
|
server_stacks[result[0]] = result[1]
|
||||||
|
except BaseException:
|
||||||
|
# Callers can bound readiness/reload with a timeout. If cancellation
|
||||||
|
# interrupts a later server, ownership of earlier connections has not
|
||||||
|
# transferred yet, so roll the whole batch back before propagating it.
|
||||||
|
for name in attempted_names:
|
||||||
|
_unregister_server_tools(registry, name)
|
||||||
|
try:
|
||||||
|
await _close_mcp_connections(server_stacks)
|
||||||
|
except BaseException as cleanup_exc:
|
||||||
|
logger.debug("MCP batch rollback cleanup error (can be ignored): {}", cleanup_exc)
|
||||||
|
raise
|
||||||
|
|
||||||
return server_stacks
|
return server_stacks
|
||||||
|
|
||||||
@@ -1253,53 +1328,158 @@ def session_extra(metadata: Mapping[str, Any] | None) -> dict[str, Any]:
|
|||||||
return {"mcp_presets": mcp_presets} if isinstance(mcp_presets, list) and mcp_presets else {}
|
return {"mcp_presets": mcp_presets} if isinstance(mcp_presets, list) and mcp_presets else {}
|
||||||
|
|
||||||
|
|
||||||
async def connect_missing_servers(state: Any, registry: ToolRegistry) -> None:
|
def _configured_servers(config: Config) -> dict[str, MCPServerConfig]:
|
||||||
"""Connect configured MCP servers that are not currently live."""
|
from nanobot.agent.plugins import agent_plugin_mcp_servers
|
||||||
async with _reload_lock(state):
|
|
||||||
if getattr(state, "_mcp_closing", False):
|
return agent_plugin_mcp_servers(
|
||||||
return
|
config.workspace_path,
|
||||||
missing_servers = {
|
config.tools.mcp_servers,
|
||||||
name: cfg for name, cfg in state._mcp_servers.items() if name not in state._mcp_stacks
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _load_current_servers() -> dict[str, MCPServerConfig]:
|
||||||
|
from nanobot.config.loader import load_config, resolve_config_env_vars
|
||||||
|
|
||||||
|
return _configured_servers(resolve_config_env_vars(load_config()))
|
||||||
|
|
||||||
|
|
||||||
|
class MCPProvider:
|
||||||
|
"""Own configured MCP connections and their dynamic tool registrations."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
servers: Mapping[str, MCPServerConfig],
|
||||||
|
registry: ToolRegistry,
|
||||||
|
*,
|
||||||
|
server_loader: MCPServerLoader | None = None,
|
||||||
|
) -> None:
|
||||||
|
self._servers = dict(servers)
|
||||||
|
self._registry = registry
|
||||||
|
self._server_loader = server_loader or _load_current_servers
|
||||||
|
self._connections: dict[str, MCPConnection] = {}
|
||||||
|
self._runtime_statuses: dict[str, MCPRuntimeStatus] = {}
|
||||||
|
self._lock = asyncio.Lock()
|
||||||
|
self._closing = False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_config(
|
||||||
|
cls,
|
||||||
|
config: Config,
|
||||||
|
registry: ToolRegistry,
|
||||||
|
*,
|
||||||
|
server_loader: MCPServerLoader | None = None,
|
||||||
|
) -> MCPProvider:
|
||||||
|
return cls(
|
||||||
|
_configured_servers(config),
|
||||||
|
registry,
|
||||||
|
server_loader=server_loader,
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def configured_server_names(self) -> set[str]:
|
||||||
|
return set(self._servers)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def connected_server_names(self) -> set[str]:
|
||||||
|
return set(self._connections)
|
||||||
|
|
||||||
|
def runtime_status(self) -> dict[str, MCPRuntimeStatus]:
|
||||||
|
"""Return the latest connection-attempt result for configured servers."""
|
||||||
|
return {
|
||||||
|
name: status
|
||||||
|
for name, status in self._runtime_statuses.items()
|
||||||
|
if name in self._servers
|
||||||
}
|
}
|
||||||
if state._mcp_connecting or not missing_servers:
|
|
||||||
|
def _set_runtime_status(
|
||||||
|
self,
|
||||||
|
server_names: Iterable[str],
|
||||||
|
status: MCPRuntimeStatus,
|
||||||
|
) -> None:
|
||||||
|
for name in server_names:
|
||||||
|
self._runtime_statuses[name] = status
|
||||||
|
|
||||||
|
def _record_connection_result(
|
||||||
|
self,
|
||||||
|
attempted: Iterable[str],
|
||||||
|
connected: Iterable[str],
|
||||||
|
) -> None:
|
||||||
|
attempted_names = set(attempted)
|
||||||
|
connected_names = set(connected)
|
||||||
|
self._set_runtime_status(connected_names, "connected")
|
||||||
|
self._set_runtime_status(attempted_names - connected_names, "failed")
|
||||||
|
|
||||||
|
async def connect(self) -> None:
|
||||||
|
"""Connect configured servers that are not currently live."""
|
||||||
|
async with self._lock:
|
||||||
|
if self._closing:
|
||||||
return
|
return
|
||||||
state._mcp_connecting = True
|
configured_missing = {
|
||||||
|
name: cfg
|
||||||
|
for name, cfg in self._servers.items()
|
||||||
|
if name not in self._connections
|
||||||
|
}
|
||||||
|
oauth_servers = {
|
||||||
|
name: cfg
|
||||||
|
for name, cfg in configured_missing.items()
|
||||||
|
if cfg.auth == "oauth"
|
||||||
|
}
|
||||||
|
authorization_pending: set[str] = set()
|
||||||
|
if oauth_servers:
|
||||||
|
from nanobot.agent.tools.mcp_oauth import mcp_oauth_has_credentials
|
||||||
|
|
||||||
|
authorization_pending = {
|
||||||
|
name
|
||||||
|
for name, cfg in oauth_servers.items()
|
||||||
|
if not mcp_oauth_has_credentials(name, cfg.url)
|
||||||
|
}
|
||||||
|
for name in authorization_pending:
|
||||||
|
self._runtime_statuses.pop(name, None)
|
||||||
|
missing_servers = {
|
||||||
|
name: cfg
|
||||||
|
for name, cfg in configured_missing.items()
|
||||||
|
if name not in authorization_pending
|
||||||
|
}
|
||||||
|
if not missing_servers:
|
||||||
|
return
|
||||||
|
self._set_runtime_status(missing_servers, "connecting")
|
||||||
try:
|
try:
|
||||||
connected = await connect_mcp_servers(missing_servers, registry)
|
connected = await connect_mcp_servers(missing_servers, self._registry)
|
||||||
if getattr(state, "_mcp_closing", False):
|
if self._closing:
|
||||||
for connection in connected.values():
|
await _close_mcp_connections(connected)
|
||||||
await connection.aclose()
|
|
||||||
return
|
return
|
||||||
state._mcp_stacks.update(connected)
|
self._connections.update(connected)
|
||||||
_attach_reconnect_handlers(state, registry, connected)
|
self._record_connection_result(missing_servers, connected)
|
||||||
|
self._attach_reconnect_handlers(connected)
|
||||||
if connected:
|
if connected:
|
||||||
logger.info("MCP connected servers: {}", sorted(connected))
|
logger.info("MCP connected servers: {}", sorted(connected))
|
||||||
else:
|
else:
|
||||||
logger.warning("No MCP servers connected successfully (will retry next message)")
|
logger.warning(
|
||||||
|
"No MCP servers connected successfully "
|
||||||
|
"(will retry on the next readiness check)"
|
||||||
|
)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
|
self._set_runtime_status(missing_servers, "failed")
|
||||||
if task_is_cancelling():
|
if task_is_cancelling():
|
||||||
raise
|
raise
|
||||||
logger.warning("MCP connection cancelled (will retry next message)")
|
logger.warning(
|
||||||
except BaseException as e:
|
"MCP connection cancelled (will retry on the next readiness check)"
|
||||||
logger.warning("Failed to connect MCP servers (will retry next message): {}", e)
|
)
|
||||||
finally:
|
except BaseException as exc:
|
||||||
state._mcp_connecting = False
|
self._set_runtime_status(missing_servers, "failed")
|
||||||
|
logger.warning(
|
||||||
|
"Failed to connect MCP servers "
|
||||||
|
"(will retry on the next readiness check): {}",
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def reload(self) -> dict[str, Any]:
|
||||||
async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
|
"""Reconcile live MCP connections with the current configuration."""
|
||||||
"""Reconcile live MCP connections with the current config file."""
|
async with self._lock:
|
||||||
async with _reload_lock(state):
|
if self._closing:
|
||||||
if getattr(state, "_mcp_closing", False):
|
return self._closing_result()
|
||||||
return {
|
|
||||||
"ok": False,
|
|
||||||
"message": "MCP connections are shutting down.",
|
|
||||||
"requires_restart": True,
|
|
||||||
}
|
|
||||||
try:
|
try:
|
||||||
from nanobot.config.loader import load_config, resolve_config_env_vars
|
next_servers = dict(self._server_loader())
|
||||||
|
|
||||||
config = resolve_config_env_vars(load_config())
|
|
||||||
next_servers = dict(config.tools.mcp_servers)
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.warning("MCP hot reload could not read config: {}", exc)
|
logger.warning("MCP hot reload could not read config: {}", exc)
|
||||||
return {
|
return {
|
||||||
@@ -1309,49 +1489,69 @@ async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
|
|||||||
"error": str(exc),
|
"error": str(exc),
|
||||||
}
|
}
|
||||||
|
|
||||||
current_servers = dict(state._mcp_servers)
|
current_servers = dict(self._servers)
|
||||||
current_names = set(current_servers)
|
current_names = set(current_servers)
|
||||||
next_names = set(next_servers)
|
next_names = set(next_servers)
|
||||||
|
from nanobot.agent.tools.mcp_oauth import mcp_oauth_has_credentials
|
||||||
|
|
||||||
|
authorization_pending = {
|
||||||
|
name
|
||||||
|
for name, cfg in next_servers.items()
|
||||||
|
if cfg.auth == "oauth" and not mcp_oauth_has_credentials(name, cfg.url)
|
||||||
|
}
|
||||||
removed = sorted(current_names - next_names)
|
removed = sorted(current_names - next_names)
|
||||||
added = sorted(next_names - current_names)
|
added = sorted(next_names - current_names)
|
||||||
changed = sorted(
|
changed = sorted(
|
||||||
name
|
name
|
||||||
for name in current_names & next_names
|
for name in current_names & next_names
|
||||||
if _server_signature(current_servers[name]) != _server_signature(next_servers[name])
|
if _server_signature(current_servers[name])
|
||||||
|
!= _server_signature(next_servers[name])
|
||||||
)
|
)
|
||||||
|
|
||||||
tools_removed = 0
|
tools_removed = 0
|
||||||
for name in [*removed, *changed]:
|
for name in [*removed, *changed]:
|
||||||
tools_removed += _unregister_server_tools(registry, name)
|
tools_removed += _unregister_server_tools(self._registry, name)
|
||||||
await _close_server(state, name)
|
await self._close_server(name)
|
||||||
|
|
||||||
state._mcp_servers = next_servers
|
for name in [*removed, *authorization_pending]:
|
||||||
|
self._runtime_statuses.pop(name, None)
|
||||||
|
|
||||||
|
self._servers = next_servers
|
||||||
retry_missing = sorted(
|
retry_missing = sorted(
|
||||||
name
|
name
|
||||||
for name in next_names
|
for name in next_names
|
||||||
if name not in state._mcp_stacks and name not in set(added) | set(changed)
|
if name not in self._connections
|
||||||
|
and name not in set(added) | set(changed)
|
||||||
|
and name not in authorization_pending
|
||||||
|
)
|
||||||
|
to_connect_names = sorted(
|
||||||
|
(set(added) | set(changed) | set(retry_missing))
|
||||||
|
- authorization_pending
|
||||||
)
|
)
|
||||||
to_connect_names = sorted(set(added) | set(changed) | set(retry_missing))
|
|
||||||
to_connect = {name: next_servers[name] for name in to_connect_names}
|
to_connect = {name: next_servers[name] for name in to_connect_names}
|
||||||
connected: dict[str, MCPConnection] = {}
|
connected: dict[str, MCPConnection] = {}
|
||||||
if to_connect:
|
if to_connect:
|
||||||
connected = await connect_mcp_servers(to_connect, registry)
|
self._set_runtime_status(to_connect, "connecting")
|
||||||
if getattr(state, "_mcp_closing", False):
|
try:
|
||||||
for connection in connected.values():
|
connected = await connect_mcp_servers(to_connect, self._registry)
|
||||||
await connection.aclose()
|
except BaseException:
|
||||||
return {
|
self._set_runtime_status(to_connect, "failed")
|
||||||
"ok": False,
|
raise
|
||||||
"message": "MCP connections are shutting down.",
|
if self._closing:
|
||||||
"requires_restart": True,
|
await _close_mcp_connections(connected)
|
||||||
}
|
return self._closing_result()
|
||||||
state._mcp_stacks.update(connected)
|
self._connections.update(connected)
|
||||||
_attach_reconnect_handlers(state, registry, connected)
|
self._record_connection_result(to_connect, connected)
|
||||||
|
self._attach_reconnect_handlers(connected)
|
||||||
|
|
||||||
failed = sorted(set(to_connect) - set(connected))
|
failed = sorted(set(to_connect) - set(connected))
|
||||||
unchanged = not removed and not added and not changed and not retry_missing
|
unchanged = not removed and not added and not changed and not retry_missing
|
||||||
ok = not failed
|
ok = not failed
|
||||||
if failed:
|
if failed:
|
||||||
message = "MCP config reloaded, but some servers did not connect: " + ", ".join(failed)
|
message = (
|
||||||
|
"MCP config reloaded, but some servers did not connect: "
|
||||||
|
+ ", ".join(failed)
|
||||||
|
)
|
||||||
elif unchanged:
|
elif unchanged:
|
||||||
message = "MCP config is already live."
|
message = "MCP config is already live."
|
||||||
elif retry_missing and not added and not changed and not removed:
|
elif retry_missing and not added and not changed and not removed:
|
||||||
@@ -1360,7 +1560,8 @@ async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
|
|||||||
message = "MCP config reloaded without restarting nanobot."
|
message = "MCP config reloaded without restarting nanobot."
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"MCP hot reload: added={} changed={} removed={} retried={} connected={} failed={} tools_removed={}",
|
"MCP hot reload: added={} changed={} removed={} retried={} "
|
||||||
|
"connected={} failed={} tools_removed={}",
|
||||||
added,
|
added,
|
||||||
changed,
|
changed,
|
||||||
removed,
|
removed,
|
||||||
@@ -1376,114 +1577,51 @@ async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
|
|||||||
"changed": changed,
|
"changed": changed,
|
||||||
"removed": removed,
|
"removed": removed,
|
||||||
"retried": retry_missing,
|
"retried": retry_missing,
|
||||||
"connected": sorted(state._mcp_stacks),
|
"connected": sorted(self._connections),
|
||||||
"configured": sorted(state._mcp_servers),
|
"configured": sorted(self._servers),
|
||||||
"failed": failed,
|
"failed": failed,
|
||||||
"tools_removed": tools_removed,
|
"tools_removed": tools_removed,
|
||||||
"requires_restart": False,
|
"requires_restart": False,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
async def request_mcp_reload(
|
def _closing_result() -> dict[str, Any]:
|
||||||
bus: MessageBus,
|
|
||||||
*,
|
|
||||||
timeout: float = 15.0,
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
"""Ask the running agent loop to reconcile live MCP connections."""
|
|
||||||
loop = asyncio.get_running_loop()
|
|
||||||
ack: asyncio.Future[dict[str, Any]] = loop.create_future()
|
|
||||||
await bus.publish_inbound(
|
|
||||||
InboundMessage(
|
|
||||||
channel="system",
|
|
||||||
sender_id="webui-settings",
|
|
||||||
chat_id="runtime",
|
|
||||||
content=RUNTIME_CONTROL_MCP_RELOAD,
|
|
||||||
metadata={
|
|
||||||
INBOUND_META_RUNTIME_CONTROL: RUNTIME_CONTROL_MCP_RELOAD,
|
|
||||||
RUNTIME_CONTROL_ACK: ack,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
result = await asyncio.wait_for(ack, timeout=timeout)
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
return {
|
return {
|
||||||
"ok": False,
|
"ok": False,
|
||||||
"message": "MCP hot reload timed out. Restart nanobot to pick up changes.",
|
"message": "MCP connections are shutting down.",
|
||||||
"requires_restart": True,
|
|
||||||
}
|
|
||||||
return result if isinstance(cast(object, result), dict) else {
|
|
||||||
"ok": False,
|
|
||||||
"message": "MCP hot reload returned an unexpected response.",
|
|
||||||
"requires_restart": True,
|
"requires_restart": True,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def _attach_reconnect_handlers(self, server_names: Iterable[str]) -> None:
|
||||||
async def handle_runtime_control(state: Any, msg: InboundMessage, registry: ToolRegistry) -> bool:
|
async def reconnect(
|
||||||
metadata = msg.metadata if isinstance(cast(object, msg.metadata), dict) else {}
|
server_name: str,
|
||||||
control = metadata.get(INBOUND_META_RUNTIME_CONTROL)
|
tool_name: str,
|
||||||
if control != RUNTIME_CONTROL_MCP_RELOAD:
|
stale_tool: Tool,
|
||||||
return False
|
) -> Tool | None:
|
||||||
|
return await self._refresh_terminated_server(
|
||||||
ack = metadata.get(RUNTIME_CONTROL_ACK)
|
|
||||||
try:
|
|
||||||
result = await reload_servers(state, registry)
|
|
||||||
except Exception as exc:
|
|
||||||
logger.exception("MCP hot reload failed")
|
|
||||||
result = {
|
|
||||||
"ok": False,
|
|
||||||
"message": "MCP hot reload failed. Restart nanobot to pick up changes.",
|
|
||||||
"requires_restart": True,
|
|
||||||
"error": str(exc),
|
|
||||||
}
|
|
||||||
if isinstance(ack, asyncio.Future) and not ack.done():
|
|
||||||
cast(asyncio.Future[dict[str, Any]], ack).set_result(result)
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
def _reload_lock(state: Any) -> asyncio.Lock:
|
|
||||||
try:
|
|
||||||
return _RELOAD_LOCKS[state]
|
|
||||||
except KeyError:
|
|
||||||
lock = asyncio.Lock()
|
|
||||||
_RELOAD_LOCKS[state] = lock
|
|
||||||
return lock
|
|
||||||
|
|
||||||
|
|
||||||
def _attach_reconnect_handlers(
|
|
||||||
state: Any,
|
|
||||||
registry: ToolRegistry,
|
|
||||||
server_names: Mapping[str, Any] | set[str] | list[str] | tuple[str, ...],
|
|
||||||
) -> None:
|
|
||||||
async def reconnect(server_name: str, tool_name: str, stale_tool: Tool) -> Tool | None:
|
|
||||||
return await _refresh_terminated_server(
|
|
||||||
state,
|
|
||||||
registry,
|
|
||||||
server_name,
|
server_name,
|
||||||
tool_name,
|
tool_name,
|
||||||
stale_tool,
|
stale_tool,
|
||||||
)
|
)
|
||||||
|
|
||||||
for server_name in server_names:
|
for server_name in server_names:
|
||||||
for tool_name in list(registry.tool_names):
|
for tool_name in list(self._registry.tool_names):
|
||||||
tool = registry.get(tool_name)
|
tool = self._registry.get(tool_name)
|
||||||
if not _tool_belongs_to_server(tool, tool_name, server_name):
|
if not _tool_belongs_to_server(tool, tool_name, server_name):
|
||||||
continue
|
continue
|
||||||
if isinstance(tool, _MCPWrapperBase):
|
if isinstance(tool, _MCPWrapperBase):
|
||||||
tool.set_reconnect_handler(reconnect)
|
tool.set_reconnect_handler(reconnect)
|
||||||
|
|
||||||
|
|
||||||
async def _refresh_terminated_server(
|
async def _refresh_terminated_server(
|
||||||
state: Any,
|
self,
|
||||||
registry: ToolRegistry,
|
|
||||||
server_name: str,
|
server_name: str,
|
||||||
tool_name: str,
|
tool_name: str,
|
||||||
stale_tool: Tool,
|
stale_tool: Tool,
|
||||||
) -> Tool | None:
|
) -> Tool | None:
|
||||||
async with _reload_lock(state):
|
async with self._lock:
|
||||||
if getattr(state, "_mcp_closing", False):
|
if self._closing:
|
||||||
return None
|
return None
|
||||||
cfg = state._mcp_servers.get(server_name)
|
cfg = self._servers.get(server_name)
|
||||||
if cfg is None:
|
if cfg is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"MCP server '{}' session terminated but is no longer configured",
|
"MCP server '{}' session terminated but is no longer configured",
|
||||||
@@ -1491,29 +1629,56 @@ async def _refresh_terminated_server(
|
|||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
current_tool = registry.get(tool_name)
|
current_tool = self._registry.get(tool_name)
|
||||||
if (
|
if (
|
||||||
current_tool is not None
|
current_tool is not None
|
||||||
and current_tool is not stale_tool
|
and current_tool is not stale_tool
|
||||||
and server_name in state._mcp_stacks
|
and server_name in self._connections
|
||||||
):
|
):
|
||||||
return current_tool
|
return current_tool
|
||||||
|
|
||||||
logger.warning("MCP server '{}' session terminated; refreshing connection", server_name)
|
logger.warning(
|
||||||
_unregister_server_tools(registry, server_name)
|
"MCP server '{}' session terminated; refreshing connection",
|
||||||
await _close_server(state, server_name)
|
server_name,
|
||||||
|
)
|
||||||
|
_unregister_server_tools(self._registry, server_name)
|
||||||
|
await self._close_server(server_name)
|
||||||
|
|
||||||
connected = await connect_mcp_servers({server_name: cfg}, registry)
|
self._set_runtime_status({server_name}, "connecting")
|
||||||
if getattr(state, "_mcp_closing", False):
|
connected = await connect_mcp_servers(
|
||||||
for connection in connected.values():
|
{server_name: cfg},
|
||||||
await connection.aclose()
|
self._registry,
|
||||||
|
)
|
||||||
|
if self._closing:
|
||||||
|
await _close_mcp_connections(connected)
|
||||||
return None
|
return None
|
||||||
state._mcp_stacks.update(connected)
|
self._connections.update(connected)
|
||||||
_attach_reconnect_handlers(state, registry, connected)
|
self._record_connection_result({server_name}, connected)
|
||||||
|
self._attach_reconnect_handlers(connected)
|
||||||
if server_name not in connected:
|
if server_name not in connected:
|
||||||
logger.warning("MCP server '{}' reconnect failed after session termination", server_name)
|
logger.warning(
|
||||||
|
"MCP server '{}' reconnect failed after session termination",
|
||||||
|
server_name,
|
||||||
|
)
|
||||||
return None
|
return None
|
||||||
return registry.get(tool_name)
|
return self._registry.get(tool_name)
|
||||||
|
|
||||||
|
async def _close_server(self, server_name: str) -> None:
|
||||||
|
connection = self._connections.pop(server_name, None)
|
||||||
|
if connection is None:
|
||||||
|
return
|
||||||
|
await _close_mcp_connection(server_name, connection)
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
"""Close every connection while excluding reconnect and hot reload."""
|
||||||
|
self._closing = True
|
||||||
|
async with self._lock:
|
||||||
|
connections = dict(self._connections)
|
||||||
|
self._connections.clear()
|
||||||
|
self._runtime_statuses.clear()
|
||||||
|
for name in self._servers:
|
||||||
|
_unregister_server_tools(self._registry, name)
|
||||||
|
await _close_mcp_connections(connections)
|
||||||
|
|
||||||
|
|
||||||
def _server_signature(cfg: Any) -> Any:
|
def _server_signature(cfg: Any) -> Any:
|
||||||
@@ -1540,34 +1705,3 @@ def _unregister_server_tools(registry: ToolRegistry, server_name: str) -> int:
|
|||||||
registry.unregister(tool_name)
|
registry.unregister(tool_name)
|
||||||
removed += 1
|
removed += 1
|
||||||
return removed
|
return removed
|
||||||
|
|
||||||
|
|
||||||
async def _close_server(state: Any, server_name: str) -> None:
|
|
||||||
stack = state._mcp_stacks.pop(server_name, None)
|
|
||||||
if stack is None:
|
|
||||||
return
|
|
||||||
try:
|
|
||||||
await stack.aclose()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
if task_is_cancelling():
|
|
||||||
raise
|
|
||||||
logger.debug("MCP server '{}' cleanup error (can be ignored)", server_name)
|
|
||||||
except (RuntimeError, BaseExceptionGroup):
|
|
||||||
logger.debug("MCP server '{}' cleanup error (can be ignored)", server_name)
|
|
||||||
|
|
||||||
|
|
||||||
async def close_mcp_servers(state: Any) -> None:
|
|
||||||
"""Close every MCP connection while excluding reconnect and hot reload."""
|
|
||||||
state._mcp_closing = True
|
|
||||||
async with _reload_lock(state):
|
|
||||||
connections = list(state._mcp_stacks.items())
|
|
||||||
state._mcp_stacks.clear()
|
|
||||||
for name, connection in connections:
|
|
||||||
try:
|
|
||||||
await connection.aclose()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
if task_is_cancelling():
|
|
||||||
raise
|
|
||||||
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
|
|
||||||
except (RuntimeError, BaseExceptionGroup):
|
|
||||||
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
|
|
||||||
|
|||||||
@@ -0,0 +1,401 @@
|
|||||||
|
"""OAuth support for remote MCP servers.
|
||||||
|
|
||||||
|
This module intentionally owns MCP OAuth end to end. Provider OAuth has a
|
||||||
|
different lifecycle and storage contract, so sharing a higher-level workflow
|
||||||
|
would couple unrelated extension boundaries.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import secrets
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from contextlib import suppress
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, TypedDict, cast
|
||||||
|
|
||||||
|
from filelock import FileLock
|
||||||
|
from loguru import logger
|
||||||
|
from mcp.client.auth import OAuthClientProvider
|
||||||
|
from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken
|
||||||
|
from pydantic import AnyHttpUrl, AnyUrl
|
||||||
|
|
||||||
|
from nanobot.config.paths import get_data_dir
|
||||||
|
from nanobot.utils.helpers import _write_text_atomic # pyright: ignore[reportPrivateUsage]
|
||||||
|
|
||||||
|
MCP_OAUTH_CALLBACK_PATH = "/auth/mcp/callback"
|
||||||
|
_STORE_VERSION = 1
|
||||||
|
_STORE_LOCK_TIMEOUT_S = 15
|
||||||
|
_DEFAULT_REDIRECT_URI = f"http://127.0.0.1{MCP_OAUTH_CALLBACK_PATH}"
|
||||||
|
_CLIENT_URI = AnyHttpUrl("https://github.com/HKUDS/nanobot")
|
||||||
|
_LOGO_URI = AnyHttpUrl(
|
||||||
|
"https://raw.githubusercontent.com/HKUDS/nanobot/main/"
|
||||||
|
"webui/public/brand/nanobot_apple_touch.png"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _StoredServer(TypedDict, total=False):
|
||||||
|
server_fingerprint: str
|
||||||
|
write_lease: str
|
||||||
|
tokens: dict[str, Any]
|
||||||
|
client_info: dict[str, Any]
|
||||||
|
redirect_uri: str
|
||||||
|
|
||||||
|
|
||||||
|
class _CredentialStore(TypedDict):
|
||||||
|
version: int
|
||||||
|
servers: dict[str, _StoredServer]
|
||||||
|
generations: dict[str, str]
|
||||||
|
|
||||||
|
|
||||||
|
class MCPAuthorizationRequiredError(RuntimeError):
|
||||||
|
"""Raised when a background MCP connection needs interactive authorization."""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class MCPOAuthHandlers:
|
||||||
|
"""Browser callbacks supplied only for a user-initiated OAuth attempt."""
|
||||||
|
|
||||||
|
redirect_uri: str
|
||||||
|
redirect_handler: Callable[[str], Awaitable[None]]
|
||||||
|
callback_handler: Callable[[], Awaitable[tuple[str, str | None]]]
|
||||||
|
reset_credentials: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
def _store_path() -> Path:
|
||||||
|
return get_data_dir() / "auth" / "mcp.json"
|
||||||
|
|
||||||
|
|
||||||
|
def _server_fingerprint(server_url: str) -> str:
|
||||||
|
return hashlib.sha256(server_url.strip().encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def _empty_store() -> _CredentialStore:
|
||||||
|
return {"version": _STORE_VERSION, "servers": {}, "generations": {}}
|
||||||
|
|
||||||
|
|
||||||
|
def _stored_server(value: object) -> _StoredServer | None:
|
||||||
|
if not isinstance(value, dict):
|
||||||
|
return None
|
||||||
|
raw = cast(dict[object, object], value)
|
||||||
|
entry: _StoredServer = {}
|
||||||
|
fingerprint = raw.get("server_fingerprint")
|
||||||
|
if isinstance(fingerprint, str):
|
||||||
|
entry["server_fingerprint"] = fingerprint
|
||||||
|
write_lease = raw.get("write_lease")
|
||||||
|
if isinstance(write_lease, str) and write_lease:
|
||||||
|
entry["write_lease"] = write_lease
|
||||||
|
redirect_uri = raw.get("redirect_uri")
|
||||||
|
if isinstance(redirect_uri, str):
|
||||||
|
entry["redirect_uri"] = redirect_uri
|
||||||
|
tokens = raw.get("tokens")
|
||||||
|
if isinstance(tokens, dict):
|
||||||
|
token_values = cast(dict[object, object], tokens)
|
||||||
|
if all(isinstance(key, str) for key in token_values):
|
||||||
|
entry["tokens"] = cast(dict[str, Any], token_values)
|
||||||
|
client_info = raw.get("client_info")
|
||||||
|
if isinstance(client_info, dict):
|
||||||
|
client_values = cast(dict[object, object], client_info)
|
||||||
|
if all(isinstance(key, str) for key in client_values):
|
||||||
|
entry["client_info"] = cast(dict[str, Any], client_values)
|
||||||
|
return entry
|
||||||
|
|
||||||
|
|
||||||
|
def _read_store_unlocked(path: Path) -> _CredentialStore:
|
||||||
|
try:
|
||||||
|
raw = cast(object, json.loads(path.read_text(encoding="utf-8")))
|
||||||
|
except FileNotFoundError:
|
||||||
|
return _empty_store()
|
||||||
|
except (OSError, ValueError, TypeError) as exc:
|
||||||
|
logger.warning("Could not read MCP OAuth credentials: {}", type(exc).__name__)
|
||||||
|
return _empty_store()
|
||||||
|
if not isinstance(raw, dict):
|
||||||
|
return _empty_store()
|
||||||
|
payload = cast(dict[object, object], raw)
|
||||||
|
raw_servers = payload.get("servers")
|
||||||
|
if not isinstance(raw_servers, dict):
|
||||||
|
return _empty_store()
|
||||||
|
servers: dict[str, _StoredServer] = {}
|
||||||
|
for name, value in cast(dict[object, object], raw_servers).items():
|
||||||
|
entry = _stored_server(value)
|
||||||
|
if isinstance(name, str) and entry is not None:
|
||||||
|
servers[name] = entry
|
||||||
|
generations: dict[str, str] = {}
|
||||||
|
raw_generations = payload.get("generations")
|
||||||
|
if isinstance(raw_generations, dict):
|
||||||
|
for name, value in cast(dict[object, object], raw_generations).items():
|
||||||
|
if isinstance(name, str) and isinstance(value, str) and value:
|
||||||
|
generations[name] = value
|
||||||
|
return {
|
||||||
|
"version": _STORE_VERSION,
|
||||||
|
"servers": servers,
|
||||||
|
"generations": generations,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _with_store_lock(path: Path) -> FileLock:
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
return FileLock(str(path.with_suffix(".lock")), timeout=_STORE_LOCK_TIMEOUT_S)
|
||||||
|
|
||||||
|
|
||||||
|
def _write_store_unlocked(path: Path, payload: _CredentialStore) -> None:
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
with suppress(OSError):
|
||||||
|
os.chmod(path.parent, 0o700)
|
||||||
|
_write_text_atomic(path, json.dumps(payload, indent=2, ensure_ascii=False))
|
||||||
|
with suppress(OSError):
|
||||||
|
os.chmod(path, 0o600)
|
||||||
|
|
||||||
|
|
||||||
|
class MCPOAuthStorage:
|
||||||
|
"""Persistent MCP SDK token storage, isolated by config name and server URL."""
|
||||||
|
|
||||||
|
def __init__(self, server_name: str, server_url: str) -> None:
|
||||||
|
self.server_name = server_name
|
||||||
|
self.server_fingerprint = _server_fingerprint(server_url)
|
||||||
|
self._observed_generation = self._read_generation_sync()
|
||||||
|
self._write_lease: str | None = None
|
||||||
|
|
||||||
|
def _read_generation_sync(self) -> str | None:
|
||||||
|
path = _store_path()
|
||||||
|
if not path.exists():
|
||||||
|
return None
|
||||||
|
# Writes replace the whole file atomically, so this observes either side
|
||||||
|
# of a concurrent deletion without blocking the async connection path.
|
||||||
|
return _read_store_unlocked(path)["generations"].get(self.server_name)
|
||||||
|
|
||||||
|
def _generation_is_current(self, payload: _CredentialStore) -> bool:
|
||||||
|
return payload["generations"].get(self.server_name) == self._observed_generation
|
||||||
|
|
||||||
|
def _entry_unlocked(self, payload: _CredentialStore) -> _StoredServer | None:
|
||||||
|
servers = payload["servers"]
|
||||||
|
entry = servers.get(self.server_name)
|
||||||
|
if entry is None or entry.get("server_fingerprint") != self.server_fingerprint:
|
||||||
|
return None
|
||||||
|
return entry
|
||||||
|
|
||||||
|
def _bind_entry_unlocked(
|
||||||
|
self,
|
||||||
|
payload: _CredentialStore,
|
||||||
|
*,
|
||||||
|
create: bool,
|
||||||
|
) -> tuple[_StoredServer | None, bool]:
|
||||||
|
if not self._generation_is_current(payload):
|
||||||
|
return None, False
|
||||||
|
entry = self._entry_unlocked(payload)
|
||||||
|
if self._write_lease is not None:
|
||||||
|
if entry is None or entry.get("write_lease") != self._write_lease:
|
||||||
|
return None, False
|
||||||
|
return entry, False
|
||||||
|
if entry is None:
|
||||||
|
if not create:
|
||||||
|
return None, False
|
||||||
|
self._write_lease = secrets.token_urlsafe(24)
|
||||||
|
entry = _StoredServer(
|
||||||
|
server_fingerprint=self.server_fingerprint,
|
||||||
|
write_lease=self._write_lease,
|
||||||
|
)
|
||||||
|
payload["servers"][self.server_name] = entry
|
||||||
|
return entry, True
|
||||||
|
write_lease = entry.get("write_lease")
|
||||||
|
changed = not isinstance(write_lease, str) or not write_lease
|
||||||
|
if changed:
|
||||||
|
write_lease = secrets.token_urlsafe(24)
|
||||||
|
entry["write_lease"] = write_lease
|
||||||
|
self._write_lease = write_lease
|
||||||
|
return entry, changed
|
||||||
|
|
||||||
|
def _read_entry_sync(self) -> _StoredServer | None:
|
||||||
|
path = _store_path()
|
||||||
|
with _with_store_lock(path):
|
||||||
|
payload = _read_store_unlocked(path)
|
||||||
|
entry, changed = self._bind_entry_unlocked(payload, create=False)
|
||||||
|
if changed:
|
||||||
|
_write_store_unlocked(path, payload)
|
||||||
|
return entry
|
||||||
|
|
||||||
|
def _update_entry_sync(
|
||||||
|
self,
|
||||||
|
update: Callable[[_StoredServer], None],
|
||||||
|
*,
|
||||||
|
create: bool = True,
|
||||||
|
claim: bool = False,
|
||||||
|
) -> bool:
|
||||||
|
path = _store_path()
|
||||||
|
with _with_store_lock(path):
|
||||||
|
payload = _read_store_unlocked(path)
|
||||||
|
if claim:
|
||||||
|
# A browser flow owns subsequent SDK writes until another flow
|
||||||
|
# claims the entry or the configured server is removed.
|
||||||
|
if not self._generation_is_current(payload):
|
||||||
|
logger.info(
|
||||||
|
"Ignored stale MCP OAuth credential claim for '{}'",
|
||||||
|
self.server_name,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
entry = self._entry_unlocked(payload)
|
||||||
|
if entry is None:
|
||||||
|
entry = _StoredServer(server_fingerprint=self.server_fingerprint)
|
||||||
|
payload["servers"][self.server_name] = entry
|
||||||
|
self._write_lease = secrets.token_urlsafe(24)
|
||||||
|
entry["write_lease"] = self._write_lease
|
||||||
|
else:
|
||||||
|
entry, _ = self._bind_entry_unlocked(payload, create=create)
|
||||||
|
if entry is None:
|
||||||
|
if self._write_lease is not None:
|
||||||
|
logger.info(
|
||||||
|
"Ignored stale MCP OAuth credential update for '{}'",
|
||||||
|
self.server_name,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
update(entry)
|
||||||
|
payload["version"] = _STORE_VERSION
|
||||||
|
_write_store_unlocked(path, payload)
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def get_tokens(self) -> OAuthToken | None:
|
||||||
|
entry = await asyncio.to_thread(self._read_entry_sync)
|
||||||
|
raw = entry.get("tokens") if entry is not None else None
|
||||||
|
if not isinstance(raw, dict):
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
return OAuthToken.model_validate(raw)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
logger.warning("Ignoring invalid MCP OAuth tokens for '{}'", self.server_name)
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def set_tokens(self, tokens: OAuthToken) -> None:
|
||||||
|
raw = tokens.model_dump(mode="json", exclude_none=True)
|
||||||
|
|
||||||
|
def update(entry: _StoredServer) -> None:
|
||||||
|
entry["tokens"] = raw
|
||||||
|
|
||||||
|
await asyncio.to_thread(self._update_entry_sync, update)
|
||||||
|
|
||||||
|
async def clear_tokens(self) -> None:
|
||||||
|
def update(entry: _StoredServer) -> None:
|
||||||
|
entry.pop("tokens", None)
|
||||||
|
|
||||||
|
await asyncio.to_thread(self._update_entry_sync, update, create=False)
|
||||||
|
|
||||||
|
async def get_client_info(self) -> OAuthClientInformationFull | None:
|
||||||
|
entry = await asyncio.to_thread(self._read_entry_sync)
|
||||||
|
raw = entry.get("client_info") if entry is not None else None
|
||||||
|
if not isinstance(raw, dict):
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
return OAuthClientInformationFull.model_validate(raw)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
logger.warning("Ignoring invalid MCP OAuth client info for '{}'", self.server_name)
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def set_client_info(self, client_info: OAuthClientInformationFull) -> None:
|
||||||
|
raw = client_info.model_dump(mode="json", exclude_none=True)
|
||||||
|
|
||||||
|
def update(entry: _StoredServer) -> None:
|
||||||
|
entry["client_info"] = raw
|
||||||
|
|
||||||
|
await asyncio.to_thread(self._update_entry_sync, update)
|
||||||
|
|
||||||
|
async def redirect_uri(self) -> str | None:
|
||||||
|
entry = await asyncio.to_thread(self._read_entry_sync)
|
||||||
|
value = entry.get("redirect_uri") if entry is not None else None
|
||||||
|
return value if isinstance(value, str) and value else None
|
||||||
|
|
||||||
|
async def prepare_redirect_uri(self, redirect_uri: str, *, reset: bool = False) -> None:
|
||||||
|
def update(entry: _StoredServer) -> None:
|
||||||
|
changed = entry.get("redirect_uri") != redirect_uri
|
||||||
|
if reset:
|
||||||
|
entry.pop("tokens", None)
|
||||||
|
entry.pop("client_info", None)
|
||||||
|
elif changed:
|
||||||
|
# Dynamic registrations bind a client to its redirect URI.
|
||||||
|
entry.pop("client_info", None)
|
||||||
|
entry["redirect_uri"] = redirect_uri
|
||||||
|
|
||||||
|
claimed = await asyncio.to_thread(self._update_entry_sync, update, claim=True)
|
||||||
|
if not claimed:
|
||||||
|
raise MCPAuthorizationRequiredError("MCP authorization was cancelled")
|
||||||
|
|
||||||
|
def has_credentials(self) -> bool:
|
||||||
|
entry = self._read_entry_sync()
|
||||||
|
raw_tokens = entry.get("tokens") if entry is not None else None
|
||||||
|
if not isinstance(raw_tokens, dict):
|
||||||
|
return False
|
||||||
|
tokens = cast(dict[str, object], raw_tokens)
|
||||||
|
access_token = tokens.get("access_token")
|
||||||
|
return isinstance(access_token, str) and bool(access_token)
|
||||||
|
|
||||||
|
|
||||||
|
async def _missing_callback() -> tuple[str, str | None]:
|
||||||
|
raise MCPAuthorizationRequiredError("MCP server requires browser authorization")
|
||||||
|
|
||||||
|
|
||||||
|
async def create_mcp_oauth_auth(
|
||||||
|
server_name: str,
|
||||||
|
server_url: str,
|
||||||
|
handlers: MCPOAuthHandlers | None = None,
|
||||||
|
) -> OAuthClientProvider:
|
||||||
|
"""Build the official MCP SDK OAuth provider for one configured server."""
|
||||||
|
storage = MCPOAuthStorage(server_name, server_url)
|
||||||
|
if handlers is not None:
|
||||||
|
await storage.prepare_redirect_uri(
|
||||||
|
handlers.redirect_uri,
|
||||||
|
reset=handlers.reset_credentials,
|
||||||
|
)
|
||||||
|
redirect_uri = handlers.redirect_uri
|
||||||
|
redirect_handler = handlers.redirect_handler
|
||||||
|
callback_handler = handlers.callback_handler
|
||||||
|
else:
|
||||||
|
if not await asyncio.to_thread(storage.has_credentials):
|
||||||
|
# Do not perform discovery or dynamic registration from a background
|
||||||
|
# startup. Interactive OAuth begins only after an explicit user action.
|
||||||
|
raise MCPAuthorizationRequiredError("MCP server requires browser authorization")
|
||||||
|
redirect_uri = await storage.redirect_uri() or _DEFAULT_REDIRECT_URI
|
||||||
|
|
||||||
|
async def authorization_required(_authorization_url: str) -> None:
|
||||||
|
await storage.clear_tokens()
|
||||||
|
raise MCPAuthorizationRequiredError("MCP server requires browser authorization")
|
||||||
|
|
||||||
|
redirect_handler = authorization_required
|
||||||
|
callback_handler = _missing_callback
|
||||||
|
|
||||||
|
metadata = OAuthClientMetadata(
|
||||||
|
redirect_uris=[AnyUrl(redirect_uri)],
|
||||||
|
token_endpoint_auth_method="none",
|
||||||
|
client_name="nanobot",
|
||||||
|
client_uri=_CLIENT_URI,
|
||||||
|
logo_uri=_LOGO_URI,
|
||||||
|
software_id="https://github.com/HKUDS/nanobot",
|
||||||
|
)
|
||||||
|
return OAuthClientProvider(
|
||||||
|
server_url,
|
||||||
|
metadata,
|
||||||
|
storage,
|
||||||
|
redirect_handler=redirect_handler,
|
||||||
|
callback_handler=callback_handler,
|
||||||
|
timeout=300,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def mcp_oauth_has_credentials(server_name: str, server_url: str) -> bool:
|
||||||
|
"""Return whether this exact configured MCP instance has an access token."""
|
||||||
|
return MCPOAuthStorage(server_name, server_url).has_credentials()
|
||||||
|
|
||||||
|
|
||||||
|
def delete_mcp_oauth_credentials(server_name: str) -> bool:
|
||||||
|
"""Delete credentials for one config name without touching other MCP instances."""
|
||||||
|
path = _store_path()
|
||||||
|
with _with_store_lock(path):
|
||||||
|
payload = _read_store_unlocked(path)
|
||||||
|
servers = payload["servers"]
|
||||||
|
removed = servers.pop(server_name, None) is not None
|
||||||
|
# Rotate even when no entry exists so a flow created before removal cannot
|
||||||
|
# claim the name later and resurrect credentials.
|
||||||
|
payload["generations"][server_name] = secrets.token_urlsafe(24)
|
||||||
|
_write_store_unlocked(path, payload)
|
||||||
|
return removed
|
||||||
@@ -0,0 +1,319 @@
|
|||||||
|
"""Explicit runtime state boundary used by :class:`MyTool`."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Mapping
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Protocol, TypeAlias, runtime_checkable
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from nanobot.agent.subagent import SubagentManager, SubagentStatus
|
||||||
|
from nanobot.agent.tools.shell import ExecToolConfig
|
||||||
|
from nanobot.agent.tools.web import WebToolsConfig
|
||||||
|
from nanobot.config.schema import ModelPresetConfig
|
||||||
|
from nanobot.utils.llm_runtime import LLMRuntime
|
||||||
|
|
||||||
|
|
||||||
|
JsonScalar: TypeAlias = str | int | float | bool | None
|
||||||
|
JsonValue: TypeAlias = JsonScalar | list["JsonValue"] | dict[str, "JsonValue"]
|
||||||
|
|
||||||
|
|
||||||
|
RUNTIME_SNAPSHOT_KEYS = frozenset({
|
||||||
|
"model",
|
||||||
|
"model_preset",
|
||||||
|
"model_presets",
|
||||||
|
"max_iterations",
|
||||||
|
"context_window_tokens",
|
||||||
|
"workspace",
|
||||||
|
"provider_retry_mode",
|
||||||
|
"max_tool_result_chars",
|
||||||
|
"current_iteration",
|
||||||
|
"_current_iteration",
|
||||||
|
"tool_names",
|
||||||
|
"web_config",
|
||||||
|
"exec_config",
|
||||||
|
"subagents",
|
||||||
|
"_last_usage",
|
||||||
|
})
|
||||||
|
|
||||||
|
RUNTIME_COMMAND_KEYS = frozenset({
|
||||||
|
"model",
|
||||||
|
"model_preset",
|
||||||
|
"max_iterations",
|
||||||
|
"context_window_tokens",
|
||||||
|
"provider_retry_mode",
|
||||||
|
"max_tool_result_chars",
|
||||||
|
"workspace",
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class RuntimeSnapshot:
|
||||||
|
"""Detached, allowlisted values available to self-inspection."""
|
||||||
|
|
||||||
|
model: str
|
||||||
|
model_preset: str | None
|
||||||
|
model_presets: dict[str, dict[str, object]]
|
||||||
|
max_iterations: int
|
||||||
|
context_window_tokens: int
|
||||||
|
workspace: Path | str
|
||||||
|
provider_retry_mode: str
|
||||||
|
max_tool_result_chars: int
|
||||||
|
current_iteration: int
|
||||||
|
tool_names: list[str]
|
||||||
|
web_config: dict[str, object]
|
||||||
|
exec_config: dict[str, object]
|
||||||
|
subagent_statuses: dict[str, dict[str, object]]
|
||||||
|
last_usage: dict[str, int]
|
||||||
|
scratchpad: dict[str, JsonValue]
|
||||||
|
|
||||||
|
def as_mapping(self) -> Mapping[str, object]:
|
||||||
|
"""Return the fixed public names understood by ``MyTool``."""
|
||||||
|
values: dict[str, object] = {
|
||||||
|
"model": self.model,
|
||||||
|
"model_preset": self.model_preset,
|
||||||
|
"model_presets": self.model_presets,
|
||||||
|
"max_iterations": self.max_iterations,
|
||||||
|
"context_window_tokens": self.context_window_tokens,
|
||||||
|
"workspace": self.workspace,
|
||||||
|
"provider_retry_mode": self.provider_retry_mode,
|
||||||
|
"max_tool_result_chars": self.max_tool_result_chars,
|
||||||
|
"current_iteration": self.current_iteration,
|
||||||
|
"_current_iteration": self.current_iteration,
|
||||||
|
"tool_names": self.tool_names,
|
||||||
|
"web_config": self.web_config,
|
||||||
|
"exec_config": self.exec_config,
|
||||||
|
"subagents": {"_task_statuses": self.subagent_statuses},
|
||||||
|
"_last_usage": self.last_usage,
|
||||||
|
}
|
||||||
|
assert values.keys() == RUNTIME_SNAPSHOT_KEYS
|
||||||
|
return values
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class RuntimeControl(Protocol):
|
||||||
|
"""The complete runtime capability exposed to ``MyTool``."""
|
||||||
|
|
||||||
|
def snapshot(self) -> RuntimeSnapshot: ...
|
||||||
|
|
||||||
|
def set_model(self, model: str) -> LLMRuntime: ...
|
||||||
|
|
||||||
|
def set_model_preset(
|
||||||
|
self,
|
||||||
|
name: str,
|
||||||
|
*,
|
||||||
|
session_key: str | None,
|
||||||
|
) -> LLMRuntime: ...
|
||||||
|
|
||||||
|
def set_max_iterations(self, value: int) -> None: ...
|
||||||
|
|
||||||
|
def set_context_window_tokens(self, value: int) -> LLMRuntime: ...
|
||||||
|
|
||||||
|
def set_provider_retry_mode(self, value: str) -> None: ...
|
||||||
|
|
||||||
|
def set_max_tool_result_chars(self, value: int) -> None: ...
|
||||||
|
|
||||||
|
def set_workspace_display(self, value: str) -> None: ...
|
||||||
|
|
||||||
|
def set_scratchpad(self, key: str, value: JsonValue, *, max_keys: int) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
class _RuntimeControlTarget(Protocol):
|
||||||
|
"""Narrow structural dependency required by ``AgentRuntimeControl``."""
|
||||||
|
|
||||||
|
max_iterations: int
|
||||||
|
provider_retry_mode: str
|
||||||
|
max_tool_result_chars: int
|
||||||
|
web_config: WebToolsConfig
|
||||||
|
exec_config: ExecToolConfig
|
||||||
|
subagents: SubagentManager
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model(self) -> str: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model_preset(self) -> str | None: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model_presets(self) -> Mapping[str, ModelPresetConfig]: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def context_window_tokens(self) -> int: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def workspace(self) -> Path: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def current_iteration(self) -> int: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def tool_names(self) -> list[str]: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def last_usage(self) -> Mapping[str, int]: ...
|
||||||
|
|
||||||
|
def set_runtime_model(self, model: str) -> LLMRuntime: ...
|
||||||
|
|
||||||
|
def set_runtime_context_window(self, context_window_tokens: int) -> LLMRuntime: ...
|
||||||
|
|
||||||
|
def set_model_preset(self, name: str | None) -> LLMRuntime: ...
|
||||||
|
|
||||||
|
def set_session_model_preset(self, session_key: str, name: str) -> LLMRuntime: ...
|
||||||
|
|
||||||
|
|
||||||
|
class AgentRuntimeControl:
|
||||||
|
"""Allowlisted adapter from agent-loop state to ``RuntimeControl``."""
|
||||||
|
|
||||||
|
def __init__(self, target: _RuntimeControlTarget) -> None:
|
||||||
|
self.__target = target
|
||||||
|
self.__scratchpad: dict[str, JsonValue] = {}
|
||||||
|
self.__workspace_display: str | None = None
|
||||||
|
|
||||||
|
def snapshot(self) -> RuntimeSnapshot:
|
||||||
|
target = self.__target
|
||||||
|
return RuntimeSnapshot(
|
||||||
|
model=target.model,
|
||||||
|
model_preset=target.model_preset,
|
||||||
|
model_presets=_snapshot_model_presets(target.model_presets),
|
||||||
|
max_iterations=target.max_iterations,
|
||||||
|
context_window_tokens=target.context_window_tokens,
|
||||||
|
workspace=(
|
||||||
|
self.__workspace_display
|
||||||
|
if self.__workspace_display is not None
|
||||||
|
else target.workspace
|
||||||
|
),
|
||||||
|
provider_retry_mode=target.provider_retry_mode,
|
||||||
|
max_tool_result_chars=target.max_tool_result_chars,
|
||||||
|
current_iteration=target.current_iteration,
|
||||||
|
tool_names=list(target.tool_names),
|
||||||
|
web_config=_snapshot_web_config(target.web_config),
|
||||||
|
exec_config=_snapshot_exec_config(target.exec_config),
|
||||||
|
subagent_statuses=_snapshot_subagent_statuses(target.subagents),
|
||||||
|
last_usage=dict(target.last_usage),
|
||||||
|
scratchpad=_snapshot_json_mapping(self.__scratchpad),
|
||||||
|
)
|
||||||
|
|
||||||
|
def set_model(self, model: str) -> LLMRuntime:
|
||||||
|
return self.__target.set_runtime_model(model)
|
||||||
|
|
||||||
|
def set_model_preset(
|
||||||
|
self,
|
||||||
|
name: str,
|
||||||
|
*,
|
||||||
|
session_key: str | None,
|
||||||
|
) -> LLMRuntime:
|
||||||
|
if session_key is not None:
|
||||||
|
return self.__target.set_session_model_preset(session_key, name)
|
||||||
|
return self.__target.set_model_preset(name)
|
||||||
|
|
||||||
|
def set_max_iterations(self, value: int) -> None:
|
||||||
|
self.__target.max_iterations = value
|
||||||
|
self.__target.subagents.max_iterations = value
|
||||||
|
|
||||||
|
def set_context_window_tokens(self, value: int) -> LLMRuntime:
|
||||||
|
return self.__target.set_runtime_context_window(value)
|
||||||
|
|
||||||
|
def set_provider_retry_mode(self, value: str) -> None:
|
||||||
|
self.__target.provider_retry_mode = value
|
||||||
|
|
||||||
|
def set_max_tool_result_chars(self, value: int) -> None:
|
||||||
|
self.__target.max_tool_result_chars = value
|
||||||
|
|
||||||
|
def set_workspace_display(self, value: str) -> None:
|
||||||
|
"""Preserve MyTool display compatibility without changing path enforcement."""
|
||||||
|
self.__workspace_display = value
|
||||||
|
|
||||||
|
def set_scratchpad(self, key: str, value: JsonValue, *, max_keys: int) -> None:
|
||||||
|
if key not in self.__scratchpad and len(self.__scratchpad) >= max_keys:
|
||||||
|
raise ValueError(f"scratchpad is full (max {max_keys} keys)")
|
||||||
|
self.__scratchpad[key] = value
|
||||||
|
|
||||||
|
|
||||||
|
def _snapshot_model_presets(
|
||||||
|
presets: Mapping[str, ModelPresetConfig],
|
||||||
|
) -> dict[str, dict[str, object]]:
|
||||||
|
return {
|
||||||
|
name: {
|
||||||
|
"label": preset.label,
|
||||||
|
"model": preset.model,
|
||||||
|
"provider": preset.provider,
|
||||||
|
"max_tokens": preset.max_tokens,
|
||||||
|
"context_window_tokens": preset.context_window_tokens,
|
||||||
|
"temperature": preset.temperature,
|
||||||
|
"reasoning_effort": preset.reasoning_effort,
|
||||||
|
}
|
||||||
|
for name, preset in presets.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _snapshot_web_config(config: WebToolsConfig) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"enable": config.enable,
|
||||||
|
# Proxy URLs may embed credentials. Presence is enough for diagnosis.
|
||||||
|
"proxy": "<configured>" if config.proxy else config.proxy,
|
||||||
|
"user_agent": config.user_agent,
|
||||||
|
"search": {
|
||||||
|
"provider": config.search.provider,
|
||||||
|
"base_url": config.search.base_url,
|
||||||
|
"max_results": config.search.max_results,
|
||||||
|
"timeout": config.search.timeout,
|
||||||
|
},
|
||||||
|
"fetch": {
|
||||||
|
"use_jina_reader": config.fetch.use_jina_reader,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _snapshot_exec_config(config: ExecToolConfig) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"enable": config.enable,
|
||||||
|
"timeout": config.timeout,
|
||||||
|
"path_prepend": config.path_prepend,
|
||||||
|
"path_append": config.path_append,
|
||||||
|
"sandbox": config.sandbox,
|
||||||
|
"sandbox_ro_binds": list(config.sandbox_ro_binds),
|
||||||
|
"sandbox_rw_binds": list(config.sandbox_rw_binds),
|
||||||
|
"allowed_env_keys": list(config.allowed_env_keys),
|
||||||
|
"allow_patterns": list(config.allow_patterns),
|
||||||
|
"deny_patterns": list(config.deny_patterns),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _snapshot_subagent_statuses(
|
||||||
|
manager: SubagentManager,
|
||||||
|
) -> dict[str, dict[str, object]]:
|
||||||
|
return {
|
||||||
|
task_id: _snapshot_subagent_status(status)
|
||||||
|
for task_id, status in manager.runtime_statuses().items()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _snapshot_subagent_status(status: SubagentStatus) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"task_id": status.task_id,
|
||||||
|
"label": status.label,
|
||||||
|
"task_description": status.task_description,
|
||||||
|
"started_at": status.started_at,
|
||||||
|
"phase": status.phase,
|
||||||
|
"iteration": status.iteration,
|
||||||
|
"tool_events": [dict(event) for event in status.tool_events],
|
||||||
|
"usage": dict(status.usage),
|
||||||
|
"stop_reason": status.stop_reason,
|
||||||
|
"error": status.error,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _snapshot_json_mapping(values: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||||
|
return {key: _snapshot_json_value(value) for key, value in values.items()}
|
||||||
|
|
||||||
|
|
||||||
|
def _snapshot_json_value(value: JsonValue) -> JsonValue:
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [_snapshot_json_value(item) for item in value]
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return {
|
||||||
|
key: _snapshot_json_value(item)
|
||||||
|
for key, item in value.items()
|
||||||
|
}
|
||||||
|
return value
|
||||||
@@ -1,76 +0,0 @@
|
|||||||
"""RuntimeState protocol: agent loop state exposed to MyTool."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import TYPE_CHECKING, Any, Protocol
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from nanobot.agent.subagent import SubagentManager
|
|
||||||
from nanobot.agent.tools.shell import ExecToolConfig
|
|
||||||
from nanobot.agent.tools.web import WebToolsConfig
|
|
||||||
from nanobot.utils.llm_runtime import LLMRuntime
|
|
||||||
|
|
||||||
|
|
||||||
class RuntimeState(Protocol):
|
|
||||||
"""Minimum contract that MyTool requires from its runtime state provider.
|
|
||||||
|
|
||||||
In practice, this is always satisfied by ``AgentLoop``. MyTool also
|
|
||||||
accesses arbitrary attributes dynamically (via ``getattr`` / ``setattr``)
|
|
||||||
for dot-path inspection and modification; those paths are validated at
|
|
||||||
runtime rather than by this protocol.
|
|
||||||
"""
|
|
||||||
|
|
||||||
@property
|
|
||||||
def model(self) -> str: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def max_iterations(self) -> int: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def current_iteration(self) -> int: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def tool_names(self) -> list[str]: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def workspace(self) -> Path: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def provider_retry_mode(self) -> str: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def max_tool_result_chars(self) -> int: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def context_window_tokens(self) -> int: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def web_config(self) -> WebToolsConfig: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def exec_config(self) -> ExecToolConfig: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def subagents(self) -> SubagentManager: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _runtime_vars(self) -> dict[str, Any]: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _last_usage(self) -> dict[str, int]: ...
|
|
||||||
|
|
||||||
def _sync_subagent_runtime_limits(self) -> None: ...
|
|
||||||
|
|
||||||
def set_runtime_model(self, model: str) -> LLMRuntime: ...
|
|
||||||
|
|
||||||
def set_runtime_context_window(self, context_window_tokens: int) -> LLMRuntime: ...
|
|
||||||
|
|
||||||
def set_session_model_preset(
|
|
||||||
self,
|
|
||||||
session_key: str,
|
|
||||||
name: str,
|
|
||||||
) -> LLMRuntime: ...
|
|
||||||
|
|
||||||
@property
|
|
||||||
def model_preset(self) -> str | None: ...
|
|
||||||
+197
-166
@@ -1,8 +1,7 @@
|
|||||||
"""MyTool: runtime state inspection and configuration for the agent loop."""
|
"""MyTool: runtime state inspection and configuration for the agent loop."""
|
||||||
|
|
||||||
# RuntimeState intentionally exposes a narrow set of AgentLoop internals to
|
# Tool.execute accepts heterogeneous schemas.
|
||||||
# this manually registered tool. Tool.execute accepts heterogeneous schemas.
|
# pyright: reportIncompatibleMethodOverride=false
|
||||||
# pyright: reportPrivateUsage=false, reportIncompatibleMethodOverride=false
|
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -14,7 +13,13 @@ from loguru import logger
|
|||||||
|
|
||||||
from nanobot.agent.tools.base import Tool, ToolResult
|
from nanobot.agent.tools.base import Tool, ToolResult
|
||||||
from nanobot.agent.tools.context import current_request_context, current_request_session_key
|
from nanobot.agent.tools.context import current_request_context, current_request_session_key
|
||||||
from nanobot.agent.tools.runtime_state import RuntimeState
|
from nanobot.agent.tools.runtime_control import (
|
||||||
|
RUNTIME_COMMAND_KEYS,
|
||||||
|
RUNTIME_SNAPSHOT_KEYS,
|
||||||
|
JsonValue,
|
||||||
|
RuntimeControl,
|
||||||
|
RuntimeSnapshot,
|
||||||
|
)
|
||||||
from nanobot.config_base import Base
|
from nanobot.config_base import Base
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -28,25 +33,28 @@ class MyToolConfig(Base):
|
|||||||
allow_set: bool = False
|
allow_set: bool = False
|
||||||
|
|
||||||
|
|
||||||
def _has_real_attr(obj: Any, key: str) -> bool:
|
|
||||||
"""Check if obj has a real (explicitly set) attribute, not auto-generated by mock."""
|
|
||||||
if isinstance(obj, dict):
|
|
||||||
return key in obj
|
|
||||||
d = getattr(obj, "__dict__", None)
|
|
||||||
if d is not None and key in d:
|
|
||||||
return True
|
|
||||||
for cls in type(obj).__mro__:
|
|
||||||
if key in cls.__dict__:
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def _is_subagent_status(value: object) -> TypeGuard[SubagentStatus]:
|
def _is_subagent_status(value: object) -> TypeGuard[SubagentStatus]:
|
||||||
from nanobot.agent.subagent import SubagentStatus
|
from nanobot.agent.subagent import SubagentStatus
|
||||||
|
|
||||||
return isinstance(value, SubagentStatus)
|
return isinstance(value, SubagentStatus)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_subagent_status_snapshot(value: object) -> TypeGuard[Mapping[str, object]]:
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
return False
|
||||||
|
return all(
|
||||||
|
field in value
|
||||||
|
for field in ("task_id", "label", "task_description", "started_at", "phase")
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_string_mapping(value: object) -> TypeGuard[Mapping[str, object]]:
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
return False
|
||||||
|
mapping = cast(Mapping[object, object], value)
|
||||||
|
return all(isinstance(key, str) for key in mapping)
|
||||||
|
|
||||||
|
|
||||||
class MyTool(Tool):
|
class MyTool(Tool):
|
||||||
"""Check and set the agent loop's runtime configuration."""
|
"""Check and set the agent loop's runtime configuration."""
|
||||||
|
|
||||||
@@ -70,7 +78,7 @@ class MyTool(Tool):
|
|||||||
"runner", "sessions", "consolidator",
|
"runner", "sessions", "consolidator",
|
||||||
"dream", "auto_compact", "context", "commands",
|
"dream", "auto_compact", "context", "commands",
|
||||||
# Sensitive runtime state (credentials, message routing, task tracking)
|
# Sensitive runtime state (credentials, message routing, task tracking)
|
||||||
"_mcp_servers", "_mcp_stacks", "_pending_queues",
|
"_pending_queues",
|
||||||
"_session_locks", "_active_tasks", "_background_tasks",
|
"_session_locks", "_active_tasks", "_background_tasks",
|
||||||
# Security boundaries (inspect + modify both blocked)
|
# Security boundaries (inspect + modify both blocked)
|
||||||
"restrict_to_workspace", "channels_config",
|
"restrict_to_workspace", "channels_config",
|
||||||
@@ -79,7 +87,10 @@ class MyTool(Tool):
|
|||||||
|
|
||||||
READ_ONLY = frozenset({
|
READ_ONLY = frozenset({
|
||||||
"subagents", # observable but replacing it would break the system
|
"subagents", # observable but replacing it would break the system
|
||||||
|
"tool_names",
|
||||||
|
"current_iteration",
|
||||||
"_current_iteration", # updated by runner only
|
"_current_iteration", # updated by runner only
|
||||||
|
"_last_usage",
|
||||||
"exec_config", # inspect allowed (e.g. check sandbox), modify blocked
|
"exec_config", # inspect allowed (e.g. check sandbox), modify blocked
|
||||||
"web_config", # inspect allowed (e.g. check enable), modify blocked
|
"web_config", # inspect allowed (e.g. check enable), modify blocked
|
||||||
"model_presets", # config-derived catalog; changes require config reload
|
"model_presets", # config-derived catalog; changes require config reload
|
||||||
@@ -103,13 +114,6 @@ class MyTool(Tool):
|
|||||||
"private_key", "access_token", "refresh_token", "auth",
|
"private_key", "access_token", "refresh_token", "auth",
|
||||||
})
|
})
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _is_sensitive_field_name(cls, name: str) -> bool:
|
|
||||||
lowered = name.lower()
|
|
||||||
return lowered in cls._SENSITIVE_NAMES or any(
|
|
||||||
part in cls._SENSITIVE_NAMES for part in lowered.split("_")
|
|
||||||
)
|
|
||||||
|
|
||||||
RESTRICTED: dict[str, dict[str, Any]] = {
|
RESTRICTED: dict[str, dict[str, Any]] = {
|
||||||
"max_iterations": {"type": int, "min": 1, "max": 100},
|
"max_iterations": {"type": int, "min": 1, "max": 100},
|
||||||
"context_window_tokens": {"type": int, "min": 4096, "max": 1_000_000},
|
"context_window_tokens": {"type": int, "min": 4096, "max": 1_000_000},
|
||||||
@@ -123,15 +127,15 @@ class MyTool(Tool):
|
|||||||
"context_window_tokens",
|
"context_window_tokens",
|
||||||
})
|
})
|
||||||
|
|
||||||
def __init__(self, runtime_state: RuntimeState, modify_allowed: bool = True) -> None:
|
def __init__(self, runtime_control: RuntimeControl, modify_allowed: bool = True) -> None:
|
||||||
self._runtime_state = runtime_state
|
self._runtime_control = runtime_control
|
||||||
self._modify_allowed = modify_allowed
|
self._modify_allowed = modify_allowed
|
||||||
|
|
||||||
def __deepcopy__(self, memo: dict[int, Any]) -> MyTool:
|
def __deepcopy__(self, memo: dict[int, Any]) -> MyTool:
|
||||||
cls = self.__class__
|
cls = self.__class__
|
||||||
result = cls.__new__(cls)
|
result = cls.__new__(cls)
|
||||||
memo[id(self)] = result
|
memo[id(self)] = result
|
||||||
result._runtime_state = self._runtime_state
|
result._runtime_control = self._runtime_control
|
||||||
result._modify_allowed = self._modify_allowed
|
result._modify_allowed = self._modify_allowed
|
||||||
return result
|
return result
|
||||||
|
|
||||||
@@ -208,9 +212,12 @@ class MyTool(Tool):
|
|||||||
# Path resolution
|
# Path resolution
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
def _resolve_path(self, path: str) -> tuple[Any, str | None]:
|
def _resolve_path(
|
||||||
|
self,
|
||||||
|
snapshot: RuntimeSnapshot,
|
||||||
|
path: str,
|
||||||
|
) -> tuple[object | None, str | None]:
|
||||||
parts = path.split(".")
|
parts = path.split(".")
|
||||||
obj: Any = self._runtime_state
|
|
||||||
for part in parts:
|
for part in parts:
|
||||||
if part in self._DENIED_ATTRS or part.startswith("__"):
|
if part in self._DENIED_ATTRS or part.startswith("__"):
|
||||||
return None, f"'{part}' is not accessible"
|
return None, f"'{part}' is not accessible"
|
||||||
@@ -218,17 +225,13 @@ class MyTool(Tool):
|
|||||||
return None, f"'{part}' is not accessible"
|
return None, f"'{part}' is not accessible"
|
||||||
if part.lower() in self._SENSITIVE_NAMES:
|
if part.lower() in self._SENSITIVE_NAMES:
|
||||||
return None, f"'{part}' is not accessible"
|
return None, f"'{part}' is not accessible"
|
||||||
try:
|
obj: object = snapshot.as_mapping()
|
||||||
if isinstance(obj, Mapping):
|
for part in parts:
|
||||||
mapping = cast(Mapping[str, Any], obj)
|
if not _is_string_mapping(obj):
|
||||||
if part in mapping:
|
return None, f"'{part}' not found"
|
||||||
obj = mapping[part]
|
if part not in obj:
|
||||||
else:
|
|
||||||
return None, f"'{part}' not found in mapping"
|
return None, f"'{part}' not found in mapping"
|
||||||
else:
|
obj = obj[part]
|
||||||
obj = getattr(obj, part)
|
|
||||||
except (KeyError, AttributeError) as e:
|
|
||||||
return None, f"'{part}' not found: {e}"
|
|
||||||
return obj, None
|
return obj, None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -242,20 +245,48 @@ class MyTool(Tool):
|
|||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _format_status(st: "SubagentStatus", indent: str = " ") -> str:
|
def _format_status(
|
||||||
elapsed = time.monotonic() - st.started_at
|
st: "SubagentStatus | Mapping[str, object]",
|
||||||
tool_summary = ", ".join(
|
indent: str = " ",
|
||||||
f"{e.get('name', '?')}({e.get('status', '?')})" for e in st.tool_events[-5:]
|
) -> str:
|
||||||
) or "none"
|
if isinstance(st, Mapping):
|
||||||
|
started_at = st.get("started_at", time.monotonic())
|
||||||
|
raw_events = st.get("tool_events", [])
|
||||||
|
phase = st.get("phase", "unknown")
|
||||||
|
iteration = st.get("iteration", 0)
|
||||||
|
usage = st.get("usage", {})
|
||||||
|
error = st.get("error")
|
||||||
|
stop_reason = st.get("stop_reason")
|
||||||
|
else:
|
||||||
|
started_at = st.started_at
|
||||||
|
raw_events = st.tool_events
|
||||||
|
phase = st.phase
|
||||||
|
iteration = st.iteration
|
||||||
|
usage = st.usage
|
||||||
|
error = st.error
|
||||||
|
stop_reason = st.stop_reason
|
||||||
|
elapsed = time.monotonic() - (
|
||||||
|
float(started_at) if isinstance(started_at, (int, float)) else time.monotonic()
|
||||||
|
)
|
||||||
|
tool_events = cast(list[object], raw_events) if isinstance(raw_events, list) else []
|
||||||
|
tool_summaries: list[str] = []
|
||||||
|
for raw_event in tool_events[-5:]:
|
||||||
|
if not isinstance(raw_event, Mapping):
|
||||||
|
continue
|
||||||
|
event = cast(Mapping[str, object], raw_event)
|
||||||
|
tool_summaries.append(
|
||||||
|
f"{event.get('name', '?')}({event.get('status', '?')})"
|
||||||
|
)
|
||||||
|
tool_summary = ", ".join(tool_summaries) or "none"
|
||||||
lines = [
|
lines = [
|
||||||
f"{indent}phase: {st.phase}, iteration: {st.iteration}, elapsed: {elapsed:.1f}s",
|
f"{indent}phase: {phase}, iteration: {iteration}, elapsed: {elapsed:.1f}s",
|
||||||
f"{indent}tools: {tool_summary}",
|
f"{indent}tools: {tool_summary}",
|
||||||
f"{indent}usage: {st.usage or 'n/a'}",
|
f"{indent}usage: {usage or 'n/a'}",
|
||||||
]
|
]
|
||||||
if st.error:
|
if error:
|
||||||
lines.append(f"{indent}error: {st.error}")
|
lines.append(f"{indent}error: {error}")
|
||||||
if st.stop_reason:
|
if stop_reason:
|
||||||
lines.append(f"{indent}stop_reason: {st.stop_reason}")
|
lines.append(f"{indent}stop_reason: {stop_reason}")
|
||||||
return "\n".join(lines)
|
return "\n".join(lines)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -264,29 +295,38 @@ class MyTool(Tool):
|
|||||||
header = f"Subagent [{val.task_id}] '{val.label}'"
|
header = f"Subagent [{val.task_id}] '{val.label}'"
|
||||||
detail = MyTool._format_status(val, " ")
|
detail = MyTool._format_status(val, " ")
|
||||||
return f"{header}\n task: {val.task_description}\n{detail}"
|
return f"{header}\n task: {val.task_description}\n{detail}"
|
||||||
# SubagentManager: delegate to its _task_statuses dict
|
if _is_subagent_status_snapshot(val):
|
||||||
task_statuses = getattr(val, "_task_statuses", None)
|
header = f"Subagent [{val['task_id']}] '{val['label']}'"
|
||||||
if isinstance(task_statuses, dict):
|
detail = MyTool._format_status(val, " ")
|
||||||
return MyTool._format_value(task_statuses, key)
|
return f"{header}\n task: {val['task_description']}\n{detail}"
|
||||||
if isinstance(val, Mapping):
|
if isinstance(val, Mapping):
|
||||||
mapping = cast(Mapping[object, object], val)
|
mapping = cast(Mapping[object, object], val)
|
||||||
else:
|
else:
|
||||||
mapping = None
|
mapping = None
|
||||||
|
if mapping and set(mapping) == {"_task_statuses"}:
|
||||||
|
task_statuses = mapping["_task_statuses"]
|
||||||
|
if isinstance(task_statuses, Mapping):
|
||||||
|
return MyTool._format_value(task_statuses, key)
|
||||||
if (
|
if (
|
||||||
mapping
|
mapping
|
||||||
and _is_subagent_status(next(iter(mapping.values())))
|
and (
|
||||||
|
_is_subagent_status(next(iter(mapping.values())))
|
||||||
|
or _is_subagent_status_snapshot(next(iter(mapping.values())))
|
||||||
|
)
|
||||||
):
|
):
|
||||||
status_mapping: Mapping[object, SubagentStatus] = cast(Any, mapping)
|
|
||||||
prefix = f"{key}: " if key else ""
|
prefix = f"{key}: " if key else ""
|
||||||
lines = [f"{prefix}{len(status_mapping)} subagent(s):"]
|
lines = [f"{prefix}{len(mapping)} subagent(s):"]
|
||||||
for tid, st in status_mapping.items():
|
for tid, st in mapping.items():
|
||||||
|
if _is_subagent_status(st):
|
||||||
detail = MyTool._format_status(st, " ")
|
detail = MyTool._format_status(st, " ")
|
||||||
lines.append(f" [{tid}] '{st.label}'\n{detail}")
|
label = st.label
|
||||||
|
elif _is_subagent_status_snapshot(st):
|
||||||
|
detail = MyTool._format_status(st, " ")
|
||||||
|
label = st.get("label", "?")
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
lines.append(f" [{tid}] '{label}'\n{detail}")
|
||||||
return "\n".join(lines)
|
return "\n".join(lines)
|
||||||
dynamic_value = cast(Any, val)
|
|
||||||
if hasattr(dynamic_value, "tool_names"):
|
|
||||||
tool_names: Any = getattr(dynamic_value, "tool_names")
|
|
||||||
return f"tools: {len(tool_names)} registered — {tool_names}"
|
|
||||||
# Scalar types — repr is fine
|
# Scalar types — repr is fine
|
||||||
if isinstance(val, (str, int, float, bool, type(None))):
|
if isinstance(val, (str, int, float, bool, type(None))):
|
||||||
r = repr(val)
|
r = repr(val)
|
||||||
@@ -311,32 +351,6 @@ class MyTool(Tool):
|
|||||||
return f"{key}: [{len(sequence)} items]" if key else f"[{len(sequence)} items]"
|
return f"{key}: [{len(sequence)} items]" if key else f"[{len(sequence)} items]"
|
||||||
r = repr(sequence)
|
r = repr(sequence)
|
||||||
return f"{key}: {r}" if key else r
|
return f"{key}: {r}" if key else r
|
||||||
# Complex object — small Pydantic models: show values; others: show field names for navigation
|
|
||||||
value_type = type(cast(object, val))
|
|
||||||
cls_name = value_type.__name__
|
|
||||||
model_fields = cast(object, getattr(value_type, "model_fields", None))
|
|
||||||
if isinstance(model_fields, Mapping) and model_fields:
|
|
||||||
fields = list(cast(Mapping[str, object], model_fields).keys())
|
|
||||||
if len(fields) <= 8:
|
|
||||||
# Small config objects: show field=value pairs
|
|
||||||
pairs: list[str] = []
|
|
||||||
for f in fields:
|
|
||||||
fv = getattr(val, f, "?")
|
|
||||||
if MyTool._is_sensitive_field_name(f):
|
|
||||||
continue
|
|
||||||
if isinstance(fv, (str, int, float, bool, type(None))):
|
|
||||||
pairs.append(f"{f}={fv!r}")
|
|
||||||
else:
|
|
||||||
pairs.append(f"{f}=<{type(fv).__name__}>")
|
|
||||||
preview = ", ".join(pairs)
|
|
||||||
return f"{key}: {preview}" if key else preview
|
|
||||||
else:
|
|
||||||
attributes = cast(dict[str, Any], getattr(val, "__dict__", {}))
|
|
||||||
fields = [name for name in attributes if not name.startswith("__")]
|
|
||||||
if fields:
|
|
||||||
preview = ", ".join(str(f) for f in fields[:20])
|
|
||||||
suffix = ", ..." if len(fields) > 20 else ""
|
|
||||||
return f"{key}: <{cls_name}> [{preview}{suffix}]" if key else f"<{cls_name}> [{preview}{suffix}]"
|
|
||||||
r = repr(val)
|
r = repr(val)
|
||||||
return f"{key}: {r}" if key else r
|
return f"{key}: {r}" if key else r
|
||||||
|
|
||||||
@@ -366,7 +380,12 @@ class MyTool(Tool):
|
|||||||
runtime = request_ctx.runtime if request_ctx is not None else None
|
runtime = request_ctx.runtime if request_ctx is not None else None
|
||||||
if runtime is None or key not in self._MODEL_RUNTIME_FIELDS:
|
if runtime is None or key not in self._MODEL_RUNTIME_FIELDS:
|
||||||
return False, None
|
return False, None
|
||||||
return True, getattr(runtime, key)
|
values: dict[str, object] = {
|
||||||
|
"model": runtime.model,
|
||||||
|
"model_preset": runtime.model_preset,
|
||||||
|
"context_window_tokens": runtime.context_window_tokens,
|
||||||
|
}
|
||||||
|
return True, values[key]
|
||||||
|
|
||||||
def _inspect(self, key: str | None) -> str:
|
def _inspect(self, key: str | None) -> str:
|
||||||
if not key:
|
if not key:
|
||||||
@@ -375,62 +394,64 @@ class MyTool(Tool):
|
|||||||
request_ctx = current_request_context()
|
request_ctx = current_request_context()
|
||||||
if request_ctx is None:
|
if request_ctx is None:
|
||||||
return ToolResult.error("Error: current request context is unavailable")
|
return ToolResult.error("Error: current request context is unavailable")
|
||||||
|
request_values: dict[str, str | None] = {
|
||||||
|
"channel": request_ctx.channel,
|
||||||
|
"chat_id": request_ctx.chat_id,
|
||||||
|
"sender_id": request_ctx.sender_id,
|
||||||
|
}
|
||||||
if key == "request":
|
if key == "request":
|
||||||
return self._format_value(
|
return self._format_value(request_values, key)
|
||||||
{field: getattr(request_ctx, field) for field in self._REQUEST_FIELDS},
|
|
||||||
key,
|
|
||||||
)
|
|
||||||
field = key.removeprefix("request.")
|
field = key.removeprefix("request.")
|
||||||
if field not in self._REQUEST_FIELDS:
|
if field not in self._REQUEST_FIELDS:
|
||||||
return ToolResult.error(f"Error: '{key}' not found")
|
return ToolResult.error(f"Error: '{key}' not found")
|
||||||
return self._format_value(getattr(request_ctx, field), key)
|
return self._format_value(request_values[field], key)
|
||||||
if "." not in key:
|
if "." not in key:
|
||||||
found, value = self._current_runtime_value(key)
|
found, value = self._current_runtime_value(key)
|
||||||
if found:
|
if found:
|
||||||
return self._format_value(value, key)
|
return self._format_value(value, key)
|
||||||
|
snapshot = self._runtime_control.snapshot()
|
||||||
top = key.split(".")[0]
|
top = key.split(".")[0]
|
||||||
if top in self._DENIED_ATTRS or top.startswith("__"):
|
if top in self._DENIED_ATTRS or top.startswith("__"):
|
||||||
return ToolResult.error(f"Error: '{top}' is not accessible")
|
return ToolResult.error(f"Error: '{top}' is not accessible")
|
||||||
obj, err = self._resolve_path(key)
|
obj, err = self._resolve_path(snapshot, key)
|
||||||
if err:
|
if err:
|
||||||
# "scratchpad" alias for _runtime_vars
|
|
||||||
if key == "scratchpad":
|
if key == "scratchpad":
|
||||||
rv = self._runtime_state._runtime_vars
|
return (
|
||||||
return self._format_value(rv, "scratchpad") if rv else "scratchpad is empty"
|
self._format_value(snapshot.scratchpad, "scratchpad")
|
||||||
# Fallback: check _runtime_vars for simple keys stored by modify
|
if snapshot.scratchpad
|
||||||
if "." not in key and key in self._runtime_state._runtime_vars:
|
else "scratchpad is empty"
|
||||||
return self._format_value(self._runtime_state._runtime_vars[key], key)
|
)
|
||||||
|
if "." not in key and key in snapshot.scratchpad:
|
||||||
|
return self._format_value(snapshot.scratchpad[key], key)
|
||||||
return ToolResult.error(f"Error: {err}")
|
return ToolResult.error(f"Error: {err}")
|
||||||
# Guard against mock auto-generated attributes
|
|
||||||
if "." not in key and not _has_real_attr(self._runtime_state, key):
|
|
||||||
if key in self._runtime_state._runtime_vars:
|
|
||||||
return self._format_value(self._runtime_state._runtime_vars[key], key)
|
|
||||||
return ToolResult.error(f"Error: '{key}' not found")
|
|
||||||
return self._format_value(obj, key)
|
return self._format_value(obj, key)
|
||||||
|
|
||||||
def _inspect_all(self) -> str:
|
def _inspect_all(self) -> str:
|
||||||
state = self._runtime_state
|
snapshot = self._runtime_control.snapshot()
|
||||||
|
values = snapshot.as_mapping()
|
||||||
parts: list[str] = []
|
parts: list[str] = []
|
||||||
# RESTRICTED keys
|
|
||||||
for k in self.RESTRICTED:
|
for k in self.RESTRICTED:
|
||||||
found, value = self._current_runtime_value(k)
|
found, value = self._current_runtime_value(k)
|
||||||
parts.append(self._format_value(value if found else getattr(state, k, None), k))
|
parts.append(self._format_value(value if found else values[k], k))
|
||||||
found, value = self._current_runtime_value("model_preset")
|
found, value = self._current_runtime_value("model_preset")
|
||||||
parts.append(self._format_value(
|
parts.append(self._format_value(
|
||||||
value if found else state.model_preset,
|
value if found else snapshot.model_preset,
|
||||||
"model_preset",
|
"model_preset",
|
||||||
))
|
))
|
||||||
# Other useful top-level keys shown in description
|
for k in (
|
||||||
for k in ("workspace", "provider_retry_mode", "max_tool_result_chars", "_current_iteration", "web_config", "exec_config", "workspace_sandbox", "subagents"):
|
"workspace",
|
||||||
if _has_real_attr(state, k):
|
"provider_retry_mode",
|
||||||
parts.append(self._format_value(getattr(state, k, None), k))
|
"max_tool_result_chars",
|
||||||
# Token usage
|
"_current_iteration",
|
||||||
usage = state._last_usage
|
"web_config",
|
||||||
if usage:
|
"exec_config",
|
||||||
parts.append(self._format_value(usage, "_last_usage"))
|
"subagents",
|
||||||
rv = state._runtime_vars
|
):
|
||||||
if rv:
|
parts.append(self._format_value(values[k], k))
|
||||||
parts.append(self._format_value(rv, "scratchpad"))
|
if snapshot.last_usage:
|
||||||
|
parts.append(self._format_value(snapshot.last_usage, "_last_usage"))
|
||||||
|
if snapshot.scratchpad:
|
||||||
|
parts.append(self._format_value(snapshot.scratchpad, "scratchpad"))
|
||||||
return "\n".join(parts)
|
return "\n".join(parts)
|
||||||
|
|
||||||
# -- modify --
|
# -- modify --
|
||||||
@@ -454,48 +475,49 @@ class MyTool(Tool):
|
|||||||
if leaf.lower() in self._SENSITIVE_NAMES:
|
if leaf.lower() in self._SENSITIVE_NAMES:
|
||||||
self._audit("modify", f"BLOCKED sensitive leaf '{leaf}'")
|
self._audit("modify", f"BLOCKED sensitive leaf '{leaf}'")
|
||||||
return ToolResult.error(f"Error: '{leaf}' is not accessible")
|
return ToolResult.error(f"Error: '{leaf}' is not accessible")
|
||||||
parent, err = self._resolve_path(parent_path)
|
snapshot = self._runtime_control.snapshot()
|
||||||
|
_parent, err = self._resolve_path(snapshot, parent_path)
|
||||||
if err:
|
if err:
|
||||||
return ToolResult.error(f"Error: {err}")
|
return ToolResult.error(f"Error: {err}")
|
||||||
if isinstance(parent, dict):
|
self._audit("modify", f"READ_ONLY {key}")
|
||||||
parent[leaf] = value
|
return ToolResult.error(f"Error: '{key}' is read-only and cannot be modified")
|
||||||
else:
|
|
||||||
setattr(parent, leaf, value)
|
|
||||||
self._audit("modify", f"{key} = {value!r}")
|
|
||||||
return f"Set {key} = {value!r}"
|
|
||||||
if key == "model_preset":
|
if key == "model_preset":
|
||||||
return self._modify_model_preset(value)
|
return self._modify_model_preset(value)
|
||||||
if key in self.RESTRICTED:
|
if key in self.RESTRICTED:
|
||||||
return self._modify_restricted(key, value)
|
return self._modify_restricted(key, value)
|
||||||
return self._modify_free(key, value)
|
if key in RUNTIME_COMMAND_KEYS:
|
||||||
|
return self._modify_runtime_setting(key, value)
|
||||||
|
if key in RUNTIME_SNAPSHOT_KEYS:
|
||||||
|
self._audit("modify", f"READ_ONLY {key}")
|
||||||
|
return ToolResult.error(f"Error: '{key}' is read-only and cannot be modified")
|
||||||
|
return self._modify_scratchpad(key, value)
|
||||||
|
|
||||||
def _modify_model_preset(self, value: Any) -> str:
|
def _modify_model_preset(self, value: Any) -> str:
|
||||||
if not isinstance(value, str) or not value.strip():
|
if not isinstance(value, str) or not value.strip():
|
||||||
return ToolResult.error("Error: 'model_preset' must be a non-empty string")
|
return ToolResult.error("Error: 'model_preset' must be a non-empty string")
|
||||||
name = value.strip()
|
name = value.strip()
|
||||||
session_key = current_request_session_key()
|
session_key = current_request_session_key()
|
||||||
if session_key:
|
old = self._runtime_control.snapshot().model_preset
|
||||||
try:
|
try:
|
||||||
runtime = self._runtime_state.set_session_model_preset(
|
runtime = self._runtime_control.set_model_preset(
|
||||||
session_key,
|
|
||||||
name,
|
name,
|
||||||
|
session_key=session_key,
|
||||||
)
|
)
|
||||||
except (KeyError, ValueError) as exc:
|
except (KeyError, ValueError) as exc:
|
||||||
message = str(exc.args[0]) if exc.args else str(exc)
|
message = str(exc.args[0]) if exc.args else str(exc)
|
||||||
punctuation = "" if message.endswith((".", "!", "?")) else "."
|
punctuation = "" if message.endswith((".", "!", "?")) else "."
|
||||||
return ToolResult.error(f"Error: {message}{punctuation}")
|
return ToolResult.error(f"Error: {message}{punctuation}")
|
||||||
|
if session_key:
|
||||||
self._audit("modify", f"model_preset = {name!r}")
|
self._audit("modify", f"model_preset = {name!r}")
|
||||||
return (
|
return (
|
||||||
f"Set model_preset = {name!r} for the next turn; "
|
f"Set model_preset = {name!r} for the next turn; "
|
||||||
f"model will be {runtime.model!r}; "
|
f"model will be {runtime.model!r}; "
|
||||||
f"context_window_tokens will be {runtime.context_window_tokens!r}"
|
f"context_window_tokens will be {runtime.context_window_tokens!r}"
|
||||||
)
|
)
|
||||||
result = self._modify_free("model_preset", name)
|
self._audit("modify", f"model_preset: {old!r} -> {name!r}")
|
||||||
if isinstance(result, ToolResult) and result.is_error:
|
|
||||||
return result if result.endswith((".", "!", "?")) else ToolResult.error(f"{result}.")
|
|
||||||
return (
|
return (
|
||||||
f"{result}; model is now {self._runtime_state.model!r}; "
|
f"Set model_preset = {name!r} (was {old!r}); model is now {runtime.model!r}; "
|
||||||
f"context_window_tokens is now {self._runtime_state.context_window_tokens!r}"
|
f"context_window_tokens is now {runtime.context_window_tokens!r}"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _modify_restricted(self, key: str, value: Any) -> str:
|
def _modify_restricted(self, key: str, value: Any) -> str:
|
||||||
@@ -508,7 +530,7 @@ class MyTool(Tool):
|
|||||||
value = expected(value)
|
value = expected(value)
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
return ToolResult.error(f"Error: '{key}' must be {expected.__name__}, got {type(value).__name__}")
|
return ToolResult.error(f"Error: '{key}' must be {expected.__name__}, got {type(value).__name__}")
|
||||||
old = getattr(self._runtime_state, key)
|
old = self._runtime_control.snapshot().as_mapping()[key]
|
||||||
if "min" in spec and value < spec["min"]:
|
if "min" in spec and value < spec["min"]:
|
||||||
return ToolResult.error(f"Error: '{key}' must be >= {spec['min']}")
|
return ToolResult.error(f"Error: '{key}' must be >= {spec['min']}")
|
||||||
if "max" in spec and value > spec["max"]:
|
if "max" in spec and value > spec["max"]:
|
||||||
@@ -521,41 +543,46 @@ class MyTool(Tool):
|
|||||||
"during an active session; use a configured model_preset"
|
"during an active session; use a configured model_preset"
|
||||||
)
|
)
|
||||||
if key == "model":
|
if key == "model":
|
||||||
self._runtime_state.set_runtime_model(cast(str, value))
|
self._runtime_control.set_model(cast(str, value))
|
||||||
elif key == "context_window_tokens":
|
elif key == "context_window_tokens":
|
||||||
self._runtime_state.set_runtime_context_window(cast(int, value))
|
self._runtime_control.set_context_window_tokens(cast(int, value))
|
||||||
else:
|
else:
|
||||||
setattr(self._runtime_state, key, value)
|
self._runtime_control.set_max_iterations(cast(int, value))
|
||||||
if key == "max_iterations" and hasattr(
|
|
||||||
self._runtime_state,
|
|
||||||
"_sync_subagent_runtime_limits",
|
|
||||||
):
|
|
||||||
self._runtime_state._sync_subagent_runtime_limits()
|
|
||||||
self._audit("modify", f"{key}: {old!r} -> {value!r}")
|
self._audit("modify", f"{key}: {old!r} -> {value!r}")
|
||||||
return f"Set {key} = {value!r} (was {old!r})"
|
return f"Set {key} = {value!r} (was {old!r})"
|
||||||
|
|
||||||
def _modify_free(self, key: str, value: Any) -> str:
|
def _modify_runtime_setting(self, key: str, value: Any) -> str:
|
||||||
if _has_real_attr(self._runtime_state, key):
|
old = self._runtime_control.snapshot().as_mapping()[key]
|
||||||
old = getattr(self._runtime_state, key)
|
if key == "workspace":
|
||||||
if isinstance(old, (str, int, float, bool)):
|
if not isinstance(value, str):
|
||||||
old_t: type[Any] = type(old)
|
return ToolResult.error(
|
||||||
|
f"Error: 'workspace' expects str, got {type(value).__name__}"
|
||||||
|
)
|
||||||
|
self._runtime_control.set_workspace_display(value)
|
||||||
|
self._audit("modify", f"workspace: {old!r} -> {value!r}")
|
||||||
|
return f"Set workspace = {value!r} (was {old!r})"
|
||||||
|
old_t = type(old)
|
||||||
new_t = cast(type[Any], type(value))
|
new_t = cast(type[Any], type(value))
|
||||||
if old_t is float and new_t is int:
|
if old_t is float and new_t is int:
|
||||||
pass # int → float coercion allowed
|
pass
|
||||||
elif old_t is not new_t:
|
elif old_t is not new_t:
|
||||||
self._audit(
|
self._audit(
|
||||||
"modify",
|
"modify",
|
||||||
f"REJECTED type mismatch {key}: expects {old_t.__name__}, got {new_t.__name__}",
|
f"REJECTED type mismatch {key}: expects {old_t.__name__}, got {new_t.__name__}",
|
||||||
)
|
)
|
||||||
return ToolResult.error(f"Error: '{key}' expects {old_t.__name__}, got {new_t.__name__}")
|
return ToolResult.error(
|
||||||
try:
|
f"Error: '{key}' expects {old_t.__name__}, got {new_t.__name__}"
|
||||||
setattr(self._runtime_state, key, value)
|
)
|
||||||
except (ValueError, KeyError) as e:
|
if key == "provider_retry_mode":
|
||||||
message = str(e.args[0] if isinstance(e, KeyError) and e.args else e).strip('"')
|
self._runtime_control.set_provider_retry_mode(cast(str, value))
|
||||||
self._audit("modify", f"REJECTED {key}: {message}")
|
elif key == "max_tool_result_chars":
|
||||||
return ToolResult.error(f"Error: {message}")
|
self._runtime_control.set_max_tool_result_chars(cast(int, value))
|
||||||
|
else:
|
||||||
|
raise AssertionError(f"Unhandled runtime command: {key}")
|
||||||
self._audit("modify", f"{key}: {old!r} -> {value!r}")
|
self._audit("modify", f"{key}: {old!r} -> {value!r}")
|
||||||
return f"Set {key} = {value!r} (was {old!r})"
|
return f"Set {key} = {value!r} (was {old!r})"
|
||||||
|
|
||||||
|
def _modify_scratchpad(self, key: str, value: Any) -> str:
|
||||||
if callable(value):
|
if callable(value):
|
||||||
self._audit("modify", f"REJECTED callable {key}")
|
self._audit("modify", f"REJECTED callable {key}")
|
||||||
return ToolResult.error("Error: cannot store callable values")
|
return ToolResult.error("Error: cannot store callable values")
|
||||||
@@ -563,12 +590,16 @@ class MyTool(Tool):
|
|||||||
if err:
|
if err:
|
||||||
self._audit("modify", f"REJECTED {key}: {err}")
|
self._audit("modify", f"REJECTED {key}: {err}")
|
||||||
return ToolResult.error(f"Error: {err}")
|
return ToolResult.error(f"Error: {err}")
|
||||||
if key not in self._runtime_state._runtime_vars and len(self._runtime_state._runtime_vars) >= self._MAX_RUNTIME_KEYS:
|
try:
|
||||||
|
self._runtime_control.set_scratchpad(
|
||||||
|
key,
|
||||||
|
cast(JsonValue, value),
|
||||||
|
max_keys=self._MAX_RUNTIME_KEYS,
|
||||||
|
)
|
||||||
|
except ValueError as exc:
|
||||||
self._audit("modify", f"REJECTED {key}: max keys ({self._MAX_RUNTIME_KEYS}) reached")
|
self._audit("modify", f"REJECTED {key}: max keys ({self._MAX_RUNTIME_KEYS}) reached")
|
||||||
return ToolResult.error(f"Error: scratchpad is full (max {self._MAX_RUNTIME_KEYS} keys). Remove unused keys first.")
|
return ToolResult.error(f"Error: {exc}. Remove unused keys first.")
|
||||||
old = self._runtime_state._runtime_vars.get(key)
|
self._audit("modify", f"scratchpad.{key} = {value!r}")
|
||||||
self._runtime_state._runtime_vars[key] = value
|
|
||||||
self._audit("modify", f"scratchpad.{key}: {old!r} -> {value!r}")
|
|
||||||
return f"Set scratchpad.{key} = {value!r}"
|
return f"Set scratchpad.{key} = {value!r}"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
+218
-14
@@ -5,6 +5,7 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
import shlex
|
||||||
import shutil
|
import shutil
|
||||||
import signal
|
import signal
|
||||||
import subprocess
|
import subprocess
|
||||||
@@ -12,7 +13,8 @@ import sys
|
|||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path, PureWindowsPath
|
from pathlib import Path, PureWindowsPath
|
||||||
from typing import Any
|
from typing import Any, Protocol, cast
|
||||||
|
from urllib.parse import unquote
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
@@ -42,6 +44,17 @@ from nanobot.security.workspace_access import current_scope_allows_loopback, cur
|
|||||||
from nanobot.security.workspace_policy import is_path_within
|
from nanobot.security.workspace_policy import is_path_within
|
||||||
|
|
||||||
_IS_WINDOWS = sys.platform == "win32"
|
_IS_WINDOWS = sys.platform == "win32"
|
||||||
|
_PROCESS_TREE_OWNER_ATTR = "_nanobot_process_tree_owner"
|
||||||
|
|
||||||
|
|
||||||
|
class _ProcessTreeOwner(Protocol):
|
||||||
|
creation_flags: int
|
||||||
|
|
||||||
|
def assign_and_resume(self, pid: int) -> None: ...
|
||||||
|
|
||||||
|
def release(self) -> None: ...
|
||||||
|
|
||||||
|
def terminate(self) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
def _reap_pid(pid: int) -> None:
|
def _reap_pid(pid: int) -> None:
|
||||||
@@ -326,6 +339,7 @@ class ExecTool(Tool):
|
|||||||
prepared.env,
|
prepared.env,
|
||||||
prepared.shell_program,
|
prepared.shell_program,
|
||||||
prepared.login,
|
prepared.login,
|
||||||
|
process_tree=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -334,10 +348,10 @@ class ExecTool(Tool):
|
|||||||
timeout=prepared.timeout,
|
timeout=prepared.timeout,
|
||||||
)
|
)
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
await self._kill_process(process)
|
await self._kill_process_tree(process)
|
||||||
return ToolResult.error(f"Error: Command timed out after {prepared.timeout} seconds")
|
return ToolResult.error(f"Error: Command timed out after {prepared.timeout} seconds")
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
await self._kill_process(process)
|
await self._kill_process_tree(process)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
# Safety-net reap: asyncio *should* have reaped the child via
|
# Safety-net reap: asyncio *should* have reaped the child via
|
||||||
@@ -368,13 +382,14 @@ class ExecTool(Tool):
|
|||||||
+ result[-half:]
|
+ result[-half:]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self._release_process_tree(process)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# Kill and reap the child if it was spawned but an unexpected
|
# Kill and reap the child if it was spawned but an unexpected
|
||||||
# error prevented communicate() from completing.
|
# error prevented communicate() from completing.
|
||||||
if process is not None:
|
if process is not None:
|
||||||
await self._kill_process(process)
|
await self._kill_process_tree(process)
|
||||||
return ToolResult.error(f"Error executing command: {str(e)}")
|
return ToolResult.error(f"Error executing command: {str(e)}")
|
||||||
|
|
||||||
async def _execute_session(
|
async def _execute_session(
|
||||||
@@ -537,22 +552,31 @@ class ExecTool(Tool):
|
|||||||
) -> asyncio.subprocess.Process:
|
) -> asyncio.subprocess.Process:
|
||||||
"""Launch *command* in a platform-appropriate shell."""
|
"""Launch *command* in a platform-appropriate shell."""
|
||||||
if _IS_WINDOWS:
|
if _IS_WINDOWS:
|
||||||
|
windows_job = None
|
||||||
|
process = None
|
||||||
|
creation_flags = 0
|
||||||
|
if process_tree and sys.platform == "win32":
|
||||||
|
windows_job = ExecTool._create_windows_job()
|
||||||
|
creation_flags = windows_job.creation_flags
|
||||||
# Default to PowerShell so single-line and multi-line commands
|
# Default to PowerShell so single-line and multi-line commands
|
||||||
# share the same shell semantics. cmd.exe is reachable via the
|
# share the same shell semantics. cmd.exe is reachable via the
|
||||||
# explicit shell="cmd" parameter (see _resolve_shell).
|
# explicit shell="cmd" parameter (see _resolve_shell).
|
||||||
default_program = shutil.which("pwsh") or shutil.which("powershell") or "powershell"
|
default_program = shutil.which("pwsh") or shutil.which("powershell") or "powershell"
|
||||||
program = shell_program or default_program
|
program = shell_program or default_program
|
||||||
program_name = PureWindowsPath(program).name.lower()
|
program_name = PureWindowsPath(program).name.lower()
|
||||||
|
try:
|
||||||
if program_name in ("cmd", "cmd.exe"):
|
if program_name in ("cmd", "cmd.exe"):
|
||||||
cmd_env = {**env, "COMSPEC": program}
|
cmd_env = {**env, "COMSPEC": program}
|
||||||
return await asyncio.create_subprocess_shell(
|
process = await asyncio.create_subprocess_shell(
|
||||||
command,
|
command,
|
||||||
stdin=stdin,
|
stdin=stdin,
|
||||||
stdout=asyncio.subprocess.PIPE,
|
stdout=asyncio.subprocess.PIPE,
|
||||||
stderr=asyncio.subprocess.PIPE,
|
stderr=asyncio.subprocess.PIPE,
|
||||||
cwd=cwd,
|
cwd=cwd,
|
||||||
env=cmd_env,
|
env=cmd_env,
|
||||||
|
creationflags=creation_flags,
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
command = ExecTool._normalize_powershell_command(command)
|
command = ExecTool._normalize_powershell_command(command)
|
||||||
command = (
|
command = (
|
||||||
"[Console]::OutputEncoding = [System.Text.UTF8Encoding]::new($false)\n"
|
"[Console]::OutputEncoding = [System.Text.UTF8Encoding]::new($false)\n"
|
||||||
@@ -561,14 +585,25 @@ class ExecTool(Tool):
|
|||||||
f"{command}\n"
|
f"{command}\n"
|
||||||
"if ($LASTEXITCODE -ne $null) { exit $LASTEXITCODE }"
|
"if ($LASTEXITCODE -ne $null) { exit $LASTEXITCODE }"
|
||||||
)
|
)
|
||||||
return await asyncio.create_subprocess_exec(
|
process = await asyncio.create_subprocess_exec(
|
||||||
program, "-NoProfile", "-NonInteractive", "-Command", command,
|
program, "-NoProfile", "-NonInteractive", "-Command", command,
|
||||||
stdin=stdin,
|
stdin=stdin,
|
||||||
stdout=asyncio.subprocess.PIPE,
|
stdout=asyncio.subprocess.PIPE,
|
||||||
stderr=asyncio.subprocess.PIPE,
|
stderr=asyncio.subprocess.PIPE,
|
||||||
cwd=cwd,
|
cwd=cwd,
|
||||||
env=env,
|
env=env,
|
||||||
|
creationflags=creation_flags,
|
||||||
)
|
)
|
||||||
|
if windows_job is not None:
|
||||||
|
windows_job.assign_and_resume(process.pid)
|
||||||
|
setattr(process, _PROCESS_TREE_OWNER_ATTR, windows_job)
|
||||||
|
return process
|
||||||
|
except BaseException:
|
||||||
|
if windows_job is not None:
|
||||||
|
windows_job.terminate()
|
||||||
|
if process is not None:
|
||||||
|
await ExecTool._kill_process(process)
|
||||||
|
raise
|
||||||
shell_program = shell_program or shutil.which("bash") or "/bin/bash"
|
shell_program = shell_program or shutil.which("bash") or "/bin/bash"
|
||||||
args: list[str] = [shell_program]
|
args: list[str] = [shell_program]
|
||||||
shell_name = Path(shell_program).name.lower()
|
shell_name = Path(shell_program).name.lower()
|
||||||
@@ -687,11 +722,12 @@ class ExecTool(Tool):
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
async def _kill_process_tree(process: asyncio.subprocess.Process) -> None:
|
async def _kill_process_tree(process: asyncio.subprocess.Process) -> None:
|
||||||
"""Kill a session process and descendants, then reap the root process."""
|
"""Kill a session process and descendants, then reap the root process."""
|
||||||
if process.returncode is not None:
|
owner = ExecTool._process_tree_owner(process)
|
||||||
_reap_pid(process.pid)
|
|
||||||
return
|
|
||||||
try:
|
try:
|
||||||
if _IS_WINDOWS:
|
if owner is not None:
|
||||||
|
owner.terminate()
|
||||||
|
elif _IS_WINDOWS:
|
||||||
|
if process.returncode is None:
|
||||||
with suppress(OSError, asyncio.TimeoutError):
|
with suppress(OSError, asyncio.TimeoutError):
|
||||||
await asyncio.wait_for(
|
await asyncio.wait_for(
|
||||||
asyncio.to_thread(
|
asyncio.to_thread(
|
||||||
@@ -715,8 +751,36 @@ class ExecTool(Tool):
|
|||||||
with suppress(asyncio.TimeoutError):
|
with suppress(asyncio.TimeoutError):
|
||||||
await asyncio.wait_for(process.wait(), timeout=5.0)
|
await asyncio.wait_for(process.wait(), timeout=5.0)
|
||||||
finally:
|
finally:
|
||||||
|
if owner is not None:
|
||||||
|
ExecTool._drop_process_tree_owner(process)
|
||||||
_reap_pid(process.pid)
|
_reap_pid(process.pid)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _process_tree_owner(
|
||||||
|
process: asyncio.subprocess.Process,
|
||||||
|
) -> _ProcessTreeOwner | None:
|
||||||
|
# _spawn is the only writer for this private ownership marker.
|
||||||
|
return cast(_ProcessTreeOwner | None, vars(process).get(_PROCESS_TREE_OWNER_ATTR))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _create_windows_job() -> _ProcessTreeOwner:
|
||||||
|
from nanobot.agent.tools._windows_job import WindowsJob
|
||||||
|
|
||||||
|
return WindowsJob.create()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _drop_process_tree_owner(process: asyncio.subprocess.Process) -> None:
|
||||||
|
with suppress(AttributeError):
|
||||||
|
delattr(process, _PROCESS_TREE_OWNER_ATTR)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _release_process_tree(process: asyncio.subprocess.Process) -> None:
|
||||||
|
owner = ExecTool._process_tree_owner(process)
|
||||||
|
if owner is None:
|
||||||
|
return
|
||||||
|
owner.release()
|
||||||
|
ExecTool._drop_process_tree_owner(process)
|
||||||
|
|
||||||
def _build_env(self) -> dict[str, str]:
|
def _build_env(self) -> dict[str, str]:
|
||||||
"""Build a minimal environment for subprocess execution.
|
"""Build a minimal environment for subprocess execution.
|
||||||
|
|
||||||
@@ -826,12 +890,27 @@ class ExecTool(Tool):
|
|||||||
for raw in self._extract_absolute_paths(cmd):
|
for raw in self._extract_absolute_paths(cmd):
|
||||||
try:
|
try:
|
||||||
expanded = os.path.expandvars(raw.strip())
|
expanded = os.path.expandvars(raw.strip())
|
||||||
|
# Python's expanduser() intentionally does not implement
|
||||||
|
# shell directory-stack forms. ``~+`` is the active cwd,
|
||||||
|
# while ``~-`` and indexed forms can resolve outside it;
|
||||||
|
# normalize the former and fail closed on the latter.
|
||||||
|
if expanded == "~+":
|
||||||
|
p = cwd_path
|
||||||
|
elif expanded.startswith("~+/"):
|
||||||
|
p = (cwd_path / expanded[3:]).resolve()
|
||||||
|
elif re.match(r"^~(?:-|[+-]\d+)(?:/|$)", expanded):
|
||||||
|
return ToolResult.error(
|
||||||
|
"Error: Command blocked by safety guard "
|
||||||
|
"(path outside working dir)"
|
||||||
|
+ _WORKSPACE_BOUNDARY_NOTE
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
p = Path(expanded).expanduser().resolve()
|
||||||
# Match against the un-resolved path first. On Linux,
|
# Match against the un-resolved path first. On Linux,
|
||||||
# /dev/stderr is a symlink to /proc/self/fd/2 and
|
# /dev/stderr is a symlink to /proc/self/fd/2 and
|
||||||
# ``Path.resolve()`` would mask the device-file intent.
|
# ``Path.resolve()`` would mask the device-file intent.
|
||||||
if self._is_benign_device_path(expanded):
|
if self._is_benign_device_path(expanded):
|
||||||
continue
|
continue
|
||||||
p = Path(expanded).expanduser().resolve()
|
|
||||||
except Exception:
|
except Exception:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -948,9 +1027,134 @@ class ExecTool(Tool):
|
|||||||
r"(?<![A-Za-z])(?:[A-Za-z]:[^\s\"'|><;]*|\\\\[^\s\"'|><;]+(?:\\[^\s\"'|><;]+)*)",
|
r"(?<![A-Za-z])(?:[A-Za-z]:[^\s\"'|><;]*|\\\\[^\s\"'|><;]+(?:\\[^\s\"'|><;]+)*)",
|
||||||
command
|
command
|
||||||
)
|
)
|
||||||
posix_paths = re.findall(r"(?:^|[\s|>='\"])(/[^\s\"'>;|<]+)", command) # POSIX: /absolute only
|
try:
|
||||||
home_paths = re.findall(r"(?:^|[\s>='\"])(~[/+][^\s\"'>;|<]*)", command) # POSIX/Windows home shortcut: ~/ or ~+
|
lexer = shlex.shlex(command, posix=True, punctuation_chars="();<>|&")
|
||||||
return win_paths + posix_paths + home_paths
|
lexer.whitespace_split = True
|
||||||
|
lexer.commenters = ""
|
||||||
|
tokens = list(lexer)
|
||||||
|
except ValueError:
|
||||||
|
# Keep malformed quoting fail-closed. The shell will normally reject
|
||||||
|
# it too, but a conservative raw scan must not turn it into a bypass.
|
||||||
|
tokens = [command]
|
||||||
|
|
||||||
|
paths = [*win_paths]
|
||||||
|
seen = set(win_paths)
|
||||||
|
for index, token in enumerate(tokens):
|
||||||
|
for path in ExecTool._extract_posix_paths_from_token(token):
|
||||||
|
if path not in seen:
|
||||||
|
paths.append(path)
|
||||||
|
seen.add(path)
|
||||||
|
if index > 0 and tokens[index - 1] in {"-c", "-lc", "--command"}:
|
||||||
|
for path in ExecTool._extract_absolute_paths(token):
|
||||||
|
if path not in seen:
|
||||||
|
paths.append(path)
|
||||||
|
seen.add(path)
|
||||||
|
return paths
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_posix_paths_from_token(token: str) -> list[str]:
|
||||||
|
"""Extract local POSIX/home paths from one shell-decoded token.
|
||||||
|
|
||||||
|
``shlex`` separates real grouping/redirection operators while preserving
|
||||||
|
parentheses and spaces that were quoted or escaped as part of a path.
|
||||||
|
Embedded scripts (for example ``sh -c \"cat /tmp/x\"``) still need a
|
||||||
|
small boundary scan. Colons are not general boundaries: treating them
|
||||||
|
as such misclassifies URLs, ``host:/remote`` and ``C:/Windows``. They
|
||||||
|
are considered only inside a syntactically valid assignment, where
|
||||||
|
shells expand each colon-delimited tilde component.
|
||||||
|
"""
|
||||||
|
paths: list[str] = []
|
||||||
|
for match in re.finditer(
|
||||||
|
r"file://(?:[^/\s\"']+)?(/[^\s\"'<>|;&]*)",
|
||||||
|
token,
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
):
|
||||||
|
uri_prefix = token[: match.start()]
|
||||||
|
raw_path = match.group(1)
|
||||||
|
if uri_prefix.count("(") > uri_prefix.count(")"):
|
||||||
|
raw_path = raw_path.split(")", 1)[0]
|
||||||
|
if uri_prefix.count("{") > uri_prefix.count("}"):
|
||||||
|
raw_path = raw_path.split(",", 1)[0].split("}", 1)[0]
|
||||||
|
raw_path = raw_path.split("?", 1)[0].split("#", 1)[0]
|
||||||
|
if raw_path:
|
||||||
|
paths.append(unquote(raw_path))
|
||||||
|
boundary_chars = frozenset(" \t\r\n=({,<>|;&\"'")
|
||||||
|
i = 0
|
||||||
|
while i < len(token):
|
||||||
|
is_posix = token[i] == "/"
|
||||||
|
home_match = re.match(
|
||||||
|
r"~(?:[+-](?:\d+)?|[A-Za-z0-9_.@-]+)?(?=/|:|$)",
|
||||||
|
token[i:],
|
||||||
|
)
|
||||||
|
is_home = home_match is not None
|
||||||
|
if not is_posix and not is_home:
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
prefix = token[:i]
|
||||||
|
parameter_default = (
|
||||||
|
i >= 2 and token[i - 2] == ":" and token[i - 1] in "-+?="
|
||||||
|
)
|
||||||
|
word_start = max(
|
||||||
|
(prefix.rfind(char) for char in " \t\r\n<>|;&"),
|
||||||
|
default=-1,
|
||||||
|
) + 1
|
||||||
|
word_prefix = prefix[word_start:]
|
||||||
|
assignment_component = bool(
|
||||||
|
re.fullmatch(
|
||||||
|
r"(?:[A-Za-z_][A-Za-z0-9_]*|--?[A-Za-z0-9_.-]+)="
|
||||||
|
r"(?:[^:=\s]*:)*",
|
||||||
|
word_prefix,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
at_boundary = i == 0 or token[i - 1] in boundary_chars
|
||||||
|
if is_home:
|
||||||
|
# A shell word beginning with ``~`` is a separate shlex token.
|
||||||
|
# Mid-token expansion is valid only after ``=`` or a colon in
|
||||||
|
# an assignment. This avoids PromQL/Loki ``=~`` and ``|~``
|
||||||
|
# match operators while covering PATH-like values.
|
||||||
|
at_boundary = i == 0 or assignment_component
|
||||||
|
if not at_boundary and not parameter_default:
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
if re.search(r"[A-Za-z][A-Za-z0-9+.-]*://", word_prefix) or re.match(
|
||||||
|
r"(?:[^/:=\s]+@)?[^/:=\s]+:$",
|
||||||
|
word_prefix,
|
||||||
|
):
|
||||||
|
# HTTP-style URL path/query fragments and scp-style remote paths
|
||||||
|
# are not local filesystem references. ``file://`` paths were
|
||||||
|
# decoded above. Windows drive paths are already captured by the
|
||||||
|
# platform-specific expression above.
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
assignment_value = assignment_component
|
||||||
|
if i == 0 or assignment_value:
|
||||||
|
end = len(token)
|
||||||
|
if assignment_value:
|
||||||
|
separator = token.find(":", i)
|
||||||
|
if separator >= 0:
|
||||||
|
end = separator
|
||||||
|
elif token[i - 1] in {"'", '"'}:
|
||||||
|
quote = token[i - 1]
|
||||||
|
closing = token.find(quote, i)
|
||||||
|
end = len(token) if closing < 0 else closing
|
||||||
|
else:
|
||||||
|
end_chars = set(" \t\r\n\"'<>|;&")
|
||||||
|
if prefix.count("(") > prefix.count(")"):
|
||||||
|
end_chars.add(")")
|
||||||
|
if prefix.count("{") > prefix.count("}"):
|
||||||
|
end_chars.update({",", "}"})
|
||||||
|
end = i
|
||||||
|
while end < len(token) and token[end] not in end_chars:
|
||||||
|
end += 1
|
||||||
|
|
||||||
|
candidate = token[i:end]
|
||||||
|
if candidate:
|
||||||
|
paths.append(candidate)
|
||||||
|
i = max(end, i + 1)
|
||||||
|
return paths
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _normalize_bind_roots(paths: list[str] | None) -> list[Path]:
|
def _normalize_bind_roots(paths: list[str] | None) -> list[Path]:
|
||||||
|
|||||||
+120
-17
@@ -11,7 +11,7 @@ import os
|
|||||||
import re
|
import re
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
from urllib.parse import quote, urljoin, urlparse
|
from urllib.parse import parse_qsl, quote, urljoin, urlparse
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -148,6 +148,59 @@ def _unsafe_url_request_error(exc: BaseException) -> str | None:
|
|||||||
return str(exc) if isinstance(exc, UnsafeURLRequestError) else None
|
return str(exc) if isinstance(exc, UnsafeURLRequestError) else None
|
||||||
|
|
||||||
|
|
||||||
|
# Forwarding a URL to the remote Jina reader discloses it to a third party, so
|
||||||
|
# URLs that embed credential material (userinfo, signed-URL parameters, token
|
||||||
|
# or key query values) must never leave the machine. Matching is by parameter
|
||||||
|
# name: over-matching only costs the local readability fallback, while
|
||||||
|
# under-matching leaks a secret.
|
||||||
|
_CREDENTIAL_QUERY_PARAMS = frozenset({
|
||||||
|
"access_token", "api-key", "api-token", "apikey", "api_key", "api_token",
|
||||||
|
"auth", "authorization", "client_assertion", "client_secret", "code",
|
||||||
|
"credential", "credentials", "id_token", "jwt", "key", "password",
|
||||||
|
"passwd", "private_key", "pwd", "refresh_token", "samlresponse", "secret",
|
||||||
|
"session_id", "session_token", "sessionid", "sig", "signature", "sso_token",
|
||||||
|
"ticket", "token",
|
||||||
|
})
|
||||||
|
_CREDENTIAL_QUERY_PREFIXES = ("x-amz-", "x-goog-")
|
||||||
|
|
||||||
|
|
||||||
|
def _url_carries_credentials(url: str) -> bool:
|
||||||
|
try:
|
||||||
|
parsed = urlparse(url)
|
||||||
|
except ValueError:
|
||||||
|
return True
|
||||||
|
if parsed.username is not None or parsed.password is not None:
|
||||||
|
return True
|
||||||
|
# Some frameworks still accept semicolons as query separators. Treating
|
||||||
|
# them as separators here may over-match a value, but the safe consequence
|
||||||
|
# is only using the local extractor instead of disclosing a credential.
|
||||||
|
query = parsed.query.replace(";", "&")
|
||||||
|
for name, _value in parse_qsl(query, keep_blank_values=True):
|
||||||
|
lowered = name.strip().lower()
|
||||||
|
if lowered in _CREDENTIAL_QUERY_PARAMS or lowered.startswith(_CREDENTIAL_QUERY_PREFIXES):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _redact_url_for_log(url: str) -> str:
|
||||||
|
"""Return only a URL's origin, excluding userinfo, path, query, and fragment."""
|
||||||
|
try:
|
||||||
|
parsed = urlparse(url)
|
||||||
|
hostname = parsed.hostname
|
||||||
|
if not parsed.scheme or hostname is None:
|
||||||
|
return "<redacted URL>"
|
||||||
|
if ":" in hostname:
|
||||||
|
hostname = f"[{hostname}]"
|
||||||
|
try:
|
||||||
|
port = parsed.port
|
||||||
|
except ValueError:
|
||||||
|
port = None
|
||||||
|
authority = f"{hostname}:{port}" if port is not None else hostname
|
||||||
|
return f"{parsed.scheme}://{authority}"
|
||||||
|
except ValueError:
|
||||||
|
return "<redacted URL>"
|
||||||
|
|
||||||
|
|
||||||
async def _get_with_safe_redirects(
|
async def _get_with_safe_redirects(
|
||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
url: str,
|
url: str,
|
||||||
@@ -191,13 +244,14 @@ async def _stream_with_safe_redirects(
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
url: str,
|
url: str,
|
||||||
headers: dict[str, str] | None = None,
|
headers: dict[str, str] | None = None,
|
||||||
) -> tuple[httpx.Response | None, Any | None, str | None]:
|
) -> tuple[httpx.Response | None, Any | None, str | None, bool]:
|
||||||
"""Open a streamed response while validating every redirect target first."""
|
"""Open a streamed response while validating every redirect target first."""
|
||||||
current_url = url
|
current_url = url
|
||||||
|
chain_carries_credentials = _url_carries_credentials(url)
|
||||||
for _ in range(MAX_REDIRECTS + 1):
|
for _ in range(MAX_REDIRECTS + 1):
|
||||||
is_valid, error_msg, _ = _resolve_url_safe(current_url)
|
is_valid, error_msg, _ = _resolve_url_safe(current_url)
|
||||||
if not is_valid:
|
if not is_valid:
|
||||||
return None, None, f"Redirect blocked: {error_msg}"
|
return None, None, f"Redirect blocked: {error_msg}", chain_carries_credentials
|
||||||
|
|
||||||
stream = client.stream(
|
stream = client.stream(
|
||||||
"GET",
|
"GET",
|
||||||
@@ -210,26 +264,39 @@ async def _stream_with_safe_redirects(
|
|||||||
except httpx.RequestError as exc:
|
except httpx.RequestError as exc:
|
||||||
unsafe_error = _unsafe_url_request_error(exc)
|
unsafe_error = _unsafe_url_request_error(exc)
|
||||||
if unsafe_error is not None:
|
if unsafe_error is not None:
|
||||||
return None, None, f"Redirect blocked: {unsafe_error}"
|
return (
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
f"Redirect blocked: {unsafe_error}",
|
||||||
|
chain_carries_credentials,
|
||||||
|
)
|
||||||
raise
|
raise
|
||||||
is_redirect = 300 <= response.status_code < 400
|
is_redirect = 300 <= response.status_code < 400
|
||||||
if not is_redirect:
|
if not is_redirect:
|
||||||
return response, stream, None
|
return response, stream, None, chain_carries_credentials
|
||||||
|
|
||||||
location = response.headers.get("location")
|
location = response.headers.get("location")
|
||||||
if not location:
|
if not location:
|
||||||
return response, stream, None
|
return response, stream, None, chain_carries_credentials
|
||||||
|
|
||||||
next_url = urljoin(str(response.url), location)
|
next_url = urljoin(str(response.url), location)
|
||||||
|
chain_carries_credentials = (
|
||||||
|
chain_carries_credentials or _url_carries_credentials(next_url)
|
||||||
|
)
|
||||||
is_valid, error_msg = _validate_url_safe(next_url)
|
is_valid, error_msg = _validate_url_safe(next_url)
|
||||||
if not is_valid:
|
if not is_valid:
|
||||||
await stream.__aexit__(None, None, None)
|
await stream.__aexit__(None, None, None)
|
||||||
return None, None, f"Redirect blocked: {error_msg}"
|
return None, None, f"Redirect blocked: {error_msg}", chain_carries_credentials
|
||||||
|
|
||||||
await stream.__aexit__(None, None, None)
|
await stream.__aexit__(None, None, None)
|
||||||
current_url = next_url
|
current_url = next_url
|
||||||
|
|
||||||
return None, None, f"Too many redirects: exceeded limit of {MAX_REDIRECTS}"
|
return (
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
f"Too many redirects: exceeded limit of {MAX_REDIRECTS}",
|
||||||
|
chain_carries_credentials,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _format_results(query: str, items: list[dict[str, Any]], n: int) -> str:
|
def _format_results(query: str, items: list[dict[str, Any]], n: int) -> str:
|
||||||
@@ -1043,20 +1110,26 @@ class WebFetchTool(Tool):
|
|||||||
if not is_valid:
|
if not is_valid:
|
||||||
return json.dumps({"error": f"URL validation failed: {error_msg}", "url": url}, ensure_ascii=False)
|
return json.dumps({"error": f"URL validation failed: {error_msg}", "url": url}, ensure_ascii=False)
|
||||||
|
|
||||||
# Detect and fetch images directly to avoid Jina's textual image captioning
|
# Detect and fetch images directly to avoid Jina's textual image captioning.
|
||||||
|
# This local preflight also proves that no credential-bearing URL occurs
|
||||||
|
# in the redirect chain before the original URL may be sent to Jina.
|
||||||
|
jina_remote_safe = False
|
||||||
try:
|
try:
|
||||||
async with httpx.AsyncClient(
|
async with httpx.AsyncClient(
|
||||||
**_fetch_client_kwargs(self.proxy, 15.0),
|
**_fetch_client_kwargs(self.proxy, 15.0),
|
||||||
) as client:
|
) as client:
|
||||||
r, stream, redirect_error = await _stream_with_safe_redirects(
|
r, stream, redirect_error, chain_carries_credentials = (
|
||||||
|
await _stream_with_safe_redirects(
|
||||||
client,
|
client,
|
||||||
url,
|
url,
|
||||||
headers={"User-Agent": self.user_agent},
|
headers={"User-Agent": self.user_agent},
|
||||||
)
|
)
|
||||||
|
)
|
||||||
if redirect_error:
|
if redirect_error:
|
||||||
return json.dumps({"error": redirect_error, "url": url}, ensure_ascii=False)
|
return json.dumps({"error": redirect_error, "url": url}, ensure_ascii=False)
|
||||||
if r is None:
|
if r is None:
|
||||||
return json.dumps({"error": "Fetch failed", "url": url}, ensure_ascii=False)
|
return json.dumps({"error": "Fetch failed", "url": url}, ensure_ascii=False)
|
||||||
|
jina_remote_safe = not chain_carries_credentials
|
||||||
|
|
||||||
try:
|
try:
|
||||||
ctype = r.headers.get("content-type", "")
|
ctype = r.headers.get("content-type", "")
|
||||||
@@ -1071,10 +1144,14 @@ class WebFetchTool(Tool):
|
|||||||
unsafe_error = _unsafe_url_request_error(e)
|
unsafe_error = _unsafe_url_request_error(e)
|
||||||
if unsafe_error is not None:
|
if unsafe_error is not None:
|
||||||
return json.dumps({"error": f"URL validation failed: {unsafe_error}", "url": url}, ensure_ascii=False)
|
return json.dumps({"error": f"URL validation failed: {unsafe_error}", "url": url}, ensure_ascii=False)
|
||||||
logger.debug("Pre-fetch image detection failed for {}: {}", url, e)
|
logger.debug(
|
||||||
|
"Pre-fetch image detection failed for {} ({})",
|
||||||
|
_redact_url_for_log(url),
|
||||||
|
type(e).__name__,
|
||||||
|
)
|
||||||
|
|
||||||
result = None
|
result = None
|
||||||
if self.config.use_jina_reader:
|
if self.config.use_jina_reader and jina_remote_safe:
|
||||||
result = await self._fetch_jina(url, max_chars)
|
result = await self._fetch_jina(url, max_chars)
|
||||||
if result is None:
|
if result is None:
|
||||||
result = await self._fetch_readability(url, extract_mode, max_chars)
|
result = await self._fetch_readability(url, extract_mode, max_chars)
|
||||||
@@ -1082,13 +1159,23 @@ class WebFetchTool(Tool):
|
|||||||
|
|
||||||
async def _fetch_jina(self, url: str, max_chars: int) -> str | None:
|
async def _fetch_jina(self, url: str, max_chars: int) -> str | None:
|
||||||
"""Try fetching via Jina Reader API. Returns None on failure."""
|
"""Try fetching via Jina Reader API. Returns None on failure."""
|
||||||
|
if _url_carries_credentials(url):
|
||||||
|
logger.debug(
|
||||||
|
"Skipping Jina Reader for {}: URL carries credential material",
|
||||||
|
_redact_url_for_log(url),
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
# httpx already drops the fragment when building the request; strip it
|
||||||
|
# explicitly so client-side-only data (OAuth implicit flows put tokens
|
||||||
|
# there) stays out of this path even if the transport changes.
|
||||||
|
forwarded_url = url.split("#", 1)[0]
|
||||||
try:
|
try:
|
||||||
headers = {"Accept": "application/json", "User-Agent": self.user_agent}
|
headers = {"Accept": "application/json", "User-Agent": self.user_agent}
|
||||||
jina_key = os.environ.get("JINA_API_KEY", "")
|
jina_key = os.environ.get("JINA_API_KEY", "")
|
||||||
if jina_key:
|
if jina_key:
|
||||||
headers["Authorization"] = f"Bearer {jina_key}"
|
headers["Authorization"] = f"Bearer {jina_key}"
|
||||||
async with httpx.AsyncClient(proxy=self.proxy, timeout=20.0) as client:
|
async with httpx.AsyncClient(proxy=self.proxy, timeout=20.0) as client:
|
||||||
r = await client.get(f"https://r.jina.ai/{url}", headers=headers)
|
r = await client.get(f"https://r.jina.ai/{forwarded_url}", headers=headers)
|
||||||
if r.status_code == 429:
|
if r.status_code == 429:
|
||||||
logger.debug("Jina Reader rate limited, falling back to readability")
|
logger.debug("Jina Reader rate limited, falling back to readability")
|
||||||
return None
|
return None
|
||||||
@@ -1113,7 +1200,11 @@ class WebFetchTool(Tool):
|
|||||||
"untrusted": True, "text": text,
|
"untrusted": True, "text": text,
|
||||||
}, ensure_ascii=False)
|
}, ensure_ascii=False)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug("Jina Reader failed for {}, falling back to readability: {}", url, e)
|
logger.debug(
|
||||||
|
"Jina Reader failed for {}, falling back to readability ({})",
|
||||||
|
_redact_url_for_log(url),
|
||||||
|
type(e).__name__,
|
||||||
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _fetch_readability(self, url: str, extract_mode: str, max_chars: int) -> Any:
|
async def _fetch_readability(self, url: str, extract_mode: str, max_chars: int) -> Any:
|
||||||
@@ -1144,7 +1235,11 @@ class WebFetchTool(Tool):
|
|||||||
text = self._extract_readable_html(r.text, extract_mode)
|
text = self._extract_readable_html(r.text, extract_mode)
|
||||||
extractor = "readability"
|
extractor = "readability"
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("Readability failed for {}, using raw HTML fallback: {}", url, e)
|
logger.warning(
|
||||||
|
"Readability failed for {}, using raw HTML fallback ({})",
|
||||||
|
_redact_url_for_log(url),
|
||||||
|
type(e).__name__,
|
||||||
|
)
|
||||||
text, extractor = _normalize(_strip_tags(r.text)), "html"
|
text, extractor = _normalize(_strip_tags(r.text)), "html"
|
||||||
else:
|
else:
|
||||||
text, extractor = r.text, "raw"
|
text, extractor = r.text, "raw"
|
||||||
@@ -1160,10 +1255,18 @@ class WebFetchTool(Tool):
|
|||||||
"untrusted": True, "text": text,
|
"untrusted": True, "text": text,
|
||||||
}, ensure_ascii=False)
|
}, ensure_ascii=False)
|
||||||
except httpx.ProxyError as e:
|
except httpx.ProxyError as e:
|
||||||
logger.exception("WebFetch proxy error for {}", url)
|
logger.warning(
|
||||||
|
"WebFetch proxy error for {} ({})",
|
||||||
|
_redact_url_for_log(url),
|
||||||
|
type(e).__name__,
|
||||||
|
)
|
||||||
return json.dumps({"error": f"Proxy error: {e}", "url": url}, ensure_ascii=False)
|
return json.dumps({"error": f"Proxy error: {e}", "url": url}, ensure_ascii=False)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception("WebFetch error for {}", url)
|
logger.warning(
|
||||||
|
"WebFetch error for {} ({})",
|
||||||
|
_redact_url_for_log(url),
|
||||||
|
type(e).__name__,
|
||||||
|
)
|
||||||
return json.dumps({"error": str(e), "url": url}, ensure_ascii=False)
|
return json.dumps({"error": str(e), "url": url}, ensure_ascii=False)
|
||||||
|
|
||||||
def _extract_readable_html(self, html_content: str, extract_mode: str) -> str:
|
def _extract_readable_html(self, html_content: str, extract_mode: str) -> str:
|
||||||
|
|||||||
+21
-8
@@ -48,6 +48,7 @@ _AGENT_LOOP_KEY = web.AppKey[Any]("agent_loop")
|
|||||||
_MODEL_NAME_KEY = web.AppKey[str]("model_name")
|
_MODEL_NAME_KEY = web.AppKey[str]("model_name")
|
||||||
_REQUEST_TIMEOUT_KEY = web.AppKey[float]("request_timeout")
|
_REQUEST_TIMEOUT_KEY = web.AppKey[float]("request_timeout")
|
||||||
_SESSION_LOCKS_KEY = web.AppKey[dict[str, asyncio.Lock]]("session_locks")
|
_SESSION_LOCKS_KEY = web.AppKey[dict[str, asyncio.Lock]]("session_locks")
|
||||||
|
_PREPARE_AGENT_KEY = web.AppKey[Callable[[], Awaitable[None]] | None]("prepare_agent")
|
||||||
_MISSING = object()
|
_MISSING = object()
|
||||||
|
|
||||||
|
|
||||||
@@ -66,6 +67,17 @@ def _app_value(
|
|||||||
return app.get(legacy_key, default)
|
return app.get(legacy_key, default)
|
||||||
|
|
||||||
|
|
||||||
|
async def _prepare_agent(app: Any) -> None:
|
||||||
|
prepare: Callable[[], Awaitable[None]] | None = _app_value(
|
||||||
|
app,
|
||||||
|
_PREPARE_AGENT_KEY,
|
||||||
|
"prepare_agent",
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if prepare is not None:
|
||||||
|
await prepare()
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Response helpers
|
# Response helpers
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -346,8 +358,9 @@ async def handle_chat_completions(request: web.Request) -> web.Response | web.St
|
|||||||
nonlocal stream_failed
|
nonlocal stream_failed
|
||||||
try:
|
try:
|
||||||
async with session_lock:
|
async with session_lock:
|
||||||
response = await asyncio.wait_for(
|
async with asyncio.timeout(timeout_s):
|
||||||
agent_loop.process_direct(
|
await _prepare_agent(request.app)
|
||||||
|
response = await agent_loop.process_direct(
|
||||||
content=text,
|
content=text,
|
||||||
media=media_paths if media_paths else None,
|
media=media_paths if media_paths else None,
|
||||||
session_key=session_key,
|
session_key=session_key,
|
||||||
@@ -355,8 +368,6 @@ async def handle_chat_completions(request: web.Request) -> web.Response | web.St
|
|||||||
chat_id=API_CHAT_ID,
|
chat_id=API_CHAT_ID,
|
||||||
on_stream=_on_stream,
|
on_stream=_on_stream,
|
||||||
on_stream_end=_on_stream_end,
|
on_stream_end=_on_stream_end,
|
||||||
),
|
|
||||||
timeout=timeout_s,
|
|
||||||
)
|
)
|
||||||
if not emitted_content:
|
if not emitted_content:
|
||||||
response_text = _response_text(response)
|
response_text = _response_text(response)
|
||||||
@@ -390,15 +401,14 @@ async def handle_chat_completions(request: web.Request) -> web.Response | web.St
|
|||||||
try:
|
try:
|
||||||
async with session_lock:
|
async with session_lock:
|
||||||
try:
|
try:
|
||||||
response = await asyncio.wait_for(
|
async with asyncio.timeout(timeout_s):
|
||||||
agent_loop.process_direct(
|
await _prepare_agent(request.app)
|
||||||
|
response = await agent_loop.process_direct(
|
||||||
content=text,
|
content=text,
|
||||||
media=media_paths if media_paths else None,
|
media=media_paths if media_paths else None,
|
||||||
session_key=session_key,
|
session_key=session_key,
|
||||||
channel="api",
|
channel="api",
|
||||||
chat_id=API_CHAT_ID,
|
chat_id=API_CHAT_ID,
|
||||||
),
|
|
||||||
timeout=timeout_s,
|
|
||||||
)
|
)
|
||||||
response_text = _response_text(response)
|
response_text = _response_text(response)
|
||||||
if not response_text or not response_text.strip():
|
if not response_text or not response_text.strip():
|
||||||
@@ -452,6 +462,7 @@ def create_app(
|
|||||||
model_name: str = "nanobot",
|
model_name: str = "nanobot",
|
||||||
request_timeout: float = 120.0,
|
request_timeout: float = 120.0,
|
||||||
api_key: str = "",
|
api_key: str = "",
|
||||||
|
prepare_agent: Callable[[], Awaitable[None]] | None = None,
|
||||||
) -> web.Application:
|
) -> web.Application:
|
||||||
"""Create the aiohttp application.
|
"""Create the aiohttp application.
|
||||||
|
|
||||||
@@ -460,12 +471,14 @@ def create_app(
|
|||||||
model_name: Model name reported in responses.
|
model_name: Model name reported in responses.
|
||||||
request_timeout: Per-request timeout in seconds.
|
request_timeout: Per-request timeout in seconds.
|
||||||
api_key: Optional API key for Bearer-token authentication on API routes.
|
api_key: Optional API key for Bearer-token authentication on API routes.
|
||||||
|
prepare_agent: Optional application-owned readiness callback run before each turn.
|
||||||
"""
|
"""
|
||||||
app = web.Application(client_max_size=20 * 1024 * 1024) # 20MB for base64 images
|
app = web.Application(client_max_size=20 * 1024 * 1024) # 20MB for base64 images
|
||||||
app[_AGENT_LOOP_KEY] = agent_loop
|
app[_AGENT_LOOP_KEY] = agent_loop
|
||||||
app[_MODEL_NAME_KEY] = model_name
|
app[_MODEL_NAME_KEY] = model_name
|
||||||
app[_REQUEST_TIMEOUT_KEY] = request_timeout
|
app[_REQUEST_TIMEOUT_KEY] = request_timeout
|
||||||
app[_SESSION_LOCKS_KEY] = {} # per-user locks, keyed by session_key
|
app[_SESSION_LOCKS_KEY] = {} # per-user locks, keyed by session_key
|
||||||
|
app[_PREPARE_AGENT_KEY] = prepare_agent
|
||||||
|
|
||||||
@web.middleware
|
@web.middleware
|
||||||
async def auth_middleware(
|
async def auth_middleware(
|
||||||
|
|||||||
+99
-19
@@ -20,6 +20,7 @@ from urllib.parse import urlparse
|
|||||||
import httpx
|
import httpx
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
|
from nanobot.agent.skills import parse_skill_metadata, valid_skill_metadata
|
||||||
from nanobot.apps.protocol import app_manifest, compact_dict
|
from nanobot.apps.protocol import app_manifest, compact_dict
|
||||||
from nanobot.config.paths import get_runtime_subdir
|
from nanobot.config.paths import get_runtime_subdir
|
||||||
from nanobot.security.workspace_policy import is_path_within
|
from nanobot.security.workspace_policy import is_path_within
|
||||||
@@ -27,6 +28,7 @@ from nanobot.security.workspace_policy import is_path_within
|
|||||||
CLI_ANYTHING_REGISTRY_URL = "https://hkuds.github.io/CLI-Anything/registry.json"
|
CLI_ANYTHING_REGISTRY_URL = "https://hkuds.github.io/CLI-Anything/registry.json"
|
||||||
CLI_ANYTHING_PUBLIC_REGISTRY_URL = "https://hkuds.github.io/CLI-Anything/public_registry.json"
|
CLI_ANYTHING_PUBLIC_REGISTRY_URL = "https://hkuds.github.io/CLI-Anything/public_registry.json"
|
||||||
CLI_ANYTHING_RAW_BASE = "https://raw.githubusercontent.com/HKUDS/CLI-Anything/main"
|
CLI_ANYTHING_RAW_BASE = "https://raw.githubusercontent.com/HKUDS/CLI-Anything/main"
|
||||||
|
AGENT_PLUGIN_SCHEMA = "https://agent-plugins.org/schemas/1.0.0/plugin.schema.json"
|
||||||
NANOBOT_EXTENSION_REGISTRY_URL = "https://raw.githubusercontent.com/Re-bin/nanobot-extension/main/registry.json"
|
NANOBOT_EXTENSION_REGISTRY_URL = "https://raw.githubusercontent.com/Re-bin/nanobot-extension/main/registry.json"
|
||||||
NANOBOT_EXTENSION_RAW_BASE = "https://raw.githubusercontent.com/Re-bin/nanobot-extension/main"
|
NANOBOT_EXTENSION_RAW_BASE = "https://raw.githubusercontent.com/Re-bin/nanobot-extension/main"
|
||||||
_CATALOG_SOURCES = (
|
_CATALOG_SOURCES = (
|
||||||
@@ -210,11 +212,27 @@ def _as_object_dict(value: object) -> dict[str, Any] | None:
|
|||||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||||
|
|
||||||
|
|
||||||
def _safe_skill_name(name: str) -> str:
|
def _skill_name(name: str, *, legacy: bool = False) -> str:
|
||||||
clean = _SAFE_NAME_RE.sub("-", name.lower()).strip("-")
|
clean = _SAFE_NAME_RE.sub("-", name.lower()).strip("-")
|
||||||
|
if not legacy:
|
||||||
|
clean = clean.replace("_", "-")
|
||||||
return f"cli-app-{clean or 'app'}"
|
return f"cli-app-{clean or 'app'}"
|
||||||
|
|
||||||
|
|
||||||
|
def _plugin_skill_relative_path(name: str) -> str:
|
||||||
|
skill_name = _skill_name(name)
|
||||||
|
return f"plugins/{skill_name}/skills/{skill_name}/SKILL.md"
|
||||||
|
|
||||||
|
|
||||||
|
def cli_app_skill_relative_path(workspace: Path, name: str) -> str:
|
||||||
|
"""Return a CLI App's skill path, including the legacy location."""
|
||||||
|
canonical = _plugin_skill_relative_path(name)
|
||||||
|
legacy = f"skills/{_skill_name(name, legacy=True)}/SKILL.md"
|
||||||
|
if not (workspace / canonical).is_file() and (workspace / legacy).is_file():
|
||||||
|
return legacy
|
||||||
|
return canonical
|
||||||
|
|
||||||
|
|
||||||
def _has_shell_meta(command: str) -> bool:
|
def _has_shell_meta(command: str) -> bool:
|
||||||
return any(char in command for char in _SHELL_META_CHARS)
|
return any(char in command for char in _SHELL_META_CHARS)
|
||||||
|
|
||||||
@@ -442,6 +460,16 @@ class CliAppManager:
|
|||||||
"""Return registry names explicitly installed through CLI Apps."""
|
"""Return registry names explicitly installed through CLI Apps."""
|
||||||
return sorted(str(name) for name in self._load_installed())
|
return sorted(str(name) for name in self._load_installed())
|
||||||
|
|
||||||
|
def installed_skill_aliases(self) -> dict[str, str]:
|
||||||
|
"""Map pre-plugin CLI App skill names to their portable identities."""
|
||||||
|
aliases: dict[str, str] = {}
|
||||||
|
for name in self.installed_names():
|
||||||
|
legacy = _skill_name(name, legacy=True)
|
||||||
|
canonical = _skill_name(name)
|
||||||
|
if legacy != canonical:
|
||||||
|
aliases[legacy] = canonical
|
||||||
|
return aliases
|
||||||
|
|
||||||
def _fetch_registry(
|
def _fetch_registry(
|
||||||
self,
|
self,
|
||||||
url: str,
|
url: str,
|
||||||
@@ -613,7 +641,7 @@ class CliAppManager:
|
|||||||
"name": installed_name,
|
"name": installed_name,
|
||||||
"entry_point": entry_point,
|
"entry_point": entry_point,
|
||||||
"source": str(data.get("source") or ""),
|
"source": str(data.get("source") or ""),
|
||||||
"skill": f"skills/{_safe_skill_name(installed_name)}/SKILL.md",
|
"skill": cli_app_skill_relative_path(self.workspace, installed_name),
|
||||||
"tool": "run_cli_app",
|
"tool": "run_cli_app",
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
@@ -639,9 +667,6 @@ class CliAppManager:
|
|||||||
install_cmd = str(app.get("install_cmd") or "")
|
install_cmd = str(app.get("install_cmd") or "")
|
||||||
return not _has_shell_meta(install_cmd)
|
return not _has_shell_meta(install_cmd)
|
||||||
|
|
||||||
def _skill_path(self, name: str) -> Path:
|
|
||||||
return self.workspace / "skills" / _safe_skill_name(name) / "SKILL.md"
|
|
||||||
|
|
||||||
def _app_payload(
|
def _app_payload(
|
||||||
self,
|
self,
|
||||||
app: dict[str, Any],
|
app: dict[str, Any],
|
||||||
@@ -677,7 +702,7 @@ class CliAppManager:
|
|||||||
"status": status,
|
"status": status,
|
||||||
"logo_url": logo_url,
|
"logo_url": logo_url,
|
||||||
"brand_color": brand_color,
|
"brand_color": brand_color,
|
||||||
"skill_installed": self._skill_path(name).is_file(),
|
"skill_installed": (self.workspace / cli_app_skill_relative_path(self.workspace, name)).is_file(),
|
||||||
"manifest": self._manifest_payload(app, logo_url=logo_url, brand_color=brand_color),
|
"manifest": self._manifest_payload(app, logo_url=logo_url, brand_color=brand_color),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -713,7 +738,8 @@ class CliAppManager:
|
|||||||
name = str(app["name"])
|
name = str(app["name"])
|
||||||
entry_point = str(app.get("entry_point") or "")
|
entry_point = str(app.get("entry_point") or "")
|
||||||
strategy = self._strategy(app)
|
strategy = self._strategy(app)
|
||||||
skill_path = f"skills/{_safe_skill_name(name)}/SKILL.md"
|
skill_path = _plugin_skill_relative_path(name)
|
||||||
|
plugin_path = f"plugins/{_skill_name(name)}"
|
||||||
capabilities = [
|
capabilities = [
|
||||||
compact_dict({
|
compact_dict({
|
||||||
"type": "cli",
|
"type": "cli",
|
||||||
@@ -726,13 +752,13 @@ class CliAppManager:
|
|||||||
install = compact_dict({
|
install = compact_dict({
|
||||||
"supported": install_supported,
|
"supported": install_supported,
|
||||||
"strategy": strategy,
|
"strategy": strategy,
|
||||||
"managed_paths": [skill_path],
|
"managed_paths": [plugin_path],
|
||||||
"verification": ["entry_point_available"] if entry_point else [],
|
"verification": ["entry_point_available"] if entry_point else [],
|
||||||
})
|
})
|
||||||
remove = compact_dict({
|
remove = compact_dict({
|
||||||
"supported": strategy != "unsupported",
|
"supported": strategy != "unsupported",
|
||||||
"strategy": strategy,
|
"strategy": strategy,
|
||||||
"managed_paths": [skill_path],
|
"managed_paths": [plugin_path],
|
||||||
"verification": (
|
"verification": (
|
||||||
["package_manager_ok", "entry_point_absent", "managed_paths_absent"]
|
["package_manager_ok", "entry_point_absent", "managed_paths_absent"]
|
||||||
if strategy not in {"bundled", "unsupported"}
|
if strategy not in {"bundled", "unsupported"}
|
||||||
@@ -964,6 +990,35 @@ class CliAppManager:
|
|||||||
return None
|
return None
|
||||||
raise CliAppError("this CLI app uses an unsupported install strategy")
|
raise CliAppError("this CLI app uses an unsupported install strategy")
|
||||||
|
|
||||||
|
def _subprocess_env(self) -> dict[str, str]:
|
||||||
|
"""Minimal env for CLI app subprocesses — no API keys or secrets.
|
||||||
|
|
||||||
|
Mirrors the shell tool's allowlist so installed apps cannot read
|
||||||
|
provider credentials from the parent process environment.
|
||||||
|
"""
|
||||||
|
if sys.platform == "win32":
|
||||||
|
sr = os.environ.get("SYSTEMROOT", r"C:\Windows")
|
||||||
|
env = {
|
||||||
|
"SYSTEMROOT": sr,
|
||||||
|
"COMSPEC": os.environ.get("COMSPEC", f"{sr}\\system32\\cmd.exe"),
|
||||||
|
"USERPROFILE": os.environ.get("USERPROFILE", ""),
|
||||||
|
"HOMEDRIVE": os.environ.get("HOMEDRIVE", "C:"),
|
||||||
|
"HOMEPATH": os.environ.get("HOMEPATH", "\\"),
|
||||||
|
"TEMP": os.environ.get("TEMP", f"{sr}\\Temp"),
|
||||||
|
"TMP": os.environ.get("TMP", f"{sr}\\Temp"),
|
||||||
|
"PATHEXT": os.environ.get("PATHEXT", ".COM;.EXE;.BAT;.CMD"),
|
||||||
|
"PATH": os.environ.get("PATH", f"{sr}\\system32;{sr}"),
|
||||||
|
"PYTHONUNBUFFERED": "1",
|
||||||
|
}
|
||||||
|
return env
|
||||||
|
return {
|
||||||
|
"HOME": os.environ.get("HOME", "/tmp"),
|
||||||
|
"LANG": os.environ.get("LANG", "C.UTF-8"),
|
||||||
|
"TERM": os.environ.get("TERM", "dumb"),
|
||||||
|
"PATH": os.environ.get("PATH", "/usr/bin:/bin"),
|
||||||
|
"PYTHONUNBUFFERED": "1",
|
||||||
|
}
|
||||||
|
|
||||||
def _run_argv(self, argv: list[str], *, timeout: int) -> subprocess.CompletedProcess[str]:
|
def _run_argv(self, argv: list[str], *, timeout: int) -> subprocess.CompletedProcess[str]:
|
||||||
command = subprocess.list2cmdline(argv)
|
command = subprocess.list2cmdline(argv)
|
||||||
logger.info("CLI Apps: running {}", command)
|
logger.info("CLI Apps: running {}", command)
|
||||||
@@ -974,6 +1029,7 @@ class CliAppManager:
|
|||||||
encoding="utf-8",
|
encoding="utf-8",
|
||||||
errors="replace",
|
errors="replace",
|
||||||
timeout=timeout,
|
timeout=timeout,
|
||||||
|
env=self._subprocess_env(),
|
||||||
)
|
)
|
||||||
logger.info("CLI Apps: command exited with code {}: {}", result.returncode, command)
|
logger.info("CLI Apps: command exited with code {}: {}", result.returncode, command)
|
||||||
output = (result.stderr or result.stdout or "").strip()
|
output = (result.stderr or result.stdout or "").strip()
|
||||||
@@ -1032,11 +1088,10 @@ class CliAppManager:
|
|||||||
name = str(app.get("name") or "unknown")
|
name = str(app.get("name") or "unknown")
|
||||||
display = str(app.get("display_name") or name)
|
display = str(app.get("display_name") or name)
|
||||||
entry = str(app.get("entry_point") or f"cli-anything-{name}")
|
entry = str(app.get("entry_point") or f"cli-anything-{name}")
|
||||||
description = _catalog_description(app) or f"Use {display} from nanobot."
|
description = (_catalog_description(app) or f"Use {display} from nanobot.")[:1024]
|
||||||
return f"""---
|
return f"""---
|
||||||
name: {_safe_skill_name(name)}
|
name: {_skill_name(name)}
|
||||||
description: >-
|
description: {json.dumps(description, ensure_ascii=False)}
|
||||||
{description}
|
|
||||||
---
|
---
|
||||||
|
|
||||||
# {display}
|
# {display}
|
||||||
@@ -1056,10 +1111,17 @@ Prefer machine-readable output when the CLI supports `--json`.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def _with_nanobot_skill_note(self, content: str, app: dict[str, Any]) -> str:
|
def _with_nanobot_skill_note(self, content: str, app: dict[str, Any]) -> str:
|
||||||
|
name = str(app.get("name") or "unknown")
|
||||||
|
skill_name = _skill_name(name)
|
||||||
|
metadata = parse_skill_metadata(content)
|
||||||
|
if metadata is None or not valid_skill_metadata(metadata | {"name": skill_name}, skill_name):
|
||||||
|
content = self._fallback_skill(app)
|
||||||
|
content, replaced = re.subn(r"(?m)^name\s*:.*$", f"name: {skill_name}", content, count=1)
|
||||||
|
if not replaced:
|
||||||
|
content = content.replace("---\n", f"---\nname: {skill_name}\n", 1)
|
||||||
marker = "<!-- nanobot-cli-app-note -->"
|
marker = "<!-- nanobot-cli-app-note -->"
|
||||||
if marker in content:
|
if marker in content:
|
||||||
return content
|
return content
|
||||||
name = str(app.get("name") or "unknown")
|
|
||||||
note = f"""{marker}
|
note = f"""{marker}
|
||||||
## Nanobot execution
|
## Nanobot execution
|
||||||
|
|
||||||
@@ -1073,24 +1135,42 @@ Use the `run_cli_app` tool with `name="{name}"` for command execution. Do not in
|
|||||||
return note + "\n" + content
|
return note + "\n" + content
|
||||||
|
|
||||||
def install_skill(self, app: dict[str, Any]) -> Path:
|
def install_skill(self, app: dict[str, Any]) -> Path:
|
||||||
path = self._skill_path(str(app["name"]))
|
name = str(app["name"])
|
||||||
|
path = self.workspace / _plugin_skill_relative_path(name)
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
content = self._fetch_skill_content(app) or self._fallback_skill(app)
|
content = self._fetch_skill_content(app) or self._fallback_skill(app)
|
||||||
content = self._with_nanobot_skill_note(content, app)
|
content = self._with_nanobot_skill_note(content, app)
|
||||||
path.write_text(content, encoding="utf-8")
|
path.write_text(content, encoding="utf-8")
|
||||||
|
plugin_root = path.parents[2]
|
||||||
|
manifest = compact_dict({
|
||||||
|
"$schema": AGENT_PLUGIN_SCHEMA,
|
||||||
|
"name": _skill_name(str(app["name"])),
|
||||||
|
"version": str(app.get("version") or ""),
|
||||||
|
"description": _catalog_description(app),
|
||||||
|
})
|
||||||
|
_write_json(plugin_root / "plugin.json", manifest)
|
||||||
|
legacy_dir = self.workspace / "skills" / _skill_name(str(app["name"]), legacy=True)
|
||||||
|
if legacy_dir.is_dir():
|
||||||
|
shutil.rmtree(legacy_dir)
|
||||||
return path
|
return path
|
||||||
|
|
||||||
def remove_skill(self, name: str) -> None:
|
def remove_skill(self, name: str) -> None:
|
||||||
skill_dir = self._skill_path(name).parent
|
plugin_root = (self.workspace / _plugin_skill_relative_path(name)).parents[2]
|
||||||
if skill_dir.is_dir():
|
if plugin_root.is_dir():
|
||||||
shutil.rmtree(skill_dir)
|
shutil.rmtree(plugin_root)
|
||||||
|
legacy_dir = self.workspace / "skills" / _skill_name(name, legacy=True)
|
||||||
|
if legacy_dir.is_dir():
|
||||||
|
shutil.rmtree(legacy_dir)
|
||||||
|
|
||||||
def _record_installed(self, app: dict[str, Any]) -> dict[str, Any]:
|
def _record_installed(self, app: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
from nanobot.agent.plugins import set_agent_plugin_enabled
|
||||||
|
|
||||||
installed = self._load_installed()
|
installed = self._load_installed()
|
||||||
entry = self._installed_entry(app)
|
entry = self._installed_entry(app)
|
||||||
installed[str(app["name"])] = entry
|
installed[str(app["name"])] = entry
|
||||||
self._save_installed(installed)
|
self._save_installed(installed)
|
||||||
self.install_skill(app)
|
self.install_skill(app)
|
||||||
|
set_agent_plugin_enabled(self.workspace, _skill_name(str(app["name"])), True)
|
||||||
return entry
|
return entry
|
||||||
|
|
||||||
def install(self, name: str) -> dict[str, Any]:
|
def install(self, name: str) -> dict[str, Any]:
|
||||||
@@ -1381,7 +1461,7 @@ Use the `run_cli_app` tool with `name="{name}"` for command execution. Do not in
|
|||||||
encoding="utf-8",
|
encoding="utf-8",
|
||||||
errors="replace",
|
errors="replace",
|
||||||
timeout=effective_timeout,
|
timeout=effective_timeout,
|
||||||
env=os.environ.copy(),
|
env=self._subprocess_env(),
|
||||||
)
|
)
|
||||||
except subprocess.TimeoutExpired:
|
except subprocess.TimeoutExpired:
|
||||||
return f"CLI app '{name}' timed out after {effective_timeout}s"
|
return f"CLI app '{name}' timed out after {effective_timeout}s"
|
||||||
|
|||||||
@@ -20,6 +20,8 @@ def runtime_lines_for_request(
|
|||||||
"""Return CLI App annotations from an immutable request snapshot."""
|
"""Return CLI App annotations from an immutable request snapshot."""
|
||||||
structured = metadata.get("cli_apps") if isinstance(metadata, Mapping) else None
|
structured = metadata.get("cli_apps") if isinstance(metadata, Mapping) else None
|
||||||
if isinstance(structured, list):
|
if isinstance(structured, list):
|
||||||
|
from nanobot.apps.cli.service import cli_app_skill_relative_path
|
||||||
|
|
||||||
structured_items = cast(list[Any], structured)
|
structured_items = cast(list[Any], structured)
|
||||||
mentions = [
|
mentions = [
|
||||||
cast(Mapping[str, Any], item) for item in structured_items
|
cast(Mapping[str, Any], item) for item in structured_items
|
||||||
@@ -32,7 +34,7 @@ def runtime_lines_for_request(
|
|||||||
f"@{str(item['name']).strip().lower()} "
|
f"@{str(item['name']).strip().lower()} "
|
||||||
f"(installed; tool=run_cli_app; "
|
f"(installed; tool=run_cli_app; "
|
||||||
f"entry_point={str(item.get('entry_point') or 'unknown')}; "
|
f"entry_point={str(item.get('entry_point') or 'unknown')}; "
|
||||||
f"skill=skills/cli-app-{str(item['name']).strip().lower()}/SKILL.md). "
|
f"skill={cli_app_skill_relative_path(workspace, str(item['name']))}). "
|
||||||
"Read the skill when useful, then run this app with `run_cli_app`; do not bypass it with shell."
|
"Read the skill when useful, then run this app with `run_cli_app`; do not bypass it with shell."
|
||||||
for item in mentions
|
for item in mentions
|
||||||
if str(item.get("name") or "").strip()
|
if str(item.get("name") or "").strip()
|
||||||
|
|||||||
@@ -16,7 +16,6 @@ OUTBOUND_META_AGENT_UI = "_agent_ui"
|
|||||||
# loop to update runtime state without going through a user session.
|
# loop to update runtime state without going through a user session.
|
||||||
INBOUND_META_RUNTIME_CONTROL = "_runtime_control"
|
INBOUND_META_RUNTIME_CONTROL = "_runtime_control"
|
||||||
RUNTIME_CONTROL_ACK = "_ack"
|
RUNTIME_CONTROL_ACK = "_ack"
|
||||||
RUNTIME_CONTROL_MCP_RELOAD = "mcp_reload"
|
|
||||||
RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD = "image_generation_reload"
|
RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD = "image_generation_reload"
|
||||||
RUNTIME_CONTROL_SESSION_DISCARD = "session_discard"
|
RUNTIME_CONTROL_SESSION_DISCARD = "session_discard"
|
||||||
|
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import type {
|
|||||||
NanobotFeatureInfo,
|
NanobotFeatureInfo,
|
||||||
NanobotFeaturesPayload,
|
NanobotFeaturesPayload,
|
||||||
} from "@/lib/types";
|
} from "@/lib/types";
|
||||||
|
import { useClient } from "@/providers/ClientProvider";
|
||||||
|
|
||||||
import { FeishuConnectFlow } from "./FeishuConnectFlow";
|
import { FeishuConnectFlow } from "./FeishuConnectFlow";
|
||||||
|
|
||||||
@@ -33,7 +34,6 @@ export function FeishuAssistantsPanel({
|
|||||||
|
|
||||||
return (
|
return (
|
||||||
<ChannelInstancesPanel
|
<ChannelInstancesPanel
|
||||||
token={token}
|
|
||||||
feature={feature}
|
feature={feature}
|
||||||
showBrandLogos={showBrandLogos}
|
showBrandLogos={showBrandLogos}
|
||||||
chatAppsDocsUrl={chatAppsDocsUrl}
|
chatAppsDocsUrl={chatAppsDocsUrl}
|
||||||
@@ -92,6 +92,7 @@ function FeishuInstanceAction({
|
|||||||
instance: NanobotChannelInstanceInfo;
|
instance: NanobotChannelInstanceInfo;
|
||||||
onFeaturesUpdate: (payload: NanobotFeaturesPayload) => void;
|
onFeaturesUpdate: (payload: NanobotFeaturesPayload) => void;
|
||||||
}) {
|
}) {
|
||||||
|
const { client } = useClient();
|
||||||
const { t } = useTranslation();
|
const { t } = useTranslation();
|
||||||
const tx = channelTranslator(t, "feishu");
|
const tx = channelTranslator(t, "feishu");
|
||||||
const [busy, setBusy] = useState(false);
|
const [busy, setBusy] = useState(false);
|
||||||
@@ -114,7 +115,7 @@ function FeishuInstanceAction({
|
|||||||
setError(null);
|
setError(null);
|
||||||
try {
|
try {
|
||||||
onFeaturesUpdate(
|
onFeaturesUpdate(
|
||||||
await enableNanobotFeature(token, "feishu", { instanceId: instance.id }),
|
await enableNanobotFeature(client, "feishu", { instanceId: instance.id }),
|
||||||
);
|
);
|
||||||
} catch (err) {
|
} catch (err) {
|
||||||
setError((err as Error).message);
|
setError((err as Error).message);
|
||||||
|
|||||||
@@ -1,7 +1,13 @@
|
|||||||
|
import { lazy } from "react";
|
||||||
|
|
||||||
import type { ChannelUiContribution } from "@/channel-plugins/types";
|
import type { ChannelUiContribution } from "@/channel-plugins/types";
|
||||||
import { chatAppGuideUrl } from "@/components/settings/channels/catalog";
|
import { chatAppGuideUrl } from "@/components/settings/channels/catalog";
|
||||||
|
|
||||||
import { FeishuAssistantsPanel } from "./FeishuAssistantsPanel";
|
const FeishuAssistantsPanel = lazy(() =>
|
||||||
|
import("./FeishuAssistantsPanel").then(({ FeishuAssistantsPanel: component }) => ({
|
||||||
|
default: component,
|
||||||
|
})),
|
||||||
|
);
|
||||||
|
|
||||||
export default {
|
export default {
|
||||||
Panel: FeishuAssistantsPanel,
|
Panel: FeishuAssistantsPanel,
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import hashlib
|
import hashlib
|
||||||
import inspect
|
import inspect
|
||||||
from collections.abc import Callable, Iterable
|
from collections.abc import Awaitable, Callable, Iterable, Mapping
|
||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any, cast
|
from typing import TYPE_CHECKING, Any, cast
|
||||||
@@ -100,9 +100,17 @@ class ChannelManager:
|
|||||||
webui_static_dist: bool = True,
|
webui_static_dist: bool = True,
|
||||||
webui_runtime_surface: str = "browser",
|
webui_runtime_surface: str = "browser",
|
||||||
webui_runtime_capabilities: dict[str, Any] | None = None,
|
webui_runtime_capabilities: dict[str, Any] | None = None,
|
||||||
|
webui_mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
|
||||||
|
webui_mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
|
||||||
webui_skill_state_action: Callable[[set[str]], None] | None = None,
|
webui_skill_state_action: Callable[[set[str]], None] | None = None,
|
||||||
|
config_path: Path | None = None,
|
||||||
):
|
):
|
||||||
|
if config_path is None:
|
||||||
|
from nanobot.config.loader import get_config_path
|
||||||
|
|
||||||
|
config_path = get_config_path()
|
||||||
self.config = config
|
self.config = config
|
||||||
|
self._config_path = config_path.expanduser().resolve(strict=False)
|
||||||
self.bus = bus
|
self.bus = bus
|
||||||
self._session_manager = session_manager
|
self._session_manager = session_manager
|
||||||
self._cron_service = cron_service
|
self._cron_service = cron_service
|
||||||
@@ -113,6 +121,8 @@ class ChannelManager:
|
|||||||
self._webui_static_dist = webui_static_dist
|
self._webui_static_dist = webui_static_dist
|
||||||
self._webui_runtime_surface = webui_runtime_surface
|
self._webui_runtime_surface = webui_runtime_surface
|
||||||
self._webui_runtime_capabilities = dict(webui_runtime_capabilities or {})
|
self._webui_runtime_capabilities = dict(webui_runtime_capabilities or {})
|
||||||
|
self._webui_mcp_runtime_status = webui_mcp_runtime_status
|
||||||
|
self._webui_mcp_reload = webui_mcp_reload
|
||||||
self._webui_skill_state_action = webui_skill_state_action
|
self._webui_skill_state_action = webui_skill_state_action
|
||||||
self.channels: dict[str, BaseChannel] = {}
|
self.channels: dict[str, BaseChannel] = {}
|
||||||
self._channel_owners: dict[str, str] = {}
|
self._channel_owners: dict[str, str] = {}
|
||||||
@@ -170,6 +180,7 @@ class ChannelManager:
|
|||||||
static_dist_path=static_path,
|
static_dist_path=static_path,
|
||||||
workspace_path=workspace,
|
workspace_path=workspace,
|
||||||
default_restrict_to_workspace=self.config.tools.restrict_to_workspace,
|
default_restrict_to_workspace=self.config.tools.restrict_to_workspace,
|
||||||
|
config_path=self._config_path,
|
||||||
disabled_skills=set(self.config.agents.defaults.disabled_skills),
|
disabled_skills=set(self.config.agents.defaults.disabled_skills),
|
||||||
runtime_model_name=self._webui_runtime_model_name,
|
runtime_model_name=self._webui_runtime_model_name,
|
||||||
runtime_surface=self._webui_runtime_surface,
|
runtime_surface=self._webui_runtime_surface,
|
||||||
@@ -180,6 +191,8 @@ class ChannelManager:
|
|||||||
local_trigger_pending_ids=self._webui_local_trigger_pending_ids,
|
local_trigger_pending_ids=self._webui_local_trigger_pending_ids,
|
||||||
channel_feature_action=self.apply_channel_feature_action,
|
channel_feature_action=self.apply_channel_feature_action,
|
||||||
channel_runtime_status=self.get_status,
|
channel_runtime_status=self.get_status,
|
||||||
|
mcp_runtime_status=self._webui_mcp_runtime_status,
|
||||||
|
mcp_reload=self._webui_mcp_reload,
|
||||||
skill_state_action=self._webui_skill_state_action,
|
skill_state_action=self._webui_skill_state_action,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -968,6 +968,11 @@ class MatrixChannel(BaseChannel):
|
|||||||
meta["thread_reply_to_event_id"] = reply_to
|
meta["thread_reply_to_event_id"] = reply_to
|
||||||
return meta
|
return meta
|
||||||
|
|
||||||
|
def _thread_session_key(self, room_id: str, event: RoomMessage) -> str | None:
|
||||||
|
if not (root_id := self._event_thread_root_id(event)):
|
||||||
|
return None
|
||||||
|
return f"{self.name}:{room_id}:thread:{root_id}"
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _build_thread_relates_to(metadata: dict[str, Any] | None) -> dict[str, Any] | None:
|
def _build_thread_relates_to(metadata: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||||
if not metadata:
|
if not metadata:
|
||||||
@@ -1171,6 +1176,7 @@ class MatrixChannel(BaseChannel):
|
|||||||
await self._handle_message(
|
await self._handle_message(
|
||||||
sender_id=event.sender, chat_id=room.room_id,
|
sender_id=event.sender, chat_id=room.room_id,
|
||||||
content=event.body, metadata=self._base_metadata(room, event),
|
content=event.body, metadata=self._base_metadata(room, event),
|
||||||
|
session_key=self._thread_session_key(room.room_id, event),
|
||||||
is_dm=self._is_direct_room(room),
|
is_dm=self._is_direct_room(room),
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -1209,6 +1215,7 @@ class MatrixChannel(BaseChannel):
|
|||||||
content="\n".join(parts),
|
content="\n".join(parts),
|
||||||
media=[attachment["path"]] if attachment else [],
|
media=[attachment["path"]] if attachment else [],
|
||||||
metadata=meta,
|
metadata=meta,
|
||||||
|
session_key=self._thread_session_key(room.room_id, event),
|
||||||
is_dm=self._is_direct_room(room),
|
is_dm=self._is_direct_room(room),
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|||||||
@@ -971,6 +971,81 @@ async def test_on_message_sets_thread_metadata_when_threaded_event() -> None:
|
|||||||
assert metadata["thread_root_event_id"] == "$root1"
|
assert metadata["thread_root_event_id"] == "$root1"
|
||||||
assert metadata["thread_reply_to_event_id"] == "$reply1"
|
assert metadata["thread_reply_to_event_id"] == "$reply1"
|
||||||
assert metadata["event_id"] == "$reply1"
|
assert metadata["event_id"] == "$reply1"
|
||||||
|
assert handled[0]["session_key"] == "matrix:!room:matrix.org:thread:$root1"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_keeps_matrix_thread_sessions_independent() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
handled: list[dict[str, object]] = []
|
||||||
|
|
||||||
|
async def _fake_handle_message(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = _fake_handle_message # type: ignore[method-assign]
|
||||||
|
|
||||||
|
room = SimpleNamespace(room_id="!room:matrix.org", display_name="Test room", member_count=3)
|
||||||
|
|
||||||
|
def _thread_event(body: str, event_id: str, root_id: str) -> SimpleNamespace:
|
||||||
|
return SimpleNamespace(
|
||||||
|
sender="@alice:matrix.org",
|
||||||
|
body=body,
|
||||||
|
event_id=event_id,
|
||||||
|
source={
|
||||||
|
"content": {
|
||||||
|
"m.relates_to": {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": root_id,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._on_message(room, _thread_event("Plan the wedding", "$reply1", "$root1"))
|
||||||
|
await channel._on_message(room, _thread_event("Pick a gift", "$reply2", "$root1"))
|
||||||
|
await channel._on_message(room, _thread_event("/new", "$reply3", "$root2"))
|
||||||
|
|
||||||
|
assert [message["chat_id"] for message in handled] == [
|
||||||
|
"!room:matrix.org",
|
||||||
|
"!room:matrix.org",
|
||||||
|
"!room:matrix.org",
|
||||||
|
]
|
||||||
|
assert [message["session_key"] for message in handled] == [
|
||||||
|
"matrix:!room:matrix.org:thread:$root1",
|
||||||
|
"matrix:!room:matrix.org:thread:$root1",
|
||||||
|
"matrix:!room:matrix.org:thread:$root2",
|
||||||
|
]
|
||||||
|
assert handled[2]["content"] == "/new"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_keeps_non_threaded_room_session() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
handled: list[dict[str, object]] = []
|
||||||
|
|
||||||
|
async def _fake_handle_message(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = _fake_handle_message # type: ignore[method-assign]
|
||||||
|
|
||||||
|
room = SimpleNamespace(room_id="!room:matrix.org", display_name="Test room", member_count=3)
|
||||||
|
event = SimpleNamespace(
|
||||||
|
sender="@alice:matrix.org",
|
||||||
|
body="Hello",
|
||||||
|
event_id="$event1",
|
||||||
|
source={"content": {}},
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._on_message(room, event)
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["session_key"] is None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -1076,6 +1151,7 @@ async def test_on_media_message_sets_thread_metadata_when_threaded_event(
|
|||||||
assert metadata["thread_root_event_id"] == "$root1"
|
assert metadata["thread_root_event_id"] == "$root1"
|
||||||
assert metadata["thread_reply_to_event_id"] == "$event1"
|
assert metadata["thread_reply_to_event_id"] == "$event1"
|
||||||
assert metadata["event_id"] == "$event1"
|
assert metadata["event_id"] == "$event1"
|
||||||
|
assert handled[0]["session_key"] == "matrix:!room:matrix.org:thread:$root1"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@@ -373,6 +373,13 @@ class WebSocketChannel(BaseChannel):
|
|||||||
self._conn_default: dict[ServerConnection, str] = {}
|
self._conn_default: dict[ServerConnection, str] = {}
|
||||||
# Connections authenticated with a one-time token from /webui/bootstrap.
|
# Connections authenticated with a one-time token from /webui/bootstrap.
|
||||||
self._webui_connections: set[ServerConnection] = set()
|
self._webui_connections: set[ServerConnection] = set()
|
||||||
|
# Request/reply mutations aren't replayed across reconnects. Tasks may
|
||||||
|
# finish after a client-side deadline so an already-started mutation
|
||||||
|
# isn't ambiguously cancelled halfway through.
|
||||||
|
self._webui_request_tasks: dict[
|
||||||
|
tuple[ServerConnection, str],
|
||||||
|
asyncio.Task[None],
|
||||||
|
] = {}
|
||||||
self._stop_event: asyncio.Event | None = None
|
self._stop_event: asyncio.Event | None = None
|
||||||
self._server_task: asyncio.Task[None] | None = None
|
self._server_task: asyncio.Task[None] | None = None
|
||||||
|
|
||||||
@@ -523,6 +530,10 @@ class WebSocketChannel(BaseChannel):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.warning("failed to send {} event: {}", event, e)
|
self.logger.warning("failed to send {} event: {}", event, e)
|
||||||
|
|
||||||
|
async def _broadcast_webui_event(self, event: str, **fields: Any) -> None:
|
||||||
|
for connection in tuple(self._webui_connections):
|
||||||
|
await self._send_event(connection, event, **fields)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def default_config(cls) -> dict[str, Any]:
|
def default_config(cls) -> dict[str, Any]:
|
||||||
return WebSocketConfig().model_dump(by_alias=True)
|
return WebSocketConfig().model_dump(by_alias=True)
|
||||||
@@ -758,6 +769,9 @@ class WebSocketChannel(BaseChannel):
|
|||||||
) -> None:
|
) -> None:
|
||||||
"""Route one typed inbound envelope (``new_chat`` / ``attach`` / ``message``)."""
|
"""Route one typed inbound envelope (``new_chat`` / ``attach`` / ``message``)."""
|
||||||
t = envelope.get("type")
|
t = envelope.get("type")
|
||||||
|
if t == "webui_request":
|
||||||
|
await self._start_webui_request(connection, envelope)
|
||||||
|
return
|
||||||
if t == "new_chat":
|
if t == "new_chat":
|
||||||
new_id = str(uuid.uuid4())
|
new_id = str(uuid.uuid4())
|
||||||
scope = await self._workspace_scope_or_error(
|
scope = await self._workspace_scope_or_error(
|
||||||
@@ -838,7 +852,7 @@ class WebSocketChannel(BaseChannel):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
await asyncio.to_thread(
|
saved_state = await asyncio.to_thread(
|
||||||
write_webui_sidebar_state,
|
write_webui_sidebar_state,
|
||||||
cast(dict[str, Any], state),
|
cast(dict[str, Any], state),
|
||||||
)
|
)
|
||||||
@@ -849,6 +863,11 @@ class WebSocketChannel(BaseChannel):
|
|||||||
detail="invalid_sidebar_state",
|
detail="invalid_sidebar_state",
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
await self._broadcast_webui_event(
|
||||||
|
"sidebar_state_updated",
|
||||||
|
state=saved_state,
|
||||||
|
)
|
||||||
|
return
|
||||||
if t == "set_workspace_scope":
|
if t == "set_workspace_scope":
|
||||||
cid = envelope.get("chat_id")
|
cid = envelope.get("chat_id")
|
||||||
if not _is_valid_chat_id(cid):
|
if not _is_valid_chat_id(cid):
|
||||||
@@ -1105,6 +1124,157 @@ class WebSocketChannel(BaseChannel):
|
|||||||
return
|
return
|
||||||
await self._send_event(connection, "error", detail=f"unknown type: {t!r}")
|
await self._send_event(connection, "error", detail=f"unknown type: {t!r}")
|
||||||
|
|
||||||
|
async def _start_webui_request(
|
||||||
|
self,
|
||||||
|
connection: ServerConnection,
|
||||||
|
envelope: dict[str, Any],
|
||||||
|
) -> None:
|
||||||
|
request_id = envelope.get("request_id")
|
||||||
|
if not isinstance(request_id, str) or re.fullmatch(
|
||||||
|
r"[A-Za-z0-9._:-]{1,128}",
|
||||||
|
request_id,
|
||||||
|
) is None:
|
||||||
|
await self._send_event(
|
||||||
|
connection,
|
||||||
|
"error",
|
||||||
|
detail="invalid webui request_id",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if connection not in self._webui_connections:
|
||||||
|
await self._send_webui_response(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
status=403,
|
||||||
|
message="access_denied",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
action = envelope.get("action")
|
||||||
|
payload = envelope.get("payload")
|
||||||
|
if not isinstance(action, str) or re.fullmatch(
|
||||||
|
r"[a-z][a-z0-9_.]{0,127}",
|
||||||
|
action,
|
||||||
|
) is None:
|
||||||
|
await self._send_webui_response(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
status=400,
|
||||||
|
message="invalid WebUI mutation action",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
await self._send_webui_response(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
status=400,
|
||||||
|
message="WebUI mutation payload must be an object",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
key = (connection, request_id)
|
||||||
|
if key in self._webui_request_tasks:
|
||||||
|
await self._send_webui_response(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
status=409,
|
||||||
|
message="duplicate WebUI request_id",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
task = asyncio.create_task(
|
||||||
|
self._complete_webui_request(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
action,
|
||||||
|
cast(dict[str, Any], payload),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self._webui_request_tasks[key] = task
|
||||||
|
|
||||||
|
async def _complete_webui_request(
|
||||||
|
self,
|
||||||
|
connection: ServerConnection,
|
||||||
|
request_id: str,
|
||||||
|
action: str,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
) -> None:
|
||||||
|
try:
|
||||||
|
response = await self._http_router.dispatch_webui_mutation(
|
||||||
|
connection,
|
||||||
|
action,
|
||||||
|
payload,
|
||||||
|
)
|
||||||
|
status = response.status_code
|
||||||
|
body = bytes(response.body).decode("utf-8", errors="replace").strip()
|
||||||
|
if 200 <= status < 300:
|
||||||
|
try:
|
||||||
|
result = json.loads(body)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
await self._send_webui_response(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
status=502,
|
||||||
|
message="WebUI mutation returned an invalid response",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if action == "sidebar.update" and isinstance(result, dict):
|
||||||
|
await self._broadcast_webui_event(
|
||||||
|
"sidebar_state_updated",
|
||||||
|
state=result,
|
||||||
|
)
|
||||||
|
await self._send_webui_response(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
result=result,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
await self._send_webui_response(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
status=status,
|
||||||
|
message=body or response.reason_phrase,
|
||||||
|
)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("WebUI mutation '{}' failed", action)
|
||||||
|
await self._send_webui_response(
|
||||||
|
connection,
|
||||||
|
request_id,
|
||||||
|
status=500,
|
||||||
|
message="WebUI mutation failed",
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
self._webui_request_tasks.pop((connection, request_id), None)
|
||||||
|
|
||||||
|
async def _send_webui_response(
|
||||||
|
self,
|
||||||
|
connection: ServerConnection,
|
||||||
|
request_id: str,
|
||||||
|
*,
|
||||||
|
result: Any = None,
|
||||||
|
status: int | None = None,
|
||||||
|
message: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
if status is None:
|
||||||
|
await self._send_event(
|
||||||
|
connection,
|
||||||
|
"webui_response",
|
||||||
|
request_id=request_id,
|
||||||
|
ok=True,
|
||||||
|
result=result,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
await self._send_event(
|
||||||
|
connection,
|
||||||
|
"webui_response",
|
||||||
|
request_id=request_id,
|
||||||
|
ok=False,
|
||||||
|
error={
|
||||||
|
"status": status,
|
||||||
|
"message": message or "WebUI mutation failed",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
async def _workspace_scope_or_error(
|
async def _workspace_scope_or_error(
|
||||||
self,
|
self,
|
||||||
connection: ServerConnection,
|
connection: ServerConnection,
|
||||||
@@ -1145,6 +1315,12 @@ class WebSocketChannel(BaseChannel):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.warning("server task error during shutdown: {}", e)
|
self.logger.warning("server task error during shutdown: {}", e)
|
||||||
self._server_task = None
|
self._server_task = None
|
||||||
|
mutation_tasks = tuple(self._webui_request_tasks.values())
|
||||||
|
for task in mutation_tasks:
|
||||||
|
task.cancel()
|
||||||
|
if mutation_tasks:
|
||||||
|
await asyncio.gather(*mutation_tasks, return_exceptions=True)
|
||||||
|
self._webui_request_tasks.clear()
|
||||||
self._subs.clear()
|
self._subs.clear()
|
||||||
self._conn_chats.clear()
|
self._conn_chats.clear()
|
||||||
self._conn_default.clear()
|
self._conn_default.clear()
|
||||||
|
|||||||
@@ -3,12 +3,16 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
|
import uuid
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from types import SimpleNamespace
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
import websockets
|
import websockets
|
||||||
|
from websockets.datastructures import Headers
|
||||||
from websockets.exceptions import ConnectionClosed
|
from websockets.exceptions import ConnectionClosed
|
||||||
from websockets.frames import Close
|
from websockets.frames import Close
|
||||||
|
|
||||||
@@ -42,6 +46,12 @@ from nanobot.security.workspace_access import WORKSPACE_SCOPE_METADATA_KEY
|
|||||||
from nanobot.session import webui_turns as wth
|
from nanobot.session import webui_turns as wth
|
||||||
from nanobot.session.manager import SessionManager
|
from nanobot.session.manager import SessionManager
|
||||||
from nanobot.webui.gateway_services import GatewayServices, build_gateway_services
|
from nanobot.webui.gateway_services import GatewayServices, build_gateway_services
|
||||||
|
from nanobot.webui.http_utils import (
|
||||||
|
http_error as _http_error,
|
||||||
|
)
|
||||||
|
from nanobot.webui.http_utils import (
|
||||||
|
http_json_response as _http_json_response,
|
||||||
|
)
|
||||||
from nanobot.webui.http_utils import (
|
from nanobot.webui.http_utils import (
|
||||||
issue_route_secret_matches as _issue_route_secret_matches,
|
issue_route_secret_matches as _issue_route_secret_matches,
|
||||||
)
|
)
|
||||||
@@ -119,6 +129,46 @@ def _basic_handler(bus: Any, **kw: Any) -> GatewayServices:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _connect_when_ready(url: str) -> Any:
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
return await websockets.connect(url)
|
||||||
|
except OSError:
|
||||||
|
await asyncio.sleep(0.02)
|
||||||
|
|
||||||
|
|
||||||
|
async def _webui_mutate(
|
||||||
|
client: Any,
|
||||||
|
action: str,
|
||||||
|
payload: dict[str, Any] | None = None,
|
||||||
|
) -> httpx.Response:
|
||||||
|
request_id = f"test-{uuid.uuid4().hex}"
|
||||||
|
await client.send(json.dumps({
|
||||||
|
"type": "webui_request",
|
||||||
|
"request_id": request_id,
|
||||||
|
"action": action,
|
||||||
|
"payload": payload or {},
|
||||||
|
}))
|
||||||
|
while True:
|
||||||
|
envelope = json.loads(await asyncio.wait_for(client.recv(), timeout=5))
|
||||||
|
if envelope.get("event") != "webui_response":
|
||||||
|
continue
|
||||||
|
if envelope.get("request_id") != request_id:
|
||||||
|
continue
|
||||||
|
if envelope.get("ok") is True:
|
||||||
|
status = 200
|
||||||
|
body = envelope.get("result")
|
||||||
|
else:
|
||||||
|
error = envelope.get("error") or {}
|
||||||
|
status = int(error.get("status") or 500)
|
||||||
|
body = {"error": str(error.get("message") or "WebUI mutation failed")}
|
||||||
|
return httpx.Response(
|
||||||
|
status,
|
||||||
|
json=body,
|
||||||
|
request=httpx.Request("WS", "http://nanobot.local/webui-mutation"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_stop_treats_cancelled_server_task_as_shutdown() -> None:
|
async def test_stop_treats_cancelled_server_task_as_shutdown() -> None:
|
||||||
channel = _ch(MessageBus())
|
channel = _ch(MessageBus())
|
||||||
@@ -667,7 +717,7 @@ async def test_token_issue_route_requires_secret_when_static_token_configured(bu
|
|||||||
bus,
|
bus,
|
||||||
port=port,
|
port=port,
|
||||||
token="static-token",
|
token="static-token",
|
||||||
tokenIssuePath="/auth/token",
|
tokenIssuePath="/custom-token",
|
||||||
websocketRequiresToken=True,
|
websocketRequiresToken=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -675,15 +725,16 @@ async def test_token_issue_route_requires_secret_when_static_token_configured(bu
|
|||||||
await asyncio.sleep(0.3)
|
await asyncio.sleep(0.3)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
denied = await _http_get(f"http://127.0.0.1:{port}/auth/token")
|
denied = await _http_get(f"http://127.0.0.1:{port}/custom-token")
|
||||||
assert denied.status_code == 401
|
assert denied.status_code == 401
|
||||||
|
|
||||||
allowed = await _http_get(
|
allowed = await _http_get(
|
||||||
f"http://127.0.0.1:{port}/auth/token",
|
f"http://127.0.0.1:{port}/custom-token",
|
||||||
headers={"Authorization": "Bearer static-token"},
|
headers={"Authorization": "Bearer static-token"},
|
||||||
)
|
)
|
||||||
assert allowed.status_code == 200
|
assert allowed.status_code == 200
|
||||||
assert allowed.json()["token"].startswith("nbwt_")
|
assert allowed.json()["token"].startswith("nbwt_")
|
||||||
|
assert allowed.headers["Cache-Control"] == "no-store"
|
||||||
finally:
|
finally:
|
||||||
await channel.stop()
|
await channel.stop()
|
||||||
await server_task
|
await server_task
|
||||||
@@ -857,6 +908,98 @@ def test_only_bootstrap_tokens_mark_webui_connections(bus: MagicMock) -> None:
|
|||||||
assert client_connection not in channel._webui_connections
|
assert client_connection not in channel._webui_connections
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticated_webui_request_returns_correlated_success(bus: MagicMock) -> None:
|
||||||
|
channel = _ch(bus)
|
||||||
|
conn = AsyncMock()
|
||||||
|
channel._webui_connections.add(conn)
|
||||||
|
channel.gateway.http.dispatch_webui_mutation = AsyncMock(
|
||||||
|
return_value=_http_json_response({"saved": True})
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
await channel._dispatch_envelope(
|
||||||
|
conn,
|
||||||
|
"webui-client",
|
||||||
|
{
|
||||||
|
"type": "webui_request",
|
||||||
|
"request_id": "request-1",
|
||||||
|
"action": "settings.provider.update",
|
||||||
|
"payload": {"provider": "openrouter", "apiKey": "secret"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
await asyncio.gather(*tuple(channel._webui_request_tasks.values()))
|
||||||
|
|
||||||
|
channel.gateway.http.dispatch_webui_mutation.assert_awaited_once_with(
|
||||||
|
conn,
|
||||||
|
"settings.provider.update",
|
||||||
|
{"provider": "openrouter", "apiKey": "secret"},
|
||||||
|
)
|
||||||
|
assert json.loads(conn.send.await_args.args[0]) == {
|
||||||
|
"event": "webui_response",
|
||||||
|
"request_id": "request-1",
|
||||||
|
"ok": True,
|
||||||
|
"result": {"saved": True},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_webui_request_returns_correlated_route_error(bus: MagicMock) -> None:
|
||||||
|
channel = _ch(bus)
|
||||||
|
conn = AsyncMock()
|
||||||
|
channel._webui_connections.add(conn)
|
||||||
|
channel.gateway.http.dispatch_webui_mutation = AsyncMock(
|
||||||
|
return_value=_http_error(400, "invalid settings payload")
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._dispatch_envelope(
|
||||||
|
conn,
|
||||||
|
"webui-client",
|
||||||
|
{
|
||||||
|
"type": "webui_request",
|
||||||
|
"request_id": "request-2",
|
||||||
|
"action": "settings.agent.update",
|
||||||
|
"payload": {},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
await asyncio.gather(*tuple(channel._webui_request_tasks.values()))
|
||||||
|
|
||||||
|
assert json.loads(conn.send.await_args.args[0]) == {
|
||||||
|
"event": "webui_response",
|
||||||
|
"request_id": "request-2",
|
||||||
|
"ok": False,
|
||||||
|
"error": {"status": 400, "message": "invalid settings payload"},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_webui_request_requires_bootstrap_authenticated_connection(
|
||||||
|
bus: MagicMock,
|
||||||
|
) -> None:
|
||||||
|
channel = _ch(bus)
|
||||||
|
conn = AsyncMock()
|
||||||
|
channel.gateway.http.dispatch_webui_mutation = AsyncMock()
|
||||||
|
|
||||||
|
await channel._dispatch_envelope(
|
||||||
|
conn,
|
||||||
|
"static-token-client",
|
||||||
|
{
|
||||||
|
"type": "webui_request",
|
||||||
|
"request_id": "request-3",
|
||||||
|
"action": "settings.agent.update",
|
||||||
|
"payload": {},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
channel.gateway.http.dispatch_webui_mutation.assert_not_awaited()
|
||||||
|
assert json.loads(conn.send.await_args.args[0]) == {
|
||||||
|
"event": "webui_response",
|
||||||
|
"request_id": "request-3",
|
||||||
|
"ok": False,
|
||||||
|
"error": {"status": 403, "message": "access_denied"},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_webui_persists_sidebar_state_larger_than_http_request_line(
|
async def test_webui_persists_sidebar_state_larger_than_http_request_line(
|
||||||
bus: MagicMock,
|
bus: MagicMock,
|
||||||
@@ -866,23 +1009,90 @@ async def test_webui_persists_sidebar_state_larger_than_http_request_line(
|
|||||||
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
|
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
|
||||||
channel = _ch(bus)
|
channel = _ch(bus)
|
||||||
conn = AsyncMock()
|
conn = AsyncMock()
|
||||||
|
conn.request = SimpleNamespace(headers=Headers())
|
||||||
channel._webui_connections.add(conn)
|
channel._webui_connections.add(conn)
|
||||||
session_order = [f"websocket:{index:04d}-{'x' * 48}" for index in range(160)]
|
session_order = [f"websocket:{index:04d}-{'x' * 48}" for index in range(160)]
|
||||||
|
request_id = "sidebar-large-state"
|
||||||
envelope = {
|
envelope = {
|
||||||
"type": "set_sidebar_state",
|
"type": "webui_request",
|
||||||
"state": {
|
"request_id": request_id,
|
||||||
|
"action": "sidebar.update",
|
||||||
|
"payload": {"state": {
|
||||||
"session_order": session_order,
|
"session_order": session_order,
|
||||||
"view": {"sort": "manual"},
|
"view": {"sort": "manual"},
|
||||||
},
|
}},
|
||||||
}
|
}
|
||||||
assert len(json.dumps(envelope).encode()) > 8_192
|
assert len(json.dumps(envelope).encode()) > 8_192
|
||||||
|
|
||||||
await channel._dispatch_envelope(conn, "webui-client", envelope)
|
await channel._dispatch_envelope(conn, "webui-client", envelope)
|
||||||
|
await asyncio.gather(*tuple(channel._webui_request_tasks.values()))
|
||||||
|
|
||||||
saved = json.loads((tmp_path / "webui" / "sidebar-state.json").read_text(encoding="utf-8"))
|
saved = json.loads((tmp_path / "webui" / "sidebar-state.json").read_text(encoding="utf-8"))
|
||||||
assert saved["session_order"] == session_order
|
assert saved["session_order"] == session_order
|
||||||
assert saved["view"]["sort"] == "manual"
|
assert saved["view"]["sort"] == "manual"
|
||||||
conn.send.assert_not_awaited()
|
assert json.loads(conn.send.await_args.args[0]) == {
|
||||||
|
"event": "webui_response",
|
||||||
|
"request_id": request_id,
|
||||||
|
"ok": True,
|
||||||
|
"result": saved,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_webui_sidebar_state_update_broadcasts_workbench_to_other_devices(
|
||||||
|
bus: MagicMock,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
|
||||||
|
channel = _ch(bus)
|
||||||
|
source = AsyncMock()
|
||||||
|
source.request = SimpleNamespace(headers=Headers())
|
||||||
|
other_device = AsyncMock()
|
||||||
|
channel._webui_connections.update({source, other_device})
|
||||||
|
request_id = "sidebar-workbench-state"
|
||||||
|
|
||||||
|
await channel._dispatch_envelope(
|
||||||
|
source,
|
||||||
|
"webui-client",
|
||||||
|
{
|
||||||
|
"type": "webui_request",
|
||||||
|
"request_id": request_id,
|
||||||
|
"action": "sidebar.update",
|
||||||
|
"payload": {
|
||||||
|
"state": {
|
||||||
|
"workbench": {
|
||||||
|
"version": 1,
|
||||||
|
"tabs": {
|
||||||
|
"tab:websocket:a": {
|
||||||
|
"explicit": True,
|
||||||
|
"title": "Research",
|
||||||
|
"paneKeys": ["websocket:a", "websocket:b"],
|
||||||
|
"layoutPaneKeys": ["websocket:b", "websocket:a"],
|
||||||
|
"layout": "columns",
|
||||||
|
"splitRatios": [0.35],
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
await asyncio.gather(*tuple(channel._webui_request_tasks.values()))
|
||||||
|
|
||||||
|
event = json.loads(other_device.send.await_args.args[0])
|
||||||
|
assert event["event"] == "sidebar_state_updated"
|
||||||
|
assert event["state"]["workbench"]["tabs"]["tab:websocket:a"]["paneKeys"] == [
|
||||||
|
"websocket:a",
|
||||||
|
"websocket:b",
|
||||||
|
]
|
||||||
|
assert event["state"]["workbench"]["tabs"]["tab:websocket:a"]["layoutPaneKeys"] == [
|
||||||
|
"websocket:b",
|
||||||
|
"websocket:a",
|
||||||
|
]
|
||||||
|
assert event["state"]["workbench"]["tabs"]["tab:websocket:a"]["splitRatios"] == [
|
||||||
|
0.35
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -2721,10 +2931,13 @@ async def test_end_to_end_client_receives_ready_and_agent_sees_inbound(bus: Magi
|
|||||||
channel = _ch(bus, port=port)
|
channel = _ch(bus, port=port)
|
||||||
|
|
||||||
server_task = asyncio.create_task(channel.start())
|
server_task = asyncio.create_task(channel.start())
|
||||||
await asyncio.sleep(0.3)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
async with websockets.connect(f"ws://127.0.0.1:{port}/ws?client_id=tester") as client:
|
client = await asyncio.wait_for(
|
||||||
|
_connect_when_ready(f"ws://127.0.0.1:{port}/ws?client_id=tester"),
|
||||||
|
timeout=5,
|
||||||
|
)
|
||||||
|
async with client:
|
||||||
ready_raw = await client.recv()
|
ready_raw = await client.recv()
|
||||||
ready = json.loads(ready_raw)
|
ready = json.loads(ready_raw)
|
||||||
assert ready["event"] == "ready"
|
assert ready["event"] == "ready"
|
||||||
@@ -2887,7 +3100,15 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
server_task = asyncio.create_task(channel.start())
|
server_task = asyncio.create_task(channel.start())
|
||||||
await asyncio.sleep(0.3)
|
await asyncio.sleep(0.3)
|
||||||
|
|
||||||
|
webui_client = None
|
||||||
try:
|
try:
|
||||||
|
webui_token = channel.gateway.tokens.issue_token(300, audience="webui")
|
||||||
|
webui_client = await websockets.connect(
|
||||||
|
f"ws://127.0.0.1:{port}/ws?token={webui_token}&client_id=settings-test"
|
||||||
|
)
|
||||||
|
ready = json.loads(await asyncio.wait_for(webui_client.recv(), timeout=5))
|
||||||
|
assert ready["event"] == "ready"
|
||||||
|
|
||||||
settings = await _http_get(
|
settings = await _http_get(
|
||||||
f"http://127.0.0.1:{port}/api/settings",
|
f"http://127.0.0.1:{port}/api/settings",
|
||||||
headers={"Authorization": "Bearer tok"},
|
headers={"Authorization": "Bearer tok"},
|
||||||
@@ -2971,11 +3192,14 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert unknown_api.status_code == 404
|
assert unknown_api.status_code == 404
|
||||||
assert "<!doctype html>" not in unknown_api.text.lower()
|
assert "<!doctype html>" not in unknown_api.text.lower()
|
||||||
|
|
||||||
provider_updated = await _http_get(
|
provider_updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/provider/update?provider=openrouter"
|
"settings.provider.update",
|
||||||
"&api_key=sk-or-test&api_base=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1",
|
{
|
||||||
headers={"Authorization": "Bearer tok"},
|
"provider": "openrouter",
|
||||||
|
"apiKey": "sk-or-test",
|
||||||
|
"apiBase": "https://openrouter.ai/api/v1",
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert provider_updated.status_code == 200
|
assert provider_updated.status_code == 200
|
||||||
provider_body = provider_updated.json()
|
provider_body = provider_updated.json()
|
||||||
@@ -2985,11 +3209,9 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert provider_body["image_generation"]["provider_configured"] is True
|
assert provider_body["image_generation"]["provider_configured"] is True
|
||||||
assert "sk-or-test" not in provider_updated.text
|
assert "sk-or-test" not in provider_updated.text
|
||||||
|
|
||||||
custom_provider_created = await _http_get(
|
custom_provider_created = await _webui_mutate(
|
||||||
f"http://127.0.0.1:{port}/api/settings/provider/create",
|
webui_client,
|
||||||
headers={
|
"settings.provider.create",
|
||||||
"Authorization": "Bearer tok",
|
|
||||||
"X-Nanobot-Provider-Values": json.dumps(
|
|
||||||
{
|
{
|
||||||
"name": "Company Gateway",
|
"name": "Company Gateway",
|
||||||
"apiBase": "https://gateway.example/v1",
|
"apiBase": "https://gateway.example/v1",
|
||||||
@@ -2999,8 +3221,6 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
"extraQuery": json.dumps({"api-version": "2026-01-01"}),
|
"extraQuery": json.dumps({"api-version": "2026-01-01"}),
|
||||||
"proxy": "http://127.0.0.1:7890",
|
"proxy": "http://127.0.0.1:7890",
|
||||||
"thinkingStyle": "enable_thinking",
|
"thinkingStyle": "enable_thinking",
|
||||||
}
|
|
||||||
),
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
assert custom_provider_created.status_code == 200
|
assert custom_provider_created.status_code == 200
|
||||||
@@ -3015,11 +3235,10 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
}
|
}
|
||||||
assert "sk-company" not in custom_provider_created.text
|
assert "sk-company" not in custom_provider_created.text
|
||||||
|
|
||||||
local_provider_updated = await _http_get(
|
local_provider_updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/provider/update?provider=atomic_chat"
|
"settings.provider.update",
|
||||||
"&api_base=http%3A%2F%2Flocalhost%3A1337%2Fv1",
|
{"provider": "atomic_chat", "apiBase": "http://localhost:1337/v1"},
|
||||||
headers={"Authorization": "Bearer tok"},
|
|
||||||
)
|
)
|
||||||
assert local_provider_updated.status_code == 200
|
assert local_provider_updated.status_code == 200
|
||||||
local_provider_body = local_provider_updated.json()
|
local_provider_body = local_provider_updated.json()
|
||||||
@@ -3029,38 +3248,44 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert local_provider_rows["atomic_chat"]["configured"] is True
|
assert local_provider_rows["atomic_chat"]["configured"] is True
|
||||||
assert "localhost:1337" in local_provider_updated.text
|
assert "localhost:1337" in local_provider_updated.text
|
||||||
|
|
||||||
updated = await _http_get(
|
updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/update?model=atomic_chat/test"
|
"settings.agent.update",
|
||||||
"&provider=atomic_chat&timezone=Asia%2FShanghai"
|
{
|
||||||
"&bot_name=Nano&bot_icon=N&tool_hint_max_length=120",
|
"model": "atomic_chat/test",
|
||||||
headers={"Authorization": "Bearer tok"},
|
"provider": "atomic_chat",
|
||||||
|
"timezone": "Asia/Shanghai",
|
||||||
|
"tool_hint_max_length": 120,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert updated.status_code == 200
|
assert updated.status_code == 200
|
||||||
updated_body = updated.json()
|
updated_body = updated.json()
|
||||||
assert updated_body["requires_restart"] is True
|
assert updated_body["requires_restart"] is True
|
||||||
assert updated_body["restart_required_sections"] == ["runtime"]
|
assert updated_body["restart_required_sections"] == ["runtime"]
|
||||||
|
|
||||||
preset_updated = await _http_get(
|
preset_updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/update?model_preset=deep",
|
"settings.agent.update",
|
||||||
headers={"Authorization": "Bearer tok"},
|
{"model_preset": "deep"},
|
||||||
)
|
)
|
||||||
assert preset_updated.status_code == 200
|
assert preset_updated.status_code == 200
|
||||||
assert preset_updated.json()["agent"]["model"] == "anthropic/claude-opus-4-5"
|
assert preset_updated.json()["agent"]["model"] == "anthropic/claude-opus-4-5"
|
||||||
|
|
||||||
bad_preset = await _http_get(
|
bad_preset = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/update?model_preset=missing",
|
"settings.agent.update",
|
||||||
headers={"Authorization": "Bearer tok"},
|
{"model_preset": "missing"},
|
||||||
)
|
)
|
||||||
assert bad_preset.status_code == 400
|
assert bad_preset.status_code == 400
|
||||||
|
|
||||||
created_preset = await _http_get(
|
created_preset = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/model-configurations/create"
|
"settings.model_configuration.create",
|
||||||
"?label=Fast%20writing&provider=openai&model=openai%2Fgpt-4.1-mini",
|
{
|
||||||
headers={"Authorization": "Bearer tok"},
|
"label": "Fast writing",
|
||||||
|
"provider": "openai",
|
||||||
|
"model": "openai/gpt-4.1-mini",
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert created_preset.status_code == 200
|
assert created_preset.status_code == 200
|
||||||
created_body = created_preset.json()
|
created_body = created_preset.json()
|
||||||
@@ -3074,11 +3299,15 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert created_presets["fast-writing"]["label"] == "Fast writing"
|
assert created_presets["fast-writing"]["label"] == "Fast writing"
|
||||||
assert created_presets["fast-writing"]["provider"] == "openai"
|
assert created_presets["fast-writing"]["provider"] == "openai"
|
||||||
|
|
||||||
updated_preset = await _http_get(
|
updated_preset = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/model-configurations/update"
|
"settings.model_configuration.update",
|
||||||
"?name=fast-writing&label=Codex&provider=openai&model=openai%2Fgpt-5.5",
|
{
|
||||||
headers={"Authorization": "Bearer tok"},
|
"name": "fast-writing",
|
||||||
|
"label": "Codex",
|
||||||
|
"provider": "openai",
|
||||||
|
"model": "openai/gpt-5.5",
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert updated_preset.status_code == 200
|
assert updated_preset.status_code == 200
|
||||||
updated_preset_body = updated_preset.json()
|
updated_preset_body = updated_preset.json()
|
||||||
@@ -3089,11 +3318,10 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
}
|
}
|
||||||
assert updated_presets["fast-writing"]["label"] == "Codex"
|
assert updated_presets["fast-writing"]["label"] == "Codex"
|
||||||
|
|
||||||
call_order_updated = await _http_get(
|
call_order_updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/model-call-order/update"
|
"settings.model_call_order.update",
|
||||||
"?order=%5B%22fast-writing%22%2C%22deep%22%5D",
|
{"order": ["fast-writing", "deep"]},
|
||||||
headers={"Authorization": "Bearer tok"},
|
|
||||||
)
|
)
|
||||||
assert call_order_updated.status_code == 200
|
assert call_order_updated.status_code == 200
|
||||||
call_order_body = call_order_updated.json()
|
call_order_body = call_order_updated.json()
|
||||||
@@ -3101,20 +3329,27 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert call_order_body["agent"]["model"] == "openai/gpt-5.5"
|
assert call_order_body["agent"]["model"] == "openai/gpt-5.5"
|
||||||
assert call_order_body["model_call_order"] == ["fast-writing", "deep"]
|
assert call_order_body["model_call_order"] == ["fast-writing", "deep"]
|
||||||
|
|
||||||
duplicate_preset = await _http_get(
|
duplicate_preset = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/model-configurations/create"
|
"settings.model_configuration.create",
|
||||||
"?label=Fast%20writing&provider=openai&model=openai%2Fgpt-4.1-mini",
|
{
|
||||||
headers={"Authorization": "Bearer tok"},
|
"label": "Fast writing",
|
||||||
|
"provider": "openai",
|
||||||
|
"model": "openai/gpt-4.1-mini",
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert duplicate_preset.status_code == 409
|
assert duplicate_preset.status_code == 409
|
||||||
|
|
||||||
search_updated = await _http_get(
|
search_updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/web-search/update?provider=searxng"
|
"settings.web_search.update",
|
||||||
"&base_url=https%3A%2F%2Fsearch.example.com"
|
{
|
||||||
"&max_results=8&timeout=45&use_jina_reader=false",
|
"provider": "searxng",
|
||||||
headers={"Authorization": "Bearer tok"},
|
"base_url": "https://search.example.com",
|
||||||
|
"max_results": 8,
|
||||||
|
"timeout": 45,
|
||||||
|
"use_jina_reader": False,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert search_updated.status_code == 200
|
assert search_updated.status_code == 200
|
||||||
search_body = search_updated.json()
|
search_body = search_updated.json()
|
||||||
@@ -3126,10 +3361,13 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert search_body["web_search"]["max_results"] == 8
|
assert search_body["web_search"]["max_results"] == 8
|
||||||
assert search_body["web"]["fetch"]["use_jina_reader"] is False
|
assert search_body["web"]["fetch"]["use_jina_reader"] is False
|
||||||
|
|
||||||
network_safety_updated = await _http_get(
|
network_safety_updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/network-safety/update?webui_allow_local_service_access=false&webui_default_access_mode=full",
|
"settings.network_safety.update",
|
||||||
headers={"Authorization": "Bearer tok"},
|
{
|
||||||
|
"webui_allow_local_service_access": False,
|
||||||
|
"webui_default_access_mode": "full",
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert network_safety_updated.status_code == 200
|
assert network_safety_updated.status_code == 200
|
||||||
network_safety_body = network_safety_updated.json()
|
network_safety_body = network_safety_updated.json()
|
||||||
@@ -3139,13 +3377,17 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert network_safety_body["advanced"]["webui_default_access_mode"] == "full"
|
assert network_safety_body["advanced"]["webui_default_access_mode"] == "full"
|
||||||
assert network_safety_body["advanced"]["private_service_protection_enabled"] is True
|
assert network_safety_body["advanced"]["private_service_protection_enabled"] is True
|
||||||
|
|
||||||
image_updated = await _http_get(
|
image_updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/image-generation/update?enabled=true"
|
"settings.image_generation.update",
|
||||||
"&provider=openrouter&model=openai%2Fgpt-image-1"
|
{
|
||||||
"&default_aspect_ratio=16%3A9&default_image_size=2K"
|
"enabled": True,
|
||||||
"&max_images_per_turn=3",
|
"provider": "openrouter",
|
||||||
headers={"Authorization": "Bearer tok"},
|
"model": "openai/gpt-image-1",
|
||||||
|
"default_aspect_ratio": "16:9",
|
||||||
|
"default_image_size": "2K",
|
||||||
|
"max_images_per_turn": 3,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert image_updated.status_code == 200
|
assert image_updated.status_code == 200
|
||||||
image_body = image_updated.json()
|
image_body = image_updated.json()
|
||||||
@@ -3157,11 +3399,14 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert image_body["image_generation"]["default_image_size"] == "2K"
|
assert image_body["image_generation"]["default_image_size"] == "2K"
|
||||||
assert image_body["image_generation"]["max_images_per_turn"] == 3
|
assert image_body["image_generation"]["max_images_per_turn"] == 3
|
||||||
|
|
||||||
image_provider_updated = await _http_get(
|
image_provider_updated = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/provider/update?provider=openrouter"
|
"settings.provider.update",
|
||||||
"&api_key=sk-or-next&api_base=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1",
|
{
|
||||||
headers={"Authorization": "Bearer tok"},
|
"provider": "openrouter",
|
||||||
|
"apiKey": "sk-or-next",
|
||||||
|
"apiBase": "https://openrouter.ai/api/v1",
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert image_provider_updated.status_code == 200
|
assert image_provider_updated.status_code == 200
|
||||||
assert image_provider_updated.json()["requires_restart"] is True
|
assert image_provider_updated.json()["requires_restart"] is True
|
||||||
@@ -3169,17 +3414,17 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert "sk-or-next" not in image_provider_updated.text
|
assert "sk-or-next" not in image_provider_updated.text
|
||||||
assert image_reload.await_count == 2
|
assert image_reload.await_count == 2
|
||||||
|
|
||||||
bad_web = await _http_get(
|
bad_web = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/web-search/update?provider=duckduckgo&max_results=99",
|
"settings.web_search.update",
|
||||||
headers={"Authorization": "Bearer tok"},
|
{"provider": "duckduckgo", "max_results": 99},
|
||||||
)
|
)
|
||||||
assert bad_web.status_code == 400
|
assert bad_web.status_code == 400
|
||||||
|
|
||||||
bad_image = await _http_get(
|
bad_image = await _webui_mutate(
|
||||||
"http://127.0.0.1:"
|
webui_client,
|
||||||
f"{port}/api/settings/image-generation/update?provider=missing",
|
"settings.image_generation.update",
|
||||||
headers={"Authorization": "Bearer tok"},
|
{"provider": "missing"},
|
||||||
)
|
)
|
||||||
assert bad_image.status_code == 400
|
assert bad_image.status_code == 400
|
||||||
|
|
||||||
@@ -3216,6 +3461,8 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert saved.tools.image_generation.default_image_size == "2K"
|
assert saved.tools.image_generation.default_image_size == "2K"
|
||||||
assert saved.tools.image_generation.max_images_per_turn == 3
|
assert saved.tools.image_generation.max_images_per_turn == 3
|
||||||
finally:
|
finally:
|
||||||
|
if webui_client is not None:
|
||||||
|
await webui_client.close()
|
||||||
await channel.stop()
|
await channel.stop()
|
||||||
await server_task
|
await server_task
|
||||||
|
|
||||||
@@ -3248,11 +3495,17 @@ async def test_image_settings_hot_reload_without_restart(
|
|||||||
channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300
|
channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300
|
||||||
server_task = asyncio.create_task(channel.start())
|
server_task = asyncio.create_task(channel.start())
|
||||||
await asyncio.sleep(0.3)
|
await asyncio.sleep(0.3)
|
||||||
|
webui_client = None
|
||||||
try:
|
try:
|
||||||
response = await _http_get(
|
webui_token = channel.gateway.tokens.issue_token(300, audience="webui")
|
||||||
f"http://127.0.0.1:{port}/api/settings/image-generation/update"
|
webui_client = await websockets.connect(
|
||||||
"?enabled=true&provider=openrouter&model=openai%2Fgpt-image-1",
|
f"ws://127.0.0.1:{port}/ws?token={webui_token}&client_id=image-reload-test"
|
||||||
headers={"Authorization": "Bearer tok"},
|
)
|
||||||
|
assert json.loads(await webui_client.recv())["event"] == "ready"
|
||||||
|
response = await _webui_mutate(
|
||||||
|
webui_client,
|
||||||
|
"settings.image_generation.update",
|
||||||
|
{"enabled": True, "provider": "openrouter", "model": "openai/gpt-image-1"},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -3260,6 +3513,8 @@ async def test_image_settings_hot_reload_without_restart(
|
|||||||
assert response.json()["restart_required_sections"] == []
|
assert response.json()["restart_required_sections"] == []
|
||||||
image_reload.assert_awaited_once_with(bus)
|
image_reload.assert_awaited_once_with(bus)
|
||||||
finally:
|
finally:
|
||||||
|
if webui_client is not None:
|
||||||
|
await webui_client.close()
|
||||||
await channel.stop()
|
await channel.stop()
|
||||||
await server_task
|
await server_task
|
||||||
|
|
||||||
@@ -3291,17 +3546,25 @@ async def test_image_settings_fall_back_to_restart_when_hot_reload_fails(
|
|||||||
channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300
|
channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300
|
||||||
server_task = asyncio.create_task(channel.start())
|
server_task = asyncio.create_task(channel.start())
|
||||||
await asyncio.sleep(0.3)
|
await asyncio.sleep(0.3)
|
||||||
|
webui_client = None
|
||||||
try:
|
try:
|
||||||
response = await _http_get(
|
webui_token = channel.gateway.tokens.issue_token(300, audience="webui")
|
||||||
f"http://127.0.0.1:{port}/api/settings/image-generation/update"
|
webui_client = await websockets.connect(
|
||||||
"?enabled=true&provider=openrouter&model=openai%2Fgpt-image-1",
|
f"ws://127.0.0.1:{port}/ws?token={webui_token}&client_id=image-fallback-test"
|
||||||
headers={"Authorization": "Bearer tok"},
|
)
|
||||||
|
assert json.loads(await webui_client.recv())["event"] == "ready"
|
||||||
|
response = await _webui_mutate(
|
||||||
|
webui_client,
|
||||||
|
"settings.image_generation.update",
|
||||||
|
{"enabled": True, "provider": "openrouter", "model": "openai/gpt-image-1"},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
assert response.json()["requires_restart"] is True
|
assert response.json()["requires_restart"] is True
|
||||||
assert response.json()["restart_required_sections"] == ["image"]
|
assert response.json()["restart_required_sections"] == ["image"]
|
||||||
finally:
|
finally:
|
||||||
|
if webui_client is not None:
|
||||||
|
await webui_client.close()
|
||||||
await channel.stop()
|
await channel.stop()
|
||||||
await server_task
|
await server_task
|
||||||
|
|
||||||
@@ -3541,6 +3804,7 @@ async def test_token_issue_rejects_when_at_capacity(bus: MagicMock) -> None:
|
|||||||
headers={"Authorization": "Bearer s"},
|
headers={"Authorization": "Bearer s"},
|
||||||
)
|
)
|
||||||
assert resp.status_code == 429
|
assert resp.status_code == 429
|
||||||
|
assert resp.headers["Cache-Control"] == "no-store"
|
||||||
data = resp.json()
|
data = resp.json()
|
||||||
assert "error" in data
|
assert "error" in data
|
||||||
finally:
|
finally:
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -22,6 +22,7 @@ class WeixinConnectSession:
|
|||||||
channel: WeixinChannel
|
channel: WeixinChannel
|
||||||
current_poll_base_url: str
|
current_poll_base_url: str
|
||||||
refresh_count: int
|
refresh_count: int
|
||||||
|
force: bool
|
||||||
created_wall: float
|
created_wall: float
|
||||||
deadline: float
|
deadline: float
|
||||||
last_error: str | None = None
|
last_error: str | None = None
|
||||||
@@ -72,7 +73,7 @@ class WeixinConnectStore:
|
|||||||
|
|
||||||
channel.connect_open_client()
|
channel.connect_open_client()
|
||||||
try:
|
try:
|
||||||
qrcode_id, qr_url = await channel.connect_fetch_qr_code()
|
qrcode_id, qr_url = await channel.connect_fetch_qr_code(force=force)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
await self._close_channel(channel)
|
await self._close_channel(channel)
|
||||||
raise ChannelConnectError(
|
raise ChannelConnectError(
|
||||||
@@ -89,6 +90,7 @@ class WeixinConnectStore:
|
|||||||
channel=channel,
|
channel=channel,
|
||||||
current_poll_base_url=channel.connect_base_url,
|
current_poll_base_url=channel.connect_base_url,
|
||||||
refresh_count=0,
|
refresh_count=0,
|
||||||
|
force=force,
|
||||||
created_wall=now_wall,
|
created_wall=now_wall,
|
||||||
deadline=time.monotonic() + 600,
|
deadline=time.monotonic() + 600,
|
||||||
)
|
)
|
||||||
@@ -187,7 +189,7 @@ class WeixinConnectStore:
|
|||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
session.qrcode_id, session.qr_url = (
|
session.qrcode_id, session.qr_url = (
|
||||||
await session.channel.connect_fetch_qr_code()
|
await session.channel.connect_fetch_qr_code(force=session.force)
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self._sessions.pop(session_id, None)
|
self._sessions.pop(session_id, None)
|
||||||
@@ -204,6 +206,17 @@ class WeixinConnectStore:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if status == "binded_redirect":
|
if status == "binded_redirect":
|
||||||
|
if session.force:
|
||||||
|
self._sessions.pop(session_id, None)
|
||||||
|
await self._close_channel(session.channel)
|
||||||
|
return {
|
||||||
|
"session_id": session_id,
|
||||||
|
"status": "failed",
|
||||||
|
"message": (
|
||||||
|
"Unable to complete a new WeChat login. "
|
||||||
|
"Start again and scan with the account you want to connect."
|
||||||
|
),
|
||||||
|
}
|
||||||
if not session.channel.connect_load_state():
|
if not session.channel.connect_load_state():
|
||||||
self._sessions.pop(session_id, None)
|
self._sessions.pop(session_id, None)
|
||||||
await self._close_channel(session.channel)
|
await self._close_channel(session.channel)
|
||||||
@@ -234,7 +247,7 @@ class WeixinConnectStore:
|
|||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
session.qrcode_id, session.qr_url = (
|
session.qrcode_id, session.qr_url = (
|
||||||
await session.channel.connect_fetch_qr_code()
|
await session.channel.connect_fetch_qr_code(force=session.force)
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self._sessions.pop(session_id, None)
|
self._sessions.pop(session_id, None)
|
||||||
|
|||||||
@@ -726,9 +726,9 @@ class WeixinChannel(BaseChannel):
|
|||||||
break
|
break
|
||||||
return tokens
|
return tokens
|
||||||
|
|
||||||
async def _fetch_qr_code(self) -> tuple[str, str]:
|
async def _fetch_qr_code(self, *, force: bool = False) -> tuple[str, str]:
|
||||||
"""Fetch a fresh QR code. Returns (qrcode_id, scan_url)."""
|
"""Fetch a QR code without existing credentials when forced."""
|
||||||
local_tokens = self._local_token_list()
|
local_tokens = [] if force else self._local_token_list()
|
||||||
data = await self._api_post(
|
data = await self._api_post(
|
||||||
"ilink/bot/get_bot_qrcode?bot_type=3",
|
"ilink/bot/get_bot_qrcode?bot_type=3",
|
||||||
{"local_token_list": local_tokens},
|
{"local_token_list": local_tokens},
|
||||||
@@ -755,11 +755,11 @@ class WeixinChannel(BaseChannel):
|
|||||||
raise RuntimeError(f"Failed to get QR code from WeChat API: {data}")
|
raise RuntimeError(f"Failed to get QR code from WeChat API: {data}")
|
||||||
return qrcode_id, (qrcode_img_content or qrcode_id)
|
return qrcode_id, (qrcode_img_content or qrcode_id)
|
||||||
|
|
||||||
async def _qr_login(self) -> bool:
|
async def _qr_login(self, *, force: bool = False) -> bool:
|
||||||
"""Perform QR code login flow. Returns True on success."""
|
"""Perform QR login; forced flows accept only newly confirmed credentials."""
|
||||||
try:
|
try:
|
||||||
refresh_count = 0
|
refresh_count = 0
|
||||||
qrcode_id, scan_url = await self._fetch_qr_code()
|
qrcode_id, scan_url = await self._fetch_qr_code(force=force)
|
||||||
self._print_qr_code(scan_url)
|
self._print_qr_code(scan_url)
|
||||||
current_poll_base_url = self.config.base_url
|
current_poll_base_url = self.config.base_url
|
||||||
verify_code = ""
|
verify_code = ""
|
||||||
@@ -825,11 +825,16 @@ class WeixinChannel(BaseChannel):
|
|||||||
if refresh_count > MAX_QR_REFRESH_COUNT:
|
if refresh_count > MAX_QR_REFRESH_COUNT:
|
||||||
self.logger.warning("WeChat verification failed too many times")
|
self.logger.warning("WeChat verification failed too many times")
|
||||||
return False
|
return False
|
||||||
qrcode_id, scan_url = await self._fetch_qr_code()
|
qrcode_id, scan_url = await self._fetch_qr_code(force=force)
|
||||||
current_poll_base_url = self.config.base_url
|
current_poll_base_url = self.config.base_url
|
||||||
self._print_qr_code(scan_url)
|
self._print_qr_code(scan_url)
|
||||||
continue
|
continue
|
||||||
elif status == "binded_redirect":
|
elif status == "binded_redirect":
|
||||||
|
if force:
|
||||||
|
self.logger.error(
|
||||||
|
"Forced WeChat login returned an existing binding without new credentials"
|
||||||
|
)
|
||||||
|
return False
|
||||||
if self._token or self._load_state():
|
if self._token or self._load_state():
|
||||||
self.logger.info("WeChat account is already connected")
|
self.logger.info("WeChat account is already connected")
|
||||||
return True
|
return True
|
||||||
@@ -846,7 +851,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
MAX_QR_REFRESH_COUNT,
|
MAX_QR_REFRESH_COUNT,
|
||||||
)
|
)
|
||||||
return False
|
return False
|
||||||
qrcode_id, scan_url = await self._fetch_qr_code()
|
qrcode_id, scan_url = await self._fetch_qr_code(force=force)
|
||||||
current_poll_base_url = self.config.base_url
|
current_poll_base_url = self.config.base_url
|
||||||
verify_code = ""
|
verify_code = ""
|
||||||
self._print_qr_code(scan_url)
|
self._print_qr_code(scan_url)
|
||||||
@@ -893,8 +898,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
self._client = self._new_http_client(httpx.Timeout(60, connect=30))
|
self._client = self._new_http_client(httpx.Timeout(60, connect=30))
|
||||||
self._running = True
|
self._running = True
|
||||||
|
|
||||||
async def connect_fetch_qr_code(self) -> tuple[str, str]:
|
async def connect_fetch_qr_code(self, *, force: bool = False) -> tuple[str, str]:
|
||||||
return await self._fetch_qr_code()
|
return await self._fetch_qr_code(force=force)
|
||||||
|
|
||||||
async def connect_poll_qr_code(
|
async def connect_poll_qr_code(
|
||||||
self,
|
self,
|
||||||
@@ -947,14 +952,14 @@ class WeixinChannel(BaseChannel):
|
|||||||
if force:
|
if force:
|
||||||
self._token = ""
|
self._token = ""
|
||||||
self._get_updates_buf = ""
|
self._get_updates_buf = ""
|
||||||
if self._token or self._load_state():
|
if self._token or (not force and self._load_state()):
|
||||||
return True
|
return True
|
||||||
|
|
||||||
# Initialize HTTP client for the login flow
|
# Initialize HTTP client for the login flow
|
||||||
self._client = self._new_http_client(httpx.Timeout(60, connect=30))
|
self._client = self._new_http_client(httpx.Timeout(60, connect=30))
|
||||||
self._running = True # Enable polling loop in _qr_login()
|
self._running = True # Enable polling loop in _qr_login()
|
||||||
try:
|
try:
|
||||||
return await self._qr_login()
|
return await self._qr_login(force=force)
|
||||||
finally:
|
finally:
|
||||||
self._running = False
|
self._running = False
|
||||||
if self._client:
|
if self._client:
|
||||||
|
|||||||
@@ -25,7 +25,9 @@ async def test_weixin_connect_store_saves_confirmed_qr_login(
|
|||||||
)
|
)
|
||||||
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
||||||
|
|
||||||
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
|
async def fake_fetch_qr_code(
|
||||||
|
self: WeixinChannel, **_kwargs: Any
|
||||||
|
) -> tuple[str, str]:
|
||||||
return "qr-1", "https://qr.example/1"
|
return "qr-1", "https://qr.example/1"
|
||||||
|
|
||||||
async def fake_api_get_with_base(
|
async def fake_api_get_with_base(
|
||||||
@@ -86,14 +88,31 @@ async def test_weixin_reconnect_keeps_existing_account_until_scan_succeeds(
|
|||||||
)
|
)
|
||||||
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
||||||
|
|
||||||
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
|
observed_force: list[bool] = []
|
||||||
return "qr-reconnect", "https://qr.example/reconnect"
|
|
||||||
|
async def fake_fetch_qr_code(
|
||||||
|
self: WeixinChannel,
|
||||||
|
*,
|
||||||
|
force: bool = False,
|
||||||
|
) -> tuple[str, str]:
|
||||||
|
observed_force.append(force)
|
||||||
|
return f"qr-reconnect-{len(observed_force)}", "https://qr.example/reconnect"
|
||||||
|
|
||||||
|
async def fake_api_get_with_base(
|
||||||
|
self: WeixinChannel,
|
||||||
|
**_kwargs: Any,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
return {"status": "expired"}
|
||||||
|
|
||||||
monkeypatch.setattr(WeixinChannel, "_fetch_qr_code", fake_fetch_qr_code)
|
monkeypatch.setattr(WeixinChannel, "_fetch_qr_code", fake_fetch_qr_code)
|
||||||
|
monkeypatch.setattr(WeixinChannel, "_api_get_with_base", fake_api_get_with_base)
|
||||||
|
|
||||||
store = WeixinConnectStore()
|
store = WeixinConnectStore()
|
||||||
started = await store.start(force=True)
|
started = await store.start(force=True)
|
||||||
|
refreshed = await store.poll(started["session_id"])
|
||||||
|
|
||||||
|
assert refreshed["status"] == "pending"
|
||||||
|
assert observed_force == [True, True]
|
||||||
assert json.loads(state_file.read_text(encoding="utf-8")) == existing
|
assert json.loads(state_file.read_text(encoding="utf-8")) == existing
|
||||||
cancelled = await store.cancel(started["session_id"])
|
cancelled = await store.cancel(started["session_id"])
|
||||||
assert cancelled["status"] == "cancelled"
|
assert cancelled["status"] == "cancelled"
|
||||||
@@ -116,7 +135,9 @@ async def test_weixin_cancel_wins_over_inflight_confirmation(
|
|||||||
poll_started = asyncio.Event()
|
poll_started = asyncio.Event()
|
||||||
release_poll = asyncio.Event()
|
release_poll = asyncio.Event()
|
||||||
|
|
||||||
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
|
async def fake_fetch_qr_code(
|
||||||
|
self: WeixinChannel, **_kwargs: Any
|
||||||
|
) -> tuple[str, str]:
|
||||||
return "qr-cancel", "https://qr.example/cancel"
|
return "qr-cancel", "https://qr.example/cancel"
|
||||||
|
|
||||||
async def fake_api_get_with_base(
|
async def fake_api_get_with_base(
|
||||||
@@ -162,7 +183,9 @@ async def test_weixin_connect_store_handles_verification_code(
|
|||||||
)
|
)
|
||||||
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
||||||
|
|
||||||
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
|
async def fake_fetch_qr_code(
|
||||||
|
self: WeixinChannel, **_kwargs: Any
|
||||||
|
) -> tuple[str, str]:
|
||||||
return "qr-verify", "https://qr.example/verify"
|
return "qr-verify", "https://qr.example/verify"
|
||||||
|
|
||||||
responses = [
|
responses = [
|
||||||
@@ -204,7 +227,7 @@ async def test_weixin_connect_store_handles_verification_code(
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_weixin_connect_store_treats_existing_binding_as_success(
|
async def test_weixin_connect_store_rejects_existing_binding_during_forced_login(
|
||||||
tmp_path,
|
tmp_path,
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -221,7 +244,12 @@ async def test_weixin_connect_store_treats_existing_binding_as_success(
|
|||||||
)
|
)
|
||||||
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
||||||
|
|
||||||
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
|
async def fake_fetch_qr_code(
|
||||||
|
self: WeixinChannel,
|
||||||
|
*,
|
||||||
|
force: bool = False,
|
||||||
|
) -> tuple[str, str]:
|
||||||
|
assert force is True
|
||||||
return "qr-existing", "https://qr.example/existing"
|
return "qr-existing", "https://qr.example/existing"
|
||||||
|
|
||||||
async def fake_api_get_with_base(
|
async def fake_api_get_with_base(
|
||||||
@@ -237,8 +265,8 @@ async def test_weixin_connect_store_treats_existing_binding_as_success(
|
|||||||
started = await store.start(force=True)
|
started = await store.start(force=True)
|
||||||
completed = await store.poll(started["session_id"])
|
completed = await store.poll(started["session_id"])
|
||||||
|
|
||||||
assert completed["status"] == "succeeded"
|
assert completed["status"] == "failed"
|
||||||
assert "already connected" in completed["message"]
|
assert "new WeChat login" in completed["message"]
|
||||||
assert json.loads((state_dir / "account.json").read_text())["token"] == "working-token"
|
assert json.loads((state_dir / "account.json").read_text())["token"] == "working-token"
|
||||||
|
|
||||||
|
|
||||||
@@ -255,7 +283,9 @@ async def test_weixin_connect_store_rejects_existing_binding_without_local_crede
|
|||||||
)
|
)
|
||||||
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
||||||
|
|
||||||
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
|
async def fake_fetch_qr_code(
|
||||||
|
self: WeixinChannel, **_kwargs: Any
|
||||||
|
) -> tuple[str, str]:
|
||||||
return "qr-missing", "https://qr.example/missing"
|
return "qr-missing", "https://qr.example/missing"
|
||||||
|
|
||||||
async def fake_api_get_with_base(
|
async def fake_api_get_with_base(
|
||||||
@@ -268,7 +298,7 @@ async def test_weixin_connect_store_rejects_existing_binding_without_local_crede
|
|||||||
monkeypatch.setattr(WeixinChannel, "_api_get_with_base", fake_api_get_with_base)
|
monkeypatch.setattr(WeixinChannel, "_api_get_with_base", fake_api_get_with_base)
|
||||||
|
|
||||||
store = WeixinConnectStore()
|
store = WeixinConnectStore()
|
||||||
started = await store.start(force=True)
|
started = await store.start(force=False)
|
||||||
completed = await store.poll(started["session_id"])
|
completed = await store.poll(started["session_id"])
|
||||||
|
|
||||||
assert completed["status"] == "failed"
|
assert completed["status"] == "failed"
|
||||||
|
|||||||
@@ -196,6 +196,86 @@ def test_save_state_with_empty_runtime_token_preserves_persisted_account(tmp_pat
|
|||||||
assert json.loads((tmp_path / "account.json").read_text()) == persisted
|
assert json.loads((tmp_path / "account.json").read_text()) == persisted
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_login_force_ignores_persisted_account_through_qr_flow(tmp_path) -> None:
|
||||||
|
persisted = {
|
||||||
|
"token": "persisted-token",
|
||||||
|
"get_updates_buf": "persisted-cursor",
|
||||||
|
"context_tokens": {"wx-user": "ctx-persisted"},
|
||||||
|
"typing_tickets": {"wx-user": {"ticket": "ticket-persisted"}},
|
||||||
|
"base_url": "https://persisted.example",
|
||||||
|
}
|
||||||
|
channel = WeixinChannel(
|
||||||
|
WeixinConfig(
|
||||||
|
enabled=True,
|
||||||
|
allow_from=["*"],
|
||||||
|
token="configured-token",
|
||||||
|
state_dir=str(tmp_path),
|
||||||
|
),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
(tmp_path / "account.json").write_text(
|
||||||
|
json.dumps(persisted),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
channel._print_qr_code = lambda _url: None
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
||||||
|
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"status": "expired"},
|
||||||
|
{"status": "binded_redirect"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
ok = await channel.login(force=True)
|
||||||
|
|
||||||
|
assert ok is False
|
||||||
|
assert [call.args[1]["local_token_list"] for call in channel._api_post.await_args_list] == [
|
||||||
|
[],
|
||||||
|
[],
|
||||||
|
]
|
||||||
|
assert channel._token == ""
|
||||||
|
assert channel._get_updates_buf == ""
|
||||||
|
assert channel._context_tokens == {}
|
||||||
|
assert channel._typing_tickets == {}
|
||||||
|
assert channel.config.base_url == "https://ilinkai.weixin.qq.com"
|
||||||
|
assert json.loads((tmp_path / "account.json").read_text()) == persisted
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_login_without_force_reuses_persisted_account(tmp_path) -> None:
|
||||||
|
channel = WeixinChannel(
|
||||||
|
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
(tmp_path / "account.json").write_text(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"token": "persisted-token",
|
||||||
|
"get_updates_buf": "persisted-cursor",
|
||||||
|
"context_tokens": {"wx-user": "ctx-persisted"},
|
||||||
|
"base_url": "https://persisted.example",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
channel._qr_login = AsyncMock(return_value=False)
|
||||||
|
|
||||||
|
ok = await channel.login(force=False)
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
channel._qr_login.assert_not_awaited()
|
||||||
|
assert channel._token == "persisted-token"
|
||||||
|
assert channel._get_updates_buf == "persisted-cursor"
|
||||||
|
assert channel._context_tokens == {"wx-user": "ctx-persisted"}
|
||||||
|
assert channel.config.base_url == "https://persisted.example"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_process_message_deduplicates_inbound_ids() -> None:
|
async def test_process_message_deduplicates_inbound_ids() -> None:
|
||||||
channel, bus = _make_channel()
|
channel, bus = _make_channel()
|
||||||
|
|||||||
@@ -27,33 +27,16 @@ import type {
|
|||||||
NanobotFeatureInfo,
|
NanobotFeatureInfo,
|
||||||
} from "@/lib/types";
|
} from "@/lib/types";
|
||||||
import { cn } from "@/lib/utils";
|
import { cn } from "@/lib/utils";
|
||||||
|
import { useClient } from "@/providers/ClientProvider";
|
||||||
|
|
||||||
import {
|
import {
|
||||||
WEIXIN_AUTH_EXPIRED_MESSAGE,
|
WEIXIN_AUTH_EXPIRED_MESSAGE,
|
||||||
WeixinConnectFlow,
|
WeixinConnectFlow,
|
||||||
} from "./WeixinConnectFlow";
|
} from "./WeixinConnectFlow";
|
||||||
|
import {
|
||||||
export const WEIXIN_PRIMARY_FIELD_KEYS = [
|
WEIXIN_ADVANCED_FIELD_KEYS,
|
||||||
"channels.weixin.sendProgress",
|
WEIXIN_PRIMARY_FIELD_KEYS,
|
||||||
"channels.weixin.sendToolHints",
|
} from "./presentation";
|
||||||
"channels.weixin.streaming",
|
|
||||||
] as const;
|
|
||||||
|
|
||||||
export const WEIXIN_ADVANCED_FIELD_KEYS = [
|
|
||||||
"channels.weixin.allowFrom",
|
|
||||||
"channels.weixin.token",
|
|
||||||
"channels.weixin.replyProgressMessages",
|
|
||||||
"channels.weixin.replyProgressMaxMessages",
|
|
||||||
"channels.weixin.contextMessageBudget",
|
|
||||||
"channels.weixin.blockStreaming",
|
|
||||||
"channels.weixin.blockStreamingMinChars",
|
|
||||||
"channels.weixin.blockStreamingMaxMessages",
|
|
||||||
"channels.weixin.baseUrl",
|
|
||||||
"channels.weixin.cdnBaseUrl",
|
|
||||||
"channels.weixin.routeTag",
|
|
||||||
"channels.weixin.stateDir",
|
|
||||||
"channels.weixin.pollTimeout",
|
|
||||||
] as const;
|
|
||||||
|
|
||||||
export function WeixinPanel({
|
export function WeixinPanel({
|
||||||
token,
|
token,
|
||||||
@@ -64,6 +47,7 @@ export function WeixinPanel({
|
|||||||
onAction,
|
onAction,
|
||||||
onFeaturesUpdate,
|
onFeaturesUpdate,
|
||||||
}: ChannelPluginPanelProps) {
|
}: ChannelPluginPanelProps) {
|
||||||
|
const { client } = useClient();
|
||||||
const { t, i18n } = useTranslation();
|
const { t, i18n } = useTranslation();
|
||||||
const tx = (key: string, fallback: string) => t(key, { defaultValue: fallback });
|
const tx = (key: string, fallback: string) => t(key, { defaultValue: fallback });
|
||||||
const channelTx = channelTranslator(t, "weixin");
|
const channelTx = channelTranslator(t, "weixin");
|
||||||
@@ -150,7 +134,7 @@ export function WeixinPanel({
|
|||||||
setSaveState("idle");
|
setSaveState("idle");
|
||||||
try {
|
try {
|
||||||
const payload = await configureChannel(
|
const payload = await configureChannel(
|
||||||
context.token,
|
client,
|
||||||
"weixin",
|
"weixin",
|
||||||
channelValuesForSave(editableFieldsRef.current, values),
|
channelValuesForSave(editableFieldsRef.current, values),
|
||||||
{ enable: context.enabled },
|
{ enable: context.enabled },
|
||||||
@@ -168,7 +152,7 @@ export function WeixinPanel({
|
|||||||
} finally {
|
} finally {
|
||||||
setSaving(false);
|
setSaving(false);
|
||||||
}
|
}
|
||||||
}, []);
|
}, [client]);
|
||||||
|
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -1,12 +1,21 @@
|
|||||||
|
import { lazy } from "react";
|
||||||
|
|
||||||
import type { ChannelUiContribution } from "@/channel-plugins/types";
|
import type { ChannelUiContribution } from "@/channel-plugins/types";
|
||||||
import { chatAppGuideUrl } from "@/components/settings/channels/catalog";
|
import { chatAppGuideUrl } from "@/components/settings/channels/catalog";
|
||||||
|
|
||||||
import { WeixinConnectFlow } from "./WeixinConnectFlow";
|
|
||||||
import {
|
import {
|
||||||
WEIXIN_ADVANCED_FIELD_KEYS,
|
WEIXIN_ADVANCED_FIELD_KEYS,
|
||||||
WEIXIN_PRIMARY_FIELD_KEYS,
|
WEIXIN_PRIMARY_FIELD_KEYS,
|
||||||
WeixinPanel,
|
} from "./presentation";
|
||||||
} from "./WeixinPanel";
|
|
||||||
|
const WeixinPanel = lazy(() =>
|
||||||
|
import("./WeixinPanel").then(({ WeixinPanel: component }) => ({ default: component })),
|
||||||
|
);
|
||||||
|
const WeixinConnectFlow = lazy(() =>
|
||||||
|
import("./WeixinConnectFlow").then(({ WeixinConnectFlow: component }) => ({
|
||||||
|
default: component,
|
||||||
|
})),
|
||||||
|
);
|
||||||
|
|
||||||
export default {
|
export default {
|
||||||
Panel: WeixinPanel,
|
Panel: WeixinPanel,
|
||||||
|
|||||||
@@ -0,0 +1,21 @@
|
|||||||
|
export const WEIXIN_PRIMARY_FIELD_KEYS = [
|
||||||
|
"channels.weixin.sendProgress",
|
||||||
|
"channels.weixin.sendToolHints",
|
||||||
|
"channels.weixin.streaming",
|
||||||
|
] as const;
|
||||||
|
|
||||||
|
export const WEIXIN_ADVANCED_FIELD_KEYS = [
|
||||||
|
"channels.weixin.allowFrom",
|
||||||
|
"channels.weixin.token",
|
||||||
|
"channels.weixin.replyProgressMessages",
|
||||||
|
"channels.weixin.replyProgressMaxMessages",
|
||||||
|
"channels.weixin.contextMessageBudget",
|
||||||
|
"channels.weixin.blockStreaming",
|
||||||
|
"channels.weixin.blockStreamingMinChars",
|
||||||
|
"channels.weixin.blockStreamingMaxMessages",
|
||||||
|
"channels.weixin.baseUrl",
|
||||||
|
"channels.weixin.cdnBaseUrl",
|
||||||
|
"channels.weixin.routeTag",
|
||||||
|
"channels.weixin.stateDir",
|
||||||
|
"channels.weixin.pollTimeout",
|
||||||
|
] as const;
|
||||||
+17
-2
@@ -13,6 +13,8 @@ from rich.console import Console
|
|||||||
from nanobot import __logo__
|
from nanobot import __logo__
|
||||||
from nanobot.agent.hooks import create_file_edit_activity_hook
|
from nanobot.agent.hooks import create_file_edit_activity_hook
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.mcp import MCPProvider
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.bus.outbound_events import (
|
from nanobot.bus.outbound_events import (
|
||||||
StreamDeltaEvent,
|
StreamDeltaEvent,
|
||||||
StreamedResponseEvent,
|
StreamedResponseEvent,
|
||||||
@@ -84,6 +86,8 @@ def agent(
|
|||||||
# Create cron service with workspace-scoped store
|
# Create cron service with workspace-scoped store
|
||||||
cron_store_path = runtime_config.workspace_path / "cron" / "jobs.json"
|
cron_store_path = runtime_config.workspace_path / "cron" / "jobs.json"
|
||||||
cron = CronService(cron_store_path)
|
cron = CronService(cron_store_path)
|
||||||
|
tools = ToolRegistry()
|
||||||
|
mcp_provider = MCPProvider.from_config(runtime_config, tools)
|
||||||
|
|
||||||
_set_nanobot_logs(logs)
|
_set_nanobot_logs(logs)
|
||||||
|
|
||||||
@@ -95,6 +99,7 @@ def agent(
|
|||||||
cron_service=cron,
|
cron_service=cron,
|
||||||
image_generation_provider_configs=image_gen_provider_configs(runtime_config),
|
image_generation_provider_configs=image_gen_provider_configs(runtime_config),
|
||||||
hook_factories=[create_file_edit_activity_hook],
|
hook_factories=[create_file_edit_activity_hook],
|
||||||
|
tool_registry=tools,
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
_print_agent_start_error(exc)
|
_print_agent_start_error(exc)
|
||||||
@@ -106,6 +111,12 @@ def agent(
|
|||||||
render_markdown=False,
|
render_markdown=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def _close_runtime() -> None:
|
||||||
|
try:
|
||||||
|
await agent_loop.aclose()
|
||||||
|
finally:
|
||||||
|
await mcp_provider.aclose()
|
||||||
|
|
||||||
# Shared reference for progress callbacks
|
# Shared reference for progress callbacks
|
||||||
_thinking: ThinkingSpinner | None = None
|
_thinking: ThinkingSpinner | None = None
|
||||||
|
|
||||||
@@ -149,6 +160,8 @@ def agent(
|
|||||||
if message:
|
if message:
|
||||||
# Single message mode — direct call, no bus needed
|
# Single message mode — direct call, no bus needed
|
||||||
async def run_once() -> None:
|
async def run_once() -> None:
|
||||||
|
try:
|
||||||
|
await mcp_provider.connect()
|
||||||
renderer = StreamRenderer(
|
renderer = StreamRenderer(
|
||||||
render_markdown=markdown,
|
render_markdown=markdown,
|
||||||
bot_name=runtime_config.agents.defaults.bot_name,
|
bot_name=runtime_config.agents.defaults.bot_name,
|
||||||
@@ -172,7 +185,8 @@ def agent(
|
|||||||
metadata=response.metadata if response else None,
|
metadata=response.metadata if response else None,
|
||||||
**print_kwargs,
|
**print_kwargs,
|
||||||
)
|
)
|
||||||
await agent_loop.close_mcp()
|
finally:
|
||||||
|
await _close_runtime()
|
||||||
|
|
||||||
asyncio.run(run_once())
|
asyncio.run(run_once())
|
||||||
else:
|
else:
|
||||||
@@ -209,6 +223,7 @@ def agent(
|
|||||||
signal.signal(signal.SIGPIPE, signal.SIG_IGN)
|
signal.signal(signal.SIGPIPE, signal.SIG_IGN)
|
||||||
|
|
||||||
async def run_interactive() -> None:
|
async def run_interactive() -> None:
|
||||||
|
await mcp_provider.connect()
|
||||||
bus_task = asyncio.create_task(agent_loop.run())
|
bus_task = asyncio.create_task(agent_loop.run())
|
||||||
turn_done = asyncio.Event()
|
turn_done = asyncio.Event()
|
||||||
turn_done.set()
|
turn_done.set()
|
||||||
@@ -347,6 +362,6 @@ def agent(
|
|||||||
agent_loop.stop()
|
agent_loop.stop()
|
||||||
outbound_task.cancel()
|
outbound_task.cancel()
|
||||||
await asyncio.gather(bus_task, outbound_task, return_exceptions=True)
|
await asyncio.gather(bus_task, outbound_task, return_exceptions=True)
|
||||||
await agent_loop.close_mcp()
|
await _close_runtime()
|
||||||
|
|
||||||
asyncio.run(run_interactive())
|
asyncio.run(run_interactive())
|
||||||
|
|||||||
+49
-2
@@ -49,6 +49,8 @@ from nanobot import __logo__, __version__ # noqa: E402
|
|||||||
from nanobot import optional_features as feature_support # noqa: E402
|
from nanobot import optional_features as feature_support # noqa: E402
|
||||||
from nanobot.agent.hooks import create_file_edit_activity_hook # noqa: E402
|
from nanobot.agent.hooks import create_file_edit_activity_hook # noqa: E402
|
||||||
from nanobot.agent.loop import AgentLoop # noqa: E402
|
from nanobot.agent.loop import AgentLoop # noqa: E402
|
||||||
|
from nanobot.agent.tools.mcp import MCPProvider # noqa: E402
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry # noqa: E402
|
||||||
from nanobot.cli import terminal as cli_terminal # noqa: E402
|
from nanobot.cli import terminal as cli_terminal # noqa: E402
|
||||||
from nanobot.cli.agent import agent # noqa: E402
|
from nanobot.cli.agent import agent # noqa: E402
|
||||||
from nanobot.cli.gateway import create_gateway_app # noqa: E402
|
from nanobot.cli.gateway import create_gateway_app # noqa: E402
|
||||||
@@ -351,12 +353,15 @@ def serve(
|
|||||||
sync_workspace_templates(runtime_config.workspace_path)
|
sync_workspace_templates(runtime_config.workspace_path)
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
session_manager = SessionManager(runtime_config.workspace_path)
|
session_manager = SessionManager(runtime_config.workspace_path)
|
||||||
|
tools = ToolRegistry()
|
||||||
|
mcp_provider = MCPProvider.from_config(runtime_config, tools)
|
||||||
try:
|
try:
|
||||||
agent_loop = AgentLoop.from_config(
|
agent_loop = AgentLoop.from_config(
|
||||||
runtime_config, bus,
|
runtime_config, bus,
|
||||||
session_manager=session_manager,
|
session_manager=session_manager,
|
||||||
image_generation_provider_configs=image_gen_provider_configs(runtime_config),
|
image_generation_provider_configs=image_gen_provider_configs(runtime_config),
|
||||||
hook_factories=[create_file_edit_activity_hook],
|
hook_factories=[create_file_edit_activity_hook],
|
||||||
|
tool_registry=tools,
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
console.print(f"[red]Error: {exc}[/red]")
|
console.print(f"[red]Error: {exc}[/red]")
|
||||||
@@ -378,13 +383,17 @@ def serve(
|
|||||||
api_app = create_app(
|
api_app = create_app(
|
||||||
agent_loop, model_name=model_name, request_timeout=timeout,
|
agent_loop, model_name=model_name, request_timeout=timeout,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
|
prepare_agent=mcp_provider.connect,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def on_startup(_app: Any) -> None:
|
async def on_startup(_app: Any) -> None:
|
||||||
await agent_loop._connect_mcp()
|
await mcp_provider.connect()
|
||||||
|
|
||||||
async def on_cleanup(_app: Any) -> None:
|
async def on_cleanup(_app: Any) -> None:
|
||||||
await agent_loop.close_mcp()
|
try:
|
||||||
|
await agent_loop.aclose()
|
||||||
|
finally:
|
||||||
|
await mcp_provider.aclose()
|
||||||
|
|
||||||
api_app.on_startup.append(on_startup)
|
api_app.on_startup.append(on_startup)
|
||||||
api_app.on_cleanup.append(on_cleanup)
|
api_app.on_cleanup.append(on_cleanup)
|
||||||
@@ -431,6 +440,44 @@ app.add_typer(
|
|||||||
app.command(name="agent")(agent)
|
app.command(name="agent")(agent)
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# Session Commands
|
||||||
|
# ============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
sessions_app = typer.Typer(help="Manage persisted session history")
|
||||||
|
app.add_typer(sessions_app, name="sessions")
|
||||||
|
|
||||||
|
|
||||||
|
@sessions_app.command("restore-workspace")
|
||||||
|
def sessions_restore_workspace(
|
||||||
|
config: str | None = typer.Option(None, "--config", "-c", help="Path to config file"),
|
||||||
|
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
|
||||||
|
) -> None:
|
||||||
|
"""Copy sessions back into the workspace before downgrading nanobot."""
|
||||||
|
from nanobot.session.manager import SessionManager
|
||||||
|
|
||||||
|
runtime_config = _load_runtime_config(config, workspace)
|
||||||
|
data_dir = runtime_config.runtime_data_dir
|
||||||
|
manager = SessionManager(
|
||||||
|
runtime_config.workspace_path,
|
||||||
|
sessions_root=data_dir / "sessions" if data_dir is not None else None,
|
||||||
|
)
|
||||||
|
result = manager.restore_sessions_to_workspace()
|
||||||
|
console.print(
|
||||||
|
f"Restored {result.restored} session file(s) to "
|
||||||
|
f"{escape(str(runtime_config.workspace_path / 'sessions'))}; "
|
||||||
|
f"{result.unchanged} already matched."
|
||||||
|
)
|
||||||
|
if result.conflicts:
|
||||||
|
console.print(
|
||||||
|
"[red]Rollback is incomplete: existing or invalid files require manual review.[/red]"
|
||||||
|
)
|
||||||
|
for path in result.conflicts:
|
||||||
|
console.print(Text(f"- {path}", style="red"))
|
||||||
|
raise typer.Exit(1)
|
||||||
|
|
||||||
|
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
# Channel Commands
|
# Channel Commands
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
|
|||||||
@@ -14,6 +14,8 @@ from rich.console import Console
|
|||||||
from nanobot import __logo__, __version__
|
from nanobot import __logo__, __version__
|
||||||
from nanobot.agent.hooks import create_file_edit_activity_hook
|
from nanobot.agent.hooks import create_file_edit_activity_hook
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.mcp import MCPProvider
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.cli import terminal as cli_terminal
|
from nanobot.cli import terminal as cli_terminal
|
||||||
from nanobot.cli.runtime_config import _migrate_cron_store
|
from nanobot.cli.runtime_config import _migrate_cron_store
|
||||||
from nanobot.cli.webui_support import (
|
from nanobot.cli.webui_support import (
|
||||||
@@ -233,6 +235,7 @@ def _print_gateway_health_endpoint(host: str, port: int) -> None:
|
|||||||
|
|
||||||
async def _close_gateway_runtime(
|
async def _close_gateway_runtime(
|
||||||
agent: AgentLoop,
|
agent: AgentLoop,
|
||||||
|
mcp_provider: MCPProvider,
|
||||||
channels: Any,
|
channels: Any,
|
||||||
tasks: list[asyncio.Task[Any]],
|
tasks: list[asyncio.Task[Any]],
|
||||||
runtime_tasks: asyncio.Future[list[Any]] | None,
|
runtime_tasks: asyncio.Future[list[Any]] | None,
|
||||||
@@ -240,18 +243,13 @@ async def _close_gateway_runtime(
|
|||||||
task_wait_timeout: float = 15.0,
|
task_wait_timeout: float = 15.0,
|
||||||
close_timeout: float = 15.0,
|
close_timeout: float = 15.0,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Cancel runtime tasks, then deterministically close agent resources.
|
"""Cancel runtime tasks, then deterministically close application resources.
|
||||||
|
|
||||||
Order matters: runtime tasks (including the agent loop and any in-flight
|
Order matters: runtime tasks (including the agent loop and any in-flight
|
||||||
turn) are cancelled and awaited -- bounded -- before exec sessions,
|
turn) are cancelled and awaited -- bounded -- before the loop-owned resources
|
||||||
subagents, and MCP servers are torn down, so no active turn is using a
|
and the application-owned MCP provider are torn down. The final close is
|
||||||
shared resource when it closes. The final close is bounded and idempotent:
|
bounded and idempotent, so it also covers a cancelled or incomplete loop
|
||||||
the agent loop's own finally also calls ``close_mcp()``, so this runs again
|
cleanup without leaving subprocess transports alive past ``loop.close()``.
|
||||||
as a no-op when that path already completed, and as the guaranteed final
|
|
||||||
close when it was skipped or cut short (which previously left asyncio
|
|
||||||
subprocess transports alive past ``loop.close()``, producing
|
|
||||||
"RuntimeError: Event loop is closed" noise and potentially orphaned
|
|
||||||
processes at interpreter exit).
|
|
||||||
"""
|
"""
|
||||||
# Some SDKs swallow task cancellation while attempting to reconnect.
|
# Some SDKs swallow task cancellation while attempting to reconnect.
|
||||||
# Close channel transports before waiting for their runners to exit.
|
# Close channel transports before waiting for their runners to exit.
|
||||||
@@ -272,10 +270,14 @@ async def _close_gateway_runtime(
|
|||||||
task.cancel()
|
task.cancel()
|
||||||
if runtime_tasks is not None and not runtime_tasks.done():
|
if runtime_tasks is not None and not runtime_tasks.done():
|
||||||
runtime_tasks.cancel()
|
runtime_tasks.cancel()
|
||||||
|
for label, close in (
|
||||||
|
("agent", agent.aclose),
|
||||||
|
("MCP provider", mcp_provider.aclose),
|
||||||
|
):
|
||||||
try:
|
try:
|
||||||
await asyncio.wait_for(agent.close_mcp(), timeout=close_timeout)
|
await asyncio.wait_for(close(), timeout=close_timeout)
|
||||||
except BaseException as exc: # noqa: BLE001 - shutdown must proceed
|
except BaseException as exc: # noqa: BLE001 - shutdown must proceed
|
||||||
logger.warning("Gateway shutdown: agent resource cleanup incomplete: {}", exc)
|
logger.warning("Gateway shutdown: {} cleanup incomplete: {}", label, exc)
|
||||||
# Retrieving an already-finished gather prevents noisy unhandled exceptions,
|
# Retrieving an already-finished gather prevents noisy unhandled exceptions,
|
||||||
# but never wait for it here: its children were bounded individually above.
|
# but never wait for it here: its children were bounded individually above.
|
||||||
if runtime_tasks is not None and runtime_tasks.done():
|
if runtime_tasks is not None and runtime_tasks.done():
|
||||||
@@ -414,6 +416,9 @@ def _run_gateway(
|
|||||||
route_policy=WebuiTurnRoutePolicy(session_manager),
|
route_policy=WebuiTurnRoutePolicy(session_manager),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
tools = ToolRegistry()
|
||||||
|
mcp_provider = MCPProvider.from_config(config, tools)
|
||||||
|
|
||||||
# Create agent with cron service
|
# Create agent with cron service
|
||||||
agent = AgentLoop.from_config(
|
agent = AgentLoop.from_config(
|
||||||
config, bus,
|
config, bus,
|
||||||
@@ -431,6 +436,7 @@ def _run_gateway(
|
|||||||
hooks=[TokenUsageHook(timezone_name=config.agents.defaults.timezone)],
|
hooks=[TokenUsageHook(timezone_name=config.agents.defaults.timezone)],
|
||||||
local_trigger_store=trigger_store,
|
local_trigger_store=trigger_store,
|
||||||
hook_factories=[create_file_edit_activity_hook],
|
hook_factories=[create_file_edit_activity_hook],
|
||||||
|
tool_registry=tools,
|
||||||
)
|
)
|
||||||
def _schedule_webui_background(awaitable: Awaitable[None]) -> None:
|
def _schedule_webui_background(awaitable: Awaitable[None]) -> None:
|
||||||
agent.schedule_background(cast(Coroutine[Any, Any, None], awaitable))
|
agent.schedule_background(cast(Coroutine[Any, Any, None], awaitable))
|
||||||
@@ -512,6 +518,7 @@ def _run_gateway(
|
|||||||
prompt, last_cursor = result
|
prompt, last_cursor = result
|
||||||
key = dream_session_key()
|
key = dream_session_key()
|
||||||
dream_runtime = agent.dream_runtime()
|
dream_runtime = agent.dream_runtime()
|
||||||
|
await mcp_provider.connect()
|
||||||
resp = await agent.process_direct(
|
resp = await agent.process_direct(
|
||||||
prompt,
|
prompt,
|
||||||
session_key=key,
|
session_key=key,
|
||||||
@@ -589,6 +596,7 @@ def _run_gateway(
|
|||||||
if isinstance(message_tool, MessageTool):
|
if isinstance(message_tool, MessageTool):
|
||||||
suppress_token = message_tool.set_suppress_delivery(True)
|
suppress_token = message_tool.set_suppress_delivery(True)
|
||||||
try:
|
try:
|
||||||
|
await mcp_provider.connect()
|
||||||
resp = await agent.process_direct(
|
resp = await agent.process_direct(
|
||||||
prompt,
|
prompt,
|
||||||
session_key="heartbeat",
|
session_key="heartbeat",
|
||||||
@@ -668,7 +676,10 @@ def _run_gateway(
|
|||||||
webui_static_dist=webui_static_dist,
|
webui_static_dist=webui_static_dist,
|
||||||
webui_runtime_surface=webui_runtime_surface,
|
webui_runtime_surface=webui_runtime_surface,
|
||||||
webui_runtime_capabilities=webui_runtime_capabilities,
|
webui_runtime_capabilities=webui_runtime_capabilities,
|
||||||
|
webui_mcp_runtime_status=mcp_provider.runtime_status,
|
||||||
|
webui_mcp_reload=mcp_provider.reload,
|
||||||
webui_skill_state_action=_webui_skill_state_action,
|
webui_skill_state_action=_webui_skill_state_action,
|
||||||
|
config_path=Path(config_path),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _pick_heartbeat_target() -> tuple[str, str]:
|
def _pick_heartbeat_target() -> tuple[str, str]:
|
||||||
@@ -842,6 +853,13 @@ def _run_gateway(
|
|||||||
await cron.start()
|
await cron.start()
|
||||||
# Re-read once on first admission to close the watcher subscription window.
|
# Re-read once on first admission to close the watcher subscription window.
|
||||||
agent.runtime_resolver.invalidate()
|
agent.runtime_resolver.invalidate()
|
||||||
|
async def _run_agent() -> None:
|
||||||
|
try:
|
||||||
|
await mcp_provider.connect()
|
||||||
|
await agent.run()
|
||||||
|
finally:
|
||||||
|
await mcp_provider.aclose()
|
||||||
|
|
||||||
tasks = [
|
tasks = [
|
||||||
asyncio.create_task(
|
asyncio.create_task(
|
||||||
watch_config_file(
|
watch_config_file(
|
||||||
@@ -850,7 +868,7 @@ def _run_gateway(
|
|||||||
),
|
),
|
||||||
name="nanobot-config-watcher",
|
name="nanobot-config-watcher",
|
||||||
),
|
),
|
||||||
asyncio.create_task(agent.run(), name="nanobot-agent-loop"),
|
asyncio.create_task(_run_agent(), name="nanobot-agent-loop"),
|
||||||
asyncio.create_task(channels.start_all(), name="nanobot-channels"),
|
asyncio.create_task(channels.start_all(), name="nanobot-channels"),
|
||||||
asyncio.create_task(
|
asyncio.create_task(
|
||||||
run_local_trigger_queue(
|
run_local_trigger_queue(
|
||||||
@@ -908,7 +926,13 @@ def _run_gateway(
|
|||||||
agent.stop()
|
agent.stop()
|
||||||
# Cancel runtime tasks first, then deterministically close
|
# Cancel runtime tasks first, then deterministically close
|
||||||
# exec/MCP resources while the event loop is still alive.
|
# exec/MCP resources while the event loop is still alive.
|
||||||
await _close_gateway_runtime(agent, channels, tasks, runtime_tasks)
|
await _close_gateway_runtime(
|
||||||
|
agent,
|
||||||
|
mcp_provider,
|
||||||
|
channels,
|
||||||
|
tasks,
|
||||||
|
runtime_tasks,
|
||||||
|
)
|
||||||
# Flush all cached sessions to durable storage before exit.
|
# Flush all cached sessions to durable storage before exit.
|
||||||
# This prevents data loss on filesystems with write-back
|
# This prevents data loss on filesystems with write-back
|
||||||
# caching (rclone VFS, NFS, FUSE mounts, etc.).
|
# caching (rclone VFS, NFS, FUSE mounts, etc.).
|
||||||
|
|||||||
@@ -75,6 +75,7 @@ def load_config(config_path: Path | None = None) -> Config:
|
|||||||
summary="Environment-based configuration is invalid.",
|
summary="Environment-based configuration is invalid.",
|
||||||
issues=validation_issues(exc),
|
issues=validation_issues(exc),
|
||||||
) from exc
|
) from exc
|
||||||
|
config.bind_source_path(path)
|
||||||
_apply_ssrf_whitelist(config)
|
_apply_ssrf_whitelist(config)
|
||||||
return config
|
return config
|
||||||
|
|
||||||
@@ -130,6 +131,7 @@ def load_config(config_path: Path | None = None) -> Config:
|
|||||||
issues=issues,
|
issues=issues,
|
||||||
) from exc
|
) from exc
|
||||||
|
|
||||||
|
config.bind_source_path(path)
|
||||||
_apply_ssrf_whitelist(config)
|
_apply_ssrf_whitelist(config)
|
||||||
return config
|
return config
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast
|
from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast
|
||||||
|
|
||||||
from pydantic import AliasChoices, ConfigDict, Field, field_validator, model_validator
|
from pydantic import AliasChoices, ConfigDict, Field, PrivateAttr, field_validator, model_validator
|
||||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||||
|
|
||||||
from nanobot.config.timezone import detect_system_timezone
|
from nanobot.config.timezone import detect_system_timezone
|
||||||
@@ -373,6 +373,7 @@ class MCPServerConfig(Base):
|
|||||||
"""MCP server connection configuration (stdio or HTTP)."""
|
"""MCP server connection configuration (stdio or HTTP)."""
|
||||||
|
|
||||||
type: Literal["stdio", "sse", "streamableHttp"] | None = None # auto-detected if omitted
|
type: Literal["stdio", "sse", "streamableHttp"] | None = None # auto-detected if omitted
|
||||||
|
auth: Literal["oauth"] | None = None # Remote MCP OAuth; tokens are stored outside config
|
||||||
command: str = "" # Stdio: command to run (e.g. "npx")
|
command: str = "" # Stdio: command to run (e.g. "npx")
|
||||||
args: list[str] = Field(default_factory=list) # Stdio: command arguments
|
args: list[str] = Field(default_factory=list) # Stdio: command arguments
|
||||||
env: dict[str, str] = Field(default_factory=dict) # Stdio: extra env vars
|
env: dict[str, str] = Field(default_factory=dict) # Stdio: extra env vars
|
||||||
@@ -430,6 +431,8 @@ class ToolsConfig(Base):
|
|||||||
class Config(BaseSettings):
|
class Config(BaseSettings):
|
||||||
"""Root configuration for nanobot."""
|
"""Root configuration for nanobot."""
|
||||||
|
|
||||||
|
_source_path: Path | None = PrivateAttr(default=None)
|
||||||
|
|
||||||
agents: AgentsConfig = Field(default_factory=AgentsConfig)
|
agents: AgentsConfig = Field(default_factory=AgentsConfig)
|
||||||
channels: ChannelsConfig = Field(default_factory=ChannelsConfig)
|
channels: ChannelsConfig = Field(default_factory=ChannelsConfig)
|
||||||
transcription: TranscriptionConfig = Field(default_factory=TranscriptionConfig)
|
transcription: TranscriptionConfig = Field(default_factory=TranscriptionConfig)
|
||||||
@@ -448,6 +451,15 @@ class Config(BaseSettings):
|
|||||||
_resolve_tool_config_refs()
|
_resolve_tool_config_refs()
|
||||||
super().__init__(**values)
|
super().__init__(**values)
|
||||||
|
|
||||||
|
def bind_source_path(self, path: Path) -> None:
|
||||||
|
"""Record the config file that owns instance-level runtime data."""
|
||||||
|
self._source_path = path.expanduser().resolve(strict=False)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def runtime_data_dir(self) -> Path | None:
|
||||||
|
"""Return the active instance data directory when loaded from a config path."""
|
||||||
|
return self._source_path.parent if self._source_path is not None else None
|
||||||
|
|
||||||
@model_validator(mode="after")
|
@model_validator(mode="after")
|
||||||
def _validate_model_preset(self) -> "Config":
|
def _validate_model_preset(self) -> "Config":
|
||||||
if "default" in self.model_presets:
|
if "default" in self.model_presets:
|
||||||
|
|||||||
+24
-4
@@ -10,6 +10,8 @@ from typing import Any
|
|||||||
from nanobot.agent.hook import AgentHook, SDKCaptureHook
|
from nanobot.agent.hook import AgentHook, SDKCaptureHook
|
||||||
from nanobot.agent.hooks import create_file_edit_activity_hook
|
from nanobot.agent.hooks import create_file_edit_activity_hook
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.mcp import MCPProvider
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.config.schema import Config
|
from nanobot.config.schema import Config
|
||||||
from nanobot.providers.image_generation import image_gen_provider_configs
|
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||||
from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient
|
from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient
|
||||||
@@ -71,9 +73,16 @@ class Nanobot:
|
|||||||
print(result.content)
|
print(result.content)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
loop: AgentLoop,
|
||||||
|
*,
|
||||||
|
config: Config | None = None,
|
||||||
|
mcp_provider: MCPProvider | None = None,
|
||||||
|
) -> None:
|
||||||
self._loop = loop
|
self._loop = loop
|
||||||
self._config = config
|
self._config = config
|
||||||
|
self._mcp_provider = mcp_provider
|
||||||
self.sessions = SessionClient(loop)
|
self.sessions = SessionClient(loop)
|
||||||
self.memory = MemoryClient(loop)
|
self.memory = MemoryClient(loop)
|
||||||
self.runtime = RuntimeClient(loop)
|
self.runtime = RuntimeClient(loop)
|
||||||
@@ -120,12 +129,15 @@ class Nanobot:
|
|||||||
elif model_preset is not None:
|
elif model_preset is not None:
|
||||||
config.agents.defaults.model_preset = model_preset
|
config.agents.defaults.model_preset = model_preset
|
||||||
|
|
||||||
|
tools = ToolRegistry()
|
||||||
|
mcp_provider = MCPProvider.from_config(config, tools)
|
||||||
loop = AgentLoop.from_config(
|
loop = AgentLoop.from_config(
|
||||||
config,
|
config,
|
||||||
image_generation_provider_configs=image_gen_provider_configs(config),
|
image_generation_provider_configs=image_gen_provider_configs(config),
|
||||||
hook_factories=[create_file_edit_activity_hook],
|
hook_factories=[create_file_edit_activity_hook],
|
||||||
|
tool_registry=tools,
|
||||||
)
|
)
|
||||||
return cls(loop, config=config)
|
return cls(loop, config=config, mcp_provider=mcp_provider)
|
||||||
|
|
||||||
async def run(
|
async def run(
|
||||||
self,
|
self,
|
||||||
@@ -178,6 +190,8 @@ class Nanobot:
|
|||||||
)
|
)
|
||||||
if runtime is not None:
|
if runtime is not None:
|
||||||
kwargs["runtime"] = runtime
|
kwargs["runtime"] = runtime
|
||||||
|
if self._mcp_provider is not None:
|
||||||
|
await self._mcp_provider.connect()
|
||||||
response = await self._loop.process_direct(
|
response = await self._loop.process_direct(
|
||||||
message,
|
message,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
@@ -259,6 +273,8 @@ class Nanobot:
|
|||||||
if override_runtime is not None:
|
if override_runtime is not None:
|
||||||
kwargs["runtime"] = override_runtime
|
kwargs["runtime"] = override_runtime
|
||||||
try:
|
try:
|
||||||
|
if self._mcp_provider is not None:
|
||||||
|
await self._mcp_provider.connect()
|
||||||
response = await self._loop.process_direct(
|
response = await self._loop.process_direct(
|
||||||
message,
|
message,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
@@ -327,8 +343,12 @@ class Nanobot:
|
|||||||
await run.aclose()
|
await run.aclose()
|
||||||
|
|
||||||
async def aclose(self) -> None:
|
async def aclose(self) -> None:
|
||||||
"""Release resources held by this instance (MCP connections, etc.)."""
|
"""Release resources held by this instance."""
|
||||||
await self._loop.close_mcp()
|
try:
|
||||||
|
await self._loop.aclose()
|
||||||
|
finally:
|
||||||
|
if self._mcp_provider is not None:
|
||||||
|
await self._mcp_provider.aclose()
|
||||||
|
|
||||||
async def __aenter__(self) -> Nanobot:
|
async def __aenter__(self) -> Nanobot:
|
||||||
return self
|
return self
|
||||||
|
|||||||
@@ -56,6 +56,8 @@ if TYPE_CHECKING:
|
|||||||
# that ``unittest.mock.patch`` can find and replace it.
|
# that ``unittest.mock.patch`` can find and replace it.
|
||||||
AsyncOpenAI: Any = None
|
AsyncOpenAI: Any = None
|
||||||
|
|
||||||
|
_GEMINI_SKIP_THOUGHT_SIGNATURE = "skip_thought_signature_validator"
|
||||||
|
|
||||||
|
|
||||||
def _is_hosted_web_search_type(value: object) -> bool:
|
def _is_hosted_web_search_type(value: object) -> bool:
|
||||||
return isinstance(value, str) and (
|
return isinstance(value, str) and (
|
||||||
@@ -447,6 +449,28 @@ def _merge_unique_list(base: object, override: object) -> object:
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _merge_chat_extra_body(
|
||||||
|
kwargs: dict[str, Any],
|
||||||
|
extra_body: dict[str, Any],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Merge configured Chat Completions fields without clobbering tools."""
|
||||||
|
regular_extra = {key: value for key, value in extra_body.items() if key != "tools"}
|
||||||
|
merged = dict(kwargs)
|
||||||
|
if regular_extra:
|
||||||
|
existing = kwargs.get("extra_body", {})
|
||||||
|
merged["extra_body"] = _deep_merge(existing, regular_extra)
|
||||||
|
|
||||||
|
if "tools" in extra_body:
|
||||||
|
current_tools = kwargs.get("tools")
|
||||||
|
configured_tools = extra_body["tools"]
|
||||||
|
if isinstance(current_tools, list) and isinstance(configured_tools, list):
|
||||||
|
merged["tools"] = [*current_tools, *configured_tools]
|
||||||
|
else:
|
||||||
|
merged["tools"] = configured_tools
|
||||||
|
|
||||||
|
return merged
|
||||||
|
|
||||||
|
|
||||||
def _merge_responses_extra_body(
|
def _merge_responses_extra_body(
|
||||||
body: dict[str, Any],
|
body: dict[str, Any],
|
||||||
extra_body: dict[str, Any],
|
extra_body: dict[str, Any],
|
||||||
@@ -501,9 +525,6 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
self._proxy = proxy or None
|
self._proxy = proxy or None
|
||||||
self._native_compaction_available = True
|
self._native_compaction_available = True
|
||||||
|
|
||||||
if api_key and spec and spec.env_key:
|
|
||||||
self._setup_env(api_key, api_base)
|
|
||||||
|
|
||||||
effective_base = api_base or (spec.default_api_base if spec else None) or None
|
effective_base = api_base or (spec.default_api_base if spec else None) or None
|
||||||
self._effective_base = effective_base
|
self._effective_base = effective_base
|
||||||
self._default_headers = {"x-session-affinity": uuid.uuid4().hex}
|
self._default_headers = {"x-session-affinity": uuid.uuid4().hex}
|
||||||
@@ -596,20 +617,6 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
raise RuntimeError("OpenAI client initialization did not produce a client")
|
raise RuntimeError("OpenAI client initialization did not produce a client")
|
||||||
return self._client
|
return self._client
|
||||||
|
|
||||||
def _setup_env(self, api_key: str, api_base: str | None) -> None:
|
|
||||||
"""Set environment variables based on provider spec."""
|
|
||||||
spec = self._spec
|
|
||||||
if not spec or not spec.env_key:
|
|
||||||
return
|
|
||||||
if spec.is_gateway:
|
|
||||||
os.environ[spec.env_key] = api_key
|
|
||||||
else:
|
|
||||||
os.environ.setdefault(spec.env_key, api_key)
|
|
||||||
effective_base = api_base or spec.default_api_base
|
|
||||||
for env_name, env_val in spec.env_extras:
|
|
||||||
resolved = env_val.replace("{api_key}", api_key).replace("{api_base}", effective_base)
|
|
||||||
os.environ.setdefault(env_name, resolved)
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _apply_cache_control(
|
def _apply_cache_control(
|
||||||
cls,
|
cls,
|
||||||
@@ -685,6 +692,8 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
if strip_reasoning:
|
if strip_reasoning:
|
||||||
for msg in sanitized:
|
for msg in sanitized:
|
||||||
msg.pop("reasoning_content", None)
|
msg.pop("reasoning_content", None)
|
||||||
|
if self._spec and self._spec.name == "gemini":
|
||||||
|
sanitized = self._ensure_gemini_thought_signatures(sanitized)
|
||||||
|
|
||||||
def map_id(value: Any) -> Any:
|
def map_id(value: Any) -> Any:
|
||||||
if not isinstance(value, str):
|
if not isinstance(value, str):
|
||||||
@@ -762,6 +771,81 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
clean["content"] = self._coerce_content_to_string(clean.get("content"))
|
clean["content"] = self._coerce_content_to_string(clean.get("content"))
|
||||||
return self._enforce_role_alternation(sanitized)
|
return self._enforce_role_alternation(sanitized)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _gemini_thought_signature(tool_call: dict[str, Any]) -> str | None:
|
||||||
|
"""Return Gemini's thought signature attached to a tool call, if any.
|
||||||
|
|
||||||
|
Gemini's OpenAI-compatible endpoint returns tool calls with an
|
||||||
|
``extra_content`` field: ``{"google": {"thought_signature": "..."}}``.
|
||||||
|
nanobot preserves it through the parse -> serialize round-trip so
|
||||||
|
replayed calls stay valid. Calls produced by other providers (e.g.
|
||||||
|
after a mid-conversation model switch) carry no signature.
|
||||||
|
"""
|
||||||
|
extra = tool_call.get("extra_content")
|
||||||
|
if not isinstance(extra, dict):
|
||||||
|
return None
|
||||||
|
google = cast(dict[str, Any], extra).get("google")
|
||||||
|
if not isinstance(google, dict):
|
||||||
|
return None
|
||||||
|
signature = cast(dict[str, Any], google).get("thought_signature")
|
||||||
|
if isinstance(signature, str) and signature:
|
||||||
|
return signature
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _ensure_gemini_thought_signatures(
|
||||||
|
self, messages: list[dict[str, Any]]
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Keep migrated tool history wire-valid without losing tool context.
|
||||||
|
|
||||||
|
Gemini requires the first call in each function-call step to carry a
|
||||||
|
thought signature. Native parallel calls intentionally leave later
|
||||||
|
calls unsigned, so they must remain in their original order. For a
|
||||||
|
fully unsigned step imported from another provider, Google documents
|
||||||
|
``skip_thought_signature_validator`` as a last-resort migration value.
|
||||||
|
"""
|
||||||
|
kept: list[dict[str, Any]] = []
|
||||||
|
for msg in messages:
|
||||||
|
role = msg.get("role")
|
||||||
|
calls = msg.get("tool_calls")
|
||||||
|
if role != "assistant" or not isinstance(calls, list) or not calls:
|
||||||
|
kept.append(msg)
|
||||||
|
continue
|
||||||
|
|
||||||
|
call_values = cast(list[object], calls)
|
||||||
|
typed_calls = [
|
||||||
|
cast(dict[str, Any], tool_call)
|
||||||
|
for tool_call in call_values
|
||||||
|
if isinstance(tool_call, dict)
|
||||||
|
]
|
||||||
|
if not typed_calls:
|
||||||
|
if msg.get("content"):
|
||||||
|
clean = dict(msg)
|
||||||
|
clean.pop("tool_calls", None)
|
||||||
|
kept.append(clean)
|
||||||
|
continue
|
||||||
|
|
||||||
|
clean_calls = typed_calls
|
||||||
|
if self._gemini_thought_signature(typed_calls[0]) is None:
|
||||||
|
first = dict(typed_calls[0])
|
||||||
|
extra_value = first.get("extra_content")
|
||||||
|
extra = dict(cast(dict[str, Any], extra_value)) if isinstance(
|
||||||
|
extra_value, dict
|
||||||
|
) else {}
|
||||||
|
google_value = extra.get("google")
|
||||||
|
google = dict(cast(dict[str, Any], google_value)) if isinstance(
|
||||||
|
google_value, dict
|
||||||
|
) else {}
|
||||||
|
google["thought_signature"] = _GEMINI_SKIP_THOUGHT_SIGNATURE
|
||||||
|
extra["google"] = google
|
||||||
|
first["extra_content"] = extra
|
||||||
|
clean_calls = [first, *typed_calls[1:]]
|
||||||
|
|
||||||
|
if clean_calls != call_values:
|
||||||
|
msg = dict(msg)
|
||||||
|
msg["tool_calls"] = clean_calls
|
||||||
|
kept.append(msg)
|
||||||
|
return kept
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Build kwargs
|
# Build kwargs
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@@ -968,14 +1052,11 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
if msg.get("role") == "assistant" and "reasoning_content" not in msg:
|
if msg.get("role") == "assistant" and "reasoning_content" not in msg:
|
||||||
msg["reasoning_content"] = ""
|
msg["reasoning_content"] = ""
|
||||||
|
|
||||||
# Merge user-configured extra_body last so it can override or
|
# Merge user-configured extra_body last so ordinary fields can override
|
||||||
# extend provider-specific defaults (e.g. chat_template_kwargs,
|
# provider defaults. Keep configured tools at the top level: the SDK
|
||||||
# guided_json, repetition_penalty). Uses recursive merge so
|
# otherwise lets extra_body.tools replace nanobot's generated functions.
|
||||||
# nested dicts like {"chat_template_kwargs": {"enable_thinking": false}}
|
|
||||||
# do not clobber sibling keys already set by thinking-style logic.
|
|
||||||
if self._extra_body:
|
if self._extra_body:
|
||||||
existing = kwargs.get("extra_body", {})
|
kwargs = _merge_chat_extra_body(kwargs, self._extra_body)
|
||||||
kwargs["extra_body"] = _deep_merge(existing, self._extra_body)
|
|
||||||
|
|
||||||
return kwargs
|
return kwargs
|
||||||
|
|
||||||
@@ -1156,7 +1237,8 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
self._sanitize_empty_content(sanitized_state.pending_messages)
|
self._sanitize_empty_content(sanitized_state.pending_messages)
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
preserve_reasoning = bool(self._spec and self._spec.name == "deepseek")
|
is_deepseek = bool(self._spec and self._spec.name == "deepseek")
|
||||||
|
preserve_reasoning = is_deepseek
|
||||||
instructions, input_items, replayed = prepare_responses_input(
|
instructions, input_items, replayed = prepare_responses_input(
|
||||||
sanitized_messages,
|
sanitized_messages,
|
||||||
state=sanitized_state,
|
state=sanitized_state,
|
||||||
@@ -1192,7 +1274,7 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
|
|
||||||
if not self._supports_temperature(model_name, reasoning_effort) and not preserve_reasoning:
|
if not self._supports_temperature(model_name, reasoning_effort) and not preserve_reasoning:
|
||||||
body["include"] = ["reasoning.encrypted_content"]
|
body["include"] = ["reasoning.encrypted_content"]
|
||||||
if reasoning_effort and reasoning_effort.lower() != "none":
|
if reasoning_effort and (reasoning_effort.lower() != "none" or is_deepseek):
|
||||||
body["reasoning"] = {"effort": reasoning_effort}
|
body["reasoning"] = {"effort": reasoning_effort}
|
||||||
if replayed and "gpt-5.6" in model_name.lower():
|
if replayed and "gpt-5.6" in model_name.lower():
|
||||||
body.setdefault("reasoning", {})["context"] = "all_turns"
|
body.setdefault("reasoning", {})["context"] = "all_turns"
|
||||||
|
|||||||
@@ -112,8 +112,7 @@ class ProviderSpec:
|
|||||||
implicit_reasoning_models: tuple[str, ...] = ()
|
implicit_reasoning_models: tuple[str, ...] = ()
|
||||||
|
|
||||||
# Models that expose the OpenAI Responses wire format. This is model-level
|
# Models that expose the OpenAI Responses wire format. This is model-level
|
||||||
# because providers may add Responses support incrementally (DeepSeek V4
|
# because providers may add Responses support incrementally.
|
||||||
# Flash is supported before V4 Pro).
|
|
||||||
responses_models: tuple[str, ...] = ()
|
responses_models: tuple[str, ...] = ()
|
||||||
|
|
||||||
# Provider-hosted Responses tools sent unless extraBody.tools explicitly
|
# Provider-hosted Responses tools sent unless extraBody.tools explicitly
|
||||||
@@ -482,7 +481,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
|
|||||||
backend="openai_compat",
|
backend="openai_compat",
|
||||||
default_api_base="https://api.deepseek.com",
|
default_api_base="https://api.deepseek.com",
|
||||||
thinking_style="thinking_type",
|
thinking_style="thinking_type",
|
||||||
responses_models=("deepseek-v4-flash",),
|
responses_models=("deepseek-v4-flash", "deepseek-v4-pro"),
|
||||||
responses_default_tools=("web_search",),
|
responses_default_tools=("web_search",),
|
||||||
),
|
),
|
||||||
# Gemini: Google's OpenAI-compatible endpoint
|
# Gemini: Google's OpenAI-compatible endpoint
|
||||||
|
|||||||
+477
-5
@@ -2,9 +2,12 @@
|
|||||||
|
|
||||||
import base64
|
import base64
|
||||||
import errno
|
import errno
|
||||||
|
import hashlib
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
import secrets
|
||||||
|
import stat
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
@@ -14,9 +17,10 @@ from pathlib import Path
|
|||||||
from typing import Any, Callable, Collection, Protocol, TypedDict, cast
|
from typing import Any, Callable, Collection, Protocol, TypedDict, cast
|
||||||
from weakref import WeakValueDictionary
|
from weakref import WeakValueDictionary
|
||||||
|
|
||||||
|
from filelock import FileLock
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.config.paths import get_legacy_sessions_dir
|
from nanobot.config.paths import get_legacy_sessions_dir, get_runtime_subdir
|
||||||
from nanobot.providers.base import ProviderConversationState
|
from nanobot.providers.base import ProviderConversationState
|
||||||
from nanobot.runtime_context import (
|
from nanobot.runtime_context import (
|
||||||
RUNTIME_CONTEXT_HISTORY_META,
|
RUNTIME_CONTEXT_HISTORY_META,
|
||||||
@@ -57,6 +61,11 @@ _FORK_VOLATILE_METADATA_KEYS = {
|
|||||||
"title",
|
"title",
|
||||||
"title_user_edited",
|
"title_user_edited",
|
||||||
}
|
}
|
||||||
|
_WORKSPACE_STATE_DIR = ".nanobot"
|
||||||
|
_WORKSPACE_ID_FILE = "workspace-id"
|
||||||
|
_WORKSPACE_ID_RE = re.compile(r"^[0-9a-f]{32}$")
|
||||||
|
_SESSION_MIGRATION_LOCK_TIMEOUT_SECONDS = 30
|
||||||
|
_COPY_CHUNK_SIZE = 1024 * 1024
|
||||||
|
|
||||||
|
|
||||||
def _json_object(value: object) -> dict[str, Any]:
|
def _json_object(value: object) -> dict[str, Any]:
|
||||||
@@ -503,6 +512,23 @@ class SessionInfo(TypedDict):
|
|||||||
path: str
|
path: str
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _SessionFileSnapshot:
|
||||||
|
digest: str
|
||||||
|
size: int
|
||||||
|
mtime_ns: int
|
||||||
|
updated_at: float
|
||||||
|
device: int
|
||||||
|
inode: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SessionRestoreResult:
|
||||||
|
restored: int
|
||||||
|
unchanged: int
|
||||||
|
conflicts: tuple[Path, ...]
|
||||||
|
|
||||||
|
|
||||||
class SessionStore(Protocol):
|
class SessionStore(Protocol):
|
||||||
def load(self, key: str) -> Session | None: ...
|
def load(self, key: str) -> Session | None: ...
|
||||||
|
|
||||||
@@ -520,9 +546,445 @@ class SessionStore(Protocol):
|
|||||||
class JsonlSessionStore:
|
class JsonlSessionStore:
|
||||||
"""JSONL implementation of session persistence."""
|
"""JSONL implementation of session persistence."""
|
||||||
|
|
||||||
def __init__(self, workspace: Path):
|
def __init__(self, workspace: Path, *, sessions_root: Path | None = None):
|
||||||
self.sessions_dir = ensure_dir(workspace / "sessions")
|
canonical_workspace = Path(workspace).expanduser().resolve(strict=False)
|
||||||
|
ensure_dir(canonical_workspace)
|
||||||
|
root = (
|
||||||
|
Path(sessions_root).expanduser().resolve(strict=False)
|
||||||
|
if sessions_root is not None
|
||||||
|
else get_runtime_subdir("sessions").resolve(strict=False)
|
||||||
|
)
|
||||||
|
if root == canonical_workspace or root.is_relative_to(canonical_workspace):
|
||||||
|
raise RuntimeError(
|
||||||
|
"session storage must be outside the agent workspace; "
|
||||||
|
"move --config outside --workspace or choose a nested workspace directory"
|
||||||
|
)
|
||||||
|
ensure_dir(root)
|
||||||
|
with suppress(OSError):
|
||||||
|
os.chmod(root, 0o700)
|
||||||
|
self.workspace = canonical_workspace
|
||||||
|
self._migration_lock = FileLock(
|
||||||
|
str(root / ".workspace-migration.lock"),
|
||||||
|
timeout=_SESSION_MIGRATION_LOCK_TIMEOUT_SECONDS,
|
||||||
|
)
|
||||||
|
with self._migration_lock:
|
||||||
|
workspace_id = self._load_or_create_workspace_id(canonical_workspace, root)
|
||||||
|
workspace_id = self._claim_workspace_namespace(
|
||||||
|
root,
|
||||||
|
canonical_workspace,
|
||||||
|
workspace_id,
|
||||||
|
)
|
||||||
|
self.sessions_dir = ensure_dir(root / workspace_id)
|
||||||
self.legacy_sessions_dir = get_legacy_sessions_dir()
|
self.legacy_sessions_dir = get_legacy_sessions_dir()
|
||||||
|
self._migrate_from_workspace(canonical_workspace)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _fsync_directory(path: Path) -> None:
|
||||||
|
with suppress(PermissionError, NotImplementedError):
|
||||||
|
fd = os.open(path, os.O_RDONLY)
|
||||||
|
try:
|
||||||
|
os.fsync(fd)
|
||||||
|
except OSError as exc:
|
||||||
|
if exc.errno != errno.EINVAL:
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
os.close(fd)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _write_text_atomic(cls, path: Path, content: str, *, mode: int = 0o600) -> None:
|
||||||
|
tmp = path.with_name(f".{path.name}.{secrets.token_hex(8)}.tmp")
|
||||||
|
try:
|
||||||
|
with open(tmp, "x", encoding="utf-8") as handle:
|
||||||
|
os.chmod(tmp, mode)
|
||||||
|
handle.write(content)
|
||||||
|
handle.flush()
|
||||||
|
os.fsync(handle.fileno())
|
||||||
|
os.replace(tmp, path)
|
||||||
|
cls._fsync_directory(path.parent)
|
||||||
|
finally:
|
||||||
|
tmp.unlink(missing_ok=True)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _read_workspace_id(cls, marker: Path) -> str:
|
||||||
|
if marker.is_symlink():
|
||||||
|
raise RuntimeError(f"workspace identity marker must not be a symlink: {marker}")
|
||||||
|
value = marker.read_text(encoding="utf-8").strip()
|
||||||
|
if not _WORKSPACE_ID_RE.fullmatch(value):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"workspace identity marker is invalid: {marker}; "
|
||||||
|
"restore its original 32-character identifier before starting nanobot"
|
||||||
|
)
|
||||||
|
return value
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _workspace_id_path(workspace: Path) -> Path:
|
||||||
|
state_dir = workspace / _WORKSPACE_STATE_DIR
|
||||||
|
if state_dir.is_symlink():
|
||||||
|
raise RuntimeError(f"workspace state directory must not be a symlink: {state_dir}")
|
||||||
|
ensure_dir(state_dir)
|
||||||
|
return state_dir / _WORKSPACE_ID_FILE
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _find_workspace_namespace(cls, workspace: Path, root: Path) -> str | None:
|
||||||
|
"""Recover an identity marker removed by cleanup at the same workspace path."""
|
||||||
|
matches: list[str] = []
|
||||||
|
for sessions_dir in root.iterdir():
|
||||||
|
if (
|
||||||
|
not _WORKSPACE_ID_RE.fullmatch(sessions_dir.name)
|
||||||
|
or sessions_dir.is_symlink()
|
||||||
|
or not sessions_dir.is_dir()
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
marker = sessions_dir / ".workspace"
|
||||||
|
if marker.is_symlink() or not marker.is_file():
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
recorded = Path(marker.read_text(encoding="utf-8").strip()).expanduser()
|
||||||
|
recorded = recorded.resolve(strict=False)
|
||||||
|
same_workspace = recorded == workspace or (
|
||||||
|
recorded.exists() and recorded.samefile(workspace)
|
||||||
|
)
|
||||||
|
except (OSError, UnicodeError, ValueError):
|
||||||
|
continue
|
||||||
|
if same_workspace:
|
||||||
|
matches.append(sessions_dir.name)
|
||||||
|
if len(matches) > 1:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"multiple session namespaces claim workspace {workspace}; "
|
||||||
|
"remove the stale namespace marker before starting nanobot"
|
||||||
|
)
|
||||||
|
return matches[0] if matches else None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _load_or_create_workspace_id(cls, workspace: Path, root: Path) -> str:
|
||||||
|
marker = cls._workspace_id_path(workspace)
|
||||||
|
if marker.exists() or marker.is_symlink():
|
||||||
|
return cls._read_workspace_id(marker)
|
||||||
|
|
||||||
|
recovered = cls._find_workspace_namespace(workspace, root)
|
||||||
|
if recovered is not None:
|
||||||
|
cls._write_text_atomic(marker, f"{recovered}\n")
|
||||||
|
return recovered
|
||||||
|
|
||||||
|
workspace_id = secrets.token_hex(16)
|
||||||
|
flags = os.O_WRONLY | os.O_CREAT | os.O_EXCL
|
||||||
|
flags |= getattr(os, "O_NOFOLLOW", 0)
|
||||||
|
try:
|
||||||
|
fd = os.open(marker, flags, 0o600)
|
||||||
|
except FileExistsError:
|
||||||
|
return cls._read_workspace_id(marker)
|
||||||
|
try:
|
||||||
|
payload = f"{workspace_id}\n".encode("ascii")
|
||||||
|
view = memoryview(payload)
|
||||||
|
while view:
|
||||||
|
written = os.write(fd, view)
|
||||||
|
view = view[written:]
|
||||||
|
os.fsync(fd)
|
||||||
|
except BaseException:
|
||||||
|
with suppress(OSError):
|
||||||
|
marker.unlink()
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
os.close(fd)
|
||||||
|
cls._fsync_directory(marker.parent)
|
||||||
|
return workspace_id
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _replace_workspace_id(cls, workspace: Path, workspace_id: str) -> None:
|
||||||
|
cls._write_text_atomic(cls._workspace_id_path(workspace), f"{workspace_id}\n")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _write_workspace_marker(cls, sessions_dir: Path, workspace: Path) -> None:
|
||||||
|
cls._write_text_atomic(sessions_dir / ".workspace", f"{workspace}\n")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _claim_workspace_namespace(
|
||||||
|
cls,
|
||||||
|
root: Path,
|
||||||
|
workspace: Path,
|
||||||
|
workspace_id: str,
|
||||||
|
) -> str:
|
||||||
|
"""Bind a stable workspace ID, rotating copied live workspaces apart."""
|
||||||
|
for _attempt in range(3):
|
||||||
|
sessions_dir = root / workspace_id
|
||||||
|
marker = sessions_dir / ".workspace"
|
||||||
|
if sessions_dir.is_symlink():
|
||||||
|
raise RuntimeError(f"session namespace must not be a symlink: {sessions_dir}")
|
||||||
|
if not sessions_dir.exists():
|
||||||
|
ensure_dir(sessions_dir)
|
||||||
|
cls._write_workspace_marker(sessions_dir, workspace)
|
||||||
|
return workspace_id
|
||||||
|
if marker.is_symlink():
|
||||||
|
raise RuntimeError(f"session workspace marker must not be a symlink: {marker}")
|
||||||
|
if not marker.exists():
|
||||||
|
if any(sessions_dir.iterdir()):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"session namespace has data but no workspace marker: {sessions_dir}"
|
||||||
|
)
|
||||||
|
cls._write_workspace_marker(sessions_dir, workspace)
|
||||||
|
return workspace_id
|
||||||
|
|
||||||
|
recorded_text = marker.read_text(encoding="utf-8").strip()
|
||||||
|
if not recorded_text:
|
||||||
|
raise RuntimeError(f"session workspace marker is empty: {marker}")
|
||||||
|
recorded = Path(recorded_text).expanduser().resolve(strict=False)
|
||||||
|
if recorded == workspace:
|
||||||
|
return workspace_id
|
||||||
|
try:
|
||||||
|
same_workspace = recorded.exists() and recorded.samefile(workspace)
|
||||||
|
except OSError:
|
||||||
|
same_workspace = False
|
||||||
|
if same_workspace:
|
||||||
|
cls._write_workspace_marker(sessions_dir, workspace)
|
||||||
|
return workspace_id
|
||||||
|
if not recorded.exists():
|
||||||
|
# The identity marker travelled with a renamed or moved workspace.
|
||||||
|
cls._write_workspace_marker(sessions_dir, workspace)
|
||||||
|
return workspace_id
|
||||||
|
|
||||||
|
# Both paths exist and are different: this is a copy, not a move.
|
||||||
|
workspace_id = secrets.token_hex(16)
|
||||||
|
cls._replace_workspace_id(workspace, workspace_id)
|
||||||
|
|
||||||
|
raise RuntimeError(f"could not allocate an isolated session namespace for {workspace}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _session_file_snapshot(path: Path) -> _SessionFileSnapshot | None:
|
||||||
|
flags = os.O_RDONLY | getattr(os, "O_NOFOLLOW", 0)
|
||||||
|
try:
|
||||||
|
fd = os.open(path, flags)
|
||||||
|
except OSError:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
before = os.fstat(fd)
|
||||||
|
if not stat.S_ISREG(before.st_mode):
|
||||||
|
return None
|
||||||
|
digest = hashlib.sha256()
|
||||||
|
saw_record = False
|
||||||
|
updated_at: float | None = None
|
||||||
|
with os.fdopen(fd, "rb", closefd=False) as handle:
|
||||||
|
for raw_line in handle:
|
||||||
|
digest.update(raw_line)
|
||||||
|
if not raw_line.strip():
|
||||||
|
continue
|
||||||
|
value: object = json.loads(raw_line.decode("utf-8"))
|
||||||
|
data = _json_object(value)
|
||||||
|
saw_record = True
|
||||||
|
if data.get("_type") == "metadata":
|
||||||
|
raw_updated_at = cast(object, data.get("updated_at"))
|
||||||
|
if isinstance(raw_updated_at, str) and raw_updated_at:
|
||||||
|
updated_at = datetime.fromisoformat(raw_updated_at).timestamp()
|
||||||
|
after = os.fstat(fd)
|
||||||
|
if (
|
||||||
|
not saw_record
|
||||||
|
or before.st_dev != after.st_dev
|
||||||
|
or before.st_ino != after.st_ino
|
||||||
|
or before.st_size != after.st_size
|
||||||
|
or before.st_mtime_ns != after.st_mtime_ns
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
return _SessionFileSnapshot(
|
||||||
|
digest=digest.hexdigest(),
|
||||||
|
size=after.st_size,
|
||||||
|
mtime_ns=after.st_mtime_ns,
|
||||||
|
updated_at=(updated_at if updated_at is not None else after.st_mtime_ns / 1e9),
|
||||||
|
device=after.st_dev,
|
||||||
|
inode=after.st_ino,
|
||||||
|
)
|
||||||
|
except (OSError, UnicodeDecodeError, json.JSONDecodeError, ValueError, TypeError):
|
||||||
|
return None
|
||||||
|
finally:
|
||||||
|
os.close(fd)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _prepare_copy(
|
||||||
|
cls,
|
||||||
|
src: Path,
|
||||||
|
dst_dir: Path,
|
||||||
|
snapshot: _SessionFileSnapshot,
|
||||||
|
) -> Path:
|
||||||
|
tmp = dst_dir / f".{src.name}.{secrets.token_hex(8)}.tmp"
|
||||||
|
flags = os.O_RDONLY | getattr(os, "O_NOFOLLOW", 0)
|
||||||
|
src_fd = os.open(src, flags)
|
||||||
|
try:
|
||||||
|
before = os.fstat(src_fd)
|
||||||
|
if (
|
||||||
|
before.st_dev != snapshot.device
|
||||||
|
or before.st_ino != snapshot.inode
|
||||||
|
or before.st_size != snapshot.size
|
||||||
|
or before.st_mtime_ns != snapshot.mtime_ns
|
||||||
|
):
|
||||||
|
raise OSError("session source changed before migration")
|
||||||
|
digest = hashlib.sha256()
|
||||||
|
size = 0
|
||||||
|
with os.fdopen(src_fd, "rb", closefd=False) as source, open(tmp, "xb") as target:
|
||||||
|
os.chmod(tmp, 0o600)
|
||||||
|
while chunk := source.read(_COPY_CHUNK_SIZE):
|
||||||
|
digest.update(chunk)
|
||||||
|
size += len(chunk)
|
||||||
|
target.write(chunk)
|
||||||
|
target.flush()
|
||||||
|
os.fsync(target.fileno())
|
||||||
|
after = os.fstat(src_fd)
|
||||||
|
if (
|
||||||
|
digest.hexdigest() != snapshot.digest
|
||||||
|
or size != snapshot.size
|
||||||
|
or after.st_dev != snapshot.device
|
||||||
|
or after.st_ino != snapshot.inode
|
||||||
|
or after.st_size != snapshot.size
|
||||||
|
or after.st_mtime_ns != snapshot.mtime_ns
|
||||||
|
):
|
||||||
|
raise OSError("session source changed during migration")
|
||||||
|
return tmp
|
||||||
|
except BaseException:
|
||||||
|
tmp.unlink(missing_ok=True)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
os.close(src_fd)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _install_snapshot(
|
||||||
|
cls,
|
||||||
|
src: Path,
|
||||||
|
dst: Path,
|
||||||
|
snapshot: _SessionFileSnapshot,
|
||||||
|
) -> None:
|
||||||
|
tmp = cls._prepare_copy(src, dst.parent, snapshot)
|
||||||
|
try:
|
||||||
|
os.replace(tmp, dst)
|
||||||
|
cls._fsync_directory(dst.parent)
|
||||||
|
installed = cls._session_file_snapshot(dst)
|
||||||
|
if installed is None or installed.digest != snapshot.digest:
|
||||||
|
raise OSError(f"session migration verification failed: {dst}")
|
||||||
|
finally:
|
||||||
|
tmp.unlink(missing_ok=True)
|
||||||
|
|
||||||
|
def _archive_conflict(
|
||||||
|
self,
|
||||||
|
src: Path,
|
||||||
|
snapshot: _SessionFileSnapshot,
|
||||||
|
label: str,
|
||||||
|
) -> Path:
|
||||||
|
conflict_dir = ensure_dir(self.sessions_dir / ".migration-conflicts")
|
||||||
|
conflict = conflict_dir / (
|
||||||
|
f"{src.stem}.{label}.{snapshot.digest[:12]}.{secrets.token_hex(4)}.jsonl"
|
||||||
|
)
|
||||||
|
self._install_snapshot(src, conflict, snapshot)
|
||||||
|
return conflict
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _remove_migrated_source(
|
||||||
|
cls,
|
||||||
|
src: Path,
|
||||||
|
snapshot: _SessionFileSnapshot,
|
||||||
|
) -> bool:
|
||||||
|
try:
|
||||||
|
current = src.stat(follow_symlinks=False)
|
||||||
|
if (
|
||||||
|
current.st_dev != snapshot.device
|
||||||
|
or current.st_ino != snapshot.inode
|
||||||
|
or current.st_size != snapshot.size
|
||||||
|
or current.st_mtime_ns != snapshot.mtime_ns
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
src.unlink()
|
||||||
|
cls._fsync_directory(src.parent)
|
||||||
|
return True
|
||||||
|
except OSError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _migrate_from_workspace(self, workspace: Path) -> None:
|
||||||
|
"""Durably copy legacy sessions out of the workspace, then remove the source."""
|
||||||
|
old_dir = workspace / "sessions"
|
||||||
|
if old_dir.is_symlink() or not old_dir.is_dir():
|
||||||
|
if old_dir.is_symlink():
|
||||||
|
logger.warning("Skipping symlinked legacy sessions directory: {}", old_dir)
|
||||||
|
return
|
||||||
|
for src in old_dir.glob("*.jsonl"):
|
||||||
|
if src.is_symlink() or not src.is_file():
|
||||||
|
logger.warning("Skipping unsafe legacy session file: {}", src)
|
||||||
|
continue
|
||||||
|
dst = self.sessions_dir / src.name
|
||||||
|
source_snapshot = self._session_file_snapshot(src)
|
||||||
|
if source_snapshot is None:
|
||||||
|
logger.warning("Skipping invalid or changing legacy session file: {}", src)
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
destination_snapshot = self._session_file_snapshot(dst) if dst.exists() else None
|
||||||
|
if dst.exists() and destination_snapshot is None:
|
||||||
|
logger.warning(
|
||||||
|
"Keeping legacy session because destination is invalid: {}",
|
||||||
|
dst,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
if destination_snapshot is None:
|
||||||
|
self._install_snapshot(src, dst, source_snapshot)
|
||||||
|
elif destination_snapshot.digest == source_snapshot.digest:
|
||||||
|
pass
|
||||||
|
elif source_snapshot.updated_at > destination_snapshot.updated_at:
|
||||||
|
archived = self._archive_conflict(dst, destination_snapshot, "destination")
|
||||||
|
self._install_snapshot(src, dst, source_snapshot)
|
||||||
|
logger.warning("Archived older session migration conflict at {}", archived)
|
||||||
|
else:
|
||||||
|
archived = self._archive_conflict(src, source_snapshot, "workspace")
|
||||||
|
logger.warning("Archived older session migration conflict at {}", archived)
|
||||||
|
|
||||||
|
installed = self._session_file_snapshot(dst)
|
||||||
|
if installed is None:
|
||||||
|
raise OSError(f"session migration destination is unreadable: {dst}")
|
||||||
|
selected_digest = (
|
||||||
|
source_snapshot.digest
|
||||||
|
if destination_snapshot is None
|
||||||
|
or source_snapshot.updated_at > destination_snapshot.updated_at
|
||||||
|
else destination_snapshot.digest
|
||||||
|
)
|
||||||
|
if installed.digest != selected_digest:
|
||||||
|
raise OSError(f"session migration selected unexpected data: {dst}")
|
||||||
|
if not self._remove_migrated_source(src, source_snapshot):
|
||||||
|
logger.warning(
|
||||||
|
"Session migrated but legacy source changed or could not be removed: {}",
|
||||||
|
src,
|
||||||
|
)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.warning("Failed to migrate session {}: {}", src, exc)
|
||||||
|
|
||||||
|
def restore_to_workspace(self) -> SessionRestoreResult:
|
||||||
|
"""Copy canonical sessions back for an explicit downgrade or rollback."""
|
||||||
|
restored = 0
|
||||||
|
unchanged = 0
|
||||||
|
conflicts: list[Path] = []
|
||||||
|
old_dir = self.workspace / "sessions"
|
||||||
|
if old_dir.is_symlink():
|
||||||
|
raise RuntimeError(f"refusing to restore into symlinked sessions directory: {old_dir}")
|
||||||
|
ensure_dir(old_dir)
|
||||||
|
|
||||||
|
with self._migration_lock:
|
||||||
|
for src in self.sessions_dir.glob("*.jsonl"):
|
||||||
|
if self.session_key_from_path(src) is None:
|
||||||
|
continue
|
||||||
|
source_snapshot = self._session_file_snapshot(src)
|
||||||
|
if source_snapshot is None:
|
||||||
|
conflicts.append(src)
|
||||||
|
continue
|
||||||
|
dst = old_dir / src.name
|
||||||
|
if dst.exists():
|
||||||
|
destination_snapshot = self._session_file_snapshot(dst)
|
||||||
|
if (
|
||||||
|
destination_snapshot is not None
|
||||||
|
and destination_snapshot.digest == source_snapshot.digest
|
||||||
|
):
|
||||||
|
unchanged += 1
|
||||||
|
else:
|
||||||
|
conflicts.append(dst)
|
||||||
|
continue
|
||||||
|
self._install_snapshot(src, dst, source_snapshot)
|
||||||
|
restored += 1
|
||||||
|
return SessionRestoreResult(
|
||||||
|
restored=restored,
|
||||||
|
unchanged=unchanged,
|
||||||
|
conflicts=tuple(conflicts),
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def safe_key(key: str) -> str:
|
def safe_key(key: str) -> str:
|
||||||
@@ -991,9 +1453,15 @@ class JsonlSessionStore:
|
|||||||
class SessionManager:
|
class SessionManager:
|
||||||
"""Manage session identity, caching, retention, and persistence."""
|
"""Manage session identity, caching, retention, and persistence."""
|
||||||
|
|
||||||
def __init__(self, workspace: Path, *, store: SessionStore | None = None):
|
def __init__(
|
||||||
|
self,
|
||||||
|
workspace: Path,
|
||||||
|
*,
|
||||||
|
store: SessionStore | None = None,
|
||||||
|
sessions_root: Path | None = None,
|
||||||
|
):
|
||||||
self.workspace = workspace
|
self.workspace = workspace
|
||||||
self._jsonl_store = JsonlSessionStore(workspace)
|
self._jsonl_store = JsonlSessionStore(workspace, sessions_root=sessions_root)
|
||||||
self._store: SessionStore = store if store is not None else self._jsonl_store
|
self._store: SessionStore = store if store is not None else self._jsonl_store
|
||||||
self.sessions_dir = self._jsonl_store.sessions_dir
|
self.sessions_dir = self._jsonl_store.sessions_dir
|
||||||
self.legacy_sessions_dir = self._jsonl_store.legacy_sessions_dir
|
self.legacy_sessions_dir = self._jsonl_store.legacy_sessions_dir
|
||||||
@@ -1159,6 +1627,10 @@ class SessionManager:
|
|||||||
self.invalidate(key)
|
self.invalidate(key)
|
||||||
return self._store.delete(key)
|
return self._store.delete(key)
|
||||||
|
|
||||||
|
def restore_sessions_to_workspace(self) -> SessionRestoreResult:
|
||||||
|
"""Restore session files to the pre-relocation path for an explicit rollback."""
|
||||||
|
return self._jsonl_store.restore_to_workspace()
|
||||||
|
|
||||||
def fork_session_before_user_index(
|
def fork_session_before_user_index(
|
||||||
self,
|
self,
|
||||||
source_key: str,
|
source_key: str,
|
||||||
|
|||||||
@@ -11,23 +11,39 @@ Two free services, no API keys needed.
|
|||||||
|
|
||||||
## wttr.in (primary)
|
## wttr.in (primary)
|
||||||
|
|
||||||
Quick one-liner:
|
Choose one request that matches the user's scope. Do not fetch current
|
||||||
|
conditions separately when a today or forecast request already includes them.
|
||||||
|
|
||||||
|
Platform notes:
|
||||||
|
- On Windows PowerShell, use `curl.exe`; bare `curl` may resolve to
|
||||||
|
`Invoke-WebRequest`.
|
||||||
|
- On macOS and Linux, use `curl`.
|
||||||
|
|
||||||
|
Current conditions only:
|
||||||
```bash
|
```bash
|
||||||
curl -s "wttr.in/London?format=3"
|
curl -s "https://wttr.in/London?format=3"
|
||||||
# Output: London: ⛅️ +8°C
|
# Output: London: ⛅️ +8°C
|
||||||
```
|
```
|
||||||
|
|
||||||
Compact format:
|
Custom current conditions format:
|
||||||
```bash
|
```bash
|
||||||
curl -s "wttr.in/London?format=%l:+%c+%t+%h+%w"
|
curl -s "https://wttr.in/London?format=%l:+%c+%t+%h+%w"
|
||||||
# Output: London: ⛅️ +8°C 71% ↙5km/h
|
# Output: London: ⛅️ +8°C 71% ↙5km/h
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Today's weather, including current conditions (use this single request for
|
||||||
|
questions about today's weather):
|
||||||
|
```bash
|
||||||
|
curl -s "https://wttr.in/London?1&m"
|
||||||
|
```
|
||||||
|
|
||||||
Full forecast:
|
Full forecast:
|
||||||
```bash
|
```bash
|
||||||
curl -s "wttr.in/London?T"
|
curl -s "https://wttr.in/London?T&m"
|
||||||
```
|
```
|
||||||
|
|
||||||
|
On Windows PowerShell, replace `curl` with `curl.exe` in the commands above.
|
||||||
|
|
||||||
Format codes: `%c` condition · `%t` temp · `%h` humidity · `%w` wind · `%l` location · `%m` moon
|
Format codes: `%c` condition · `%t` temp · `%h` humidity · `%w` wind · `%l` location · `%m` moon
|
||||||
|
|
||||||
Tips:
|
Tips:
|
||||||
@@ -35,7 +51,8 @@ Tips:
|
|||||||
- Airport codes: `wttr.in/JFK`
|
- Airport codes: `wttr.in/JFK`
|
||||||
- Units: `?m` (metric) `?u` (USCS)
|
- Units: `?m` (metric) `?u` (USCS)
|
||||||
- Today only: `?1` · Current only: `?0`
|
- Today only: `?1` · Current only: `?0`
|
||||||
- PNG: `curl -s "wttr.in/Berlin.png" -o /tmp/weather.png`
|
- PNG (macOS/Linux): `curl -s "https://wttr.in/Berlin.png" -o weather.png`
|
||||||
|
- PNG (Windows PowerShell): `curl.exe -s "https://wttr.in/Berlin.png" -o weather.png`
|
||||||
|
|
||||||
## Open-Meteo (fallback, JSON)
|
## Open-Meteo (fallback, JSON)
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
|
from pathlib import Path
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
from nanobot.apps.cli import CliAppError, CliAppManager, CliAppsRuntimeConfig
|
from nanobot.apps.cli import CliAppError, CliAppManager, CliAppsRuntimeConfig
|
||||||
@@ -89,8 +90,8 @@ def _query_first(query: QueryParams, key: str) -> str | None:
|
|||||||
return values[0] if values else None
|
return values[0] if values else None
|
||||||
|
|
||||||
|
|
||||||
def _manager() -> CliAppManager:
|
def _manager(config_path: Path | None = None) -> CliAppManager:
|
||||||
config = load_config()
|
config = load_config(config_path) if config_path is not None else load_config()
|
||||||
cli_cfg = config.tools.cli_apps
|
cli_cfg = config.tools.cli_apps
|
||||||
return CliAppManager(
|
return CliAppManager(
|
||||||
workspace=config.workspace_path,
|
workspace=config.workspace_path,
|
||||||
@@ -102,8 +103,12 @@ def _manager() -> CliAppManager:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
async def cli_apps_payload(*, installed_only: bool = False) -> dict[str, Any]:
|
async def cli_apps_payload(
|
||||||
manager = _manager()
|
*,
|
||||||
|
installed_only: bool = False,
|
||||||
|
config_path: Path | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
manager = _manager(config_path) if config_path is not None else _manager()
|
||||||
if installed_only:
|
if installed_only:
|
||||||
return manager.installed_payload()
|
return manager.installed_payload()
|
||||||
payload = manager.payload(cache_only=True)
|
payload = manager.payload(cache_only=True)
|
||||||
@@ -118,11 +123,16 @@ async def cli_apps_payload(*, installed_only: bool = False) -> dict[str, Any]:
|
|||||||
return payload
|
return payload
|
||||||
|
|
||||||
|
|
||||||
def cli_apps_action(action: str, query: QueryParams) -> dict[str, Any]:
|
def cli_apps_action(
|
||||||
|
action: str,
|
||||||
|
query: QueryParams,
|
||||||
|
*,
|
||||||
|
config_path: Path | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
name = (_query_first(query, "name") or "").strip()
|
name = (_query_first(query, "name") or "").strip()
|
||||||
if not name:
|
if not name:
|
||||||
raise CliAppError("missing CLI app name")
|
raise CliAppError("missing CLI app name")
|
||||||
manager = _manager()
|
manager = _manager(config_path) if config_path is not None else _manager()
|
||||||
if action == "install":
|
if action == "install":
|
||||||
return manager.install(name)
|
return manager.install(name)
|
||||||
if action == "update":
|
if action == "update":
|
||||||
|
|||||||
@@ -2,15 +2,18 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Awaitable, Mapping
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any, Callable
|
from typing import TYPE_CHECKING, Any, Callable
|
||||||
|
|
||||||
from loguru import logger as default_logger
|
from loguru import logger as default_logger
|
||||||
|
|
||||||
|
from nanobot.config.loader import get_config_path
|
||||||
from nanobot.webui.gateway_tokens import GatewayTokenStore
|
from nanobot.webui.gateway_tokens import GatewayTokenStore
|
||||||
from nanobot.webui.ingress_policy import DEFAULT_WEBUI_INGRESS_POLICY, WebUIIngressPolicy
|
from nanobot.webui.ingress_policy import DEFAULT_WEBUI_INGRESS_POLICY, WebUIIngressPolicy
|
||||||
from nanobot.webui.media_gateway import WebUIMediaGateway
|
from nanobot.webui.media_gateway import WebUIMediaGateway
|
||||||
|
from nanobot.webui.settings_services import WebUISettingsServices
|
||||||
from nanobot.webui.temporary_chats import WebUITemporaryChats
|
from nanobot.webui.temporary_chats import WebUITemporaryChats
|
||||||
from nanobot.webui.transcript import WebUITranscriptRecorder
|
from nanobot.webui.transcript import WebUITranscriptRecorder
|
||||||
from nanobot.webui.workspaces import WebUIWorkspaceController
|
from nanobot.webui.workspaces import WebUIWorkspaceController
|
||||||
@@ -29,6 +32,7 @@ class GatewayServices:
|
|||||||
"""Explicit dependencies shared by WebSocket transport and HTTP routes."""
|
"""Explicit dependencies shared by WebSocket transport and HTTP routes."""
|
||||||
|
|
||||||
http: GatewayHTTPHandler
|
http: GatewayHTTPHandler
|
||||||
|
settings: WebUISettingsServices
|
||||||
tokens: GatewayTokenStore
|
tokens: GatewayTokenStore
|
||||||
media: WebUIMediaGateway
|
media: WebUIMediaGateway
|
||||||
ingress: WebUIIngressPolicy
|
ingress: WebUIIngressPolicy
|
||||||
@@ -50,6 +54,7 @@ def build_gateway_services(
|
|||||||
static_dist_path: Path | None,
|
static_dist_path: Path | None,
|
||||||
workspace_path: Path,
|
workspace_path: Path,
|
||||||
default_restrict_to_workspace: bool,
|
default_restrict_to_workspace: bool,
|
||||||
|
config_path: Path | None = None,
|
||||||
runtime_model_name: Callable[[], str | None] | None,
|
runtime_model_name: Callable[[], str | None] | None,
|
||||||
runtime_surface: str,
|
runtime_surface: str,
|
||||||
runtime_capabilities_overrides: dict[str, Any] | None,
|
runtime_capabilities_overrides: dict[str, Any] | None,
|
||||||
@@ -60,9 +65,12 @@ def build_gateway_services(
|
|||||||
local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
||||||
channel_feature_action: Callable[..., Any] | None = None,
|
channel_feature_action: Callable[..., Any] | None = None,
|
||||||
channel_runtime_status: Callable[[], dict[str, Any]] | None = None,
|
channel_runtime_status: Callable[[], dict[str, Any]] | None = None,
|
||||||
|
mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
|
||||||
|
mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
|
||||||
skill_state_action: Callable[[set[str]], None] | None = None,
|
skill_state_action: Callable[[set[str]], None] | None = None,
|
||||||
logger: Any = default_logger,
|
logger: Any = default_logger,
|
||||||
) -> GatewayServices:
|
) -> GatewayServices:
|
||||||
|
settings = WebUISettingsServices.create(config_path or get_config_path())
|
||||||
tokens = GatewayTokenStore()
|
tokens = GatewayTokenStore()
|
||||||
ingress = DEFAULT_WEBUI_INGRESS_POLICY
|
ingress = DEFAULT_WEBUI_INGRESS_POLICY
|
||||||
minimum_frame_bytes = ingress.minimum_full_policy_frame_bytes()
|
minimum_frame_bytes = ingress.minimum_full_policy_frame_bytes()
|
||||||
@@ -102,6 +110,7 @@ def build_gateway_services(
|
|||||||
media=media,
|
media=media,
|
||||||
ingress=ingress,
|
ingress=ingress,
|
||||||
workspaces=workspaces,
|
workspaces=workspaces,
|
||||||
|
settings=settings,
|
||||||
skills_workspace_path=workspace_path,
|
skills_workspace_path=workspace_path,
|
||||||
disabled_skills=disabled_skills,
|
disabled_skills=disabled_skills,
|
||||||
cron_service=cron_service,
|
cron_service=cron_service,
|
||||||
@@ -110,11 +119,14 @@ def build_gateway_services(
|
|||||||
local_trigger_pending_ids=local_trigger_pending_ids,
|
local_trigger_pending_ids=local_trigger_pending_ids,
|
||||||
channel_feature_action=channel_feature_action,
|
channel_feature_action=channel_feature_action,
|
||||||
channel_runtime_status=channel_runtime_status,
|
channel_runtime_status=channel_runtime_status,
|
||||||
|
mcp_runtime_status=mcp_runtime_status,
|
||||||
|
mcp_reload=mcp_reload,
|
||||||
skill_state_action=skill_state_action,
|
skill_state_action=skill_state_action,
|
||||||
log=logger,
|
log=logger,
|
||||||
)
|
)
|
||||||
return GatewayServices(
|
return GatewayServices(
|
||||||
http=http,
|
http=http,
|
||||||
|
settings=settings,
|
||||||
tokens=tokens,
|
tokens=tokens,
|
||||||
media=media,
|
media=media,
|
||||||
ingress=ingress,
|
ingress=ingress,
|
||||||
|
|||||||
@@ -100,6 +100,7 @@ def http_json_response(
|
|||||||
*,
|
*,
|
||||||
status: int = 200,
|
status: int = 200,
|
||||||
accept_encoding: str | None = None,
|
accept_encoding: str | None = None,
|
||||||
|
extra_headers: list[tuple[str, str]] | None = None,
|
||||||
) -> Response:
|
) -> Response:
|
||||||
body = json.dumps(data, ensure_ascii=False).encode("utf-8")
|
body = json.dumps(data, ensure_ascii=False).encode("utf-8")
|
||||||
headers = [
|
headers = [
|
||||||
@@ -112,6 +113,8 @@ def http_json_response(
|
|||||||
if len(body) >= _JSON_GZIP_MIN_BYTES and accepts_gzip(accept_encoding):
|
if len(body) >= _JSON_GZIP_MIN_BYTES and accepts_gzip(accept_encoding):
|
||||||
body = gzip.compress(body, compresslevel=_JSON_GZIP_LEVEL, mtime=0)
|
body = gzip.compress(body, compresslevel=_JSON_GZIP_LEVEL, mtime=0)
|
||||||
headers.append(("Content-Encoding", "gzip"))
|
headers.append(("Content-Encoding", "gzip"))
|
||||||
|
if extra_headers:
|
||||||
|
headers.extend(extra_headers)
|
||||||
headers.append(("Content-Length", str(len(body))))
|
headers.append(("Content-Length", str(len(body))))
|
||||||
reason = http.HTTPStatus(status).phrase
|
reason = http.HTTPStatus(status).phrase
|
||||||
return Response(status, reason, Headers(headers), body)
|
return Response(status, reason, Headers(headers), body)
|
||||||
|
|||||||
@@ -0,0 +1,415 @@
|
|||||||
|
"""Gateway-owned browser authorization flows for remote MCP servers."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import re
|
||||||
|
import secrets
|
||||||
|
import time
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from contextlib import suppress
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
from urllib.parse import SplitResult, parse_qs, urlsplit, urlunsplit
|
||||||
|
|
||||||
|
from nanobot.agent.tools.mcp import MCPConnection, connect_mcp_servers
|
||||||
|
from nanobot.agent.tools.mcp_oauth import MCP_OAUTH_CALLBACK_PATH, MCPOAuthHandlers
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
|
from nanobot.config.schema import MCPServerConfig
|
||||||
|
from nanobot.security.network import validate_url_target
|
||||||
|
from nanobot.webui.http_utils import is_loopback_host
|
||||||
|
|
||||||
|
McpReload = Callable[[], Awaitable[dict[str, Any]]]
|
||||||
|
_FLOW_TTL_S = 300
|
||||||
|
_START_WAIT_S = 20
|
||||||
|
_OAUTH_ERROR_RE = re.compile(r"^[a-zA-Z0-9_.-]{1,80}$")
|
||||||
|
|
||||||
|
|
||||||
|
class McpOAuthError(Exception):
|
||||||
|
"""Safe WebUI error for an MCP OAuth request."""
|
||||||
|
|
||||||
|
def __init__(self, message: str, status: int = 400) -> None:
|
||||||
|
super().__init__(message)
|
||||||
|
self.message = message
|
||||||
|
self.status = status
|
||||||
|
|
||||||
|
|
||||||
|
class _OAuthCallbackError(RuntimeError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _McpOAuthFlow:
|
||||||
|
flow_id: str
|
||||||
|
name: str
|
||||||
|
cfg: MCPServerConfig
|
||||||
|
redirect_uri: str
|
||||||
|
manual_callback: bool
|
||||||
|
expires_at: float
|
||||||
|
authorization_ready: asyncio.Event = field(default_factory=asyncio.Event)
|
||||||
|
callback_result: asyncio.Future[tuple[str, str | None]] | None = None
|
||||||
|
task: asyncio.Task[bool] | None = None
|
||||||
|
authorization_url: str | None = None
|
||||||
|
state: str | None = None
|
||||||
|
callback_received: bool = False
|
||||||
|
error: str | None = None
|
||||||
|
reload_result: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_mcp_oauth_redirect_uri(redirect_uri: str) -> tuple[str, SplitResult, int | None]:
|
||||||
|
cleaned = redirect_uri.strip()
|
||||||
|
parsed = urlsplit(cleaned)
|
||||||
|
try:
|
||||||
|
port = parsed.port
|
||||||
|
except ValueError as exc:
|
||||||
|
raise McpOAuthError("Invalid MCP OAuth callback URL") from exc
|
||||||
|
if (
|
||||||
|
not parsed.netloc
|
||||||
|
or not parsed.hostname
|
||||||
|
or parsed.path != MCP_OAUTH_CALLBACK_PATH
|
||||||
|
or parsed.query
|
||||||
|
or parsed.fragment
|
||||||
|
or parsed.username is not None
|
||||||
|
or parsed.password is not None
|
||||||
|
):
|
||||||
|
raise McpOAuthError("Invalid MCP OAuth callback URL")
|
||||||
|
return cleaned, parsed, port
|
||||||
|
|
||||||
|
|
||||||
|
def validate_mcp_oauth_redirect_uri(redirect_uri: str) -> str:
|
||||||
|
"""Allow HTTPS callbacks, plus loopback HTTP for a local gateway."""
|
||||||
|
cleaned, parsed, _port = _parse_mcp_oauth_redirect_uri(redirect_uri)
|
||||||
|
if parsed.scheme == "https":
|
||||||
|
return cleaned
|
||||||
|
if parsed.scheme == "http" and is_loopback_host(parsed.netloc):
|
||||||
|
return cleaned
|
||||||
|
raise McpOAuthError("MCP OAuth callbacks must use HTTPS or localhost")
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_mcp_oauth_redirect_uri(redirect_uri: str) -> tuple[str, bool]:
|
||||||
|
"""Use a pasteable loopback callback when a remote WebUI is served over HTTP."""
|
||||||
|
cleaned, parsed, port = _parse_mcp_oauth_redirect_uri(redirect_uri)
|
||||||
|
if parsed.scheme != "http" or is_loopback_host(parsed.netloc):
|
||||||
|
return validate_mcp_oauth_redirect_uri(cleaned), False
|
||||||
|
|
||||||
|
loopback = "127.0.0.1" if port is None else f"127.0.0.1:{port}"
|
||||||
|
manual_redirect_uri = urlunsplit(("http", loopback, parsed.path, "", ""))
|
||||||
|
return validate_mcp_oauth_redirect_uri(manual_redirect_uri), True
|
||||||
|
|
||||||
|
|
||||||
|
class McpOAuthManager:
|
||||||
|
"""Own short-lived browser flows while the gateway process is running."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._flows: dict[str, _McpOAuthFlow] = {}
|
||||||
|
self._states: dict[str, str] = {}
|
||||||
|
|
||||||
|
async def start(
|
||||||
|
self,
|
||||||
|
name: str,
|
||||||
|
cfg: MCPServerConfig,
|
||||||
|
redirect_uri: str,
|
||||||
|
*,
|
||||||
|
reload_mcp: McpReload,
|
||||||
|
reset_credentials: bool = False,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
self._prune()
|
||||||
|
redirect_uri, manual_callback = prepare_mcp_oauth_redirect_uri(redirect_uri)
|
||||||
|
await self._cancel_name(name)
|
||||||
|
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
now = time.monotonic()
|
||||||
|
flow = _McpOAuthFlow(
|
||||||
|
flow_id=secrets.token_urlsafe(24),
|
||||||
|
name=name,
|
||||||
|
cfg=cfg,
|
||||||
|
redirect_uri=redirect_uri,
|
||||||
|
manual_callback=manual_callback,
|
||||||
|
expires_at=now + _FLOW_TTL_S,
|
||||||
|
callback_result=loop.create_future(),
|
||||||
|
)
|
||||||
|
self._flows[flow.flow_id] = flow
|
||||||
|
handlers = MCPOAuthHandlers(
|
||||||
|
redirect_uri=redirect_uri,
|
||||||
|
redirect_handler=lambda url: self._receive_authorization_url(flow, url),
|
||||||
|
callback_handler=lambda: self._wait_for_callback(flow),
|
||||||
|
reset_credentials=reset_credentials,
|
||||||
|
)
|
||||||
|
flow.task = asyncio.create_task(
|
||||||
|
self._connect_and_reload(flow, handlers, reload_mcp),
|
||||||
|
name=f"mcp-oauth:{name}",
|
||||||
|
)
|
||||||
|
|
||||||
|
ready_waiter = asyncio.create_task(flow.authorization_ready.wait())
|
||||||
|
try:
|
||||||
|
await asyncio.wait(
|
||||||
|
{ready_waiter, flow.task},
|
||||||
|
timeout=_START_WAIT_S,
|
||||||
|
return_when=asyncio.FIRST_COMPLETED,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
ready_waiter.cancel()
|
||||||
|
with suppress(asyncio.CancelledError):
|
||||||
|
await ready_waiter
|
||||||
|
return self._payload(flow)
|
||||||
|
|
||||||
|
async def status(self, flow_id: str) -> dict[str, Any]:
|
||||||
|
self._prune()
|
||||||
|
flow = self._flow(flow_id)
|
||||||
|
return self._payload(flow)
|
||||||
|
|
||||||
|
def submit_callback(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
state: str,
|
||||||
|
code: str | None,
|
||||||
|
error: str | None,
|
||||||
|
) -> str:
|
||||||
|
self._prune()
|
||||||
|
flow_id = self._states.pop(state, None)
|
||||||
|
if flow_id is None:
|
||||||
|
raise McpOAuthError("This MCP authorization request has expired", status=410)
|
||||||
|
flow = self._flow(flow_id)
|
||||||
|
callback_result = flow.callback_result
|
||||||
|
if callback_result is None or callback_result.done():
|
||||||
|
raise McpOAuthError("This MCP authorization callback was already used", status=409)
|
||||||
|
|
||||||
|
flow.callback_received = True
|
||||||
|
if error:
|
||||||
|
safe_error = error if _OAUTH_ERROR_RE.fullmatch(error) else "authorization_failed"
|
||||||
|
flow.error = f"Authorization was not completed ({safe_error})."
|
||||||
|
callback_result.set_exception(_OAuthCallbackError(flow.error))
|
||||||
|
raise McpOAuthError(flow.error)
|
||||||
|
elif not code or len(code) > 8192:
|
||||||
|
flow.error = "The MCP server did not return an authorization code."
|
||||||
|
callback_result.set_exception(_OAuthCallbackError(flow.error))
|
||||||
|
raise McpOAuthError(flow.error)
|
||||||
|
else:
|
||||||
|
callback_result.set_result((code, state))
|
||||||
|
return flow.name
|
||||||
|
|
||||||
|
def submit_callback_url(self, *, flow_id: str, callback_url: str) -> dict[str, Any]:
|
||||||
|
"""Complete a flow from a full browser callback URL pasted into the WebUI."""
|
||||||
|
self._prune()
|
||||||
|
flow = self._flow(flow_id)
|
||||||
|
parsed = urlsplit(callback_url.strip())
|
||||||
|
expected = urlsplit(flow.redirect_uri)
|
||||||
|
if (
|
||||||
|
not parsed.query
|
||||||
|
or parsed.fragment
|
||||||
|
or parsed.username is not None
|
||||||
|
or parsed.password is not None
|
||||||
|
or parsed.scheme != expected.scheme
|
||||||
|
or parsed.netloc != expected.netloc
|
||||||
|
or parsed.path != expected.path
|
||||||
|
):
|
||||||
|
raise McpOAuthError(
|
||||||
|
"Paste the complete callback URL from the browser address bar."
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
query = parse_qs(parsed.query, keep_blank_values=True, max_num_fields=16)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise McpOAuthError(
|
||||||
|
"Paste the complete callback URL from the browser address bar."
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
states = query.get("state", [])
|
||||||
|
state = states[0] if len(states) == 1 else ""
|
||||||
|
if not state or state != flow.state:
|
||||||
|
raise McpOAuthError(
|
||||||
|
"This callback belongs to a different or expired authorization request. "
|
||||||
|
"Start again.",
|
||||||
|
status=410,
|
||||||
|
)
|
||||||
|
|
||||||
|
codes = query.get("code", [])
|
||||||
|
errors = query.get("error", [])
|
||||||
|
if len(codes) > 1 or len(errors) > 1 or (codes and errors):
|
||||||
|
raise McpOAuthError(
|
||||||
|
"Paste the complete callback URL from the browser address bar."
|
||||||
|
)
|
||||||
|
code = codes[0] if len(codes) == 1 else None
|
||||||
|
error = errors[0] if len(errors) == 1 else None
|
||||||
|
if (not code and not error) or (code is not None and len(code) > 8192):
|
||||||
|
raise McpOAuthError(
|
||||||
|
"Paste the complete callback URL from the browser address bar."
|
||||||
|
)
|
||||||
|
|
||||||
|
self.submit_callback(state=state, code=code, error=error)
|
||||||
|
return self._payload(flow)
|
||||||
|
|
||||||
|
async def cancel(self, flow_id: str) -> dict[str, Any]:
|
||||||
|
self._prune()
|
||||||
|
flow = self._flow(flow_id)
|
||||||
|
await self._cancel_flow(flow)
|
||||||
|
return self._payload(flow)
|
||||||
|
|
||||||
|
async def _receive_authorization_url(
|
||||||
|
self,
|
||||||
|
flow: _McpOAuthFlow,
|
||||||
|
authorization_url: str,
|
||||||
|
) -> None:
|
||||||
|
parsed = urlsplit(authorization_url)
|
||||||
|
if (
|
||||||
|
parsed.scheme != "https"
|
||||||
|
or not parsed.netloc
|
||||||
|
or parsed.username is not None
|
||||||
|
or parsed.password is not None
|
||||||
|
or parsed.fragment
|
||||||
|
):
|
||||||
|
flow.error = "The MCP server returned an unsafe authorization URL."
|
||||||
|
raise McpOAuthError(flow.error)
|
||||||
|
ok, _error = validate_url_target(authorization_url)
|
||||||
|
if not ok:
|
||||||
|
flow.error = "The MCP server returned an unsafe authorization URL."
|
||||||
|
raise McpOAuthError(flow.error)
|
||||||
|
states = parse_qs(parsed.query).get("state", [])
|
||||||
|
state = states[0] if len(states) == 1 else ""
|
||||||
|
if not state or len(state) > 512:
|
||||||
|
flow.error = "The MCP server returned an invalid authorization URL."
|
||||||
|
raise McpOAuthError(flow.error)
|
||||||
|
if state in self._states:
|
||||||
|
flow.error = "The MCP server reused an OAuth state value."
|
||||||
|
raise McpOAuthError(flow.error)
|
||||||
|
flow.authorization_url = authorization_url
|
||||||
|
flow.state = state
|
||||||
|
self._states[state] = flow.flow_id
|
||||||
|
flow.authorization_ready.set()
|
||||||
|
|
||||||
|
async def _wait_for_callback(self, flow: _McpOAuthFlow) -> tuple[str, str | None]:
|
||||||
|
callback_result = flow.callback_result
|
||||||
|
if callback_result is None:
|
||||||
|
raise _OAuthCallbackError("MCP OAuth callback is unavailable")
|
||||||
|
remaining = max(0.1, flow.expires_at - time.monotonic())
|
||||||
|
try:
|
||||||
|
return await asyncio.wait_for(asyncio.shield(callback_result), timeout=remaining)
|
||||||
|
except asyncio.TimeoutError as exc:
|
||||||
|
flow.error = "MCP authorization timed out."
|
||||||
|
raise _OAuthCallbackError(flow.error) from exc
|
||||||
|
|
||||||
|
async def _connect(self, flow: _McpOAuthFlow, handlers: MCPOAuthHandlers) -> bool:
|
||||||
|
connections: dict[str, MCPConnection] = {}
|
||||||
|
try:
|
||||||
|
connections = await connect_mcp_servers(
|
||||||
|
{flow.name: flow.cfg},
|
||||||
|
ToolRegistry(),
|
||||||
|
oauth_handlers={flow.name: handlers},
|
||||||
|
)
|
||||||
|
succeeded = flow.name in connections
|
||||||
|
if not succeeded and flow.error is None:
|
||||||
|
flow.error = "Could not complete the MCP OAuth connection."
|
||||||
|
return succeeded
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception:
|
||||||
|
if flow.error is None:
|
||||||
|
flow.error = "Could not complete the MCP OAuth connection."
|
||||||
|
return False
|
||||||
|
finally:
|
||||||
|
for connection in connections.values():
|
||||||
|
with suppress(Exception):
|
||||||
|
await connection.aclose()
|
||||||
|
|
||||||
|
async def _connect_and_reload(
|
||||||
|
self,
|
||||||
|
flow: _McpOAuthFlow,
|
||||||
|
handlers: MCPOAuthHandlers,
|
||||||
|
reload_mcp: McpReload,
|
||||||
|
) -> bool:
|
||||||
|
succeeded = await self._connect(flow, handlers)
|
||||||
|
if not succeeded:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
flow.reload_result = await reload_mcp()
|
||||||
|
failed = flow.reload_result.get("failed")
|
||||||
|
if (
|
||||||
|
not flow.reload_result.get("ok")
|
||||||
|
and not flow.reload_result.get("requires_restart")
|
||||||
|
and isinstance(failed, list)
|
||||||
|
and flow.name in failed
|
||||||
|
):
|
||||||
|
flow.reload_result = await reload_mcp()
|
||||||
|
except Exception:
|
||||||
|
flow.reload_result = {
|
||||||
|
"ok": False,
|
||||||
|
"message": "Signed in, but nanobot could not activate the MCP tools.",
|
||||||
|
"requires_restart": True,
|
||||||
|
}
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _flow(self, flow_id: str) -> _McpOAuthFlow:
|
||||||
|
flow = self._flows.get(flow_id)
|
||||||
|
if flow is None:
|
||||||
|
raise McpOAuthError("Unknown or expired MCP OAuth flow", status=404)
|
||||||
|
return flow
|
||||||
|
|
||||||
|
def _payload(self, flow: _McpOAuthFlow) -> dict[str, Any]:
|
||||||
|
task = flow.task
|
||||||
|
connected = flow.reload_result.get("connected") if flow.reload_result is not None else None
|
||||||
|
if task is not None and task.cancelled():
|
||||||
|
status = "cancelled"
|
||||||
|
elif task is not None and task.done():
|
||||||
|
try:
|
||||||
|
succeeded = task.result()
|
||||||
|
except Exception:
|
||||||
|
succeeded = False
|
||||||
|
if not succeeded:
|
||||||
|
status = "failed"
|
||||||
|
elif flow.reload_result is None:
|
||||||
|
status = "authorized"
|
||||||
|
elif flow.reload_result.get("ok") or (
|
||||||
|
isinstance(connected, list) and flow.name in connected
|
||||||
|
):
|
||||||
|
status = "connected"
|
||||||
|
else:
|
||||||
|
status = "authorized"
|
||||||
|
elif flow.callback_received:
|
||||||
|
status = "connecting"
|
||||||
|
elif flow.authorization_url:
|
||||||
|
status = "authorization_required"
|
||||||
|
else:
|
||||||
|
status = "starting"
|
||||||
|
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"flow_id": flow.flow_id,
|
||||||
|
"name": flow.name,
|
||||||
|
"status": status,
|
||||||
|
"expires_in": max(0, int(flow.expires_at - time.monotonic())),
|
||||||
|
}
|
||||||
|
if flow.manual_callback:
|
||||||
|
payload["completion_input"] = "callback_url"
|
||||||
|
if flow.authorization_url and status == "authorization_required":
|
||||||
|
payload["authorization_url"] = flow.authorization_url
|
||||||
|
if flow.error:
|
||||||
|
payload["error"] = flow.error
|
||||||
|
if flow.reload_result is not None:
|
||||||
|
payload["hot_reload"] = flow.reload_result
|
||||||
|
return payload
|
||||||
|
|
||||||
|
async def _cancel_name(self, name: str) -> None:
|
||||||
|
for flow in list(self._flows.values()):
|
||||||
|
if flow.name == name and flow.task is not None and not flow.task.done():
|
||||||
|
await self._cancel_flow(flow)
|
||||||
|
|
||||||
|
async def _cancel_flow(self, flow: _McpOAuthFlow) -> None:
|
||||||
|
if flow.state:
|
||||||
|
self._states.pop(flow.state, None)
|
||||||
|
task = flow.task
|
||||||
|
if task is not None and not task.done():
|
||||||
|
task.cancel()
|
||||||
|
with suppress(BaseException):
|
||||||
|
await task
|
||||||
|
|
||||||
|
def _prune(self) -> None:
|
||||||
|
now = time.monotonic()
|
||||||
|
for flow_id, flow in list(self._flows.items()):
|
||||||
|
if flow.expires_at > now:
|
||||||
|
continue
|
||||||
|
if flow.state:
|
||||||
|
self._states.pop(flow.state, None)
|
||||||
|
if flow.task is not None and not flow.task.done():
|
||||||
|
flow.task.cancel()
|
||||||
|
callback_result = flow.callback_result
|
||||||
|
if callback_result is not None and not callback_result.done():
|
||||||
|
callback_result.cancel()
|
||||||
|
self._flows.pop(flow_id, None)
|
||||||
@@ -14,8 +14,17 @@ from contextlib import suppress
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Literal, Mapping, cast
|
from typing import TYPE_CHECKING, Any, Literal, Mapping, cast
|
||||||
|
|
||||||
|
from nanobot.agent.plugins import (
|
||||||
|
AgentPlugin,
|
||||||
|
discover_agent_plugins,
|
||||||
|
set_agent_plugin_enabled,
|
||||||
|
)
|
||||||
|
from nanobot.agent.tools.mcp_oauth import (
|
||||||
|
delete_mcp_oauth_credentials,
|
||||||
|
mcp_oauth_has_credentials,
|
||||||
|
)
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.apps.protocol import app_manifest, compact_dict
|
from nanobot.apps.protocol import app_manifest, compact_dict
|
||||||
from nanobot.config.loader import load_config, resolve_config_env_vars, save_config
|
from nanobot.config.loader import load_config, resolve_config_env_vars, save_config
|
||||||
@@ -25,6 +34,9 @@ from nanobot.utils.helpers import ensure_dir
|
|||||||
|
|
||||||
QueryParams = dict[str, list[str]]
|
QueryParams = dict[str, list[str]]
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from nanobot.webui.settings_services import WebUISettingsConfig
|
||||||
|
|
||||||
_MCP_PRESET_NAME_RE = re.compile(r"^[a-z0-9][a-z0-9_-]{0,63}$", re.IGNORECASE)
|
_MCP_PRESET_NAME_RE = re.compile(r"^[a-z0-9][a-z0-9_-]{0,63}$", re.IGNORECASE)
|
||||||
_SECRET_QUERY_RE = re.compile(
|
_SECRET_QUERY_RE = re.compile(
|
||||||
r"([?&](?:[^=&]*(?:api[_-]?key|token|secret|password|bearer)[^=&]*)=)[^&#\s]+",
|
r"([?&](?:[^=&]*(?:api[_-]?key|token|secret|password|bearer)[^=&]*)=)[^&#\s]+",
|
||||||
@@ -44,12 +56,13 @@ _MCP_ATTACHMENT_KEYS = (
|
|||||||
"status",
|
"status",
|
||||||
"configured",
|
"configured",
|
||||||
)
|
)
|
||||||
_MAX_TEST_TOOLS = 16
|
|
||||||
_DEFAULT_TEST_TIMEOUT = 20
|
_DEFAULT_TEST_TIMEOUT = 20
|
||||||
_DEFAULT_CUSTOM_TIMEOUT = 30
|
_DEFAULT_CUSTOM_TIMEOUT = 30
|
||||||
_CUSTOM_ACTIONS = {"custom", "import", "import-cursor", "tools"}
|
_CUSTOM_ACTIONS = {"custom", "import", "import-cursor", "tools"}
|
||||||
|
_MCP_RUNTIME_STATUSES = {"connecting", "connected", "failed"}
|
||||||
|
|
||||||
McpReload = Callable[[], Awaitable[dict[str, Any]]]
|
McpReload = Callable[[], Awaitable[dict[str, Any]]]
|
||||||
|
McpRuntimeStatus = Callable[[], Mapping[str, str]]
|
||||||
|
|
||||||
|
|
||||||
class McpPresetError(Exception):
|
class McpPresetError(Exception):
|
||||||
@@ -334,6 +347,63 @@ MCP_PRESETS: tuple[McpPreset, ...] = (
|
|||||||
),
|
),
|
||||||
note="Requires Figma Desktop Dev Mode MCP to be running locally.",
|
note="Requires Figma Desktop Dev Mode MCP to be running locally.",
|
||||||
),
|
),
|
||||||
|
McpPreset(
|
||||||
|
name="xmind",
|
||||||
|
display_name="Xmind",
|
||||||
|
category="productivity",
|
||||||
|
description="Create, read, and edit cloud mind maps through Xmind.",
|
||||||
|
docs_url="https://xmind.com/user-guide/xmind-mcp",
|
||||||
|
transport="streamableHttp",
|
||||||
|
install_supported=True,
|
||||||
|
brand_domain="xmind.com",
|
||||||
|
brand_color="#F4B41A",
|
||||||
|
requires="Xmind account",
|
||||||
|
server=MCPServerConfig(
|
||||||
|
type="streamableHttp",
|
||||||
|
auth="oauth",
|
||||||
|
url="https://app.xmind.com/api/mcp",
|
||||||
|
tool_timeout=60,
|
||||||
|
),
|
||||||
|
note="Connects securely in your browser with Xmind OAuth.",
|
||||||
|
),
|
||||||
|
McpPreset(
|
||||||
|
name="notion",
|
||||||
|
display_name="Notion",
|
||||||
|
category="productivity",
|
||||||
|
description="Read and update your Notion workspace through Notion MCP.",
|
||||||
|
docs_url="https://developers.notion.com/guides/mcp/get-started-with-mcp",
|
||||||
|
transport="streamableHttp",
|
||||||
|
install_supported=True,
|
||||||
|
brand_domain="notion.so",
|
||||||
|
brand_color="#111111",
|
||||||
|
requires="Notion account",
|
||||||
|
server=MCPServerConfig(
|
||||||
|
type="streamableHttp",
|
||||||
|
auth="oauth",
|
||||||
|
url="https://mcp.notion.com/mcp",
|
||||||
|
tool_timeout=60,
|
||||||
|
),
|
||||||
|
note="Connects securely in your browser with Notion OAuth.",
|
||||||
|
),
|
||||||
|
McpPreset(
|
||||||
|
name="linear",
|
||||||
|
display_name="Linear",
|
||||||
|
category="productivity",
|
||||||
|
description="Find and manage Linear issues, projects, and comments.",
|
||||||
|
docs_url="https://linear.app/docs/mcp",
|
||||||
|
transport="streamableHttp",
|
||||||
|
install_supported=True,
|
||||||
|
brand_domain="linear.app",
|
||||||
|
brand_color="#5E6AD2",
|
||||||
|
requires="Linear account",
|
||||||
|
server=MCPServerConfig(
|
||||||
|
type="streamableHttp",
|
||||||
|
auth="oauth",
|
||||||
|
url="https://mcp.linear.app/mcp",
|
||||||
|
tool_timeout=60,
|
||||||
|
),
|
||||||
|
note="Connects securely in your browser with Linear OAuth.",
|
||||||
|
),
|
||||||
McpPreset(
|
McpPreset(
|
||||||
name="github",
|
name="github",
|
||||||
display_name="GitHub",
|
display_name="GitHub",
|
||||||
@@ -654,6 +724,8 @@ def _status_for(preset: McpPreset, cfg: MCPServerConfig | None) -> str:
|
|||||||
return "not_installed" if preset.install_supported else "coming_soon"
|
return "not_installed" if preset.install_supported else "coming_soon"
|
||||||
if any(field.required and not _field_configured(field, cfg) for field in preset.fields):
|
if any(field.required and not _field_configured(field, cfg) for field in preset.fields):
|
||||||
return "missing_credentials"
|
return "missing_credentials"
|
||||||
|
if cfg.auth == "oauth" and not mcp_oauth_has_credentials(preset.name, cfg.url):
|
||||||
|
return "authorization_required"
|
||||||
if cfg.command and not _command_available(cfg.command):
|
if cfg.command and not _command_available(cfg.command):
|
||||||
return "missing_dependency"
|
return "missing_dependency"
|
||||||
return "configured"
|
return "configured"
|
||||||
@@ -699,6 +771,7 @@ def _preset_manifest(preset: McpPreset, *, logo_url: str) -> dict[str, Any]:
|
|||||||
compact_dict({
|
compact_dict({
|
||||||
"type": "mcp",
|
"type": "mcp",
|
||||||
"transport": preset.transport,
|
"transport": preset.transport,
|
||||||
|
"auth": server.auth if server and server.auth else None,
|
||||||
"command": server.command if server and server.command else None,
|
"command": server.command if server and server.command else None,
|
||||||
"args": list(server.args) if server and server.command else None,
|
"args": list(server.args) if server and server.command else None,
|
||||||
"url": _connection_summary(server) if server and server.url else None,
|
"url": _connection_summary(server) if server and server.url else None,
|
||||||
@@ -749,6 +822,7 @@ def _custom_manifest(name: str, cfg: MCPServerConfig) -> dict[str, Any]:
|
|||||||
compact_dict({
|
compact_dict({
|
||||||
"type": "mcp",
|
"type": "mcp",
|
||||||
"transport": transport,
|
"transport": transport,
|
||||||
|
"auth": cfg.auth,
|
||||||
"command": cfg.command or None,
|
"command": cfg.command or None,
|
||||||
"url": _connection_summary(cfg) if cfg.url else None,
|
"url": _connection_summary(cfg) if cfg.url else None,
|
||||||
})
|
})
|
||||||
@@ -776,7 +850,7 @@ def _custom_manifest(name: str, cfg: MCPServerConfig) -> dict[str, Any]:
|
|||||||
def _preset_payload(preset: McpPreset, configured_servers: dict[str, MCPServerConfig]) -> dict[str, Any]:
|
def _preset_payload(preset: McpPreset, configured_servers: dict[str, MCPServerConfig]) -> dict[str, Any]:
|
||||||
cfg = configured_servers.get(preset.name)
|
cfg = configured_servers.get(preset.name)
|
||||||
status = _status_for(preset, cfg)
|
status = _status_for(preset, cfg)
|
||||||
configured = cfg is not None and status not in {"missing_credentials"}
|
configured = cfg is not None and status not in {"missing_credentials", "authorization_required"}
|
||||||
logo_url = _favicon_url(preset.brand_domain)
|
logo_url = _favicon_url(preset.brand_domain)
|
||||||
return {
|
return {
|
||||||
"name": preset.name,
|
"name": preset.name,
|
||||||
@@ -785,6 +859,7 @@ def _preset_payload(preset: McpPreset, configured_servers: dict[str, MCPServerCo
|
|||||||
"description": preset.description,
|
"description": preset.description,
|
||||||
"docs_url": preset.docs_url,
|
"docs_url": preset.docs_url,
|
||||||
"transport": preset.transport,
|
"transport": preset.transport,
|
||||||
|
"auth": (cfg.auth if cfg is not None else (preset.server.auth if preset.server else None)),
|
||||||
"requires": preset.requires,
|
"requires": preset.requires,
|
||||||
"note": preset.note,
|
"note": preset.note,
|
||||||
"install_supported": preset.install_supported,
|
"install_supported": preset.install_supported,
|
||||||
@@ -811,7 +886,11 @@ def _custom_payload(
|
|||||||
transport = cfg.type
|
transport = cfg.type
|
||||||
if not transport:
|
if not transport:
|
||||||
transport = "stdio" if cfg.command else ("sse" if cfg.url.rstrip("/").endswith("/sse") else "streamableHttp")
|
transport = "stdio" if cfg.command else ("sse" if cfg.url.rstrip("/").endswith("/sse") else "streamableHttp")
|
||||||
|
if cfg.auth == "oauth" and not mcp_oauth_has_credentials(name, cfg.url):
|
||||||
|
status = "authorization_required"
|
||||||
|
else:
|
||||||
status = "missing_dependency" if cfg.command and not _command_available(cfg.command) else "configured"
|
status = "missing_dependency" if cfg.command and not _command_available(cfg.command) else "configured"
|
||||||
|
configured = status != "authorization_required"
|
||||||
return {
|
return {
|
||||||
"name": name,
|
"name": name,
|
||||||
"display_name": name,
|
"display_name": name,
|
||||||
@@ -819,12 +898,13 @@ def _custom_payload(
|
|||||||
"description": "Custom MCP server from nanobot config.",
|
"description": "Custom MCP server from nanobot config.",
|
||||||
"docs_url": "",
|
"docs_url": "",
|
||||||
"transport": transport,
|
"transport": transport,
|
||||||
|
"auth": cfg.auth,
|
||||||
"requires": "",
|
"requires": "",
|
||||||
"note": "",
|
"note": "",
|
||||||
"install_supported": True,
|
"install_supported": True,
|
||||||
"installed": True,
|
"installed": True,
|
||||||
"configured": True,
|
"configured": configured,
|
||||||
"available": _config_available(cfg),
|
"available": configured and _config_available(cfg),
|
||||||
"status": status,
|
"status": status,
|
||||||
"logo_url": None,
|
"logo_url": None,
|
||||||
"brand_color": "#64748B",
|
"brand_color": "#64748B",
|
||||||
@@ -837,12 +917,38 @@ def _custom_payload(
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _agent_plugin_payload(plugin: AgentPlugin) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"name": f"plugin-{plugin.name}",
|
||||||
|
"display_name": plugin.display_name,
|
||||||
|
"category": plugin.category,
|
||||||
|
"description": plugin.description or "Agent Plugin",
|
||||||
|
"docs_url": plugin.repository,
|
||||||
|
"transport": "stdio",
|
||||||
|
"requires": ", ".join(plugin.permissions),
|
||||||
|
"note": "",
|
||||||
|
"install_supported": False,
|
||||||
|
"installed": True,
|
||||||
|
"configured": True,
|
||||||
|
"enabled": plugin.enabled,
|
||||||
|
"available": plugin.enabled,
|
||||||
|
"status": "enabled" if plugin.enabled else "disabled",
|
||||||
|
"logo_url": plugin.logo,
|
||||||
|
"brand_color": plugin.accent_color,
|
||||||
|
"required_fields": [],
|
||||||
|
"connection_summary": ", ".join(plugin.mcp_servers),
|
||||||
|
"source": "agent-plugin",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def mcp_presets_payload(
|
def mcp_presets_payload(
|
||||||
*,
|
*,
|
||||||
last_action: dict[str, Any] | None = None,
|
last_action: dict[str, Any] | None = None,
|
||||||
tool_preview: Mapping[str, list[str]] | None = None,
|
tool_preview: Mapping[str, list[str]] | None = None,
|
||||||
|
runtime_status: Mapping[str, str] | None = None,
|
||||||
|
config_path: Path | None = None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
config = load_config()
|
config = load_config(config_path) if config_path is not None else load_config()
|
||||||
known = _known_preset_names()
|
known = _known_preset_names()
|
||||||
preset_rows = [
|
preset_rows = [
|
||||||
_preset_payload(preset, config.tools.mcp_servers)
|
_preset_payload(preset, config.tools.mcp_servers)
|
||||||
@@ -854,13 +960,51 @@ def mcp_presets_payload(
|
|||||||
for name, cfg in sorted(config.tools.mcp_servers.items())
|
for name, cfg in sorted(config.tools.mcp_servers.items())
|
||||||
if name not in known
|
if name not in known
|
||||||
]
|
]
|
||||||
|
existing_names = {str(row["name"]) for row in (*preset_rows, *custom_rows)}
|
||||||
|
plugin_rows = [
|
||||||
|
_agent_plugin_payload(plugin)
|
||||||
|
for plugin in discover_agent_plugins(config.workspace_path)
|
||||||
|
if f"plugin-{plugin.name}" not in existing_names
|
||||||
|
]
|
||||||
payload: dict[str, Any] = {
|
payload: dict[str, Any] = {
|
||||||
"presets": [*preset_rows, *custom_rows],
|
"presets": [*preset_rows, *custom_rows, *plugin_rows],
|
||||||
"installed_count": len(config.tools.mcp_servers),
|
"installed_count": len(config.tools.mcp_servers)
|
||||||
|
+ sum(int(row["enabled"]) for row in plugin_rows),
|
||||||
}
|
}
|
||||||
if last_action is not None:
|
if last_action is not None:
|
||||||
payload["last_action"] = last_action
|
payload["last_action"] = last_action
|
||||||
|
return attach_mcp_runtime_status(payload, runtime_status)
|
||||||
|
|
||||||
|
|
||||||
|
def attach_mcp_runtime_status(
|
||||||
|
payload: dict[str, Any],
|
||||||
|
runtime_status: Mapping[str, str] | None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Project safe, connection-attempt state onto configured MCP rows."""
|
||||||
|
if runtime_status is None:
|
||||||
return payload
|
return payload
|
||||||
|
projected = dict(payload)
|
||||||
|
raw_rows: object = payload.get("presets", [])
|
||||||
|
preset_rows = cast(list[object], raw_rows) if isinstance(raw_rows, list) else []
|
||||||
|
rows: list[Any] = []
|
||||||
|
for raw_row in preset_rows:
|
||||||
|
if not isinstance(raw_row, dict):
|
||||||
|
rows.append(raw_row)
|
||||||
|
continue
|
||||||
|
row = dict(cast(dict[str, Any], raw_row))
|
||||||
|
name = row.get("name")
|
||||||
|
status = runtime_status.get(name) if isinstance(name, str) else None
|
||||||
|
if (
|
||||||
|
status in _MCP_RUNTIME_STATUSES
|
||||||
|
and row.get("installed") is True
|
||||||
|
and row.get("configured") is True
|
||||||
|
):
|
||||||
|
row["runtime_status"] = status
|
||||||
|
else:
|
||||||
|
row.pop("runtime_status", None)
|
||||||
|
rows.append(row)
|
||||||
|
projected["presets"] = rows
|
||||||
|
return projected
|
||||||
|
|
||||||
|
|
||||||
def _display_name_for(name: str, preset: McpPreset | None = None) -> str:
|
def _display_name_for(name: str, preset: McpPreset | None = None) -> str:
|
||||||
@@ -893,6 +1037,7 @@ def _server_action_message(action: str, name: str, *, ok: bool = True) -> dict[s
|
|||||||
"import-cursor": "Imported",
|
"import-cursor": "Imported",
|
||||||
"tools": "Updated tools for",
|
"tools": "Updated tools for",
|
||||||
"remove": "Removed",
|
"remove": "Removed",
|
||||||
|
"reconnect": "Retried connection for",
|
||||||
}.get(action, "Updated")
|
}.get(action, "Updated")
|
||||||
payload: dict[str, Any] = {
|
payload: dict[str, Any] = {
|
||||||
"ok": ok,
|
"ok": ok,
|
||||||
@@ -907,6 +1052,24 @@ def _server_action_message(action: str, name: str, *, ok: bool = True) -> dict[s
|
|||||||
return payload
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
def mcp_reconnect_action(
|
||||||
|
query: QueryParams,
|
||||||
|
*,
|
||||||
|
config_path: Path | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Validate a configured server before asking the live runtime to retry it."""
|
||||||
|
name = _validated_server_name((_query_first(query, "name") or "").strip())
|
||||||
|
config = load_config(config_path) if config_path is not None else load_config()
|
||||||
|
if name not in config.tools.mcp_servers:
|
||||||
|
raise McpPresetError("unknown MCP server", status=404)
|
||||||
|
payload = mcp_presets_payload(
|
||||||
|
last_action=_server_action_message("reconnect", name),
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
|
payload["requires_restart"] = True
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
def _scrub_test_error(text: str) -> str:
|
def _scrub_test_error(text: str) -> str:
|
||||||
scrubbed = _SECRET_QUERY_RE.sub(r"\1<redacted>", text.strip())
|
scrubbed = _SECRET_QUERY_RE.sub(r"\1<redacted>", text.strip())
|
||||||
scrubbed = _SECRET_ASSIGNMENT_RE.sub(r"\1<redacted>", scrubbed)
|
scrubbed = _SECRET_ASSIGNMENT_RE.sub(r"\1<redacted>", scrubbed)
|
||||||
@@ -928,8 +1091,12 @@ async def _close_mcp_stacks(stacks: Mapping[str, Any]) -> None:
|
|||||||
await stack.aclose()
|
await stack.aclose()
|
||||||
|
|
||||||
|
|
||||||
async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
|
async def mcp_presets_test_action(
|
||||||
"""Connect to an enabled MCP preset and report its tool surface."""
|
query: QueryParams,
|
||||||
|
*,
|
||||||
|
config_path: Path | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Connect to an enabled MCP preset and report its complete tool surface."""
|
||||||
from nanobot.agent.tools.mcp import connect_mcp_servers
|
from nanobot.agent.tools.mcp import connect_mcp_servers
|
||||||
|
|
||||||
name = (_query_first(query, "name") or "").strip()
|
name = (_query_first(query, "name") or "").strip()
|
||||||
@@ -941,16 +1108,22 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
|
|||||||
display_name = _display_name_for(name, preset)
|
display_name = _display_name_for(name, preset)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
config = resolve_config_env_vars(load_config())
|
config = resolve_config_env_vars(
|
||||||
|
load_config(config_path) if config_path is not None else load_config(),
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
return mcp_presets_payload(last_action={
|
return mcp_presets_payload(
|
||||||
|
last_action={
|
||||||
"ok": False,
|
"ok": False,
|
||||||
"message": _scrub_test_error(str(exc)),
|
"message": _scrub_test_error(str(exc)),
|
||||||
"error": _scrub_test_error(str(exc)),
|
"error": _scrub_test_error(str(exc)),
|
||||||
"tool_count": 0,
|
"tool_count": 0,
|
||||||
"tool_names": [],
|
"tool_names": [],
|
||||||
"checked_at": _checked_at(),
|
"checked_at": _checked_at(),
|
||||||
})
|
},
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
|
|
||||||
cfg = config.tools.mcp_servers.get(name)
|
cfg = config.tools.mcp_servers.get(name)
|
||||||
if cfg is None:
|
if cfg is None:
|
||||||
@@ -968,7 +1141,7 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
|
|||||||
"tool_names": [],
|
"tool_names": [],
|
||||||
"checked_at": _checked_at(),
|
"checked_at": _checked_at(),
|
||||||
}
|
}
|
||||||
return mcp_presets_payload(last_action=last_action)
|
return mcp_presets_payload(last_action=last_action, config_path=config_path)
|
||||||
|
|
||||||
if cfg.command and not _command_available(cfg.command):
|
if cfg.command and not _command_available(cfg.command):
|
||||||
last_action = {
|
last_action = {
|
||||||
@@ -979,13 +1152,14 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
|
|||||||
"tool_names": [],
|
"tool_names": [],
|
||||||
"checked_at": _checked_at(),
|
"checked_at": _checked_at(),
|
||||||
}
|
}
|
||||||
return mcp_presets_payload(last_action=last_action)
|
return mcp_presets_payload(last_action=last_action, config_path=config_path)
|
||||||
|
|
||||||
registry = ToolRegistry()
|
registry = ToolRegistry()
|
||||||
stacks: dict[str, Any] = {}
|
stacks: dict[str, Any] = {}
|
||||||
|
inspection_cfg = cfg.model_copy(update={"enabled_tools": ["*"]})
|
||||||
try:
|
try:
|
||||||
stacks = await asyncio.wait_for(
|
stacks = await asyncio.wait_for(
|
||||||
connect_mcp_servers({name: cfg}, registry),
|
connect_mcp_servers({name: inspection_cfg}, registry),
|
||||||
timeout=_test_timeout(cfg),
|
timeout=_test_timeout(cfg),
|
||||||
)
|
)
|
||||||
tool_prefix = f"mcp_{name}_"
|
tool_prefix = f"mcp_{name}_"
|
||||||
@@ -1004,7 +1178,7 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
|
|||||||
else f"{display_name} connected, but reported no tools."
|
else f"{display_name} connected, but reported no tools."
|
||||||
),
|
),
|
||||||
"tool_count": len(tool_names),
|
"tool_count": len(tool_names),
|
||||||
"tool_names": tool_names[:_MAX_TEST_TOOLS],
|
"tool_names": tool_names,
|
||||||
"checked_at": _checked_at(),
|
"checked_at": _checked_at(),
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
@@ -1040,7 +1214,11 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
|
|||||||
|
|
||||||
tool_names = last_action.get("tool_names", [])
|
tool_names = last_action.get("tool_names", [])
|
||||||
preview = {name: tool_names} if tool_names else None
|
preview = {name: tool_names} if tool_names else None
|
||||||
return mcp_presets_payload(last_action=last_action, tool_preview=preview)
|
return mcp_presets_payload(
|
||||||
|
last_action=last_action,
|
||||||
|
tool_preview=preview,
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _parse_json_value(raw: str | None, *, fallback: Any) -> Any:
|
def _parse_json_value(raw: str | None, *, fallback: Any) -> Any:
|
||||||
@@ -1109,6 +1287,32 @@ def _normalize_transport(value: str | None, *, command: str = "", url: str = "")
|
|||||||
return normalized # type: ignore[return-value]
|
return normalized # type: ignore[return-value]
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_auth(
|
||||||
|
value: object,
|
||||||
|
*,
|
||||||
|
transport: Literal["stdio", "sse", "streamableHttp"],
|
||||||
|
url: str,
|
||||||
|
headers: Mapping[str, str],
|
||||||
|
) -> Literal["oauth"] | None:
|
||||||
|
raw = str(value or "").strip().lower()
|
||||||
|
if not raw and url and not headers:
|
||||||
|
normalized_url = url.rstrip("/")
|
||||||
|
if any(
|
||||||
|
preset.server is not None
|
||||||
|
and preset.server.auth == "oauth"
|
||||||
|
and preset.server.url.rstrip("/") == normalized_url
|
||||||
|
for preset in MCP_PRESETS
|
||||||
|
):
|
||||||
|
raw = "oauth"
|
||||||
|
if not raw:
|
||||||
|
return None
|
||||||
|
if raw != "oauth":
|
||||||
|
raise McpPresetError("unsupported MCP auth type")
|
||||||
|
if transport == "stdio":
|
||||||
|
raise McpPresetError("MCP OAuth requires a remote HTTP transport")
|
||||||
|
return "oauth"
|
||||||
|
|
||||||
|
|
||||||
def _validated_server_name(name: str) -> str:
|
def _validated_server_name(name: str) -> str:
|
||||||
if not name or _MCP_PRESET_NAME_RE.match(name) is None:
|
if not name or _MCP_PRESET_NAME_RE.match(name) is None:
|
||||||
raise McpPresetError("invalid MCP server name")
|
raise McpPresetError("invalid MCP server name")
|
||||||
@@ -1124,6 +1328,13 @@ def _custom_server_from_query(query: QueryParams) -> tuple[str, MCPServerConfig]
|
|||||||
raise McpPresetError("stdio MCP servers require a command")
|
raise McpPresetError("stdio MCP servers require a command")
|
||||||
if transport in {"sse", "streamableHttp"} and not url:
|
if transport in {"sse", "streamableHttp"} and not url:
|
||||||
raise McpPresetError("remote MCP servers require a URL")
|
raise McpPresetError("remote MCP servers require a URL")
|
||||||
|
headers = _parse_string_map(_query_first(query, "headers"))
|
||||||
|
auth = _normalize_auth(
|
||||||
|
_query_first(query, "auth"),
|
||||||
|
transport=transport,
|
||||||
|
url=url,
|
||||||
|
headers=headers,
|
||||||
|
)
|
||||||
raw_timeout = (_query_first(query, "tool_timeout") or "").strip()
|
raw_timeout = (_query_first(query, "tool_timeout") or "").strip()
|
||||||
tool_timeout = _DEFAULT_CUSTOM_TIMEOUT
|
tool_timeout = _DEFAULT_CUSTOM_TIMEOUT
|
||||||
if raw_timeout:
|
if raw_timeout:
|
||||||
@@ -1133,12 +1344,13 @@ def _custom_server_from_query(query: QueryParams) -> tuple[str, MCPServerConfig]
|
|||||||
raise McpPresetError("tool_timeout must be an integer") from exc
|
raise McpPresetError("tool_timeout must be an integer") from exc
|
||||||
cfg = MCPServerConfig(
|
cfg = MCPServerConfig(
|
||||||
type=transport,
|
type=transport,
|
||||||
|
auth=auth,
|
||||||
command=command if transport == "stdio" else "",
|
command=command if transport == "stdio" else "",
|
||||||
args=_parse_string_list(_query_first(query, "args")),
|
args=_parse_string_list(_query_first(query, "args")),
|
||||||
env=_parse_string_map(_query_first(query, "env")),
|
env=_parse_string_map(_query_first(query, "env")),
|
||||||
cwd=(_query_first(query, "cwd") or "").strip() if transport == "stdio" else "",
|
cwd=(_query_first(query, "cwd") or "").strip() if transport == "stdio" else "",
|
||||||
url=url if transport in {"sse", "streamableHttp"} else "",
|
url=url if transport in {"sse", "streamableHttp"} else "",
|
||||||
headers=_parse_string_map(_query_first(query, "headers")),
|
headers=headers,
|
||||||
tool_timeout=tool_timeout,
|
tool_timeout=tool_timeout,
|
||||||
enabled_tools=_parse_enabled_tools(_query_first(query, "enabled_tools")),
|
enabled_tools=_parse_enabled_tools(_query_first(query, "enabled_tools")),
|
||||||
)
|
)
|
||||||
@@ -1183,6 +1395,13 @@ def _mcp_server_config(name: str, raw: Any) -> tuple[str, MCPServerConfig]:
|
|||||||
headers = cast(dict[object, object], headers_value)
|
headers = cast(dict[object, object], headers_value)
|
||||||
if not all(isinstance(k, str) and isinstance(v, str) for k, v in headers.items()):
|
if not all(isinstance(k, str) and isinstance(v, str) for k, v in headers.items()):
|
||||||
raise McpPresetError(f"MCP server '{server_name}' headers must be a string object")
|
raise McpPresetError(f"MCP server '{server_name}' headers must be a string object")
|
||||||
|
typed_headers = cast(dict[str, str], headers)
|
||||||
|
auth = _normalize_auth(
|
||||||
|
server.get("auth"),
|
||||||
|
transport=transport,
|
||||||
|
url=url,
|
||||||
|
headers=typed_headers,
|
||||||
|
)
|
||||||
if not isinstance(enabled_tools_value, list):
|
if not isinstance(enabled_tools_value, list):
|
||||||
enabled_tools_value = ["*"]
|
enabled_tools_value = ["*"]
|
||||||
else:
|
else:
|
||||||
@@ -1191,12 +1410,13 @@ def _mcp_server_config(name: str, raw: Any) -> tuple[str, MCPServerConfig]:
|
|||||||
enabled_tools_value = ["*"]
|
enabled_tools_value = ["*"]
|
||||||
return server_name, MCPServerConfig(
|
return server_name, MCPServerConfig(
|
||||||
type=transport,
|
type=transport,
|
||||||
|
auth=auth,
|
||||||
command=command if transport == "stdio" else "",
|
command=command if transport == "stdio" else "",
|
||||||
args=cast(list[str], args),
|
args=cast(list[str], args),
|
||||||
env=cast(dict[str, str], env),
|
env=cast(dict[str, str], env),
|
||||||
cwd=cwd if transport == "stdio" else "",
|
cwd=cwd if transport == "stdio" else "",
|
||||||
url=url if transport in {"sse", "streamableHttp"} else "",
|
url=url if transport in {"sse", "streamableHttp"} else "",
|
||||||
headers=cast(dict[str, str], headers),
|
headers=typed_headers,
|
||||||
tool_timeout=timeout_int,
|
tool_timeout=timeout_int,
|
||||||
enabled_tools=cast(list[str], enabled_tools_value),
|
enabled_tools=cast(list[str], enabled_tools_value),
|
||||||
)
|
)
|
||||||
@@ -1221,24 +1441,54 @@ def _import_mcp_servers(raw_json: str | None) -> dict[str, MCPServerConfig]:
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
def custom_mcp_action(action: str, query: QueryParams) -> dict[str, Any]:
|
def _oauth_credentials_replaced(
|
||||||
config = load_config()
|
previous: MCPServerConfig | None,
|
||||||
|
replacement: MCPServerConfig,
|
||||||
|
) -> bool:
|
||||||
|
if previous is None or previous.auth != "oauth":
|
||||||
|
return False
|
||||||
|
return replacement.auth != "oauth" or replacement.url != previous.url
|
||||||
|
|
||||||
|
|
||||||
|
def custom_mcp_action(
|
||||||
|
action: str,
|
||||||
|
query: QueryParams,
|
||||||
|
*,
|
||||||
|
config_path: Path | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
config = load_config(config_path) if config_path is not None else load_config()
|
||||||
if action == "custom":
|
if action == "custom":
|
||||||
name, cfg = _custom_server_from_query(query)
|
name, cfg = _custom_server_from_query(query)
|
||||||
|
delete_credentials = _oauth_credentials_replaced(config.tools.mcp_servers.get(name), cfg)
|
||||||
config.tools.mcp_servers[name] = cfg
|
config.tools.mcp_servers[name] = cfg
|
||||||
save_config(config)
|
save_config(config, config_path)
|
||||||
payload = mcp_presets_payload(last_action=_server_action_message(action, name))
|
if delete_credentials:
|
||||||
|
delete_mcp_oauth_credentials(name)
|
||||||
|
payload = mcp_presets_payload(
|
||||||
|
last_action=_server_action_message(action, name),
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
payload["requires_restart"] = True
|
payload["requires_restart"] = True
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
if action in {"import", "import-cursor"}:
|
if action in {"import", "import-cursor"}:
|
||||||
servers = _import_mcp_servers(_query_first(query, "config"))
|
servers = _import_mcp_servers(_query_first(query, "config"))
|
||||||
|
delete_credentials = [
|
||||||
|
name
|
||||||
|
for name, cfg in servers.items()
|
||||||
|
if _oauth_credentials_replaced(config.tools.mcp_servers.get(name), cfg)
|
||||||
|
]
|
||||||
config.tools.mcp_servers.update(servers)
|
config.tools.mcp_servers.update(servers)
|
||||||
save_config(config)
|
save_config(config, config_path)
|
||||||
payload = mcp_presets_payload(last_action={
|
for name in delete_credentials:
|
||||||
|
delete_mcp_oauth_credentials(name)
|
||||||
|
payload = mcp_presets_payload(
|
||||||
|
last_action={
|
||||||
"ok": True,
|
"ok": True,
|
||||||
"message": f"Imported {len(servers)} MCP server(s).",
|
"message": f"Imported {len(servers)} MCP server(s).",
|
||||||
})
|
},
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
payload["requires_restart"] = True
|
payload["requires_restart"] = True
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
@@ -1249,29 +1499,61 @@ def custom_mcp_action(action: str, query: QueryParams) -> dict[str, Any]:
|
|||||||
raise McpPresetError("unknown MCP server", status=404)
|
raise McpPresetError("unknown MCP server", status=404)
|
||||||
cfg.enabled_tools = _parse_enabled_tools(_query_first(query, "enabled_tools"))
|
cfg.enabled_tools = _parse_enabled_tools(_query_first(query, "enabled_tools"))
|
||||||
config.tools.mcp_servers[name] = cfg
|
config.tools.mcp_servers[name] = cfg
|
||||||
save_config(config)
|
save_config(config, config_path)
|
||||||
payload = mcp_presets_payload(last_action=_server_action_message(action, name))
|
payload = mcp_presets_payload(
|
||||||
|
last_action=_server_action_message(action, name),
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
payload["requires_restart"] = True
|
payload["requires_restart"] = True
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
raise McpPresetError(f"unknown MCP action '{action}'", status=404)
|
raise McpPresetError(f"unknown MCP action '{action}'", status=404)
|
||||||
|
|
||||||
|
|
||||||
def mcp_presets_action(action: str, query: QueryParams) -> dict[str, Any]:
|
def ensure_mcp_oauth_server(
|
||||||
|
query: QueryParams,
|
||||||
|
*,
|
||||||
|
config_path: Path | None = None,
|
||||||
|
) -> tuple[str, MCPServerConfig]:
|
||||||
|
"""Materialize an OAuth preset on first click and return its saved config."""
|
||||||
|
name = _validated_server_name((_query_first(query, "name") or "").strip())
|
||||||
|
config = load_config(config_path) if config_path is not None else load_config()
|
||||||
|
cfg = config.tools.mcp_servers.get(name)
|
||||||
|
if cfg is None:
|
||||||
|
preset = _preset_by_name(name)
|
||||||
|
if preset.server is None or preset.server.auth != "oauth":
|
||||||
|
raise McpPresetError("MCP server does not support browser authorization", status=409)
|
||||||
|
cfg = _materialize_server(preset, query, None)
|
||||||
|
config.tools.mcp_servers[name] = cfg
|
||||||
|
save_config(config, config_path)
|
||||||
|
if cfg.auth != "oauth" or cfg.type not in {"sse", "streamableHttp"} or not cfg.url:
|
||||||
|
raise McpPresetError("MCP server is not configured for OAuth", status=409)
|
||||||
|
return name, cfg
|
||||||
|
|
||||||
|
|
||||||
|
def mcp_presets_action(
|
||||||
|
action: str,
|
||||||
|
query: QueryParams,
|
||||||
|
*,
|
||||||
|
config_path: Path | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
name = (_query_first(query, "name") or "").strip()
|
name = (_query_first(query, "name") or "").strip()
|
||||||
if not name:
|
if not name:
|
||||||
raise McpPresetError("missing MCP preset name")
|
raise McpPresetError("missing MCP preset name")
|
||||||
preset = _preset_by_name_optional(name)
|
preset = _preset_by_name_optional(name)
|
||||||
|
|
||||||
config = load_config()
|
config = load_config(config_path) if config_path is not None else load_config()
|
||||||
existing = config.tools.mcp_servers.get(name)
|
existing = config.tools.mcp_servers.get(name)
|
||||||
|
|
||||||
if action == "enable":
|
if action == "enable":
|
||||||
if preset is None:
|
if preset is None:
|
||||||
raise McpPresetError("unknown MCP preset", status=404)
|
raise McpPresetError("unknown MCP preset", status=404)
|
||||||
config.tools.mcp_servers[preset.name] = _materialize_server(preset, query, existing)
|
config.tools.mcp_servers[preset.name] = _materialize_server(preset, query, existing)
|
||||||
save_config(config)
|
save_config(config, config_path)
|
||||||
payload = mcp_presets_payload(last_action=_action_message(action, preset))
|
payload = mcp_presets_payload(
|
||||||
|
last_action=_action_message(action, preset),
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
payload["requires_restart"] = True
|
payload["requires_restart"] = True
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
@@ -1287,7 +1569,8 @@ def mcp_presets_action(action: str, query: QueryParams) -> dict[str, Any]:
|
|||||||
except OSError as exc:
|
except OSError as exc:
|
||||||
cleanup_error = str(exc)
|
cleanup_error = str(exc)
|
||||||
del config.tools.mcp_servers[name]
|
del config.tools.mcp_servers[name]
|
||||||
save_config(config)
|
save_config(config, config_path)
|
||||||
|
delete_mcp_oauth_credentials(name)
|
||||||
last_action = (
|
last_action = (
|
||||||
_action_message(action, preset)
|
_action_message(action, preset)
|
||||||
if preset is not None
|
if preset is not None
|
||||||
@@ -1303,7 +1586,10 @@ def mcp_presets_action(action: str, query: QueryParams) -> dict[str, Any]:
|
|||||||
f"{last_action['message']} Could not remove managed runtime files: {cleanup_error}"
|
f"{last_action['message']} Could not remove managed runtime files: {cleanup_error}"
|
||||||
)
|
)
|
||||||
last_action["verification_failed"] = ["managed_paths_absent"]
|
last_action["verification_failed"] = ["managed_paths_absent"]
|
||||||
payload = mcp_presets_payload(last_action=last_action)
|
payload = mcp_presets_payload(
|
||||||
|
last_action=last_action,
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
payload["requires_restart"] = True
|
payload["requires_restart"] = True
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
@@ -1339,16 +1625,64 @@ async def mcp_presets_settings_action(
|
|||||||
query: QueryParams,
|
query: QueryParams,
|
||||||
*,
|
*,
|
||||||
reload_mcp: McpReload | None = None,
|
reload_mcp: McpReload | None = None,
|
||||||
|
mcp_runtime_status: McpRuntimeStatus | None = None,
|
||||||
|
config: WebUISettingsConfig | None = None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""Run a WebUI MCP preset action and hot-reload the agent when config changes."""
|
"""Run a WebUI MCP preset action and hot-reload the agent when config changes."""
|
||||||
|
config_path = config.path if config is not None else None
|
||||||
if action is None:
|
if action is None:
|
||||||
return mcp_presets_payload()
|
return mcp_presets_payload(
|
||||||
|
runtime_status=mcp_runtime_status() if mcp_runtime_status is not None else None,
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
|
name = (_query_first(query, "name") or "").strip()
|
||||||
|
if name.startswith("plugin-"):
|
||||||
|
plugin_config = load_config(config_path) if config_path is not None else load_config()
|
||||||
|
plugin_name = name.removeprefix("plugin-")
|
||||||
|
plugins = discover_agent_plugins(plugin_config.workspace_path)
|
||||||
|
plugin = next((item for item in plugins if item.name == plugin_name), None)
|
||||||
|
if name not in plugin_config.tools.mcp_servers and plugin is not None:
|
||||||
|
if action not in {"enable", "disable"}:
|
||||||
|
raise McpPresetError("Agent Plugins support enable and disable actions only")
|
||||||
|
await asyncio.to_thread(
|
||||||
|
set_agent_plugin_enabled,
|
||||||
|
plugin_config.workspace_path,
|
||||||
|
plugin_name,
|
||||||
|
action == "enable",
|
||||||
|
)
|
||||||
|
verb = "enabled" if action == "enable" else "disabled"
|
||||||
|
payload = mcp_presets_payload(
|
||||||
|
last_action={"ok": True, "message": f"{plugin.display_name} {verb}."},
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
|
if reload_mcp is not None:
|
||||||
|
payload = attach_mcp_hot_reload_result(payload, await reload_mcp())
|
||||||
|
return payload
|
||||||
if action == "test":
|
if action == "test":
|
||||||
return await mcp_presets_test_action(query)
|
payload = await mcp_presets_test_action(query, config_path=config_path)
|
||||||
if action in _CUSTOM_ACTIONS:
|
return attach_mcp_runtime_status(
|
||||||
|
payload,
|
||||||
|
mcp_runtime_status() if mcp_runtime_status is not None else None,
|
||||||
|
)
|
||||||
|
if action == "reconnect":
|
||||||
|
payload = await asyncio.to_thread(
|
||||||
|
mcp_reconnect_action,
|
||||||
|
query,
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
|
elif config is not None:
|
||||||
|
operation = custom_mcp_action if action in _CUSTOM_ACTIONS else mcp_presets_action
|
||||||
|
payload = await asyncio.to_thread(
|
||||||
|
config.run_serialized,
|
||||||
|
lambda path: operation(action, query, config_path=path),
|
||||||
|
)
|
||||||
|
elif action in _CUSTOM_ACTIONS:
|
||||||
payload = await asyncio.to_thread(custom_mcp_action, action, query)
|
payload = await asyncio.to_thread(custom_mcp_action, action, query)
|
||||||
else:
|
else:
|
||||||
payload = await asyncio.to_thread(mcp_presets_action, action, query)
|
payload = await asyncio.to_thread(mcp_presets_action, action, query)
|
||||||
if reload_mcp is not None:
|
if reload_mcp is not None:
|
||||||
payload = attach_mcp_hot_reload_result(payload, await reload_mcp())
|
payload = attach_mcp_hot_reload_result(payload, await reload_mcp())
|
||||||
return payload
|
return attach_mcp_runtime_status(
|
||||||
|
payload,
|
||||||
|
mcp_runtime_status() if mcp_runtime_status is not None else None,
|
||||||
|
)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Nanobot optional feature helpers for WebUI Settings."""
|
"""Nanobot optional feature helpers for WebUI Settings."""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from nanobot.channels.registry import load_channel_plugin
|
from nanobot.channels.registry import load_channel_plugin
|
||||||
@@ -15,9 +16,14 @@ from nanobot.webui.http_utils import query_first
|
|||||||
QueryParams = dict[str, list[str]]
|
QueryParams = dict[str, list[str]]
|
||||||
|
|
||||||
|
|
||||||
def nanobot_features_payload() -> dict[str, Any]:
|
def nanobot_features_payload(*, config_path: Path | None = None) -> dict[str, Any]:
|
||||||
|
if config_path is None:
|
||||||
return optional_features_payload()
|
return optional_features_payload()
|
||||||
|
|
||||||
|
from nanobot.config.loader import load_config
|
||||||
|
|
||||||
|
return optional_features_payload(config=load_config(config_path))
|
||||||
|
|
||||||
|
|
||||||
def nanobot_feature_instance_target(query: QueryParams) -> str | None:
|
def nanobot_feature_instance_target(query: QueryParams) -> str | None:
|
||||||
"""Preserve the difference between a global action and an explicit instance."""
|
"""Preserve the difference between a global action and an explicit instance."""
|
||||||
@@ -32,13 +38,19 @@ def nanobot_features_action(
|
|||||||
query: QueryParams,
|
query: QueryParams,
|
||||||
*,
|
*,
|
||||||
allow_install: bool = True,
|
allow_install: bool = True,
|
||||||
|
config_path: Path | None = None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
name = (query_first(query, "name") or "").strip()
|
name = (query_first(query, "name") or "").strip()
|
||||||
instance_id = nanobot_feature_instance_target(query)
|
instance_id = nanobot_feature_instance_target(query)
|
||||||
if not name:
|
if not name:
|
||||||
raise OptionalFeatureError("missing feature name")
|
raise OptionalFeatureError("missing feature name")
|
||||||
if action == "enable":
|
if action == "enable":
|
||||||
return enable_optional_feature(name, allow_install=allow_install, instance_id=instance_id)
|
return enable_optional_feature(
|
||||||
|
name,
|
||||||
|
config_path=config_path,
|
||||||
|
allow_install=allow_install,
|
||||||
|
instance_id=instance_id,
|
||||||
|
)
|
||||||
if action == "disable":
|
if action == "disable":
|
||||||
try:
|
try:
|
||||||
plugin = load_channel_plugin(name)
|
plugin = load_channel_plugin(name)
|
||||||
@@ -50,5 +62,9 @@ def nanobot_features_action(
|
|||||||
f"Use `nanobot plugins disable {name}` from a terminal if you need to disable it.",
|
f"Use `nanobot plugins disable {name}` from a terminal if you need to disable it.",
|
||||||
status=400,
|
status=400,
|
||||||
)
|
)
|
||||||
return disable_optional_feature(name, instance_id=instance_id)
|
return disable_optional_feature(
|
||||||
|
name,
|
||||||
|
config_path=config_path,
|
||||||
|
instance_id=instance_id,
|
||||||
|
)
|
||||||
raise OptionalFeatureError(f"unknown feature action '{action}'", status=404)
|
raise OptionalFeatureError(f"unknown feature action '{action}'", status=404)
|
||||||
|
|||||||
+241
-2150
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,804 @@
|
|||||||
|
"""Capability settings domain logic for Web, media, network, and API features."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING, Any, Awaitable, Callable, TypedDict
|
||||||
|
|
||||||
|
from nanobot.agent.tools.web import SEARCH_PROVIDER_OPTIONS
|
||||||
|
from nanobot.api.runtime import ApiRuntime, ApiStartOptions
|
||||||
|
from nanobot.audio.transcription import resolve_transcription_config
|
||||||
|
from nanobot.audio.transcription_registry import (
|
||||||
|
resolve_transcription_provider,
|
||||||
|
transcription_provider_names,
|
||||||
|
)
|
||||||
|
from nanobot.config.schema import Config
|
||||||
|
from nanobot.optional_features import (
|
||||||
|
OptionalFeatureError,
|
||||||
|
extra_installed,
|
||||||
|
optional_dependency_groups,
|
||||||
|
)
|
||||||
|
from nanobot.providers.image_generation import (
|
||||||
|
get_image_gen_provider,
|
||||||
|
image_gen_provider_names,
|
||||||
|
)
|
||||||
|
from nanobot.providers.registry import find_by_name
|
||||||
|
from nanobot.security.network import is_loopback_host
|
||||||
|
from nanobot.webui.settings_contracts import (
|
||||||
|
QueryParams,
|
||||||
|
SettingsRequest,
|
||||||
|
SettingsRouteResult,
|
||||||
|
WebUISettingsError,
|
||||||
|
parse_bool,
|
||||||
|
query_first,
|
||||||
|
query_first_alias,
|
||||||
|
)
|
||||||
|
from nanobot.webui.settings_models import (
|
||||||
|
OAuthStatusReader,
|
||||||
|
mask_secret_hint,
|
||||||
|
provider_configured_for_settings,
|
||||||
|
)
|
||||||
|
from nanobot.webui.workspaces import (
|
||||||
|
read_webui_default_access_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from nanobot.webui.settings_services import WebUISettingsServices
|
||||||
|
|
||||||
|
SettingsOperation = Callable[..., dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class CapabilitySettingsOperations:
|
||||||
|
update_web_search: SettingsOperation
|
||||||
|
update_api: SettingsOperation
|
||||||
|
update_image: SettingsOperation
|
||||||
|
update_transcription: SettingsOperation
|
||||||
|
update_network: SettingsOperation
|
||||||
|
nanobot_features_action: SettingsOperation
|
||||||
|
api_runtime: Callable[[], ApiRuntime]
|
||||||
|
reload_image: Callable[[], Awaitable[dict[str, Any]]]
|
||||||
|
|
||||||
|
|
||||||
|
class CapabilitySettingsPayload(TypedDict):
|
||||||
|
web_search: dict[str, Any]
|
||||||
|
web: dict[str, Any]
|
||||||
|
api: dict[str, Any]
|
||||||
|
observability: dict[str, Any]
|
||||||
|
image_generation: dict[str, Any]
|
||||||
|
transcription: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
_WEB_SEARCH_PROVIDER_OPTIONS = SEARCH_PROVIDER_OPTIONS
|
||||||
|
_WEB_SEARCH_PROVIDER_BY_NAME = {
|
||||||
|
provider["name"]: provider for provider in _WEB_SEARCH_PROVIDER_OPTIONS
|
||||||
|
}
|
||||||
|
_IMAGE_GENERATION_ASPECT_RATIOS = {
|
||||||
|
"1:1",
|
||||||
|
"3:4",
|
||||||
|
"9:16",
|
||||||
|
"4:3",
|
||||||
|
"16:9",
|
||||||
|
"3:2",
|
||||||
|
"2:3",
|
||||||
|
"21:9",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _image_generation_provider_rows(
|
||||||
|
config: Config,
|
||||||
|
*,
|
||||||
|
oauth_status: OAuthStatusReader,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
rows: list[dict[str, Any]] = []
|
||||||
|
for name in image_gen_provider_names():
|
||||||
|
image_provider = get_image_gen_provider(name)
|
||||||
|
spec = find_by_name(name)
|
||||||
|
provider_config = getattr(config.providers, name, None)
|
||||||
|
configured = (
|
||||||
|
provider_configured_for_settings(spec, provider_config, oauth_status)
|
||||||
|
if spec is not None and provider_config is not None
|
||||||
|
else bool(getattr(provider_config, "api_key", None))
|
||||||
|
)
|
||||||
|
rows.append(
|
||||||
|
{
|
||||||
|
"name": name,
|
||||||
|
"label": spec.label if spec is not None else name,
|
||||||
|
"configured": configured,
|
||||||
|
"auth_type": "oauth" if spec is not None and spec.is_oauth else "api_key",
|
||||||
|
"api_key_hint": mask_secret_hint(getattr(provider_config, "api_key", None)),
|
||||||
|
"api_base": getattr(provider_config, "api_base", None),
|
||||||
|
"default_api_base": (
|
||||||
|
spec.default_api_base if spec and spec.default_api_base else None
|
||||||
|
),
|
||||||
|
"models": list(image_provider.model_options) if image_provider else [],
|
||||||
|
"default_model": (
|
||||||
|
image_provider.model_options[0]
|
||||||
|
if image_provider and image_provider.model_options
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return rows
|
||||||
|
|
||||||
|
|
||||||
|
def _transcription_provider_rows(config: Config) -> list[dict[str, Any]]:
|
||||||
|
rows: list[dict[str, Any]] = []
|
||||||
|
for name in transcription_provider_names():
|
||||||
|
spec = find_by_name(name)
|
||||||
|
provider_config = getattr(config.providers, name, None)
|
||||||
|
rows.append(
|
||||||
|
{
|
||||||
|
"name": name,
|
||||||
|
"label": spec.label if spec is not None else name,
|
||||||
|
"configured": bool(getattr(provider_config, "api_key", None)),
|
||||||
|
"api_key_hint": mask_secret_hint(getattr(provider_config, "api_key", None)),
|
||||||
|
"api_base": getattr(provider_config, "api_base", None),
|
||||||
|
"default_api_base": (
|
||||||
|
spec.default_api_base if spec and spec.default_api_base else None
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return rows
|
||||||
|
|
||||||
|
|
||||||
|
def capability_settings_payload(
|
||||||
|
config: Config,
|
||||||
|
*,
|
||||||
|
oauth_status: OAuthStatusReader,
|
||||||
|
) -> CapabilitySettingsPayload:
|
||||||
|
search_config = config.tools.web.search
|
||||||
|
image_config = config.tools.image_generation
|
||||||
|
transcription = resolve_transcription_config(config)
|
||||||
|
search_provider = (
|
||||||
|
search_config.provider
|
||||||
|
if search_config.provider in _WEB_SEARCH_PROVIDER_BY_NAME
|
||||||
|
else "duckduckgo"
|
||||||
|
)
|
||||||
|
image_providers = _image_generation_provider_rows(config, oauth_status=oauth_status)
|
||||||
|
selected_image_provider = next(
|
||||||
|
(
|
||||||
|
provider
|
||||||
|
for provider in image_providers
|
||||||
|
if provider["name"] == image_config.provider
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"web_search": {
|
||||||
|
"provider": search_provider,
|
||||||
|
"api_key_hint": mask_secret_hint(search_config.api_key),
|
||||||
|
"base_url": search_config.base_url or None,
|
||||||
|
"max_results": search_config.max_results,
|
||||||
|
"timeout": search_config.timeout,
|
||||||
|
"providers": list(_WEB_SEARCH_PROVIDER_OPTIONS),
|
||||||
|
},
|
||||||
|
"web": {
|
||||||
|
"enable": config.tools.web.enable,
|
||||||
|
"proxy": config.tools.web.proxy,
|
||||||
|
"user_agent": config.tools.web.user_agent,
|
||||||
|
"search": {
|
||||||
|
"max_results": search_config.max_results,
|
||||||
|
"timeout": search_config.timeout,
|
||||||
|
},
|
||||||
|
"fetch": {
|
||||||
|
"use_jina_reader": config.tools.web.fetch.use_jina_reader,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"api": {
|
||||||
|
"host": config.api.host,
|
||||||
|
"port": config.api.port,
|
||||||
|
"timeout": config.api.timeout,
|
||||||
|
"api_key_hint": mask_secret_hint(config.api.api_key),
|
||||||
|
},
|
||||||
|
"observability": {
|
||||||
|
"provider": "langfuse",
|
||||||
|
"configured": bool(
|
||||||
|
os.environ.get("LANGFUSE_SECRET_KEY")
|
||||||
|
and os.environ.get("LANGFUSE_PUBLIC_KEY")
|
||||||
|
),
|
||||||
|
"base_url": os.environ.get("LANGFUSE_BASE_URL")
|
||||||
|
or "https://cloud.langfuse.com",
|
||||||
|
},
|
||||||
|
"image_generation": {
|
||||||
|
"enabled": image_config.enabled,
|
||||||
|
"provider": image_config.provider,
|
||||||
|
"provider_configured": bool(
|
||||||
|
selected_image_provider and selected_image_provider["configured"]
|
||||||
|
),
|
||||||
|
"model": image_config.model,
|
||||||
|
"default_aspect_ratio": image_config.default_aspect_ratio,
|
||||||
|
"default_image_size": image_config.default_image_size,
|
||||||
|
"max_images_per_turn": image_config.max_images_per_turn,
|
||||||
|
"save_dir": image_config.save_dir,
|
||||||
|
"providers": image_providers,
|
||||||
|
},
|
||||||
|
"transcription": {
|
||||||
|
"enabled": transcription.enabled,
|
||||||
|
"provider": transcription.provider,
|
||||||
|
"provider_configured": transcription.configured,
|
||||||
|
"model": transcription.model,
|
||||||
|
"language": transcription.language,
|
||||||
|
"max_duration_sec": transcription.max_duration_sec,
|
||||||
|
"max_upload_mb": transcription.max_upload_mb,
|
||||||
|
"providers": _transcription_provider_rows(config),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def update_network_safety_settings(
|
||||||
|
config: Config,
|
||||||
|
query: QueryParams,
|
||||||
|
) -> tuple[bool, str | None]:
|
||||||
|
raw_allow = (
|
||||||
|
query_first_alias(
|
||||||
|
query,
|
||||||
|
"webui_allow_local_service_access",
|
||||||
|
"webuiAllowLocalServiceAccess",
|
||||||
|
)
|
||||||
|
or query_first_alias(
|
||||||
|
query,
|
||||||
|
"allow_local_preview_access",
|
||||||
|
"allowLocalPreviewAccess",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
raw_default_access_mode = query_first_alias(
|
||||||
|
query,
|
||||||
|
"webui_default_access_mode",
|
||||||
|
"webuiDefaultAccessMode",
|
||||||
|
)
|
||||||
|
if raw_allow is None and raw_default_access_mode is None:
|
||||||
|
raise WebUISettingsError(
|
||||||
|
"webui_allow_local_service_access or webui_default_access_mode is required"
|
||||||
|
)
|
||||||
|
|
||||||
|
changed = False
|
||||||
|
if raw_allow is not None:
|
||||||
|
allow_local = parse_bool(raw_allow, "webui_allow_local_service_access")
|
||||||
|
if config.tools.webui_allow_local_service_access != allow_local:
|
||||||
|
config.tools.webui_allow_local_service_access = allow_local
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
default_access_mode: str | None = None
|
||||||
|
if raw_default_access_mode is not None:
|
||||||
|
default_access_mode = raw_default_access_mode.strip().lower()
|
||||||
|
if default_access_mode == "restricted":
|
||||||
|
default_access_mode = "default"
|
||||||
|
if default_access_mode not in {"default", "full"}:
|
||||||
|
raise WebUISettingsError(
|
||||||
|
"webui_default_access_mode must be default or full"
|
||||||
|
)
|
||||||
|
return changed, default_access_mode
|
||||||
|
|
||||||
|
|
||||||
|
def update_web_search_settings(config: Config, query: QueryParams) -> tuple[bool, bool]:
|
||||||
|
provider_name = (query_first(query, "provider") or "").strip().lower()
|
||||||
|
provider_option = _WEB_SEARCH_PROVIDER_BY_NAME.get(provider_name)
|
||||||
|
if provider_option is None:
|
||||||
|
raise WebUISettingsError("unknown web search provider")
|
||||||
|
|
||||||
|
search_config = config.tools.web.search
|
||||||
|
web_config = config.tools.web
|
||||||
|
previous_provider = search_config.provider
|
||||||
|
changed = False
|
||||||
|
restart_required = False
|
||||||
|
|
||||||
|
def set_search_value(attr: str, value: object) -> None:
|
||||||
|
nonlocal changed
|
||||||
|
if getattr(search_config, attr) != value:
|
||||||
|
setattr(search_config, attr, value)
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
def set_fetch_value(attr: str, value: object) -> None:
|
||||||
|
nonlocal changed
|
||||||
|
if getattr(web_config.fetch, attr) != value:
|
||||||
|
setattr(web_config.fetch, attr, value)
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
if search_config.provider != provider_name:
|
||||||
|
search_config.provider = provider_name
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
credential = provider_option["credential"]
|
||||||
|
if credential == "none":
|
||||||
|
set_search_value("api_key", "")
|
||||||
|
set_search_value("base_url", "")
|
||||||
|
elif credential == "base_url":
|
||||||
|
base_url = query_first_alias(query, "base_url", "baseUrl")
|
||||||
|
base_url = base_url.strip() if base_url is not None else None
|
||||||
|
if not base_url and previous_provider == provider_name and search_config.base_url:
|
||||||
|
base_url = search_config.base_url
|
||||||
|
if not base_url:
|
||||||
|
raise WebUISettingsError("base_url is required")
|
||||||
|
set_search_value("base_url", base_url)
|
||||||
|
set_search_value("api_key", "")
|
||||||
|
elif credential in {"api_key", "optional_api_key"}:
|
||||||
|
raw_api_key = query_first_alias(query, "api_key", "apiKey")
|
||||||
|
api_key = raw_api_key.strip() if raw_api_key is not None else None
|
||||||
|
if api_key is None and previous_provider == provider_name and search_config.api_key:
|
||||||
|
api_key = search_config.api_key
|
||||||
|
if credential == "api_key" and not api_key:
|
||||||
|
raise WebUISettingsError("api_key is required")
|
||||||
|
set_search_value("api_key", api_key or "")
|
||||||
|
set_search_value("base_url", "")
|
||||||
|
else:
|
||||||
|
raise WebUISettingsError("unknown web search credential type")
|
||||||
|
|
||||||
|
max_results = query_first_alias(query, "max_results", "maxResults")
|
||||||
|
if max_results is not None:
|
||||||
|
try:
|
||||||
|
parsed = int(max_results)
|
||||||
|
except ValueError:
|
||||||
|
raise WebUISettingsError("max_results must be an integer") from None
|
||||||
|
if parsed < 1 or parsed > 10:
|
||||||
|
raise WebUISettingsError("max_results must be between 1 and 10")
|
||||||
|
set_search_value("max_results", parsed)
|
||||||
|
|
||||||
|
timeout = query_first(query, "timeout")
|
||||||
|
if timeout is not None:
|
||||||
|
try:
|
||||||
|
parsed_timeout = int(timeout)
|
||||||
|
except ValueError:
|
||||||
|
raise WebUISettingsError("timeout must be an integer") from None
|
||||||
|
if parsed_timeout < 1 or parsed_timeout > 120:
|
||||||
|
raise WebUISettingsError("timeout must be between 1 and 120")
|
||||||
|
set_search_value("timeout", parsed_timeout)
|
||||||
|
|
||||||
|
use_jina_reader = query_first_alias(query, "use_jina_reader", "useJinaReader")
|
||||||
|
if use_jina_reader is not None:
|
||||||
|
previous_jina_reader = web_config.fetch.use_jina_reader
|
||||||
|
set_fetch_value("use_jina_reader", parse_bool(use_jina_reader, "use_jina_reader"))
|
||||||
|
if web_config.fetch.use_jina_reader != previous_jina_reader:
|
||||||
|
restart_required = True
|
||||||
|
return changed, restart_required
|
||||||
|
|
||||||
|
|
||||||
|
def update_api_settings(config: Config, query: QueryParams) -> None:
|
||||||
|
"""Update the managed OpenAI-compatible API configuration."""
|
||||||
|
api = config.api
|
||||||
|
host = query_first(query, "host")
|
||||||
|
if host is not None:
|
||||||
|
host = host.strip()
|
||||||
|
if not host:
|
||||||
|
raise WebUISettingsError("host is required")
|
||||||
|
api.host = host
|
||||||
|
|
||||||
|
port = query_first(query, "port")
|
||||||
|
if port is not None:
|
||||||
|
try:
|
||||||
|
parsed_port = int(port)
|
||||||
|
except ValueError:
|
||||||
|
raise WebUISettingsError("port must be an integer") from None
|
||||||
|
if parsed_port < 1 or parsed_port > 65535:
|
||||||
|
raise WebUISettingsError("port must be between 1 and 65535")
|
||||||
|
api.port = parsed_port
|
||||||
|
|
||||||
|
timeout = query_first(query, "timeout")
|
||||||
|
if timeout is not None:
|
||||||
|
try:
|
||||||
|
parsed_timeout = float(timeout)
|
||||||
|
except ValueError:
|
||||||
|
raise WebUISettingsError("timeout must be a number") from None
|
||||||
|
if parsed_timeout < 1 or parsed_timeout > 3600:
|
||||||
|
raise WebUISettingsError("timeout must be between 1 and 3600")
|
||||||
|
api.timeout = parsed_timeout
|
||||||
|
|
||||||
|
api_key = query_first_alias(query, "api_key", "apiKey")
|
||||||
|
if api_key is not None:
|
||||||
|
api.api_key = api_key.strip()
|
||||||
|
if not is_loopback_host(api.host) and not api.api_key.strip():
|
||||||
|
raise WebUISettingsError(
|
||||||
|
"an API key is required when the API is available on the network"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def update_image_generation_settings(
|
||||||
|
config: Config,
|
||||||
|
query: QueryParams,
|
||||||
|
*,
|
||||||
|
oauth_status: OAuthStatusReader,
|
||||||
|
) -> bool:
|
||||||
|
image_config = config.tools.image_generation
|
||||||
|
changed = False
|
||||||
|
|
||||||
|
provider_name = query_first(query, "provider")
|
||||||
|
if provider_name is not None:
|
||||||
|
provider_name = provider_name.strip().lower()
|
||||||
|
if not provider_name:
|
||||||
|
raise WebUISettingsError("image generation provider is required")
|
||||||
|
if get_image_gen_provider(provider_name) is None:
|
||||||
|
raise WebUISettingsError("unknown image generation provider")
|
||||||
|
if image_config.provider != provider_name:
|
||||||
|
image_config.provider = provider_name
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
enabled = query_first(query, "enabled")
|
||||||
|
if enabled is not None:
|
||||||
|
parsed_enabled = parse_bool(enabled, "enabled")
|
||||||
|
if image_config.enabled != parsed_enabled:
|
||||||
|
image_config.enabled = parsed_enabled
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
model = query_first(query, "model")
|
||||||
|
if model is not None:
|
||||||
|
model = model.strip()
|
||||||
|
if not model:
|
||||||
|
raise WebUISettingsError("image generation model is required")
|
||||||
|
if len(model) > 200:
|
||||||
|
raise WebUISettingsError("image generation model is too long")
|
||||||
|
if image_config.model != model:
|
||||||
|
image_config.model = model
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
default_aspect_ratio = query_first_alias(
|
||||||
|
query,
|
||||||
|
"default_aspect_ratio",
|
||||||
|
"defaultAspectRatio",
|
||||||
|
)
|
||||||
|
if default_aspect_ratio is not None:
|
||||||
|
default_aspect_ratio = default_aspect_ratio.strip()
|
||||||
|
if default_aspect_ratio not in _IMAGE_GENERATION_ASPECT_RATIOS:
|
||||||
|
raise WebUISettingsError("unsupported image generation aspect ratio")
|
||||||
|
if image_config.default_aspect_ratio != default_aspect_ratio:
|
||||||
|
image_config.default_aspect_ratio = default_aspect_ratio
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
default_image_size = query_first_alias(
|
||||||
|
query,
|
||||||
|
"default_image_size",
|
||||||
|
"defaultImageSize",
|
||||||
|
)
|
||||||
|
if default_image_size is not None:
|
||||||
|
default_image_size = default_image_size.strip()
|
||||||
|
if not default_image_size:
|
||||||
|
raise WebUISettingsError("default image size is required")
|
||||||
|
if len(default_image_size) > 32 or not all(
|
||||||
|
char.isascii() and (char.isalnum() or char in {"x", "X", ":", "-", "_"})
|
||||||
|
for char in default_image_size
|
||||||
|
):
|
||||||
|
raise WebUISettingsError("unsupported image generation size")
|
||||||
|
if image_config.default_image_size != default_image_size:
|
||||||
|
image_config.default_image_size = default_image_size
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
max_images_per_turn = query_first_alias(
|
||||||
|
query,
|
||||||
|
"max_images_per_turn",
|
||||||
|
"maxImagesPerTurn",
|
||||||
|
)
|
||||||
|
if max_images_per_turn is not None:
|
||||||
|
try:
|
||||||
|
parsed_max = int(max_images_per_turn)
|
||||||
|
except ValueError:
|
||||||
|
raise WebUISettingsError("max_images_per_turn must be an integer") from None
|
||||||
|
if parsed_max < 1 or parsed_max > 8:
|
||||||
|
raise WebUISettingsError("max_images_per_turn must be between 1 and 8")
|
||||||
|
if image_config.max_images_per_turn != parsed_max:
|
||||||
|
image_config.max_images_per_turn = parsed_max
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
if image_config.enabled:
|
||||||
|
selected_provider = next(
|
||||||
|
(
|
||||||
|
provider
|
||||||
|
for provider in _image_generation_provider_rows(
|
||||||
|
config,
|
||||||
|
oauth_status=oauth_status,
|
||||||
|
)
|
||||||
|
if provider["name"] == image_config.provider
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if not selected_provider or not selected_provider["configured"]:
|
||||||
|
raise WebUISettingsError("image generation provider is not configured")
|
||||||
|
return changed
|
||||||
|
|
||||||
|
|
||||||
|
def update_transcription_settings(config: Config, query: QueryParams) -> bool:
|
||||||
|
transcription = config.transcription
|
||||||
|
changed = False
|
||||||
|
|
||||||
|
enabled = query_first(query, "enabled")
|
||||||
|
if enabled is not None:
|
||||||
|
parsed_enabled = parse_bool(enabled, "enabled")
|
||||||
|
if transcription.enabled != parsed_enabled:
|
||||||
|
transcription.enabled = parsed_enabled
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
provider = query_first(query, "provider")
|
||||||
|
if provider is not None:
|
||||||
|
provider = provider.strip().lower()
|
||||||
|
provider_spec = resolve_transcription_provider(provider)
|
||||||
|
if provider_spec is None:
|
||||||
|
raise WebUISettingsError("unknown transcription provider")
|
||||||
|
provider = provider_spec.name
|
||||||
|
if transcription.provider != provider:
|
||||||
|
transcription.provider = provider
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
model = query_first(query, "model")
|
||||||
|
if model is not None:
|
||||||
|
model = model.strip() or None
|
||||||
|
if model is not None and len(model) > 200:
|
||||||
|
raise WebUISettingsError("transcription model is too long")
|
||||||
|
if transcription.model != model:
|
||||||
|
transcription.model = model
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
language = query_first(query, "language")
|
||||||
|
if language is not None:
|
||||||
|
language = language.strip().lower() or None
|
||||||
|
if language is not None and not re.fullmatch(r"[a-z]{2,3}", language):
|
||||||
|
raise WebUISettingsError(
|
||||||
|
"transcription language must be 2-3 lowercase letters"
|
||||||
|
)
|
||||||
|
if transcription.language != language:
|
||||||
|
transcription.language = language
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
max_duration_sec = query_first_alias(query, "max_duration_sec", "maxDurationSec")
|
||||||
|
if max_duration_sec is not None:
|
||||||
|
try:
|
||||||
|
parsed_duration = int(max_duration_sec)
|
||||||
|
except ValueError:
|
||||||
|
raise WebUISettingsError("max_duration_sec must be an integer") from None
|
||||||
|
if parsed_duration < 1 or parsed_duration > 600:
|
||||||
|
raise WebUISettingsError("max_duration_sec must be between 1 and 600")
|
||||||
|
if transcription.max_duration_sec != parsed_duration:
|
||||||
|
transcription.max_duration_sec = parsed_duration
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
max_upload_mb = query_first_alias(query, "max_upload_mb", "maxUploadMb")
|
||||||
|
if max_upload_mb is not None:
|
||||||
|
try:
|
||||||
|
parsed_upload = int(max_upload_mb)
|
||||||
|
except ValueError:
|
||||||
|
raise WebUISettingsError("max_upload_mb must be an integer") from None
|
||||||
|
if parsed_upload < 1 or parsed_upload > 100:
|
||||||
|
raise WebUISettingsError("max_upload_mb must be between 1 and 100")
|
||||||
|
if transcription.max_upload_mb != parsed_upload:
|
||||||
|
transcription.max_upload_mb = parsed_upload
|
||||||
|
changed = True
|
||||||
|
return changed
|
||||||
|
|
||||||
|
|
||||||
|
def network_safety_payload(config: Config) -> dict[str, Any]:
|
||||||
|
"""Return the network-related fields embedded in the advanced DTO."""
|
||||||
|
return {
|
||||||
|
"webui_allow_local_service_access": config.tools.webui_allow_local_service_access,
|
||||||
|
"allow_local_preview_access": config.tools.webui_allow_local_service_access,
|
||||||
|
"webui_default_access_mode": read_webui_default_access_mode(),
|
||||||
|
"private_service_protection_enabled": True,
|
||||||
|
"ssrf_whitelist_count": len(config.tools.ssrf_whitelist),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def masked_api_secret(value: str) -> str | None:
|
||||||
|
value = value.strip()
|
||||||
|
if not value:
|
||||||
|
return None
|
||||||
|
return f"{value[:3]}...{value[-4:]}" if len(value) > 8 else "configured"
|
||||||
|
|
||||||
|
|
||||||
|
def api_runtime_message(message: str) -> str:
|
||||||
|
known = {
|
||||||
|
"api_exited_during_startup": "API server exited during startup. Check its log for details.",
|
||||||
|
"api_stop_timeout": "API server did not stop in time.",
|
||||||
|
"api_state_stale": "API server state was stale; try starting it again.",
|
||||||
|
}
|
||||||
|
if message in known:
|
||||||
|
return known[message]
|
||||||
|
if message.startswith("api_"):
|
||||||
|
return f"API server {message.removeprefix('api_').replace('_', ' ')}"
|
||||||
|
return message.replace("_", " ")
|
||||||
|
|
||||||
|
|
||||||
|
def api_service_payload(
|
||||||
|
settings: WebUISettingsServices,
|
||||||
|
runtime: ApiRuntime,
|
||||||
|
*,
|
||||||
|
last_action: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
config = settings.config.load()
|
||||||
|
status = runtime.status()
|
||||||
|
extras = optional_dependency_groups()
|
||||||
|
connect_host = (
|
||||||
|
"127.0.0.1" if config.api.host in {"0.0.0.0", "::"} else config.api.host
|
||||||
|
)
|
||||||
|
payload = {
|
||||||
|
"installed": extra_installed("api", extras.get("api")),
|
||||||
|
"running": status.running,
|
||||||
|
"managed": status.running,
|
||||||
|
"host": config.api.host,
|
||||||
|
"port": config.api.port,
|
||||||
|
"timeout": config.api.timeout,
|
||||||
|
"api_key_hint": masked_api_secret(config.api.api_key),
|
||||||
|
"endpoint": f"http://{connect_host}:{config.api.port}/v1",
|
||||||
|
"command": "nanobot serve",
|
||||||
|
"log_path": str(status.log_path),
|
||||||
|
}
|
||||||
|
if last_action:
|
||||||
|
payload["last_action"] = last_action
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
class CapabilitySettingsHandler:
|
||||||
|
"""Handle capability commands after transport authentication and decoding."""
|
||||||
|
|
||||||
|
def __init__(self, settings: WebUISettingsServices, logger: Any) -> None:
|
||||||
|
self.settings = settings
|
||||||
|
self.logger = logger
|
||||||
|
|
||||||
|
async def handle(
|
||||||
|
self,
|
||||||
|
action: str,
|
||||||
|
request: SettingsRequest,
|
||||||
|
operations: CapabilitySettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
if action == "api-status":
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
api_service_payload(self.settings, operations.api_runtime())
|
||||||
|
)
|
||||||
|
if action == "api-start":
|
||||||
|
return await self._start_api(request, operations)
|
||||||
|
if action == "api-stop":
|
||||||
|
return await self._stop_api(operations)
|
||||||
|
|
||||||
|
mutation = {
|
||||||
|
"web-search-update": (
|
||||||
|
operations.update_web_search,
|
||||||
|
"browser",
|
||||||
|
False,
|
||||||
|
),
|
||||||
|
"transcription-update": (
|
||||||
|
operations.update_transcription,
|
||||||
|
None,
|
||||||
|
False,
|
||||||
|
),
|
||||||
|
"network-update": (
|
||||||
|
operations.update_network,
|
||||||
|
"runtime",
|
||||||
|
False,
|
||||||
|
),
|
||||||
|
"image-update": (
|
||||||
|
operations.update_image,
|
||||||
|
"image",
|
||||||
|
True,
|
||||||
|
),
|
||||||
|
}.get(action)
|
||||||
|
if mutation is None:
|
||||||
|
return SettingsRouteResult.failure(404, "unknown settings action")
|
||||||
|
|
||||||
|
operation, section, apply_image_reload = mutation
|
||||||
|
try:
|
||||||
|
payload = self.settings.mutate(operation, request.query)
|
||||||
|
except WebUISettingsError as exc:
|
||||||
|
return SettingsRouteResult.failure(exc.status, exc.message)
|
||||||
|
if apply_image_reload:
|
||||||
|
payload, image_restart_cleared = await self.apply_image_runtime_change(
|
||||||
|
payload,
|
||||||
|
operations.reload_image,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
image_restart_cleared = False
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
payload,
|
||||||
|
decorate_restart=True,
|
||||||
|
restart_section=section,
|
||||||
|
clear_restart_section=("image" if image_restart_cleared else None),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def apply_image_runtime_change(
|
||||||
|
self,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
reload_image: Callable[[], Awaitable[dict[str, Any]]],
|
||||||
|
) -> tuple[dict[str, Any], bool]:
|
||||||
|
"""Hot-apply image settings, preserving restart fallback on failure."""
|
||||||
|
if not payload.get("requires_restart"):
|
||||||
|
return payload, False
|
||||||
|
try:
|
||||||
|
result = await reload_image()
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("failed to hot-reload image generation settings")
|
||||||
|
return payload, False
|
||||||
|
|
||||||
|
applied = bool(result.get("ok")) and not result.get("requires_restart")
|
||||||
|
updated = dict(payload)
|
||||||
|
updated["requires_restart"] = not applied
|
||||||
|
if not applied:
|
||||||
|
self.logger.warning(
|
||||||
|
"image generation settings were saved but require restart: {}",
|
||||||
|
result.get("message") or "hot reload failed",
|
||||||
|
)
|
||||||
|
return updated, applied
|
||||||
|
|
||||||
|
async def _start_api(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
operations: CapabilitySettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
api_key = (request.payload or {}).get("api_key")
|
||||||
|
if api_key is not None and not isinstance(api_key, str):
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
400,
|
||||||
|
"API service API key must be a string",
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
await asyncio.to_thread(
|
||||||
|
self.settings.mutate,
|
||||||
|
operations.nanobot_features_action,
|
||||||
|
"enable",
|
||||||
|
{"name": ["api"]},
|
||||||
|
allow_install=self._allow_feature_package_install(request),
|
||||||
|
)
|
||||||
|
self.settings.mutate(operations.update_api, request.query)
|
||||||
|
config = self.settings.config.load()
|
||||||
|
runtime = operations.api_runtime()
|
||||||
|
options = ApiStartOptions(
|
||||||
|
host=config.api.host,
|
||||||
|
port=config.api.port,
|
||||||
|
workspace=str(config.workspace_path),
|
||||||
|
config_path=str(self.settings.config.path),
|
||||||
|
)
|
||||||
|
current = runtime.status()
|
||||||
|
result = await asyncio.to_thread(
|
||||||
|
runtime.restart if current.running else runtime.start_background,
|
||||||
|
options,
|
||||||
|
)
|
||||||
|
if not result.ok:
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
500,
|
||||||
|
api_runtime_message(result.message),
|
||||||
|
)
|
||||||
|
except (WebUISettingsError, OptionalFeatureError) as exc:
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
getattr(exc, "status", 400),
|
||||||
|
getattr(exc, "message", str(exc)),
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
self.logger.exception("failed to start managed API service")
|
||||||
|
return SettingsRouteResult.failure(500, str(exc))
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
api_service_payload(
|
||||||
|
self.settings,
|
||||||
|
operations.api_runtime(),
|
||||||
|
last_action="started",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _stop_api(
|
||||||
|
self,
|
||||||
|
operations: CapabilitySettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
runtime = operations.api_runtime()
|
||||||
|
try:
|
||||||
|
result = await asyncio.to_thread(runtime.stop)
|
||||||
|
except Exception as exc:
|
||||||
|
self.logger.exception("failed to stop managed API service")
|
||||||
|
return SettingsRouteResult.failure(500, str(exc))
|
||||||
|
if not result.ok and result.message != "api_not_running":
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
500,
|
||||||
|
api_runtime_message(result.message),
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
api_service_payload(
|
||||||
|
self.settings,
|
||||||
|
operations.api_runtime(),
|
||||||
|
last_action="stopped",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _allow_feature_package_install(self, request: SettingsRequest) -> bool:
|
||||||
|
if request.local_browser:
|
||||||
|
return True
|
||||||
|
try:
|
||||||
|
return bool(
|
||||||
|
self.settings.config.load().tools.webui_allow_remote_package_install
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("failed to load remote package install policy")
|
||||||
|
return False
|
||||||
@@ -0,0 +1,82 @@
|
|||||||
|
"""Stable request and error contracts shared by WebUI settings domains."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
QueryParams = dict[str, list[str]]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SettingsRequest:
|
||||||
|
"""Transport-neutral input decoded by the settings route facade."""
|
||||||
|
|
||||||
|
query: QueryParams
|
||||||
|
payload: dict[str, Any] | None = None
|
||||||
|
local_browser: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SettingsRouteResult:
|
||||||
|
"""Transport-neutral result returned by a settings domain handler."""
|
||||||
|
|
||||||
|
payload: dict[str, Any] | None = None
|
||||||
|
status: int = 200
|
||||||
|
error: str | None = None
|
||||||
|
decorate_restart: bool = False
|
||||||
|
restart_section: str | None = None
|
||||||
|
clear_restart_section: str | None = None
|
||||||
|
restart_payload_key: str | None = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def success(
|
||||||
|
cls,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
*,
|
||||||
|
decorate_restart: bool = False,
|
||||||
|
restart_section: str | None = None,
|
||||||
|
clear_restart_section: str | None = None,
|
||||||
|
restart_payload_key: str | None = None,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
return cls(
|
||||||
|
payload=payload,
|
||||||
|
decorate_restart=decorate_restart,
|
||||||
|
restart_section=restart_section,
|
||||||
|
clear_restart_section=clear_restart_section,
|
||||||
|
restart_payload_key=restart_payload_key,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def failure(cls, status: int, error: str) -> SettingsRouteResult:
|
||||||
|
return cls(status=status, error=error)
|
||||||
|
|
||||||
|
|
||||||
|
class WebUISettingsError(ValueError):
|
||||||
|
"""User-facing settings validation failure."""
|
||||||
|
|
||||||
|
def __init__(self, message: str, *, status: int = 400) -> None:
|
||||||
|
super().__init__(message)
|
||||||
|
self.message = message
|
||||||
|
self.status = status
|
||||||
|
|
||||||
|
|
||||||
|
def query_first(query: QueryParams, key: str) -> str | None:
|
||||||
|
values = query.get(key)
|
||||||
|
return values[0] if values else None
|
||||||
|
|
||||||
|
|
||||||
|
def query_first_alias(query: QueryParams, snake: str, camel: str) -> str | None:
|
||||||
|
value = query_first(query, snake)
|
||||||
|
return query_first(query, camel) if value is None else value
|
||||||
|
|
||||||
|
|
||||||
|
def query_has_alias(query: QueryParams, snake: str, camel: str) -> bool:
|
||||||
|
return snake in query or camel in query
|
||||||
|
|
||||||
|
|
||||||
|
def parse_bool(value: str, field: str) -> bool:
|
||||||
|
normalized = value.strip().lower()
|
||||||
|
if normalized not in {"1", "0", "true", "false", "yes", "no"}:
|
||||||
|
raise WebUISettingsError(f"{field} must be boolean")
|
||||||
|
return normalized in {"1", "true", "yes"}
|
||||||
File diff suppressed because it is too large
Load Diff
+583
-1006
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,148 @@
|
|||||||
|
"""Gateway-owned state for the WebUI settings surface."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import threading
|
||||||
|
from collections.abc import Callable
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, TypeVar
|
||||||
|
|
||||||
|
from nanobot.config.loader import load_config, save_config
|
||||||
|
from nanobot.config.schema import Config
|
||||||
|
|
||||||
|
_T = TypeVar("_T")
|
||||||
|
_WEBUI_OAUTH_MAX_FLOWS = 8
|
||||||
|
|
||||||
|
|
||||||
|
class WebUISettingsConfig:
|
||||||
|
"""Instance-scoped config access with serialized read-modify-write operations."""
|
||||||
|
|
||||||
|
def __init__(self, config_path: Path) -> None:
|
||||||
|
self.path = config_path.expanduser().resolve(strict=False)
|
||||||
|
self._lock = threading.RLock()
|
||||||
|
|
||||||
|
def load(self) -> Config:
|
||||||
|
"""Load this gateway's config without consulting the process-global path."""
|
||||||
|
with self._lock:
|
||||||
|
return load_config(self.path)
|
||||||
|
|
||||||
|
def update(self, mutation: Callable[[Config], _T]) -> _T:
|
||||||
|
"""Apply and atomically persist one in-process read-modify-write operation."""
|
||||||
|
with self._lock:
|
||||||
|
config = load_config(self.path)
|
||||||
|
result = mutation(config)
|
||||||
|
save_config(config, self.path)
|
||||||
|
return result
|
||||||
|
|
||||||
|
def run_serialized(self, operation: Callable[[Path], _T]) -> _T:
|
||||||
|
"""Run a path-aware read-modify-write operation under the instance lock."""
|
||||||
|
with self._lock:
|
||||||
|
return operation(self.path)
|
||||||
|
|
||||||
|
|
||||||
|
class WebUIOAuthFlowRegistry:
|
||||||
|
"""Bounded, thread-safe OAuth flows owned by one gateway instance."""
|
||||||
|
|
||||||
|
def __init__(self, *, max_flows: int = _WEBUI_OAUTH_MAX_FLOWS) -> None:
|
||||||
|
if max_flows < 1:
|
||||||
|
raise ValueError("max_flows must be at least one")
|
||||||
|
self._max_flows = max_flows
|
||||||
|
self._flows: dict[str, tuple[str, Any]] = {}
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def register(self, provider_name: str, flow_id: str, flow: Any) -> None:
|
||||||
|
discarded: list[Any] = []
|
||||||
|
with self._lock:
|
||||||
|
for existing_id, (_provider_name, existing) in list(self._flows.items()):
|
||||||
|
if existing.expired:
|
||||||
|
discarded.append(self._flows.pop(existing_id)[1])
|
||||||
|
while len(self._flows) >= self._max_flows:
|
||||||
|
oldest_id = next(iter(self._flows))
|
||||||
|
discarded.append(self._flows.pop(oldest_id)[1])
|
||||||
|
self._flows[flow_id] = (provider_name, flow)
|
||||||
|
for existing in discarded:
|
||||||
|
existing.cancel()
|
||||||
|
|
||||||
|
def get(self, provider_name: str, flow_id: str) -> Any | None:
|
||||||
|
with self._lock:
|
||||||
|
registered = self._flows.get(flow_id)
|
||||||
|
if registered is None or registered[0] != provider_name:
|
||||||
|
return None
|
||||||
|
flow = registered[1]
|
||||||
|
if not flow.expired:
|
||||||
|
return flow
|
||||||
|
self._flows.pop(flow_id, None)
|
||||||
|
flow.cancel()
|
||||||
|
return None
|
||||||
|
|
||||||
|
def remove(
|
||||||
|
self,
|
||||||
|
provider_name: str,
|
||||||
|
flow_id: str,
|
||||||
|
flow: Any,
|
||||||
|
*,
|
||||||
|
cancel: bool = True,
|
||||||
|
) -> None:
|
||||||
|
with self._lock:
|
||||||
|
registered = self._flows.get(flow_id)
|
||||||
|
if (
|
||||||
|
registered is not None
|
||||||
|
and registered[0] == provider_name
|
||||||
|
and registered[1] is flow
|
||||||
|
):
|
||||||
|
self._flows.pop(flow_id)
|
||||||
|
if cancel:
|
||||||
|
flow.cancel()
|
||||||
|
|
||||||
|
def clear(self, provider_name: str) -> None:
|
||||||
|
with self._lock:
|
||||||
|
flow_ids = [
|
||||||
|
flow_id
|
||||||
|
for flow_id, (registered_provider, _flow) in self._flows.items()
|
||||||
|
if registered_provider == provider_name
|
||||||
|
]
|
||||||
|
flows = [self._flows.pop(flow_id)[1] for flow_id in flow_ids]
|
||||||
|
for flow in flows:
|
||||||
|
flow.cancel()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class WebUISettingsServices:
|
||||||
|
"""Settings dependencies composed once for a gateway instance."""
|
||||||
|
|
||||||
|
config: WebUISettingsConfig
|
||||||
|
oauth_flows: WebUIOAuthFlowRegistry
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, config_path: Path) -> WebUISettingsServices:
|
||||||
|
return cls(
|
||||||
|
config=WebUISettingsConfig(config_path),
|
||||||
|
oauth_flows=WebUIOAuthFlowRegistry(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def read(
|
||||||
|
self,
|
||||||
|
operation: Callable[..., _T],
|
||||||
|
/,
|
||||||
|
*args: Any,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> _T:
|
||||||
|
"""Run a settings read against this gateway's explicit config path."""
|
||||||
|
return operation(*args, config_path=self.config.path, **kwargs)
|
||||||
|
|
||||||
|
def mutate(
|
||||||
|
self,
|
||||||
|
operation: Callable[..., _T],
|
||||||
|
/,
|
||||||
|
*args: Any,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> _T:
|
||||||
|
"""Serialize a path-aware settings read-modify-write operation."""
|
||||||
|
return self.config.run_serialized(
|
||||||
|
lambda config_path: operation(
|
||||||
|
*args,
|
||||||
|
config_path=config_path,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
)
|
||||||
@@ -0,0 +1,959 @@
|
|||||||
|
"""System and channel settings domain logic."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import inspect
|
||||||
|
import re
|
||||||
|
import time
|
||||||
|
from collections.abc import Callable, Iterable, Mapping
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Any, TypedDict, cast
|
||||||
|
from zoneinfo import ZoneInfo
|
||||||
|
|
||||||
|
from nanobot.channels._setup import channel_setup_spec
|
||||||
|
from nanobot.channels.connect import ChannelConnectError
|
||||||
|
from nanobot.channels.contracts import (
|
||||||
|
RouteFieldType,
|
||||||
|
channel_instance_config,
|
||||||
|
channel_update_instance_config,
|
||||||
|
)
|
||||||
|
from nanobot.config.schema import Config
|
||||||
|
from nanobot.optional_features import OptionalFeatureError, with_channel_runtime_status
|
||||||
|
from nanobot.security.workspace_access import workspace_sandbox_status
|
||||||
|
from nanobot.webui.settings_capabilities import network_safety_payload
|
||||||
|
from nanobot.webui.settings_contracts import (
|
||||||
|
QueryParams,
|
||||||
|
SettingsRequest,
|
||||||
|
SettingsRouteResult,
|
||||||
|
WebUISettingsError,
|
||||||
|
query_first,
|
||||||
|
query_first_alias,
|
||||||
|
)
|
||||||
|
from nanobot.webui.token_usage import token_usage_payload
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from nanobot.webui.settings_services import WebUISettingsServices
|
||||||
|
|
||||||
|
LoadChannelPlugin = Callable[[str], Any]
|
||||||
|
ListPendingPairings = Callable[[], Iterable[dict[str, Any]]]
|
||||||
|
SettingsOperation = Callable[..., Any]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SystemSettingsOperations:
|
||||||
|
cli_apps_payload: SettingsOperation
|
||||||
|
cli_apps_action: SettingsOperation
|
||||||
|
nanobot_features_payload: SettingsOperation
|
||||||
|
nanobot_features_action: SettingsOperation
|
||||||
|
nanobot_feature_instance_target: SettingsOperation
|
||||||
|
validate_channel_config: SettingsOperation
|
||||||
|
load_channel_plugin: LoadChannelPlugin
|
||||||
|
list_pending: ListPendingPairings
|
||||||
|
approve_code: SettingsOperation
|
||||||
|
deny_code: SettingsOperation
|
||||||
|
mcp_presets_action: SettingsOperation
|
||||||
|
reload_mcp: SettingsOperation
|
||||||
|
mcp_runtime_status: Callable[[], Mapping[str, str]] | None
|
||||||
|
check_for_update: SettingsOperation
|
||||||
|
channel_feature_action: SettingsOperation | None = None
|
||||||
|
channel_runtime_status: Callable[[], dict[str, Any]] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class SystemSettingsPayload(TypedDict):
|
||||||
|
runtime: dict[str, Any]
|
||||||
|
usage: dict[str, Any]
|
||||||
|
advanced: dict[str, Any]
|
||||||
|
version: dict[str, Any]
|
||||||
|
docs: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
_DOCS_STABLE_VERSION_RE = re.compile(r"^\d+\.\d+\.\d+(?:\.post\d+)?$")
|
||||||
|
_DOCS_LATEST_URL = "https://nanobot.wiki/docs/latest"
|
||||||
|
_SKIP_FIELD = object()
|
||||||
|
|
||||||
|
|
||||||
|
def docs_version(version: str) -> str:
|
||||||
|
"""Map package versions to the matching public docs path."""
|
||||||
|
normalized = version.strip()
|
||||||
|
if _DOCS_STABLE_VERSION_RE.fullmatch(normalized):
|
||||||
|
return normalized
|
||||||
|
return "latest"
|
||||||
|
|
||||||
|
|
||||||
|
def docs_payload(version: str) -> dict[str, Any]:
|
||||||
|
selected_version = docs_version(version)
|
||||||
|
base_url = f"https://nanobot.wiki/docs/{selected_version}"
|
||||||
|
return {
|
||||||
|
"version": selected_version,
|
||||||
|
"base_url": base_url,
|
||||||
|
"chat_apps_url": f"{base_url}/getting-started/chat-apps",
|
||||||
|
"latest_url": _DOCS_LATEST_URL,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def system_settings_payload(
|
||||||
|
config: Config,
|
||||||
|
*,
|
||||||
|
config_path: Path,
|
||||||
|
version: str,
|
||||||
|
) -> SystemSettingsPayload:
|
||||||
|
defaults = config.agents.defaults
|
||||||
|
exec_config = config.tools.exec
|
||||||
|
sandbox_status = workspace_sandbox_status(
|
||||||
|
restrict_to_workspace=config.tools.restrict_to_workspace,
|
||||||
|
workspace=config.workspace_path,
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"runtime": {
|
||||||
|
"config_path": str(config_path.expanduser()),
|
||||||
|
"workspace_path": str(config.workspace_path),
|
||||||
|
"gateway_host": config.gateway.host,
|
||||||
|
"gateway_port": config.gateway.port,
|
||||||
|
"heartbeat": {
|
||||||
|
"enabled": config.gateway.heartbeat.enabled,
|
||||||
|
"interval_s": config.gateway.heartbeat.interval_s,
|
||||||
|
"keep_recent_messages": config.gateway.heartbeat.keep_recent_messages,
|
||||||
|
},
|
||||||
|
"dream": {
|
||||||
|
"schedule": defaults.dream.describe_schedule(),
|
||||||
|
},
|
||||||
|
"unified_session": defaults.unified_session,
|
||||||
|
},
|
||||||
|
"usage": token_usage_payload(timezone_name=defaults.timezone),
|
||||||
|
"advanced": {
|
||||||
|
"restrict_to_workspace": config.tools.restrict_to_workspace,
|
||||||
|
"workspace_sandbox": sandbox_status.as_dict(),
|
||||||
|
**network_safety_payload(config),
|
||||||
|
"mcp_server_count": len(config.tools.mcp_servers),
|
||||||
|
"exec_enabled": exec_config.enable,
|
||||||
|
"exec_sandbox": exec_config.sandbox or None,
|
||||||
|
"exec_path_prepend_set": bool(exec_config.path_prepend),
|
||||||
|
"exec_path_append_set": bool(exec_config.path_append),
|
||||||
|
},
|
||||||
|
"version": {"current": version},
|
||||||
|
"docs": docs_payload(version),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def settings_usage_payload(config: Config) -> dict[str, Any]:
|
||||||
|
"""Return the lightweight token usage slice for Overview refreshes."""
|
||||||
|
return token_usage_payload(timezone_name=config.agents.defaults.timezone)
|
||||||
|
|
||||||
|
|
||||||
|
def update_agent_system_settings(config: Config, query: QueryParams) -> tuple[bool, bool]:
|
||||||
|
defaults = config.agents.defaults
|
||||||
|
changed = False
|
||||||
|
restart_required = False
|
||||||
|
|
||||||
|
timezone = query_first(query, "timezone")
|
||||||
|
if timezone is not None:
|
||||||
|
timezone = timezone.strip()
|
||||||
|
if not timezone:
|
||||||
|
raise WebUISettingsError("timezone is required")
|
||||||
|
try:
|
||||||
|
ZoneInfo(timezone)
|
||||||
|
except Exception:
|
||||||
|
raise WebUISettingsError("invalid timezone") from None
|
||||||
|
timezone_changed = defaults.timezone != timezone
|
||||||
|
if timezone_changed or defaults.timezone_mode != "manual":
|
||||||
|
defaults.timezone = timezone
|
||||||
|
defaults.timezone_mode = "manual"
|
||||||
|
changed = True
|
||||||
|
restart_required = timezone_changed
|
||||||
|
|
||||||
|
tool_hint_max_length = query_first_alias(
|
||||||
|
query,
|
||||||
|
"tool_hint_max_length",
|
||||||
|
"toolHintMaxLength",
|
||||||
|
)
|
||||||
|
if tool_hint_max_length is not None:
|
||||||
|
try:
|
||||||
|
parsed = int(tool_hint_max_length)
|
||||||
|
except ValueError:
|
||||||
|
raise WebUISettingsError(
|
||||||
|
"tool_hint_max_length must be an integer"
|
||||||
|
) from None
|
||||||
|
if parsed < 20 or parsed > 500:
|
||||||
|
raise WebUISettingsError(
|
||||||
|
"tool_hint_max_length must be between 20 and 500"
|
||||||
|
)
|
||||||
|
if defaults.tool_hint_max_length != parsed:
|
||||||
|
defaults.tool_hint_max_length = parsed
|
||||||
|
changed = True
|
||||||
|
restart_required = True
|
||||||
|
return changed, restart_required
|
||||||
|
|
||||||
|
|
||||||
|
def save_channel_config_values(
|
||||||
|
config: Config,
|
||||||
|
name: str,
|
||||||
|
raw_values: dict[str, Any],
|
||||||
|
instance_id: str = "default",
|
||||||
|
*,
|
||||||
|
load_channel_plugin: LoadChannelPlugin,
|
||||||
|
) -> list[str]:
|
||||||
|
if not name:
|
||||||
|
raise WebUISettingsError("missing channel name")
|
||||||
|
try:
|
||||||
|
plugin = load_channel_plugin(name)
|
||||||
|
except ImportError:
|
||||||
|
raise WebUISettingsError(f"unknown channel '{name}'", status=404) from None
|
||||||
|
setup_spec = channel_setup_spec(name, plugin=plugin)
|
||||||
|
if setup_spec is None:
|
||||||
|
raise WebUISettingsError(
|
||||||
|
f"channel '{name}' cannot be configured from WebUI",
|
||||||
|
status=404,
|
||||||
|
)
|
||||||
|
field_types = setup_spec.route_field_types
|
||||||
|
if not raw_values:
|
||||||
|
return []
|
||||||
|
|
||||||
|
section = getattr(config.channels, name, None)
|
||||||
|
channel_config = channel_instance_config(
|
||||||
|
plugin,
|
||||||
|
section,
|
||||||
|
instance_id=instance_id,
|
||||||
|
)
|
||||||
|
saved: list[str] = []
|
||||||
|
prefix = f"channels.{name}."
|
||||||
|
for raw_key, raw_value in raw_values.items():
|
||||||
|
if not raw_key:
|
||||||
|
raise WebUISettingsError(
|
||||||
|
"channel settings payload contains an invalid key"
|
||||||
|
)
|
||||||
|
field = raw_key[len(prefix) :] if raw_key.startswith(prefix) else raw_key
|
||||||
|
value_type = field_types.get(field)
|
||||||
|
if value_type is None:
|
||||||
|
raise WebUISettingsError(f"'{raw_key}' cannot be configured from WebUI")
|
||||||
|
value = coerce_channel_value(raw_key, raw_value, value_type)
|
||||||
|
if value is _SKIP_FIELD:
|
||||||
|
continue
|
||||||
|
assign_channel_config_value(channel_config, field, value)
|
||||||
|
saved.append(raw_key)
|
||||||
|
|
||||||
|
try:
|
||||||
|
updated_section = channel_update_instance_config(
|
||||||
|
plugin,
|
||||||
|
section,
|
||||||
|
channel_config,
|
||||||
|
instance_id=instance_id,
|
||||||
|
)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise WebUISettingsError(
|
||||||
|
f"Invalid {name} configuration: {exc}",
|
||||||
|
status=400,
|
||||||
|
) from exc
|
||||||
|
setattr(config.channels, name, updated_section)
|
||||||
|
return saved
|
||||||
|
|
||||||
|
|
||||||
|
def coerce_channel_value(
|
||||||
|
raw_key: str,
|
||||||
|
raw_value: Any,
|
||||||
|
value_type: RouteFieldType,
|
||||||
|
) -> Any:
|
||||||
|
if isinstance(value_type, tuple):
|
||||||
|
kind = value_type[0]
|
||||||
|
allowed = value_type[1]
|
||||||
|
else:
|
||||||
|
kind = value_type
|
||||||
|
allowed = None
|
||||||
|
|
||||||
|
if kind in {"string", "secret"}:
|
||||||
|
value = raw_value.strip() if isinstance(raw_value, str) else str(raw_value)
|
||||||
|
if kind == "secret" and not value:
|
||||||
|
return _SKIP_FIELD
|
||||||
|
return value
|
||||||
|
|
||||||
|
if kind == "list":
|
||||||
|
if raw_value is None:
|
||||||
|
return []
|
||||||
|
if isinstance(raw_value, str):
|
||||||
|
return [item.strip() for item in raw_value.split(",") if item.strip()]
|
||||||
|
if isinstance(raw_value, list):
|
||||||
|
return [
|
||||||
|
str(item).strip()
|
||||||
|
for item in cast(list[Any], raw_value)
|
||||||
|
if str(item).strip()
|
||||||
|
]
|
||||||
|
raise WebUISettingsError(f"'{raw_key}' must be a comma-separated list")
|
||||||
|
|
||||||
|
if kind == "int":
|
||||||
|
if raw_value in (None, ""):
|
||||||
|
return _SKIP_FIELD
|
||||||
|
try:
|
||||||
|
return int(raw_value)
|
||||||
|
except (TypeError, ValueError) as exc:
|
||||||
|
raise WebUISettingsError(f"'{raw_key}' must be a number") from exc
|
||||||
|
|
||||||
|
if kind == "bool":
|
||||||
|
if isinstance(raw_value, bool):
|
||||||
|
return raw_value
|
||||||
|
value = str(raw_value).strip().lower()
|
||||||
|
if value in {"true", "1", "yes", "on"}:
|
||||||
|
return True
|
||||||
|
if value in {"false", "0", "no", "off"}:
|
||||||
|
return False
|
||||||
|
raise WebUISettingsError(f"'{raw_key}' must be true or false")
|
||||||
|
|
||||||
|
if kind == "enum":
|
||||||
|
value = raw_value.strip() if isinstance(raw_value, str) else str(raw_value)
|
||||||
|
if not value:
|
||||||
|
return _SKIP_FIELD
|
||||||
|
if allowed is None or value not in allowed:
|
||||||
|
options = ", ".join(sorted(allowed or ()))
|
||||||
|
raise WebUISettingsError(f"'{raw_key}' must be one of: {options}")
|
||||||
|
return value
|
||||||
|
|
||||||
|
raise WebUISettingsError(f"'{raw_key}' has an unsupported field type")
|
||||||
|
|
||||||
|
|
||||||
|
def assign_channel_config_value(
|
||||||
|
channel_config: dict[str, Any],
|
||||||
|
field: str,
|
||||||
|
value: Any,
|
||||||
|
) -> None:
|
||||||
|
target = channel_config
|
||||||
|
parts = field.split(".")
|
||||||
|
for part in parts[:-1]:
|
||||||
|
current: object = target.get(part)
|
||||||
|
if not isinstance(current, dict):
|
||||||
|
current = {}
|
||||||
|
target[part] = current
|
||||||
|
target = cast(dict[str, Any], current)
|
||||||
|
target[parts[-1]] = value
|
||||||
|
|
||||||
|
|
||||||
|
def pairing_payload(
|
||||||
|
list_pending: ListPendingPairings,
|
||||||
|
last_action: dict[str, Any] | None = None,
|
||||||
|
*,
|
||||||
|
now: float | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
current_time = time.time() if now is None else now
|
||||||
|
requests: list[dict[str, Any]] = []
|
||||||
|
for item in list_pending():
|
||||||
|
expires_at = float(item.get("expires_at", 0) or 0)
|
||||||
|
created_at = float(item.get("created_at", 0) or 0)
|
||||||
|
requests.append(
|
||||||
|
{
|
||||||
|
"code": str(item.get("code", "")),
|
||||||
|
"channel": str(item.get("channel", "")),
|
||||||
|
"sender_id": str(item.get("sender_id", "")),
|
||||||
|
"created_at_ms": int(created_at * 1000) if created_at else None,
|
||||||
|
"expires_at_ms": int(expires_at * 1000) if expires_at else None,
|
||||||
|
"expires_in_seconds": (
|
||||||
|
max(0, int(expires_at - current_time)) if expires_at else None
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
payload: dict[str, Any] = {"requests": requests}
|
||||||
|
if last_action is not None:
|
||||||
|
payload["last_action"] = last_action
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
class SystemSettingsHandler:
|
||||||
|
"""Handle channel and system commands behind a transport-neutral request DTO."""
|
||||||
|
|
||||||
|
def __init__(self, settings: WebUISettingsServices, logger: Any) -> None:
|
||||||
|
self.settings = settings
|
||||||
|
self.logger = logger
|
||||||
|
self._channel_connectors: dict[str, Any] = {}
|
||||||
|
|
||||||
|
async def handle(
|
||||||
|
self,
|
||||||
|
action: str,
|
||||||
|
request: SettingsRequest,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
*,
|
||||||
|
channel_name: str | None = None,
|
||||||
|
connect_action: str | None = None,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
if action == "cli-list":
|
||||||
|
return await self._cli_apps(request, operations)
|
||||||
|
if action.startswith("cli-"):
|
||||||
|
return await self._cli_apps_action(
|
||||||
|
request,
|
||||||
|
action.removeprefix("cli-"),
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
if action == "features-list":
|
||||||
|
return await self._features(operations)
|
||||||
|
if action in {"features-enable", "features-disable"}:
|
||||||
|
return await self._features_action(
|
||||||
|
request,
|
||||||
|
action.removeprefix("features-"),
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
if action == "channel-validate":
|
||||||
|
return await self._channel_validate(request, operations)
|
||||||
|
if action == "channel-configure":
|
||||||
|
return await self._channel_configure(request, operations)
|
||||||
|
if action == "channel-connect" and channel_name and connect_action:
|
||||||
|
return await self._channel_connect(
|
||||||
|
request,
|
||||||
|
channel_name,
|
||||||
|
connect_action,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
if action == "pairing-list":
|
||||||
|
return SettingsRouteResult.success(pairing_payload(operations.list_pending))
|
||||||
|
if action in {"pairing-approve", "pairing-deny"}:
|
||||||
|
return self._pairing_action(
|
||||||
|
request,
|
||||||
|
action.removeprefix("pairing-"),
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
if action == "mcp-list":
|
||||||
|
return await self._mcp_presets(request, None, operations)
|
||||||
|
if action.startswith("mcp-"):
|
||||||
|
return await self._mcp_presets(
|
||||||
|
request,
|
||||||
|
action.removeprefix("mcp-"),
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
if action == "version-check":
|
||||||
|
return await self._version_check(operations)
|
||||||
|
return SettingsRouteResult.failure(404, "unknown settings action")
|
||||||
|
|
||||||
|
async def _cli_apps(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
installed_only = (query_first(request.query, "installed_only") or "").lower() in {
|
||||||
|
"1",
|
||||||
|
"true",
|
||||||
|
"yes",
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
payload = await operations.cli_apps_payload(
|
||||||
|
installed_only=installed_only,
|
||||||
|
config_path=self.settings.config.path,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("failed to load CLI Apps payload")
|
||||||
|
return SettingsRouteResult.failure(500, "failed to load CLI Apps")
|
||||||
|
return SettingsRouteResult.success(payload)
|
||||||
|
|
||||||
|
async def _cli_apps_action(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
action: str,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
try:
|
||||||
|
payload = await asyncio.to_thread(
|
||||||
|
operations.cli_apps_action,
|
||||||
|
action,
|
||||||
|
request.query,
|
||||||
|
config_path=self.settings.config.path,
|
||||||
|
)
|
||||||
|
except WebUISettingsError as exc:
|
||||||
|
return SettingsRouteResult.failure(exc.status, exc.message)
|
||||||
|
except Exception as exc:
|
||||||
|
status = getattr(exc, "status", 500)
|
||||||
|
message = getattr(exc, "message", str(exc))
|
||||||
|
if status >= 500:
|
||||||
|
self.logger.exception("CLI Apps action '{}' failed", action)
|
||||||
|
return SettingsRouteResult.failure(status, message)
|
||||||
|
return SettingsRouteResult.success(payload)
|
||||||
|
|
||||||
|
async def _features(
|
||||||
|
self,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
try:
|
||||||
|
payload = await asyncio.to_thread(
|
||||||
|
operations.nanobot_features_payload,
|
||||||
|
config_path=self.settings.config.path,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("failed to load nanobot features")
|
||||||
|
return SettingsRouteResult.failure(500, "failed to load nanobot features")
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
self._with_channel_runtime_status(payload, operations)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _nanobot_features_payload(
|
||||||
|
self,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
return operations.nanobot_features_payload(config_path=self.settings.config.path)
|
||||||
|
|
||||||
|
def _nanobot_features_action(
|
||||||
|
self,
|
||||||
|
action: str,
|
||||||
|
query: QueryParams,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
*,
|
||||||
|
allow_install: bool = True,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
return self.settings.mutate(
|
||||||
|
operations.nanobot_features_action,
|
||||||
|
action,
|
||||||
|
query,
|
||||||
|
allow_install=allow_install,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _features_action(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
action: str,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
try:
|
||||||
|
payload = await asyncio.to_thread(
|
||||||
|
self._nanobot_features_action,
|
||||||
|
action,
|
||||||
|
request.query,
|
||||||
|
operations,
|
||||||
|
allow_install=(
|
||||||
|
action != "enable"
|
||||||
|
or self.allow_feature_package_install(request)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
except OptionalFeatureError as exc:
|
||||||
|
return SettingsRouteResult.failure(exc.status, exc.message)
|
||||||
|
except Exception as exc:
|
||||||
|
status = getattr(exc, "status", 500)
|
||||||
|
message = getattr(exc, "message", str(exc))
|
||||||
|
if status >= 500:
|
||||||
|
self.logger.exception(
|
||||||
|
"nanobot feature action '{}' failed",
|
||||||
|
action,
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.failure(status, message)
|
||||||
|
payload = await self._apply_feature_runtime_change(
|
||||||
|
action,
|
||||||
|
request.query,
|
||||||
|
payload,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
payload = self._with_channel_runtime_status(payload, operations)
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
payload,
|
||||||
|
decorate_restart=True,
|
||||||
|
restart_section="runtime",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _with_channel_runtime_status(
|
||||||
|
self,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if operations.channel_runtime_status is None:
|
||||||
|
return payload
|
||||||
|
try:
|
||||||
|
return with_channel_runtime_status(
|
||||||
|
payload,
|
||||||
|
operations.channel_runtime_status(),
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("failed to load channel runtime status")
|
||||||
|
return payload
|
||||||
|
|
||||||
|
async def _apply_feature_runtime_change(
|
||||||
|
self,
|
||||||
|
action: str,
|
||||||
|
query: QueryParams,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if operations.channel_feature_action is None:
|
||||||
|
return payload
|
||||||
|
name = (query_first(query, "name") or "").strip()
|
||||||
|
if not name:
|
||||||
|
return payload
|
||||||
|
try:
|
||||||
|
instance_id = operations.nanobot_feature_instance_target(query)
|
||||||
|
result = operations.channel_feature_action(action, name, instance_id)
|
||||||
|
if inspect.isawaitable(result):
|
||||||
|
result = await result
|
||||||
|
except Exception as exc:
|
||||||
|
self.logger.exception("failed to apply channel '{}' without restart", name)
|
||||||
|
return self.feature_runtime_fallback(
|
||||||
|
payload,
|
||||||
|
message=(
|
||||||
|
f"{name} channel config was saved, but hot reload failed: {exc}"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
if not isinstance(result, dict):
|
||||||
|
return payload
|
||||||
|
result = cast(dict[str, Any], result)
|
||||||
|
if not result.get("handled"):
|
||||||
|
return payload
|
||||||
|
|
||||||
|
updated = dict(payload)
|
||||||
|
updated["requires_restart"] = bool(result.get("requires_restart"))
|
||||||
|
message = result.get("message")
|
||||||
|
if isinstance(message, str) and message:
|
||||||
|
last_action = dict(updated.get("last_action") or {})
|
||||||
|
previous = last_action.get("message")
|
||||||
|
last_action["message"] = (
|
||||||
|
f"{previous}. {message}"
|
||||||
|
if isinstance(previous, str) and previous
|
||||||
|
else message
|
||||||
|
)
|
||||||
|
last_action["hot_reload"] = not updated["requires_restart"]
|
||||||
|
if "ok" in result:
|
||||||
|
last_action["ok"] = bool(result["ok"])
|
||||||
|
updated["last_action"] = last_action
|
||||||
|
return updated
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def feature_runtime_fallback(
|
||||||
|
payload: dict[str, Any],
|
||||||
|
*,
|
||||||
|
message: str,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
updated = dict(payload)
|
||||||
|
updated["requires_restart"] = True
|
||||||
|
last_action = dict(updated.get("last_action") or {})
|
||||||
|
previous = last_action.get("message")
|
||||||
|
last_action["message"] = (
|
||||||
|
f"{previous}. {message}"
|
||||||
|
if isinstance(previous, str) and previous
|
||||||
|
else message
|
||||||
|
)
|
||||||
|
last_action["hot_reload"] = False
|
||||||
|
updated["last_action"] = last_action
|
||||||
|
return updated
|
||||||
|
|
||||||
|
async def _channel_configure(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
name = (query_first(request.query, "name") or "").strip()
|
||||||
|
instance_id = (
|
||||||
|
query_first(request.query, "instance_id") or "default"
|
||||||
|
).strip()
|
||||||
|
enable = (query_first(request.query, "enable") or "").strip().lower() in {
|
||||||
|
"1",
|
||||||
|
"true",
|
||||||
|
"yes",
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
saved = await asyncio.to_thread(
|
||||||
|
self._save_channel_config_values,
|
||||||
|
name,
|
||||||
|
self.parse_channel_values(request),
|
||||||
|
instance_id,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
except WebUISettingsError as exc:
|
||||||
|
return SettingsRouteResult.failure(exc.status, exc.message)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("failed to save channel '{}' settings", name)
|
||||||
|
return SettingsRouteResult.failure(500, "failed to save channel settings")
|
||||||
|
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"name": name,
|
||||||
|
"saved": True,
|
||||||
|
"saved_keys": saved,
|
||||||
|
}
|
||||||
|
if not enable:
|
||||||
|
features = await asyncio.to_thread(
|
||||||
|
self._nanobot_features_payload,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
payload["nanobot_features"] = self._with_channel_runtime_status(
|
||||||
|
features,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
payload,
|
||||||
|
decorate_restart=True,
|
||||||
|
restart_section="runtime",
|
||||||
|
restart_payload_key="nanobot_features",
|
||||||
|
)
|
||||||
|
|
||||||
|
feature_query = {"name": [name]}
|
||||||
|
if instance_id:
|
||||||
|
feature_query["instance_id"] = [instance_id]
|
||||||
|
try:
|
||||||
|
features = await asyncio.to_thread(
|
||||||
|
self._nanobot_features_action,
|
||||||
|
"enable",
|
||||||
|
feature_query,
|
||||||
|
operations,
|
||||||
|
allow_install=self.allow_feature_package_install(request),
|
||||||
|
)
|
||||||
|
except OptionalFeatureError as exc:
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
exc.status,
|
||||||
|
f"Settings saved, but {exc.message}",
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
self.logger.exception(
|
||||||
|
"failed to enable channel '{}' after settings save",
|
||||||
|
name,
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
500,
|
||||||
|
f"Settings saved, but enabling {name} failed: {exc}",
|
||||||
|
)
|
||||||
|
|
||||||
|
features = await self._apply_feature_runtime_change(
|
||||||
|
"enable",
|
||||||
|
feature_query,
|
||||||
|
features,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
payload["nanobot_features"] = self._with_channel_runtime_status(
|
||||||
|
features,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
payload,
|
||||||
|
decorate_restart=True,
|
||||||
|
restart_section="runtime",
|
||||||
|
restart_payload_key="nanobot_features",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _channel_validate(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
name = (query_first(request.query, "name") or "").strip()
|
||||||
|
instance_id = (
|
||||||
|
query_first(request.query, "instance_id") or "default"
|
||||||
|
).strip()
|
||||||
|
try:
|
||||||
|
payload = await asyncio.to_thread(
|
||||||
|
operations.validate_channel_config,
|
||||||
|
name,
|
||||||
|
self.parse_channel_values(request),
|
||||||
|
instance_id=instance_id,
|
||||||
|
)
|
||||||
|
except WebUISettingsError as exc:
|
||||||
|
return SettingsRouteResult.failure(exc.status, exc.message)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("failed to validate channel '{}' settings", name)
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
500,
|
||||||
|
"failed to validate channel settings",
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.success(payload)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def parse_channel_values(request: SettingsRequest) -> dict[str, Any]:
|
||||||
|
if request.payload is None or "values" not in request.payload:
|
||||||
|
return {}
|
||||||
|
values = request.payload.get("values")
|
||||||
|
if not isinstance(values, dict):
|
||||||
|
raise WebUISettingsError(
|
||||||
|
"channel settings payload must be a JSON object"
|
||||||
|
)
|
||||||
|
return cast(dict[str, Any], values)
|
||||||
|
|
||||||
|
def _save_channel_config_values(
|
||||||
|
self,
|
||||||
|
name: str,
|
||||||
|
raw_values: dict[str, Any],
|
||||||
|
instance_id: str,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> list[str]:
|
||||||
|
return self.settings.config.update(
|
||||||
|
lambda config: save_channel_config_values(
|
||||||
|
config,
|
||||||
|
name,
|
||||||
|
raw_values,
|
||||||
|
instance_id,
|
||||||
|
load_channel_plugin=operations.load_channel_plugin,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _channel_connect(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
channel_name: str,
|
||||||
|
action: str,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
try:
|
||||||
|
connector = self._channel_connectors.get(channel_name)
|
||||||
|
if connector is None:
|
||||||
|
plugin = operations.load_channel_plugin(channel_name)
|
||||||
|
connector = plugin.load_connector()
|
||||||
|
self._channel_connectors[channel_name] = connector
|
||||||
|
except ImportError:
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
404,
|
||||||
|
f"channel '{channel_name}' does not support connect",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
payload = await connector.handle(action, request.query)
|
||||||
|
except ChannelConnectError as exc:
|
||||||
|
return SettingsRouteResult.failure(exc.status, exc.message)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception(
|
||||||
|
"failed to run {} WebUI connect action for {}",
|
||||||
|
action,
|
||||||
|
channel_name,
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
500,
|
||||||
|
f"failed to {action} {channel_name} connection",
|
||||||
|
)
|
||||||
|
|
||||||
|
if payload.get("status") != "succeeded":
|
||||||
|
return SettingsRouteResult.success(payload)
|
||||||
|
payload = await self._with_channel_connect_success(
|
||||||
|
request,
|
||||||
|
channel_name,
|
||||||
|
payload,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
payload,
|
||||||
|
decorate_restart=True,
|
||||||
|
restart_section="runtime",
|
||||||
|
restart_payload_key="nanobot_features",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _with_channel_connect_success(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
channel_name: str,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
target = {"name": [channel_name]}
|
||||||
|
if payload.get("instance_id"):
|
||||||
|
target["instance_id"] = [str(payload["instance_id"])]
|
||||||
|
try:
|
||||||
|
features = await asyncio.to_thread(
|
||||||
|
self._nanobot_features_action,
|
||||||
|
"enable",
|
||||||
|
target,
|
||||||
|
operations,
|
||||||
|
allow_install=self.allow_feature_package_install(request),
|
||||||
|
)
|
||||||
|
except OptionalFeatureError as exc:
|
||||||
|
features = self.feature_runtime_fallback(
|
||||||
|
self._nanobot_features_payload(operations),
|
||||||
|
message=(
|
||||||
|
f"{channel_name} connected, but enabling channel support failed: "
|
||||||
|
f"{exc.message}"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
features = await self._apply_feature_runtime_change(
|
||||||
|
"enable",
|
||||||
|
target,
|
||||||
|
features,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
updated = dict(payload)
|
||||||
|
updated["nanobot_features"] = self._with_channel_runtime_status(
|
||||||
|
features,
|
||||||
|
operations,
|
||||||
|
)
|
||||||
|
return updated
|
||||||
|
|
||||||
|
def allow_feature_package_install(self, request: SettingsRequest) -> bool:
|
||||||
|
if request.local_browser:
|
||||||
|
return True
|
||||||
|
try:
|
||||||
|
return bool(
|
||||||
|
self.settings.config.load().tools.webui_allow_remote_package_install
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("failed to load remote package install policy")
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _pairing_action(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
action: str,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
code = (query_first(request.query, "code") or "").strip()
|
||||||
|
if not code:
|
||||||
|
return SettingsRouteResult.failure(400, "Missing pairing code")
|
||||||
|
if action == "approve":
|
||||||
|
result = operations.approve_code(code)
|
||||||
|
if result is None:
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
404,
|
||||||
|
"Pairing code not found or expired",
|
||||||
|
)
|
||||||
|
channel, sender_id = result
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
pairing_payload(
|
||||||
|
operations.list_pending,
|
||||||
|
{
|
||||||
|
"ok": True,
|
||||||
|
"action": "approve",
|
||||||
|
"message": f"Approved {sender_id} for {channel}",
|
||||||
|
"channel": channel,
|
||||||
|
"sender_id": sender_id,
|
||||||
|
"code": code,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if not operations.deny_code(code):
|
||||||
|
return SettingsRouteResult.failure(
|
||||||
|
404,
|
||||||
|
"Pairing code not found or expired",
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
pairing_payload(
|
||||||
|
operations.list_pending,
|
||||||
|
{
|
||||||
|
"ok": True,
|
||||||
|
"action": "deny",
|
||||||
|
"message": f"Denied pairing code {code}",
|
||||||
|
"code": code,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _mcp_presets(
|
||||||
|
self,
|
||||||
|
request: SettingsRequest,
|
||||||
|
action: str | None,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
try:
|
||||||
|
payload = await operations.mcp_presets_action(
|
||||||
|
action,
|
||||||
|
request.query,
|
||||||
|
reload_mcp=operations.reload_mcp,
|
||||||
|
mcp_runtime_status=operations.mcp_runtime_status,
|
||||||
|
config=self.settings.config,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
status = getattr(exc, "status", 500)
|
||||||
|
message = getattr(exc, "message", str(exc))
|
||||||
|
if status >= 500:
|
||||||
|
self.logger.exception(
|
||||||
|
"MCP preset action '{}' failed",
|
||||||
|
action or "list",
|
||||||
|
)
|
||||||
|
return SettingsRouteResult.failure(status, message)
|
||||||
|
return SettingsRouteResult.success(
|
||||||
|
payload,
|
||||||
|
decorate_restart=action is not None,
|
||||||
|
restart_section="runtime" if action is not None else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _version_check(
|
||||||
|
self,
|
||||||
|
operations: SystemSettingsOperations,
|
||||||
|
) -> SettingsRouteResult:
|
||||||
|
try:
|
||||||
|
update_info = await asyncio.to_thread(operations.check_for_update)
|
||||||
|
except Exception:
|
||||||
|
self.logger.exception("version check failed")
|
||||||
|
return SettingsRouteResult.failure(500, "version check failed")
|
||||||
|
return SettingsRouteResult.success({"updateAvailable": update_info})
|
||||||
@@ -8,7 +8,9 @@ does not modify agent sessions.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
import math
|
||||||
import os
|
import os
|
||||||
|
import threading
|
||||||
import time
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
@@ -24,8 +26,11 @@ _MAX_MAP_ITEMS = 2_000
|
|||||||
_MAX_KEY_LEN = 512
|
_MAX_KEY_LEN = 512
|
||||||
_MAX_TITLE_LEN = 160
|
_MAX_TITLE_LEN = 160
|
||||||
_MAX_TAG_LEN = 40
|
_MAX_TAG_LEN = 40
|
||||||
|
_MAX_WORKBENCH_PANES = 4
|
||||||
_ALLOWED_DENSITIES = {"comfortable", "compact"}
|
_ALLOWED_DENSITIES = {"comfortable", "compact"}
|
||||||
_ALLOWED_SORTS = {"updated_desc", "created_desc", "title_asc", "manual"}
|
_ALLOWED_SORTS = {"updated_desc", "created_desc", "title_asc", "manual"}
|
||||||
|
_ALLOWED_WORKBENCH_LAYOUTS = {"columns", "rows", "grid", "bsp", "main-stack"}
|
||||||
|
_SIDEBAR_STATE_WRITE_LOCK = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
def webui_sidebar_state_path() -> Path:
|
def webui_sidebar_state_path() -> Path:
|
||||||
@@ -42,6 +47,7 @@ def default_webui_sidebar_state() -> dict[str, Any]:
|
|||||||
"project_name_overrides": {},
|
"project_name_overrides": {},
|
||||||
"tags_by_key": {},
|
"tags_by_key": {},
|
||||||
"collapsed_groups": {},
|
"collapsed_groups": {},
|
||||||
|
"workbench": {"version": 1, "tabs": {}},
|
||||||
"view": {
|
"view": {
|
||||||
"density": "comfortable",
|
"density": "comfortable",
|
||||||
"show_previews": False,
|
"show_previews": False,
|
||||||
@@ -76,6 +82,20 @@ def _clean_string_list(value: Any, *, max_len: int = _MAX_KEY_LEN) -> list[str]:
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_split_ratios(value: Any) -> list[float]:
|
||||||
|
if not isinstance(value, list):
|
||||||
|
return []
|
||||||
|
ratios: list[float] = []
|
||||||
|
for raw_ratio in cast(list[Any], value)[: _MAX_WORKBENCH_PANES - 1]:
|
||||||
|
if isinstance(raw_ratio, bool) or not isinstance(raw_ratio, (int, float)):
|
||||||
|
continue
|
||||||
|
ratio = float(raw_ratio)
|
||||||
|
if not math.isfinite(ratio):
|
||||||
|
continue
|
||||||
|
ratios.append(round(min(0.95, max(0.05, ratio)), 4))
|
||||||
|
return ratios
|
||||||
|
|
||||||
|
|
||||||
def _clean_bool_map(value: Any) -> dict[str, bool]:
|
def _clean_bool_map(value: Any) -> dict[str, bool]:
|
||||||
if not isinstance(value, dict):
|
if not isinstance(value, dict):
|
||||||
return {}
|
return {}
|
||||||
@@ -131,8 +151,56 @@ def _clean_view(value: Any) -> dict[str, Any]:
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_workbench(value: Any) -> dict[str, Any]:
|
||||||
|
if not isinstance(value, dict):
|
||||||
|
return {"version": 1, "tabs": {}}
|
||||||
|
workbench = cast(dict[str, Any], value)
|
||||||
|
if workbench.get("version") != 1:
|
||||||
|
return {"version": 1, "tabs": {}}
|
||||||
|
raw_tabs = workbench.get("tabs")
|
||||||
|
if not isinstance(raw_tabs, dict):
|
||||||
|
return {"version": 1, "tabs": {}}
|
||||||
|
|
||||||
|
tabs: dict[str, dict[str, Any]] = {}
|
||||||
|
claimed_panes: set[str] = set()
|
||||||
|
for raw_tab_key, raw_tab in list(cast(dict[Any, Any], raw_tabs).items())[:_MAX_MAP_ITEMS]:
|
||||||
|
tab_key = _clean_string(raw_tab_key)
|
||||||
|
if tab_key is None or not isinstance(raw_tab, dict):
|
||||||
|
continue
|
||||||
|
tab = cast(dict[str, Any], raw_tab)
|
||||||
|
pane_keys = [
|
||||||
|
key
|
||||||
|
for key in _clean_string_list(tab.get("paneKeys"))
|
||||||
|
if key not in claimed_panes
|
||||||
|
][:_MAX_WORKBENCH_PANES]
|
||||||
|
if not pane_keys:
|
||||||
|
continue
|
||||||
|
explicit = tab.get("explicit") is True
|
||||||
|
if not explicit and len(pane_keys) == 1:
|
||||||
|
continue
|
||||||
|
requested_layout_pane_keys = [
|
||||||
|
key for key in _clean_string_list(tab.get("layoutPaneKeys")) if key in pane_keys
|
||||||
|
]
|
||||||
|
layout_pane_keys = requested_layout_pane_keys + [
|
||||||
|
key for key in pane_keys if key not in requested_layout_pane_keys
|
||||||
|
]
|
||||||
|
claimed_panes.update(pane_keys)
|
||||||
|
raw_layout = tab.get("layout")
|
||||||
|
layout = raw_layout if raw_layout in _ALLOWED_WORKBENCH_LAYOUTS else "columns"
|
||||||
|
title = _clean_string(tab.get("title"), max_len=_MAX_TITLE_LEN)
|
||||||
|
tabs[tab_key] = {
|
||||||
|
"explicit": explicit,
|
||||||
|
"title": title,
|
||||||
|
"paneKeys": pane_keys,
|
||||||
|
"layoutPaneKeys": layout_pane_keys,
|
||||||
|
"layout": layout,
|
||||||
|
"splitRatios": _clean_split_ratios(tab.get("splitRatios")),
|
||||||
|
}
|
||||||
|
return {"version": 1, "tabs": tabs}
|
||||||
|
|
||||||
|
|
||||||
def normalize_webui_sidebar_state(raw: Any) -> dict[str, Any]:
|
def normalize_webui_sidebar_state(raw: Any) -> dict[str, Any]:
|
||||||
"""Return a schema-v1 sidebar state from any older/partial input."""
|
"""Return a validated canonical sidebar state."""
|
||||||
if not isinstance(raw, dict):
|
if not isinstance(raw, dict):
|
||||||
raw = {}
|
raw = {}
|
||||||
raw = cast(dict[str, Any], raw)
|
raw = cast(dict[str, Any], raw)
|
||||||
@@ -146,6 +214,7 @@ def normalize_webui_sidebar_state(raw: Any) -> dict[str, Any]:
|
|||||||
)
|
)
|
||||||
state["tags_by_key"] = _clean_tags_by_key(raw.get("tags_by_key"))
|
state["tags_by_key"] = _clean_tags_by_key(raw.get("tags_by_key"))
|
||||||
state["collapsed_groups"] = _clean_bool_map(raw.get("collapsed_groups"))
|
state["collapsed_groups"] = _clean_bool_map(raw.get("collapsed_groups"))
|
||||||
|
state["workbench"] = _clean_workbench(raw.get("workbench"))
|
||||||
state["view"] = _clean_view(raw.get("view"))
|
state["view"] = _clean_view(raw.get("view"))
|
||||||
updated_at = raw.get("updated_at")
|
updated_at = raw.get("updated_at")
|
||||||
state["updated_at"] = updated_at if isinstance(updated_at, str) else None
|
state["updated_at"] = updated_at if isinstance(updated_at, str) else None
|
||||||
@@ -169,6 +238,11 @@ def read_webui_sidebar_state() -> dict[str, Any]:
|
|||||||
|
|
||||||
|
|
||||||
def write_webui_sidebar_state(raw: dict[str, Any]) -> dict[str, Any]:
|
def write_webui_sidebar_state(raw: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
with _SIDEBAR_STATE_WRITE_LOCK:
|
||||||
|
return _write_webui_sidebar_state(raw)
|
||||||
|
|
||||||
|
|
||||||
|
def _write_webui_sidebar_state(raw: dict[str, Any]) -> dict[str, Any]:
|
||||||
state = normalize_webui_sidebar_state(raw)
|
state = normalize_webui_sidebar_state(raw)
|
||||||
state["updated_at"] = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
|
state["updated_at"] = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
|
||||||
encoded = json.dumps(
|
encoded = json.dumps(
|
||||||
|
|||||||
+207
-30
@@ -14,12 +14,13 @@ import json
|
|||||||
import mimetypes
|
import mimetypes
|
||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
from collections.abc import Callable
|
from collections.abc import Awaitable, Callable, Mapping
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any, cast
|
from typing import TYPE_CHECKING, Any, cast
|
||||||
from urllib.parse import unquote
|
from urllib.parse import quote, unquote, urlsplit, urlunsplit
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
from websockets.datastructures import Headers
|
||||||
from websockets.http11 import Request as WsRequest
|
from websockets.http11 import Request as WsRequest
|
||||||
from websockets.http11 import Response
|
from websockets.http11 import Response
|
||||||
|
|
||||||
@@ -118,7 +119,66 @@ from nanobot.webui.transcript import build_webui_thread_response
|
|||||||
from nanobot.webui.workspaces import WebUIWorkspaceController
|
from nanobot.webui.workspaces import WebUIWorkspaceController
|
||||||
|
|
||||||
_SLOW_WEBUI_HTTP_LOG_MS = 1_000
|
_SLOW_WEBUI_HTTP_LOG_MS = 1_000
|
||||||
_AUTOMATION_VALUES_HEADER = "X-Nanobot-Automation-Values"
|
_WEBUI_MUTATION_PAYLOAD_ATTR = "_nanobot_webui_mutation_payload"
|
||||||
|
_WEBUI_MUTATION_REQUEST_ATTR = "_nanobot_webui_mutation_request"
|
||||||
|
_NO_STORE_HEADERS = [("Cache-Control", "no-store")]
|
||||||
|
|
||||||
|
_WEBUI_MUTATION_PATHS = {
|
||||||
|
"automation.enable": "/api/webui/automations/enable",
|
||||||
|
"automation.disable": "/api/webui/automations/disable",
|
||||||
|
"automation.delete": "/api/webui/automations/delete",
|
||||||
|
"automation.run": "/api/webui/automations/run",
|
||||||
|
"automation.update": "/api/webui/automations/update",
|
||||||
|
"skill.install": "/api/webui/skills/install",
|
||||||
|
"skill.update": "/api/webui/skills/update",
|
||||||
|
"skill.delete": "/api/webui/skills/delete",
|
||||||
|
"sidebar.update": "/api/webui/sidebar-state/update",
|
||||||
|
"settings.agent.update": "/api/settings/update",
|
||||||
|
"settings.model_configuration.create": "/api/settings/model-configurations/create",
|
||||||
|
"settings.model_configuration.update": "/api/settings/model-configurations/update",
|
||||||
|
"settings.model_configuration.delete": "/api/settings/model-configurations/delete",
|
||||||
|
"settings.model_configuration.migrate": "/api/settings/model-configurations/migrate",
|
||||||
|
"settings.model_call_order.update": "/api/settings/model-call-order/update",
|
||||||
|
"settings.provider.update": "/api/settings/provider/update",
|
||||||
|
"settings.provider.create": "/api/settings/provider/create",
|
||||||
|
"settings.provider.oauth_login": "/api/settings/provider/oauth-login",
|
||||||
|
"settings.provider.oauth_complete": "/api/settings/provider/oauth-login/complete",
|
||||||
|
"settings.provider.oauth_logout": "/api/settings/provider/oauth-logout",
|
||||||
|
"settings.web_search.update": "/api/settings/web-search/update",
|
||||||
|
"settings.api_service.start": "/api/settings/api-service/start",
|
||||||
|
"settings.api_service.stop": "/api/settings/api-service/stop",
|
||||||
|
"settings.image_generation.update": "/api/settings/image-generation/update",
|
||||||
|
"settings.transcription.update": "/api/settings/transcription/update",
|
||||||
|
"settings.network_safety.update": "/api/settings/network-safety/update",
|
||||||
|
"settings.cli_app.install": "/api/settings/cli-apps/install",
|
||||||
|
"settings.cli_app.update": "/api/settings/cli-apps/update",
|
||||||
|
"settings.cli_app.uninstall": "/api/settings/cli-apps/uninstall",
|
||||||
|
"settings.cli_app.test": "/api/settings/cli-apps/test",
|
||||||
|
"settings.feature.enable": "/api/settings/nanobot-features/enable",
|
||||||
|
"settings.feature.disable": "/api/settings/nanobot-features/disable",
|
||||||
|
"settings.channel.validate": "/api/settings/channels/validate",
|
||||||
|
"settings.channel.configure": "/api/settings/channels/configure",
|
||||||
|
"settings.pairing.approve": "/api/settings/pairing/approve",
|
||||||
|
"settings.pairing.deny": "/api/settings/pairing/deny",
|
||||||
|
"settings.mcp.enable": "/api/settings/mcp-presets/enable",
|
||||||
|
"settings.mcp.disable": "/api/settings/mcp-presets/disable",
|
||||||
|
"settings.mcp.remove": "/api/settings/mcp-presets/remove",
|
||||||
|
"settings.mcp.test": "/api/settings/mcp-presets/test",
|
||||||
|
"settings.mcp.reconnect": "/api/settings/mcp-presets/reconnect",
|
||||||
|
"settings.mcp.custom": "/api/settings/mcp-presets/custom",
|
||||||
|
"settings.mcp.import": "/api/settings/mcp-presets/import",
|
||||||
|
"settings.mcp.import_cursor": "/api/settings/mcp-presets/import-cursor",
|
||||||
|
"settings.mcp.tools": "/api/settings/mcp-presets/tools",
|
||||||
|
"settings.mcp.oauth_start": "/api/settings/mcp-oauth/start",
|
||||||
|
"settings.mcp.oauth_complete": "/api/settings/mcp-oauth/complete",
|
||||||
|
"settings.mcp.oauth_cancel": "/api/settings/mcp-oauth/cancel",
|
||||||
|
}
|
||||||
|
|
||||||
|
_WEBUI_CHANNEL_CONNECT_ACTIONS = {
|
||||||
|
"settings.channel.connect.start": "start",
|
||||||
|
"settings.channel.connect.poll": "poll",
|
||||||
|
"settings.channel.connect.cancel": "cancel",
|
||||||
|
}
|
||||||
|
|
||||||
# Fix for #5190: On Windows, mimetypes.guess_type() reads the registry key
|
# Fix for #5190: On Windows, mimetypes.guess_type() reads the registry key
|
||||||
# HKEY_CLASSES_ROOT\.js\Content Type, which is commonly set to 'text/plain'
|
# HKEY_CLASSES_ROOT\.js\Content Type, which is commonly set to 'text/plain'
|
||||||
@@ -150,6 +210,7 @@ if TYPE_CHECKING:
|
|||||||
from nanobot.cron.service import CronService
|
from nanobot.cron.service import CronService
|
||||||
from nanobot.session.manager import SessionManager
|
from nanobot.session.manager import SessionManager
|
||||||
from nanobot.triggers.local_store import LocalTriggerStore
|
from nanobot.triggers.local_store import LocalTriggerStore
|
||||||
|
from nanobot.webui.settings_services import WebUISettingsServices
|
||||||
|
|
||||||
def _decode_api_key(raw_key: str) -> str | None:
|
def _decode_api_key(raw_key: str) -> str | None:
|
||||||
key = unquote(raw_key)
|
key = unquote(raw_key)
|
||||||
@@ -159,6 +220,33 @@ def _decode_api_key(raw_key: str) -> str | None:
|
|||||||
return key
|
return key
|
||||||
|
|
||||||
|
|
||||||
|
def _mutation_payload(request: WsRequest) -> dict[str, Any] | None:
|
||||||
|
payload = getattr(request, _WEBUI_MUTATION_PAYLOAD_ATTR, None)
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
return None
|
||||||
|
return cast(dict[str, Any], payload)
|
||||||
|
|
||||||
|
|
||||||
|
def _request_query(request: WsRequest) -> dict[str, list[str]]:
|
||||||
|
payload = _mutation_payload(request)
|
||||||
|
if payload is None:
|
||||||
|
return _parse_query(request.path)
|
||||||
|
query: dict[str, list[str]] = {}
|
||||||
|
for key, value in payload.items():
|
||||||
|
if not key:
|
||||||
|
continue
|
||||||
|
if isinstance(value, bool):
|
||||||
|
text = "true" if value else "false"
|
||||||
|
elif value is None:
|
||||||
|
text = ""
|
||||||
|
elif isinstance(value, (dict, list)):
|
||||||
|
text = json.dumps(value, ensure_ascii=False, separators=(",", ":"))
|
||||||
|
else:
|
||||||
|
text = str(value)
|
||||||
|
query[key] = [text]
|
||||||
|
return query
|
||||||
|
|
||||||
|
|
||||||
def _default_model_name_from_config() -> str | None:
|
def _default_model_name_from_config() -> str | None:
|
||||||
try:
|
try:
|
||||||
from nanobot.config.loader import load_config
|
from nanobot.config.loader import load_config
|
||||||
@@ -211,6 +299,7 @@ class GatewayHTTPHandler:
|
|||||||
media: WebUIMediaGateway,
|
media: WebUIMediaGateway,
|
||||||
ingress: WebUIIngressPolicy,
|
ingress: WebUIIngressPolicy,
|
||||||
workspaces: WebUIWorkspaceController,
|
workspaces: WebUIWorkspaceController,
|
||||||
|
settings: WebUISettingsServices,
|
||||||
skills_workspace_path: Path,
|
skills_workspace_path: Path,
|
||||||
disabled_skills: set[str] | None = None,
|
disabled_skills: set[str] | None = None,
|
||||||
cron_service: CronService | None = None,
|
cron_service: CronService | None = None,
|
||||||
@@ -219,6 +308,8 @@ class GatewayHTTPHandler:
|
|||||||
local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
||||||
channel_feature_action: Callable[..., Any] | None = None,
|
channel_feature_action: Callable[..., Any] | None = None,
|
||||||
channel_runtime_status: Callable[[], dict[str, Any]] | None = None,
|
channel_runtime_status: Callable[[], dict[str, Any]] | None = None,
|
||||||
|
mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
|
||||||
|
mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
|
||||||
skill_state_action: Callable[[set[str]], None] | None = None,
|
skill_state_action: Callable[[set[str]], None] | None = None,
|
||||||
log: Any = logger,
|
log: Any = logger,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -231,6 +322,7 @@ class GatewayHTTPHandler:
|
|||||||
self.media = media
|
self.media = media
|
||||||
self.ingress = ingress
|
self.ingress = ingress
|
||||||
self.workspaces = workspaces
|
self.workspaces = workspaces
|
||||||
|
self.settings = settings
|
||||||
self.skills_workspace_path = skills_workspace_path
|
self.skills_workspace_path = skills_workspace_path
|
||||||
self.disabled_skills: set[str] = (
|
self.disabled_skills: set[str] = (
|
||||||
disabled_skills if disabled_skills is not None else set()
|
disabled_skills if disabled_skills is not None else set()
|
||||||
@@ -249,6 +341,7 @@ class GatewayHTTPHandler:
|
|||||||
|
|
||||||
self._capabilities = _rc(runtime_surface, runtime_capabilities_overrides or {})
|
self._capabilities = _rc(runtime_surface, runtime_capabilities_overrides or {})
|
||||||
self.settings_routes = WebUISettingsRouter(
|
self.settings_routes = WebUISettingsRouter(
|
||||||
|
settings=settings,
|
||||||
bus=bus,
|
bus=bus,
|
||||||
logger=self._log,
|
logger=self._log,
|
||||||
check_api_token=self.check_api_token,
|
check_api_token=self.check_api_token,
|
||||||
@@ -259,6 +352,9 @@ class GatewayHTTPHandler:
|
|||||||
runtime_capabilities=self._capabilities,
|
runtime_capabilities=self._capabilities,
|
||||||
channel_feature_action=channel_feature_action,
|
channel_feature_action=channel_feature_action,
|
||||||
channel_runtime_status=channel_runtime_status,
|
channel_runtime_status=channel_runtime_status,
|
||||||
|
mcp_runtime_status=mcp_runtime_status,
|
||||||
|
mcp_reload=mcp_reload,
|
||||||
|
mcp_oauth_redirect_uri=self._mcp_oauth_redirect_uri,
|
||||||
)
|
)
|
||||||
|
|
||||||
def workspace_controls_available(self, connection: Any) -> bool:
|
def workspace_controls_available(self, connection: Any) -> bool:
|
||||||
@@ -285,11 +381,86 @@ class GatewayHTTPHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
if self._is_webui_mutation_path(got):
|
||||||
|
return _http_error(
|
||||||
|
405,
|
||||||
|
"WebUI mutations require an authenticated WebSocket",
|
||||||
|
)
|
||||||
response = await self._dispatch_resolved(connection, request, got)
|
response = await self._dispatch_resolved(connection, request, got)
|
||||||
return response
|
return response
|
||||||
finally:
|
finally:
|
||||||
self._log_slow_http(got, response, started)
|
self._log_slow_http(got, response, started)
|
||||||
|
|
||||||
|
async def dispatch_webui_mutation(
|
||||||
|
self,
|
||||||
|
connection: Any,
|
||||||
|
action: str,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
) -> Response:
|
||||||
|
"""Run one explicitly allowlisted mutation for an authenticated WebUI socket."""
|
||||||
|
path = self._webui_mutation_path(action, payload)
|
||||||
|
if isinstance(path, Response):
|
||||||
|
return path
|
||||||
|
|
||||||
|
source_request = getattr(connection, "request", None)
|
||||||
|
source_headers = getattr(source_request, "headers", None)
|
||||||
|
if source_headers is None:
|
||||||
|
headers = Headers()
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
headers = Headers(source_headers.raw_items())
|
||||||
|
except (AttributeError, TypeError):
|
||||||
|
try:
|
||||||
|
headers = Headers(source_headers)
|
||||||
|
except TypeError:
|
||||||
|
headers = Headers()
|
||||||
|
request = WsRequest(path, headers)
|
||||||
|
setattr(request, "_nanobot_trusted_proxy_authenticated", True)
|
||||||
|
setattr(request, _WEBUI_MUTATION_REQUEST_ATTR, True)
|
||||||
|
setattr(request, _WEBUI_MUTATION_PAYLOAD_ATTR, dict(payload))
|
||||||
|
response = await self._dispatch_resolved(connection, request, path)
|
||||||
|
if isinstance(response, Response):
|
||||||
|
return response
|
||||||
|
return _http_error(404, "WebUI mutation action not found")
|
||||||
|
|
||||||
|
def _is_webui_mutation_path(self, path: str) -> bool:
|
||||||
|
if self.settings_routes.is_mutation_path(path):
|
||||||
|
return True
|
||||||
|
if re.match(r"^/api/sessions/[^/]+/delete$", path):
|
||||||
|
return True
|
||||||
|
if re.match(r"^/api/webui/automations/(enable|disable|delete|run|update)$", path):
|
||||||
|
return True
|
||||||
|
return path in {
|
||||||
|
"/api/webui/skills/install",
|
||||||
|
"/api/webui/skills/update",
|
||||||
|
"/api/webui/skills/delete",
|
||||||
|
"/api/webui/sidebar-state/update",
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _webui_mutation_path(
|
||||||
|
action: str,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
) -> str | Response:
|
||||||
|
path = _WEBUI_MUTATION_PATHS.get(action)
|
||||||
|
if path is not None:
|
||||||
|
return path
|
||||||
|
if action == "session.delete":
|
||||||
|
key = payload.get("key")
|
||||||
|
if not isinstance(key, str) or not key.strip():
|
||||||
|
return _http_error(400, "missing session key")
|
||||||
|
return f"/api/sessions/{quote(key, safe='')}/delete"
|
||||||
|
connect_action = _WEBUI_CHANNEL_CONNECT_ACTIONS.get(action)
|
||||||
|
if connect_action is not None:
|
||||||
|
channel = payload.get("channel")
|
||||||
|
if not isinstance(channel, str) or re.fullmatch(
|
||||||
|
r"[A-Za-z0-9_-]{1,64}",
|
||||||
|
channel,
|
||||||
|
) is None:
|
||||||
|
return _http_error(400, "invalid channel name")
|
||||||
|
return f"/api/settings/channels/{channel}/connect/{connect_action}"
|
||||||
|
return _http_error(404, "unknown WebUI mutation action")
|
||||||
|
|
||||||
async def _dispatch_resolved(
|
async def _dispatch_resolved(
|
||||||
self,
|
self,
|
||||||
connection: Any,
|
connection: Any,
|
||||||
@@ -377,9 +548,16 @@ class GatewayHTTPHandler:
|
|||||||
"too many outstanding issued tokens ({}), rejecting issuance",
|
"too many outstanding issued tokens ({}), rejecting issuance",
|
||||||
len(self.tokens.issued_tokens),
|
len(self.tokens.issued_tokens),
|
||||||
)
|
)
|
||||||
return _http_json_response({"error": "too many outstanding tokens"}, status=429)
|
return _http_json_response(
|
||||||
|
{"error": "too many outstanding tokens"},
|
||||||
|
status=429,
|
||||||
|
extra_headers=_NO_STORE_HEADERS,
|
||||||
|
)
|
||||||
token_value = self.tokens.issue_token(self.config.token_ttl_s)
|
token_value = self.tokens.issue_token(self.config.token_ttl_s)
|
||||||
return _http_json_response(token_response_payload(token_value, self.config.token_ttl_s))
|
return _http_json_response(
|
||||||
|
token_response_payload(token_value, self.config.token_ttl_s),
|
||||||
|
extra_headers=_NO_STORE_HEADERS,
|
||||||
|
)
|
||||||
|
|
||||||
# -- Bootstrap ----------------------------------------------------------
|
# -- Bootstrap ----------------------------------------------------------
|
||||||
|
|
||||||
@@ -409,7 +587,7 @@ class GatewayHTTPHandler:
|
|||||||
"runtime_surface": self._runtime_surface,
|
"runtime_surface": self._runtime_surface,
|
||||||
"runtime_capabilities": self._capabilities,
|
"runtime_capabilities": self._capabilities,
|
||||||
}
|
}
|
||||||
return _http_json_response(payload)
|
return _http_json_response(payload, extra_headers=_NO_STORE_HEADERS)
|
||||||
|
|
||||||
api_token_allowed = bool(secret) or is_local_browser
|
api_token_allowed = bool(secret) or is_local_browser
|
||||||
if not self.tokens.can_issue(include_api_token=api_token_allowed):
|
if not self.tokens.can_issue(include_api_token=api_token_allowed):
|
||||||
@@ -417,6 +595,7 @@ class GatewayHTTPHandler:
|
|||||||
json.dumps({"error": "too many outstanding tokens"}).encode("utf-8"),
|
json.dumps({"error": "too many outstanding tokens"}).encode("utf-8"),
|
||||||
status=429,
|
status=429,
|
||||||
content_type="application/json; charset=utf-8",
|
content_type="application/json; charset=utf-8",
|
||||||
|
extra_headers=_NO_STORE_HEADERS,
|
||||||
)
|
)
|
||||||
token = self.tokens.issue_token(self.config.token_ttl_s, audience="webui")
|
token = self.tokens.issue_token(self.config.token_ttl_s, audience="webui")
|
||||||
api_token = (
|
api_token = (
|
||||||
@@ -441,7 +620,7 @@ class GatewayHTTPHandler:
|
|||||||
}
|
}
|
||||||
if api_token is not None:
|
if api_token is not None:
|
||||||
payload["api_token"] = api_token
|
payload["api_token"] = api_token
|
||||||
return _http_json_response(payload)
|
return _http_json_response(payload, extra_headers=_NO_STORE_HEADERS)
|
||||||
|
|
||||||
def _bootstrap_ws_url(self, request: Any) -> str:
|
def _bootstrap_ws_url(self, request: Any) -> str:
|
||||||
headers = getattr(request, "headers", {}) or {}
|
headers = getattr(request, "headers", {}) or {}
|
||||||
@@ -457,6 +636,14 @@ class GatewayHTTPHandler:
|
|||||||
expected_path = _normalize_config_path(self.config.path)
|
expected_path = _normalize_config_path(self.config.path)
|
||||||
return f"{scheme}://{host}{expected_path}"
|
return f"{scheme}://{host}{expected_path}"
|
||||||
|
|
||||||
|
def _mcp_oauth_redirect_uri(self, request: WsRequest) -> str:
|
||||||
|
"""Derive the browser callback from the same public origin as WebSocket bootstrap."""
|
||||||
|
from nanobot.agent.tools.mcp_oauth import MCP_OAUTH_CALLBACK_PATH
|
||||||
|
|
||||||
|
public_ws_url = urlsplit(self._bootstrap_ws_url(request))
|
||||||
|
scheme = "https" if public_ws_url.scheme == "wss" else "http"
|
||||||
|
return urlunsplit((scheme, public_ws_url.netloc, MCP_OAUTH_CALLBACK_PATH, "", ""))
|
||||||
|
|
||||||
# -- Session routes -----------------------------------------------------
|
# -- Session routes -----------------------------------------------------
|
||||||
|
|
||||||
async def _dispatch_session_routes(self, request: WsRequest, got: str) -> Response | None:
|
async def _dispatch_session_routes(self, request: WsRequest, got: str) -> Response | None:
|
||||||
@@ -646,7 +833,7 @@ class GatewayHTTPHandler:
|
|||||||
return _http_error(400, "invalid session key")
|
return _http_error(400, "invalid session key")
|
||||||
if not _is_websocket_channel_session_key(decoded_key):
|
if not _is_websocket_channel_session_key(decoded_key):
|
||||||
return _http_error(404, "session not found")
|
return _http_error(404, "session not found")
|
||||||
query = _parse_query(request.path)
|
query = _request_query(request)
|
||||||
delete_automations = (_query_first(query, "delete_automations") or "").lower()
|
delete_automations = (_query_first(query, "delete_automations") or "").lower()
|
||||||
automation_jobs = session_automation_jobs(
|
automation_jobs = session_automation_jobs(
|
||||||
self.cron_service,
|
self.cron_service,
|
||||||
@@ -742,7 +929,7 @@ class GatewayHTTPHandler:
|
|||||||
if self.cron_service is None and self.local_trigger_store is None:
|
if self.cron_service is None and self.local_trigger_store is None:
|
||||||
return _http_error(503, "automation service unavailable")
|
return _http_error(503, "automation service unavailable")
|
||||||
|
|
||||||
query = _parse_query(request.path)
|
query = _request_query(request)
|
||||||
job_id = (_query_first(query, "id") or _query_first(query, "job_id") or "").strip()
|
job_id = (_query_first(query, "id") or _query_first(query, "job_id") or "").strip()
|
||||||
if not job_id:
|
if not job_id:
|
||||||
return _http_error(400, "missing automation id")
|
return _http_error(400, "missing automation id")
|
||||||
@@ -974,7 +1161,7 @@ class GatewayHTTPHandler:
|
|||||||
if self._skill_install_lock.locked():
|
if self._skill_install_lock.locked():
|
||||||
return _http_error(409, "another skill installation is already in progress")
|
return _http_error(409, "another skill installation is already in progress")
|
||||||
|
|
||||||
query = _parse_query(request.path)
|
query = _request_query(request)
|
||||||
provider = _query_first(query, "provider") or "skills_sh"
|
provider = _query_first(query, "provider") or "skills_sh"
|
||||||
source = _query_first(query, "source") or ""
|
source = _query_first(query, "source") or ""
|
||||||
skill_id = _query_first(query, "skill") or ""
|
skill_id = _query_first(query, "skill") or ""
|
||||||
@@ -1015,7 +1202,7 @@ class GatewayHTTPHandler:
|
|||||||
def _handle_webui_skill_update(self, request: WsRequest) -> Response:
|
def _handle_webui_skill_update(self, request: WsRequest) -> Response:
|
||||||
if not self.check_api_token(request):
|
if not self.check_api_token(request):
|
||||||
return _http_error(401, "Unauthorized")
|
return _http_error(401, "Unauthorized")
|
||||||
query = _parse_query(request.path)
|
query = _request_query(request)
|
||||||
name = _query_first(query, "name") or ""
|
name = _query_first(query, "name") or ""
|
||||||
raw_enabled = (_query_first(query, "enabled") or "").lower()
|
raw_enabled = (_query_first(query, "enabled") or "").lower()
|
||||||
if raw_enabled not in {"true", "false"}:
|
if raw_enabled not in {"true", "false"}:
|
||||||
@@ -1047,7 +1234,7 @@ class GatewayHTTPHandler:
|
|||||||
return _http_error(401, "Unauthorized")
|
return _http_error(401, "Unauthorized")
|
||||||
if not _is_local_browser_request(connection, request.headers):
|
if not _is_local_browser_request(connection, request.headers):
|
||||||
return _http_error(403, "remote skill deletion is disabled")
|
return _http_error(403, "remote skill deletion is disabled")
|
||||||
name = _query_first(_parse_query(request.path), "name") or ""
|
name = _query_first(_request_query(request), "name") or ""
|
||||||
try:
|
try:
|
||||||
action = delete_webui_skill(
|
action = delete_webui_skill(
|
||||||
self.skills_workspace_path,
|
self.skills_workspace_path,
|
||||||
@@ -1094,18 +1281,14 @@ class GatewayHTTPHandler:
|
|||||||
def _handle_webui_sidebar_state_update(self, request: WsRequest) -> Response:
|
def _handle_webui_sidebar_state_update(self, request: WsRequest) -> Response:
|
||||||
if not self.check_api_token(request):
|
if not self.check_api_token(request):
|
||||||
return _http_error(401, "Unauthorized")
|
return _http_error(401, "Unauthorized")
|
||||||
query = _parse_query(request.path)
|
payload = _mutation_payload(request)
|
||||||
raw_state = _query_first(query, "state")
|
state_value = payload.get("state") if payload is not None else None
|
||||||
if raw_state is None:
|
if state_value is None:
|
||||||
return _http_error(400, "missing state")
|
return _http_error(400, "missing state")
|
||||||
try:
|
if not isinstance(state_value, dict):
|
||||||
decoded = json.loads(raw_state)
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
return _http_error(400, "state must be JSON")
|
|
||||||
if not isinstance(decoded, dict):
|
|
||||||
return _http_error(400, "state must be an object")
|
return _http_error(400, "state must be an object")
|
||||||
try:
|
try:
|
||||||
state = write_webui_sidebar_state(cast(dict[str, Any], decoded))
|
state = write_webui_sidebar_state(cast(dict[str, Any], state_value))
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
return _http_error(400, str(e))
|
return _http_error(400, str(e))
|
||||||
except OSError:
|
except OSError:
|
||||||
@@ -1174,16 +1357,10 @@ class GatewayHTTPHandler:
|
|||||||
|
|
||||||
|
|
||||||
def _automation_values_from_request(request: WsRequest) -> dict[str, Any] | None:
|
def _automation_values_from_request(request: WsRequest) -> dict[str, Any] | None:
|
||||||
raw = _case_insensitive_header(request.headers, _AUTOMATION_VALUES_HEADER)
|
payload = _mutation_payload(request)
|
||||||
if not raw:
|
if payload is None or "values" not in payload:
|
||||||
return {}
|
return {}
|
||||||
try:
|
values = payload.get("values")
|
||||||
values = json.loads(raw)
|
|
||||||
except Exception:
|
|
||||||
try:
|
|
||||||
values = json.loads(unquote(raw))
|
|
||||||
except Exception:
|
|
||||||
return None
|
|
||||||
return cast(dict[str, Any], values) if isinstance(values, dict) else None
|
return cast(dict[str, Any], values) if isinstance(values, dict) else None
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ dependencies = [
|
|||||||
"readability-lxml>=0.8.4,<1.0.0",
|
"readability-lxml>=0.8.4,<1.0.0",
|
||||||
"lxml-html-clean>=0.4.0,<1.0.0",
|
"lxml-html-clean>=0.4.0,<1.0.0",
|
||||||
"rich>=14.0.0,<15.0.0",
|
"rich>=14.0.0,<15.0.0",
|
||||||
|
"qrcode[pil]>=8.0",
|
||||||
"croniter>=6.0.0,<7.0.0",
|
"croniter>=6.0.0,<7.0.0",
|
||||||
"prompt-toolkit>=3.0.50,<4.0.0",
|
"prompt-toolkit>=3.0.50,<4.0.0",
|
||||||
"questionary>=2.0.0,<3.0.0",
|
"questionary>=2.0.0,<3.0.0",
|
||||||
|
|||||||
@@ -46,7 +46,6 @@ def make_loop(
|
|||||||
context_window_tokens: int = 128_000,
|
context_window_tokens: int = 128_000,
|
||||||
session_ttl_minutes: int = 0,
|
session_ttl_minutes: int = 0,
|
||||||
unified_session: bool = False,
|
unified_session: bool = False,
|
||||||
mcp_servers: dict | None = None,
|
|
||||||
tools_config=None,
|
tools_config=None,
|
||||||
model_presets: dict | None = None,
|
model_presets: dict | None = None,
|
||||||
hooks: list | None = None,
|
hooks: list | None = None,
|
||||||
@@ -72,8 +71,6 @@ def make_loop(
|
|||||||
session_ttl_minutes=session_ttl_minutes,
|
session_ttl_minutes=session_ttl_minutes,
|
||||||
unified_session=unified_session,
|
unified_session=unified_session,
|
||||||
)
|
)
|
||||||
if mcp_servers is not None:
|
|
||||||
kwargs["mcp_servers"] = mcp_servers
|
|
||||||
if tools_config is not None:
|
if tools_config is not None:
|
||||||
kwargs["tools_config"] = tools_config
|
kwargs["tools_config"] = tools_config
|
||||||
if model_presets is not None:
|
if model_presets is not None:
|
||||||
|
|||||||
@@ -0,0 +1,370 @@
|
|||||||
|
import json
|
||||||
|
import shutil
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.agent import plugins as agent_plugins
|
||||||
|
from nanobot.agent.plugins import (
|
||||||
|
AGENT_PLUGIN_MCP_SCHEMA,
|
||||||
|
AGENT_PLUGIN_SCHEMA,
|
||||||
|
agent_plugin_mcp_servers,
|
||||||
|
discover_agent_plugins,
|
||||||
|
enabled_agent_plugin_skill_dirs,
|
||||||
|
enabled_agent_plugin_skills,
|
||||||
|
set_agent_plugin_enabled,
|
||||||
|
)
|
||||||
|
from nanobot.agent.skills import SkillsLoader
|
||||||
|
from nanobot.agent.tools.context import ToolContext
|
||||||
|
from nanobot.agent.tools.filesystem import ReadFileTool, WriteFileTool
|
||||||
|
from nanobot.config.schema import ToolsConfig
|
||||||
|
from nanobot.security.workspace_access import (
|
||||||
|
bind_workspace_scope,
|
||||||
|
reset_workspace_scope,
|
||||||
|
validate_workspace_scope_payload,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _isolate_plugin_state(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
monkeypatch.setattr(
|
||||||
|
agent_plugins, "get_config_path", lambda: tmp_path / "config" / "config.json"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _write_json(path: Path, value: object) -> None:
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
path.write_text(json.dumps(value), encoding="utf-8")
|
||||||
|
|
||||||
|
|
||||||
|
def _manifest(name: str, **fields: object) -> dict[str, object]:
|
||||||
|
return {"$schema": AGENT_PLUGIN_SCHEMA, "name": name, **fields}
|
||||||
|
|
||||||
|
|
||||||
|
def _plugin(workspace: Path, name: str = "demo", **fields: object) -> Path:
|
||||||
|
root = workspace / "plugins" / name
|
||||||
|
_write_json(root / "plugin.json", _manifest(name, **fields))
|
||||||
|
return root
|
||||||
|
|
||||||
|
|
||||||
|
def _skill(root: Path, name: str, frontmatter: str | None = None, body: str = "") -> Path:
|
||||||
|
path = root / name
|
||||||
|
path.mkdir(parents=True)
|
||||||
|
metadata = frontmatter or f"name: {name}\ndescription: Plugin skill."
|
||||||
|
(path / "SKILL.md").write_text(f"---\n{metadata}\n---\n\n{body}\n", encoding="utf-8")
|
||||||
|
return path
|
||||||
|
|
||||||
|
|
||||||
|
def _loaded_skills(workspace: Path) -> list[str]:
|
||||||
|
return [name for name, _ in enabled_agent_plugin_skills(workspace)]
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_skill_lifecycle_and_precedence(tmp_path: Path) -> None:
|
||||||
|
plugin = _plugin(tmp_path)
|
||||||
|
_skill(
|
||||||
|
plugin / "skills",
|
||||||
|
"shared",
|
||||||
|
"name: shared\ndescription: Plugin version.\nalways: true",
|
||||||
|
"Plugin body.",
|
||||||
|
)
|
||||||
|
_skill(tmp_path / "builtin", "shared", body="Built-in body.")
|
||||||
|
workspace_skill = _skill(
|
||||||
|
tmp_path / "skills", "shared", "name: shared\ndescription: Workspace version."
|
||||||
|
)
|
||||||
|
loader = SkillsLoader(tmp_path, builtin_skills_dir=tmp_path / "builtin")
|
||||||
|
|
||||||
|
assert [entry["source"] for entry in loader.list_skills()] == ["workspace"]
|
||||||
|
assert "Workspace version" in (loader.load_skill("shared") or "")
|
||||||
|
set_agent_plugin_enabled(tmp_path, "demo", True)
|
||||||
|
assert [entry["source"] for entry in loader.list_skills()] == ["workspace"]
|
||||||
|
|
||||||
|
shutil.rmtree(workspace_skill)
|
||||||
|
assert [entry["source"] for entry in loader.list_skills()] == ["plugin"]
|
||||||
|
assert loader.get_explicitly_invoked_skills("Use $shared") == ["shared"]
|
||||||
|
assert loader.get_always_skills() == ["shared"]
|
||||||
|
assert "Plugin body" in (loader.load_skill("shared") or "")
|
||||||
|
assert "`demo/skills/shared/SKILL.md`" in loader.build_skills_summary()
|
||||||
|
|
||||||
|
set_agent_plugin_enabled(tmp_path, "demo", False)
|
||||||
|
assert [entry["source"] for entry in loader.list_skills()] == ["builtin"]
|
||||||
|
assert "Built-in body" in (loader.load_skill("shared") or "")
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_skills_are_direct_valid_and_contained(tmp_path: Path) -> None:
|
||||||
|
plugin = _plugin(tmp_path)
|
||||||
|
skills = plugin / "skills"
|
||||||
|
_skill(skills, "direct")
|
||||||
|
_skill(skills / "group", "nested")
|
||||||
|
for name, frontmatter in (
|
||||||
|
("wrong-directory", "name: another\ndescription: Mismatch."),
|
||||||
|
("missing-description", "name: missing-description"),
|
||||||
|
("Bad-Name", "name: Bad-Name\ndescription: Invalid name."),
|
||||||
|
):
|
||||||
|
_skill(skills, name, frontmatter)
|
||||||
|
outside = _skill(tmp_path / "outside", "escaped")
|
||||||
|
try:
|
||||||
|
(skills / "escaped").symlink_to(outside, target_is_directory=True)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
set_agent_plugin_enabled(tmp_path, "demo", True)
|
||||||
|
assert _loaded_skills(tmp_path) == ["direct"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("manifest", "valid"),
|
||||||
|
[
|
||||||
|
({"$schema": "https://agent-plugins.org/schemas/2.0.0/plugin.schema.json", "name": "demo"}, False),
|
||||||
|
(_manifest("Bad-Name"), False),
|
||||||
|
(_manifest("demo", futureField=True, extensions="invalid but non-fatal"), True),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_plugin_manifest_boundary(tmp_path: Path, manifest: object, valid: bool) -> None:
|
||||||
|
_write_json(tmp_path / "plugins" / "candidate" / "plugin.json", manifest)
|
||||||
|
assert bool(discover_agent_plugins(tmp_path)) is valid
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_logo_is_validated_and_contained(tmp_path: Path) -> None:
|
||||||
|
extension = {"extensions": {"dev.nanobot": {"logo": "./assets/icon.png"}}}
|
||||||
|
plugin = _plugin(tmp_path, "demo", **extension)
|
||||||
|
icon = plugin / "assets" / "icon.png"
|
||||||
|
icon.parent.mkdir()
|
||||||
|
icon.write_bytes(b"\x89PNG\r\n\x1a\nlogo")
|
||||||
|
escaped = _plugin(tmp_path, "escaped", **extension)
|
||||||
|
(escaped / "assets").mkdir()
|
||||||
|
try:
|
||||||
|
(escaped / "assets" / "icon.png").symlink_to(icon)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert {plugin.name: plugin.logo for plugin in discover_agent_plugins(tmp_path)} == {
|
||||||
|
"demo": "data:image/png;base64,iVBORw0KGgpsb2dv",
|
||||||
|
"escaped": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_mcp_requires_explicit_enable(tmp_path: Path) -> None:
|
||||||
|
plugin = _plugin(tmp_path, "desktop")
|
||||||
|
executable = plugin / "bin" / "server"
|
||||||
|
executable.parent.mkdir()
|
||||||
|
executable.write_text("#!/bin/sh\n", encoding="utf-8")
|
||||||
|
_write_json(
|
||||||
|
plugin / "mcp.json",
|
||||||
|
{
|
||||||
|
"$schema": AGENT_PLUGIN_MCP_SCHEMA,
|
||||||
|
"mcpServers": {
|
||||||
|
"desktop": {
|
||||||
|
"type": "stdio",
|
||||||
|
"command": "./bin/server",
|
||||||
|
"args": ["--data", "${PLUGIN_DATA}/state"],
|
||||||
|
"cwd": "${PLUGIN_ROOT}",
|
||||||
|
},
|
||||||
|
"public-http": {"type": "streamable-http", "url": "http://example.com/mcp"},
|
||||||
|
"escape": {"type": "stdio", "command": "../outside"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert agent_plugin_mcp_servers(tmp_path) == {}
|
||||||
|
set_agent_plugin_enabled(tmp_path, "desktop", True)
|
||||||
|
server = agent_plugin_mcp_servers(tmp_path)["desktop"]
|
||||||
|
assert (server.command, server.cwd, server.env["PLUGIN_ROOT"]) == (
|
||||||
|
str(executable),
|
||||||
|
str(plugin),
|
||||||
|
str(plugin),
|
||||||
|
)
|
||||||
|
assert server.args[1].endswith("/state")
|
||||||
|
set_agent_plugin_enabled(tmp_path, "desktop", False)
|
||||||
|
assert agent_plugin_mcp_servers(tmp_path) == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_mcp_namespaces_cannot_shadow_plugin_identities(tmp_path: Path) -> None:
|
||||||
|
single = _plugin(tmp_path, "foo-bar")
|
||||||
|
multi = _plugin(tmp_path, "foo")
|
||||||
|
for root, servers in (
|
||||||
|
(single, {"main": {"type": "stdio", "command": "echo", "args": ["single"]}}),
|
||||||
|
(
|
||||||
|
multi,
|
||||||
|
{
|
||||||
|
"bar": {"type": "stdio", "command": "echo", "args": ["multi"]},
|
||||||
|
"other": {"type": "stdio", "command": "echo"},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
):
|
||||||
|
_write_json(
|
||||||
|
root / "mcp.json",
|
||||||
|
{"$schema": AGENT_PLUGIN_MCP_SCHEMA, "mcpServers": servers},
|
||||||
|
)
|
||||||
|
set_agent_plugin_enabled(tmp_path, "foo-bar", True)
|
||||||
|
set_agent_plugin_enabled(tmp_path, "foo", True)
|
||||||
|
|
||||||
|
servers = agent_plugin_mcp_servers(tmp_path)
|
||||||
|
|
||||||
|
assert set(servers) == {"foo-bar", "foo--bar", "foo--other"}
|
||||||
|
assert servers["foo-bar"].args == ["single"]
|
||||||
|
assert servers["foo--bar"].args == ["multi"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_restricted_project_can_read_only_enabled_plugin_skill(
|
||||||
|
tmp_path: Path, monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
agent_workspace = tmp_path / "agent"
|
||||||
|
project = tmp_path / "project"
|
||||||
|
project.mkdir()
|
||||||
|
plugin = _plugin(agent_workspace)
|
||||||
|
skill = _skill(plugin / "skills", "demo-skill")
|
||||||
|
resource = skill / "reference.md"
|
||||||
|
resource.write_text("plugin reference", encoding="utf-8")
|
||||||
|
ctx = ToolContext(
|
||||||
|
config=ToolsConfig(restrict_to_workspace=True),
|
||||||
|
workspace=str(agent_workspace),
|
||||||
|
)
|
||||||
|
read_tool = ReadFileTool.create(ctx)
|
||||||
|
write_tool = WriteFileTool.create(ctx)
|
||||||
|
set_agent_plugin_enabled(agent_workspace, "demo", True)
|
||||||
|
activation_checks = 0
|
||||||
|
activation_marker = agent_plugins._activation_marker
|
||||||
|
|
||||||
|
def count_activation_checks(plugin: agent_plugins.AgentPlugin) -> str | None:
|
||||||
|
nonlocal activation_checks
|
||||||
|
activation_checks += 1
|
||||||
|
return activation_marker(plugin)
|
||||||
|
|
||||||
|
monkeypatch.setattr(agent_plugins, "_activation_marker", count_activation_checks)
|
||||||
|
scope = validate_workspace_scope_payload(
|
||||||
|
{"project_path": str(project), "access_mode": "restricted"},
|
||||||
|
default_workspace=agent_workspace,
|
||||||
|
default_restrict_to_workspace=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
token = bind_workspace_scope(scope)
|
||||||
|
try:
|
||||||
|
read_result = await read_tool.execute(path=str(resource))
|
||||||
|
repeated_read_result = await read_tool.execute(path=str(resource))
|
||||||
|
write_result = await write_tool.execute(path=str(resource), content="changed")
|
||||||
|
set_agent_plugin_enabled(agent_workspace, "demo", False)
|
||||||
|
disabled_result = await read_tool.execute(path=str(resource))
|
||||||
|
finally:
|
||||||
|
reset_workspace_scope(token)
|
||||||
|
|
||||||
|
assert "plugin reference" in read_result
|
||||||
|
assert "File unchanged since last read" in repeated_read_result
|
||||||
|
assert activation_checks == 1
|
||||||
|
assert "outside allowed directory" in write_result
|
||||||
|
assert "outside allowed directory" in disabled_result
|
||||||
|
assert resource.read_text(encoding="utf-8") == "plugin reference"
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_state_symlink_cannot_escape_config_root(
|
||||||
|
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||||
|
) -> None:
|
||||||
|
config = tmp_path / "config"
|
||||||
|
config.mkdir()
|
||||||
|
outside = tmp_path / "outside"
|
||||||
|
outside.mkdir()
|
||||||
|
try:
|
||||||
|
(config / "plugin-data").symlink_to(outside, target_is_directory=True)
|
||||||
|
except OSError as exc:
|
||||||
|
pytest.skip(f"directory symlink unavailable: {exc}")
|
||||||
|
monkeypatch.setattr(agent_plugins, "get_config_path", lambda: config / "config.json")
|
||||||
|
_plugin(tmp_path, "desktop")
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="escapes its parent"):
|
||||||
|
set_agent_plugin_enabled(tmp_path, "desktop", True)
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_activation_requires_one_stable_package_identity(tmp_path: Path) -> None:
|
||||||
|
roots = [tmp_path / "plugins" / directory for directory in ("first", "second")]
|
||||||
|
for root, marker in zip(roots, ("trusted", "replacement"), strict=True):
|
||||||
|
_write_json(root / "plugin.json", _manifest("duplicate"))
|
||||||
|
_write_json(
|
||||||
|
root / "mcp.json",
|
||||||
|
{
|
||||||
|
"$schema": AGENT_PLUGIN_MCP_SCHEMA,
|
||||||
|
"mcpServers": {
|
||||||
|
"server": {"type": "stdio", "command": "echo", "args": [marker]}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert discover_agent_plugins(tmp_path) == []
|
||||||
|
with pytest.raises(ValueError, match="unknown Agent Plugin"):
|
||||||
|
set_agent_plugin_enabled(tmp_path, "duplicate", True)
|
||||||
|
|
||||||
|
shutil.rmtree(roots[1])
|
||||||
|
set_agent_plugin_enabled(tmp_path, "duplicate", True)
|
||||||
|
assert discover_agent_plugins(tmp_path)[0].enabled is True
|
||||||
|
|
||||||
|
moved = tmp_path / "plugins" / "moved"
|
||||||
|
roots[0].rename(moved)
|
||||||
|
assert discover_agent_plugins(tmp_path)[0].enabled is False
|
||||||
|
assert agent_plugin_mcp_servers(tmp_path) == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_legacy_path_activation_is_upgraded_to_package_fingerprint(tmp_path: Path) -> None:
|
||||||
|
plugin = _plugin(tmp_path)
|
||||||
|
set_agent_plugin_enabled(tmp_path, "demo", True)
|
||||||
|
marker = next((tmp_path / "config" / "plugin-data").glob("*/demo/enabled"))
|
||||||
|
marker.write_text(str(plugin), encoding="utf-8")
|
||||||
|
|
||||||
|
assert discover_agent_plugins(tmp_path)[0].enabled is True
|
||||||
|
assert marker.read_text(encoding="utf-8").startswith('{"fingerprint":')
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_activation_does_not_survive_in_place_contract_replacement(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
plugin = _plugin(tmp_path, "desktop")
|
||||||
|
mcp = plugin / "mcp.json"
|
||||||
|
|
||||||
|
def write_server(marker: str) -> None:
|
||||||
|
_write_json(
|
||||||
|
mcp,
|
||||||
|
{
|
||||||
|
"$schema": AGENT_PLUGIN_MCP_SCHEMA,
|
||||||
|
"mcpServers": {
|
||||||
|
"server": {"type": "stdio", "command": "echo", "args": [marker]}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
write_server("trusted")
|
||||||
|
set_agent_plugin_enabled(tmp_path, "desktop", True)
|
||||||
|
assert agent_plugin_mcp_servers(tmp_path)["desktop"].args == ["trusted"]
|
||||||
|
|
||||||
|
write_server("replacement")
|
||||||
|
|
||||||
|
assert discover_agent_plugins(tmp_path)[0].enabled is False
|
||||||
|
assert agent_plugin_mcp_servers(tmp_path) == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_plugin_activation_does_not_survive_in_place_code_replacement(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
plugin = _plugin(tmp_path, "desktop")
|
||||||
|
_skill(plugin / "skills", "demo")
|
||||||
|
executable = plugin / "server.py"
|
||||||
|
executable.write_text("print('trusted')\n", encoding="utf-8")
|
||||||
|
_write_json(
|
||||||
|
plugin / "mcp.json",
|
||||||
|
{
|
||||||
|
"$schema": AGENT_PLUGIN_MCP_SCHEMA,
|
||||||
|
"mcpServers": {
|
||||||
|
"server": {
|
||||||
|
"type": "stdio",
|
||||||
|
"command": "python",
|
||||||
|
"args": ["${PLUGIN_ROOT}/server.py"],
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
set_agent_plugin_enabled(tmp_path, "desktop", True)
|
||||||
|
assert discover_agent_plugins(tmp_path)[0].enabled is True
|
||||||
|
assert enabled_agent_plugin_skill_dirs(tmp_path) == (plugin / "skills" / "demo",)
|
||||||
|
|
||||||
|
executable.write_text("print('replacement')\n", encoding="utf-8")
|
||||||
|
|
||||||
|
assert discover_agent_plugins(tmp_path)[0].enabled is False
|
||||||
|
assert enabled_agent_plugin_skill_dirs(tmp_path) == ()
|
||||||
|
assert agent_plugin_mcp_servers(tmp_path) == {}
|
||||||
@@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.bus.events import InboundMessage
|
from nanobot.bus.events import InboundMessage
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.command import CommandContext
|
from nanobot.command import CommandContext
|
||||||
@@ -193,7 +194,11 @@ class TestIdleScanThrottling:
|
|||||||
})
|
})
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
provider.get_default_model.return_value = "test-model"
|
provider.get_default_model.return_value = "test-model"
|
||||||
loop = AgentLoop.from_config(config, provider=provider)
|
loop = AgentLoop.from_config(
|
||||||
|
config,
|
||||||
|
tool_registry=ToolRegistry(),
|
||||||
|
provider=provider,
|
||||||
|
)
|
||||||
loop.auto_compact.check_expired = MagicMock()
|
loop.auto_compact.check_expired = MagicMock()
|
||||||
|
|
||||||
loop._check_expired_sessions_if_due()
|
loop._check_expired_sessions_if_due()
|
||||||
@@ -310,7 +315,7 @@ class TestAutoCompact:
|
|||||||
assert loop.auto_compact._is_expired(ts) is True
|
assert loop.auto_compact._is_expired(ts) is True
|
||||||
ts2 = datetime.now() - timedelta(minutes=14, seconds=59)
|
ts2 = datetime.now() - timedelta(minutes=14, seconds=59)
|
||||||
assert loop.auto_compact._is_expired(ts2) is False
|
assert loop.auto_compact._is_expired(ts2) is False
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_is_expired_string_timestamp(self, tmp_path):
|
async def test_is_expired_string_timestamp(self, tmp_path):
|
||||||
@@ -320,7 +325,7 @@ class TestAutoCompact:
|
|||||||
assert loop.auto_compact._is_expired(ts) is True
|
assert loop.auto_compact._is_expired(ts) is True
|
||||||
assert loop.auto_compact._is_expired(None) is False
|
assert loop.auto_compact._is_expired(None) is False
|
||||||
assert loop.auto_compact._is_expired("") is False
|
assert loop.auto_compact._is_expired("") is False
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_check_expired_only_archives_expired_sessions(self, tmp_path):
|
async def test_check_expired_only_archives_expired_sessions(self, tmp_path):
|
||||||
@@ -343,7 +348,7 @@ class TestAutoCompact:
|
|||||||
active_after = loop.sessions.get_or_create("cli:active")
|
active_after = loop.sessions.get_or_create("cli:active")
|
||||||
assert len(active_after.messages) == 1
|
assert len(active_after.messages) == 1
|
||||||
assert active_after.messages[0]["content"] == "recent"
|
assert active_after.messages[0]["content"] == "recent"
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_archives_full_tail_without_deleting_history(self, tmp_path):
|
async def test_auto_compact_archives_full_tail_without_deleting_history(self, tmp_path):
|
||||||
@@ -367,7 +372,7 @@ class TestAutoCompact:
|
|||||||
assert len(visible) == loop.auto_compact._RECENT_SUFFIX_MESSAGES
|
assert len(visible) == loop.auto_compact._RECENT_SUFFIX_MESSAGES
|
||||||
assert visible[0]["content"] == "msg user 2"
|
assert visible[0]["content"] == "msg user 2"
|
||||||
assert visible[-1]["content"] == "msg assistant 5"
|
assert visible[-1]["content"] == "msg assistant 5"
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_extends_recent_suffix_to_user_turn(self, tmp_path):
|
async def test_auto_compact_extends_recent_suffix_to_user_turn(self, tmp_path):
|
||||||
@@ -398,7 +403,7 @@ class TestAutoCompact:
|
|||||||
for m in visible
|
for m in visible
|
||||||
for tc in (m.get("tool_calls") or [])
|
for tc in (m.get("tool_calls") or [])
|
||||||
)
|
)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_stores_summary(self, tmp_path):
|
async def test_auto_compact_stores_summary(self, tmp_path):
|
||||||
@@ -422,7 +427,7 @@ class TestAutoCompact:
|
|||||||
assert len(session_after.get_history(max_messages=12)) == (
|
assert len(session_after.get_history(max_messages=12)) == (
|
||||||
loop.auto_compact._RECENT_SUFFIX_MESSAGES
|
loop.auto_compact._RECENT_SUFFIX_MESSAGES
|
||||||
)
|
)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_empty_session(self, tmp_path):
|
async def test_auto_compact_empty_session(self, tmp_path):
|
||||||
@@ -436,7 +441,7 @@ class TestAutoCompact:
|
|||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert len(session_after.messages) == 0
|
assert len(session_after.messages) == 0
|
||||||
assert "cli:test" not in loop.auto_compact._summaries
|
assert "cli:test" not in loop.auto_compact._summaries
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_respects_last_consolidated(self, tmp_path):
|
async def test_auto_compact_respects_last_consolidated(self, tmp_path):
|
||||||
@@ -455,7 +460,7 @@ class TestAutoCompact:
|
|||||||
await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime())
|
await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime())
|
||||||
|
|
||||||
assert len(archived_messages) == 10
|
assert len(archived_messages) == 10
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
|
|
||||||
class TestAutoCompactIdleDetection:
|
class TestAutoCompactIdleDetection:
|
||||||
@@ -474,7 +479,7 @@ class TestAutoCompactIdleDetection:
|
|||||||
|
|
||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert any(m["content"] == "old message" for m in session_after.messages)
|
assert any(m["content"] == "old message" for m in session_after.messages)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_triggers_on_idle(self, tmp_path):
|
async def test_auto_compact_triggers_on_idle(self, tmp_path):
|
||||||
@@ -503,7 +508,7 @@ class TestAutoCompactIdleDetection:
|
|||||||
for m in session_after.get_history(max_messages=len(session_after.messages))
|
for m in session_after.get_history(max_messages=len(session_after.messages))
|
||||||
)
|
)
|
||||||
assert any(m["content"] == "new msg" for m in session_after.messages)
|
assert any(m["content"] == "new msg" for m in session_after.messages)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_no_auto_compact_when_active(self, tmp_path):
|
async def test_no_auto_compact_when_active(self, tmp_path):
|
||||||
@@ -517,7 +522,7 @@ class TestAutoCompactIdleDetection:
|
|||||||
|
|
||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert any(m["content"] == "recent message" for m in session_after.messages)
|
assert any(m["content"] == "recent message" for m in session_after.messages)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_does_not_affect_priority_commands(self, tmp_path):
|
async def test_auto_compact_does_not_affect_priority_commands(self, tmp_path):
|
||||||
@@ -540,7 +545,7 @@ class TestAutoCompactIdleDetection:
|
|||||||
# Session should be untouched since priority commands skip _process_message
|
# Session should be untouched since priority commands skip _process_message
|
||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert any(m["content"] == "old message" for m in session_after.messages)
|
assert any(m["content"] == "old message" for m in session_after.messages)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_with_slash_new(self, tmp_path):
|
async def test_auto_compact_with_slash_new(self, tmp_path):
|
||||||
@@ -562,7 +567,7 @@ class TestAutoCompactIdleDetection:
|
|||||||
|
|
||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert len(session_after.messages) == 0
|
assert len(session_after.messages) == 0
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_shortcut_command_persisted_with_command_flag(self, tmp_path):
|
async def test_shortcut_command_persisted_with_command_flag(self, tmp_path):
|
||||||
@@ -581,7 +586,7 @@ class TestAutoCompactIdleDetection:
|
|||||||
assert session_after.messages[1]["role"] == "assistant"
|
assert session_after.messages[1]["role"] == "assistant"
|
||||||
assert session_after.messages[1].get("_command") is True
|
assert session_after.messages[1].get("_command") is True
|
||||||
assert AgentLoop._PENDING_USER_TURN_KEY not in session_after.metadata
|
assert AgentLoop._PENDING_USER_TURN_KEY not in session_after.metadata
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_shortcut_command_excluded_from_get_history(self, tmp_path):
|
async def test_shortcut_command_excluded_from_get_history(self, tmp_path):
|
||||||
@@ -597,7 +602,7 @@ class TestAutoCompactIdleDetection:
|
|||||||
assert len(history) == 2
|
assert len(history) == 2
|
||||||
assert all(m["content"] != "/help" for m in history)
|
assert all(m["content"] != "/help" for m in history)
|
||||||
assert all(m["content"] != "help text" for m in history)
|
assert all(m["content"] != "help text" for m in history)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
|
|
||||||
class TestAutoCompactSystemMessages:
|
class TestAutoCompactSystemMessages:
|
||||||
@@ -628,7 +633,7 @@ class TestAutoCompactSystemMessages:
|
|||||||
m["content"] == "old user 0"
|
m["content"] == "old user 0"
|
||||||
for m in session_after.get_history(max_messages=len(session_after.messages))
|
for m in session_after.get_history(max_messages=len(session_after.messages))
|
||||||
)
|
)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
|
|
||||||
class TestAutoCompactEdgeCases:
|
class TestAutoCompactEdgeCases:
|
||||||
@@ -656,7 +661,7 @@ class TestAutoCompactEdgeCases:
|
|||||||
# "(nothing)" summary should not be stored
|
# "(nothing)" summary should not be stored
|
||||||
assert "cli:test" not in loop.auto_compact._summaries
|
assert "cli:test" not in loop.auto_compact._summaries
|
||||||
|
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_archive_failure_preserves_raw_history(self, tmp_path):
|
async def test_auto_compact_archive_failure_preserves_raw_history(self, tmp_path):
|
||||||
@@ -677,7 +682,7 @@ class TestAutoCompactEdgeCases:
|
|||||||
loop.auto_compact._RECENT_SUFFIX_MESSAGES
|
loop.auto_compact._RECENT_SUFFIX_MESSAGES
|
||||||
)
|
)
|
||||||
|
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_preserves_runtime_checkpoint_before_check(self, tmp_path):
|
async def test_auto_compact_preserves_runtime_checkpoint_before_check(self, tmp_path):
|
||||||
@@ -709,7 +714,7 @@ class TestAutoCompactEdgeCases:
|
|||||||
assert any(m["content"] == "previous message" for m in session_after.messages)
|
assert any(m["content"] == "previous message" for m in session_after.messages)
|
||||||
assert any(m["content"] == "interrupted response" for m in session_after.messages)
|
assert any(m["content"] == "interrupted response" for m in session_after.messages)
|
||||||
|
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
|
|
||||||
class TestAutoCompactIntegration:
|
class TestAutoCompactIntegration:
|
||||||
@@ -779,7 +784,7 @@ class TestAutoCompactIntegration:
|
|||||||
# The new message should be processed (response exists)
|
# The new message should be processed (response exists)
|
||||||
assert response is not None
|
assert response is not None
|
||||||
|
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runtime_context_markers_not_persisted_for_multi_paragraph_turn(self, tmp_path):
|
async def test_runtime_context_markers_not_persisted_for_multi_paragraph_turn(self, tmp_path):
|
||||||
@@ -807,7 +812,7 @@ class TestAutoCompactIntegration:
|
|||||||
content = str(persisted.get("content", ""))
|
content = str(persisted.get("content", ""))
|
||||||
assert "[Runtime Context" not in content
|
assert "[Runtime Context" not in content
|
||||||
assert "[/Runtime Context]" not in content
|
assert "[/Runtime Context]" not in content
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
|
|
||||||
class TestProactiveAutoCompact:
|
class TestProactiveAutoCompact:
|
||||||
@@ -870,7 +875,7 @@ class TestProactiveAutoCompact:
|
|||||||
|
|
||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert len(session_after.messages) == 1
|
assert len(session_after.messages) == 1
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_proactive_archive_on_idle_tick(self, tmp_path):
|
async def test_proactive_archive_on_idle_tick(self, tmp_path):
|
||||||
@@ -897,7 +902,7 @@ class TestProactiveAutoCompact:
|
|||||||
entry = loop.auto_compact._summaries.get("cli:test")
|
entry = loop.auto_compact._summaries.get("cli:test")
|
||||||
assert entry is not None
|
assert entry is not None
|
||||||
assert entry[0] == "User chatted about old things."
|
assert entry[0] == "User chatted about old things."
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_proactive_archive_skips_dream_sessions(self, tmp_path):
|
async def test_proactive_archive_skips_dream_sessions(self, tmp_path):
|
||||||
@@ -918,7 +923,7 @@ class TestProactiveAutoCompact:
|
|||||||
assert _fake_compact.state["count"] == 0
|
assert _fake_compact.state["count"] == 0
|
||||||
assert "dream:20260602-155256" not in loop.auto_compact._archiving
|
assert "dream:20260602-155256" not in loop.auto_compact._archiving
|
||||||
assert "dream:20260602-155256" not in loop.auto_compact._summaries
|
assert "dream:20260602-155256" not in loop.auto_compact._summaries
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_no_proactive_archive_when_active(self, tmp_path):
|
async def test_no_proactive_archive_when_active(self, tmp_path):
|
||||||
@@ -932,7 +937,7 @@ class TestProactiveAutoCompact:
|
|||||||
|
|
||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert len(session_after.messages) == 1
|
assert len(session_after.messages) == 1
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_no_duplicate_archive(self, tmp_path):
|
async def test_no_duplicate_archive(self, tmp_path):
|
||||||
@@ -968,7 +973,7 @@ class TestProactiveAutoCompact:
|
|||||||
# Clean up
|
# Clean up
|
||||||
block_forever.set()
|
block_forever.set()
|
||||||
await _drain_background_tasks(loop)
|
await _drain_background_tasks(loop)
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_proactive_archive_error_does_not_block(self, tmp_path):
|
async def test_proactive_archive_error_does_not_block(self, tmp_path):
|
||||||
@@ -989,7 +994,7 @@ class TestProactiveAutoCompact:
|
|||||||
|
|
||||||
# Key should be removed from _archiving (finally block)
|
# Key should be removed from _archiving (finally block)
|
||||||
assert "cli:test" not in loop.auto_compact._archiving
|
assert "cli:test" not in loop.auto_compact._archiving
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_proactive_archive_skips_empty_sessions(self, tmp_path):
|
async def test_proactive_archive_skips_empty_sessions(self, tmp_path):
|
||||||
@@ -1005,7 +1010,7 @@ class TestProactiveAutoCompact:
|
|||||||
|
|
||||||
# Empty session should not produce a summary
|
# Empty session should not produce a summary
|
||||||
assert "cli:test" not in loop.auto_compact._summaries
|
assert "cli:test" not in loop.auto_compact._summaries
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_skip_expired_session_with_active_agent_task(self, tmp_path):
|
async def test_skip_expired_session_with_active_agent_task(self, tmp_path):
|
||||||
@@ -1026,7 +1031,7 @@ class TestProactiveAutoCompact:
|
|||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert len(session_after.messages) == 12 # All messages preserved
|
assert len(session_after.messages) == 12 # All messages preserved
|
||||||
|
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_archive_after_active_task_completes(self, tmp_path):
|
async def test_archive_after_active_task_completes(self, tmp_path):
|
||||||
@@ -1047,7 +1052,7 @@ class TestProactiveAutoCompact:
|
|||||||
# Second tick: task completed, should archive
|
# Second tick: task completed, should archive
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
assert _fake_compact.state["count"] == 1
|
assert _fake_compact.state["count"] == 1
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_partial_active_set_only_archives_inactive_expired(self, tmp_path):
|
async def test_partial_active_set_only_archives_inactive_expired(self, tmp_path):
|
||||||
@@ -1083,7 +1088,7 @@ class TestProactiveAutoCompact:
|
|||||||
assert len(s2_after.messages) == 12 # Preserved
|
assert len(s2_after.messages) == 12 # Preserved
|
||||||
s3_after = loop.sessions.get_or_create("cli:recent")
|
s3_after = loop.sessions.get_or_create("cli:recent")
|
||||||
assert len(s3_after.messages) == 1 # Preserved
|
assert len(s3_after.messages) == 1 # Preserved
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_no_reschedule_after_successful_archive(self, tmp_path):
|
async def test_no_reschedule_after_successful_archive(self, tmp_path):
|
||||||
@@ -1104,7 +1109,7 @@ class TestProactiveAutoCompact:
|
|||||||
# Second tick: should NOT re-schedule because the session has no removable tail.
|
# Second tick: should NOT re-schedule because the session has no removable tail.
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
assert _fake_compact.state["count"] == 1 # Still 1, not re-scheduled
|
assert _fake_compact.state["count"] == 1 # Still 1, not re-scheduled
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_empty_session_does_not_schedule_idle_compact(self, tmp_path):
|
async def test_empty_session_does_not_schedule_idle_compact(self, tmp_path):
|
||||||
@@ -1124,7 +1129,7 @@ class TestProactiveAutoCompact:
|
|||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
assert _fake_compact.state["count"] == 0
|
assert _fake_compact.state["count"] == 0
|
||||||
assert "cli:test" not in loop.auto_compact._summaries
|
assert "cli:test" not in loop.auto_compact._summaries
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_session_can_be_compacted_again_after_new_messages(self, tmp_path):
|
async def test_session_can_be_compacted_again_after_new_messages(self, tmp_path):
|
||||||
@@ -1155,7 +1160,7 @@ class TestProactiveAutoCompact:
|
|||||||
# Second compact cycle should succeed
|
# Second compact cycle should succeed
|
||||||
await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime())
|
await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime())
|
||||||
assert _fake_compact.state["count"] == 2
|
assert _fake_compact.state["count"] == 2
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
|
|
||||||
class TestSummaryPersistence:
|
class TestSummaryPersistence:
|
||||||
@@ -1182,7 +1187,7 @@ class TestSummaryPersistence:
|
|||||||
assert meta is not None
|
assert meta is not None
|
||||||
assert meta["text"] == "User said hello."
|
assert meta["text"] == "User said hello."
|
||||||
assert "last_active" in meta
|
assert "last_active" in meta
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_summary_recovered_after_restart(self, tmp_path):
|
async def test_summary_recovered_after_restart(self, tmp_path):
|
||||||
@@ -1218,7 +1223,7 @@ class TestSummaryPersistence:
|
|||||||
assert "Previous conversation summary" in summary
|
assert "Previous conversation summary" in summary
|
||||||
# _last_summary persists in metadata for restart survival.
|
# _last_summary persists in metadata for restart survival.
|
||||||
assert "_last_summary" in reloaded.metadata
|
assert "_last_summary" in reloaded.metadata
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_metadata_persists_for_restart(self, tmp_path):
|
async def test_metadata_persists_for_restart(self, tmp_path):
|
||||||
@@ -1246,7 +1251,7 @@ class TestSummaryPersistence:
|
|||||||
assert "Summary." in summary2
|
assert "Summary." in summary2
|
||||||
# _last_summary persists in metadata for restart survival.
|
# _last_summary persists in metadata for restart survival.
|
||||||
assert "_last_summary" in reloaded.metadata
|
assert "_last_summary" in reloaded.metadata
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_metadata_cleanup_on_inmemory_path(self, tmp_path):
|
async def test_metadata_cleanup_on_inmemory_path(self, tmp_path):
|
||||||
@@ -1272,7 +1277,7 @@ class TestSummaryPersistence:
|
|||||||
assert summary is not None
|
assert summary is not None
|
||||||
# _last_summary persists in metadata for restart survival.
|
# _last_summary persists in metadata for restart survival.
|
||||||
assert "_last_summary" in reloaded.metadata
|
assert "_last_summary" in reloaded.metadata
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_new_summary_overrides_old(self, tmp_path):
|
async def test_new_summary_overrides_old(self, tmp_path):
|
||||||
@@ -1314,7 +1319,7 @@ class TestSummaryPersistence:
|
|||||||
_, summary2 = loop.auto_compact.prepare_session(reloaded, "cli:test")
|
_, summary2 = loop.auto_compact.prepare_session(reloaded, "cli:test")
|
||||||
assert summary2 is not None
|
assert summary2 is not None
|
||||||
assert "Second summary." in summary2
|
assert "Second summary." in summary2
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_new_command_clears_last_summary(self, tmp_path):
|
async def test_new_command_clears_last_summary(self, tmp_path):
|
||||||
@@ -1342,4 +1347,4 @@ class TestSummaryPersistence:
|
|||||||
# After /new, metadata should no longer contain _last_summary
|
# After /new, metadata should no longer contain _last_summary
|
||||||
fresh = loop.sessions.get_or_create("cli:test")
|
fresh = loop.sessions.get_or_create("cli:test")
|
||||||
assert "_last_summary" not in fresh.metadata
|
assert "_last_summary" not in fresh.metadata
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
||||||
|
|
||||||
|
|
||||||
|
def test_weather_skill_uses_windows_safe_single_today_request() -> None:
|
||||||
|
content = (BUILTIN_SKILLS_DIR / "weather" / "SKILL.md").read_text(encoding="utf-8")
|
||||||
|
normalized = " ".join(content.split())
|
||||||
|
|
||||||
|
assert "On Windows PowerShell, use `curl.exe`" in normalized
|
||||||
|
assert "bare `curl` may resolve to `Invoke-WebRequest`" in normalized
|
||||||
|
assert "https://wttr.in/London?1&m" in content
|
||||||
|
assert 'curl.exe -s "https://wttr.in/Berlin.png" -o weather.png' in content
|
||||||
|
assert "/tmp/weather.png" not in content
|
||||||
|
assert (
|
||||||
|
"Do not fetch current conditions separately when a today or forecast "
|
||||||
|
"request already includes them."
|
||||||
|
) in normalized
|
||||||
@@ -538,7 +538,7 @@ class TestNewCommandArchival:
|
|||||||
session_after = loop.sessions.get_or_create("cli:test")
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
assert len(session_after.messages) == 0
|
assert len(session_after.messages) == 0
|
||||||
|
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
assert call_count == 1
|
assert call_count == 1
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -572,7 +572,7 @@ class TestNewCommandArchival:
|
|||||||
assert response is not None
|
assert response is not None
|
||||||
assert "new session started" in response.content.lower()
|
assert "new session started" in response.content.lower()
|
||||||
|
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
assert archived_count == 3
|
assert archived_count == 3
|
||||||
assert archived_session_key == "cli:test"
|
assert archived_session_key == "cli:test"
|
||||||
|
|
||||||
@@ -603,8 +603,8 @@ class TestNewCommandArchival:
|
|||||||
assert loop.sessions.get_or_create("cli:test").messages == []
|
assert loop.sessions.get_or_create("cli:test").messages == []
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_close_mcp_drains_background_tasks(self, tmp_path: Path) -> None:
|
async def test_aclose_drains_background_tasks(self, tmp_path: Path) -> None:
|
||||||
"""close_mcp waits for background tasks to complete."""
|
"""aclose waits for background tasks to complete."""
|
||||||
from nanobot.bus.events import InboundMessage
|
from nanobot.bus.events import InboundMessage
|
||||||
|
|
||||||
loop = self._make_loop(tmp_path)
|
loop = self._make_loop(tmp_path)
|
||||||
@@ -632,5 +632,5 @@ class TestNewCommandArchival:
|
|||||||
|
|
||||||
assert not archived.is_set()
|
assert not archived.is_set()
|
||||||
release_archive.set()
|
release_archive.set()
|
||||||
await loop.close_mcp()
|
await loop.aclose()
|
||||||
assert archived.is_set()
|
assert archived.is_set()
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from unittest.mock import patch
|
|||||||
|
|
||||||
from nanobot.providers.base import ToolCallRequest
|
from nanobot.providers.base import ToolCallRequest
|
||||||
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
|
from nanobot.providers.registry import ProviderSpec
|
||||||
|
|
||||||
GEMINI_EXTRA = {"google": {"thought_signature": "sig-abc-123"}}
|
GEMINI_EXTRA = {"google": {"thought_signature": "sig-abc-123"}}
|
||||||
|
|
||||||
@@ -243,3 +244,251 @@ def test_stale_extra_content_in_tool_calls_survives_sanitize() -> None:
|
|||||||
sanitized = provider._sanitize_messages(messages)
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
assert sanitized[1]["tool_calls"][0]["extra_content"] == GEMINI_EXTRA
|
assert sanitized[1]["tool_calls"][0]["extra_content"] == GEMINI_EXTRA
|
||||||
|
|
||||||
|
|
||||||
|
# ── Replay to Gemini: preserve or backfill thought signatures ─────────
|
||||||
|
|
||||||
|
def _gemini_provider() -> OpenAICompatProvider:
|
||||||
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
||||||
|
return OpenAICompatProvider(
|
||||||
|
spec=ProviderSpec(
|
||||||
|
name="gemini", keywords=("gemini",), env_key="GEMINI_API_KEY"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _tool_call(tc_id: str, name: str, *, signed: bool = False) -> dict:
|
||||||
|
tc: dict = {
|
||||||
|
"id": tc_id,
|
||||||
|
"type": "function",
|
||||||
|
"function": {"name": name, "arguments": "{}"},
|
||||||
|
}
|
||||||
|
if signed:
|
||||||
|
tc["extra_content"] = GEMINI_EXTRA
|
||||||
|
return tc
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemini_backfills_unsigned_tool_calls_and_keeps_results() -> None:
|
||||||
|
"""Cross-provider history stays intact and receives the documented fallback."""
|
||||||
|
provider = _gemini_provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "check the sensor"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": "On it.",
|
||||||
|
"tool_calls": [_tool_call("default_api:exec", "exec")],
|
||||||
|
},
|
||||||
|
{"role": "tool", "content": "done", "tool_call_id": "default_api:exec"},
|
||||||
|
{"role": "user", "content": "thanks"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
assert [m["role"] for m in sanitized] == ["user", "assistant", "tool", "user"]
|
||||||
|
call = sanitized[1]["tool_calls"][0]
|
||||||
|
assert call["extra_content"]["google"]["thought_signature"] == (
|
||||||
|
"skip_thought_signature_validator"
|
||||||
|
)
|
||||||
|
assert sanitized[2]["tool_call_id"] == call["id"]
|
||||||
|
assert sanitized[2]["content"] == "done"
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemini_preserves_parallel_calls_when_only_first_is_signed() -> None:
|
||||||
|
"""Gemini signs only the first native parallel call; all calls must replay."""
|
||||||
|
provider = _gemini_provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "do both"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [
|
||||||
|
_tool_call("call_signed", "read_file", signed=True),
|
||||||
|
_tool_call("default_api:exec", "exec"),
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{"role": "tool", "content": "file contents", "tool_call_id": "call_signed"},
|
||||||
|
{"role": "tool", "content": "done", "tool_call_id": "default_api:exec"},
|
||||||
|
{"role": "user", "content": "thanks"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
assert [m["role"] for m in sanitized] == [
|
||||||
|
"user",
|
||||||
|
"assistant",
|
||||||
|
"tool",
|
||||||
|
"tool",
|
||||||
|
"user",
|
||||||
|
]
|
||||||
|
calls = sanitized[1]["tool_calls"]
|
||||||
|
assert len(calls) == 2
|
||||||
|
assert calls[0]["extra_content"] == GEMINI_EXTRA
|
||||||
|
assert sanitized[2]["tool_call_id"] == calls[0]["id"]
|
||||||
|
assert sanitized[2]["content"] == "file contents"
|
||||||
|
assert "extra_content" not in calls[1]
|
||||||
|
assert sanitized[3]["tool_call_id"] == calls[1]["id"]
|
||||||
|
assert sanitized[3]["content"] == "done"
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemini_backfills_only_first_cross_provider_parallel_call() -> None:
|
||||||
|
provider = _gemini_provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "do both"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [
|
||||||
|
_tool_call("call_1", "read_file"),
|
||||||
|
_tool_call("call_2", "exec"),
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{"role": "tool", "content": "file contents", "tool_call_id": "call_1"},
|
||||||
|
{"role": "tool", "content": "done", "tool_call_id": "call_2"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
calls = sanitized[1]["tool_calls"]
|
||||||
|
assert len(calls) == 2
|
||||||
|
assert calls[0]["extra_content"]["google"]["thought_signature"] == (
|
||||||
|
"skip_thought_signature_validator"
|
||||||
|
)
|
||||||
|
assert "extra_content" not in calls[1]
|
||||||
|
assert [message["content"] for message in sanitized[2:]] == ["file contents", "done"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemini_requires_signature_on_first_parallel_call() -> None:
|
||||||
|
provider = _gemini_provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "do both"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [
|
||||||
|
_tool_call("call_1", "read_file"),
|
||||||
|
_tool_call("call_2", "exec", signed=True),
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{"role": "tool", "content": "contents", "tool_call_id": "call_1"},
|
||||||
|
{"role": "tool", "content": "done", "tool_call_id": "call_2"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
calls = sanitized[1]["tool_calls"]
|
||||||
|
assert calls[0]["extra_content"]["google"]["thought_signature"] == (
|
||||||
|
"skip_thought_signature_validator"
|
||||||
|
)
|
||||||
|
assert calls[1]["extra_content"] == GEMINI_EXTRA
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemini_replay_preserves_signed_tool_calls() -> None:
|
||||||
|
"""A pure Gemini-origin history replays unchanged (signature intact)."""
|
||||||
|
provider = _gemini_provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "hi"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [_tool_call("call_1", "get_weather", signed=True)],
|
||||||
|
},
|
||||||
|
{"role": "tool", "content": "sunny", "tool_call_id": "call_1"},
|
||||||
|
{"role": "user", "content": "thanks"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
assert [m["role"] for m in sanitized] == ["user", "assistant", "tool", "user"]
|
||||||
|
calls = sanitized[1]["tool_calls"]
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert calls[0]["extra_content"] == GEMINI_EXTRA
|
||||||
|
assert sanitized[2]["tool_call_id"] == calls[0]["id"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_gemini_provider_keeps_unsigned_tool_calls() -> None:
|
||||||
|
"""The filter is Gemini-scoped: other providers still replay unsigned calls."""
|
||||||
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
||||||
|
provider = OpenAICompatProvider()
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "hi"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [_tool_call("default_api:exec", "exec")],
|
||||||
|
},
|
||||||
|
{"role": "tool", "content": "done", "tool_call_id": "default_api:exec"},
|
||||||
|
{"role": "user", "content": "thanks"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
assert len(sanitized[1]["tool_calls"]) == 1
|
||||||
|
assert sanitized[2]["role"] == "tool"
|
||||||
|
assert sanitized[2]["tool_call_id"] == sanitized[1]["tool_calls"][0]["id"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemini_drops_malformed_tool_call_entries_without_crashing() -> None:
|
||||||
|
provider = _gemini_provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "hi"},
|
||||||
|
{"role": "assistant", "content": None, "tool_calls": [None]},
|
||||||
|
{"role": "user", "content": "continue"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
assert not any(message.get("tool_calls") for message in sanitized)
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemini_matches_duplicate_tool_ids_by_call_instance() -> None:
|
||||||
|
provider = _gemini_provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "old request"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [_tool_call("reused", "old_tool")],
|
||||||
|
},
|
||||||
|
{"role": "tool", "content": "old result", "tool_call_id": "reused"},
|
||||||
|
{"role": "user", "content": "new request"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [_tool_call("reused", "new_tool", signed=True)],
|
||||||
|
},
|
||||||
|
{"role": "tool", "content": "new result", "tool_call_id": "reused"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
assert any(message.get("content") == "old result" for message in sanitized)
|
||||||
|
assert any(message.get("content") == "new result" for message in sanitized)
|
||||||
|
calls = [
|
||||||
|
call
|
||||||
|
for message in sanitized
|
||||||
|
for call in message.get("tool_calls", [])
|
||||||
|
]
|
||||||
|
assert len(calls) == 2
|
||||||
|
assert calls[0]["function"]["name"] == "old_tool"
|
||||||
|
assert calls[0]["extra_content"]["google"]["thought_signature"] == (
|
||||||
|
"skip_thought_signature_validator"
|
||||||
|
)
|
||||||
|
assert calls[1]["function"]["name"] == "new_tool"
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemini_backfill_does_not_mutate_caller_history() -> None:
|
||||||
|
provider = _gemini_provider()
|
||||||
|
call = _tool_call("call_1", "read_file")
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "read it"},
|
||||||
|
{"role": "assistant", "content": None, "tool_calls": [call]},
|
||||||
|
{"role": "tool", "content": "contents", "tool_call_id": "call_1"},
|
||||||
|
]
|
||||||
|
|
||||||
|
sanitized = provider._sanitize_messages(messages)
|
||||||
|
|
||||||
|
assert "extra_content" not in call
|
||||||
|
assert sanitized[1]["tool_calls"][0]["extra_content"]["google"][
|
||||||
|
"thought_signature"
|
||||||
|
] == "skip_thought_signature_validator"
|
||||||
|
|||||||
@@ -93,7 +93,6 @@ async def test_process_direct_websocket_clears_run_status(tmp_path) -> None:
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_process_direct_reuses_existing_session_lock(tmp_path) -> None:
|
async def test_process_direct_reuses_existing_session_lock(tmp_path) -> None:
|
||||||
loop = _make_loop(tmp_path)
|
loop = _make_loop(tmp_path)
|
||||||
loop._connect_mcp = AsyncMock()
|
|
||||||
session_key = "api:fixed"
|
session_key = "api:fixed"
|
||||||
lock = loop._session_locks.setdefault(session_key, asyncio.Lock())
|
lock = loop._session_locks.setdefault(session_key, asyncio.Lock())
|
||||||
await lock.acquire()
|
await lock.acquire()
|
||||||
|
|||||||
@@ -1519,13 +1519,11 @@ async def test_run_agent_loop_goal_continue_message_reads_latest_metadata(
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_process_direct_rejects_reserved_system_channel(tmp_path: Path) -> None:
|
async def test_process_direct_rejects_reserved_system_channel(tmp_path: Path) -> None:
|
||||||
loop = _make_full_loop(tmp_path)
|
loop = _make_full_loop(tmp_path)
|
||||||
loop._connect_mcp = AsyncMock() # type: ignore[method-assign]
|
|
||||||
loop._process_message = AsyncMock(return_value=None) # type: ignore[method-assign]
|
loop._process_message = AsyncMock(return_value=None) # type: ignore[method-assign]
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="reserved for internal messages"):
|
with pytest.raises(ValueError, match="reserved for internal messages"):
|
||||||
await loop.process_direct("external input", channel="system")
|
await loop.process_direct("external input", channel="system")
|
||||||
|
|
||||||
loop._connect_mcp.assert_not_awaited()
|
|
||||||
loop._process_message.assert_not_awaited()
|
loop._process_message.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
@@ -1534,7 +1532,6 @@ async def test_process_direct_skip_user_persist_does_not_save_retry_user(
|
|||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
) -> None:
|
) -> None:
|
||||||
loop = _make_full_loop(tmp_path)
|
loop = _make_full_loop(tmp_path)
|
||||||
loop._connect_mcp = AsyncMock()
|
|
||||||
session = loop.sessions.get_or_create("api:default")
|
session = loop.sessions.get_or_create("api:default")
|
||||||
session.add_message("user", "hello")
|
session.add_message("user", "hello")
|
||||||
session.add_message("assistant", "previous empty-response attempt")
|
session.add_message("assistant", "previous empty-response attempt")
|
||||||
|
|||||||
@@ -123,8 +123,7 @@ async def test_session_discard_control_cancels_active_turn(tmp_path, monkeypatch
|
|||||||
await asyncio.sleep(0)
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
loop.provider.chat_with_retry = AsyncMock(side_effect=block_provider)
|
loop.provider.chat_with_retry = AsyncMock(side_effect=block_provider)
|
||||||
monkeypatch.setattr(loop, "_connect_mcp", AsyncMock())
|
monkeypatch.setattr(loop, "aclose", AsyncMock())
|
||||||
monkeypatch.setattr(loop, "close_mcp", AsyncMock())
|
|
||||||
terminate_exec_sessions = AsyncMock(return_value=1)
|
terminate_exec_sessions = AsyncMock(return_value=1)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
loop._exec_session_manager,
|
loop._exec_session_manager,
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
import inspect
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
@@ -11,8 +12,10 @@ from nanobot.agent.tools.context import (
|
|||||||
current_request_context,
|
current_request_context,
|
||||||
reset_request_context,
|
reset_request_context,
|
||||||
)
|
)
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.bus.events import InboundMessage
|
from nanobot.bus.events import InboundMessage
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.config.schema import Config
|
||||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
from nanobot.session.turn_continuation import INTERNAL_CONTINUATION_META
|
from nanobot.session.turn_continuation import INTERNAL_CONTINUATION_META
|
||||||
|
|
||||||
@@ -56,6 +59,51 @@ class _Tools:
|
|||||||
return (self.tool, arguments, None) if name == "cron" else (None, arguments, None)
|
return (self.tool, arguments, None) if name == "cron" else (None, arguments, None)
|
||||||
|
|
||||||
|
|
||||||
|
def test_loop_registers_default_tools_in_injected_registry(tmp_path: Path) -> None:
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
registry = ToolRegistry()
|
||||||
|
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=MessageBus(),
|
||||||
|
provider=provider,
|
||||||
|
workspace=tmp_path,
|
||||||
|
tool_registry=registry,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert loop.tools is registry
|
||||||
|
assert registry.has("read_file")
|
||||||
|
|
||||||
|
|
||||||
|
def _config_for_loop(tmp_path: Path) -> Config:
|
||||||
|
return Config.model_validate({"agents": {"defaults": {"workspace": str(tmp_path)}}})
|
||||||
|
|
||||||
|
|
||||||
|
def _provider_for_loop() -> MagicMock:
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
return provider
|
||||||
|
|
||||||
|
|
||||||
|
def test_loop_from_config_requires_caller_owned_registry(tmp_path: Path) -> None:
|
||||||
|
signature = inspect.signature(AgentLoop.from_config)
|
||||||
|
|
||||||
|
with pytest.raises(TypeError, match="tool_registry"):
|
||||||
|
signature.bind(_config_for_loop(tmp_path))
|
||||||
|
|
||||||
|
|
||||||
|
def test_loop_from_config_uses_caller_owned_registry(tmp_path: Path) -> None:
|
||||||
|
registry = ToolRegistry()
|
||||||
|
loop = AgentLoop.from_config(
|
||||||
|
_config_for_loop(tmp_path),
|
||||||
|
tool_registry=registry,
|
||||||
|
provider=_provider_for_loop(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert loop.tools is registry
|
||||||
|
assert loop.tools.has("read_file")
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) -> None:
|
async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) -> None:
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Tests for MCP connection lifecycle in AgentLoop."""
|
"""Tests for the application-owned MCP provider lifecycle."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -6,7 +6,7 @@ import asyncio
|
|||||||
from contextlib import AsyncExitStack
|
from contextlib import AsyncExitStack
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
import anyio
|
import anyio
|
||||||
import pytest
|
import pytest
|
||||||
@@ -15,11 +15,10 @@ from mcp.shared.exceptions import McpError
|
|||||||
from mcp.shared.message import SessionMessage
|
from mcp.shared.message import SessionMessage
|
||||||
from mcp.types import ErrorData
|
from mcp.types import ErrorData
|
||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
|
||||||
from nanobot.agent.tools import mcp as mcp_runtime
|
from nanobot.agent.tools import mcp as mcp_runtime
|
||||||
from nanobot.agent.tools.base import Tool
|
from nanobot.agent.tools.base import Tool
|
||||||
from nanobot.agent.tools.mcp import MCPResourceWrapper, MCPToolWrapper
|
from nanobot.agent.tools.mcp import MCPProvider, MCPResourceWrapper, MCPToolWrapper
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.config.loader import load_config, save_config
|
from nanobot.config.loader import load_config, save_config
|
||||||
from nanobot.config.schema import MCPServerConfig
|
from nanobot.config.schema import MCPServerConfig
|
||||||
|
|
||||||
@@ -74,18 +73,20 @@ class _FakeMcpTool(Tool):
|
|||||||
return "ok"
|
return "ok"
|
||||||
|
|
||||||
|
|
||||||
def _make_loop(tmp_path, *, mcp_servers: dict | None = None) -> AgentLoop:
|
def _stdio_server(command: str = "test-mcp") -> MCPServerConfig:
|
||||||
bus = MessageBus()
|
return MCPServerConfig(type="stdio", command=command)
|
||||||
provider = MagicMock()
|
|
||||||
provider.get_default_model.return_value = "test-model"
|
|
||||||
provider.generation.max_tokens = 4096
|
def _make_provider(
|
||||||
return AgentLoop(
|
*,
|
||||||
bus=bus,
|
mcp_servers: dict[str, MCPServerConfig] | None = None,
|
||||||
provider=provider,
|
) -> tuple[MCPProvider, ToolRegistry]:
|
||||||
workspace=tmp_path,
|
registry = ToolRegistry()
|
||||||
model="test-model",
|
provider = MCPProvider(
|
||||||
mcp_servers=mcp_servers or {"test": object()},
|
mcp_servers if mcp_servers is not None else {"test": _stdio_server()},
|
||||||
|
registry,
|
||||||
)
|
)
|
||||||
|
return provider, registry
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -140,7 +141,7 @@ async def test_owned_mcp_connection_closes_from_its_owner_task():
|
|||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch):
|
async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch):
|
||||||
loop = _make_loop(tmp_path)
|
provider, _registry = _make_provider()
|
||||||
attempts = 0
|
attempts = 0
|
||||||
|
|
||||||
async def _fake_connect(_servers, _registry):
|
async def _fake_connect(_servers, _registry):
|
||||||
@@ -150,20 +151,46 @@ async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch
|
|||||||
|
|
||||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
|
|
||||||
await loop._connect_mcp()
|
await provider.connect()
|
||||||
await loop._connect_mcp()
|
await provider.connect()
|
||||||
|
|
||||||
assert attempts == 2
|
assert attempts == 2
|
||||||
assert loop._mcp_stacks == {}
|
assert provider.connected_server_names == set()
|
||||||
|
assert provider.runtime_status() == {"test": "failed"}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_agent_loop_run_closes_mcp_from_connection_owner_task(
|
async def test_connect_mcp_does_not_report_failure_before_oauth_authorization(
|
||||||
tmp_path,
|
tmp_path,
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
):
|
):
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"playwright": object()})
|
cfg = MCPServerConfig(
|
||||||
connected = asyncio.Event()
|
type="streamableHttp",
|
||||||
|
auth="oauth",
|
||||||
|
url="https://mcp.example.com/mcp",
|
||||||
|
)
|
||||||
|
provider, _registry = _make_provider(mcp_servers={"oauth-app": cfg})
|
||||||
|
connect = AsyncMock()
|
||||||
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", connect)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.agent.tools.mcp_oauth.mcp_oauth_has_credentials",
|
||||||
|
lambda _name, _url: False,
|
||||||
|
)
|
||||||
|
|
||||||
|
await provider.connect()
|
||||||
|
|
||||||
|
connect.assert_not_awaited()
|
||||||
|
assert provider.runtime_status() == {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mcp_provider_closes_connections_independently_from_agent_loop(
|
||||||
|
tmp_path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
):
|
||||||
|
provider, registry = _make_provider(
|
||||||
|
mcp_servers={"playwright": _stdio_server("playwright")}
|
||||||
|
)
|
||||||
owner_tasks: list[asyncio.Task | None] = []
|
owner_tasks: list[asyncio.Task | None] = []
|
||||||
closed_tasks: list[asyncio.Task | None] = []
|
closed_tasks: list[asyncio.Task | None] = []
|
||||||
|
|
||||||
@@ -178,40 +205,38 @@ async def test_agent_loop_run_closes_mcp_from_connection_owner_task(
|
|||||||
|
|
||||||
async def _fake_connect(servers, _registry):
|
async def _fake_connect(servers, _registry):
|
||||||
stacks = {name: _OwnerCheckedStack() for name in servers}
|
stacks = {name: _OwnerCheckedStack() for name in servers}
|
||||||
connected.set()
|
|
||||||
return stacks
|
return stacks
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
|
|
||||||
task = asyncio.create_task(loop.run())
|
await provider.connect()
|
||||||
await asyncio.wait_for(connected.wait(), timeout=1)
|
registry.register(_FakeMcpTool("mcp_playwright_search"))
|
||||||
loop.stop()
|
await provider.aclose()
|
||||||
task.cancel()
|
|
||||||
await asyncio.gather(task, return_exceptions=True)
|
|
||||||
|
|
||||||
assert owner_tasks
|
assert owner_tasks
|
||||||
assert closed_tasks == owner_tasks
|
assert closed_tasks == owner_tasks
|
||||||
assert loop._mcp_stacks == {}
|
assert provider.connected_server_names == set()
|
||||||
|
assert registry.get("mcp_playwright_search") is None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_close_server_ignores_server_cancelled_error(tmp_path):
|
async def test_close_server_ignores_server_cancelled_error(tmp_path):
|
||||||
loop = _make_loop(tmp_path)
|
provider, _registry = _make_provider()
|
||||||
|
|
||||||
class _ServerCancelledStack:
|
class _ServerCancelledStack:
|
||||||
async def aclose(self) -> None:
|
async def aclose(self) -> None:
|
||||||
raise asyncio.CancelledError()
|
raise asyncio.CancelledError()
|
||||||
|
|
||||||
loop._mcp_stacks = {"test": _ServerCancelledStack()}
|
provider._connections = {"test": _ServerCancelledStack()}
|
||||||
|
|
||||||
await mcp_runtime._close_server(loop, "test")
|
await provider._close_server("test")
|
||||||
|
|
||||||
assert loop._mcp_stacks == {}
|
assert provider.connected_server_names == set()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_close_mcp_servers_continues_after_server_cancelled_error(tmp_path):
|
async def test_provider_close_continues_after_server_cancelled_error(tmp_path):
|
||||||
loop = _make_loop(tmp_path)
|
provider, _registry = _make_provider()
|
||||||
closed: list[str] = []
|
closed: list[str] = []
|
||||||
|
|
||||||
class _ServerCancelledStack:
|
class _ServerCancelledStack:
|
||||||
@@ -222,21 +247,53 @@ async def test_close_mcp_servers_continues_after_server_cancelled_error(tmp_path
|
|||||||
async def aclose(self) -> None:
|
async def aclose(self) -> None:
|
||||||
closed.append("second")
|
closed.append("second")
|
||||||
|
|
||||||
loop._mcp_stacks = {
|
provider._connections = {
|
||||||
"first": _ServerCancelledStack(),
|
"first": _ServerCancelledStack(),
|
||||||
"second": _TrackedStack(),
|
"second": _TrackedStack(),
|
||||||
}
|
}
|
||||||
|
|
||||||
await mcp_runtime.close_mcp_servers(loop)
|
await provider.aclose()
|
||||||
|
|
||||||
assert closed == ["second"]
|
assert closed == ["second"]
|
||||||
assert loop._mcp_stacks == {}
|
assert provider.connected_server_names == set()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_provider_close_finishes_other_connections_before_propagating_cancellation(
|
||||||
|
tmp_path,
|
||||||
|
):
|
||||||
|
provider, _registry = _make_provider()
|
||||||
|
started = asyncio.Event()
|
||||||
|
closed: list[str] = []
|
||||||
|
|
||||||
|
class _BlockingStack:
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
started.set()
|
||||||
|
await asyncio.Event().wait()
|
||||||
|
|
||||||
|
class _TrackedStack:
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
closed.append("second")
|
||||||
|
|
||||||
|
provider._connections = {
|
||||||
|
"first": _BlockingStack(),
|
||||||
|
"second": _TrackedStack(),
|
||||||
|
}
|
||||||
|
task = asyncio.create_task(provider.aclose())
|
||||||
|
await asyncio.wait_for(started.wait(), timeout=1)
|
||||||
|
|
||||||
|
task.cancel()
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
assert closed == ["second"]
|
||||||
|
assert provider.connected_server_names == set()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize("close_all", [False, True], ids=["single", "all"])
|
@pytest.mark.parametrize("close_all", [False, True], ids=["single", "all"])
|
||||||
async def test_mcp_cleanup_re_raises_external_cancellation(tmp_path, close_all: bool):
|
async def test_mcp_cleanup_re_raises_external_cancellation(tmp_path, close_all: bool):
|
||||||
loop = _make_loop(tmp_path)
|
provider, _registry = _make_provider()
|
||||||
started = asyncio.Event()
|
started = asyncio.Event()
|
||||||
|
|
||||||
class _BlockingStack:
|
class _BlockingStack:
|
||||||
@@ -244,12 +301,12 @@ async def test_mcp_cleanup_re_raises_external_cancellation(tmp_path, close_all:
|
|||||||
started.set()
|
started.set()
|
||||||
await asyncio.Event().wait()
|
await asyncio.Event().wait()
|
||||||
|
|
||||||
loop._mcp_stacks = {"test": _BlockingStack()}
|
provider._connections = {"test": _BlockingStack()}
|
||||||
|
|
||||||
if close_all:
|
if close_all:
|
||||||
task = asyncio.create_task(mcp_runtime.close_mcp_servers(loop))
|
task = asyncio.create_task(provider.aclose())
|
||||||
else:
|
else:
|
||||||
task = asyncio.create_task(mcp_runtime._close_server(loop, "test"))
|
task = asyncio.create_task(provider._close_server("test"))
|
||||||
await asyncio.wait_for(started.wait(), timeout=1)
|
await asyncio.wait_for(started.wait(), timeout=1)
|
||||||
task.cancel()
|
task.cancel()
|
||||||
|
|
||||||
@@ -287,41 +344,38 @@ async def test_reload_mcp_servers_adds_and_removes_tools_without_restart(
|
|||||||
return stacks
|
return stacks
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
loop = _make_loop(tmp_path, mcp_servers={})
|
provider, registry = _make_provider(mcp_servers={})
|
||||||
|
|
||||||
added = await mcp_runtime.reload_servers(loop, loop.tools)
|
added = await provider.reload()
|
||||||
|
|
||||||
assert added["ok"] is True
|
assert added["ok"] is True
|
||||||
assert added["added"] == ["browserbase"]
|
assert added["added"] == ["browserbase"]
|
||||||
assert loop.tools.has("mcp_browserbase_navigate")
|
assert registry.has("mcp_browserbase_navigate")
|
||||||
assert "browserbase" in loop._mcp_stacks
|
assert provider.connected_server_names == {"browserbase"}
|
||||||
|
|
||||||
config = load_config()
|
config = load_config()
|
||||||
del config.tools.mcp_servers["browserbase"]
|
del config.tools.mcp_servers["browserbase"]
|
||||||
save_config(config)
|
save_config(config)
|
||||||
|
|
||||||
removed = await mcp_runtime.reload_servers(loop, loop.tools)
|
removed = await provider.reload()
|
||||||
|
|
||||||
assert removed["ok"] is True
|
assert removed["ok"] is True
|
||||||
assert removed["removed"] == ["browserbase"]
|
assert removed["removed"] == ["browserbase"]
|
||||||
assert not loop.tools.has("mcp_browserbase_navigate")
|
assert not registry.has("mcp_browserbase_navigate")
|
||||||
assert "browserbase" not in loop._mcp_stacks
|
assert provider.connected_server_names == set()
|
||||||
assert closed == ["browserbase"]
|
assert closed == ["browserbase"]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_request_mcp_reload_reaches_runtime_control_without_restart(
|
async def test_reload_is_a_direct_provider_operation_without_an_agent_loop(
|
||||||
tmp_path,
|
tmp_path,
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
):
|
):
|
||||||
config_path = tmp_path / "config.json"
|
browserbase = MCPServerConfig(
|
||||||
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
|
||||||
config = load_config()
|
|
||||||
config.tools.mcp_servers["browserbase"] = MCPServerConfig(
|
|
||||||
type="stdio",
|
type="stdio",
|
||||||
command="browserbase-mcp",
|
command="browserbase-mcp",
|
||||||
)
|
)
|
||||||
save_config(config)
|
configured: dict[str, MCPServerConfig] = {"browserbase": browserbase}
|
||||||
|
|
||||||
closed: list[str] = []
|
closed: list[str] = []
|
||||||
|
|
||||||
@@ -339,37 +393,68 @@ async def test_request_mcp_reload_reaches_runtime_control_without_restart(
|
|||||||
return stacks
|
return stacks
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
loop = _make_loop(tmp_path, mcp_servers={})
|
registry = ToolRegistry()
|
||||||
|
provider = MCPProvider({}, registry, server_loader=lambda: configured)
|
||||||
|
|
||||||
async def _handle_one_runtime_control() -> None:
|
result = await provider.reload()
|
||||||
msg = await loop.bus.consume_inbound()
|
|
||||||
handled = await mcp_runtime.handle_runtime_control(loop, msg, loop.tools)
|
|
||||||
assert handled is True
|
|
||||||
|
|
||||||
consumer = asyncio.create_task(_handle_one_runtime_control())
|
|
||||||
result = await mcp_runtime.request_mcp_reload(loop.bus, timeout=2.0)
|
|
||||||
await consumer
|
|
||||||
|
|
||||||
assert result["ok"] is True
|
assert result["ok"] is True
|
||||||
assert result["added"] == ["browserbase"]
|
assert result["added"] == ["browserbase"]
|
||||||
assert result["requires_restart"] is False
|
assert result["requires_restart"] is False
|
||||||
assert loop.tools.has("mcp_browserbase_navigate")
|
assert registry.has("mcp_browserbase_navigate")
|
||||||
|
|
||||||
config = load_config()
|
configured = {}
|
||||||
del config.tools.mcp_servers["browserbase"]
|
|
||||||
save_config(config)
|
|
||||||
|
|
||||||
consumer = asyncio.create_task(_handle_one_runtime_control())
|
result = await provider.reload()
|
||||||
result = await mcp_runtime.request_mcp_reload(loop.bus, timeout=2.0)
|
|
||||||
await consumer
|
|
||||||
|
|
||||||
assert result["ok"] is True
|
assert result["ok"] is True
|
||||||
assert result["removed"] == ["browserbase"]
|
assert result["removed"] == ["browserbase"]
|
||||||
assert result["requires_restart"] is False
|
assert result["requires_restart"] is False
|
||||||
assert not loop.tools.has("mcp_browserbase_navigate")
|
assert not registry.has("mcp_browserbase_navigate")
|
||||||
assert closed == ["browserbase"]
|
assert closed == ["browserbase"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reload_timeout_marks_attempted_server_failed_and_allows_retry(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
):
|
||||||
|
server = _stdio_server("slow-mcp")
|
||||||
|
started = asyncio.Event()
|
||||||
|
attempts = 0
|
||||||
|
|
||||||
|
async def _fake_connect(servers, _registry):
|
||||||
|
nonlocal attempts
|
||||||
|
attempts += 1
|
||||||
|
if attempts == 1:
|
||||||
|
started.set()
|
||||||
|
await asyncio.Event().wait()
|
||||||
|
stack = AsyncExitStack()
|
||||||
|
await stack.__aenter__()
|
||||||
|
return {name: stack for name in servers}
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
|
provider = MCPProvider(
|
||||||
|
{"test": server},
|
||||||
|
ToolRegistry(),
|
||||||
|
server_loader=lambda: {"test": server},
|
||||||
|
)
|
||||||
|
|
||||||
|
reload_task = asyncio.create_task(provider.reload())
|
||||||
|
await asyncio.wait_for(started.wait(), timeout=1.0)
|
||||||
|
with pytest.raises(asyncio.TimeoutError):
|
||||||
|
await asyncio.wait_for(reload_task, timeout=0.01)
|
||||||
|
|
||||||
|
assert provider.connected_server_names == set()
|
||||||
|
assert provider.runtime_status() == {"test": "failed"}
|
||||||
|
|
||||||
|
result = await provider.reload()
|
||||||
|
|
||||||
|
assert result["ok"] is True
|
||||||
|
assert provider.connected_server_names == {"test"}
|
||||||
|
assert provider.runtime_status() == {"test": "connected"}
|
||||||
|
await provider.aclose()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_reload_mcp_servers_retries_configured_server_without_live_stack(
|
async def test_reload_mcp_servers_retries_configured_server_without_live_stack(
|
||||||
tmp_path,
|
tmp_path,
|
||||||
@@ -394,16 +479,64 @@ async def test_reload_mcp_servers_retries_configured_server_without_live_stack(
|
|||||||
return stacks
|
return stacks
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"browserbase": config.tools.mcp_servers["browserbase"]})
|
provider, registry = _make_provider(
|
||||||
|
mcp_servers={"browserbase": config.tools.mcp_servers["browserbase"]}
|
||||||
|
)
|
||||||
|
|
||||||
result = await mcp_runtime.reload_servers(loop, loop.tools)
|
result = await provider.reload()
|
||||||
|
|
||||||
assert result["ok"] is True
|
assert result["ok"] is True
|
||||||
assert result["added"] == []
|
assert result["added"] == []
|
||||||
assert result["changed"] == []
|
assert result["changed"] == []
|
||||||
assert result["retried"] == ["browserbase"]
|
assert result["retried"] == ["browserbase"]
|
||||||
assert loop.tools.has("mcp_browserbase_navigate")
|
assert registry.has("mcp_browserbase_navigate")
|
||||||
await loop.close_mcp()
|
await provider.aclose()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reload_mcp_servers_skips_oauth_server_waiting_for_authorization(
|
||||||
|
tmp_path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
):
|
||||||
|
config_path = tmp_path / "config.json"
|
||||||
|
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
||||||
|
config = load_config()
|
||||||
|
notion = MCPServerConfig(
|
||||||
|
type="streamableHttp",
|
||||||
|
auth="oauth",
|
||||||
|
url="https://mcp.notion.test/mcp",
|
||||||
|
)
|
||||||
|
linear = MCPServerConfig(
|
||||||
|
type="streamableHttp",
|
||||||
|
auth="oauth",
|
||||||
|
url="https://mcp.linear.test/mcp",
|
||||||
|
)
|
||||||
|
config.tools.mcp_servers.update({"notion": notion, "linear": linear})
|
||||||
|
save_config(config)
|
||||||
|
|
||||||
|
attempted: list[str] = []
|
||||||
|
|
||||||
|
async def _fake_connect(servers, _registry):
|
||||||
|
attempted.extend(servers)
|
||||||
|
stack = AsyncExitStack()
|
||||||
|
await stack.__aenter__()
|
||||||
|
return {"linear": stack}
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.agent.tools.mcp_oauth.mcp_oauth_has_credentials",
|
||||||
|
lambda name, _url: name == "linear",
|
||||||
|
)
|
||||||
|
provider, _registry = _make_provider(mcp_servers={"notion": notion})
|
||||||
|
|
||||||
|
result = await provider.reload()
|
||||||
|
|
||||||
|
assert attempted == ["linear"]
|
||||||
|
assert result["ok"] is True
|
||||||
|
assert result["failed"] == []
|
||||||
|
assert result["retried"] == []
|
||||||
|
assert result["connected"] == ["linear"]
|
||||||
|
await provider.aclose()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -411,7 +544,9 @@ async def test_mcp_tool_reconnects_after_session_terminated(
|
|||||||
tmp_path,
|
tmp_path,
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
):
|
):
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"remote": object()})
|
provider, registry = _make_provider(
|
||||||
|
mcp_servers={"remote": _stdio_server("remote")}
|
||||||
|
)
|
||||||
closed: list[str] = []
|
closed: list[str] = []
|
||||||
sessions: list[Any] = []
|
sessions: list[Any] = []
|
||||||
connect_count = 0
|
connect_count = 0
|
||||||
@@ -454,8 +589,8 @@ async def test_mcp_tool_reconnects_after_session_terminated(
|
|||||||
|
|
||||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
|
|
||||||
await loop._connect_mcp()
|
await provider.connect()
|
||||||
old_tool = loop.tools.get("mcp_remote_quote")
|
old_tool = registry.get("mcp_remote_quote")
|
||||||
assert isinstance(old_tool, MCPToolWrapper)
|
assert isinstance(old_tool, MCPToolWrapper)
|
||||||
|
|
||||||
output = await old_tool.execute(symbol="AAPL")
|
output = await old_tool.execute(symbol="AAPL")
|
||||||
@@ -465,8 +600,8 @@ async def test_mcp_tool_reconnects_after_session_terminated(
|
|||||||
assert closed == ["remote"]
|
assert closed == ["remote"]
|
||||||
assert sessions[0].call_count == 1
|
assert sessions[0].call_count == 1
|
||||||
assert sessions[1].call_count == 1
|
assert sessions[1].call_count == 1
|
||||||
assert "remote" in loop._mcp_stacks
|
assert provider.connected_server_names == {"remote"}
|
||||||
assert loop.tools.get("mcp_remote_quote") is not old_tool
|
assert registry.get("mcp_remote_quote") is not old_tool
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -474,7 +609,9 @@ async def test_mcp_reconnect_handler_uses_sanitized_server_prefix(
|
|||||||
tmp_path,
|
tmp_path,
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
):
|
):
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"remote_": object()})
|
provider, registry = _make_provider(
|
||||||
|
mcp_servers={"remote_": _stdio_server("remote")}
|
||||||
|
)
|
||||||
connect_count = 0
|
connect_count = 0
|
||||||
|
|
||||||
class _FakeSession:
|
class _FakeSession:
|
||||||
@@ -507,15 +644,15 @@ async def test_mcp_reconnect_handler_uses_sanitized_server_prefix(
|
|||||||
|
|
||||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
|
|
||||||
await loop._connect_mcp()
|
await provider.connect()
|
||||||
old_tool = loop.tools.get("mcp_remote_quote")
|
old_tool = registry.get("mcp_remote_quote")
|
||||||
assert isinstance(old_tool, MCPToolWrapper)
|
assert isinstance(old_tool, MCPToolWrapper)
|
||||||
|
|
||||||
output = await old_tool.execute()
|
output = await old_tool.execute()
|
||||||
|
|
||||||
assert output == "recovered"
|
assert output == "recovered"
|
||||||
assert connect_count == 2
|
assert connect_count == 2
|
||||||
assert loop.tools.get("mcp_remote_quote") is not old_tool
|
assert registry.get("mcp_remote_quote") is not old_tool
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -523,7 +660,9 @@ async def test_concurrent_mcp_reconnect_reuses_fresh_session(
|
|||||||
tmp_path,
|
tmp_path,
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
):
|
):
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"remote": object()})
|
provider, registry = _make_provider(
|
||||||
|
mcp_servers={"remote": _stdio_server("remote")}
|
||||||
|
)
|
||||||
closed: list[str] = []
|
closed: list[str] = []
|
||||||
connect_count = 0
|
connect_count = 0
|
||||||
|
|
||||||
@@ -567,9 +706,9 @@ async def test_concurrent_mcp_reconnect_reuses_fresh_session(
|
|||||||
|
|
||||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||||
|
|
||||||
await loop._connect_mcp()
|
await provider.connect()
|
||||||
old_alpha = loop.tools.get("mcp_remote_resource_alpha")
|
old_alpha = registry.get("mcp_remote_resource_alpha")
|
||||||
old_beta = loop.tools.get("mcp_remote_resource_beta")
|
old_beta = registry.get("mcp_remote_resource_beta")
|
||||||
assert isinstance(old_alpha, MCPResourceWrapper)
|
assert isinstance(old_alpha, MCPResourceWrapper)
|
||||||
assert isinstance(old_beta, MCPResourceWrapper)
|
assert isinstance(old_beta, MCPResourceWrapper)
|
||||||
|
|
||||||
|
|||||||
@@ -15,15 +15,13 @@ import asyncio
|
|||||||
import multiprocessing
|
import multiprocessing
|
||||||
import socket
|
import socket
|
||||||
import time
|
import time
|
||||||
from unittest.mock import MagicMock
|
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
|
||||||
from nanobot.agent.tools import mcp as mcp_module
|
from nanobot.agent.tools import mcp as mcp_module
|
||||||
from nanobot.agent.tools.mcp import MCPToolWrapper
|
from nanobot.agent.tools.mcp import MCPProvider, MCPToolWrapper
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.config.schema import MCPServerConfig
|
from nanobot.config.schema import MCPServerConfig
|
||||||
from nanobot.security import network as security_network
|
from nanobot.security import network as security_network
|
||||||
|
|
||||||
@@ -113,18 +111,9 @@ def mcp_server_url():
|
|||||||
process.join(timeout=2.0)
|
process.join(timeout=2.0)
|
||||||
|
|
||||||
|
|
||||||
def _make_loop(tmp_path, *, mcp_servers: dict) -> AgentLoop:
|
def _make_provider(*, mcp_servers: dict) -> tuple[MCPProvider, ToolRegistry]:
|
||||||
bus = MessageBus()
|
registry = ToolRegistry()
|
||||||
provider = MagicMock()
|
return MCPProvider(mcp_servers, registry), registry
|
||||||
provider.get_default_model.return_value = "test-model"
|
|
||||||
provider.generation.max_tokens = 4096
|
|
||||||
return AgentLoop(
|
|
||||||
bus=bus,
|
|
||||||
provider=provider,
|
|
||||||
workspace=tmp_path,
|
|
||||||
model="test-model",
|
|
||||||
mcp_servers=mcp_servers,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(autouse=True)
|
@pytest.fixture(autouse=True)
|
||||||
@@ -170,12 +159,12 @@ async def test_mcp_reconnect_after_session_timeout(tmp_path, mcp_server_url):
|
|||||||
tool_timeout=_TOOL_TIMEOUT_SECONDS,
|
tool_timeout=_TOOL_TIMEOUT_SECONDS,
|
||||||
enabled_tools=["*"],
|
enabled_tools=["*"],
|
||||||
)
|
)
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
|
provider, registry = _make_provider(mcp_servers={"repro": cfg})
|
||||||
|
|
||||||
await asyncio.create_task(loop._connect_mcp())
|
await asyncio.create_task(provider.connect())
|
||||||
assert "repro" in loop._mcp_stacks
|
assert provider.connected_server_names == {"repro"}
|
||||||
|
|
||||||
tool = loop.tools.get("mcp_repro_greet")
|
tool = registry.get("mcp_repro_greet")
|
||||||
assert isinstance(tool, MCPToolWrapper)
|
assert isinstance(tool, MCPToolWrapper)
|
||||||
|
|
||||||
output = await asyncio.create_task(tool.execute(name="first"))
|
output = await asyncio.create_task(tool.execute(name="first"))
|
||||||
@@ -187,7 +176,7 @@ async def test_mcp_reconnect_after_session_timeout(tmp_path, mcp_server_url):
|
|||||||
output = await asyncio.create_task(tool.execute(name="second"))
|
output = await asyncio.create_task(tool.execute(name="second"))
|
||||||
assert "Hello, second" in output
|
assert "Hello, second" in output
|
||||||
|
|
||||||
await asyncio.create_task(loop.close_mcp())
|
await asyncio.create_task(provider.aclose())
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -203,10 +192,10 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(
|
|||||||
tool_timeout=_TOOL_TIMEOUT_SECONDS,
|
tool_timeout=_TOOL_TIMEOUT_SECONDS,
|
||||||
enabled_tools=["*"],
|
enabled_tools=["*"],
|
||||||
)
|
)
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
|
provider, registry = _make_provider(mcp_servers={"repro": cfg})
|
||||||
|
|
||||||
await asyncio.create_task(loop._connect_mcp())
|
await asyncio.create_task(provider.connect())
|
||||||
tool = loop.tools.get("mcp_repro_greet")
|
tool = registry.get("mcp_repro_greet")
|
||||||
assert isinstance(tool, MCPToolWrapper)
|
assert isinstance(tool, MCPToolWrapper)
|
||||||
|
|
||||||
await asyncio.create_task(tool.execute(name="first"))
|
await asyncio.create_task(tool.execute(name="first"))
|
||||||
@@ -224,7 +213,7 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(
|
|||||||
monkeypatch.setattr(mcp_module, "connect_mcp_servers", gated_connect)
|
monkeypatch.setattr(mcp_module, "connect_mcp_servers", gated_connect)
|
||||||
call_task = asyncio.create_task(tool.execute(name="second"))
|
call_task = asyncio.create_task(tool.execute(name="second"))
|
||||||
await asyncio.wait_for(reconnect_started.wait(), timeout=5)
|
await asyncio.wait_for(reconnect_started.wait(), timeout=5)
|
||||||
close_task = asyncio.create_task(loop.close_mcp())
|
close_task = asyncio.create_task(provider.aclose())
|
||||||
await asyncio.sleep(0)
|
await asyncio.sleep(0)
|
||||||
finish_reconnect.set()
|
finish_reconnect.set()
|
||||||
|
|
||||||
@@ -245,4 +234,4 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(
|
|||||||
unhandled.append(exc)
|
unhandled.append(exc)
|
||||||
|
|
||||||
assert not unhandled, f"Unhandled exception leaked during reconnect/shutdown: {unhandled[0]}"
|
assert not unhandled, f"Unhandled exception leaked during reconnect/shutdown: {unhandled[0]}"
|
||||||
assert loop._mcp_stacks == {}
|
assert provider.connected_server_names == set()
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.bus.runtime_events import RuntimeModelChanged
|
from nanobot.bus.runtime_events import RuntimeModelChanged
|
||||||
from nanobot.config.errors import ConfigLoadError
|
from nanobot.config.errors import ConfigLoadError
|
||||||
@@ -312,7 +313,11 @@ def test_settings_context_window_refreshes_runtime_state(
|
|||||||
def loader(*, preset_name: str | None = None) -> ProviderSnapshot:
|
def loader(*, preset_name: str | None = None) -> ProviderSnapshot:
|
||||||
return load_provider_snapshot(config_path, preset_name=preset_name)
|
return load_provider_snapshot(config_path, preset_name=preset_name)
|
||||||
|
|
||||||
loop = AgentLoop.from_config(config, provider_snapshot_loader=loader)
|
loop = AgentLoop.from_config(
|
||||||
|
config,
|
||||||
|
tool_registry=ToolRegistry(),
|
||||||
|
provider_snapshot_loader=loader,
|
||||||
|
)
|
||||||
|
|
||||||
payload = update_agent_settings({"context_window_tokens": ["262144"]})
|
payload = update_agent_settings({"context_window_tokens": ["262144"]})
|
||||||
loop.runtime_resolver.invalidate()
|
loop.runtime_resolver.invalidate()
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ import pytest
|
|||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.agent.tools.context import RequestContext, request_context
|
from nanobot.agent.tools.context import RequestContext, request_context
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
|
from nanobot.agent.tools.runtime_control import AgentRuntimeControl
|
||||||
from nanobot.agent.tools.self import MyTool
|
from nanobot.agent.tools.self import MyTool
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.config.schema import ModelPresetConfig
|
from nanobot.config.schema import ModelPresetConfig
|
||||||
@@ -34,6 +36,13 @@ def _make_loop(tmp_path, presets=None, active_preset=None):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _my_tool(loop: AgentLoop) -> MyTool:
|
||||||
|
return MyTool(
|
||||||
|
runtime_control=AgentRuntimeControl(loop),
|
||||||
|
modify_allowed=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_model_preset_getter_none_when_not_set(tmp_path) -> None:
|
def test_model_preset_getter_none_when_not_set(tmp_path) -> None:
|
||||||
loop = _make_loop(tmp_path)
|
loop = _make_loop(tmp_path)
|
||||||
assert loop.model_preset is None
|
assert loop.model_preset is None
|
||||||
@@ -240,7 +249,7 @@ def test_self_tool_inspect_shows_model_preset(tmp_path) -> None:
|
|||||||
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
||||||
}
|
}
|
||||||
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
|
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
|
||||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
tool = _my_tool(loop)
|
||||||
output = tool._inspect_all()
|
output = tool._inspect_all()
|
||||||
assert "model_preset: 'fast'" in output
|
assert "model_preset: 'fast'" in output
|
||||||
|
|
||||||
@@ -250,7 +259,7 @@ def test_self_tool_set_model_preset_via_modify(tmp_path) -> None:
|
|||||||
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
||||||
}
|
}
|
||||||
loop = _make_loop(tmp_path, presets=presets)
|
loop = _make_loop(tmp_path, presets=presets)
|
||||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
tool = _my_tool(loop)
|
||||||
result = tool._modify("model_preset", "fast")
|
result = tool._modify("model_preset", "fast")
|
||||||
assert "Error" not in result
|
assert "Error" not in result
|
||||||
assert loop.model_preset == "fast"
|
assert loop.model_preset == "fast"
|
||||||
@@ -263,7 +272,7 @@ def test_self_tool_set_model_preset_switches_back_to_default(tmp_path) -> None:
|
|||||||
"fast": ModelPresetConfig(model="openai/gpt-4.1", context_window_tokens=32_768),
|
"fast": ModelPresetConfig(model="openai/gpt-4.1", context_window_tokens=32_768),
|
||||||
}
|
}
|
||||||
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
|
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
|
||||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
result = tool._modify("model_preset", "default")
|
result = tool._modify("model_preset", "default")
|
||||||
|
|
||||||
@@ -280,7 +289,7 @@ def test_self_tool_set_model_preset_unknown_lists_available(tmp_path) -> None:
|
|||||||
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
||||||
}
|
}
|
||||||
loop = _make_loop(tmp_path, presets=presets)
|
loop = _make_loop(tmp_path, presets=presets)
|
||||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
result = tool._modify("model_preset", "missing")
|
result = tool._modify("model_preset", "missing")
|
||||||
|
|
||||||
@@ -295,7 +304,7 @@ def test_self_tool_sets_model_preset_for_current_session(tmp_path) -> None:
|
|||||||
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
||||||
}
|
}
|
||||||
loop = _make_loop(tmp_path, presets=presets)
|
loop = _make_loop(tmp_path, presets=presets)
|
||||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
with request_context(RequestContext(
|
with request_context(RequestContext(
|
||||||
channel="cli",
|
channel="cli",
|
||||||
@@ -318,7 +327,7 @@ def test_self_tool_reports_session_preset_provider_configuration_error(tmp_path)
|
|||||||
loop.set_session_model_preset = MagicMock(
|
loop.set_session_model_preset = MagicMock(
|
||||||
side_effect=ValueError("No API key configured for provider 'openai'.")
|
side_effect=ValueError("No API key configured for provider 'openai'.")
|
||||||
)
|
)
|
||||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
with request_context(RequestContext(
|
with request_context(RequestContext(
|
||||||
channel="cli",
|
channel="cli",
|
||||||
@@ -343,7 +352,7 @@ def test_self_tool_rejects_instance_runtime_changes_in_session(
|
|||||||
value: object,
|
value: object,
|
||||||
) -> None:
|
) -> None:
|
||||||
loop = _make_loop(tmp_path)
|
loop = _make_loop(tmp_path)
|
||||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
tool = _my_tool(loop)
|
||||||
session = loop.sessions.get_or_create("cli:one")
|
session = loop.sessions.get_or_create("cli:one")
|
||||||
|
|
||||||
with request_context(RequestContext(
|
with request_context(RequestContext(
|
||||||
@@ -366,7 +375,7 @@ def test_self_tool_set_model_clears_active_preset(tmp_path) -> None:
|
|||||||
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
|
||||||
}
|
}
|
||||||
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
|
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
|
||||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
tool = _my_tool(loop)
|
||||||
result = tool._modify("model", "anthropic/claude-opus-4-5")
|
result = tool._modify("model", "anthropic/claude-opus-4-5")
|
||||||
assert "Error" not in result
|
assert "Error" not in result
|
||||||
assert loop.model_preset is None
|
assert loop.model_preset is None
|
||||||
@@ -382,7 +391,7 @@ def test_from_config_injects_default_preset(tmp_path) -> None:
|
|||||||
})
|
})
|
||||||
fake_provider = _provider("openai/gpt-4.1")
|
fake_provider = _provider("openai/gpt-4.1")
|
||||||
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
|
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
|
||||||
loop = AgentLoop.from_config(config)
|
loop = AgentLoop.from_config(config, tool_registry=ToolRegistry())
|
||||||
assert loop.model == "openai/gpt-4.1"
|
assert loop.model == "openai/gpt-4.1"
|
||||||
assert loop.model_preset is None
|
assert loop.model_preset is None
|
||||||
assert "default" in loop.model_presets
|
assert "default" in loop.model_presets
|
||||||
@@ -399,7 +408,7 @@ def test_from_config_static_preset_loader_does_not_enable_hot_reload(tmp_path) -
|
|||||||
})
|
})
|
||||||
fake_provider = _provider("openai/gpt-4.1")
|
fake_provider = _provider("openai/gpt-4.1")
|
||||||
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
|
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
|
||||||
loop = AgentLoop.from_config(config)
|
loop = AgentLoop.from_config(config, tool_registry=ToolRegistry())
|
||||||
default_runtime = loop.runtime_resolver.runtime
|
default_runtime = loop.runtime_resolver.runtime
|
||||||
resolved = loop.runtime_resolver.resolve_preset("fast")
|
resolved = loop.runtime_resolver.resolve_preset("fast")
|
||||||
assert resolved.model == "openai/gpt-4.1-mini"
|
assert resolved.model == "openai/gpt-4.1-mini"
|
||||||
|
|||||||
@@ -266,6 +266,7 @@ def test_disabled_skills_excluded_from_list(tmp_path: Path) -> None:
|
|||||||
assert len(entries) == 1
|
assert len(entries) == 1
|
||||||
assert entries[0]["name"] == "beta"
|
assert entries[0]["name"] == "beta"
|
||||||
assert entries[0]["path"] == str(beta_path)
|
assert entries[0]["path"] == str(beta_path)
|
||||||
|
assert loader.load_skill("alpha") is None
|
||||||
|
|
||||||
|
|
||||||
def test_disabled_skills_empty_set_no_effect(tmp_path: Path) -> None:
|
def test_disabled_skills_empty_set_no_effect(tmp_path: Path) -> None:
|
||||||
|
|||||||
@@ -56,7 +56,7 @@ class TestHandleStop:
|
|||||||
assert "No active task" in out.content
|
assert "No active task" in out.content
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_close_mcp_cancels_active_turn_before_resources(self):
|
async def test_aclose_cancels_active_turn_before_resources(self):
|
||||||
loop, _bus = _make_loop()
|
loop, _bus = _make_loop()
|
||||||
events: list[str] = []
|
events: list[str] = []
|
||||||
|
|
||||||
@@ -76,14 +76,13 @@ class TestHandleStop:
|
|||||||
|
|
||||||
loop.subagents.close = close_subagents
|
loop.subagents.close = close_subagents
|
||||||
loop._exec_session_manager.close_all = AsyncMock()
|
loop._exec_session_manager.close_all = AsyncMock()
|
||||||
with patch("nanobot.agent.loop.agent_context.close_mcp", AsyncMock()):
|
await loop.aclose()
|
||||||
await loop.close_mcp()
|
|
||||||
|
|
||||||
assert events == ["turn_cancelled", "resources_closed"]
|
assert events == ["turn_cancelled", "resources_closed"]
|
||||||
assert task.cancelled()
|
assert task.cancelled()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_close_mcp_serializes_duplicate_cleanup(self):
|
async def test_aclose_serializes_duplicate_cleanup(self):
|
||||||
loop, _bus = _make_loop()
|
loop, _bus = _make_loop()
|
||||||
entered = asyncio.Event()
|
entered = asyncio.Event()
|
||||||
release = asyncio.Event()
|
release = asyncio.Event()
|
||||||
@@ -100,10 +99,9 @@ class TestHandleStop:
|
|||||||
|
|
||||||
loop.subagents.close = close_subagents
|
loop.subagents.close = close_subagents
|
||||||
loop._exec_session_manager.close_all = AsyncMock()
|
loop._exec_session_manager.close_all = AsyncMock()
|
||||||
with patch("nanobot.agent.loop.agent_context.close_mcp", AsyncMock()):
|
first = asyncio.create_task(loop.aclose())
|
||||||
first = asyncio.create_task(loop.close_mcp())
|
|
||||||
await entered.wait()
|
await entered.wait()
|
||||||
second = asyncio.create_task(loop.close_mcp())
|
second = asyncio.create_task(loop.aclose())
|
||||||
await asyncio.sleep(0)
|
await asyncio.sleep(0)
|
||||||
assert not second.done()
|
assert not second.done()
|
||||||
release.set()
|
release.set()
|
||||||
@@ -172,8 +170,7 @@ class TestDispatch:
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_run_logs_and_continues_after_leaked_cancelled_error(self, monkeypatch):
|
async def test_run_logs_and_continues_after_leaked_cancelled_error(self, monkeypatch):
|
||||||
loop, bus = _make_loop()
|
loop, bus = _make_loop()
|
||||||
loop._connect_mcp = AsyncMock()
|
loop.aclose = AsyncMock()
|
||||||
loop.close_mcp = AsyncMock()
|
|
||||||
loop.auto_compact.check_expired = MagicMock()
|
loop.auto_compact.check_expired = MagicMock()
|
||||||
warnings: list[str] = []
|
warnings: list[str] = []
|
||||||
calls = 0
|
calls = 0
|
||||||
|
|||||||
@@ -0,0 +1,225 @@
|
|||||||
|
"""Contract and security regressions for the MyTool runtime boundary."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.runtime_control import (
|
||||||
|
RUNTIME_COMMAND_KEYS,
|
||||||
|
RUNTIME_SNAPSHOT_KEYS,
|
||||||
|
AgentRuntimeControl,
|
||||||
|
RuntimeControl,
|
||||||
|
)
|
||||||
|
from nanobot.agent.tools.self import MyTool, MyToolConfig
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.config.schema import ToolsConfig
|
||||||
|
|
||||||
|
|
||||||
|
def _make_loop(tmp_path: Path, *, allow_set: bool = False) -> AgentLoop:
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
tools_config = ToolsConfig(my=MyToolConfig(allow_set=allow_set))
|
||||||
|
return AgentLoop(
|
||||||
|
bus=MessageBus(),
|
||||||
|
provider=provider,
|
||||||
|
workspace=tmp_path,
|
||||||
|
model="test-model",
|
||||||
|
tools_config=tools_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _my_tool(loop: AgentLoop) -> MyTool:
|
||||||
|
tool = loop.tools.get("my")
|
||||||
|
assert isinstance(tool, MyTool)
|
||||||
|
return tool
|
||||||
|
|
||||||
|
|
||||||
|
def test_agent_loop_assembles_my_tool_with_runtime_control(tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
|
assert isinstance(tool._runtime_control, RuntimeControl)
|
||||||
|
assert isinstance(tool._runtime_control, AgentRuntimeControl)
|
||||||
|
assert tool._runtime_control is not loop
|
||||||
|
assert not hasattr(tool, "_runtime_state")
|
||||||
|
|
||||||
|
|
||||||
|
def test_runtime_snapshot_has_exact_allowlist_and_redacts_secrets(tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
loop.web_config.search.api_key = "search-secret"
|
||||||
|
loop.web_config.proxy = "http://proxy-user:proxy-secret@proxy.example"
|
||||||
|
loop.unlisted_secret = "loop-secret"
|
||||||
|
|
||||||
|
snapshot = _my_tool(loop)._runtime_control.snapshot()
|
||||||
|
values = snapshot.as_mapping()
|
||||||
|
|
||||||
|
assert frozenset(values) == RUNTIME_SNAPSHOT_KEYS
|
||||||
|
assert RUNTIME_COMMAND_KEYS == frozenset({
|
||||||
|
"model",
|
||||||
|
"model_preset",
|
||||||
|
"max_iterations",
|
||||||
|
"context_window_tokens",
|
||||||
|
"provider_retry_mode",
|
||||||
|
"max_tool_result_chars",
|
||||||
|
"workspace",
|
||||||
|
})
|
||||||
|
assert "provider" not in values
|
||||||
|
assert "sessions" not in values
|
||||||
|
assert "restrict_to_workspace" not in values
|
||||||
|
assert "unlisted_secret" not in values
|
||||||
|
rendered = repr(values)
|
||||||
|
assert "search-secret" not in rendered
|
||||||
|
assert "proxy-secret" not in rendered
|
||||||
|
assert "loop-secret" not in rendered
|
||||||
|
assert snapshot.web_config["proxy"] == "<configured>"
|
||||||
|
|
||||||
|
|
||||||
|
def test_runtime_snapshot_is_detached_from_mutable_config(tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
control = _my_tool(loop)._runtime_control
|
||||||
|
snapshot = control.snapshot()
|
||||||
|
search = snapshot.web_config["search"]
|
||||||
|
assert isinstance(search, dict)
|
||||||
|
|
||||||
|
search["provider"] = "mutated"
|
||||||
|
snapshot.exec_config["allow_patterns"] = ["mutated"]
|
||||||
|
snapshot.tool_names.append("mutated")
|
||||||
|
|
||||||
|
refreshed = control.snapshot()
|
||||||
|
refreshed_search = refreshed.web_config["search"]
|
||||||
|
assert isinstance(refreshed_search, dict)
|
||||||
|
assert refreshed_search["provider"] == loop.web_config.search.provider
|
||||||
|
assert refreshed.exec_config["allow_patterns"] == loop.exec_config.allow_patterns
|
||||||
|
assert "mutated" not in refreshed.tool_names
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_unlisted_loop_attributes_cannot_be_read_or_modified(tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path, allow_set=True)
|
||||||
|
loop.unlisted_control_plane = "internal-secret"
|
||||||
|
original_workspace_root = loop.workspace_scopes.default_workspace
|
||||||
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
|
inspected = await tool.execute(action="check", key="unlisted_control_plane")
|
||||||
|
modified = await tool.execute(
|
||||||
|
action="set",
|
||||||
|
key="unlisted_control_plane",
|
||||||
|
value="scratch-value",
|
||||||
|
)
|
||||||
|
nested = await tool.execute(
|
||||||
|
action="set",
|
||||||
|
key="workspace_scopes.default_workspace",
|
||||||
|
value="elsewhere",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "internal-secret" not in inspected
|
||||||
|
assert "not found" in inspected
|
||||||
|
assert modified == "Set scratchpad.unlisted_control_plane = 'scratch-value'"
|
||||||
|
assert loop.unlisted_control_plane == "internal-secret"
|
||||||
|
assert "Error" in nested
|
||||||
|
assert loop.workspace_scopes.default_workspace == original_workspace_root
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_default_allow_set_and_public_parameter_schema_are_unchanged(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
|
assert ToolsConfig().my.allow_set is False
|
||||||
|
assert tool.parameters == {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"action": {
|
||||||
|
"type": "string",
|
||||||
|
"enum": ["check", "set"],
|
||||||
|
"description": "Action to perform",
|
||||||
|
},
|
||||||
|
"key": {
|
||||||
|
"type": "string",
|
||||||
|
"description": (
|
||||||
|
"Dot-path for check/set. Examples: 'max_iterations', 'workspace', "
|
||||||
|
"'provider_retry_mode'. Use 'request.channel', 'request.chat_id', or "
|
||||||
|
"'request.sender_id' for current routing metadata. Use 'model_preset' "
|
||||||
|
"to switch named model presets. For check without key, shows all "
|
||||||
|
"config values."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
"value": {
|
||||||
|
"description": (
|
||||||
|
"New value (for set). Type must match target (int for "
|
||||||
|
"max_iterations/context_window_tokens, str for model/model_preset)."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["action"],
|
||||||
|
}
|
||||||
|
assert "READ-ONLY MODE" in tool.description
|
||||||
|
result = await tool.execute(action="set", key="max_iterations", value=80)
|
||||||
|
assert result == "Error: set is disabled (tools.my.allow_set is false)"
|
||||||
|
assert loop.max_iterations != 80
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_allowlisted_commands_preserve_runtime_side_effects(tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path, allow_set=True)
|
||||||
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
|
max_iterations = await tool.execute(
|
||||||
|
action="set",
|
||||||
|
key="max_iterations",
|
||||||
|
value=80,
|
||||||
|
)
|
||||||
|
retry_mode = await tool.execute(
|
||||||
|
action="set",
|
||||||
|
key="provider_retry_mode",
|
||||||
|
value="persistent",
|
||||||
|
)
|
||||||
|
scratchpad = await tool.execute(
|
||||||
|
action="set",
|
||||||
|
key="preference",
|
||||||
|
value={"concise": True},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert max_iterations == "Set max_iterations = 80 (was 200)"
|
||||||
|
assert retry_mode == "Set provider_retry_mode = 'persistent' (was 'standard')"
|
||||||
|
assert scratchpad == "Set scratchpad.preference = {'concise': True}"
|
||||||
|
assert loop.max_iterations == 80
|
||||||
|
assert loop.subagents.max_iterations == 80
|
||||||
|
assert loop.provider_retry_mode == "persistent"
|
||||||
|
assert tool._runtime_control.snapshot().scratchpad == {
|
||||||
|
"preference": {"concise": True},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_registry_exposes_unchanged_my_tool_actions(tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path, allow_set=True)
|
||||||
|
|
||||||
|
checked = await loop.tools.execute("my", {"action": "check", "key": "model"})
|
||||||
|
changed = await loop.tools.execute(
|
||||||
|
"my",
|
||||||
|
{"action": "set", "key": "max_iterations", "value": 80},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert checked == "model: 'test-model'"
|
||||||
|
assert changed == "Set max_iterations = 80 (was 200)"
|
||||||
|
assert loop.max_iterations == 80
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_workspace_display_command_cannot_change_path_enforcement(tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path, allow_set=True)
|
||||||
|
tool = _my_tool(loop)
|
||||||
|
|
||||||
|
result = await tool.execute(action="set", key="workspace", value="elsewhere")
|
||||||
|
|
||||||
|
assert "Set workspace" in result
|
||||||
|
assert tool._runtime_control.snapshot().workspace == "elsewhere"
|
||||||
|
assert loop.workspace == tmp_path
|
||||||
|
assert loop.workspace_scopes.default_workspace == tmp_path
|
||||||
@@ -8,10 +8,12 @@ from types import MappingProxyType
|
|||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from pydantic import BaseModel
|
|
||||||
|
|
||||||
from nanobot.agent.tools.context import RequestContext, request_context
|
from nanobot.agent.tools.context import RequestContext, request_context
|
||||||
|
from nanobot.agent.tools.runtime_control import AgentRuntimeControl
|
||||||
from nanobot.agent.tools.self import MyTool
|
from nanobot.agent.tools.self import MyTool
|
||||||
|
from nanobot.agent.tools.shell import ExecToolConfig
|
||||||
|
from nanobot.agent.tools.web import WebSearchConfig, WebToolsConfig
|
||||||
from nanobot.config.schema import ModelPresetConfig
|
from nanobot.config.schema import ModelPresetConfig
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -27,13 +29,16 @@ def _make_mock_loop(**overrides):
|
|||||||
loop.workspace = Path("/tmp/workspace")
|
loop.workspace = Path("/tmp/workspace")
|
||||||
loop.restrict_to_workspace = False
|
loop.restrict_to_workspace = False
|
||||||
loop._start_time = 1000.0
|
loop._start_time = 1000.0
|
||||||
loop.exec_config = MagicMock()
|
loop.exec_config = ExecToolConfig()
|
||||||
loop.channels_config = MagicMock()
|
loop.channels_config = MagicMock()
|
||||||
loop._last_usage = {"prompt_tokens": 100, "completion_tokens": 50}
|
loop._last_usage = {"prompt_tokens": 100, "completion_tokens": 50}
|
||||||
loop._runtime_vars = {}
|
loop.last_usage = loop._last_usage
|
||||||
loop._current_iteration = 0
|
loop._current_iteration = 0
|
||||||
|
loop.current_iteration = loop._current_iteration
|
||||||
loop.provider_retry_mode = "standard"
|
loop.provider_retry_mode = "standard"
|
||||||
loop.max_tool_result_chars = 16000
|
loop.max_tool_result_chars = 16000
|
||||||
|
loop.model_preset = None
|
||||||
|
loop.model_presets = {}
|
||||||
loop._concurrency_gate = None
|
loop._concurrency_gate = None
|
||||||
loop._unified_session = False
|
loop._unified_session = False
|
||||||
loop._extra_hooks = []
|
loop._extra_hooks = []
|
||||||
@@ -45,9 +50,7 @@ def _make_mock_loop(**overrides):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# web_config mock — needed for check tests
|
# web_config mock — needed for check tests
|
||||||
loop.web_config = MagicMock()
|
loop.web_config = WebToolsConfig()
|
||||||
loop.web_config.enable = True
|
|
||||||
loop.web_config.search = MagicMock()
|
|
||||||
loop.web_config.search.api_key = "sk-secret-key-12345"
|
loop.web_config.search.api_key = "sk-secret-key-12345"
|
||||||
|
|
||||||
# Tools registry mock
|
# Tools registry mock
|
||||||
@@ -55,10 +58,13 @@ def _make_mock_loop(**overrides):
|
|||||||
loop.tools.tool_names = ["read_file", "write_file", "exec", "web_search", "self"]
|
loop.tools.tool_names = ["read_file", "write_file", "exec", "web_search", "self"]
|
||||||
loop.tools.has.side_effect = lambda n: n in loop.tools.tool_names
|
loop.tools.has.side_effect = lambda n: n in loop.tools.tool_names
|
||||||
loop.tools.get.return_value = None
|
loop.tools.get.return_value = None
|
||||||
|
loop.tool_names = loop.tools.tool_names
|
||||||
|
|
||||||
# SubagentManager mock
|
# SubagentManager mock
|
||||||
loop.subagents = MagicMock()
|
loop.subagents = MagicMock()
|
||||||
loop.subagents._running_tasks = {"abc123": MagicMock(done=MagicMock(return_value=False))}
|
loop.subagents._running_tasks = {"abc123": MagicMock(done=MagicMock(return_value=False))}
|
||||||
|
loop.subagents._task_statuses = {}
|
||||||
|
loop.subagents.runtime_statuses.side_effect = lambda: loop.subagents._task_statuses
|
||||||
loop.subagents.get_running_count = MagicMock(return_value=1)
|
loop.subagents.get_running_count = MagicMock(return_value=1)
|
||||||
|
|
||||||
for k, v in overrides.items():
|
for k, v in overrides.items():
|
||||||
@@ -67,10 +73,10 @@ def _make_mock_loop(**overrides):
|
|||||||
return loop
|
return loop
|
||||||
|
|
||||||
|
|
||||||
def _make_tool(runtime_state=None):
|
def _make_tool(loop=None):
|
||||||
if runtime_state is None:
|
if loop is None:
|
||||||
runtime_state = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
return MyTool(runtime_state=runtime_state)
|
return MyTool(runtime_control=AgentRuntimeControl(loop))
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -87,10 +93,10 @@ class TestInspectSummary:
|
|||||||
assert "context_window_tokens: 65536" in result
|
assert "context_window_tokens: 65536" in result
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inspect_includes_runtime_vars(self):
|
async def test_inspect_includes_scratchpad(self):
|
||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop._runtime_vars = {"task": "review"}
|
tool = _make_tool(loop=loop)
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool._runtime_control.set_scratchpad("task", "review", max_keys=64)
|
||||||
result = await tool.execute(action="check")
|
result = await tool.execute(action="check")
|
||||||
assert "task" in result
|
assert "task" in result
|
||||||
|
|
||||||
@@ -150,9 +156,7 @@ class TestInspectPathNavigation:
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inspect_config_subfield(self):
|
async def test_inspect_config_subfield(self):
|
||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop.web_config = MagicMock()
|
tool = _make_tool(loop=loop)
|
||||||
loop.web_config.enable = True
|
|
||||||
tool = _make_tool(runtime_state=loop)
|
|
||||||
result = await tool.execute(action="check", key="web_config.enable")
|
result = await tool.execute(action="check", key="web_config.enable")
|
||||||
assert "True" in result
|
assert "True" in result
|
||||||
|
|
||||||
@@ -160,7 +164,7 @@ class TestInspectPathNavigation:
|
|||||||
async def test_inspect_dict_key_via_dotpath(self):
|
async def test_inspect_dict_key_via_dotpath(self):
|
||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop._last_usage = {"prompt_tokens": 100, "completion_tokens": 50}
|
loop._last_usage = {"prompt_tokens": 100, "completion_tokens": 50}
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="check", key="_last_usage.prompt_tokens")
|
result = await tool.execute(action="check", key="_last_usage.prompt_tokens")
|
||||||
assert "100" in result
|
assert "100" in result
|
||||||
|
|
||||||
@@ -179,20 +183,16 @@ class TestInspectPathNavigation:
|
|||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inspect_nested_config_redacts_sensitive_scalar_fields(self):
|
async def test_inspect_nested_config_redacts_sensitive_scalar_fields(self):
|
||||||
class SearchConfig(BaseModel):
|
|
||||||
provider: str = "tavily"
|
|
||||||
api_key: str = "sk-test-secret"
|
|
||||||
base_url: str = ""
|
|
||||||
max_results: int = 5
|
|
||||||
|
|
||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop.web_config = MagicMock()
|
loop.web_config.search = WebSearchConfig(
|
||||||
loop.web_config.search = SearchConfig()
|
provider="tavily",
|
||||||
|
api_key="sk-test-secret",
|
||||||
|
)
|
||||||
tool = _make_tool(loop)
|
tool = _make_tool(loop)
|
||||||
|
|
||||||
result = await tool.execute(action="check", key="web_config.search")
|
result = await tool.execute(action="check", key="web_config.search")
|
||||||
|
|
||||||
assert "provider='tavily'" in result
|
assert "tavily" in result
|
||||||
assert "sk-test-secret" not in result
|
assert "sk-test-secret" not in result
|
||||||
assert "api_key" not in result.lower()
|
assert "api_key" not in result.lower()
|
||||||
|
|
||||||
@@ -209,14 +209,14 @@ class TestModifyRestricted:
|
|||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="max_iterations", value=80)
|
result = await tool.execute(action="set", key="max_iterations", value=80)
|
||||||
assert "Set max_iterations = 80" in result
|
assert "Set max_iterations = 80" in result
|
||||||
assert tool._runtime_state.max_iterations == 80
|
assert tool._runtime_control.snapshot().max_iterations == 80
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_restricted_out_of_range(self):
|
async def test_modify_restricted_out_of_range(self):
|
||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="max_iterations", value=0)
|
result = await tool.execute(action="set", key="max_iterations", value=0)
|
||||||
assert "Error" in result
|
assert "Error" in result
|
||||||
assert tool._runtime_state.max_iterations == 40
|
assert tool._runtime_control.snapshot().max_iterations == 40
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_restricted_max_exceeded(self):
|
async def test_modify_restricted_max_exceeded(self):
|
||||||
@@ -241,12 +241,12 @@ class TestModifyRestricted:
|
|||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="max_iterations", value="80")
|
result = await tool.execute(action="set", key="max_iterations", value="80")
|
||||||
assert "Set max_iterations" in result
|
assert "Set max_iterations" in result
|
||||||
assert tool._runtime_state.max_iterations == 80
|
assert tool._runtime_control.snapshot().max_iterations == 80
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_context_window_valid(self):
|
async def test_modify_context_window_valid(self):
|
||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="set", key="context_window_tokens", value=131072)
|
result = await tool.execute(action="set", key="context_window_tokens", value=131072)
|
||||||
assert "Set context_window_tokens" in result
|
assert "Set context_window_tokens" in result
|
||||||
assert loop.context_window_tokens == 131072
|
assert loop.context_window_tokens == 131072
|
||||||
@@ -324,15 +324,15 @@ class TestModifyFree:
|
|||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="provider_retry_mode", value="persistent")
|
result = await tool.execute(action="set", key="provider_retry_mode", value="persistent")
|
||||||
assert "Set provider_retry_mode" in result
|
assert "Set provider_retry_mode" in result
|
||||||
assert tool._runtime_state.provider_retry_mode == "persistent"
|
assert tool._runtime_control.snapshot().provider_retry_mode == "persistent"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_new_key_stores_in_runtime_vars(self):
|
async def test_modify_new_key_stores_in_scratchpad(self):
|
||||||
"""Modifying a non-existing attribute should store in _runtime_vars."""
|
"""Modifying an unknown key should store it in the scratchpad."""
|
||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="my_custom_var", value="hello")
|
result = await tool.execute(action="set", key="my_custom_var", value="hello")
|
||||||
assert "my_custom_var" in result
|
assert "my_custom_var" in result
|
||||||
assert tool._runtime_state._runtime_vars["my_custom_var"] == "hello"
|
assert tool._runtime_control.snapshot().scratchpad["my_custom_var"] == "hello"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_rejects_callable(self):
|
async def test_modify_rejects_callable(self):
|
||||||
@@ -351,14 +351,14 @@ class TestModifyFree:
|
|||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="items", value=[1, 2, 3])
|
result = await tool.execute(action="set", key="items", value=[1, 2, 3])
|
||||||
assert result == "Set scratchpad.items = [1, 2, 3]"
|
assert result == "Set scratchpad.items = [1, 2, 3]"
|
||||||
assert tool._runtime_state._runtime_vars["items"] == [1, 2, 3]
|
assert tool._runtime_control.snapshot().scratchpad["items"] == [1, 2, 3]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_allows_dict(self):
|
async def test_modify_allows_dict(self):
|
||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="data", value={"a": 1})
|
result = await tool.execute(action="set", key="data", value={"a": 1})
|
||||||
assert result == "Set scratchpad.data = {'a': 1}"
|
assert result == "Set scratchpad.data = {'a': 1}"
|
||||||
assert tool._runtime_state._runtime_vars["data"] == {"a": 1}
|
assert tool._runtime_control.snapshot().scratchpad["data"] == {"a": 1}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_whitespace_key_rejected(self):
|
async def test_modify_whitespace_key_rejected(self):
|
||||||
@@ -396,7 +396,7 @@ class TestModifyFree:
|
|||||||
result = await tool.execute(action="set", key="provider_retry_mode", value=42)
|
result = await tool.execute(action="set", key="provider_retry_mode", value=42)
|
||||||
assert "Error" in result
|
assert "Error" in result
|
||||||
assert "str" in result
|
assert "str" in result
|
||||||
assert tool._runtime_state.provider_retry_mode == "standard"
|
assert tool._runtime_control.snapshot().provider_retry_mode == "standard"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_existing_int_attr_wrong_type_rejected(self):
|
async def test_modify_existing_int_attr_wrong_type_rejected(self):
|
||||||
@@ -404,7 +404,7 @@ class TestModifyFree:
|
|||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="max_tool_result_chars", value="big")
|
result = await tool.execute(action="set", key="max_tool_result_chars", value="big")
|
||||||
assert "Error" in result
|
assert "Error" in result
|
||||||
assert tool._runtime_state.max_tool_result_chars == 16000
|
assert tool._runtime_control.snapshot().max_tool_result_chars == 16000
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -486,25 +486,12 @@ class TestModifyOpen:
|
|||||||
assert "protected" in result
|
assert "protected" in result
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_workspace_allowed(self):
|
async def test_modify_workspace_preserves_display_compatibility(self):
|
||||||
"""workspace was READONLY in v1, now freely modifiable."""
|
"""The compatibility value is isolated from filesystem security boundaries."""
|
||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
result = await tool.execute(action="set", key="workspace", value="/new/path")
|
result = await tool.execute(action="set", key="workspace", value="/new/path")
|
||||||
assert "Set workspace" in result
|
assert "Set workspace" in result
|
||||||
|
assert tool._runtime_control.snapshot().workspace == "/new/path"
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_modify_mcp_servers_blocked(self):
|
|
||||||
"""_mcp_servers contains API credentials — must be blocked."""
|
|
||||||
tool = _make_tool()
|
|
||||||
result = await tool.execute(action="set", key="_mcp_servers", value={"evil": "leaked"})
|
|
||||||
assert "protected" in result
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_modify_mcp_stacks_blocked(self):
|
|
||||||
"""_mcp_stacks holds connection handles — must be blocked."""
|
|
||||||
tool = _make_tool()
|
|
||||||
result = await tool.execute(action="set", key="_mcp_stacks", value={})
|
|
||||||
assert "protected" in result
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_pending_queues_blocked(self):
|
async def test_modify_pending_queues_blocked(self):
|
||||||
@@ -534,13 +521,6 @@ class TestModifyOpen:
|
|||||||
result = await tool.execute(action="set", key="_background_tasks", value=[])
|
result = await tool.execute(action="set", key="_background_tasks", value=[])
|
||||||
assert "protected" in result
|
assert "protected" in result
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_inspect_mcp_servers_blocked(self):
|
|
||||||
"""_mcp_servers contains credentials — check must be blocked too."""
|
|
||||||
tool = _make_tool()
|
|
||||||
result = await tool.execute(action="check", key="_mcp_servers")
|
|
||||||
assert "not accessible" in result
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_wrapped_denied(self):
|
async def test_modify_wrapped_denied(self):
|
||||||
"""__wrapped__ allows decorator bypass — must be denied."""
|
"""__wrapped__ allows decorator bypass — must be denied."""
|
||||||
@@ -584,28 +564,28 @@ class TestUnknownAction:
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# runtime_vars limits (from code review)
|
# scratchpad limits
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
class TestRuntimeVarsLimits:
|
class TestScratchpadLimits:
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runtime_vars_rejects_at_max_keys(self):
|
async def test_scratchpad_rejects_at_max_keys(self):
|
||||||
loop = _make_mock_loop()
|
tool = _make_tool()
|
||||||
loop._runtime_vars = {f"key_{i}": i for i in range(64)}
|
for i in range(64):
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool._runtime_control.set_scratchpad(f"key_{i}", i, max_keys=64)
|
||||||
result = await tool.execute(action="set", key="overflow", value="data")
|
result = await tool.execute(action="set", key="overflow", value="data")
|
||||||
assert "full" in result
|
assert "full" in result
|
||||||
assert "overflow" not in loop._runtime_vars
|
assert "overflow" not in tool._runtime_control.snapshot().scratchpad
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runtime_vars_allows_update_existing_key_at_max(self):
|
async def test_scratchpad_allows_update_existing_key_at_max(self):
|
||||||
loop = _make_mock_loop()
|
tool = _make_tool()
|
||||||
loop._runtime_vars = {f"key_{i}": i for i in range(64)}
|
for i in range(64):
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool._runtime_control.set_scratchpad(f"key_{i}", i, max_keys=64)
|
||||||
result = await tool.execute(action="set", key="key_0", value="updated")
|
result = await tool.execute(action="set", key="key_0", value="updated")
|
||||||
assert "Error" not in result
|
assert "Error" not in result
|
||||||
assert loop._runtime_vars["key_0"] == "updated"
|
assert tool._runtime_control.snapshot().scratchpad["key_0"] == "updated"
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -844,7 +824,7 @@ class TestInspectTaskStatuses:
|
|||||||
usage={"prompt_tokens": 500, "completion_tokens": 100},
|
usage={"prompt_tokens": 500, "completion_tokens": 100},
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="check", key="subagents._task_statuses")
|
result = await tool.execute(action="check", key="subagents._task_statuses")
|
||||||
assert "abc12345" in result
|
assert "abc12345" in result
|
||||||
assert "read logs" in result
|
assert "read logs" in result
|
||||||
@@ -865,7 +845,7 @@ class TestInspectTaskStatuses:
|
|||||||
stop_reason="completed",
|
stop_reason="completed",
|
||||||
)
|
)
|
||||||
loop.subagents._task_statuses = {"xyz": status}
|
loop.subagents._task_statuses = {"xyz": status}
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="check", key="subagents._task_statuses.xyz")
|
result = await tool.execute(action="check", key="subagents._task_statuses.xyz")
|
||||||
assert "search code" in result
|
assert "search code" in result
|
||||||
assert "completed" in result
|
assert "completed" in result
|
||||||
@@ -879,7 +859,10 @@ class TestReadOnlyMode:
|
|||||||
|
|
||||||
def _make_readonly_tool(self):
|
def _make_readonly_tool(self):
|
||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
return MyTool(runtime_state=loop, modify_allowed=False)
|
return MyTool(
|
||||||
|
runtime_control=AgentRuntimeControl(loop),
|
||||||
|
modify_allowed=False,
|
||||||
|
)
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inspect_allowed_in_readonly(self):
|
async def test_inspect_allowed_in_readonly(self):
|
||||||
@@ -904,13 +887,13 @@ class TestReadOnlyMode:
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# runtime vars check fallback (Fix #1: cross-turn memory)
|
# scratchpad inspection
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
class TestRuntimeVarsInspectFallback:
|
class TestScratchpadInspection:
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inspect_runtime_var_after_modify(self):
|
async def test_inspect_scratchpad_value_after_modify(self):
|
||||||
"""Design doc scenario: set then check should return the value."""
|
"""Design doc scenario: set then check should return the value."""
|
||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
await tool.execute(action="set", key="user_prefers_concise", value=True)
|
await tool.execute(action="set", key="user_prefers_concise", value=True)
|
||||||
@@ -918,14 +901,14 @@ class TestRuntimeVarsInspectFallback:
|
|||||||
assert "True" in result
|
assert "True" in result
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inspect_runtime_var_string(self):
|
async def test_inspect_scratchpad_string(self):
|
||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
await tool.execute(action="set", key="current_project", value="nanobot")
|
await tool.execute(action="set", key="current_project", value="nanobot")
|
||||||
result = await tool.execute(action="check", key="current_project")
|
result = await tool.execute(action="check", key="current_project")
|
||||||
assert "nanobot" in result
|
assert "nanobot" in result
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inspect_runtime_var_dict(self):
|
async def test_inspect_scratchpad_dict(self):
|
||||||
tool = _make_tool()
|
tool = _make_tool()
|
||||||
await tool.execute(action="set", key="task_meta", value={"step": 2, "total": 5})
|
await tool.execute(action="set", key="task_meta", value={"step": 2, "total": 5})
|
||||||
result = await tool.execute(action="check", key="task_meta")
|
result = await tool.execute(action="check", key="task_meta")
|
||||||
@@ -958,7 +941,7 @@ class TestSensitiveSubFieldBlocking:
|
|||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop.some_config = MagicMock()
|
loop.some_config = MagicMock()
|
||||||
loop.some_config.password = "hunter2"
|
loop.some_config.password = "hunter2"
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="check", key="some_config.password")
|
result = await tool.execute(action="check", key="some_config.password")
|
||||||
assert "not accessible" in result
|
assert "not accessible" in result
|
||||||
|
|
||||||
@@ -967,7 +950,7 @@ class TestSensitiveSubFieldBlocking:
|
|||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop.vault = MagicMock()
|
loop.vault = MagicMock()
|
||||||
loop.vault.secret = "classified"
|
loop.vault.secret = "classified"
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="check", key="vault.secret")
|
result = await tool.execute(action="check", key="vault.secret")
|
||||||
assert "not accessible" in result
|
assert "not accessible" in result
|
||||||
|
|
||||||
@@ -976,7 +959,7 @@ class TestSensitiveSubFieldBlocking:
|
|||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop.auth_data = MagicMock()
|
loop.auth_data = MagicMock()
|
||||||
loop.auth_data.token = "jwt-payload"
|
loop.auth_data.token = "jwt-payload"
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="check", key="auth_data.token")
|
result = await tool.execute(action="check", key="auth_data.token")
|
||||||
assert "not accessible" in result
|
assert "not accessible" in result
|
||||||
|
|
||||||
@@ -992,7 +975,7 @@ class TestSensitiveSubFieldBlocking:
|
|||||||
async def test_modify_password_blocked(self):
|
async def test_modify_password_blocked(self):
|
||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop.some_config = MagicMock()
|
loop.some_config = MagicMock()
|
||||||
tool = _make_tool(runtime_state=loop)
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="set", key="some_config.password", value="evil")
|
result = await tool.execute(action="set", key="some_config.password", value="evil")
|
||||||
assert "not accessible" in result
|
assert "not accessible" in result
|
||||||
|
|
||||||
@@ -1083,8 +1066,8 @@ class TestSecurityAttributeProtection:
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_model_presets_dotpath_blocked(self):
|
async def test_modify_model_presets_dotpath_blocked(self):
|
||||||
"""The config-derived model preset catalog is inspectable but not mutable."""
|
"""The config-derived model preset catalog is inspectable but not mutable."""
|
||||||
presets = {"fast": {"model": "fast-model"}}
|
presets = {"fast": ModelPresetConfig(model="fast-model")}
|
||||||
tool = _make_tool(runtime_state=_make_mock_loop(model_presets=presets))
|
tool = _make_tool(loop=_make_mock_loop(model_presets=presets))
|
||||||
|
|
||||||
result = await tool.execute(
|
result = await tool.execute(
|
||||||
action="set",
|
action="set",
|
||||||
@@ -1093,14 +1076,14 @@ class TestSecurityAttributeProtection:
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert "read-only" in result
|
assert "read-only" in result
|
||||||
assert presets == {"fast": {"model": "fast-model"}}
|
assert presets == {"fast": ModelPresetConfig(model="fast-model")}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inspect_read_only_model_preset_dotpath(self):
|
async def test_inspect_read_only_model_preset_dotpath(self):
|
||||||
presets = MappingProxyType({
|
presets = MappingProxyType({
|
||||||
"fast": ModelPresetConfig(model="fast-model"),
|
"fast": ModelPresetConfig(model="fast-model"),
|
||||||
})
|
})
|
||||||
tool = _make_tool(runtime_state=_make_mock_loop(model_presets=presets))
|
tool = _make_tool(loop=_make_mock_loop(model_presets=presets))
|
||||||
|
|
||||||
result = await tool.execute(action="check", key="model_presets.fast.model")
|
result = await tool.execute(action="check", key="model_presets.fast.model")
|
||||||
|
|
||||||
@@ -1150,7 +1133,8 @@ class TestLastUsageInSummary:
|
|||||||
async def test_last_usage_not_shown_when_empty(self):
|
async def test_last_usage_not_shown_when_empty(self):
|
||||||
loop = _make_mock_loop()
|
loop = _make_mock_loop()
|
||||||
loop._last_usage = {}
|
loop._last_usage = {}
|
||||||
tool = _make_tool(runtime_state=loop)
|
loop.last_usage = loop._last_usage
|
||||||
|
tool = _make_tool(loop=loop)
|
||||||
result = await tool.execute(action="check")
|
result = await tool.execute(action="check")
|
||||||
assert "_last_usage" not in result
|
assert "_last_usage" not in result
|
||||||
|
|
||||||
|
|||||||
@@ -4,23 +4,23 @@ from unittest.mock import MagicMock
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.runtime_control import AgentRuntimeControl
|
||||||
from nanobot.agent.tools.self import MyTool
|
from nanobot.agent.tools.self import MyTool
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_my_tool_max_iterations_syncs_subagent_limit() -> None:
|
async def test_my_tool_max_iterations_syncs_subagent_limit(tmp_path) -> None:
|
||||||
loop = MagicMock()
|
provider = MagicMock()
|
||||||
loop.max_iterations = 40
|
provider.get_default_model.return_value = "test-model"
|
||||||
loop._runtime_vars = {}
|
loop = AgentLoop(
|
||||||
loop.subagents = MagicMock()
|
bus=MessageBus(),
|
||||||
loop.subagents.max_iterations = loop.max_iterations
|
provider=provider,
|
||||||
|
workspace=tmp_path,
|
||||||
def _sync_subagent_runtime_limits() -> None:
|
max_iterations=40,
|
||||||
loop.subagents.max_iterations = loop.max_iterations
|
)
|
||||||
|
tool = MyTool(runtime_control=AgentRuntimeControl(loop))
|
||||||
loop._sync_subagent_runtime_limits = _sync_subagent_runtime_limits
|
|
||||||
|
|
||||||
tool = MyTool(runtime_state=loop)
|
|
||||||
|
|
||||||
result = await tool.execute(action="set", key="max_iterations", value=80)
|
result = await tool.execute(action="set", key="max_iterations", value=80)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,90 @@
|
|||||||
|
"""CLI app subprocesses must not inherit API keys from the parent environ."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
from nanobot.apps.cli.service import CliAppManager
|
||||||
|
|
||||||
|
|
||||||
|
def test_subprocess_env_excludes_api_keys(monkeypatch, tmp_path) -> None:
|
||||||
|
monkeypatch.setenv("OPENAI_API_KEY", "sk-should-not-leak")
|
||||||
|
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-leak")
|
||||||
|
monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-leak")
|
||||||
|
|
||||||
|
manager = CliAppManager(workspace=tmp_path, data_dir=tmp_path / "cli-apps")
|
||||||
|
env = manager._subprocess_env()
|
||||||
|
|
||||||
|
assert "OPENAI_API_KEY" not in env
|
||||||
|
assert "ANTHROPIC_API_KEY" not in env
|
||||||
|
assert "OPENROUTER_API_KEY" not in env
|
||||||
|
assert env.get("PYTHONUNBUFFERED") == "1"
|
||||||
|
assert "PATH" in env
|
||||||
|
|
||||||
|
|
||||||
|
def test_subprocess_env_excludes_api_keys_on_windows(monkeypatch, tmp_path) -> None:
|
||||||
|
monkeypatch.setattr("nanobot.apps.cli.service.sys.platform", "win32")
|
||||||
|
monkeypatch.setenv("OPENAI_API_KEY", "sk-should-not-leak")
|
||||||
|
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-leak")
|
||||||
|
|
||||||
|
manager = CliAppManager(workspace=tmp_path, data_dir=tmp_path / "cli-apps")
|
||||||
|
env = manager._subprocess_env()
|
||||||
|
|
||||||
|
assert "OPENAI_API_KEY" not in env
|
||||||
|
assert "ANTHROPIC_API_KEY" not in env
|
||||||
|
assert env["PYTHONUNBUFFERED"] == "1"
|
||||||
|
assert env["SYSTEMROOT"]
|
||||||
|
assert all(isinstance(value, str) for value in env.values())
|
||||||
|
|
||||||
|
|
||||||
|
def test_run_passes_filtered_env(monkeypatch, tmp_path) -> None:
|
||||||
|
monkeypatch.setenv("OPENAI_API_KEY", "sk-should-not-leak")
|
||||||
|
manager = CliAppManager(workspace=tmp_path, data_dir=tmp_path / "cli-apps")
|
||||||
|
captured: dict[str, object] = {}
|
||||||
|
|
||||||
|
def fake_run(*args, **kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
|
||||||
|
class Result:
|
||||||
|
returncode = 0
|
||||||
|
stdout = "ok"
|
||||||
|
stderr = ""
|
||||||
|
|
||||||
|
return Result()
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.apps.cli.service.subprocess.run", fake_run)
|
||||||
|
monkeypatch.setattr(manager, "get_app", lambda name: {"name": name, "entry_point": "echo"})
|
||||||
|
monkeypatch.setattr(
|
||||||
|
manager,
|
||||||
|
"_load_installed",
|
||||||
|
lambda: {"echo": {"entry_point": "echo"}},
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("nanobot.apps.cli.service.shutil.which", lambda entry: "/bin/echo")
|
||||||
|
monkeypatch.setattr(manager, "_resolve_cwd", lambda *a, **k: tmp_path)
|
||||||
|
monkeypatch.setattr(manager, "_artifact_snapshot", lambda cwd: {})
|
||||||
|
monkeypatch.setattr(manager, "_changed_artifacts", lambda cwd, snap: [])
|
||||||
|
|
||||||
|
manager.run("echo", ["hi"])
|
||||||
|
|
||||||
|
env = captured.get("env")
|
||||||
|
assert isinstance(env, dict)
|
||||||
|
assert "OPENAI_API_KEY" not in env
|
||||||
|
|
||||||
|
|
||||||
|
def test_management_subprocesses_use_filtered_env(monkeypatch, tmp_path) -> None:
|
||||||
|
monkeypatch.setenv("OPENAI_API_KEY", "sk-should-not-leak")
|
||||||
|
captured: dict[str, object] = {}
|
||||||
|
|
||||||
|
def fake_run(*args, **kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
return subprocess.CompletedProcess(args[0], 0, stdout="ok", stderr="")
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.apps.cli.service.subprocess.run", fake_run)
|
||||||
|
manager = CliAppManager(workspace=tmp_path, data_dir=tmp_path / "cli-apps")
|
||||||
|
|
||||||
|
manager._run_argv(["example-cli", "--help"], timeout=5)
|
||||||
|
|
||||||
|
env = captured.get("env")
|
||||||
|
assert isinstance(env, dict)
|
||||||
|
assert "OPENAI_API_KEY" not in env
|
||||||
|
assert env["PYTHONUNBUFFERED"] == "1"
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user