mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-08 21:38:40 +03:00
Compare commits
66
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5922c4ebea | ||
|
|
eae51333ad | ||
|
|
6194a9b919 | ||
|
|
61ae869610 | ||
|
|
3eebe08dba | ||
|
|
38a5f09f02 | ||
|
|
af9f8d54b8 | ||
|
|
1391aa3d57 | ||
|
|
e00220bdb6 | ||
|
|
4dccee56a7 | ||
|
|
2d302a006e | ||
|
|
3f321179eb | ||
|
|
cda1de863e | ||
|
|
57d5276da1 | ||
|
|
30fc05c746 | ||
|
|
15dba8d080 | ||
|
|
a45884c0d3 | ||
|
|
6a8a17a380 | ||
|
|
705abff7a3 | ||
|
|
44b7bba9bd | ||
|
|
d7a73093a8 | ||
|
|
59548b0a04 | ||
|
|
fc1c8ea770 | ||
|
|
99e4d25d4c | ||
|
|
c588d56a77 | ||
|
|
7367741ac1 | ||
|
|
4e0d872588 | ||
|
|
0a5606b409 | ||
|
|
7411afa0e7 | ||
|
|
c4293a7835 | ||
|
|
40c1d83b32 | ||
|
|
0537cc1682 | ||
|
|
7e2dbdef7d | ||
|
|
c4794b82a9 | ||
|
|
d7122a13d3 | ||
|
|
d4ade8f680 | ||
|
|
28d0f8560e | ||
|
|
ba38f90832 | ||
|
|
eb3aed359f | ||
|
|
4445fcc8b9 | ||
|
|
b67205f5aa | ||
|
|
de8761f25a | ||
|
|
8708ccea86 | ||
|
|
eb0ff3ad1d | ||
|
|
c58a360b25 | ||
|
|
5bb94edc99 | ||
|
|
888d54790d | ||
|
|
48d35bd2d9 | ||
|
|
fce1550814 | ||
|
|
bf8a6e35fd | ||
|
|
f017e209da | ||
|
|
5a34504b76 | ||
|
|
af26ed0041 | ||
|
|
112f40ad67 | ||
|
|
2f323e24c1 | ||
|
|
361f31c0e4 | ||
|
|
945f208d38 | ||
|
|
c8bb04a8fe | ||
|
|
4b5de66c58 | ||
|
|
9340567f2d | ||
|
|
e5be4dac7a | ||
|
|
175b58e259 | ||
|
|
3bf8de047a | ||
|
|
400f822601 | ||
|
|
9fb9d7afcb | ||
|
|
c018c3fb6a |
@@ -49,7 +49,7 @@ body:
|
|||||||
attributes:
|
attributes:
|
||||||
label: nanobot Version
|
label: nanobot Version
|
||||||
description: Run `nanobot --version` or `pip show nanobot-ai`
|
description: Run `nanobot --version` or `pip show nanobot-ai`
|
||||||
placeholder: e.g., 0.1.5
|
placeholder: e.g., 0.2.0
|
||||||
validations:
|
validations:
|
||||||
required: true
|
required: true
|
||||||
|
|
||||||
|
|||||||
@@ -97,3 +97,4 @@ logs/
|
|||||||
tmp/
|
tmp/
|
||||||
temp/
|
temp/
|
||||||
*.tmp
|
*.tmp
|
||||||
|
exp/
|
||||||
|
|||||||
+6
-4
@@ -14,8 +14,9 @@ RUN apt-get update && \
|
|||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# Install Python dependencies first (cached layer)
|
# Install Python dependencies first (cached layer). Hatch reads the custom build
|
||||||
COPY pyproject.toml README.md LICENSE ./
|
# hook from hatch_build.py even for this metadata-only install.
|
||||||
|
COPY pyproject.toml README.md LICENSE THIRD_PARTY_NOTICES.md hatch_build.py ./
|
||||||
RUN mkdir -p nanobot bridge && touch nanobot/__init__.py && \
|
RUN mkdir -p nanobot bridge && touch nanobot/__init__.py && \
|
||||||
uv pip install --system --no-cache . && \
|
uv pip install --system --no-cache . && \
|
||||||
rm -rf nanobot bridge
|
rm -rf nanobot bridge
|
||||||
@@ -23,6 +24,7 @@ RUN mkdir -p nanobot bridge && touch nanobot/__init__.py && \
|
|||||||
# Copy the full source and install
|
# Copy the full source and install
|
||||||
COPY nanobot/ nanobot/
|
COPY nanobot/ nanobot/
|
||||||
COPY bridge/ bridge/
|
COPY bridge/ bridge/
|
||||||
|
COPY webui/ webui/
|
||||||
RUN uv pip install --system --no-cache .
|
RUN uv pip install --system --no-cache .
|
||||||
|
|
||||||
# Build the WhatsApp bridge
|
# Build the WhatsApp bridge
|
||||||
@@ -43,8 +45,8 @@ RUN sed -i 's/\r$//' /usr/local/bin/entrypoint.sh && chmod +x /usr/local/bin/ent
|
|||||||
USER nanobot
|
USER nanobot
|
||||||
ENV HOME=/home/nanobot
|
ENV HOME=/home/nanobot
|
||||||
|
|
||||||
# Gateway default port
|
# Gateway health endpoint and optional WebUI/WebSocket channel ports
|
||||||
EXPOSE 18790
|
EXPOSE 18790 8765
|
||||||
|
|
||||||
ENTRYPOINT ["entrypoint.sh"]
|
ENTRYPOINT ["entrypoint.sh"]
|
||||||
CMD ["status"]
|
CMD ["status"]
|
||||||
|
|||||||
@@ -23,6 +23,7 @@
|
|||||||
|
|
||||||
## 📢 News
|
## 📢 News
|
||||||
|
|
||||||
|
- **2026-05-15** 🚀 Released **v0.2.0** — **`/goal`** holds sustained objectives across turns, WebUI now ships inside the wheel, image generation end to end, 5 new providers with `fallback_models`, and a real agent-loop refactor. Please see [release notes](https://github.com/HKUDS/nanobot/releases/tag/v0.2.0) for details.
|
||||||
- **2026-05-14** 🎯 **`/goal`** for long-term objectives, visible multi-step progress, long-horizon missions in chat.
|
- **2026-05-14** 🎯 **`/goal`** for long-term objectives, visible multi-step progress, long-horizon missions in chat.
|
||||||
- **2026-05-13** 🧠 Streaming reasoning before answers, automatic backup models, smoother plug-in reconnects.
|
- **2026-05-13** 🧠 Streaming reasoning before answers, automatic backup models, smoother plug-in reconnects.
|
||||||
- **2026-05-12** 🎛️ Saved model presets with WebUI badge, simpler plug-in tools, quieter Feishu topic threads.
|
- **2026-05-12** 🎛️ Saved model presets with WebUI badge, simpler plug-in tools, quieter Feishu topic threads.
|
||||||
@@ -211,13 +212,13 @@ nanobot agent
|
|||||||
|
|
||||||
|
|
||||||
- Want different LLM providers, web search, MCP, security settings, or more config options? See [Configuration](./docs/configuration.md)
|
- Want different LLM providers, web search, MCP, security settings, or more config options? See [Configuration](./docs/configuration.md)
|
||||||
|
- Want to run locally? Use [Atomic Chat](./docs/configuration.md#atomic-chat-local), [vLLM](./docs/configuration.md#vllm-local-openai-compatible), [Ollama](./docs/configuration.md#ollama-local), and [others](./docs/configuration.md#local-providers).
|
||||||
- Want to run nanobot in chat apps like Telegram, Discord, WeChat or Feishu? See [Chat Apps](./docs/chat-apps.md)
|
- Want to run nanobot in chat apps like Telegram, Discord, WeChat or Feishu? See [Chat Apps](./docs/chat-apps.md)
|
||||||
- Want Docker or Linux service deployment? See [Deployment](./docs/deployment.md)
|
- Want Docker or Linux service deployment? See [Deployment](./docs/deployment.md)
|
||||||
|
|
||||||
## 🧪 WebUI (Development)
|
## 🌐 WebUI
|
||||||
|
|
||||||
> [!NOTE]
|
The WebUI ships **inside the published wheel** — no extra build step. Just enable the WebSocket channel and open it in your browser.
|
||||||
> The WebUI development workflow currently requires a source checkout and is not yet shipped together with the official packaged release. See [WebUI Document](./webui/README.md) for full WebUI development docs and build steps.
|
|
||||||
|
|
||||||
<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">
|
||||||
@@ -235,13 +236,12 @@ nanobot agent
|
|||||||
nanobot gateway
|
nanobot gateway
|
||||||
```
|
```
|
||||||
|
|
||||||
**3. Start the webui dev server**
|
**3. Open the WebUI**
|
||||||
|
|
||||||
```bash
|
Visit [`http://127.0.0.1:8765`](http://127.0.0.1:8765) in your browser. To open it from another device on your LAN, see [WebUI docs → LAN access](./webui/README.md#access-from-another-device-lan).
|
||||||
cd webui
|
|
||||||
bun install
|
> [!TIP]
|
||||||
bun run dev
|
> Working on the WebUI itself? Check out [`webui/README.md`](./webui/README.md) for the Vite dev server (HMR) workflow.
|
||||||
```
|
|
||||||
|
|
||||||
## 🏗️ Architecture
|
## 🏗️ Architecture
|
||||||
|
|
||||||
@@ -330,4 +330,4 @@ This project was started by [Xubin Ren](https://github.com/re-bin) as a personal
|
|||||||
<p align="center">
|
<p align="center">
|
||||||
<em> Thanks for visiting ✨ nanobot!</em><br><br>
|
<em> Thanks for visiting ✨ nanobot!</em><br><br>
|
||||||
<img src="https://visitor-badge.laobi.icu/badge?page_id=HKUDS.nanobot&style=for-the-badge&color=00d4ff" alt="Views">
|
<img src="https://visitor-badge.laobi.icu/badge?page_id=HKUDS.nanobot&style=for-the-badge&color=00d4ff" alt="Views">
|
||||||
</p>
|
</p>
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ services:
|
|||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
ports:
|
ports:
|
||||||
- 18790:18790
|
- 18790:18790
|
||||||
|
- 8765:8765
|
||||||
deploy:
|
deploy:
|
||||||
resources:
|
resources:
|
||||||
limits:
|
limits:
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ Start here for setup, everyday usage, and deployment.
|
|||||||
| Agent social network | [`agent-social-network.md`](./agent-social-network.md) | Join external agent communities from nanobot |
|
| Agent social network | [`agent-social-network.md`](./agent-social-network.md) | Join external agent communities from nanobot |
|
||||||
| Configuration | [`configuration.md`](./configuration.md) | Providers, tools, channels, MCP, and runtime settings |
|
| Configuration | [`configuration.md`](./configuration.md) | Providers, tools, channels, MCP, and runtime settings |
|
||||||
| Image generation | [`image-generation.md`](./image-generation.md) | Configure image providers, WebUI image mode, and generated artifacts |
|
| Image generation | [`image-generation.md`](./image-generation.md) | Configure image providers, WebUI image mode, and generated artifacts |
|
||||||
|
| WebUI | [`../webui/README.md`](../webui/README.md) | Open the bundled browser UI; LAN access; Vite dev server for contributors |
|
||||||
| Multiple instances | [`multiple-instances.md`](./multiple-instances.md) | Run isolated bots with separate configs and workspaces |
|
| Multiple instances | [`multiple-instances.md`](./multiple-instances.md) | Run isolated bots with separate configs and workspaces |
|
||||||
| CLI reference | [`cli-reference.md`](./cli-reference.md) | Core CLI commands and common entrypoints |
|
| CLI reference | [`cli-reference.md`](./cli-reference.md) | Core CLI commands and common entrypoints |
|
||||||
| In-chat commands | [`chat-commands.md`](./chat-commands.md) | Slash commands and periodic task behavior |
|
| In-chat commands | [`chat-commands.md`](./chat-commands.md) | Slash commands and periodic task behavior |
|
||||||
|
|||||||
+156
-10
@@ -26,7 +26,52 @@ Instead of storing secrets directly in `config.json`, you can use `${VAR_NAME}`
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
For **systemd** deployments, use `EnvironmentFile=` in the service unit to load variables from a file that only the deploying user can read:
|
Any string value in `config.json` can use `${VAR_NAME}`. Resolution runs once at startup, in memory only — resolved values are never written back to disk, so editing config through `nanobot onboard` or the WebUI preserves the placeholder.
|
||||||
|
|
||||||
|
If a referenced variable is unset, nanobot fails fast at startup with `ValueError: Environment variable 'NAME' referenced in config is not set`.
|
||||||
|
|
||||||
|
### More examples
|
||||||
|
|
||||||
|
**MCP servers** — both stdio `env` and HTTP `headers`:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"mcpServers": {
|
||||||
|
"github": {
|
||||||
|
"command": "npx",
|
||||||
|
"args": ["-y", "@modelcontextprotocol/server-github"],
|
||||||
|
"env": { "GITHUB_PERSONAL_ACCESS_TOKEN": "${GITHUB_TOKEN}" }
|
||||||
|
},
|
||||||
|
"remote": {
|
||||||
|
"url": "https://example.com/mcp/",
|
||||||
|
"headers": { "Authorization": "Bearer ${REMOTE_MCP_TOKEN}" }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Web search providers:**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"web": {
|
||||||
|
"search": {
|
||||||
|
"provider": "brave",
|
||||||
|
"apiKey": "${BRAVE_API_KEY}"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Loading variables at startup
|
||||||
|
|
||||||
|
Pick whatever fits your deployment — nanobot only reads `os.environ` at startup, so any mechanism that populates the process environment works.
|
||||||
|
|
||||||
|
**systemd** — use `EnvironmentFile=` in the service unit to load variables from a file that only the deploying user can read:
|
||||||
|
|
||||||
```ini
|
```ini
|
||||||
# /etc/systemd/system/nanobot.service (excerpt)
|
# /etc/systemd/system/nanobot.service (excerpt)
|
||||||
@@ -42,6 +87,35 @@ TELEGRAM_TOKEN=your-token-here
|
|||||||
IMAP_PASSWORD=your-password-here
|
IMAP_PASSWORD=your-password-here
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**Docker** — pass an env file to the locally built image (one `KEY=VALUE` per line), or use `-e KEY=value`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker run --rm --env-file=./nanobot.env \
|
||||||
|
-v ~/.nanobot:/home/nanobot/.nanobot \
|
||||||
|
nanobot agent -m "Hello"
|
||||||
|
```
|
||||||
|
|
||||||
|
**direnv** — drop a `.envrc` in your working directory and run `direnv allow`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# .envrc (auto-loaded by direnv)
|
||||||
|
export TELEGRAM_TOKEN=your-token-here
|
||||||
|
export ANTHROPIC_API_KEY=...
|
||||||
|
```
|
||||||
|
|
||||||
|
**Secret managers (1Password, Bitwarden, pass)** — wrap the process so secrets only exist as env vars for the lifetime of the run, never on disk:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 1Password — references in .env.tpl look like `op://Vault/Item/field`
|
||||||
|
op run --env-file=.env.tpl -- nanobot agent
|
||||||
|
|
||||||
|
# pass (passwordstore.org)
|
||||||
|
ANTHROPIC_API_KEY="$(pass show api/anthropic)" nanobot agent
|
||||||
|
|
||||||
|
# Bitwarden
|
||||||
|
ANTHROPIC_API_KEY="$(bw get password api/anthropic)" nanobot agent
|
||||||
|
```
|
||||||
|
|
||||||
## Providers
|
## Providers
|
||||||
|
|
||||||
> [!TIP]
|
> [!TIP]
|
||||||
@@ -60,6 +134,7 @@ IMAP_PASSWORD=your-password-here
|
|||||||
| `custom` | Any OpenAI-compatible endpoint | — |
|
| `custom` | Any OpenAI-compatible endpoint | — |
|
||||||
| `openrouter` | LLM (recommended, access to all models) | [openrouter.ai](https://openrouter.ai) |
|
| `openrouter` | LLM (recommended, access to all models) | [openrouter.ai](https://openrouter.ai) |
|
||||||
| `huggingface` | LLM (Hugging Face Inference Providers) | [huggingface.co/settings/tokens](https://huggingface.co/settings/tokens) |
|
| `huggingface` | LLM (Hugging Face Inference Providers) | [huggingface.co/settings/tokens](https://huggingface.co/settings/tokens) |
|
||||||
|
| `skywork` | LLM (Skywork / APIFree API gateway) | [apifree.ai](https://www.apifree.ai) |
|
||||||
| `volcengine` | LLM (VolcEngine, pay-per-use) | [Coding Plan](https://www.volcengine.com/activity/codingplan?utm_campaign=nanobot&utm_content=nanobot&utm_medium=devrel&utm_source=OWO&utm_term=nanobot) · [volcengine.com](https://www.volcengine.com) |
|
| `volcengine` | LLM (VolcEngine, pay-per-use) | [Coding Plan](https://www.volcengine.com/activity/codingplan?utm_campaign=nanobot&utm_content=nanobot&utm_medium=devrel&utm_source=OWO&utm_term=nanobot) · [volcengine.com](https://www.volcengine.com) |
|
||||||
| `byteplus` | LLM (VolcEngine international, pay-per-use) | [Coding Plan](https://www.byteplus.com/en/activity/codingplan?utm_campaign=nanobot&utm_content=nanobot&utm_medium=devrel&utm_source=OWO&utm_term=nanobot) · [byteplus.com](https://www.byteplus.com) |
|
| `byteplus` | LLM (VolcEngine international, pay-per-use) | [Coding Plan](https://www.byteplus.com/en/activity/codingplan?utm_campaign=nanobot&utm_content=nanobot&utm_medium=devrel&utm_source=OWO&utm_term=nanobot) · [byteplus.com](https://www.byteplus.com) |
|
||||||
| `anthropic` | LLM (Claude direct) | [console.anthropic.com](https://console.anthropic.com) |
|
| `anthropic` | LLM (Claude direct) | [console.anthropic.com](https://console.anthropic.com) |
|
||||||
@@ -78,6 +153,7 @@ IMAP_PASSWORD=your-password-here
|
|||||||
| `zhipu` | LLM (Zhipu GLM) | [open.bigmodel.cn](https://open.bigmodel.cn) |
|
| `zhipu` | LLM (Zhipu GLM) | [open.bigmodel.cn](https://open.bigmodel.cn) |
|
||||||
| `mimo` | LLM (MiMo) | [platform.xiaomimimo.com](https://platform.xiaomimimo.com) |
|
| `mimo` | LLM (MiMo) | [platform.xiaomimimo.com](https://platform.xiaomimimo.com) |
|
||||||
| `longcat` | LLM (LongCat) | [longcat.chat](https://longcat.chat/platform/docs/zh/) |
|
| `longcat` | LLM (LongCat) | [longcat.chat](https://longcat.chat/platform/docs/zh/) |
|
||||||
|
| `ant_ling` | LLM (Ant Ling / 蚂蚁百灵) | [developer.ant-ling.com](https://developer.ant-ling.com/en/docs/api-reference/openai/) |
|
||||||
| `ollama` | LLM (local, Ollama) | — |
|
| `ollama` | LLM (local, Ollama) | — |
|
||||||
| `lm_studio` | LLM (local, LM Studio) | — |
|
| `lm_studio` | LLM (local, LM Studio) | — |
|
||||||
| `atomic_chat` | LLM (local, [Atomic Chat](https://atomic.chat/)) | — |
|
| `atomic_chat` | LLM (local, [Atomic Chat](https://atomic.chat/)) | — |
|
||||||
@@ -89,6 +165,36 @@ IMAP_PASSWORD=your-password-here
|
|||||||
| `github_copilot` | LLM (GitHub Copilot, OAuth) | `nanobot provider login github-copilot` |
|
| `github_copilot` | LLM (GitHub Copilot, OAuth) | `nanobot provider login github-copilot` |
|
||||||
| `qianfan` | LLM (Baidu Qianfan) | [cloud.baidu.com](https://cloud.baidu.com/doc/qianfan/s/Hmh4suq26) |
|
| `qianfan` | LLM (Baidu Qianfan) | [cloud.baidu.com](https://cloud.baidu.com/doc/qianfan/s/Hmh4suq26) |
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>Skywork / APIFree</b></summary>
|
||||||
|
|
||||||
|
Skywork uses APIFree's OpenAI-compatible Agent API endpoint. Configure the provider
|
||||||
|
once, then use Skywork model IDs such as `skywork-ai/skyclaw-v1`.
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"skywork": {
|
||||||
|
"apiKey": "${SKYWORK_API_KEY}",
|
||||||
|
"apiBase": "https://api.apifree.ai/agent/v1"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "skywork",
|
||||||
|
"model": "skywork-ai/skyclaw-v1",
|
||||||
|
"maxTokens": 32768,
|
||||||
|
"contextWindowTokens": 131072
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
You can also reference `${APIFREE_API_KEY}` in `apiKey` if that is how your
|
||||||
|
environment names the credential.
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>AWS Bedrock (Converse API)</b></summary>
|
<summary><b>AWS Bedrock (Converse API)</b></summary>
|
||||||
|
|
||||||
@@ -370,6 +476,34 @@ Official model names include `LongCat-Flash-Chat`, `LongCat-Flash-Thinking`,
|
|||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>Ant Ling (OpenAI-compatible)</b></summary>
|
||||||
|
|
||||||
|
Ant Ling is available through nanobot's built-in OpenAI-compatible provider flow.
|
||||||
|
The default API base points to `https://api.ant-ling.com/v1`, so you usually
|
||||||
|
only need to set `apiKey`.
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"antLing": {
|
||||||
|
"apiKey": "${ANT_LING_API_KEY}"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "ant_ling",
|
||||||
|
"model": "Ling-2.6-flash"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Official OpenAI-compatible model names include `Ling-2.6-1T`,
|
||||||
|
`Ling-2.6-flash`, `Ling-2.5-1T`, `Ling-1T`, `Ring-2.5-1T`, and `Ring-1T`.
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>Custom Provider (Any OpenAI-compatible API)</b></summary>
|
<summary><b>Custom Provider (Any OpenAI-compatible API)</b></summary>
|
||||||
|
|
||||||
@@ -438,6 +572,8 @@ Some OpenAI-compatible gateways expose request-body extensions such as vLLM guid
|
|||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<a id="local-providers"></a>
|
||||||
|
<a id="ollama-local"></a>
|
||||||
<details>
|
<details>
|
||||||
<summary><b>Ollama (local)</b></summary>
|
<summary><b>Ollama (local)</b></summary>
|
||||||
|
|
||||||
@@ -503,12 +639,19 @@ ollama run llama3.2
|
|||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<a id="atomic-chat-local"></a>
|
||||||
<details>
|
<details>
|
||||||
<summary><b>Atomic Chat (local)</b></summary>
|
<summary><b>Atomic Chat (local)</b></summary>
|
||||||
|
|
||||||
[Atomic Chat](https://atomic.chat/) is a local-first desktop app that exposes an **OpenAI-compatible** HTTP API (default `http://localhost:1337/v1`). Start Atomic Chat and enable the local API server, then point nanobot at it.
|
[Atomic Chat](https://atomic.chat/) is a local-first desktop app that exposes an **OpenAI-compatible** HTTP API (default `http://localhost:1337/v1`). Use it when you want to run nanobot against a model on your own machine instead of a hosted API provider.
|
||||||
|
|
||||||
**1. Add to config** (partial — merge into `~/.nanobot/config.json`):
|
**1. Start Atomic Chat**
|
||||||
|
|
||||||
|
- Install [Atomic Chat](https://atomic.chat/) on your machine.
|
||||||
|
- Open Atomic Chat, download a model, and keep the app running. The local API is enabled by default.
|
||||||
|
- Copy the model ID exposed by the local API. For example, the model ID for `Qwen 3 32B` might be `qwen3-32b`.
|
||||||
|
|
||||||
|
**2. Add to config** (partial — merge into `~/.nanobot/config.json`):
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
@@ -521,13 +664,13 @@ ollama run llama3.2
|
|||||||
"agents": {
|
"agents": {
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"provider": "atomic_chat",
|
"provider": "atomic_chat",
|
||||||
"model": "your-model-id-from-atomic-chat"
|
"model": "qwen3-32b"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
> **Note:** Set `apiKey` to `null` if your Atomic Chat server does not require a key. If it does, set `apiKey` (or the `ATOMIC_CHAT_API_KEY` environment variable) to the value Atomic Chat expects. The `model` string must match the model id Atomic Chat exposes on its OpenAI-compatible endpoint.
|
> **Note:** Replace `qwen3-32b` with the model ID from Atomic Chat. Set `apiKey` to `null` if your Atomic Chat server does not require a key. If it does, set `apiKey` (or the `ATOMIC_CHAT_API_KEY` environment variable) to the value Atomic Chat expects.
|
||||||
|
|
||||||
> `provider: "auto"` also works when `providers.atomic_chat.apiBase` is configured, but setting `"provider": "atomic_chat"` is the clearest option.
|
> `provider: "auto"` also works when `providers.atomic_chat.apiBase` is configured, but setting `"provider": "atomic_chat"` is the clearest option.
|
||||||
|
|
||||||
@@ -608,6 +751,7 @@ docker run -d \
|
|||||||
> See the [official OVMS docs](https://docs.openvino.ai/2026/model-server/ovms_docs_llm_quickstart.html) for more details.
|
> See the [official OVMS docs](https://docs.openvino.ai/2026/model-server/ovms_docs_llm_quickstart.html) for more details.
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<a id="vllm-local-openai-compatible"></a>
|
||||||
<details>
|
<details>
|
||||||
<summary><b>vLLM (local / OpenAI-compatible)</b></summary>
|
<summary><b>vLLM (local / OpenAI-compatible)</b></summary>
|
||||||
|
|
||||||
@@ -917,7 +1061,7 @@ By default, web search uses `duckduckgo`, and it works out of the box without an
|
|||||||
"web": {
|
"web": {
|
||||||
"search": {
|
"search": {
|
||||||
"provider": "brave",
|
"provider": "brave",
|
||||||
"apiKey": "BSA..."
|
"apiKey": "${BRAVE_API_KEY}"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -931,7 +1075,7 @@ By default, web search uses `duckduckgo`, and it works out of the box without an
|
|||||||
"web": {
|
"web": {
|
||||||
"search": {
|
"search": {
|
||||||
"provider": "tavily",
|
"provider": "tavily",
|
||||||
"apiKey": "tvly-..."
|
"apiKey": "${TAVILY_API_KEY}"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -945,7 +1089,7 @@ By default, web search uses `duckduckgo`, and it works out of the box without an
|
|||||||
"web": {
|
"web": {
|
||||||
"search": {
|
"search": {
|
||||||
"provider": "jina",
|
"provider": "jina",
|
||||||
"apiKey": "jina_..."
|
"apiKey": "${JINA_API_KEY}"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -959,7 +1103,7 @@ By default, web search uses `duckduckgo`, and it works out of the box without an
|
|||||||
"web": {
|
"web": {
|
||||||
"search": {
|
"search": {
|
||||||
"provider": "kagi",
|
"provider": "kagi",
|
||||||
"apiKey": "your-kagi-api-key"
|
"apiKey": "${KAGI_API_KEY}"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -973,7 +1117,7 @@ By default, web search uses `duckduckgo`, and it works out of the box without an
|
|||||||
"web": {
|
"web": {
|
||||||
"search": {
|
"search": {
|
||||||
"provider": "olostep",
|
"provider": "olostep",
|
||||||
"apiKey": "YOUR_OLOSTEP_API_KEY"
|
"apiKey": "${OLOSTEP_API_KEY}"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1136,6 +1280,8 @@ MCP tools are automatically discovered and registered on startup. The LLM can us
|
|||||||
> [!TIP]
|
> [!TIP]
|
||||||
> For production deployments, set `"restrictToWorkspace": true` and `"tools.exec.sandbox": "bwrap"` in your config to sandbox the agent.
|
> For production deployments, set `"restrictToWorkspace": true` and `"tools.exec.sandbox": "bwrap"` in your config to sandbox the agent.
|
||||||
|
|
||||||
|
For API keys, tokens, and other secrets, see [Environment Variables for Secrets](#environment-variables-for-secrets) — avoid storing them directly in `config.json`.
|
||||||
|
|
||||||
| Option | Default | Description |
|
| Option | Default | Description |
|
||||||
|--------|---------|-------------|
|
|--------|---------|-------------|
|
||||||
| `tools.restrictToWorkspace` | `false` | When `true`, restricts **all** agent tools (shell, file read/write/edit, list) to the workspace directory. Prevents path traversal and out-of-scope access. |
|
| `tools.restrictToWorkspace` | `false` | When `true`, restricts **all** agent tools (shell, file read/write/edit, list) to the workspace directory. Prevents path traversal and out-of-scope access. |
|
||||||
|
|||||||
+26
-2
@@ -10,6 +10,18 @@
|
|||||||
> [!IMPORTANT]
|
> [!IMPORTANT]
|
||||||
> Official Docker usage currently means building from this repository with the included `Dockerfile`. Docker Hub images under third-party namespaces are not maintained or verified by HKUDS/nanobot; do not mount API keys or bot tokens into them unless you trust the publisher.
|
> Official Docker usage currently means building from this repository with the included `Dockerfile`. Docker Hub images under third-party namespaces are not maintained or verified by HKUDS/nanobot; do not mount API keys or bot tokens into them unless you trust the publisher.
|
||||||
|
|
||||||
|
> [!IMPORTANT]
|
||||||
|
> The gateway and WebSocket channel default to `host: "127.0.0.1"` in `config.json` (set in `nanobot/config/schema.py`). Docker `-p` port forwarding cannot reach a container's loopback interface, so for the host or LAN to reach the exposed ports you must set both binds to `0.0.0.0` in `~/.nanobot/config.json` before starting the container:
|
||||||
|
>
|
||||||
|
> ```json
|
||||||
|
> {
|
||||||
|
> "gateway": { "host": "0.0.0.0" },
|
||||||
|
> "channels": { "websocket": { "host": "0.0.0.0" } }
|
||||||
|
> }
|
||||||
|
> ```
|
||||||
|
>
|
||||||
|
> When `host` is `0.0.0.0`, the gateway refuses to start unless `token` or `tokenIssueSecret` is also configured on the WebSocket channel — see [`webui/README.md`](../webui/README.md) for details.
|
||||||
|
|
||||||
### Docker Compose
|
### Docker Compose
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
@@ -36,8 +48,20 @@ docker run -v ~/.nanobot:/home/nanobot/.nanobot --rm nanobot onboard
|
|||||||
# Edit config on host to add API keys
|
# Edit config on host to add API keys
|
||||||
vim ~/.nanobot/config.json
|
vim ~/.nanobot/config.json
|
||||||
|
|
||||||
# Run gateway (connects to enabled channels, e.g. Telegram/Discord/Mochat)
|
# Run gateway (connects to enabled channels, e.g. Telegram/Discord/Mochat).
|
||||||
docker run -v ~/.nanobot:/home/nanobot/.nanobot -p 18790:18790 nanobot gateway
|
# Mirrors the security caps and port mappings declared in docker-compose.yml:
|
||||||
|
# - `--cap-drop ALL --cap-add SYS_ADMIN` + unconfined apparmor/seccomp are required
|
||||||
|
# when `tools.exec.sandbox: "bwrap"` is enabled (bwrap needs CAP_SYS_ADMIN for
|
||||||
|
# user namespaces). Without them, `bwrap` exits with `clone3: Operation not permitted`.
|
||||||
|
# - `-p 8765:8765` exposes the WebSocket channel / WebUI alongside the gateway health
|
||||||
|
# endpoint on 18790.
|
||||||
|
docker run \
|
||||||
|
--cap-drop ALL --cap-add SYS_ADMIN \
|
||||||
|
--security-opt apparmor=unconfined \
|
||||||
|
--security-opt seccomp=unconfined \
|
||||||
|
-v ~/.nanobot:/home/nanobot/.nanobot \
|
||||||
|
-p 18790:18790 -p 8765:8765 \
|
||||||
|
nanobot gateway
|
||||||
|
|
||||||
# Or run a single command
|
# Or run a single command
|
||||||
docker run -v ~/.nanobot:/home/nanobot/.nanobot --rm nanobot agent -m "Hello!"
|
docker run -v ~/.nanobot:/home/nanobot/.nanobot --rm nanobot agent -m "Hello!"
|
||||||
|
|||||||
+108
-27
@@ -6,8 +6,6 @@ The feature is disabled by default. Enable it in `~/.nanobot/config.json`, confi
|
|||||||
|
|
||||||
## Quick Setup
|
## Quick Setup
|
||||||
|
|
||||||
OpenRouter example:
|
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"providers": {
|
"providers": {
|
||||||
@@ -19,34 +17,13 @@ OpenRouter example:
|
|||||||
"imageGeneration": {
|
"imageGeneration": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"provider": "openrouter",
|
"provider": "openrouter",
|
||||||
"model": "openai/gpt-5.4-image-2",
|
"model": "openai/gpt-5.4-image-2"
|
||||||
"defaultAspectRatio": "1:1",
|
|
||||||
"defaultImageSize": "1K"
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
AIHubMix example:
|
See [Provider Notes](#provider-notes) for AIHubMix, MiniMax, and Gemini configuration examples.
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"providers": {
|
|
||||||
"aihubmix": {
|
|
||||||
"apiKey": "${AIHUBMIX_API_KEY}"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"tools": {
|
|
||||||
"imageGeneration": {
|
|
||||||
"enabled": true,
|
|
||||||
"provider": "aihubmix",
|
|
||||||
"model": "gpt-image-2-free",
|
|
||||||
"defaultAspectRatio": "1:1",
|
|
||||||
"defaultImageSize": "1K"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
> [!TIP]
|
> [!TIP]
|
||||||
> Prefer environment variables for API keys. nanobot resolves `${VAR_NAME}` values from the environment at startup.
|
> Prefer environment variables for API keys. nanobot resolves `${VAR_NAME}` values from the environment at startup.
|
||||||
@@ -69,7 +46,7 @@ The WebUI hides provider storage details from the user. The agent sees the saved
|
|||||||
| Option | Type | Default | Description |
|
| Option | Type | Default | Description |
|
||||||
|--------|------|---------|-------------|
|
|--------|------|---------|-------------|
|
||||||
| `tools.imageGeneration.enabled` | boolean | `false` | Register the `generate_image` tool |
|
| `tools.imageGeneration.enabled` | boolean | `false` | Register the `generate_image` tool |
|
||||||
| `tools.imageGeneration.provider` | string | `"openrouter"` | Image provider name. Currently `openrouter` and `aihubmix` are supported |
|
| `tools.imageGeneration.provider` | string | `"openrouter"` | Image provider name. Supported values: `openrouter`, `aihubmix`, `minimax`, `gemini`, `stepfun` |
|
||||||
| `tools.imageGeneration.model` | string | `"openai/gpt-5.4-image-2"` | Provider model name |
|
| `tools.imageGeneration.model` | string | `"openai/gpt-5.4-image-2"` | Provider model name |
|
||||||
| `tools.imageGeneration.defaultAspectRatio` | string | `"1:1"` | Default ratio when the prompt/tool call does not specify one |
|
| `tools.imageGeneration.defaultAspectRatio` | string | `"1:1"` | Default ratio when the prompt/tool call does not specify one |
|
||||||
| `tools.imageGeneration.defaultImageSize` | string | `"1K"` | Default size hint, for example `1K`, `2K`, `4K`, or `1024x1024` |
|
| `tools.imageGeneration.defaultImageSize` | string | `"1K"` | Default size hint, for example `1K`, `2K`, `4K`, or `1024x1024` |
|
||||||
@@ -139,6 +116,110 @@ Configure:
|
|||||||
|
|
||||||
`quality: low` is optional. It can make free image models faster and less likely to time out, but it is not required for correctness.
|
`quality: low` is optional. It can make free image models faster and less likely to time out, but it is not required for correctness.
|
||||||
|
|
||||||
|
### MiniMax
|
||||||
|
|
||||||
|
MiniMax `image-01` supports text-to-image and reference-image (subject reference) edits. Supported aspect ratios are `1:1`, `16:9`, `4:3`, `3:2`, `2:3`, `3:4`, `9:16`, and `21:9`.
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"minimax": {
|
||||||
|
"apiKey": "${MINIMAX_API_KEY}"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "minimax",
|
||||||
|
"model": "image-01",
|
||||||
|
"defaultAspectRatio": "1:1"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Gemini
|
||||||
|
|
||||||
|
nanobot supports two Gemini image generation model families via Google's Generative Language API:
|
||||||
|
|
||||||
|
| Model | Endpoint | Reference images |
|
||||||
|
|-------|----------|-----------------|
|
||||||
|
| `imagen-4.0-generate-001` | `:predict` | Not supported by this integration |
|
||||||
|
| `gemini-2.5-flash-image` | `:generateContent` | Supported |
|
||||||
|
|
||||||
|
For reference-image edits, use a Gemini Flash image model:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"gemini": {
|
||||||
|
"apiKey": "${GEMINI_API_KEY}"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "gemini",
|
||||||
|
"model": "gemini-2.5-flash-image"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Imagen 4 supports the aspect ratios `1:1`, `9:16`, `16:9`, `3:4`, and `4:3`. Unsupported ratios are ignored and the model uses its default. The `defaultImageSize` setting has no effect on Gemini models; sizing is controlled by `defaultAspectRatio` only. Reference images passed with an Imagen model are ignored (with a warning logged).
|
||||||
|
|
||||||
|
### StepFun
|
||||||
|
|
||||||
|
StepFun (阶跃星辰) `step-image-edit-2` supports text-to-image generation. The `step-1x-medium` variant additionally supports **style-reference** image edits, where a reference image guides the visual style of the output.
|
||||||
|
|
||||||
|
Supported aspect ratios: `1:1`, `16:9`, `9:16`, `3:4`, `4:3`. Sizes are specified as `WIDTHxHEIGHT` (e.g. `1024x1024`, `1280x800`, `800x1280`).
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"stepfun": {
|
||||||
|
"apiKey": "${STEPFUN_API_KEY}"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "stepfun",
|
||||||
|
"model": "step-image-edit-2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> The StepFun provider reuses the existing `providers.stepfun` config block (the same one used for StepFun's LLM API). Set `providers.stepfun.apiKey` once and it is shared between text and image generation.
|
||||||
|
>
|
||||||
|
> When `step-image-edit-2` is used, `reference_images` are ignored (the model does not support style reference). Switch to `step-1x-medium` to use reference-image-guided generation.
|
||||||
|
|
||||||
|
#### StepPlan (Subscription)
|
||||||
|
|
||||||
|
StepPlan is StepFun's subscription tier and uses a different API base URL. The image generation endpoint path is the same — just override `apiBase`:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"stepfun": {
|
||||||
|
"apiKey": "${STEPFUN_API_KEY}",
|
||||||
|
"apiBase": "https://api.stepfun.com/step_plan/v1"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "stepfun",
|
||||||
|
"model": "step-image-edit-2"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`apiBase` takes precedence over the registry default, so with the StepPlan base URL configured, image requests are sent to `https://api.stepfun.com/step_plan/v1/images/generations` — the same path prefix used for LLM calls. The API key is shared with the standard StepFun provider.
|
||||||
|
|
||||||
## Artifacts
|
## Artifacts
|
||||||
|
|
||||||
Generated images are stored under the active nanobot instance's media directory:
|
Generated images are stored under the active nanobot instance's media directory:
|
||||||
@@ -193,7 +274,7 @@ Use the reference image. Keep the same robot and composition, change the palette
|
|||||||
|---------|-------|
|
|---------|-------|
|
||||||
| `generate_image` is not available | Set `tools.imageGeneration.enabled` to `true` and restart the gateway |
|
| `generate_image` is not available | Set `tools.imageGeneration.enabled` to `true` and restart the gateway |
|
||||||
| Missing API key error | Configure `providers.<provider>.apiKey`; if using `${VAR_NAME}`, confirm the environment variable is visible to the gateway process |
|
| Missing API key error | Configure `providers.<provider>.apiKey`; if using `${VAR_NAME}`, confirm the environment variable is visible to the gateway process |
|
||||||
| `unsupported image generation provider` | Use `openrouter` or `aihubmix` |
|
| `unsupported image generation provider` | Use `openrouter`, `aihubmix`, `minimax`, `gemini`, or `stepfun` |
|
||||||
| AIHubMix says `Incorrect model ID` | Use `model: "gpt-image-2-free"`; nanobot expands it to the required `openai/gpt-image-2-free` model path internally |
|
| AIHubMix says `Incorrect model ID` | Use `model: "gpt-image-2-free"`; nanobot expands it to the required `openai/gpt-image-2-free` model path internally |
|
||||||
| Generation times out | Try a smaller/default image size, set AIHubMix `extraBody.quality` to `"low"`, or retry later |
|
| Generation times out | Try a smaller/default image size, set AIHubMix `extraBody.quality` to `"low"`, or retry later |
|
||||||
| Reference image rejected | Reference image paths must be inside the workspace or nanobot media directory and must be valid image files |
|
| Reference image rejected | Reference image paths must be inside the workspace or nanobot media directory and must be valid image files |
|
||||||
|
|||||||
+101
@@ -0,0 +1,101 @@
|
|||||||
|
"""Hatch build hook that bundles the webui (Vite) into nanobot/web/dist.
|
||||||
|
|
||||||
|
Triggered automatically by `python -m build` (and any other hatch-driven build)
|
||||||
|
so published wheels and sdists ship a fresh webui without requiring developers
|
||||||
|
to remember `cd webui && bun run build` beforehand.
|
||||||
|
|
||||||
|
Behaviour:
|
||||||
|
|
||||||
|
- Skips for editable installs (`pip install -e .`). Editable mode is for Python
|
||||||
|
development; webui contributors use `cd webui && bun run dev` (Vite HMR) and
|
||||||
|
do not need a packaged `dist/`.
|
||||||
|
- No-op when `webui/package.json` is absent (e.g. installing from an sdist that
|
||||||
|
already contains a prebuilt `nanobot/web/dist/`).
|
||||||
|
- Skips when `NANOBOT_SKIP_WEBUI_BUILD=1` is set.
|
||||||
|
- Skips when `nanobot/web/dist/index.html` already exists, unless
|
||||||
|
`NANOBOT_FORCE_WEBUI_BUILD=1` is set.
|
||||||
|
- Uses `bun` when available, otherwise falls back to `npm`. The chosen tool
|
||||||
|
performs `install` followed by `run build`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import subprocess
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from hatchling.builders.hooks.plugin.interface import BuildHookInterface
|
||||||
|
|
||||||
|
|
||||||
|
class WebUIBuildHook(BuildHookInterface):
|
||||||
|
PLUGIN_NAME = "webui-build"
|
||||||
|
|
||||||
|
def initialize(self, version: str, build_data: dict) -> None: # noqa: D401
|
||||||
|
root = Path(self.root)
|
||||||
|
webui_dir = root / "webui"
|
||||||
|
package_json = webui_dir / "package.json"
|
||||||
|
dist_dir = root / "nanobot" / "web" / "dist"
|
||||||
|
index_html = dist_dir / "index.html"
|
||||||
|
|
||||||
|
# `pip install -e .` builds an editable wheel; skip the (slow) webui
|
||||||
|
# bundle since editable installs target Python development and webui
|
||||||
|
# work uses `bun run dev` instead.
|
||||||
|
if self.target_name == "wheel" and version == "editable":
|
||||||
|
self.app.display_info(
|
||||||
|
"[webui-build] skipped for editable install "
|
||||||
|
"(use `cd webui && bun run build` to bundle webui manually)"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if os.environ.get("NANOBOT_SKIP_WEBUI_BUILD") == "1":
|
||||||
|
self.app.display_info("[webui-build] skipped via NANOBOT_SKIP_WEBUI_BUILD=1")
|
||||||
|
return
|
||||||
|
|
||||||
|
if not package_json.is_file():
|
||||||
|
self.app.display_info(
|
||||||
|
"[webui-build] no webui/ source tree, assuming prebuilt nanobot/web/dist/"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
force = os.environ.get("NANOBOT_FORCE_WEBUI_BUILD") == "1"
|
||||||
|
if index_html.is_file() and not force:
|
||||||
|
self.app.display_info(
|
||||||
|
f"[webui-build] reusing existing build at {dist_dir} "
|
||||||
|
"(set NANOBOT_FORCE_WEBUI_BUILD=1 to rebuild)"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
runner = self._pick_runner()
|
||||||
|
if runner is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"[webui-build] neither `bun` nor `npm` is available on PATH; "
|
||||||
|
"install one or set NANOBOT_SKIP_WEBUI_BUILD=1 to bypass."
|
||||||
|
)
|
||||||
|
|
||||||
|
self.app.display_info(f"[webui-build] using {runner} to build webui")
|
||||||
|
self._run([runner, "install"], cwd=webui_dir)
|
||||||
|
self._run([runner, "run", "build"], cwd=webui_dir)
|
||||||
|
|
||||||
|
if not index_html.is_file():
|
||||||
|
raise RuntimeError(
|
||||||
|
f"[webui-build] build finished but {index_html} is missing; "
|
||||||
|
"check webui/vite.config.ts outDir."
|
||||||
|
)
|
||||||
|
self.app.display_info(f"[webui-build] webui ready at {dist_dir}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _pick_runner() -> str | None:
|
||||||
|
for candidate in ("bun", "npm"):
|
||||||
|
if shutil.which(candidate):
|
||||||
|
return candidate
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _run(self, cmd: list[str], *, cwd: Path) -> None:
|
||||||
|
self.app.display_info(f"[webui-build] $ {' '.join(cmd)} (cwd={cwd})")
|
||||||
|
try:
|
||||||
|
subprocess.run(cmd, cwd=cwd, check=True)
|
||||||
|
except subprocess.CalledProcessError as exc:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"[webui-build] command failed ({exc.returncode}): {' '.join(cmd)}"
|
||||||
|
) from exc
|
||||||
+20
-4
@@ -2,9 +2,10 @@
|
|||||||
nanobot - A lightweight AI agent framework
|
nanobot - A lightweight AI agent framework
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from importlib.metadata import PackageNotFoundError, version as _pkg_version
|
|
||||||
from pathlib import Path
|
|
||||||
import tomllib
|
import tomllib
|
||||||
|
from importlib.metadata import PackageNotFoundError
|
||||||
|
from importlib.metadata import version as _pkg_version
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
def _read_pyproject_version() -> str | None:
|
def _read_pyproject_version() -> str | None:
|
||||||
@@ -21,12 +22,27 @@ def _resolve_version() -> str:
|
|||||||
return _pkg_version("nanobot-ai")
|
return _pkg_version("nanobot-ai")
|
||||||
except PackageNotFoundError:
|
except PackageNotFoundError:
|
||||||
# Source checkouts often import nanobot without installed dist-info.
|
# Source checkouts often import nanobot without installed dist-info.
|
||||||
return _read_pyproject_version() or "0.1.5.post3"
|
return _read_pyproject_version() or "0.2.0"
|
||||||
|
|
||||||
|
|
||||||
__version__ = _resolve_version()
|
__version__ = _resolve_version()
|
||||||
__logo__ = "🐈"
|
__logo__ = "🐈"
|
||||||
|
|
||||||
from nanobot.nanobot import Nanobot, RunResult
|
_LAZY_EXPORTS = {
|
||||||
|
"Nanobot": ".nanobot",
|
||||||
|
"RunResult": ".nanobot",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def __getattr__(name: str):
|
||||||
|
module_path = _LAZY_EXPORTS.get(name)
|
||||||
|
if module_path is None:
|
||||||
|
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||||
|
from importlib import import_module
|
||||||
|
mod = import_module(module_path, __name__)
|
||||||
|
val = getattr(mod, name)
|
||||||
|
globals()[name] = val
|
||||||
|
return val
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["Nanobot", "RunResult"]
|
__all__ = ["Nanobot", "RunResult"]
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from collections.abc import Collection
|
from collections.abc import Collection
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import TYPE_CHECKING, Any, Callable, Coroutine
|
from typing import TYPE_CHECKING, Callable, Coroutine
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -37,27 +37,6 @@ class AutoCompact:
|
|||||||
def _format_summary(text: str, last_active: datetime) -> str:
|
def _format_summary(text: str, last_active: datetime) -> str:
|
||||||
return f"Previous conversation summary (last active {last_active.isoformat()}):\n{text}"
|
return f"Previous conversation summary (last active {last_active.isoformat()}):\n{text}"
|
||||||
|
|
||||||
def _split_unconsolidated(
|
|
||||||
self, session: Session,
|
|
||||||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
|
||||||
"""Split live session tail into archiveable prefix and retained recent suffix."""
|
|
||||||
tail = list(session.messages[session.last_consolidated:])
|
|
||||||
if not tail:
|
|
||||||
return [], []
|
|
||||||
|
|
||||||
probe = Session(
|
|
||||||
key=session.key,
|
|
||||||
messages=tail.copy(),
|
|
||||||
created_at=session.created_at,
|
|
||||||
updated_at=session.updated_at,
|
|
||||||
metadata={},
|
|
||||||
last_consolidated=0,
|
|
||||||
)
|
|
||||||
probe.retain_recent_legal_suffix(self._RECENT_SUFFIX_MESSAGES)
|
|
||||||
kept = probe.messages
|
|
||||||
cut = len(tail) - len(kept)
|
|
||||||
return tail[:cut], kept
|
|
||||||
|
|
||||||
def check_expired(self, schedule_background: Callable[[Coroutine], None],
|
def check_expired(self, schedule_background: Callable[[Coroutine], None],
|
||||||
active_session_keys: Collection[str] = ()) -> None:
|
active_session_keys: Collection[str] = ()) -> None:
|
||||||
"""Schedule archival for idle sessions, skipping those with in-flight agent tasks."""
|
"""Schedule archival for idle sessions, skipping those with in-flight agent tasks."""
|
||||||
@@ -74,33 +53,17 @@ class AutoCompact:
|
|||||||
|
|
||||||
async def _archive(self, key: str) -> None:
|
async def _archive(self, key: str) -> None:
|
||||||
try:
|
try:
|
||||||
self.sessions.invalidate(key)
|
summary = await self.consolidator.compact_idle_session(
|
||||||
session = self.sessions.get_or_create(key)
|
key, self._RECENT_SUFFIX_MESSAGES,
|
||||||
archive_msgs, kept_msgs = self._split_unconsolidated(session)
|
)
|
||||||
if not archive_msgs and not kept_msgs:
|
|
||||||
session.updated_at = datetime.now()
|
|
||||||
self.sessions.save(session)
|
|
||||||
return
|
|
||||||
|
|
||||||
last_active = session.updated_at
|
|
||||||
summary = ""
|
|
||||||
if archive_msgs:
|
|
||||||
summary = await self.consolidator.archive(archive_msgs) or ""
|
|
||||||
if summary and summary != "(nothing)":
|
if summary and summary != "(nothing)":
|
||||||
self._summaries[key] = (summary, last_active)
|
session = self.sessions.get_or_create(key)
|
||||||
session.metadata["_last_summary"] = {"text": summary, "last_active": last_active.isoformat()}
|
meta = session.metadata.get("_last_summary")
|
||||||
session.messages = kept_msgs
|
if isinstance(meta, dict):
|
||||||
session.last_consolidated = 0
|
self._summaries[key] = (
|
||||||
session.updated_at = datetime.now()
|
meta["text"],
|
||||||
self.sessions.save(session)
|
datetime.fromisoformat(meta["last_active"]),
|
||||||
if archive_msgs:
|
)
|
||||||
logger.info(
|
|
||||||
"Auto-compact: archived {} (archived={}, kept={}, summary={})",
|
|
||||||
key,
|
|
||||||
len(archive_msgs),
|
|
||||||
len(kept_msgs),
|
|
||||||
bool(summary),
|
|
||||||
)
|
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Auto-compact: failed for {}", key)
|
logger.exception("Auto-compact: failed for {}", key)
|
||||||
finally:
|
finally:
|
||||||
|
|||||||
@@ -39,7 +39,6 @@ class ContextBuilder:
|
|||||||
skill_names: list[str] | None = None,
|
skill_names: list[str] | None = None,
|
||||||
channel: str | None = None,
|
channel: str | None = None,
|
||||||
session_summary: str | None = None,
|
session_summary: str | None = None,
|
||||||
session_key: str | None = None,
|
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Build the system prompt from identity, bootstrap files, memory, and skills."""
|
"""Build the system prompt from identity, bootstrap files, memory, and skills."""
|
||||||
parts = [self._get_identity(channel=channel)]
|
parts = [self._get_identity(channel=channel)]
|
||||||
@@ -74,29 +73,8 @@ class ContextBuilder:
|
|||||||
if session_summary:
|
if session_summary:
|
||||||
parts.append(f"[Archived Context Summary]\n\n{session_summary}")
|
parts.append(f"[Archived Context Summary]\n\n{session_summary}")
|
||||||
|
|
||||||
# Inject P2P collaboration hint for task-scoped sessions
|
|
||||||
if session_key and session_key.startswith("task:"):
|
|
||||||
parts.append(self._p2p_collaboration_hint())
|
|
||||||
|
|
||||||
return "\n\n---\n\n".join(parts)
|
return "\n\n---\n\n".join(parts)
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _p2p_collaboration_hint() -> str:
|
|
||||||
return (
|
|
||||||
"# Multi-Agent Collaboration\n\n"
|
|
||||||
"You are part of a decentralized agent network. You can:\n"
|
|
||||||
"- Use `broadcast_task` to announce subtasks and collect BIDs\n"
|
|
||||||
"- Use `dispatch_task` to assign tasks to specific agents\n"
|
|
||||||
"- Use `poll_task_result` to check task status\n"
|
|
||||||
"- Use `report_user` to deliver final results to the user\n"
|
|
||||||
"- Use `finalize_task` to terminate tasks\n\n"
|
|
||||||
"Rules:\n"
|
|
||||||
"- Never block waiting for results. Dispatch and continue.\n"
|
|
||||||
"- If a task times out, decide whether to retry, failover, or report partial.\n"
|
|
||||||
"- Respect the user's INTERRUPT messages — they have highest priority.\n"
|
|
||||||
"- You are currently in a task-scoped session; focus on the delegated task."
|
|
||||||
)
|
|
||||||
|
|
||||||
def _get_identity(self, channel: str | None = None) -> str:
|
def _get_identity(self, channel: str | None = None) -> str:
|
||||||
"""Get the core identity section."""
|
"""Get the core identity section."""
|
||||||
workspace_path = str(self.workspace.expanduser().resolve())
|
workspace_path = str(self.workspace.expanduser().resolve())
|
||||||
@@ -176,7 +154,6 @@ class ContextBuilder:
|
|||||||
sender_id: str | None = None,
|
sender_id: str | None = None,
|
||||||
session_summary: str | None = None,
|
session_summary: str | None = None,
|
||||||
session_metadata: Mapping[str, Any] | None = None,
|
session_metadata: Mapping[str, Any] | None = None,
|
||||||
session_key: str | None = None,
|
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Build the complete message list for an LLM call."""
|
"""Build the complete message list for an LLM call."""
|
||||||
extra = goal_state_runtime_lines(session_metadata)
|
extra = goal_state_runtime_lines(session_metadata)
|
||||||
@@ -198,7 +175,7 @@ class ContextBuilder:
|
|||||||
else:
|
else:
|
||||||
merged = user_content + [{"type": "text", "text": runtime_ctx}]
|
merged = user_content + [{"type": "text", "text": runtime_ctx}]
|
||||||
messages = [
|
messages = [
|
||||||
{"role": "system", "content": self.build_system_prompt(skill_names, channel=channel, session_summary=session_summary, session_key=session_key)},
|
{"role": "system", "content": self.build_system_prompt(skill_names, channel=channel, session_summary=session_summary)},
|
||||||
*history,
|
*history,
|
||||||
]
|
]
|
||||||
if messages[-1].get("role") == current_role:
|
if messages[-1].get("role") == current_role:
|
||||||
|
|||||||
+30
-101
@@ -24,14 +24,6 @@ from nanobot.agent.runner import _MAX_INJECTIONS_PER_TURN, AgentRunner, AgentRun
|
|||||||
from nanobot.agent.subagent import SubagentManager
|
from nanobot.agent.subagent import SubagentManager
|
||||||
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.p2p import (
|
|
||||||
BroadcastTaskTool,
|
|
||||||
CheckAggregationTool,
|
|
||||||
DispatchTaskTool,
|
|
||||||
FinalizeTaskTool,
|
|
||||||
PollTaskResultTool,
|
|
||||||
ReportUserTool,
|
|
||||||
)
|
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.agent.tools.self import MyTool
|
from nanobot.agent.tools.self import MyTool
|
||||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||||
@@ -41,19 +33,20 @@ from nanobot.config.schema import AgentDefaults, ModelPresetConfig
|
|||||||
from nanobot.providers.base import LLMProvider
|
from nanobot.providers.base import LLMProvider
|
||||||
from nanobot.providers.factory import ProviderSnapshot
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
from nanobot.session.goal_state import (
|
from nanobot.session.goal_state import (
|
||||||
goal_state_ws_blob,
|
|
||||||
runner_wall_llm_timeout_s,
|
runner_wall_llm_timeout_s,
|
||||||
)
|
)
|
||||||
from nanobot.session.manager import Session, SessionManager
|
from nanobot.session.manager import Session, SessionManager
|
||||||
from nanobot.utils.artifacts import generated_image_paths_from_messages
|
from nanobot.session.webui_turns import (
|
||||||
|
WebuiTurnCoordinator,
|
||||||
|
build_bus_progress_callback,
|
||||||
|
mark_webui_session,
|
||||||
|
)
|
||||||
from nanobot.utils.document import extract_documents
|
from nanobot.utils.document import extract_documents
|
||||||
from nanobot.utils.helpers import image_placeholder_text
|
from nanobot.utils.helpers import image_placeholder_text
|
||||||
from nanobot.utils.helpers import truncate_text as truncate_text_fn
|
from nanobot.utils.helpers import truncate_text as truncate_text_fn
|
||||||
from nanobot.utils.image_generation_intent import image_generation_prompt
|
from nanobot.utils.image_generation_intent import image_generation_prompt
|
||||||
|
from nanobot.utils.llm_runtime import LLMRuntime
|
||||||
from nanobot.utils.runtime import EMPTY_FINAL_RESPONSE_MESSAGE
|
from nanobot.utils.runtime import EMPTY_FINAL_RESPONSE_MESSAGE
|
||||||
from nanobot.utils.session_attachments import merge_turn_media_into_last_assistant
|
|
||||||
from nanobot.utils.webui_titles import mark_webui_session, maybe_generate_webui_title_after_turn
|
|
||||||
from nanobot.utils.webui_turn_helpers import publish_turn_run_status
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nanobot.config.schema import (
|
from nanobot.config.schema import (
|
||||||
@@ -108,7 +101,6 @@ class TurnContext:
|
|||||||
save_skip: int = 0
|
save_skip: int = 0
|
||||||
|
|
||||||
outbound: OutboundMessage | None = None
|
outbound: OutboundMessage | None = None
|
||||||
generated_media: list[str] = field(default_factory=list)
|
|
||||||
|
|
||||||
on_progress: Callable[..., Awaitable[None]] | None = None
|
on_progress: Callable[..., Awaitable[None]] | None = None
|
||||||
on_stream: Callable[[str], Awaitable[None]] | None = None
|
on_stream: Callable[[str], Awaitable[None]] | None = None
|
||||||
@@ -144,6 +136,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
|
||||||
|
|
||||||
|
def llm_runtime(self) -> LLMRuntime:
|
||||||
|
"""Return the current provider/model pair owned by this loop."""
|
||||||
|
self._refresh_provider_snapshot()
|
||||||
|
return LLMRuntime(self.provider, self.model)
|
||||||
|
|
||||||
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||||
_PENDING_USER_TURN_KEY = "pending_user_turn"
|
_PENDING_USER_TURN_KEY = "pending_user_turn"
|
||||||
|
|
||||||
@@ -193,7 +190,6 @@ class AgentLoop:
|
|||||||
model_preset: str | None = None,
|
model_preset: str | None = None,
|
||||||
preset_snapshot_loader: preset_helpers.PresetSnapshotLoader | None = None,
|
preset_snapshot_loader: preset_helpers.PresetSnapshotLoader | None = None,
|
||||||
runtime_model_publisher: Callable[[str, str | None], None] | None = None,
|
runtime_model_publisher: Callable[[str, str | None], None] | None = None,
|
||||||
p2p_shell: Any | None = None,
|
|
||||||
):
|
):
|
||||||
from nanobot.config.schema import ToolsConfig
|
from nanobot.config.schema import ToolsConfig
|
||||||
|
|
||||||
@@ -201,7 +197,6 @@ class AgentLoop:
|
|||||||
defaults = AgentDefaults()
|
defaults = AgentDefaults()
|
||||||
self.bus = bus
|
self.bus = bus
|
||||||
self.channels_config = channels_config
|
self.channels_config = channels_config
|
||||||
self.p2p_shell = p2p_shell
|
|
||||||
self.provider = provider
|
self.provider = provider
|
||||||
self._provider_snapshot_loader = provider_snapshot_loader
|
self._provider_snapshot_loader = provider_snapshot_loader
|
||||||
self._preset_snapshot_loader = preset_snapshot_loader
|
self._preset_snapshot_loader = preset_snapshot_loader
|
||||||
@@ -247,6 +242,11 @@ 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._webui_turns = WebuiTurnCoordinator(
|
||||||
|
bus=self.bus,
|
||||||
|
sessions=self.sessions,
|
||||||
|
schedule_background=lambda coro: self._schedule_background(coro),
|
||||||
|
)
|
||||||
self.tools = ToolRegistry()
|
self.tools = 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.
|
||||||
@@ -473,22 +473,6 @@ class AgentLoop:
|
|||||||
)
|
)
|
||||||
registered.append("my")
|
registered.append("my")
|
||||||
|
|
||||||
# Register P2P tools if enabled
|
|
||||||
if self.p2p_shell:
|
|
||||||
self.tools.register(DispatchTaskTool(shell=self.p2p_shell))
|
|
||||||
self.tools.register(PollTaskResultTool(shell=self.p2p_shell))
|
|
||||||
self.tools.register(BroadcastTaskTool(shell=self.p2p_shell))
|
|
||||||
self.tools.register(CheckAggregationTool(shell=self.p2p_shell))
|
|
||||||
self.tools.register(
|
|
||||||
ReportUserTool(
|
|
||||||
send_callback=self.bus.publish_outbound,
|
|
||||||
default_channel=getattr(self.channels_config, "default_channel", ""),
|
|
||||||
default_chat_id=getattr(self.channels_config, "default_chat_id", ""),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
self.tools.register(FinalizeTaskTool(shell=self.p2p_shell, session_manager=self.sessions))
|
|
||||||
registered.append("p2p")
|
|
||||||
|
|
||||||
logger.info("Registered {} tools: {}", len(registered), registered)
|
logger.info("Registered {} tools: {}", len(registered), registered)
|
||||||
|
|
||||||
async def _connect_mcp(self) -> None:
|
async def _connect_mcp(self) -> None:
|
||||||
@@ -550,34 +534,7 @@ class AgentLoop:
|
|||||||
self, msg: InboundMessage
|
self, msg: InboundMessage
|
||||||
) -> Callable[..., Awaitable[None]]:
|
) -> Callable[..., Awaitable[None]]:
|
||||||
"""Build a progress callback that publishes to the message bus."""
|
"""Build a progress callback that publishes to the message bus."""
|
||||||
|
return build_bus_progress_callback(self.bus, msg)
|
||||||
async def _bus_progress(
|
|
||||||
content: str,
|
|
||||||
*,
|
|
||||||
tool_hint: bool = False,
|
|
||||||
tool_events: list[dict[str, Any]] | None = None,
|
|
||||||
reasoning: bool = False,
|
|
||||||
reasoning_end: bool = False,
|
|
||||||
) -> None:
|
|
||||||
meta = dict(msg.metadata or {})
|
|
||||||
meta["_progress"] = True
|
|
||||||
meta["_tool_hint"] = tool_hint
|
|
||||||
if reasoning:
|
|
||||||
meta["_reasoning_delta"] = True
|
|
||||||
if reasoning_end:
|
|
||||||
meta["_reasoning_end"] = True
|
|
||||||
if tool_events:
|
|
||||||
meta["_tool_events"] = tool_events
|
|
||||||
await self.bus.publish_outbound(
|
|
||||||
OutboundMessage(
|
|
||||||
channel=msg.channel,
|
|
||||||
chat_id=msg.chat_id,
|
|
||||||
content=content,
|
|
||||||
metadata=meta,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
return _bus_progress
|
|
||||||
|
|
||||||
async def _build_retry_wait_callback(
|
async def _build_retry_wait_callback(
|
||||||
self, msg: InboundMessage
|
self, msg: InboundMessage
|
||||||
@@ -964,38 +921,12 @@ class AgentLoop:
|
|||||||
content="", metadata=msg.metadata or {},
|
content="", metadata=msg.metadata or {},
|
||||||
))
|
))
|
||||||
if msg.channel == "websocket":
|
if msg.channel == "websocket":
|
||||||
# Signal that the turn is fully complete (all tools executed,
|
|
||||||
# final text streamed). This lets WS clients know when to
|
|
||||||
# definitively stop the loading indicator.
|
|
||||||
turn_lat = self._pending_turn_latency_ms.pop(session_key, None)
|
turn_lat = self._pending_turn_latency_ms.pop(session_key, None)
|
||||||
turn_metadata: dict[str, Any] = {**msg.metadata, "_turn_end": True}
|
await self._webui_turns.handle_turn_end(
|
||||||
if turn_lat is not None:
|
msg,
|
||||||
turn_metadata["latency_ms"] = int(turn_lat)
|
session_key=session_key,
|
||||||
sess_turn = self.sessions.get_or_create(session_key)
|
latency_ms=turn_lat,
|
||||||
turn_metadata["goal_state"] = goal_state_ws_blob(sess_turn.metadata)
|
)
|
||||||
await self.bus.publish_outbound(OutboundMessage(
|
|
||||||
channel=msg.channel, chat_id=msg.chat_id,
|
|
||||||
content="", metadata=turn_metadata,
|
|
||||||
))
|
|
||||||
if msg.metadata.get("webui") is True:
|
|
||||||
async def _generate_title_and_notify() -> None:
|
|
||||||
generated = await maybe_generate_webui_title_after_turn(
|
|
||||||
channel=msg.channel,
|
|
||||||
metadata=msg.metadata,
|
|
||||||
sessions=self.sessions,
|
|
||||||
session_key=session_key,
|
|
||||||
provider=self.provider,
|
|
||||||
model=self.model,
|
|
||||||
)
|
|
||||||
if generated:
|
|
||||||
await self.bus.publish_outbound(OutboundMessage(
|
|
||||||
channel=msg.channel,
|
|
||||||
chat_id=msg.chat_id,
|
|
||||||
content="",
|
|
||||||
metadata={**msg.metadata, "_session_updated": True},
|
|
||||||
))
|
|
||||||
|
|
||||||
self._schedule_background(_generate_title_and_notify())
|
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
logger.info("Task cancelled for session {}", session_key)
|
logger.info("Task cancelled for session {}", session_key)
|
||||||
# Preserve partial context from the interrupted turn so
|
# Preserve partial context from the interrupted turn so
|
||||||
@@ -1047,8 +978,9 @@ class AgentLoop:
|
|||||||
"Re-published {} leftover message(s) to bus for session {}",
|
"Re-published {} leftover message(s) to bus for session {}",
|
||||||
leftover, session_key,
|
leftover, session_key,
|
||||||
)
|
)
|
||||||
await publish_turn_run_status(self.bus, msg, "idle")
|
await self._webui_turns.publish_run_status(msg, "idle")
|
||||||
self._pending_turn_latency_ms.pop(session_key, None)
|
self._pending_turn_latency_ms.pop(session_key, None)
|
||||||
|
self._webui_turns.discard(session_key)
|
||||||
|
|
||||||
async def close_mcp(self) -> None:
|
async def close_mcp(self) -> None:
|
||||||
"""Drain pending background archives, then close MCP connections."""
|
"""Drain pending background archives, then close MCP connections."""
|
||||||
@@ -1259,7 +1191,6 @@ class AgentLoop:
|
|||||||
all_msgs: list[dict[str, Any]],
|
all_msgs: list[dict[str, Any]],
|
||||||
stop_reason: str,
|
stop_reason: str,
|
||||||
had_injections: bool,
|
had_injections: bool,
|
||||||
generated_media: list[str],
|
|
||||||
on_stream: Callable[[str], Awaitable[None]] | None,
|
on_stream: Callable[[str], Awaitable[None]] | None,
|
||||||
*,
|
*,
|
||||||
turn_latency_ms: int | None = None,
|
turn_latency_ms: int | None = None,
|
||||||
@@ -1283,7 +1214,6 @@ class AgentLoop:
|
|||||||
channel=msg.channel,
|
channel=msg.channel,
|
||||||
chat_id=msg.chat_id,
|
chat_id=msg.chat_id,
|
||||||
content=final_content,
|
content=final_content,
|
||||||
media=generated_media,
|
|
||||||
metadata=meta,
|
metadata=meta,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1364,6 +1294,11 @@ class AgentLoop:
|
|||||||
"include_timestamps": True,
|
"include_timestamps": True,
|
||||||
}
|
}
|
||||||
ctx.history = ctx.session.get_history(**_hist_kwargs)
|
ctx.history = ctx.session.get_history(**_hist_kwargs)
|
||||||
|
self._webui_turns.capture_title_context(
|
||||||
|
ctx.session_key,
|
||||||
|
ctx.msg,
|
||||||
|
self.llm_runtime(),
|
||||||
|
)
|
||||||
|
|
||||||
ctx.initial_messages = self._build_initial_messages(
|
ctx.initial_messages = self._build_initial_messages(
|
||||||
ctx.msg, ctx.session, ctx.history, ctx.pending_summary
|
ctx.msg, ctx.session, ctx.history, ctx.pending_summary
|
||||||
@@ -1380,7 +1315,7 @@ class AgentLoop:
|
|||||||
return "ok"
|
return "ok"
|
||||||
|
|
||||||
async def _state_run(self, ctx: TurnContext) -> str:
|
async def _state_run(self, ctx: TurnContext) -> str:
|
||||||
await publish_turn_run_status(self.bus, ctx.msg, "running")
|
await self._webui_turns.publish_run_status(ctx.msg, "running")
|
||||||
result = await self._run_agent_loop(
|
result = await self._run_agent_loop(
|
||||||
ctx.initial_messages,
|
ctx.initial_messages,
|
||||||
on_progress=ctx.on_progress,
|
on_progress=ctx.on_progress,
|
||||||
@@ -1408,11 +1343,6 @@ class AgentLoop:
|
|||||||
ctx.final_content = EMPTY_FINAL_RESPONSE_MESSAGE
|
ctx.final_content = EMPTY_FINAL_RESPONSE_MESSAGE
|
||||||
|
|
||||||
ctx.save_skip = 1 + len(ctx.history) + (1 if ctx.user_persisted_early else 0)
|
ctx.save_skip = 1 + len(ctx.history) + (1 if ctx.user_persisted_early else 0)
|
||||||
skip_msgs = ctx.all_messages[ctx.save_skip:]
|
|
||||||
ctx.generated_media = generated_image_paths_from_messages(skip_msgs)
|
|
||||||
mt = self.tools.get("message")
|
|
||||||
extra = getattr(mt, "turn_delivered_media_paths", lambda: [])() if mt else []
|
|
||||||
merge_turn_media_into_last_assistant(ctx.all_messages, ctx.generated_media, extra)
|
|
||||||
|
|
||||||
ctx.turn_latency_ms = max(0, int((time.time() - ctx.turn_wall_started_at) * 1000))
|
ctx.turn_latency_ms = max(0, int((time.time() - ctx.turn_wall_started_at) * 1000))
|
||||||
self._save_turn(
|
self._save_turn(
|
||||||
@@ -1440,7 +1370,6 @@ class AgentLoop:
|
|||||||
ctx.all_messages,
|
ctx.all_messages,
|
||||||
ctx.stop_reason,
|
ctx.stop_reason,
|
||||||
ctx.had_injections,
|
ctx.had_injections,
|
||||||
ctx.generated_media,
|
|
||||||
ctx.on_stream,
|
ctx.on_stream,
|
||||||
turn_latency_ms=ctx.turn_latency_ms,
|
turn_latency_ms=ctx.turn_latency_ms,
|
||||||
)
|
)
|
||||||
|
|||||||
+76
-1
@@ -678,11 +678,18 @@ class Consolidator:
|
|||||||
The budget reserves space for completion tokens and a safety buffer
|
The budget reserves space for completion tokens and a safety buffer
|
||||||
so the LLM request never exceeds the context window.
|
so the LLM request never exceeds the context window.
|
||||||
"""
|
"""
|
||||||
if not session.messages or self.context_window_tokens <= 0:
|
if self.context_window_tokens <= 0:
|
||||||
return
|
return
|
||||||
|
|
||||||
lock = self.get_lock(session.key)
|
lock = self.get_lock(session.key)
|
||||||
async with lock:
|
async with lock:
|
||||||
|
# Refresh session reference: AutoCompact may have replaced it.
|
||||||
|
fresh = self.sessions.get_or_create(session.key)
|
||||||
|
if fresh is not session:
|
||||||
|
session = fresh
|
||||||
|
if not session.messages:
|
||||||
|
return
|
||||||
|
|
||||||
budget = self._input_token_budget
|
budget = self._input_token_budget
|
||||||
target = int(budget * self.consolidation_ratio)
|
target = int(budget * self.consolidation_ratio)
|
||||||
last_summary = await self._consolidate_replay_overflow(
|
last_summary = await self._consolidate_replay_overflow(
|
||||||
@@ -769,6 +776,74 @@ class Consolidator:
|
|||||||
# the summary injection strategy with AutoCompact._archive().
|
# the summary injection strategy with AutoCompact._archive().
|
||||||
self._persist_last_summary(session, last_summary)
|
self._persist_last_summary(session, last_summary)
|
||||||
|
|
||||||
|
async def compact_idle_session(
|
||||||
|
self,
|
||||||
|
session_key: str,
|
||||||
|
max_suffix: int = 8,
|
||||||
|
) -> str | None:
|
||||||
|
"""Hard-truncate an idle session under the consolidation lock.
|
||||||
|
|
||||||
|
Used by AutoCompact so all session mutation goes through a single
|
||||||
|
lock-protected path. Returns the summary text on success, ``None``
|
||||||
|
if the LLM failed (raw_archive fallback), or ``""`` if there was
|
||||||
|
nothing to archive.
|
||||||
|
"""
|
||||||
|
lock = self.get_lock(session_key)
|
||||||
|
async with lock:
|
||||||
|
self.sessions.invalidate(session_key)
|
||||||
|
session = self.sessions.get_or_create(session_key)
|
||||||
|
|
||||||
|
tail = list(session.messages[session.last_consolidated:])
|
||||||
|
if not tail:
|
||||||
|
session.updated_at = datetime.now()
|
||||||
|
self.sessions.save(session)
|
||||||
|
return ""
|
||||||
|
|
||||||
|
probe = Session(
|
||||||
|
key=session.key,
|
||||||
|
messages=tail.copy(),
|
||||||
|
created_at=session.created_at,
|
||||||
|
updated_at=session.updated_at,
|
||||||
|
metadata={},
|
||||||
|
last_consolidated=0,
|
||||||
|
)
|
||||||
|
probe.retain_recent_legal_suffix(max_suffix)
|
||||||
|
kept = probe.messages
|
||||||
|
cut = len(tail) - len(kept)
|
||||||
|
archive_msgs = tail[:cut]
|
||||||
|
|
||||||
|
if not archive_msgs and not kept:
|
||||||
|
session.updated_at = datetime.now()
|
||||||
|
self.sessions.save(session)
|
||||||
|
return ""
|
||||||
|
|
||||||
|
last_active = session.updated_at
|
||||||
|
summary: str | None = ""
|
||||||
|
if archive_msgs:
|
||||||
|
summary = await self.archive(archive_msgs)
|
||||||
|
|
||||||
|
if summary and summary != "(nothing)":
|
||||||
|
session.metadata["_last_summary"] = {
|
||||||
|
"text": summary,
|
||||||
|
"last_active": last_active.isoformat(),
|
||||||
|
}
|
||||||
|
|
||||||
|
session.messages = kept
|
||||||
|
session.last_consolidated = 0
|
||||||
|
session.updated_at = datetime.now()
|
||||||
|
self.sessions.save(session)
|
||||||
|
|
||||||
|
if archive_msgs:
|
||||||
|
logger.info(
|
||||||
|
"Idle-session compact for {}: archived={}, kept={}, summary={}",
|
||||||
|
session_key,
|
||||||
|
len(archive_msgs),
|
||||||
|
len(kept),
|
||||||
|
bool(summary),
|
||||||
|
)
|
||||||
|
|
||||||
|
return summary
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Dream — heavyweight cron-scheduled memory consolidation
|
# Dream — heavyweight cron-scheduled memory consolidation
|
||||||
|
|||||||
@@ -15,6 +15,13 @@ from loguru import logger
|
|||||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
||||||
|
from nanobot.utils.file_edit_events import (
|
||||||
|
build_file_edit_end_event,
|
||||||
|
build_file_edit_error_event,
|
||||||
|
build_file_edit_start_event,
|
||||||
|
prepare_file_edit_tracker,
|
||||||
|
StreamingFileEditTracker,
|
||||||
|
)
|
||||||
from nanobot.utils.helpers import (
|
from nanobot.utils.helpers import (
|
||||||
IncrementalThinkExtractor,
|
IncrementalThinkExtractor,
|
||||||
build_assistant_message,
|
build_assistant_message,
|
||||||
@@ -26,6 +33,10 @@ from nanobot.utils.helpers import (
|
|||||||
strip_think,
|
strip_think,
|
||||||
truncate_text,
|
truncate_text,
|
||||||
)
|
)
|
||||||
|
from nanobot.utils.progress_events import (
|
||||||
|
invoke_file_edit_progress,
|
||||||
|
on_progress_accepts_file_edit_events,
|
||||||
|
)
|
||||||
from nanobot.utils.prompt_templates import render_template
|
from nanobot.utils.prompt_templates import render_template
|
||||||
from nanobot.utils.runtime import (
|
from nanobot.utils.runtime import (
|
||||||
EMPTY_FINAL_RESPONSE_MESSAGE,
|
EMPTY_FINAL_RESPONSE_MESSAGE,
|
||||||
@@ -619,6 +630,24 @@ class AgentRunner:
|
|||||||
)
|
)
|
||||||
|
|
||||||
progress_state: dict[str, bool] | None = None
|
progress_state: dict[str, bool] | None = None
|
||||||
|
live_file_edits: StreamingFileEditTracker | None = None
|
||||||
|
|
||||||
|
if (
|
||||||
|
spec.progress_callback is not None
|
||||||
|
and on_progress_accepts_file_edit_events(spec.progress_callback)
|
||||||
|
):
|
||||||
|
async def _emit_live_file_edits(events: list[dict[str, Any]]) -> None:
|
||||||
|
await invoke_file_edit_progress(spec.progress_callback, events)
|
||||||
|
|
||||||
|
live_file_edits = StreamingFileEditTracker(
|
||||||
|
workspace=spec.workspace,
|
||||||
|
tools=spec.tools,
|
||||||
|
emit=_emit_live_file_edits,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _tool_call_delta(delta: dict[str, Any]) -> None:
|
||||||
|
if live_file_edits is not None:
|
||||||
|
await live_file_edits.update(delta)
|
||||||
|
|
||||||
if wants_streaming:
|
if wants_streaming:
|
||||||
async def _stream(delta: str) -> None:
|
async def _stream(delta: str) -> None:
|
||||||
@@ -636,6 +665,7 @@ class AgentRunner:
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
on_content_delta=_stream,
|
on_content_delta=_stream,
|
||||||
on_thinking_delta=_thinking,
|
on_thinking_delta=_thinking,
|
||||||
|
on_tool_call_delta=_tool_call_delta if live_file_edits is not None else None,
|
||||||
)
|
)
|
||||||
elif wants_progress_streaming:
|
elif wants_progress_streaming:
|
||||||
stream_buf = ""
|
stream_buf = ""
|
||||||
@@ -665,6 +695,7 @@ class AgentRunner:
|
|||||||
coro = self.provider.chat_stream_with_retry(
|
coro = self.provider.chat_stream_with_retry(
|
||||||
**kwargs,
|
**kwargs,
|
||||||
on_content_delta=_stream_progress,
|
on_content_delta=_stream_progress,
|
||||||
|
on_tool_call_delta=_tool_call_delta if live_file_edits is not None else None,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
coro = self.provider.chat_with_retry(**kwargs)
|
coro = self.provider.chat_with_retry(**kwargs)
|
||||||
@@ -679,6 +710,14 @@ class AgentRunner:
|
|||||||
await coro if outer_timeout_s is None
|
await coro if outer_timeout_s is None
|
||||||
else await asyncio.wait_for(coro, timeout=outer_timeout_s)
|
else await asyncio.wait_for(coro, timeout=outer_timeout_s)
|
||||||
)
|
)
|
||||||
|
if live_file_edits is not None:
|
||||||
|
await live_file_edits.flush()
|
||||||
|
if response.should_execute_tools:
|
||||||
|
live_file_edits.apply_final_call_ids(response.tool_calls)
|
||||||
|
await live_file_edits.error_unmatched(
|
||||||
|
response.tool_calls if response.should_execute_tools else [],
|
||||||
|
"Tool call did not complete.",
|
||||||
|
)
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
if outer_timeout_s is None:
|
if outer_timeout_s is None:
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
@@ -813,6 +852,30 @@ class AgentRunner:
|
|||||||
return prep_error + hint, event, (
|
return prep_error + hint, event, (
|
||||||
RuntimeError(prep_error) if spec.fail_on_tool_error else None
|
RuntimeError(prep_error) if spec.fail_on_tool_error else None
|
||||||
)
|
)
|
||||||
|
emit_file_edit_events = (
|
||||||
|
spec.progress_callback is not None
|
||||||
|
and on_progress_accepts_file_edit_events(spec.progress_callback)
|
||||||
|
)
|
||||||
|
progress_callback = spec.progress_callback if emit_file_edit_events else None
|
||||||
|
file_edit_tracker = (
|
||||||
|
prepare_file_edit_tracker(
|
||||||
|
call_id=tool_call.id,
|
||||||
|
tool_name=tool_call.name,
|
||||||
|
tool=tool,
|
||||||
|
workspace=spec.workspace,
|
||||||
|
params=params if isinstance(params, dict) else None,
|
||||||
|
)
|
||||||
|
if progress_callback is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
if file_edit_tracker is not None and progress_callback is not None:
|
||||||
|
await invoke_file_edit_progress(
|
||||||
|
progress_callback,
|
||||||
|
[build_file_edit_start_event(
|
||||||
|
file_edit_tracker,
|
||||||
|
params if isinstance(params, dict) else None,
|
||||||
|
)],
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
if tool is not None:
|
if tool is not None:
|
||||||
result = await tool.execute(**params)
|
result = await tool.execute(**params)
|
||||||
@@ -821,6 +884,11 @@ class AgentRunner:
|
|||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
raise
|
raise
|
||||||
except BaseException as exc:
|
except BaseException as exc:
|
||||||
|
if file_edit_tracker is not None and progress_callback is not None:
|
||||||
|
await invoke_file_edit_progress(
|
||||||
|
progress_callback,
|
||||||
|
[build_file_edit_error_event(file_edit_tracker, str(exc))],
|
||||||
|
)
|
||||||
event = {
|
event = {
|
||||||
"name": tool_call.name,
|
"name": tool_call.name,
|
||||||
"status": "error",
|
"status": "error",
|
||||||
@@ -842,6 +910,11 @@ class AgentRunner:
|
|||||||
return payload, event, None
|
return payload, event, None
|
||||||
|
|
||||||
if isinstance(result, str) and result.startswith("Error"):
|
if isinstance(result, str) and result.startswith("Error"):
|
||||||
|
if file_edit_tracker is not None and progress_callback is not None:
|
||||||
|
await invoke_file_edit_progress(
|
||||||
|
progress_callback,
|
||||||
|
[build_file_edit_error_event(file_edit_tracker, result)],
|
||||||
|
)
|
||||||
event = {
|
event = {
|
||||||
"name": tool_call.name,
|
"name": tool_call.name,
|
||||||
"status": "error",
|
"status": "error",
|
||||||
@@ -860,6 +933,15 @@ class AgentRunner:
|
|||||||
return result + hint, event, RuntimeError(result)
|
return result + hint, event, RuntimeError(result)
|
||||||
return result + hint, event, None
|
return result + hint, event, None
|
||||||
|
|
||||||
|
if file_edit_tracker is not None and progress_callback is not None:
|
||||||
|
await invoke_file_edit_progress(
|
||||||
|
progress_callback,
|
||||||
|
[build_file_edit_end_event(
|
||||||
|
file_edit_tracker,
|
||||||
|
params if isinstance(params, dict) else None,
|
||||||
|
)],
|
||||||
|
)
|
||||||
|
|
||||||
detail = "" if result is None else str(result)
|
detail = "" if result is None else str(result)
|
||||||
detail = detail.replace("\n", " ").strip()
|
detail = detail.replace("\n", " ").strip()
|
||||||
if not detail:
|
if not detail:
|
||||||
|
|||||||
@@ -17,9 +17,9 @@ from nanobot.agent.tools.schema import (
|
|||||||
from nanobot.config.paths import get_media_dir
|
from nanobot.config.paths import get_media_dir
|
||||||
from nanobot.config.schema import Base
|
from nanobot.config.schema import Base
|
||||||
from nanobot.providers.image_generation import (
|
from nanobot.providers.image_generation import (
|
||||||
AIHubMixImageGenerationClient,
|
|
||||||
ImageGenerationError,
|
ImageGenerationError,
|
||||||
OpenRouterImageGenerationClient,
|
ImageGenerationProvider,
|
||||||
|
get_image_gen_provider,
|
||||||
)
|
)
|
||||||
from nanobot.utils.artifacts import (
|
from nanobot.utils.artifacts import (
|
||||||
ArtifactError,
|
ArtifactError,
|
||||||
@@ -117,27 +117,24 @@ class ImageGenerationTool(Tool):
|
|||||||
def _provider_config(self) -> ProviderConfig | None:
|
def _provider_config(self) -> ProviderConfig | None:
|
||||||
return self.provider_configs.get(self.config.provider)
|
return self.provider_configs.get(self.config.provider)
|
||||||
|
|
||||||
def _provider_client(self) -> OpenRouterImageGenerationClient | AIHubMixImageGenerationClient | None:
|
def _provider_client(self) -> ImageGenerationProvider | None:
|
||||||
provider = self._provider_config()
|
provider = self._provider_config()
|
||||||
|
cls = get_image_gen_provider(self.config.provider)
|
||||||
|
if cls is None:
|
||||||
|
return None
|
||||||
kwargs = {
|
kwargs = {
|
||||||
"api_key": provider.api_key if provider else None,
|
"api_key": provider.api_key if provider else None,
|
||||||
"api_base": provider.api_base if provider else None,
|
"api_base": provider.api_base if provider else None,
|
||||||
"extra_headers": provider.extra_headers if provider else None,
|
"extra_headers": provider.extra_headers if provider else None,
|
||||||
"extra_body": provider.extra_body if provider else None,
|
"extra_body": provider.extra_body if provider else None,
|
||||||
}
|
}
|
||||||
if self.config.provider == "openrouter":
|
return cls(**kwargs)
|
||||||
return OpenRouterImageGenerationClient(**kwargs)
|
|
||||||
if self.config.provider == "aihubmix":
|
|
||||||
return AIHubMixImageGenerationClient(**kwargs)
|
|
||||||
return None
|
|
||||||
|
|
||||||
def _missing_api_key_error(self) -> str:
|
def _missing_api_key_error(self) -> str:
|
||||||
provider = self.config.provider
|
cls = get_image_gen_provider(self.config.provider)
|
||||||
if provider == "openrouter":
|
if cls and cls.missing_key_message:
|
||||||
return "Error: OpenRouter API key is not configured. Set providers.openrouter.apiKey."
|
return f"Error: {cls.missing_key_message}"
|
||||||
if provider == "aihubmix":
|
return f"Error: {self.config.provider} API key is not configured."
|
||||||
return "Error: AIHubMix API key is not configured. Set providers.aihubmix.apiKey."
|
|
||||||
return f"Error: {provider} API key is not configured."
|
|
||||||
|
|
||||||
def _resolve_reference_image(self, value: str) -> str:
|
def _resolve_reference_image(self, value: str) -> str:
|
||||||
raw_path = Path(value).expanduser()
|
raw_path = Path(value).expanduser()
|
||||||
|
|||||||
@@ -31,8 +31,8 @@ from nanobot.config.paths import get_workspace_path
|
|||||||
media=ArraySchema(
|
media=ArraySchema(
|
||||||
StringSchema(""),
|
StringSchema(""),
|
||||||
description=(
|
description=(
|
||||||
"Optional list of existing file paths to attach for proactive or cross-channel delivery. "
|
"Optional list of existing file paths to attach. "
|
||||||
"Do not use this to resend generate_image outputs in the current chat."
|
"Use artifact paths returned by generate_image here when delivering generated images."
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
buttons=ArraySchema(
|
buttons=ArraySchema(
|
||||||
@@ -140,8 +140,8 @@ class MessageTool(Tool, ContextAware):
|
|||||||
"Do not use this for the normal reply in the current chat: answer naturally instead. "
|
"Do not use this for the normal reply in the current chat: answer naturally instead. "
|
||||||
"If channel/chat_id would target the current runtime conversation, do not call this tool "
|
"If channel/chat_id would target the current runtime conversation, do not call this tool "
|
||||||
"unless the user explicitly asked you to proactively send an existing file attachment. "
|
"unless the user explicitly asked you to proactively send an existing file attachment. "
|
||||||
"When generate_image creates images in the current chat, the final assistant reply "
|
"When generate_image creates images in the current chat, use the message tool "
|
||||||
"automatically attaches them; do not call message just to announce or resend them. "
|
"with the artifact paths in the media parameter to deliver the images to the user. "
|
||||||
"For proactive attachment delivery, use the 'media' parameter with file paths. "
|
"For proactive attachment delivery, use the 'media' parameter with file paths. "
|
||||||
"Do NOT use read_file to send files — that only reads content for your own analysis."
|
"Do NOT use read_file to send files — that only reads content for your own analysis."
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,328 +0,0 @@
|
|||||||
"""P2P tools for inter-agent task dispatch and coordination."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from typing import Any, Awaitable, Callable
|
|
||||||
|
|
||||||
from nanobot.agent.tools.base import Tool
|
|
||||||
from nanobot.bus.events import OutboundMessage
|
|
||||||
|
|
||||||
|
|
||||||
class DispatchTaskTool(Tool):
|
|
||||||
"""Asynchronously dispatch a task to another agent. Non-blocking."""
|
|
||||||
|
|
||||||
def __init__(self, shell: "P2PShell"):
|
|
||||||
self._shell = shell
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "dispatch_task"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Dispatch a task to a specific target agent. Returns immediately with a receipt. "
|
|
||||||
"The target agent will process the task independently. Use poll_task_result later to check completion. "
|
|
||||||
"Do NOT block waiting for results."
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def parameters(self) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"to": {"type": "string", "description": "Target agent ID"},
|
|
||||||
"task_description": {"type": "string", "description": "Clear description of the task"},
|
|
||||||
"parent_task_id": {"type": "string", "description": "Parent task ID for ancestry tracking"},
|
|
||||||
"deadline_seconds": {"type": "integer", "default": 300, "description": "Task deadline in seconds"},
|
|
||||||
"allow_redelegation": {"type": "boolean", "default": True, "description": "Whether the target may re-delegate"},
|
|
||||||
},
|
|
||||||
"required": ["to", "task_description"],
|
|
||||||
}
|
|
||||||
|
|
||||||
async def execute(
|
|
||||||
self,
|
|
||||||
to: str,
|
|
||||||
task_description: str,
|
|
||||||
parent_task_id: str | None = None,
|
|
||||||
deadline_seconds: int = 300,
|
|
||||||
allow_redelegation: bool = True,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> str:
|
|
||||||
result = self._shell.dispatch(
|
|
||||||
to=to,
|
|
||||||
parent_task_id=parent_task_id,
|
|
||||||
description=task_description,
|
|
||||||
deadline_seconds=deadline_seconds,
|
|
||||||
allow_redelegation=allow_redelegation,
|
|
||||||
)
|
|
||||||
if result.get("status") == "rejected":
|
|
||||||
return f"Error: dispatch rejected — {result.get('reason', 'unknown')}"
|
|
||||||
if result.get("status") == "circuit_open":
|
|
||||||
failover = result.get("failover_to")
|
|
||||||
return f"Error: circuit open for {to}. Failover candidate: {failover or 'none'}"
|
|
||||||
return (
|
|
||||||
f"Dispatched to {to}. Task ID: {result.get('task_id')}. "
|
|
||||||
f"Depth: {result.get('depth', 0)}."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class PollTaskResultTool(Tool):
|
|
||||||
"""Poll the status of a previously dispatched task."""
|
|
||||||
|
|
||||||
def __init__(self, shell: "P2PShell"):
|
|
||||||
self._shell = shell
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "poll_task_result"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Check the current status of a task you previously dispatched. "
|
|
||||||
"Returns completed, pending, timeout, failed, or not_found. "
|
|
||||||
"Call this proactively — do not wait for automatic notifications."
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def parameters(self) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"task_id": {"type": "string", "description": "Task ID returned by dispatch_task"},
|
|
||||||
},
|
|
||||||
"required": ["task_id"],
|
|
||||||
}
|
|
||||||
|
|
||||||
async def execute(self, task_id: str, **kwargs: Any) -> str:
|
|
||||||
result = self._shell.poll(task_id)
|
|
||||||
status = result.get("status")
|
|
||||||
if status == "not_found":
|
|
||||||
return f"Task {task_id} not found."
|
|
||||||
if status == "pending":
|
|
||||||
return f"Task {task_id} is pending (elapsed {result.get('elapsed', '?')}s)."
|
|
||||||
if status == "timeout":
|
|
||||||
return f"Task {task_id} timed out after {result.get('elapsed', '?')}s."
|
|
||||||
if status in ("completed", "failed", "aborted"):
|
|
||||||
from_agent = result.get("from", "unknown")
|
|
||||||
content = result.get("result", "")
|
|
||||||
preview = content[:500] + "..." if len(content) > 500 else content
|
|
||||||
return f"Task {task_id} is {status} (from {from_agent}).\n\n{preview}"
|
|
||||||
return f"Task {task_id} status: {status}"
|
|
||||||
|
|
||||||
|
|
||||||
class BroadcastTaskTool(Tool):
|
|
||||||
"""Broadcast subtasks to discover capable agents."""
|
|
||||||
|
|
||||||
def __init__(self, shell: "P2PShell"):
|
|
||||||
self._shell = shell
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "broadcast_task"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Announce subtasks to the agent network to collect BIDs. "
|
|
||||||
"Returns immediately. Use check_aggregation later to see which agents responded. "
|
|
||||||
"Each subtask should include a capability hint for matching."
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def parameters(self) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"task_id": {"type": "string", "description": "Your task identifier"},
|
|
||||||
"subtasks": {
|
|
||||||
"type": "array",
|
|
||||||
"items": {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"subtask_id": {"type": "string"},
|
|
||||||
"description": {"type": "string"},
|
|
||||||
"capability": {"type": "string", "description": "Required capability, e.g. 'web_search'"},
|
|
||||||
"budget_seconds": {"type": "integer", "default": 300},
|
|
||||||
},
|
|
||||||
"required": ["subtask_id", "description", "capability"],
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"aggregation_timeout": {"type": "integer", "default": 30, "description": "Seconds to wait for BIDs"},
|
|
||||||
},
|
|
||||||
"required": ["task_id", "subtasks"],
|
|
||||||
}
|
|
||||||
|
|
||||||
async def execute(
|
|
||||||
self,
|
|
||||||
task_id: str,
|
|
||||||
subtasks: list[dict[str, Any]],
|
|
||||||
aggregation_timeout: int = 30,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> str:
|
|
||||||
result = self._shell.broadcast(task_id, subtasks, aggregation_timeout)
|
|
||||||
invited = result.get("invited", 0)
|
|
||||||
return f"Broadcast opened for {task_id}. Invited {invited} agent(s). Use check_aggregation to collect BIDs."
|
|
||||||
|
|
||||||
|
|
||||||
class CheckAggregationTool(Tool):
|
|
||||||
"""Check the status of a broadcast aggregation window."""
|
|
||||||
|
|
||||||
def __init__(self, shell: "P2PShell"):
|
|
||||||
self._shell = shell
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "check_aggregation"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Check whether a previously broadcast task has collected enough BIDs or timed out. "
|
|
||||||
"Returns the list of responding agents and their bids, or a pending status with counts."
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def parameters(self) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"task_id": {"type": "string", "description": "Task ID used in broadcast_task"},
|
|
||||||
},
|
|
||||||
"required": ["task_id"],
|
|
||||||
}
|
|
||||||
|
|
||||||
async def execute(self, task_id: str, **kwargs: Any) -> str:
|
|
||||||
result = self._shell.check_aggregation(task_id)
|
|
||||||
status = result.get("status")
|
|
||||||
if status == "no_window":
|
|
||||||
return f"No broadcast window found for {task_id}."
|
|
||||||
if status == "pending":
|
|
||||||
received = result.get("received", 0)
|
|
||||||
expected = result.get("expected", "?")
|
|
||||||
remaining = result.get("seconds_remaining", 0)
|
|
||||||
return (
|
|
||||||
f"Aggregation pending for {task_id}: "
|
|
||||||
f"{received}/{expected} received, {remaining}s remaining."
|
|
||||||
)
|
|
||||||
if status == "closed":
|
|
||||||
entries = result.get("entries", [])
|
|
||||||
lines = [f"Aggregation closed for {task_id} ({result.get('reason', '')}):", ""]
|
|
||||||
for e in entries:
|
|
||||||
agent = e.get("from", "unknown")
|
|
||||||
sub = e.get("subtask_id", "")
|
|
||||||
lines.append(f"- {agent} bid for {sub}")
|
|
||||||
return "\n".join(lines)
|
|
||||||
return f"Unknown aggregation status for {task_id}: {status}"
|
|
||||||
|
|
||||||
|
|
||||||
class ReportUserTool(Tool):
|
|
||||||
"""Deliver a final answer to the user."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
send_callback: Callable[[OutboundMessage], Awaitable[None]] | None = None,
|
|
||||||
default_channel: str = "",
|
|
||||||
default_chat_id: str = "",
|
|
||||||
):
|
|
||||||
self._send_callback = send_callback
|
|
||||||
self._default_channel = default_channel
|
|
||||||
self._default_chat_id = default_chat_id
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "report_user"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Report the final answer to the user. Use this when you have gathered enough results. "
|
|
||||||
"Status 'partial' means some subtasks are incomplete — list them in pending_items."
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def parameters(self) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"final_answer": {"type": "string", "description": "Complete answer for the user"},
|
|
||||||
"status": {"type": "string", "enum": ["success", "partial", "failed"]},
|
|
||||||
"pending_items": {
|
|
||||||
"type": "array",
|
|
||||||
"items": {"type": "string"},
|
|
||||||
"description": "Incomplete items when status is partial",
|
|
||||||
},
|
|
||||||
"task_summary": {"type": "string", "description": "Optional brief summary"},
|
|
||||||
},
|
|
||||||
"required": ["final_answer", "status"],
|
|
||||||
}
|
|
||||||
|
|
||||||
async def execute(
|
|
||||||
self,
|
|
||||||
final_answer: str,
|
|
||||||
status: str,
|
|
||||||
pending_items: list[str] | None = None,
|
|
||||||
task_summary: str = "",
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> str:
|
|
||||||
if not self._send_callback:
|
|
||||||
return "Error: report_user not configured (no send callback)"
|
|
||||||
|
|
||||||
parts = [final_answer]
|
|
||||||
if pending_items:
|
|
||||||
parts.append(f"\n\nPending items:\n" + "\n".join(f"- {i}" for i in pending_items))
|
|
||||||
if task_summary:
|
|
||||||
parts.append(f"\n\nSummary: {task_summary}")
|
|
||||||
|
|
||||||
content = "\n".join(parts)
|
|
||||||
msg = OutboundMessage(
|
|
||||||
channel=self._default_channel,
|
|
||||||
chat_id=self._default_chat_id,
|
|
||||||
content=content,
|
|
||||||
)
|
|
||||||
await self._send_callback(msg)
|
|
||||||
return f"Reported to user (status={status})."
|
|
||||||
|
|
||||||
|
|
||||||
class FinalizeTaskTool(Tool):
|
|
||||||
"""Force-finalize a task and close its sessions."""
|
|
||||||
|
|
||||||
def __init__(self, shell: "P2PShell", session_manager: "SessionManager | None" = None):
|
|
||||||
self._shell = shell
|
|
||||||
self._session_manager = session_manager
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "finalize_task"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Terminate a task and all its subtasks. Use when the user says 'stop', "
|
|
||||||
"or when a task is fundamentally blocked. outcome can be completed, failed, or aborted."
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def parameters(self) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"task_id": {"type": "string"},
|
|
||||||
"outcome": {"type": "string", "enum": ["completed", "failed", "aborted"]},
|
|
||||||
"reason": {"type": "string", "description": "Why the task was finalized"},
|
|
||||||
},
|
|
||||||
"required": ["task_id", "outcome"],
|
|
||||||
}
|
|
||||||
|
|
||||||
async def execute(
|
|
||||||
self,
|
|
||||||
task_id: str,
|
|
||||||
outcome: str,
|
|
||||||
reason: str = "",
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> str:
|
|
||||||
self._shell.finalize(task_id, outcome, reason)
|
|
||||||
if self._session_manager:
|
|
||||||
self._session_manager.finalize_task_session(task_id)
|
|
||||||
return f"Task {task_id} finalized with outcome={outcome}."
|
|
||||||
@@ -266,6 +266,7 @@ class ExecTool(Tool):
|
|||||||
# the raw command string to COMSPEC without re-quoting.
|
# the raw command string to COMSPEC without re-quoting.
|
||||||
return await asyncio.create_subprocess_shell(
|
return await asyncio.create_subprocess_shell(
|
||||||
command,
|
command,
|
||||||
|
stdin=asyncio.subprocess.DEVNULL,
|
||||||
stdout=asyncio.subprocess.PIPE,
|
stdout=asyncio.subprocess.PIPE,
|
||||||
stderr=asyncio.subprocess.PIPE,
|
stderr=asyncio.subprocess.PIPE,
|
||||||
cwd=cwd,
|
cwd=cwd,
|
||||||
@@ -274,6 +275,7 @@ class ExecTool(Tool):
|
|||||||
bash = shutil.which("bash") or "/bin/bash"
|
bash = shutil.which("bash") or "/bin/bash"
|
||||||
return await asyncio.create_subprocess_exec(
|
return await asyncio.create_subprocess_exec(
|
||||||
bash, "-l", "-c", command,
|
bash, "-l", "-c", command,
|
||||||
|
stdin=asyncio.subprocess.DEVNULL,
|
||||||
stdout=asyncio.subprocess.PIPE,
|
stdout=asyncio.subprocess.PIPE,
|
||||||
stderr=asyncio.subprocess.PIPE,
|
stderr=asyncio.subprocess.PIPE,
|
||||||
cwd=cwd,
|
cwd=cwd,
|
||||||
|
|||||||
@@ -70,28 +70,40 @@ class ChannelManager:
|
|||||||
|
|
||||||
def _init_channels(self) -> None:
|
def _init_channels(self) -> None:
|
||||||
"""Initialize channels discovered via pkgutil scan + entry_points plugins."""
|
"""Initialize channels discovered via pkgutil scan + entry_points plugins."""
|
||||||
from nanobot.channels.registry import discover_all
|
from nanobot.channels.registry import discover_channel_names, discover_enabled
|
||||||
|
|
||||||
transcription_provider = self.config.channels.transcription_provider
|
transcription_provider = self.config.channels.transcription_provider
|
||||||
transcription_key = self._resolve_transcription_key(transcription_provider)
|
transcription_key = self._resolve_transcription_key(transcription_provider)
|
||||||
transcription_base = self._resolve_transcription_base(transcription_provider)
|
transcription_base = self._resolve_transcription_base(transcription_provider)
|
||||||
transcription_language = self.config.channels.transcription_language
|
transcription_language = self.config.channels.transcription_language
|
||||||
|
|
||||||
for name, cls in discover_all().items():
|
# Collect enabled module names first, then only import those.
|
||||||
|
# Channel configs live in ChannelsConfig's extra fields (via
|
||||||
|
# extra="allow"), so we enumerate candidates from pkgutil scan
|
||||||
|
# (cheap, no imports) and any plugin keys in __pydantic_extra__.
|
||||||
|
names = discover_channel_names()
|
||||||
|
candidate_names = set(names)
|
||||||
|
extra = getattr(self.config.channels, "__pydantic_extra__", None) or {}
|
||||||
|
candidate_names.update(extra.keys())
|
||||||
|
|
||||||
|
enabled_names: set[str] = set()
|
||||||
|
for name in candidate_names:
|
||||||
section = getattr(self.config.channels, name, None)
|
section = getattr(self.config.channels, name, None)
|
||||||
if section is None:
|
if section is None:
|
||||||
continue
|
continue
|
||||||
enabled = (
|
if (
|
||||||
section.get("enabled", False)
|
section.get("enabled", False)
|
||||||
if isinstance(section, dict)
|
if isinstance(section, dict)
|
||||||
else getattr(section, "enabled", False)
|
else getattr(section, "enabled", False)
|
||||||
)
|
):
|
||||||
if not enabled:
|
enabled_names.add(name)
|
||||||
|
|
||||||
|
for name, cls in discover_enabled(enabled_names, _names=names).items():
|
||||||
|
section = getattr(self.config.channels, name, None)
|
||||||
|
if section is None:
|
||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
kwargs: dict[str, Any] = {}
|
kwargs: dict[str, Any] = {}
|
||||||
# Only the WebSocket channel currently hosts the embedded webui
|
|
||||||
# surface; other channels stay oblivious to these knobs.
|
|
||||||
if cls.name == "websocket":
|
if cls.name == "websocket":
|
||||||
if self._session_manager is not None:
|
if self._session_manager is not None:
|
||||||
kwargs["session_manager"] = self._session_manager
|
kwargs["session_manager"] = self._session_manager
|
||||||
|
|||||||
@@ -1,5 +1,4 @@
|
|||||||
"""Auto-discovery for built-in channel modules and external plugins."""
|
"""Auto-discovery for built-in channel modules and external plugins."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import importlib
|
import importlib
|
||||||
@@ -37,12 +36,14 @@ def load_channel_class(module_name: str) -> type[BaseChannel]:
|
|||||||
raise ImportError(f"No BaseChannel subclass in nanobot.channels.{module_name}")
|
raise ImportError(f"No BaseChannel subclass in nanobot.channels.{module_name}")
|
||||||
|
|
||||||
|
|
||||||
def discover_plugins() -> dict[str, type[BaseChannel]]:
|
def discover_plugins(enabled_names: set[str] | None = None) -> dict[str, type[BaseChannel]]:
|
||||||
"""Discover external channel plugins registered via entry_points."""
|
"""Discover external channel plugins registered via entry_points."""
|
||||||
from importlib.metadata import entry_points
|
from importlib.metadata import entry_points
|
||||||
|
|
||||||
plugins: dict[str, type[BaseChannel]] = {}
|
plugins: dict[str, type[BaseChannel]] = {}
|
||||||
for ep in entry_points(group="nanobot.channels"):
|
for ep in entry_points(group="nanobot.channels"):
|
||||||
|
if enabled_names is not None and ep.name not in enabled_names:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
cls = ep.load()
|
cls = ep.load()
|
||||||
plugins[ep.name] = cls
|
plugins[ep.name] = cls
|
||||||
@@ -51,21 +52,44 @@ def discover_plugins() -> dict[str, type[BaseChannel]]:
|
|||||||
return plugins
|
return plugins
|
||||||
|
|
||||||
|
|
||||||
|
def discover_enabled(
|
||||||
|
enabled_names: set[str],
|
||||||
|
*,
|
||||||
|
_names: list[str] | None = None,
|
||||||
|
_include_all_external: bool = False,
|
||||||
|
) -> dict[str, type[BaseChannel]]:
|
||||||
|
"""Return channels whose module names are in *enabled_names*.
|
||||||
|
|
||||||
|
Uses cheap ``pkgutil.iter_modules`` to list names, then imports only
|
||||||
|
those that match — skipping the heavy third-party SDK imports of
|
||||||
|
unneeded channels.
|
||||||
|
"""
|
||||||
|
names = _names if _names is not None else discover_channel_names()
|
||||||
|
result: dict[str, type[BaseChannel]] = {}
|
||||||
|
for modname in names:
|
||||||
|
if modname not in enabled_names:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
result[modname] = load_channel_class(modname)
|
||||||
|
except ImportError as e:
|
||||||
|
logger.debug("Skipping built-in channel '{}': {}", modname, e)
|
||||||
|
|
||||||
|
external = discover_plugins(None if _include_all_external else enabled_names)
|
||||||
|
shadowed = set(external) & set(result)
|
||||||
|
if shadowed:
|
||||||
|
logger.warning("Plugin(s) shadowed by built-in channels (ignored): {}", shadowed)
|
||||||
|
if _include_all_external:
|
||||||
|
result.update({k: v for k, v in external.items() if k not in shadowed})
|
||||||
|
else:
|
||||||
|
result.update({k: v for k, v in external.items() if k not in shadowed and k in enabled_names})
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
def discover_all() -> dict[str, type[BaseChannel]]:
|
def discover_all() -> dict[str, type[BaseChannel]]:
|
||||||
"""Return all channels: built-in (pkgutil) merged with external (entry_points).
|
"""Return all channels: built-in (pkgutil) merged with external (entry_points).
|
||||||
|
|
||||||
Built-in channels take priority — an external plugin cannot shadow a built-in name.
|
Built-in channels take priority — an external plugin cannot shadow a built-in name.
|
||||||
"""
|
"""
|
||||||
builtin: dict[str, type[BaseChannel]] = {}
|
names = discover_channel_names()
|
||||||
for modname in discover_channel_names():
|
return discover_enabled(set(names), _names=names, _include_all_external=True)
|
||||||
try:
|
|
||||||
builtin[modname] = load_channel_class(modname)
|
|
||||||
except ImportError as e:
|
|
||||||
logger.debug("Skipping built-in channel '{}': {}", modname, e)
|
|
||||||
|
|
||||||
external = discover_plugins()
|
|
||||||
shadowed = set(external) & set(builtin)
|
|
||||||
if shadowed:
|
|
||||||
logger.warning("Plugin(s) shadowed by built-in channels (ignored): {}", shadowed)
|
|
||||||
|
|
||||||
return {**external, **builtin}
|
|
||||||
|
|||||||
+128
-210
@@ -37,15 +37,27 @@ from nanobot.command.builtin import builtin_command_palette
|
|||||||
from nanobot.config.paths import get_media_dir
|
from nanobot.config.paths import get_media_dir
|
||||||
from nanobot.config.schema import Base
|
from nanobot.config.schema import Base
|
||||||
from nanobot.session.goal_state import goal_state_ws_blob
|
from nanobot.session.goal_state import goal_state_ws_blob
|
||||||
|
from nanobot.session.webui_turns import websocket_turn_wall_started_at
|
||||||
from nanobot.utils.helpers import safe_filename
|
from nanobot.utils.helpers import safe_filename
|
||||||
from nanobot.utils.media_decode import (
|
from nanobot.utils.media_decode import (
|
||||||
FileSizeExceeded,
|
FileSizeExceeded,
|
||||||
save_base64_data_url,
|
save_base64_data_url,
|
||||||
)
|
)
|
||||||
from nanobot.utils.subagent_channel_display import scrub_subagent_messages_for_channel
|
from nanobot.utils.subagent_channel_display import scrub_subagent_messages_for_channel
|
||||||
from nanobot.utils.webui_thread_disk import delete_webui_thread
|
from nanobot.webui.settings_api import (
|
||||||
from nanobot.utils.webui_transcript import append_transcript_object, build_webui_thread_response
|
WebUISettingsError,
|
||||||
from nanobot.utils.webui_turn_helpers import websocket_turn_wall_started_at
|
settings_payload,
|
||||||
|
update_agent_settings,
|
||||||
|
update_image_generation_settings,
|
||||||
|
update_provider_settings,
|
||||||
|
update_web_search_settings,
|
||||||
|
)
|
||||||
|
from nanobot.webui.sidebar_state import (
|
||||||
|
read_webui_sidebar_state,
|
||||||
|
write_webui_sidebar_state,
|
||||||
|
)
|
||||||
|
from nanobot.webui.thread_disk import delete_webui_thread
|
||||||
|
from nanobot.webui.transcript import append_transcript_object, build_webui_thread_response
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nanobot.session.manager import SessionManager
|
from nanobot.session.manager import SessionManager
|
||||||
@@ -222,28 +234,6 @@ def _query_first(query: dict[str, list[str]], key: str) -> str | None:
|
|||||||
return values[0] if values else None
|
return values[0] if values else None
|
||||||
|
|
||||||
|
|
||||||
def _mask_secret_hint(secret: str | None) -> str | None:
|
|
||||||
if not secret:
|
|
||||||
return None
|
|
||||||
if len(secret) <= 8:
|
|
||||||
return "••••"
|
|
||||||
return f"{secret[:4]}••••{secret[-4:]}"
|
|
||||||
|
|
||||||
|
|
||||||
_WEB_SEARCH_PROVIDER_OPTIONS: tuple[dict[str, str], ...] = (
|
|
||||||
{"name": "duckduckgo", "label": "DuckDuckGo", "credential": "none"},
|
|
||||||
{"name": "brave", "label": "Brave Search", "credential": "api_key"},
|
|
||||||
{"name": "tavily", "label": "Tavily", "credential": "api_key"},
|
|
||||||
{"name": "searxng", "label": "SearXNG", "credential": "base_url"},
|
|
||||||
{"name": "jina", "label": "Jina", "credential": "api_key"},
|
|
||||||
{"name": "kagi", "label": "Kagi", "credential": "api_key"},
|
|
||||||
{"name": "olostep", "label": "Olostep", "credential": "api_key"},
|
|
||||||
)
|
|
||||||
_WEB_SEARCH_PROVIDER_BY_NAME = {
|
|
||||||
provider["name"]: provider for provider in _WEB_SEARCH_PROVIDER_OPTIONS
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _parse_inbound_payload(raw: str) -> str | None:
|
def _parse_inbound_payload(raw: str) -> str | None:
|
||||||
"""Parse a client frame into text; return None for empty or unrecognized content."""
|
"""Parse a client frame into text; return None for empty or unrecognized content."""
|
||||||
text = raw.strip()
|
text = raw.strip()
|
||||||
@@ -482,6 +472,7 @@ class WebSocketChannel(BaseChannel):
|
|||||||
static_dist_path.resolve() if static_dist_path is not None else None
|
static_dist_path.resolve() if static_dist_path is not None else None
|
||||||
)
|
)
|
||||||
self._runtime_model_name = runtime_model_name
|
self._runtime_model_name = runtime_model_name
|
||||||
|
self._settings_restart_sections: set[str] = set()
|
||||||
# Process-local secret used to HMAC-sign media URLs. The signed URL is
|
# Process-local secret used to HMAC-sign media URLs. The signed URL is
|
||||||
# the capability — anyone who holds a valid URL can fetch that one
|
# the capability — anyone who holds a valid URL can fetch that one
|
||||||
# file, nothing else. The secret regenerates on restart so links
|
# file, nothing else. The secret regenerates on restart so links
|
||||||
@@ -644,6 +635,12 @@ class WebSocketChannel(BaseChannel):
|
|||||||
if got == "/api/commands":
|
if got == "/api/commands":
|
||||||
return self._handle_commands(request)
|
return self._handle_commands(request)
|
||||||
|
|
||||||
|
if got == "/api/webui/sidebar-state":
|
||||||
|
return self._handle_webui_sidebar_state(request)
|
||||||
|
|
||||||
|
if got == "/api/webui/sidebar-state/update":
|
||||||
|
return self._handle_webui_sidebar_state_update(request)
|
||||||
|
|
||||||
if got == "/api/settings/update":
|
if got == "/api/settings/update":
|
||||||
return self._handle_settings_update(request)
|
return self._handle_settings_update(request)
|
||||||
|
|
||||||
@@ -653,6 +650,9 @@ class WebSocketChannel(BaseChannel):
|
|||||||
if got == "/api/settings/web-search/update":
|
if got == "/api/settings/web-search/update":
|
||||||
return self._handle_settings_web_search_update(request)
|
return self._handle_settings_web_search_update(request)
|
||||||
|
|
||||||
|
if got == "/api/settings/image-generation/update":
|
||||||
|
return self._handle_settings_image_generation_update(request)
|
||||||
|
|
||||||
m = re.match(r"^/api/sessions/([^/]+)/messages$", got)
|
m = re.match(r"^/api/sessions/([^/]+)/messages$", got)
|
||||||
if m:
|
if m:
|
||||||
return self._handle_session_messages(request, m.group(1))
|
return self._handle_session_messages(request, m.group(1))
|
||||||
@@ -764,215 +764,115 @@ class WebSocketChannel(BaseChannel):
|
|||||||
sessions = self._session_manager.list_sessions()
|
sessions = self._session_manager.list_sessions()
|
||||||
# Sidebar/chat listing for WS-backed sessions only — CLI / Slack / etc.
|
# Sidebar/chat listing for WS-backed sessions only — CLI / Slack / etc.
|
||||||
# keys are not intended for resume over this HTTP surface.
|
# keys are not intended for resume over this HTTP surface.
|
||||||
cleaned = [
|
cleaned = []
|
||||||
{k: v for k, v in s.items() if k != "path"}
|
for s in sessions:
|
||||||
for s in sessions
|
key = s.get("key")
|
||||||
if isinstance(s.get("key"), str) and s["key"].startswith("websocket:")
|
if not (isinstance(key, str) and key.startswith("websocket:")):
|
||||||
]
|
|
||||||
return _http_json_response({"sessions": cleaned})
|
|
||||||
|
|
||||||
def _settings_payload(self, *, requires_restart: bool = False) -> dict[str, Any]:
|
|
||||||
from nanobot.config.loader import get_config_path, load_config
|
|
||||||
from nanobot.providers.registry import PROVIDERS, find_by_name
|
|
||||||
|
|
||||||
config = load_config()
|
|
||||||
defaults = config.agents.defaults
|
|
||||||
provider_name = config.get_provider_name(defaults.model) or defaults.provider
|
|
||||||
provider = config.get_provider(defaults.model)
|
|
||||||
selected_provider = provider_name
|
|
||||||
if defaults.provider != "auto":
|
|
||||||
spec = find_by_name(defaults.provider)
|
|
||||||
selected_provider = spec.name if spec else provider_name
|
|
||||||
providers = []
|
|
||||||
for spec in PROVIDERS:
|
|
||||||
provider_config = getattr(config.providers, spec.name, None)
|
|
||||||
if provider_config is None or spec.is_oauth or spec.is_local:
|
|
||||||
continue
|
continue
|
||||||
providers.append(
|
row = {k: v for k, v in s.items() if k != "path"}
|
||||||
{
|
chat_id = key.split(":", 1)[1]
|
||||||
"name": spec.name,
|
started_at = websocket_turn_wall_started_at(chat_id)
|
||||||
"label": spec.label,
|
if started_at is not None:
|
||||||
"configured": bool(provider_config.api_key),
|
row["run_started_at"] = started_at
|
||||||
"api_key_hint": _mask_secret_hint(provider_config.api_key),
|
cleaned.append(row)
|
||||||
"api_base": provider_config.api_base,
|
return _http_json_response({"sessions": cleaned})
|
||||||
"default_api_base": spec.default_api_base or None,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
search_config = config.tools.web.search
|
|
||||||
search_provider = (
|
|
||||||
search_config.provider
|
|
||||||
if search_config.provider in _WEB_SEARCH_PROVIDER_BY_NAME
|
|
||||||
else "duckduckgo"
|
|
||||||
)
|
|
||||||
return {
|
|
||||||
"agent": {
|
|
||||||
"model": defaults.model,
|
|
||||||
"provider": selected_provider,
|
|
||||||
"resolved_provider": provider_name,
|
|
||||||
"has_api_key": bool(provider and provider.api_key),
|
|
||||||
},
|
|
||||||
"providers": providers,
|
|
||||||
"web_search": {
|
|
||||||
"provider": search_provider,
|
|
||||||
"api_key_hint": _mask_secret_hint(search_config.api_key),
|
|
||||||
"base_url": search_config.base_url or None,
|
|
||||||
"providers": list(_WEB_SEARCH_PROVIDER_OPTIONS),
|
|
||||||
},
|
|
||||||
"runtime": {
|
|
||||||
"config_path": str(get_config_path().expanduser()),
|
|
||||||
},
|
|
||||||
"requires_restart": requires_restart,
|
|
||||||
}
|
|
||||||
|
|
||||||
def _handle_settings(self, request: WsRequest) -> Response:
|
def _handle_settings(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")
|
||||||
return _http_json_response(self._settings_payload())
|
return _http_json_response(self._with_settings_restart_state(settings_payload()))
|
||||||
|
|
||||||
|
def _with_settings_restart_state(
|
||||||
|
self,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
*,
|
||||||
|
section: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Keep restart-required state alive for this gateway process."""
|
||||||
|
if section and payload.get("requires_restart"):
|
||||||
|
self._settings_restart_sections.add(section)
|
||||||
|
if self._settings_restart_sections:
|
||||||
|
payload = dict(payload)
|
||||||
|
payload["requires_restart"] = True
|
||||||
|
payload["restart_required_sections"] = sorted(self._settings_restart_sections)
|
||||||
|
else:
|
||||||
|
payload = dict(payload)
|
||||||
|
payload["restart_required_sections"] = []
|
||||||
|
return payload
|
||||||
|
|
||||||
def _handle_commands(self, request: WsRequest) -> Response:
|
def _handle_commands(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")
|
||||||
return _http_json_response({"commands": builtin_command_palette()})
|
return _http_json_response({"commands": builtin_command_palette()})
|
||||||
|
|
||||||
|
def _handle_webui_sidebar_state(self, request: WsRequest) -> Response:
|
||||||
|
if not self._check_api_token(request):
|
||||||
|
return _http_error(401, "Unauthorized")
|
||||||
|
return _http_json_response(read_webui_sidebar_state())
|
||||||
|
|
||||||
|
def _handle_webui_sidebar_state_update(self, request: WsRequest) -> Response:
|
||||||
|
if not self._check_api_token(request):
|
||||||
|
return _http_error(401, "Unauthorized")
|
||||||
|
query = _parse_query(request.path)
|
||||||
|
raw_state = _query_first(query, "state")
|
||||||
|
if raw_state is None:
|
||||||
|
return _http_error(400, "missing state")
|
||||||
|
try:
|
||||||
|
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")
|
||||||
|
try:
|
||||||
|
state = write_webui_sidebar_state(decoded)
|
||||||
|
except ValueError as e:
|
||||||
|
return _http_error(400, str(e))
|
||||||
|
except OSError:
|
||||||
|
self.logger.exception("failed to write webui sidebar state")
|
||||||
|
return _http_error(500, "failed to write sidebar state")
|
||||||
|
return _http_json_response(state)
|
||||||
|
|
||||||
def _handle_settings_update(self, request: WsRequest) -> Response:
|
def _handle_settings_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")
|
||||||
from nanobot.config.loader import load_config, save_config
|
|
||||||
from nanobot.providers.registry import find_by_name
|
|
||||||
|
|
||||||
query = _parse_query(request.path)
|
query = _parse_query(request.path)
|
||||||
config = load_config()
|
try:
|
||||||
defaults = config.agents.defaults
|
payload = update_agent_settings(query)
|
||||||
changed = False
|
except WebUISettingsError as e:
|
||||||
|
return _http_error(e.status, e.message)
|
||||||
model = _query_first(query, "model")
|
return _http_json_response(
|
||||||
if model is not None:
|
self._with_settings_restart_state(payload, section="runtime")
|
||||||
model = model.strip()
|
)
|
||||||
if not model:
|
|
||||||
return _http_error(400, "model is required")
|
|
||||||
if defaults.model != model:
|
|
||||||
defaults.model = model
|
|
||||||
changed = True
|
|
||||||
|
|
||||||
provider = _query_first(query, "provider")
|
|
||||||
if provider is not None:
|
|
||||||
provider = provider.strip()
|
|
||||||
if not provider:
|
|
||||||
return _http_error(400, "provider is required")
|
|
||||||
if find_by_name(provider) is None:
|
|
||||||
return _http_error(400, "unknown provider")
|
|
||||||
provider_config = getattr(config.providers, provider, None)
|
|
||||||
if provider_config is None or not provider_config.api_key:
|
|
||||||
return _http_error(400, "provider is not configured")
|
|
||||||
if defaults.provider != provider:
|
|
||||||
defaults.provider = provider
|
|
||||||
changed = True
|
|
||||||
|
|
||||||
if changed:
|
|
||||||
save_config(config)
|
|
||||||
# LLM provider/model changes are hot-reloaded by AgentLoop before each
|
|
||||||
# new turn via the provider snapshot loader, so a restart is unnecessary.
|
|
||||||
return _http_json_response(self._settings_payload(requires_restart=False))
|
|
||||||
|
|
||||||
def _handle_settings_provider_update(self, request: WsRequest) -> Response:
|
def _handle_settings_provider_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")
|
||||||
from nanobot.config.loader import load_config, save_config
|
|
||||||
from nanobot.providers.registry import find_by_name
|
|
||||||
|
|
||||||
query = _parse_query(request.path)
|
query = _parse_query(request.path)
|
||||||
provider_name = (_query_first(query, "provider") or "").strip()
|
try:
|
||||||
if not provider_name:
|
payload = update_provider_settings(query)
|
||||||
return _http_error(400, "provider is required")
|
except WebUISettingsError as e:
|
||||||
spec = find_by_name(provider_name)
|
return _http_error(e.status, e.message)
|
||||||
if spec is None or spec.is_oauth or spec.is_local:
|
return _http_json_response(self._with_settings_restart_state(payload, section="image"))
|
||||||
return _http_error(400, "unknown provider")
|
|
||||||
|
|
||||||
config = load_config()
|
|
||||||
provider_config = getattr(config.providers, spec.name, None)
|
|
||||||
if provider_config is None:
|
|
||||||
return _http_error(400, "unknown provider")
|
|
||||||
|
|
||||||
changed = False
|
|
||||||
if "api_key" in query or "apiKey" in query:
|
|
||||||
api_key = _query_first(query, "api_key")
|
|
||||||
if api_key is None:
|
|
||||||
api_key = _query_first(query, "apiKey")
|
|
||||||
api_key = (api_key or "").strip() or None
|
|
||||||
if provider_config.api_key != api_key:
|
|
||||||
provider_config.api_key = api_key
|
|
||||||
changed = True
|
|
||||||
|
|
||||||
if "api_base" in query or "apiBase" in query:
|
|
||||||
api_base = _query_first(query, "api_base")
|
|
||||||
if api_base is None:
|
|
||||||
api_base = _query_first(query, "apiBase")
|
|
||||||
api_base = (api_base or "").strip() or None
|
|
||||||
if provider_config.api_base != api_base:
|
|
||||||
provider_config.api_base = api_base
|
|
||||||
changed = True
|
|
||||||
|
|
||||||
if changed:
|
|
||||||
save_config(config)
|
|
||||||
# API key/base changes are picked up by the next provider snapshot refresh.
|
|
||||||
return _http_json_response(self._settings_payload(requires_restart=False))
|
|
||||||
|
|
||||||
def _handle_settings_web_search_update(self, request: WsRequest) -> Response:
|
def _handle_settings_web_search_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")
|
||||||
from nanobot.config.loader import load_config, save_config
|
|
||||||
|
|
||||||
query = _parse_query(request.path)
|
query = _parse_query(request.path)
|
||||||
provider_name = (_query_first(query, "provider") or "").strip().lower()
|
try:
|
||||||
provider_option = _WEB_SEARCH_PROVIDER_BY_NAME.get(provider_name)
|
payload = update_web_search_settings(query)
|
||||||
if provider_option is None:
|
except WebUISettingsError as e:
|
||||||
return _http_error(400, "unknown web search provider")
|
return _http_error(e.status, e.message)
|
||||||
|
return _http_json_response(self._with_settings_restart_state(payload, section="web"))
|
||||||
|
|
||||||
config = load_config()
|
def _handle_settings_image_generation_update(self, request: WsRequest) -> Response:
|
||||||
search_config = config.tools.web.search
|
if not self._check_api_token(request):
|
||||||
previous_provider = search_config.provider
|
return _http_error(401, "Unauthorized")
|
||||||
changed = False
|
query = _parse_query(request.path)
|
||||||
|
try:
|
||||||
def set_value(attr: str, value: str | None) -> None:
|
payload = update_image_generation_settings(query)
|
||||||
nonlocal changed
|
except WebUISettingsError as e:
|
||||||
if getattr(search_config, attr) != value:
|
return _http_error(e.status, e.message)
|
||||||
setattr(search_config, attr, value)
|
return _http_json_response(self._with_settings_restart_state(payload, section="image"))
|
||||||
changed = True
|
|
||||||
|
|
||||||
if search_config.provider != provider_name:
|
|
||||||
search_config.provider = provider_name
|
|
||||||
changed = True
|
|
||||||
|
|
||||||
credential = provider_option["credential"]
|
|
||||||
if credential == "none":
|
|
||||||
set_value("api_key", "")
|
|
||||||
set_value("base_url", "")
|
|
||||||
elif credential == "base_url":
|
|
||||||
base_url = _query_first(query, "base_url")
|
|
||||||
if base_url is None:
|
|
||||||
base_url = _query_first(query, "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:
|
|
||||||
return _http_error(400, "base_url is required")
|
|
||||||
set_value("base_url", base_url)
|
|
||||||
set_value("api_key", "")
|
|
||||||
else:
|
|
||||||
api_key = _query_first(query, "api_key")
|
|
||||||
if api_key is None:
|
|
||||||
api_key = _query_first(query, "apiKey")
|
|
||||||
api_key = api_key.strip() if api_key is not None else None
|
|
||||||
if not api_key and previous_provider == provider_name and search_config.api_key:
|
|
||||||
api_key = search_config.api_key
|
|
||||||
if not api_key:
|
|
||||||
return _http_error(400, "api_key is required")
|
|
||||||
set_value("api_key", api_key)
|
|
||||||
set_value("base_url", "")
|
|
||||||
|
|
||||||
if changed:
|
|
||||||
save_config(config)
|
|
||||||
return _http_json_response(self._settings_payload(requires_restart=False))
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _is_websocket_channel_session_key(key: str) -> bool:
|
def _is_websocket_channel_session_key(key: str) -> bool:
|
||||||
@@ -1581,6 +1481,7 @@ class WebSocketChannel(BaseChannel):
|
|||||||
if not conns:
|
if not conns:
|
||||||
if (
|
if (
|
||||||
msg.metadata.get("_progress")
|
msg.metadata.get("_progress")
|
||||||
|
or msg.metadata.get("_file_edit_events")
|
||||||
or msg.metadata.get("_turn_end")
|
or msg.metadata.get("_turn_end")
|
||||||
or msg.metadata.get("_session_updated")
|
or msg.metadata.get("_session_updated")
|
||||||
or msg.metadata.get("_goal_status")
|
or msg.metadata.get("_goal_status")
|
||||||
@@ -1613,7 +1514,22 @@ class WebSocketChannel(BaseChannel):
|
|||||||
await self.send_turn_end(msg.chat_id, latency_ms=lat_i, goal_state=gs_blob)
|
await self.send_turn_end(msg.chat_id, latency_ms=lat_i, goal_state=gs_blob)
|
||||||
return
|
return
|
||||||
if msg.metadata.get("_session_updated"):
|
if msg.metadata.get("_session_updated"):
|
||||||
await self.send_session_updated(msg.chat_id)
|
scope = msg.metadata.get("_session_update_scope")
|
||||||
|
await self.send_session_updated(
|
||||||
|
msg.chat_id,
|
||||||
|
scope=scope if isinstance(scope, str) else None,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if msg.metadata.get("_file_edit_events"):
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"event": "file_edit",
|
||||||
|
"chat_id": msg.chat_id,
|
||||||
|
"edits": msg.metadata["_file_edit_events"],
|
||||||
|
}
|
||||||
|
self._try_append_webui_transcript(msg.chat_id, payload)
|
||||||
|
raw = json.dumps(payload, ensure_ascii=False)
|
||||||
|
for connection in conns:
|
||||||
|
await self._safe_send_to(connection, raw, label=" ")
|
||||||
return
|
return
|
||||||
text = msg.content
|
text = msg.content
|
||||||
payload: dict[str, Any] = {
|
payload: dict[str, Any] = {
|
||||||
@@ -1780,12 +1696,14 @@ class WebSocketChannel(BaseChannel):
|
|||||||
for connection in conns:
|
for connection in conns:
|
||||||
await self._safe_send_to(connection, raw, label=" goal_status ")
|
await self._safe_send_to(connection, raw, label=" goal_status ")
|
||||||
|
|
||||||
async def send_session_updated(self, chat_id: str) -> None:
|
async def send_session_updated(self, chat_id: str, *, scope: str | None = None) -> None:
|
||||||
"""Notify clients that session metadata changed outside the main turn."""
|
"""Notify clients that session metadata changed outside the main turn."""
|
||||||
conns = list(self._subs.get(chat_id, ()))
|
conns = list(self._subs.get(chat_id, ()))
|
||||||
if not conns:
|
if not conns:
|
||||||
return
|
return
|
||||||
body: dict[str, Any] = {"event": "session_updated", "chat_id": chat_id}
|
body: dict[str, Any] = {"event": "session_updated", "chat_id": chat_id}
|
||||||
|
if scope:
|
||||||
|
body["scope"] = scope
|
||||||
raw = json.dumps(body, ensure_ascii=False)
|
raw = json.dumps(body, ensure_ascii=False)
|
||||||
for connection in conns:
|
for connection in conns:
|
||||||
await self._safe_send_to(connection, raw, label=" session_updated ")
|
await self._safe_send_to(connection, raw, label=" session_updated ")
|
||||||
|
|||||||
+242
-39
@@ -1,12 +1,14 @@
|
|||||||
"""CLI commands for nanobot."""
|
"""CLI commands for nanobot."""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import json
|
||||||
import os
|
import os
|
||||||
import select
|
import select
|
||||||
import signal
|
import signal
|
||||||
import sys
|
import sys
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from contextlib import nullcontext, suppress
|
from contextlib import nullcontext, suppress
|
||||||
|
from inspect import signature
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -75,7 +77,6 @@ class SafeFileHistory(FileHistory):
|
|||||||
from nanobot.cli.stream import StreamRenderer, ThinkingSpinner
|
from nanobot.cli.stream import StreamRenderer, ThinkingSpinner
|
||||||
from nanobot.config.paths import get_workspace_path, is_default_workspace
|
from nanobot.config.paths import get_workspace_path, is_default_workspace
|
||||||
from nanobot.config.schema import Config
|
from nanobot.config.schema import Config
|
||||||
from nanobot.p2p.shell import P2PShell
|
|
||||||
from nanobot.utils.helpers import sync_workspace_templates
|
from nanobot.utils.helpers import sync_workspace_templates
|
||||||
from nanobot.utils.restart import (
|
from nanobot.utils.restart import (
|
||||||
consume_restart_notice_from_env,
|
consume_restart_notice_from_env,
|
||||||
@@ -92,17 +93,8 @@ app = typer.Typer(
|
|||||||
|
|
||||||
console = Console()
|
console = Console()
|
||||||
EXIT_COMMANDS = {"exit", "quit", "/exit", "/quit", ":q"}
|
EXIT_COMMANDS = {"exit", "quit", "/exit", "/quit", ":q"}
|
||||||
|
_REASONING_SENTENCE_ENDINGS = (".", "!", "?", "。", "!", "?")
|
||||||
|
_REASONING_FLUSH_CHARS = 60
|
||||||
def _resolve_p2p(config: Config) -> P2PShell | None:
|
|
||||||
"""Resolve P2P config and create the stateless P2P shell."""
|
|
||||||
mb_cfg = config.mailbox
|
|
||||||
if not mb_cfg.enabled:
|
|
||||||
return None
|
|
||||||
return P2PShell(
|
|
||||||
agent_id=mb_cfg.agent_id,
|
|
||||||
mailboxes_root=mb_cfg.mailboxes_root,
|
|
||||||
)
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# CLI input: prompt_toolkit for editing, paste, history, and display
|
# CLI input: prompt_toolkit for editing, paste, history, and display
|
||||||
@@ -254,6 +246,35 @@ def _print_cli_progress_line(text: str, thinking: ThinkingSpinner | None, render
|
|||||||
target.print(f" [dim]↳ {text}[/dim]")
|
target.print(f" [dim]↳ {text}[/dim]")
|
||||||
|
|
||||||
|
|
||||||
|
class _ReasoningBuffer:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._text = ""
|
||||||
|
|
||||||
|
def add(self, text: str) -> str | None:
|
||||||
|
if not text:
|
||||||
|
return None
|
||||||
|
self._text += text
|
||||||
|
if self._should_flush(text):
|
||||||
|
return self.flush()
|
||||||
|
return None
|
||||||
|
|
||||||
|
def flush(self) -> str | None:
|
||||||
|
text = self._text.strip()
|
||||||
|
self._text = ""
|
||||||
|
return text or None
|
||||||
|
|
||||||
|
def clear(self) -> None:
|
||||||
|
self._text = ""
|
||||||
|
|
||||||
|
def _should_flush(self, text: str) -> bool:
|
||||||
|
stripped = text.rstrip()
|
||||||
|
return (
|
||||||
|
"\n" in text
|
||||||
|
or stripped.endswith(_REASONING_SENTENCE_ENDINGS)
|
||||||
|
or len(self._text) >= _REASONING_FLUSH_CHARS
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _print_cli_reasoning(text: str, thinking: ThinkingSpinner | None, renderer: StreamRenderer | None = None) -> None:
|
def _print_cli_reasoning(text: str, thinking: ThinkingSpinner | None, renderer: StreamRenderer | None = None) -> None:
|
||||||
"""Print reasoning/thinking content in a distinct style."""
|
"""Print reasoning/thinking content in a distinct style."""
|
||||||
if not text.strip():
|
if not text.strip():
|
||||||
@@ -266,6 +287,16 @@ def _print_cli_reasoning(text: str, thinking: ThinkingSpinner | None, renderer:
|
|||||||
target.print(f"[dim italic]✻ {text}[/dim italic]")
|
target.print(f"[dim italic]✻ {text}[/dim italic]")
|
||||||
|
|
||||||
|
|
||||||
|
def _flush_cli_reasoning(
|
||||||
|
reasoning_buffer: _ReasoningBuffer,
|
||||||
|
thinking: ThinkingSpinner | None,
|
||||||
|
renderer: StreamRenderer | None = None,
|
||||||
|
) -> None:
|
||||||
|
text = reasoning_buffer.flush()
|
||||||
|
if text:
|
||||||
|
_print_cli_reasoning(text, thinking, renderer)
|
||||||
|
|
||||||
|
|
||||||
async def _print_interactive_progress_line(text: str, thinking: ThinkingSpinner | None, renderer: StreamRenderer | None = None) -> None:
|
async def _print_interactive_progress_line(text: str, thinking: ThinkingSpinner | None, renderer: StreamRenderer | None = None) -> None:
|
||||||
"""Print an interactive progress line, pausing the spinner if needed."""
|
"""Print an interactive progress line, pausing the spinner if needed."""
|
||||||
if not text.strip():
|
if not text.strip():
|
||||||
@@ -284,6 +315,7 @@ async def _maybe_print_interactive_progress(
|
|||||||
thinking: ThinkingSpinner | None,
|
thinking: ThinkingSpinner | None,
|
||||||
channels_config: Any,
|
channels_config: Any,
|
||||||
renderer: StreamRenderer | None = None,
|
renderer: StreamRenderer | None = None,
|
||||||
|
reasoning_buffer: _ReasoningBuffer | None = None,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
metadata = msg.metadata or {}
|
metadata = msg.metadata or {}
|
||||||
if metadata.get("_retry_wait"):
|
if metadata.get("_retry_wait"):
|
||||||
@@ -293,12 +325,24 @@ async def _maybe_print_interactive_progress(
|
|||||||
if not metadata.get("_progress"):
|
if not metadata.get("_progress"):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
reasoning_buffer = reasoning_buffer or _ReasoningBuffer()
|
||||||
|
|
||||||
|
if metadata.get("_reasoning_end"):
|
||||||
|
if channels_config and not channels_config.show_reasoning:
|
||||||
|
reasoning_buffer.clear()
|
||||||
|
else:
|
||||||
|
_flush_cli_reasoning(reasoning_buffer, thinking, renderer)
|
||||||
|
return True
|
||||||
|
|
||||||
is_tool_hint = metadata.get("_tool_hint", False)
|
is_tool_hint = metadata.get("_tool_hint", False)
|
||||||
is_reasoning = metadata.get("_reasoning", False) or metadata.get("_reasoning_delta", False)
|
is_reasoning = metadata.get("_reasoning", False) or metadata.get("_reasoning_delta", False)
|
||||||
if is_reasoning:
|
if is_reasoning:
|
||||||
if channels_config and not channels_config.show_reasoning:
|
if channels_config and not channels_config.show_reasoning:
|
||||||
|
reasoning_buffer.clear()
|
||||||
return True
|
return True
|
||||||
_print_cli_reasoning(msg.content, thinking, renderer)
|
text = reasoning_buffer.add(msg.content)
|
||||||
|
if text:
|
||||||
|
_print_cli_reasoning(text, thinking, renderer)
|
||||||
return True
|
return True
|
||||||
if channels_config and is_tool_hint and not channels_config.send_tool_hints:
|
if channels_config and is_tool_hint and not channels_config.send_tool_hints:
|
||||||
return True
|
return True
|
||||||
@@ -578,6 +622,7 @@ def serve(
|
|||||||
|
|
||||||
from nanobot.api.server import create_app
|
from nanobot.api.server import create_app
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||||
from nanobot.session.manager import SessionManager
|
from nanobot.session.manager import SessionManager
|
||||||
|
|
||||||
if verbose:
|
if verbose:
|
||||||
@@ -593,17 +638,11 @@ 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)
|
||||||
p2p_shell = _resolve_p2p(runtime_config)
|
|
||||||
|
|
||||||
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,
|
||||||
p2p_shell=p2p_shell,
|
image_generation_provider_configs=image_gen_provider_configs(runtime_config),
|
||||||
image_generation_provider_configs={
|
|
||||||
"openrouter": runtime_config.providers.openrouter,
|
|
||||||
"aihubmix": runtime_config.providers.aihubmix,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
console.print(f"[red]Error: {exc}[/red]")
|
console.print(f"[red]Error: {exc}[/red]")
|
||||||
@@ -683,6 +722,7 @@ def _run_gateway(
|
|||||||
from nanobot.cron.types import CronJob
|
from nanobot.cron.types import CronJob
|
||||||
from nanobot.heartbeat.service import HeartbeatService
|
from nanobot.heartbeat.service import HeartbeatService
|
||||||
from nanobot.providers.factory import build_provider_snapshot, load_provider_snapshot
|
from nanobot.providers.factory import build_provider_snapshot, load_provider_snapshot
|
||||||
|
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||||
from nanobot.session.manager import SessionManager
|
from nanobot.session.manager import SessionManager
|
||||||
|
|
||||||
port = port if port is not None else config.gateway.port
|
port = port if port is not None else config.gateway.port
|
||||||
@@ -705,8 +745,6 @@ def _run_gateway(
|
|||||||
cron_store_path = config.workspace_path / "cron" / "jobs.json"
|
cron_store_path = config.workspace_path / "cron" / "jobs.json"
|
||||||
cron = CronService(cron_store_path)
|
cron = CronService(cron_store_path)
|
||||||
|
|
||||||
p2p_shell = _resolve_p2p(config)
|
|
||||||
|
|
||||||
# Create agent with cron service
|
# Create agent with cron service
|
||||||
agent = AgentLoop.from_config(
|
agent = AgentLoop.from_config(
|
||||||
config, bus,
|
config, bus,
|
||||||
@@ -715,10 +753,7 @@ def _run_gateway(
|
|||||||
context_window_tokens=provider_snapshot.context_window_tokens,
|
context_window_tokens=provider_snapshot.context_window_tokens,
|
||||||
cron_service=cron,
|
cron_service=cron,
|
||||||
session_manager=session_manager,
|
session_manager=session_manager,
|
||||||
image_generation_provider_configs={
|
image_generation_provider_configs=image_gen_provider_configs(config),
|
||||||
"openrouter": config.providers.openrouter,
|
|
||||||
"aihubmix": config.providers.aihubmix,
|
|
||||||
},
|
|
||||||
provider_snapshot_loader=load_provider_snapshot,
|
provider_snapshot_loader=load_provider_snapshot,
|
||||||
runtime_model_publisher=lambda model, preset: publish_runtime_model_update(
|
runtime_model_publisher=lambda model, preset: publish_runtime_model_update(
|
||||||
bus,
|
bus,
|
||||||
@@ -726,7 +761,6 @@ def _run_gateway(
|
|||||||
preset,
|
preset,
|
||||||
),
|
),
|
||||||
provider_signature=provider_snapshot.signature,
|
provider_signature=provider_snapshot.signature,
|
||||||
p2p_shell=p2p_shell,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
from nanobot.agent.loop import UNIFIED_SESSION_KEY
|
from nanobot.agent.loop import UNIFIED_SESSION_KEY
|
||||||
@@ -932,15 +966,12 @@ def _run_gateway(
|
|||||||
hb_cfg = config.gateway.heartbeat
|
hb_cfg = config.gateway.heartbeat
|
||||||
heartbeat = HeartbeatService(
|
heartbeat = HeartbeatService(
|
||||||
workspace=config.workspace_path,
|
workspace=config.workspace_path,
|
||||||
provider=agent.provider,
|
llm_runtime=agent.llm_runtime,
|
||||||
model=agent.model,
|
|
||||||
on_execute=on_heartbeat_execute,
|
on_execute=on_heartbeat_execute,
|
||||||
on_notify=on_heartbeat_notify,
|
on_notify=on_heartbeat_notify,
|
||||||
interval_s=hb_cfg.interval_s,
|
interval_s=hb_cfg.interval_s,
|
||||||
enabled=hb_cfg.enabled,
|
enabled=hb_cfg.enabled,
|
||||||
timezone=config.agents.defaults.timezone,
|
timezone=config.agents.defaults.timezone,
|
||||||
p2p_shell=p2p_shell,
|
|
||||||
bus=bus,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if channels.enabled_channels:
|
if channels.enabled_channels:
|
||||||
@@ -1089,6 +1120,7 @@ def agent(
|
|||||||
|
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.cron.service import CronService
|
from nanobot.cron.service import CronService
|
||||||
|
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||||
|
|
||||||
config = _load_runtime_config(config, workspace)
|
config = _load_runtime_config(config, workspace)
|
||||||
sync_workspace_templates(config.workspace_path)
|
sync_workspace_templates(config.workspace_path)
|
||||||
@@ -1103,8 +1135,6 @@ def agent(
|
|||||||
cron_store_path = config.workspace_path / "cron" / "jobs.json"
|
cron_store_path = config.workspace_path / "cron" / "jobs.json"
|
||||||
cron = CronService(cron_store_path)
|
cron = CronService(cron_store_path)
|
||||||
|
|
||||||
p2p_shell = _resolve_p2p(config)
|
|
||||||
|
|
||||||
if logs:
|
if logs:
|
||||||
logger.enable("nanobot")
|
logger.enable("nanobot")
|
||||||
else:
|
else:
|
||||||
@@ -1114,7 +1144,7 @@ def agent(
|
|||||||
agent_loop = AgentLoop.from_config(
|
agent_loop = AgentLoop.from_config(
|
||||||
config, bus,
|
config, bus,
|
||||||
cron_service=cron,
|
cron_service=cron,
|
||||||
p2p_shell=p2p_shell,
|
image_generation_provider_configs=image_gen_provider_configs(config),
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
console.print(f"[red]Error: {exc}[/red]")
|
console.print(f"[red]Error: {exc}[/red]")
|
||||||
@@ -1130,12 +1160,25 @@ def agent(
|
|||||||
_thinking: ThinkingSpinner | None = None
|
_thinking: ThinkingSpinner | None = None
|
||||||
|
|
||||||
def _make_progress(renderer: StreamRenderer | None = None):
|
def _make_progress(renderer: StreamRenderer | None = None):
|
||||||
|
reasoning_buffer = _ReasoningBuffer()
|
||||||
|
|
||||||
async def _cli_progress(content: str, *, tool_hint: bool = False, reasoning: bool = False, **_kwargs: Any) -> None:
|
async def _cli_progress(content: str, *, tool_hint: bool = False, reasoning: bool = False, **_kwargs: Any) -> None:
|
||||||
ch = agent_loop.channels_config
|
ch = agent_loop.channels_config
|
||||||
|
|
||||||
|
if _kwargs.get("reasoning_end"):
|
||||||
|
if ch and not ch.show_reasoning:
|
||||||
|
reasoning_buffer.clear()
|
||||||
|
else:
|
||||||
|
_flush_cli_reasoning(reasoning_buffer, _thinking, renderer)
|
||||||
|
return
|
||||||
|
|
||||||
if reasoning:
|
if reasoning:
|
||||||
if ch and not ch.show_reasoning:
|
if ch and not ch.show_reasoning:
|
||||||
|
reasoning_buffer.clear()
|
||||||
return
|
return
|
||||||
_print_cli_reasoning(content, _thinking, renderer)
|
text = reasoning_buffer.add(content)
|
||||||
|
if text:
|
||||||
|
_print_cli_reasoning(text, _thinking, renderer)
|
||||||
return
|
return
|
||||||
if ch and tool_hint and not ch.send_tool_hints:
|
if ch and tool_hint and not ch.send_tool_hints:
|
||||||
return
|
return
|
||||||
@@ -1206,6 +1249,7 @@ def agent(
|
|||||||
turn_done.set()
|
turn_done.set()
|
||||||
turn_response: list[tuple[str, dict]] = []
|
turn_response: list[tuple[str, dict]] = []
|
||||||
renderer: StreamRenderer | None = None
|
renderer: StreamRenderer | None = None
|
||||||
|
reasoning_buffer = _ReasoningBuffer()
|
||||||
|
|
||||||
async def _consume_outbound():
|
async def _consume_outbound():
|
||||||
while True:
|
while True:
|
||||||
@@ -1231,6 +1275,7 @@ def agent(
|
|||||||
renderer,
|
renderer,
|
||||||
agent_loop.channels_config,
|
agent_loop.channels_config,
|
||||||
renderer,
|
renderer,
|
||||||
|
reasoning_buffer,
|
||||||
):
|
):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -1271,6 +1316,7 @@ def agent(
|
|||||||
|
|
||||||
turn_done.clear()
|
turn_done.clear()
|
||||||
turn_response.clear()
|
turn_response.clear()
|
||||||
|
reasoning_buffer.clear()
|
||||||
renderer = StreamRenderer(
|
renderer = StreamRenderer(
|
||||||
render_markdown=markdown,
|
render_markdown=markdown,
|
||||||
bot_name=config.agents.defaults.bot_name,
|
bot_name=config.agents.defaults.bot_name,
|
||||||
@@ -1312,7 +1358,6 @@ def agent(
|
|||||||
console.print("\nGoodbye!")
|
console.print("\nGoodbye!")
|
||||||
break
|
break
|
||||||
finally:
|
finally:
|
||||||
pass
|
|
||||||
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)
|
||||||
@@ -1484,6 +1529,106 @@ def status():
|
|||||||
console.print(f"{spec.label}: {'[green]✓[/green]' if has_key else '[dim]not set[/dim]'}")
|
console.print(f"{spec.label}: {'[green]✓[/green]' if has_key else '[dim]not set[/dim]'}")
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# Config Commands
|
||||||
|
# ============================================================================
|
||||||
|
|
||||||
|
config_app = typer.Typer(help="Manage configuration")
|
||||||
|
app.add_typer(config_app, name="config")
|
||||||
|
|
||||||
|
|
||||||
|
@config_app.command("set")
|
||||||
|
def config_set(
|
||||||
|
path: str = typer.Argument(..., help="Dot path, e.g. agents.defaults.model"),
|
||||||
|
value: str = typer.Argument(..., help="Value. Use null/true/false or JSON for structured values."),
|
||||||
|
config_path: str | None = typer.Option(None, "--config", "-c", help="Path to config file"),
|
||||||
|
):
|
||||||
|
"""Set one config value by dot path."""
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
from nanobot.config.loader import get_config_path, load_config, save_config, set_config_path
|
||||||
|
from nanobot.config.schema import Config
|
||||||
|
|
||||||
|
resolved_path = Path(config_path).expanduser().resolve() if config_path else get_config_path()
|
||||||
|
if config_path:
|
||||||
|
set_config_path(resolved_path)
|
||||||
|
|
||||||
|
config = load_config(resolved_path)
|
||||||
|
parsed = _parse_config_cli_value(value)
|
||||||
|
try:
|
||||||
|
_set_config_cli_value(config, path, parsed)
|
||||||
|
validated = Config.model_validate(config.model_dump(mode="json", by_alias=True))
|
||||||
|
except (AttributeError, KeyError, TypeError, ValueError, ValidationError) as exc:
|
||||||
|
console.print(f"[red]Could not set config value:[/red] {exc}")
|
||||||
|
raise typer.Exit(1)
|
||||||
|
|
||||||
|
save_config(validated, resolved_path)
|
||||||
|
console.print(f"[green]✓[/green] Set [cyan]{path}[/cyan] = [bold]{value}[/bold]")
|
||||||
|
console.print(f"[dim]Config: {resolved_path}[/dim]")
|
||||||
|
if path in {"agents.defaults.provider", "agents.defaults.model"} and validated.agents.defaults.model_preset:
|
||||||
|
console.print(
|
||||||
|
"[yellow]! agents.defaults.model_preset is set and may override this. "
|
||||||
|
"Clear it with: nanobot config set agents.defaults.model_preset null[/yellow]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_config_cli_value(raw: str) -> Any:
|
||||||
|
lowered = raw.strip().lower()
|
||||||
|
if lowered == "null":
|
||||||
|
return None
|
||||||
|
if lowered == "true":
|
||||||
|
return True
|
||||||
|
if lowered == "false":
|
||||||
|
return False
|
||||||
|
with suppress(Exception):
|
||||||
|
return json.loads(raw)
|
||||||
|
return raw
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_config_field(obj: Any, key: str) -> str:
|
||||||
|
from pydantic import BaseModel
|
||||||
|
from pydantic.alias_generators import to_camel, to_snake
|
||||||
|
|
||||||
|
if not isinstance(obj, BaseModel):
|
||||||
|
return key
|
||||||
|
fields = type(obj).model_fields
|
||||||
|
if key in fields:
|
||||||
|
return key
|
||||||
|
normalized = to_snake(key.replace("-", "_"))
|
||||||
|
if normalized in fields:
|
||||||
|
return normalized
|
||||||
|
for name, field in fields.items():
|
||||||
|
aliases = {
|
||||||
|
to_camel(name),
|
||||||
|
str(field.alias) if field.alias else "",
|
||||||
|
str(field.serialization_alias) if field.serialization_alias else "",
|
||||||
|
}
|
||||||
|
if key in aliases:
|
||||||
|
return name
|
||||||
|
raise AttributeError(f"Unknown config path segment {key!r}")
|
||||||
|
|
||||||
|
|
||||||
|
def _set_config_cli_value(config: Any, path: str, value: Any) -> None:
|
||||||
|
parts = [part for part in path.split(".") if part]
|
||||||
|
if not parts:
|
||||||
|
raise ValueError("Config path cannot be empty.")
|
||||||
|
|
||||||
|
current = config
|
||||||
|
for raw_part in parts[:-1]:
|
||||||
|
if isinstance(current, dict):
|
||||||
|
current = current.setdefault(raw_part, {})
|
||||||
|
continue
|
||||||
|
part = _resolve_config_field(current, raw_part)
|
||||||
|
current = getattr(current, part)
|
||||||
|
|
||||||
|
leaf = parts[-1]
|
||||||
|
if isinstance(current, dict):
|
||||||
|
current[leaf] = value
|
||||||
|
return
|
||||||
|
leaf = _resolve_config_field(current, leaf)
|
||||||
|
setattr(current, leaf, value)
|
||||||
|
|
||||||
|
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
# OAuth Login
|
# OAuth Login
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@@ -1498,6 +1643,7 @@ _LOGOUT_HANDLERS: dict[str, Callable[[], None]] = {}
|
|||||||
_PROVIDER_DISPLAY: dict[str, str] = {
|
_PROVIDER_DISPLAY: dict[str, str] = {
|
||||||
"openai_codex": "OpenAI Codex",
|
"openai_codex": "OpenAI Codex",
|
||||||
"github_copilot": "GitHub Copilot",
|
"github_copilot": "GitHub Copilot",
|
||||||
|
"xai_oauth": "xAI Grok OAuth",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -1533,7 +1679,9 @@ def _resolve_oauth_provider(provider: str):
|
|||||||
|
|
||||||
@provider_app.command("login")
|
@provider_app.command("login")
|
||||||
def provider_login(
|
def provider_login(
|
||||||
provider: str = typer.Argument(..., help="OAuth provider (e.g. 'openai-codex', 'github-copilot')"),
|
provider: str = typer.Argument(..., help="OAuth provider (e.g. 'openai-codex', 'github-copilot', 'xai-oauth')"),
|
||||||
|
no_browser: bool = typer.Option(False, "--no-browser", help="Print the auth URL instead of opening a browser when supported."),
|
||||||
|
manual_paste: bool = typer.Option(False, "--manual-paste", help="Prompt for a callback URL or fallback code when supported."),
|
||||||
):
|
):
|
||||||
"""Authenticate with an OAuth provider."""
|
"""Authenticate with an OAuth provider."""
|
||||||
spec = _resolve_oauth_provider(provider)
|
spec = _resolve_oauth_provider(provider)
|
||||||
@@ -1544,12 +1692,18 @@ def provider_login(
|
|||||||
raise typer.Exit(1)
|
raise typer.Exit(1)
|
||||||
|
|
||||||
console.print(f"{__logo__} OAuth Login - {spec.label}\n")
|
console.print(f"{__logo__} OAuth Login - {spec.label}\n")
|
||||||
handler()
|
params = signature(handler).parameters
|
||||||
|
kwargs: dict[str, bool] = {}
|
||||||
|
if "no_browser" in params:
|
||||||
|
kwargs["no_browser"] = no_browser
|
||||||
|
if "manual_paste" in params:
|
||||||
|
kwargs["manual_paste"] = manual_paste
|
||||||
|
handler(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
@provider_app.command("logout")
|
@provider_app.command("logout")
|
||||||
def provider_logout(
|
def provider_logout(
|
||||||
provider: str = typer.Argument(..., help="OAuth provider (e.g. 'openai-codex', 'github-copilot')"),
|
provider: str = typer.Argument(..., help="OAuth provider (e.g. 'openai-codex', 'github-copilot', 'xai-oauth')"),
|
||||||
):
|
):
|
||||||
"""Log out from an OAuth provider."""
|
"""Log out from an OAuth provider."""
|
||||||
spec = _resolve_oauth_provider(provider)
|
spec = _resolve_oauth_provider(provider)
|
||||||
@@ -1613,6 +1767,24 @@ def _logout_github_copilot() -> None:
|
|||||||
_delete_oauth_files(storage.get_token_path(), _PROVIDER_DISPLAY["github_copilot"])
|
_delete_oauth_files(storage.get_token_path(), _PROVIDER_DISPLAY["github_copilot"])
|
||||||
|
|
||||||
|
|
||||||
|
@_register_logout("xai_oauth")
|
||||||
|
def _logout_xai_oauth() -> None:
|
||||||
|
"""Clear local OAuth credentials for xAI Grok OAuth."""
|
||||||
|
try:
|
||||||
|
from nanobot.providers.xai_oauth_provider import delete_xai_oauth_credentials
|
||||||
|
except ImportError:
|
||||||
|
console.print("[red]xAI Grok OAuth provider unavailable.[/red]")
|
||||||
|
raise typer.Exit(1)
|
||||||
|
|
||||||
|
removed_paths = delete_xai_oauth_credentials()
|
||||||
|
if not removed_paths:
|
||||||
|
console.print(f"[yellow]! No local OAuth credentials found for {_PROVIDER_DISPLAY['xai_oauth']}[/yellow]")
|
||||||
|
return
|
||||||
|
console.print(f"[green]✓ Logged out from {_PROVIDER_DISPLAY['xai_oauth']}[/green]")
|
||||||
|
for path in removed_paths:
|
||||||
|
console.print(f"[dim]Removed: {path}[/dim]")
|
||||||
|
|
||||||
|
|
||||||
def _delete_oauth_files(token_path: Path, provider_label: str) -> None:
|
def _delete_oauth_files(token_path: Path, provider_label: str) -> None:
|
||||||
"""Delete OAuth token and lock files, reporting the result."""
|
"""Delete OAuth token and lock files, reporting the result."""
|
||||||
removed_paths: list[Path] = []
|
removed_paths: list[Path] = []
|
||||||
@@ -1656,5 +1828,36 @@ def _login_github_copilot() -> None:
|
|||||||
raise typer.Exit(1)
|
raise typer.Exit(1)
|
||||||
|
|
||||||
|
|
||||||
|
@_register_login("xai_oauth")
|
||||||
|
def _login_xai_oauth(
|
||||||
|
*,
|
||||||
|
no_browser: bool = False,
|
||||||
|
manual_paste: bool = False,
|
||||||
|
) -> None:
|
||||||
|
try:
|
||||||
|
from nanobot.providers.xai_oauth_provider import login_xai_oauth_interactive
|
||||||
|
from nanobot.providers.xai_oauth_provider import DEFAULT_XAI_MODEL
|
||||||
|
|
||||||
|
console.print("[cyan]Starting xAI Grok OAuth login...[/cyan]\n")
|
||||||
|
credential = login_xai_oauth_interactive(
|
||||||
|
print_fn=lambda s: console.print(s),
|
||||||
|
prompt_fn=lambda s: typer.prompt(s),
|
||||||
|
open_browser=not no_browser,
|
||||||
|
manual_paste=manual_paste,
|
||||||
|
)
|
||||||
|
account = credential.account_id or "xAI"
|
||||||
|
storage = "OS keychain" if credential.storage == "keyring" else "private file"
|
||||||
|
console.print(f"[green]✓ Authenticated with xAI Grok OAuth[/green] [dim]{account} · {storage}[/dim]")
|
||||||
|
console.print("[dim]To use it for chat:[/dim]")
|
||||||
|
console.print("[dim] nanobot config set agents.defaults.model_preset null[/dim]")
|
||||||
|
console.print("[dim] nanobot config set agents.defaults.provider xai-oauth[/dim]")
|
||||||
|
console.print(f"[dim] nanobot config set agents.defaults.model {DEFAULT_XAI_MODEL}[/dim]")
|
||||||
|
console.print("[dim]Hosted X Search is enabled by default for xAI OAuth.[/dim]")
|
||||||
|
console.print("[dim]To disable it: nanobot config set providers.xai_oauth.x_search.enable false[/dim]")
|
||||||
|
except Exception as e:
|
||||||
|
console.print(f"[red]Authentication error: {e}[/red]")
|
||||||
|
raise typer.Exit(1)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
app()
|
app()
|
||||||
|
|||||||
+217
-1
@@ -22,7 +22,7 @@ from nanobot.cli.models import (
|
|||||||
get_model_suggestions,
|
get_model_suggestions,
|
||||||
)
|
)
|
||||||
from nanobot.config.loader import get_config_path, load_config
|
from nanobot.config.loader import get_config_path, load_config
|
||||||
from nanobot.config.schema import Config
|
from nanobot.config.schema import Config, ModelPresetConfig
|
||||||
|
|
||||||
console = Console()
|
console = Console()
|
||||||
|
|
||||||
@@ -49,6 +49,10 @@ _SELECT_FIELD_HINTS: dict[str, tuple[list[str], str]] = {
|
|||||||
|
|
||||||
_BACK_PRESSED = object() # Sentinel value for back navigation
|
_BACK_PRESSED = object() # Sentinel value for back navigation
|
||||||
|
|
||||||
|
# Cache of model-preset names populated at runtime so that field handlers can
|
||||||
|
# offer existing presets as choices (e.g. AgentDefaults.model_preset).
|
||||||
|
_MODEL_PRESET_CACHE: set[str] = set()
|
||||||
|
|
||||||
|
|
||||||
def _get_questionary():
|
def _get_questionary():
|
||||||
"""Return questionary or raise a clear error when wizard deps are unavailable."""
|
"""Return questionary or raise a clear error when wizard deps are unavailable."""
|
||||||
@@ -588,9 +592,102 @@ def _handle_context_window_field(
|
|||||||
setattr(working_model, field_name, new_value)
|
setattr(working_model, field_name, new_value)
|
||||||
|
|
||||||
|
|
||||||
|
def _handle_model_preset_field(
|
||||||
|
working_model: BaseModel, field_name: str, field_display: str, current_value: Any
|
||||||
|
) -> None:
|
||||||
|
"""Handle the 'model_preset' field with a list of existing presets."""
|
||||||
|
preset_names = sorted(_MODEL_PRESET_CACHE)
|
||||||
|
choices = ["(clear/unset)"] + preset_names
|
||||||
|
default_choice = str(current_value) if current_value else "(clear/unset)"
|
||||||
|
new_value = _select_with_back(field_display, choices, default=default_choice)
|
||||||
|
if new_value is _BACK_PRESSED:
|
||||||
|
return
|
||||||
|
if new_value == "(clear/unset)":
|
||||||
|
setattr(working_model, field_name, None)
|
||||||
|
elif new_value is not None:
|
||||||
|
setattr(working_model, field_name, new_value)
|
||||||
|
|
||||||
|
|
||||||
|
def _handle_provider_field(
|
||||||
|
working_model: BaseModel, field_name: str, field_display: str, current_value: Any
|
||||||
|
) -> None:
|
||||||
|
"""Handle the 'provider' field with a list of registered providers."""
|
||||||
|
provider_names = sorted(_get_provider_names().keys())
|
||||||
|
choices = ["auto"] + provider_names
|
||||||
|
default_choice = str(current_value) if current_value else "auto"
|
||||||
|
new_value = _select_with_back(field_display, choices, default=default_choice)
|
||||||
|
if new_value is _BACK_PRESSED:
|
||||||
|
return
|
||||||
|
if new_value is not None:
|
||||||
|
setattr(working_model, field_name, new_value)
|
||||||
|
|
||||||
|
|
||||||
|
def _handle_fallback_models_field(
|
||||||
|
working_model: BaseModel, field_name: str, field_display: str, current_value: Any
|
||||||
|
) -> None:
|
||||||
|
"""Handle the 'fallback_models' field with preset-aware list management."""
|
||||||
|
from nanobot.config.schema import InlineFallbackConfig
|
||||||
|
|
||||||
|
items: list[Any] = list(current_value) if isinstance(current_value, list) else []
|
||||||
|
preset_names = sorted(_MODEL_PRESET_CACHE)
|
||||||
|
|
||||||
|
while True:
|
||||||
|
console.clear()
|
||||||
|
console.print(f"[bold]{field_display}[/bold]")
|
||||||
|
if items:
|
||||||
|
for idx, item in enumerate(items, 1):
|
||||||
|
if isinstance(item, InlineFallbackConfig):
|
||||||
|
console.print(f" {idx}. {item.model} ({item.provider}) [inline]")
|
||||||
|
else:
|
||||||
|
console.print(f" {idx}. {item}")
|
||||||
|
else:
|
||||||
|
console.print(" [dim](empty)[/dim]")
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
choices = ["[+] Add preset"]
|
||||||
|
if items:
|
||||||
|
choices.append("[-] Remove last")
|
||||||
|
choices.append("[X] Clear all")
|
||||||
|
choices.append("[Done]")
|
||||||
|
choices.append("<- Back")
|
||||||
|
|
||||||
|
answer = _get_questionary().select(
|
||||||
|
"Manage fallback models:",
|
||||||
|
choices=choices,
|
||||||
|
qmark=">",
|
||||||
|
).ask()
|
||||||
|
|
||||||
|
if answer is None or answer == "<- Back":
|
||||||
|
return
|
||||||
|
if answer == "[Done]":
|
||||||
|
setattr(working_model, field_name, items)
|
||||||
|
return
|
||||||
|
if answer == "[+] Add preset":
|
||||||
|
if not preset_names:
|
||||||
|
console.print("[yellow]! No presets defined yet.[/yellow]")
|
||||||
|
_get_questionary().press_any_key_to_continue().ask()
|
||||||
|
continue
|
||||||
|
add_choices = [p for p in preset_names if p not in items]
|
||||||
|
if not add_choices:
|
||||||
|
console.print("[yellow]! All presets already added.[/yellow]")
|
||||||
|
_get_questionary().press_any_key_to_continue().ask()
|
||||||
|
continue
|
||||||
|
picked = _select_with_back("Select preset:", add_choices)
|
||||||
|
if picked is _BACK_PRESSED or picked is None:
|
||||||
|
continue
|
||||||
|
items.append(picked)
|
||||||
|
elif answer == "[-] Remove last" and items:
|
||||||
|
items.pop()
|
||||||
|
elif answer == "[X] Clear all" and items:
|
||||||
|
items.clear()
|
||||||
|
|
||||||
|
|
||||||
_FIELD_HANDLERS: dict[str, Any] = {
|
_FIELD_HANDLERS: dict[str, Any] = {
|
||||||
"model": _handle_model_field,
|
"model": _handle_model_field,
|
||||||
"context_window_tokens": _handle_context_window_field,
|
"context_window_tokens": _handle_context_window_field,
|
||||||
|
"model_preset": _handle_model_preset_field,
|
||||||
|
"provider": _handle_provider_field,
|
||||||
|
"fallback_models": _handle_fallback_models_field,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -757,6 +854,116 @@ def _try_auto_fill_context_window(model: BaseModel, new_model_name: str) -> None
|
|||||||
console.print("[dim](i) Could not auto-fill context window (model not in database)[/dim]")
|
console.print("[dim](i) Could not auto-fill context window (model not in database)[/dim]")
|
||||||
|
|
||||||
|
|
||||||
|
# --- Model Preset Configuration ---
|
||||||
|
|
||||||
|
|
||||||
|
def _sync_preset_cache(config: Config) -> None:
|
||||||
|
"""Synchronise the module-level preset name cache from config."""
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
_MODEL_PRESET_CACHE.update(config.model_presets.keys())
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_model_presets(config: Config) -> None:
|
||||||
|
"""Configure model presets (CRUD)."""
|
||||||
|
_sync_preset_cache(config)
|
||||||
|
|
||||||
|
def get_preset_choices() -> list[str]:
|
||||||
|
choices: list[str] = []
|
||||||
|
for name, preset in config.model_presets.items():
|
||||||
|
choices.append(f"{name} ({preset.model})")
|
||||||
|
choices.append("[+] Add new preset")
|
||||||
|
choices.append("<- Back")
|
||||||
|
return choices
|
||||||
|
|
||||||
|
last_preset_name: str | None = None
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
console.clear()
|
||||||
|
_show_section_header(
|
||||||
|
"Model Presets",
|
||||||
|
"Create, edit or delete named model presets for quick switching",
|
||||||
|
)
|
||||||
|
choices = get_preset_choices()
|
||||||
|
default_choice = None
|
||||||
|
if last_preset_name:
|
||||||
|
for c in choices:
|
||||||
|
if c.startswith(last_preset_name + " ("):
|
||||||
|
default_choice = c
|
||||||
|
break
|
||||||
|
answer = _select_with_back(
|
||||||
|
"Select preset:", choices, default=default_choice
|
||||||
|
)
|
||||||
|
|
||||||
|
if answer is _BACK_PRESSED or answer is None or answer == "<- Back":
|
||||||
|
break
|
||||||
|
|
||||||
|
assert isinstance(answer, str)
|
||||||
|
|
||||||
|
if answer == "[+] Add new preset":
|
||||||
|
name_input = _get_questionary().text(
|
||||||
|
"Preset name:",
|
||||||
|
validate=lambda t: True if t and t.strip() else "Name cannot be empty",
|
||||||
|
).ask()
|
||||||
|
if not name_input:
|
||||||
|
continue
|
||||||
|
name = name_input.strip()
|
||||||
|
if name in config.model_presets:
|
||||||
|
console.print(f"[yellow]! Preset '{name}' already exists[/yellow]")
|
||||||
|
_pause()
|
||||||
|
continue
|
||||||
|
if name == "default":
|
||||||
|
console.print("[yellow]! 'default' is reserved (auto-generated from Agent Settings)[/yellow]")
|
||||||
|
_pause()
|
||||||
|
continue
|
||||||
|
new_preset = ModelPresetConfig(model="")
|
||||||
|
updated = _configure_pydantic_model(new_preset, f"New Preset: {name}")
|
||||||
|
if updated is not None:
|
||||||
|
config.model_presets[name] = updated
|
||||||
|
_sync_preset_cache(config)
|
||||||
|
last_preset_name = name
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Editing / deleting an existing preset
|
||||||
|
preset_name = answer.split(" (", 1)[0]
|
||||||
|
preset = config.model_presets.get(preset_name)
|
||||||
|
if preset is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
last_preset_name = preset_name
|
||||||
|
|
||||||
|
choices = ["Edit", "Cancel"]
|
||||||
|
if preset_name != "default":
|
||||||
|
choices.insert(1, "Delete")
|
||||||
|
action = _select_with_back(
|
||||||
|
f"Preset: {preset_name}",
|
||||||
|
choices,
|
||||||
|
default="Edit",
|
||||||
|
)
|
||||||
|
if action is _BACK_PRESSED or action == "Cancel" or action is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if action == "Delete":
|
||||||
|
confirm = _get_questionary().confirm(
|
||||||
|
f"Delete preset '{preset_name}'?",
|
||||||
|
default=False,
|
||||||
|
).ask()
|
||||||
|
if confirm:
|
||||||
|
del config.model_presets[preset_name]
|
||||||
|
_sync_preset_cache(config)
|
||||||
|
last_preset_name = None
|
||||||
|
continue
|
||||||
|
|
||||||
|
if action == "Edit":
|
||||||
|
updated = _configure_pydantic_model(preset, f"Edit Preset: {preset_name}")
|
||||||
|
if updated is not None:
|
||||||
|
config.model_presets[preset_name] = updated
|
||||||
|
_sync_preset_cache(config)
|
||||||
|
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
console.print("\n[dim]Returning to main menu...[/dim]")
|
||||||
|
break
|
||||||
|
|
||||||
|
|
||||||
# --- Provider Configuration ---
|
# --- Provider Configuration ---
|
||||||
|
|
||||||
|
|
||||||
@@ -1043,6 +1250,12 @@ def _show_summary(config: Config) -> None:
|
|||||||
channel_rows.append((display, status))
|
channel_rows.append((display, status))
|
||||||
_print_summary_panel(channel_rows, "Chat Channels")
|
_print_summary_panel(channel_rows, "Chat Channels")
|
||||||
|
|
||||||
|
# Model Presets
|
||||||
|
preset_rows = []
|
||||||
|
for name, preset in config.model_presets.items():
|
||||||
|
preset_rows.append((name, f"{preset.model} (ctx={preset.context_window_tokens})"))
|
||||||
|
_print_summary_panel(preset_rows, "Model Presets")
|
||||||
|
|
||||||
# Settings sections
|
# Settings sections
|
||||||
for title, model in [
|
for title, model in [
|
||||||
("Agent Settings", config.agents.defaults),
|
("Agent Settings", config.agents.defaults),
|
||||||
@@ -1112,6 +1325,7 @@ def run_onboard(initial_config: Config | None = None) -> OnboardResult:
|
|||||||
|
|
||||||
original_config = base_config.model_copy(deep=True)
|
original_config = base_config.model_copy(deep=True)
|
||||||
config = base_config.model_copy(deep=True)
|
config = base_config.model_copy(deep=True)
|
||||||
|
_sync_preset_cache(config)
|
||||||
|
|
||||||
last_main_choice: str | None = None
|
last_main_choice: str | None = None
|
||||||
while True:
|
while True:
|
||||||
@@ -1123,6 +1337,7 @@ def run_onboard(initial_config: Config | None = None) -> OnboardResult:
|
|||||||
"What would you like to configure?",
|
"What would you like to configure?",
|
||||||
choices=[
|
choices=[
|
||||||
"[P] LLM Provider",
|
"[P] LLM Provider",
|
||||||
|
"[M] Model Presets",
|
||||||
"[C] Chat Channel",
|
"[C] Chat Channel",
|
||||||
"[H] Channel Common",
|
"[H] Channel Common",
|
||||||
"[A] Agent Settings",
|
"[A] Agent Settings",
|
||||||
@@ -1149,6 +1364,7 @@ def run_onboard(initial_config: Config | None = None) -> OnboardResult:
|
|||||||
|
|
||||||
_menu_dispatch = {
|
_menu_dispatch = {
|
||||||
"[P] LLM Provider": lambda: _configure_providers(config),
|
"[P] LLM Provider": lambda: _configure_providers(config),
|
||||||
|
"[M] Model Presets": lambda: _configure_model_presets(config),
|
||||||
"[C] Chat Channel": lambda: _configure_channels(config),
|
"[C] Chat Channel": lambda: _configure_channels(config),
|
||||||
"[H] Channel Common": lambda: _configure_general_settings(config, "Channel Common"),
|
"[H] Channel Common": lambda: _configure_general_settings(config, "Channel Common"),
|
||||||
"[A] Agent Settings": lambda: _configure_general_settings(config, "Agent Settings"),
|
"[A] Agent Settings": lambda: _configure_general_settings(config, "Agent Settings"),
|
||||||
|
|||||||
+28
-14
@@ -180,6 +180,28 @@ class BedrockProviderConfig(ProviderConfig):
|
|||||||
profile: str | None = None # Optional AWS shared config profile
|
profile: str | None = None # Optional AWS shared config profile
|
||||||
|
|
||||||
|
|
||||||
|
class XaiOAuthXSearchConfig(Base):
|
||||||
|
"""xAI hosted X Search configuration."""
|
||||||
|
|
||||||
|
enable: bool = True
|
||||||
|
allowed_x_handles: list[str] | None = None
|
||||||
|
excluded_x_handles: list[str] | None = None
|
||||||
|
from_date: str | None = None
|
||||||
|
to_date: str | None = None
|
||||||
|
enable_image_understanding: bool = False
|
||||||
|
enable_video_understanding: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
class XaiOAuthProviderConfig(ProviderConfig):
|
||||||
|
"""xAI OAuth provider configuration."""
|
||||||
|
|
||||||
|
x_search: XaiOAuthXSearchConfig = Field(default_factory=XaiOAuthXSearchConfig)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_default_xai_oauth_config(value: Any) -> bool:
|
||||||
|
return isinstance(value, XaiOAuthProviderConfig) and value == XaiOAuthProviderConfig()
|
||||||
|
|
||||||
|
|
||||||
class ProvidersConfig(Base):
|
class ProvidersConfig(Base):
|
||||||
"""Configuration for LLM providers."""
|
"""Configuration for LLM providers."""
|
||||||
|
|
||||||
@@ -190,6 +212,7 @@ class ProvidersConfig(Base):
|
|||||||
openai: ProviderConfig = Field(default_factory=ProviderConfig)
|
openai: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
openrouter: ProviderConfig = Field(default_factory=ProviderConfig)
|
openrouter: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
huggingface: ProviderConfig = Field(default_factory=ProviderConfig)
|
huggingface: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
|
skywork: ProviderConfig = Field(default_factory=ProviderConfig) # Skywork / APIFree API gateway
|
||||||
deepseek: ProviderConfig = Field(default_factory=ProviderConfig)
|
deepseek: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
groq: ProviderConfig = Field(default_factory=ProviderConfig)
|
groq: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
zhipu: ProviderConfig = Field(default_factory=ProviderConfig)
|
zhipu: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
@@ -207,6 +230,7 @@ class ProvidersConfig(Base):
|
|||||||
stepfun: ProviderConfig = Field(default_factory=ProviderConfig) # Step Fun (阶跃星辰)
|
stepfun: ProviderConfig = Field(default_factory=ProviderConfig) # Step Fun (阶跃星辰)
|
||||||
xiaomi_mimo: ProviderConfig = Field(default_factory=ProviderConfig) # Xiaomi MIMO (小米)
|
xiaomi_mimo: ProviderConfig = Field(default_factory=ProviderConfig) # Xiaomi MIMO (小米)
|
||||||
longcat: ProviderConfig = Field(default_factory=ProviderConfig) # LongCat
|
longcat: ProviderConfig = Field(default_factory=ProviderConfig) # LongCat
|
||||||
|
ant_ling: ProviderConfig = Field(default_factory=ProviderConfig) # Ant Ling
|
||||||
aihubmix: ProviderConfig = Field(default_factory=ProviderConfig) # AiHubMix API gateway
|
aihubmix: ProviderConfig = Field(default_factory=ProviderConfig) # AiHubMix API gateway
|
||||||
siliconflow: ProviderConfig = Field(default_factory=ProviderConfig) # SiliconFlow (硅基流动)
|
siliconflow: ProviderConfig = Field(default_factory=ProviderConfig) # SiliconFlow (硅基流动)
|
||||||
volcengine: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine (火山引擎)
|
volcengine: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine (火山引擎)
|
||||||
@@ -215,6 +239,10 @@ class ProvidersConfig(Base):
|
|||||||
byteplus_coding_plan: ProviderConfig = Field(default_factory=ProviderConfig) # BytePlus Coding Plan
|
byteplus_coding_plan: ProviderConfig = Field(default_factory=ProviderConfig) # BytePlus Coding Plan
|
||||||
openai_codex: ProviderConfig = Field(default_factory=ProviderConfig, exclude=True) # OpenAI Codex (OAuth)
|
openai_codex: ProviderConfig = Field(default_factory=ProviderConfig, exclude=True) # OpenAI Codex (OAuth)
|
||||||
github_copilot: ProviderConfig = Field(default_factory=ProviderConfig, exclude=True) # Github Copilot (OAuth)
|
github_copilot: ProviderConfig = Field(default_factory=ProviderConfig, exclude=True) # Github Copilot (OAuth)
|
||||||
|
xai_oauth: XaiOAuthProviderConfig = Field(
|
||||||
|
default_factory=XaiOAuthProviderConfig,
|
||||||
|
exclude_if=_is_default_xai_oauth_config,
|
||||||
|
) # xAI Grok OAuth
|
||||||
qianfan: ProviderConfig = Field(default_factory=ProviderConfig) # Qianfan (百度千帆)
|
qianfan: ProviderConfig = Field(default_factory=ProviderConfig) # Qianfan (百度千帆)
|
||||||
nvidia: ProviderConfig = Field(default_factory=ProviderConfig) # NVIDIA NIM (nvapi- keys)
|
nvidia: ProviderConfig = Field(default_factory=ProviderConfig) # NVIDIA NIM (nvapi- keys)
|
||||||
|
|
||||||
@@ -282,19 +310,6 @@ class ToolsConfig(Base):
|
|||||||
ssrf_whitelist: list[str] = Field(default_factory=list) # CIDR ranges to exempt from SSRF blocking (e.g. ["100.64.0.0/10"] for Tailscale)
|
ssrf_whitelist: list[str] = Field(default_factory=list) # CIDR ranges to exempt from SSRF blocking (e.g. ["100.64.0.0/10"] for Tailscale)
|
||||||
|
|
||||||
|
|
||||||
class P2PConfig(Base):
|
|
||||||
"""P2P collaboration network configuration."""
|
|
||||||
|
|
||||||
enabled: bool = False
|
|
||||||
agent_id: str = ""
|
|
||||||
description: str = ""
|
|
||||||
capabilities: list[str] = Field(default_factory=list)
|
|
||||||
allow_from: list[str] = Field(default_factory=lambda: ["*"])
|
|
||||||
max_concurrent_tasks: int = 3
|
|
||||||
poll_interval: float = 5.0
|
|
||||||
mailboxes_root: str = "~/.nanobot/mailboxes"
|
|
||||||
|
|
||||||
|
|
||||||
class Config(BaseSettings):
|
class Config(BaseSettings):
|
||||||
"""Root configuration for nanobot."""
|
"""Root configuration for nanobot."""
|
||||||
|
|
||||||
@@ -308,7 +323,6 @@ class Config(BaseSettings):
|
|||||||
default_factory=dict,
|
default_factory=dict,
|
||||||
validation_alias=AliasChoices("modelPresets", "model_presets"),
|
validation_alias=AliasChoices("modelPresets", "model_presets"),
|
||||||
)
|
)
|
||||||
mailbox: P2PConfig = Field(default_factory=P2PConfig)
|
|
||||||
|
|
||||||
@model_validator(mode="after")
|
@model_validator(mode="after")
|
||||||
def _validate_model_preset(self) -> "Config":
|
def _validate_model_preset(self) -> "Config":
|
||||||
|
|||||||
@@ -1,6 +1,18 @@
|
|||||||
"""Cron service for scheduled agent tasks."""
|
"""Cron service for scheduled agent tasks."""
|
||||||
|
|
||||||
from nanobot.cron.service import CronService
|
|
||||||
from nanobot.cron.types import CronJob, CronSchedule
|
from nanobot.cron.types import CronJob, CronSchedule
|
||||||
|
|
||||||
__all__ = ["CronService", "CronJob", "CronSchedule"]
|
__all__ = ["CronService", "CronJob", "CronSchedule"]
|
||||||
|
|
||||||
|
_LAZY = {"CronService": ".service"}
|
||||||
|
|
||||||
|
|
||||||
|
def __getattr__(name: str):
|
||||||
|
module_path = _LAZY.get(name)
|
||||||
|
if module_path is None:
|
||||||
|
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||||
|
from importlib import import_module
|
||||||
|
mod = import_module(module_path, __name__)
|
||||||
|
val = getattr(mod, name)
|
||||||
|
globals()[name] = val
|
||||||
|
return val
|
||||||
|
|||||||
@@ -4,12 +4,12 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any, Callable, Coroutine
|
from typing import Any, Callable, Coroutine
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
from nanobot.providers.base import LLMProvider
|
||||||
from nanobot.providers.base import LLMProvider
|
from nanobot.utils.llm_runtime import LLMRuntimeResolver, static_llm_runtime
|
||||||
|
|
||||||
_HEARTBEAT_TOOL = [
|
_HEARTBEAT_TOOL = [
|
||||||
{
|
{
|
||||||
@@ -53,29 +53,28 @@ class HeartbeatService:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
workspace: Path,
|
workspace: Path,
|
||||||
provider: LLMProvider,
|
provider: LLMProvider | None = None,
|
||||||
model: str,
|
model: str | None = None,
|
||||||
on_execute: Callable[[str], Coroutine[Any, Any, str]] | None = None,
|
on_execute: Callable[[str], Coroutine[Any, Any, str]] | None = None,
|
||||||
on_notify: Callable[[str], Coroutine[Any, Any, None]] | None = None,
|
on_notify: Callable[[str], Coroutine[Any, Any, None]] | None = None,
|
||||||
interval_s: int = 30 * 60,
|
interval_s: int = 30 * 60,
|
||||||
enabled: bool = True,
|
enabled: bool = True,
|
||||||
timezone: str | None = None,
|
timezone: str | None = None,
|
||||||
p2p_shell: Any | None = None,
|
llm_runtime: LLMRuntimeResolver | None = None,
|
||||||
bus: Any | None = None,
|
|
||||||
):
|
):
|
||||||
self.workspace = workspace
|
self.workspace = workspace
|
||||||
self.provider = provider
|
if llm_runtime is None:
|
||||||
self.model = model
|
if provider is None or model is None:
|
||||||
|
raise ValueError("HeartbeatService requires either llm_runtime or provider/model")
|
||||||
|
llm_runtime = static_llm_runtime(provider, model)
|
||||||
|
self._llm_runtime = llm_runtime
|
||||||
self.on_execute = on_execute
|
self.on_execute = on_execute
|
||||||
self.on_notify = on_notify
|
self.on_notify = on_notify
|
||||||
self.interval_s = interval_s
|
self.interval_s = interval_s
|
||||||
self.enabled = enabled
|
self.enabled = enabled
|
||||||
self.timezone = timezone
|
self.timezone = timezone
|
||||||
self.p2p_shell = p2p_shell
|
|
||||||
self.bus = bus
|
|
||||||
self._running = False
|
self._running = False
|
||||||
self._task: asyncio.Task | None = None
|
self._task: asyncio.Task | None = None
|
||||||
self._last_inbox_scan: float = 0.0
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def heartbeat_file(self) -> Path:
|
def heartbeat_file(self) -> Path:
|
||||||
@@ -96,7 +95,9 @@ class HeartbeatService:
|
|||||||
"""
|
"""
|
||||||
from nanobot.utils.helpers import current_time_str
|
from nanobot.utils.helpers import current_time_str
|
||||||
|
|
||||||
response = await self.provider.chat_with_retry(
|
llm = self._llm_runtime()
|
||||||
|
|
||||||
|
response = await llm.provider.chat_with_retry(
|
||||||
messages=[
|
messages=[
|
||||||
{"role": "system", "content": "You are a heartbeat agent. Call the heartbeat tool to report your decision."},
|
{"role": "system", "content": "You are a heartbeat agent. Call the heartbeat tool to report your decision."},
|
||||||
{"role": "user", "content": (
|
{"role": "user", "content": (
|
||||||
@@ -106,7 +107,7 @@ class HeartbeatService:
|
|||||||
)},
|
)},
|
||||||
],
|
],
|
||||||
tools=_HEARTBEAT_TOOL,
|
tools=_HEARTBEAT_TOOL,
|
||||||
model=self.model,
|
model=llm.model,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not response.should_execute_tools:
|
if not response.should_execute_tools:
|
||||||
@@ -190,32 +191,6 @@ class HeartbeatService:
|
|||||||
"""Execute a single heartbeat tick."""
|
"""Execute a single heartbeat tick."""
|
||||||
from nanobot.utils.evaluator import evaluate_response
|
from nanobot.utils.evaluator import evaluate_response
|
||||||
|
|
||||||
# --- P2P inbox scan ---
|
|
||||||
if self.p2p_shell and self.bus:
|
|
||||||
try:
|
|
||||||
new_msgs = self.p2p_shell.scan_new_inbox(since=self._last_inbox_scan)
|
|
||||||
if new_msgs:
|
|
||||||
self._last_inbox_scan = time.time()
|
|
||||||
from nanobot.bus.events import InboundMessage
|
|
||||||
for msg in new_msgs:
|
|
||||||
await self.bus.publish_inbound(
|
|
||||||
InboundMessage(
|
|
||||||
channel="p2p",
|
|
||||||
sender_id=msg.get("from", "unknown"),
|
|
||||||
chat_id=msg.get("task_id", ""),
|
|
||||||
content=msg.get("payload", {}).get("description", ""),
|
|
||||||
metadata={"p2p_msg": msg},
|
|
||||||
)
|
|
||||||
)
|
|
||||||
logger.info(
|
|
||||||
"Heartbeat: injected P2P task {} from {}",
|
|
||||||
msg.get("task_id", ""),
|
|
||||||
msg.get("from", "unknown"),
|
|
||||||
)
|
|
||||||
except Exception:
|
|
||||||
logger.exception("Heartbeat P2P scan failed")
|
|
||||||
|
|
||||||
# --- Legacy heartbeat file check ---
|
|
||||||
content = self._read_heartbeat_file()
|
content = self._read_heartbeat_file()
|
||||||
if not content:
|
if not content:
|
||||||
logger.debug("Heartbeat: HEARTBEAT.md missing or empty")
|
logger.debug("Heartbeat: HEARTBEAT.md missing or empty")
|
||||||
@@ -245,8 +220,9 @@ class HeartbeatService:
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
llm = self._llm_runtime()
|
||||||
should_notify = await evaluate_response(
|
should_notify = await evaluate_response(
|
||||||
response, tasks, self.provider, self.model,
|
response, tasks, llm.provider, llm.model,
|
||||||
)
|
)
|
||||||
if should_notify and self.on_notify:
|
if should_notify and self.on_notify:
|
||||||
logger.info("Heartbeat: completed, delivering response")
|
logger.info("Heartbeat: completed, delivering response")
|
||||||
|
|||||||
+2
-4
@@ -8,6 +8,7 @@ from typing import Any
|
|||||||
|
|
||||||
from nanobot.agent.hook import AgentHook, SDKCaptureHook
|
from nanobot.agent.hook import AgentHook, SDKCaptureHook
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
@@ -63,10 +64,7 @@ class Nanobot:
|
|||||||
|
|
||||||
loop = AgentLoop.from_config(
|
loop = AgentLoop.from_config(
|
||||||
config,
|
config,
|
||||||
image_generation_provider_configs={
|
image_generation_provider_configs=image_gen_provider_configs(config),
|
||||||
"openrouter": config.providers.openrouter,
|
|
||||||
"aihubmix": config.providers.aihubmix,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
return cls(loop)
|
return cls(loop)
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +0,0 @@
|
|||||||
"""P2P inter-agent coordination layer."""
|
|
||||||
|
|
||||||
from nanobot.p2p.shell import P2PShell
|
|
||||||
|
|
||||||
__all__ = ["P2PShell"]
|
|
||||||
@@ -1,426 +0,0 @@
|
|||||||
"""P2P shell: filesystem-backed inter-agent coordination.
|
|
||||||
|
|
||||||
All state is stored in the mailbox filesystem; this class is stateless.
|
|
||||||
Restarting the gateway restores all task state by scanning files.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
|
||||||
import os
|
|
||||||
import time
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Literal
|
|
||||||
|
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
|
|
||||||
class P2PShell:
|
|
||||||
"""Stateless P2P coordination shell backed by the mailbox filesystem."""
|
|
||||||
|
|
||||||
def __init__(self, agent_id: str, mailboxes_root: str):
|
|
||||||
self.agent_id = agent_id
|
|
||||||
self.root = Path(mailboxes_root).expanduser()
|
|
||||||
self.inbox = self.root / agent_id / "inbox"
|
|
||||||
self.processed = self.root / agent_id / "processed"
|
|
||||||
self.links_dir = self.root / "_links"
|
|
||||||
self.windows_dir = self.root / "_windows"
|
|
||||||
|
|
||||||
for d in (self.inbox, self.processed, self.links_dir, self.windows_dir):
|
|
||||||
d.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
# Discovery
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
|
|
||||||
def discover(self, capability: str, top_k: int = 3) -> list[dict[str, Any]]:
|
|
||||||
"""Read _registry.json and return candidates matching capability."""
|
|
||||||
registry = self._load_json(self.root / "_registry.json", default={})
|
|
||||||
candidates: list[dict[str, Any]] = []
|
|
||||||
for aid, info in registry.items():
|
|
||||||
if aid == self.agent_id:
|
|
||||||
continue
|
|
||||||
caps = info.get("capabilities", [])
|
|
||||||
if capability.lower() in " ".join(caps).lower():
|
|
||||||
candidates.append({"agent_id": aid, **info})
|
|
||||||
# Sort: idle first, then by current task load
|
|
||||||
candidates.sort(key=lambda x: (x.get("status") != "idle", x.get("current_tasks", 0)))
|
|
||||||
return candidates[:top_k]
|
|
||||||
|
|
||||||
def heartbeat(self, description: str, capabilities: list[str]) -> None:
|
|
||||||
"""Write self state into the shared _registry.json."""
|
|
||||||
registry = self._load_json(self.root / "_registry.json", default={})
|
|
||||||
registry[self.agent_id] = {
|
|
||||||
"description": description,
|
|
||||||
"capabilities": capabilities,
|
|
||||||
"status": "idle",
|
|
||||||
"last_heartbeat": int(time.time()),
|
|
||||||
"endpoint": "",
|
|
||||||
}
|
|
||||||
self._atomic_write(self.root / "_registry.json", registry)
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
# Task dispatch
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
|
|
||||||
def dispatch(
|
|
||||||
self,
|
|
||||||
to: str,
|
|
||||||
parent_task_id: str | None,
|
|
||||||
description: str,
|
|
||||||
deadline_seconds: int = 300,
|
|
||||||
allow_redelegation: bool = True,
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
"""Write a task into the target agent's inbox and return a receipt."""
|
|
||||||
task_id = (
|
|
||||||
f"{parent_task_id}.{int(time.time())}"
|
|
||||||
if parent_task_id
|
|
||||||
else f"root_{int(time.time())}"
|
|
||||||
)
|
|
||||||
|
|
||||||
depth = self._get_depth(parent_task_id) if parent_task_id else 0
|
|
||||||
if depth >= 3:
|
|
||||||
return {"status": "rejected", "reason": "max_depth_exceeded"}
|
|
||||||
|
|
||||||
if parent_task_id and self._is_ancestor(to, parent_task_id):
|
|
||||||
return {"status": "rejected", "reason": "ancestry_loop"}
|
|
||||||
|
|
||||||
if not self._circuit_allow(to):
|
|
||||||
failover = self._find_failover(to)
|
|
||||||
return {"status": "circuit_open", "failover_to": failover}
|
|
||||||
|
|
||||||
target_inbox = self.root / to / "inbox"
|
|
||||||
target_inbox.mkdir(parents=True, exist_ok=True)
|
|
||||||
if list(target_inbox.glob(f"task_{task_id}_from_{self.agent_id}_*.json")):
|
|
||||||
return {"status": "dispatched", "task_id": task_id, "note": "cached"}
|
|
||||||
|
|
||||||
ancestry = (
|
|
||||||
(self._get_ancestry(parent_task_id) + [self.agent_id])
|
|
||||||
if parent_task_id
|
|
||||||
else [self.agent_id]
|
|
||||||
)
|
|
||||||
|
|
||||||
msg: dict[str, Any] = {
|
|
||||||
"version": "p2p/v1",
|
|
||||||
"type": "task_dispatch",
|
|
||||||
"from": self.agent_id,
|
|
||||||
"to": to,
|
|
||||||
"task_id": task_id,
|
|
||||||
"ancestry": ancestry,
|
|
||||||
"depth": depth + 1,
|
|
||||||
"payload": {
|
|
||||||
"description": description,
|
|
||||||
"allow_redelegation": allow_redelegation,
|
|
||||||
},
|
|
||||||
"deadline": int(time.time()) + deadline_seconds,
|
|
||||||
"timestamp": int(time.time()),
|
|
||||||
}
|
|
||||||
|
|
||||||
path = target_inbox / f"task_{task_id}_from_{self.agent_id}_{os.urandom(4).hex()}.json"
|
|
||||||
self._atomic_write(path, msg)
|
|
||||||
logger.info("P2P dispatch: {} -> {} (task_id={})", self.agent_id, to, task_id)
|
|
||||||
return {"status": "dispatched", "task_id": task_id, "depth": depth + 1}
|
|
||||||
|
|
||||||
def poll(self, task_id: str) -> dict[str, Any]:
|
|
||||||
"""Scan inbox/processed and return task status."""
|
|
||||||
# Check processed results first
|
|
||||||
results = list(self.processed.glob(f"result_{task_id}_from_*.json"))
|
|
||||||
if results:
|
|
||||||
data = self._load_json(results[0])
|
|
||||||
payload = data.get("payload", {})
|
|
||||||
return {
|
|
||||||
"status": payload.get("outcome", "completed"),
|
|
||||||
"result": payload.get("content", ""),
|
|
||||||
"from": data["from"],
|
|
||||||
}
|
|
||||||
|
|
||||||
# Check inbox for results (not yet moved to processed)
|
|
||||||
inbox_results = list(self.inbox.glob(f"result_{task_id}_from_*.json"))
|
|
||||||
if inbox_results:
|
|
||||||
data = self._load_json(inbox_results[0])
|
|
||||||
payload = data.get("payload", {})
|
|
||||||
return {
|
|
||||||
"status": payload.get("outcome", "completed"),
|
|
||||||
"result": payload.get("content", ""),
|
|
||||||
"from": data["from"],
|
|
||||||
}
|
|
||||||
|
|
||||||
# Check inbox for pending task dispatches
|
|
||||||
pending = list(self.inbox.glob(f"task_{task_id}_from_*.json"))
|
|
||||||
if pending:
|
|
||||||
data = self._load_json(pending[0])
|
|
||||||
deadline = data.get("deadline", 0)
|
|
||||||
elapsed = int(time.time() - data["timestamp"])
|
|
||||||
if time.time() > deadline:
|
|
||||||
return {"status": "timeout", "elapsed": elapsed}
|
|
||||||
return {"status": "pending", "elapsed": elapsed}
|
|
||||||
|
|
||||||
return {"status": "not_found"}
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
# Aggregation (broadcast + check)
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
|
|
||||||
def broadcast(
|
|
||||||
self,
|
|
||||||
task_id: str,
|
|
||||||
subtasks: list[dict[str, Any]],
|
|
||||||
aggregation_timeout: int = 30,
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
"""Write bid requests to candidate agents and create a window descriptor."""
|
|
||||||
targets: list[tuple[str, str]] = [] # (subtask_id, agent_id)
|
|
||||||
for sub in subtasks:
|
|
||||||
caps = sub.get("capability", "")
|
|
||||||
found = self.discover(caps, top_k=3)
|
|
||||||
targets.extend([(sub["subtask_id"], a["agent_id"]) for a in found])
|
|
||||||
|
|
||||||
for subtask_id, target in targets:
|
|
||||||
msg: dict[str, Any] = {
|
|
||||||
"version": "p2p/v1",
|
|
||||||
"type": "bid_request",
|
|
||||||
"from": self.agent_id,
|
|
||||||
"to": target,
|
|
||||||
"task_id": task_id,
|
|
||||||
"subtask_id": subtask_id,
|
|
||||||
"payload": sub,
|
|
||||||
"deadline": int(time.time()) + aggregation_timeout,
|
|
||||||
"timestamp": int(time.time()),
|
|
||||||
}
|
|
||||||
target_inbox = self.root / target / "inbox"
|
|
||||||
target_inbox.mkdir(parents=True, exist_ok=True)
|
|
||||||
path = target_inbox / f"bid_{task_id}_{subtask_id}_from_{self.agent_id}.json"
|
|
||||||
self._atomic_write(path, msg)
|
|
||||||
|
|
||||||
window: dict[str, Any] = {
|
|
||||||
"task_id": task_id,
|
|
||||||
"mode": "bid",
|
|
||||||
"expected": len(targets),
|
|
||||||
"deadline": int(time.time()) + aggregation_timeout,
|
|
||||||
"created_at": int(time.time()),
|
|
||||||
}
|
|
||||||
self._atomic_write(self.windows_dir / f"{task_id}.json", window)
|
|
||||||
logger.info(
|
|
||||||
"P2P broadcast: {} invited {} agents for task_id={}",
|
|
||||||
self.agent_id,
|
|
||||||
len(targets),
|
|
||||||
task_id,
|
|
||||||
)
|
|
||||||
return {"status": "bidding_opened", "task_id": task_id, "invited": len(targets)}
|
|
||||||
|
|
||||||
def check_aggregation(self, task_id: str) -> dict[str, Any]:
|
|
||||||
"""Lazily check aggregation status by scanning files."""
|
|
||||||
window_path = self.windows_dir / f"{task_id}.json"
|
|
||||||
if not window_path.exists():
|
|
||||||
return {"status": "no_window"}
|
|
||||||
|
|
||||||
window = self._load_json(window_path)
|
|
||||||
mode = window.get("mode", "bid")
|
|
||||||
deadline = window.get("deadline", 0)
|
|
||||||
|
|
||||||
pattern = f"{mode}_{task_id}_*_from_*.json"
|
|
||||||
entries: list[dict[str, Any]] = []
|
|
||||||
for f in self.inbox.glob(pattern):
|
|
||||||
data = self._load_json(f)
|
|
||||||
entries.append(
|
|
||||||
{
|
|
||||||
"from": data.get("from", ""),
|
|
||||||
"subtask_id": data.get("subtask_id", ""),
|
|
||||||
"payload": data.get("payload", {}),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
is_timeout = time.time() > deadline
|
|
||||||
is_full = window.get("expected") and len(entries) >= window["expected"]
|
|
||||||
|
|
||||||
if is_timeout or is_full:
|
|
||||||
self._atomic_write(
|
|
||||||
self.processed / f"window_{task_id}.json",
|
|
||||||
{**window, "closed_at": int(time.time()), "received": len(entries)},
|
|
||||||
)
|
|
||||||
window_path.unlink(missing_ok=True)
|
|
||||||
return {
|
|
||||||
"status": "closed",
|
|
||||||
"mode": mode,
|
|
||||||
"entries": entries,
|
|
||||||
"reason": "timeout" if is_timeout else "full",
|
|
||||||
}
|
|
||||||
|
|
||||||
return {
|
|
||||||
"status": "pending",
|
|
||||||
"received": len(entries),
|
|
||||||
"expected": window.get("expected"),
|
|
||||||
"seconds_remaining": max(0, deadline - int(time.time())),
|
|
||||||
}
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
# Result reporting
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
|
|
||||||
def report_result(
|
|
||||||
self,
|
|
||||||
to: str,
|
|
||||||
task_id: str,
|
|
||||||
outcome: Literal["completed", "failed", "aborted"],
|
|
||||||
content: str,
|
|
||||||
callback: dict[str, Any] | None = None,
|
|
||||||
) -> None:
|
|
||||||
"""Worker calls this to write a result into the manager's inbox."""
|
|
||||||
msg: dict[str, Any] = {
|
|
||||||
"version": "p2p/v1",
|
|
||||||
"type": "result",
|
|
||||||
"from": self.agent_id,
|
|
||||||
"to": to,
|
|
||||||
"task_id": task_id,
|
|
||||||
"payload": {"outcome": outcome, "content": content},
|
|
||||||
"timestamp": int(time.time()),
|
|
||||||
}
|
|
||||||
if callback:
|
|
||||||
msg["callback"] = callback
|
|
||||||
target_inbox = self.root / to / "inbox"
|
|
||||||
target_inbox.mkdir(parents=True, exist_ok=True)
|
|
||||||
path = target_inbox / f"result_{task_id}_from_{self.agent_id}_{os.urandom(4).hex()}.json"
|
|
||||||
self._atomic_write(path, msg)
|
|
||||||
logger.info("P2P result: {} -> {} (task_id={}, outcome={})", self.agent_id, to, task_id, outcome)
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
# Finalization
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
|
|
||||||
def finalize(self, task_id: str, outcome: str, reason: str = "") -> None:
|
|
||||||
"""Move all task files from inbox to processed and mark outcome."""
|
|
||||||
for src in list(self.inbox.glob(f"*{task_id}*")):
|
|
||||||
data = self._load_json(src)
|
|
||||||
data.setdefault("payload", {})
|
|
||||||
data["payload"]["outcome"] = outcome
|
|
||||||
data["payload"]["reason"] = reason
|
|
||||||
dst = self.processed / src.name
|
|
||||||
self._atomic_write(dst, data)
|
|
||||||
src.unlink(missing_ok=True)
|
|
||||||
logger.info("P2P finalize: task_id={} outcome={}", task_id, outcome)
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
# Circuit breaker
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
|
|
||||||
def _circuit_allow(self, to: str) -> bool:
|
|
||||||
link = self._load_json(
|
|
||||||
self.links_dir / f"{to}.json",
|
|
||||||
default={"failures": 0, "last_failure": 0, "open": False},
|
|
||||||
)
|
|
||||||
if not link.get("open"):
|
|
||||||
return True
|
|
||||||
backoff = 300 * (2 ** max(0, link.get("failures", 0) - 3))
|
|
||||||
if time.time() - link.get("last_failure", 0) > backoff:
|
|
||||||
link["open"] = False
|
|
||||||
self._atomic_write(self.links_dir / f"{to}.json", link)
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
def record_failure(self, to: str) -> None:
|
|
||||||
link = self._load_json(
|
|
||||||
self.links_dir / f"{to}.json",
|
|
||||||
default={"failures": 0, "last_failure": 0, "open": False},
|
|
||||||
)
|
|
||||||
link["failures"] = link.get("failures", 0) + 1
|
|
||||||
link["last_failure"] = int(time.time())
|
|
||||||
if link["failures"] >= 3:
|
|
||||||
link["open"] = True
|
|
||||||
self._atomic_write(self.links_dir / f"{to}.json", link)
|
|
||||||
|
|
||||||
def record_success(self, to: str) -> None:
|
|
||||||
link = self._load_json(
|
|
||||||
self.links_dir / f"{to}.json",
|
|
||||||
default={"failures": 0, "last_failure": 0, "open": False},
|
|
||||||
)
|
|
||||||
link["failures"] = 0
|
|
||||||
link["open"] = False
|
|
||||||
self._atomic_write(self.links_dir / f"{to}.json", link)
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
# Inbox scanning (for HeartbeatService)
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
|
|
||||||
def scan_inbox(self) -> list[dict[str, Any]]:
|
|
||||||
"""Return all task_dispatch messages currently in inbox."""
|
|
||||||
messages: list[dict[str, Any]] = []
|
|
||||||
for f in sorted(self.inbox.glob("task_*_from_*.json"), key=lambda p: p.stat().st_mtime):
|
|
||||||
data = self._load_json(f)
|
|
||||||
# Skip expired tasks
|
|
||||||
if time.time() > data.get("deadline", 0):
|
|
||||||
continue
|
|
||||||
data["_filename"] = f.name
|
|
||||||
messages.append(data)
|
|
||||||
return messages
|
|
||||||
|
|
||||||
def scan_new_inbox(self, since: float | None = None) -> list[dict[str, Any]]:
|
|
||||||
"""Return inbox messages newer than the given timestamp."""
|
|
||||||
messages: list[dict[str, Any]] = []
|
|
||||||
for f in self.inbox.glob("task_*_from_*.json"):
|
|
||||||
mtime = f.stat().st_mtime
|
|
||||||
if since is not None and mtime <= since:
|
|
||||||
continue
|
|
||||||
data = self._load_json(f)
|
|
||||||
if time.time() > data.get("deadline", 0):
|
|
||||||
continue
|
|
||||||
data["_filename"] = f.name
|
|
||||||
data["_mtime"] = mtime
|
|
||||||
messages.append(data)
|
|
||||||
return sorted(messages, key=lambda x: x.get("_mtime", 0))
|
|
||||||
|
|
||||||
def mark_processed(self, filename: str) -> None:
|
|
||||||
"""Move a single inbox file to processed."""
|
|
||||||
src = self.inbox / filename
|
|
||||||
if not src.exists():
|
|
||||||
return
|
|
||||||
dst = self.processed / filename
|
|
||||||
try:
|
|
||||||
import shutil
|
|
||||||
shutil.move(str(src), str(dst))
|
|
||||||
except Exception:
|
|
||||||
logger.warning("Failed to mark processed: {}", filename)
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
# Helpers
|
|
||||||
# ------------------------------------------------------------------
|
|
||||||
|
|
||||||
def _load_json(self, path: Path, default: Any | None = None) -> Any:
|
|
||||||
if not path.exists():
|
|
||||||
return default if default is not None else {}
|
|
||||||
with open(path, "r", encoding="utf-8") as f:
|
|
||||||
return json.load(f)
|
|
||||||
|
|
||||||
def _atomic_write(self, path: Path, data: dict[str, Any]) -> None:
|
|
||||||
tmp = path.with_suffix(".tmp")
|
|
||||||
with open(tmp, "w", encoding="utf-8") as f:
|
|
||||||
json.dump(data, f, ensure_ascii=False, indent=2)
|
|
||||||
tmp.rename(path)
|
|
||||||
|
|
||||||
def _get_depth(self, task_id: str) -> int:
|
|
||||||
return task_id.count(".")
|
|
||||||
|
|
||||||
def _is_ancestor(self, agent_id: str, parent_task_id: str) -> bool:
|
|
||||||
for f in list(self.processed.glob(f"*{parent_task_id}*")) + list(
|
|
||||||
self.inbox.glob(f"*{parent_task_id}*")
|
|
||||||
):
|
|
||||||
data = self._load_json(f)
|
|
||||||
if agent_id in data.get("ancestry", []):
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
def _get_ancestry(self, task_id: str) -> list[str]:
|
|
||||||
for f in list(self.processed.glob(f"*{task_id}*")) + list(
|
|
||||||
self.inbox.glob(f"*{task_id}*")
|
|
||||||
):
|
|
||||||
data = self._load_json(f)
|
|
||||||
return data.get("ancestry", [])
|
|
||||||
return []
|
|
||||||
|
|
||||||
def _find_failover(self, to: str) -> str | None:
|
|
||||||
registry = self._load_json(self.root / "_registry.json", default={})
|
|
||||||
target_caps = registry.get(to, {}).get("capabilities", [])
|
|
||||||
for aid, info in registry.items():
|
|
||||||
if aid == to:
|
|
||||||
continue
|
|
||||||
if any(c in info.get("capabilities", []) for c in target_caps):
|
|
||||||
return aid
|
|
||||||
return None
|
|
||||||
@@ -14,6 +14,7 @@ __all__ = [
|
|||||||
"OpenAICompatProvider",
|
"OpenAICompatProvider",
|
||||||
"OpenAICodexProvider",
|
"OpenAICodexProvider",
|
||||||
"GitHubCopilotProvider",
|
"GitHubCopilotProvider",
|
||||||
|
"XaiOAuthProvider",
|
||||||
"AzureOpenAIProvider",
|
"AzureOpenAIProvider",
|
||||||
"BedrockProvider",
|
"BedrockProvider",
|
||||||
]
|
]
|
||||||
@@ -23,10 +24,23 @@ _LAZY_IMPORTS = {
|
|||||||
"OpenAICompatProvider": ".openai_compat_provider",
|
"OpenAICompatProvider": ".openai_compat_provider",
|
||||||
"OpenAICodexProvider": ".openai_codex_provider",
|
"OpenAICodexProvider": ".openai_codex_provider",
|
||||||
"GitHubCopilotProvider": ".github_copilot_provider",
|
"GitHubCopilotProvider": ".github_copilot_provider",
|
||||||
|
"XaiOAuthProvider": ".xai_oauth_provider",
|
||||||
"AzureOpenAIProvider": ".azure_openai_provider",
|
"AzureOpenAIProvider": ".azure_openai_provider",
|
||||||
"BedrockProvider": ".bedrock_provider",
|
"BedrockProvider": ".bedrock_provider",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
_LAZY_SUBMODULES = {
|
||||||
|
"anthropic_provider": ".anthropic_provider",
|
||||||
|
"openai_compat_provider": ".openai_compat_provider",
|
||||||
|
"openai_codex_provider": ".openai_codex_provider",
|
||||||
|
"github_copilot_provider": ".github_copilot_provider",
|
||||||
|
"xai_oauth_provider": ".xai_oauth_provider",
|
||||||
|
"azure_openai_provider": ".azure_openai_provider",
|
||||||
|
"bedrock_provider": ".bedrock_provider",
|
||||||
|
"factory": ".factory",
|
||||||
|
"registry": ".registry",
|
||||||
|
}
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nanobot.providers.anthropic_provider import AnthropicProvider
|
from nanobot.providers.anthropic_provider import AnthropicProvider
|
||||||
from nanobot.providers.azure_openai_provider import AzureOpenAIProvider
|
from nanobot.providers.azure_openai_provider import AzureOpenAIProvider
|
||||||
@@ -34,12 +48,18 @@ if TYPE_CHECKING:
|
|||||||
from nanobot.providers.github_copilot_provider import GitHubCopilotProvider
|
from nanobot.providers.github_copilot_provider import GitHubCopilotProvider
|
||||||
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
from nanobot.providers.openai_codex_provider import OpenAICodexProvider
|
from nanobot.providers.openai_codex_provider import OpenAICodexProvider
|
||||||
|
from nanobot.providers.xai_oauth_provider import XaiOAuthProvider
|
||||||
|
|
||||||
|
|
||||||
def __getattr__(name: str):
|
def __getattr__(name: str):
|
||||||
"""Lazily expose provider implementations without importing all backends up front."""
|
"""Lazily expose provider implementations without importing all backends up front."""
|
||||||
module_name = _LAZY_IMPORTS.get(name)
|
module_name = _LAZY_IMPORTS.get(name)
|
||||||
if module_name is None:
|
if module_name is not None:
|
||||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
module = import_module(module_name, __name__)
|
||||||
module = import_module(module_name, __name__)
|
return getattr(module, name)
|
||||||
return getattr(module, name)
|
module_name = _LAZY_SUBMODULES.get(name)
|
||||||
|
if module_name is not None:
|
||||||
|
module = import_module(module_name, __name__)
|
||||||
|
globals()[name] = module
|
||||||
|
return module
|
||||||
|
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||||
|
|||||||
@@ -590,6 +590,7 @@ class AnthropicProvider(LLMProvider):
|
|||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
kwargs = self._build_kwargs(
|
kwargs = self._build_kwargs(
|
||||||
messages, tools, model, max_tokens, temperature,
|
messages, tools, model, max_tokens, temperature,
|
||||||
@@ -598,11 +599,12 @@ class AnthropicProvider(LLMProvider):
|
|||||||
idle_timeout_s = int(os.environ.get("NANOBOT_STREAM_IDLE_TIMEOUT_S", "90"))
|
idle_timeout_s = int(os.environ.get("NANOBOT_STREAM_IDLE_TIMEOUT_S", "90"))
|
||||||
try:
|
try:
|
||||||
async with self._client.messages.stream(**kwargs) as stream:
|
async with self._client.messages.stream(**kwargs) as stream:
|
||||||
if on_content_delta or on_thinking_delta:
|
if on_content_delta or on_thinking_delta or on_tool_call_delta:
|
||||||
# Idle timeout must track *any* SSE chunk (thinking_delta,
|
# Idle timeout must track *any* SSE chunk (thinking_delta,
|
||||||
# tool JSON deltas, etc.), not only text_stream tokens.
|
# tool JSON deltas, etc.), not only text_stream tokens.
|
||||||
# Otherwise extended thinking can stall text_stream for minutes
|
# Otherwise extended thinking can stall text_stream for minutes
|
||||||
# while the connection is healthy (e.g. MiniMax Anthropic).
|
# while the connection is healthy (e.g. MiniMax Anthropic).
|
||||||
|
tool_blocks: dict[int, dict[str, str]] = {}
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
chunk = await asyncio.wait_for(
|
chunk = await asyncio.wait_for(
|
||||||
@@ -611,7 +613,22 @@ class AnthropicProvider(LLMProvider):
|
|||||||
)
|
)
|
||||||
except StopAsyncIteration:
|
except StopAsyncIteration:
|
||||||
break
|
break
|
||||||
if (
|
if chunk.type == "content_block_start":
|
||||||
|
block = getattr(chunk, "content_block", None)
|
||||||
|
if getattr(block, "type", None) == "tool_use":
|
||||||
|
index = int(getattr(chunk, "index", 0) or 0)
|
||||||
|
state = {
|
||||||
|
"call_id": str(getattr(block, "id", "") or ""),
|
||||||
|
"name": str(getattr(block, "name", "") or ""),
|
||||||
|
}
|
||||||
|
tool_blocks[index] = state
|
||||||
|
if on_tool_call_delta:
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": index,
|
||||||
|
**state,
|
||||||
|
"arguments_delta": "",
|
||||||
|
})
|
||||||
|
elif (
|
||||||
chunk.type == "content_block_delta"
|
chunk.type == "content_block_delta"
|
||||||
and getattr(chunk.delta, "type", None) == "thinking_delta"
|
and getattr(chunk.delta, "type", None) == "thinking_delta"
|
||||||
):
|
):
|
||||||
@@ -625,6 +642,20 @@ class AnthropicProvider(LLMProvider):
|
|||||||
text = getattr(chunk.delta, "text", None) or ""
|
text = getattr(chunk.delta, "text", None) or ""
|
||||||
if text and on_content_delta:
|
if text and on_content_delta:
|
||||||
await on_content_delta(text)
|
await on_content_delta(text)
|
||||||
|
elif (
|
||||||
|
chunk.type == "content_block_delta"
|
||||||
|
and getattr(chunk.delta, "type", None) == "input_json_delta"
|
||||||
|
):
|
||||||
|
partial = getattr(chunk.delta, "partial_json", None) or ""
|
||||||
|
if partial and on_tool_call_delta:
|
||||||
|
index = int(getattr(chunk, "index", 0) or 0)
|
||||||
|
state = tool_blocks.get(index, {})
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": index,
|
||||||
|
"call_id": state.get("call_id", ""),
|
||||||
|
"name": state.get("name", ""),
|
||||||
|
"arguments_delta": partial,
|
||||||
|
})
|
||||||
response = await asyncio.wait_for(
|
response = await asyncio.wait_for(
|
||||||
stream.get_final_message(),
|
stream.get_final_message(),
|
||||||
timeout=idle_timeout_s,
|
timeout=idle_timeout_s,
|
||||||
|
|||||||
@@ -158,6 +158,7 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
_ = on_thinking_delta
|
_ = on_thinking_delta
|
||||||
body = self._build_body(
|
body = self._build_body(
|
||||||
@@ -169,7 +170,7 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
try:
|
try:
|
||||||
stream = await self._client.responses.create(**body)
|
stream = await self._client.responses.create(**body)
|
||||||
content, tool_calls, finish_reason, usage, reasoning_content = (
|
content, tool_calls, finish_reason, usage, reasoning_content = (
|
||||||
await consume_sdk_stream(stream, on_content_delta)
|
await consume_sdk_stream(stream, on_content_delta, on_tool_call_delta)
|
||||||
)
|
)
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
content=content or None,
|
content=content or None,
|
||||||
|
|||||||
@@ -70,11 +70,11 @@ class LLMResponse:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def should_execute_tools(self) -> bool:
|
def should_execute_tools(self) -> bool:
|
||||||
"""Tools execute only when has_tool_calls AND finish_reason is ``tool_calls`` / ``stop``.
|
"""Tools execute only when has_tool_calls AND finish_reason is a tool-capable stop.
|
||||||
Blocks gateway-injected calls under ``refusal`` / ``content_filter`` / ``error`` (#3220)."""
|
Blocks gateway-injected calls under ``refusal`` / ``content_filter`` / ``error`` (#3220)."""
|
||||||
if not self.has_tool_calls:
|
if not self.has_tool_calls:
|
||||||
return False
|
return False
|
||||||
return self.finish_reason in ("tool_calls", "stop")
|
return self.finish_reason in ("tool_calls", "function_call", "stop")
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -112,6 +112,7 @@ class LLMProvider(ABC):
|
|||||||
"server error",
|
"server error",
|
||||||
"temporarily unavailable",
|
"temporarily unavailable",
|
||||||
"速率限制",
|
"速率限制",
|
||||||
|
"访问量过大",
|
||||||
)
|
)
|
||||||
_RETRYABLE_STATUS_CODES = frozenset({408, 409, 429})
|
_RETRYABLE_STATUS_CODES = frozenset({408, 409, 429})
|
||||||
_TRANSIENT_ERROR_KINDS = frozenset({"timeout", "connection"})
|
_TRANSIENT_ERROR_KINDS = frozenset({"timeout", "connection"})
|
||||||
@@ -500,6 +501,7 @@ class LLMProvider(ABC):
|
|||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
"""Stream a chat completion, calling *on_content_delta* for each text chunk.
|
"""Stream a chat completion, calling *on_content_delta* for each text chunk.
|
||||||
|
|
||||||
@@ -513,7 +515,7 @@ class LLMProvider(ABC):
|
|||||||
full content as a single delta. Providers that support native
|
full content as a single delta. Providers that support native
|
||||||
streaming should override this method.
|
streaming should override this method.
|
||||||
"""
|
"""
|
||||||
_ = on_thinking_delta
|
_ = on_thinking_delta, on_tool_call_delta
|
||||||
response = await self.chat(
|
response = await self.chat(
|
||||||
messages=messages, tools=tools, model=model,
|
messages=messages, tools=tools, model=model,
|
||||||
max_tokens=max_tokens, temperature=temperature,
|
max_tokens=max_tokens, temperature=temperature,
|
||||||
@@ -543,6 +545,7 @@ class LLMProvider(ABC):
|
|||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
retry_mode: str = "standard",
|
retry_mode: str = "standard",
|
||||||
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
@@ -560,6 +563,7 @@ class LLMProvider(ABC):
|
|||||||
reasoning_effort=reasoning_effort, tool_choice=tool_choice,
|
reasoning_effort=reasoning_effort, tool_choice=tool_choice,
|
||||||
on_content_delta=on_content_delta,
|
on_content_delta=on_content_delta,
|
||||||
on_thinking_delta=on_thinking_delta,
|
on_thinking_delta=on_thinking_delta,
|
||||||
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
)
|
)
|
||||||
return await self._run_with_retry(
|
return await self._run_with_retry(
|
||||||
self._safe_chat_stream,
|
self._safe_chat_stream,
|
||||||
|
|||||||
@@ -704,8 +704,9 @@ class BedrockProvider(LLMProvider):
|
|||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
_ = on_thinking_delta
|
_ = on_thinking_delta, on_tool_call_delta
|
||||||
idle_timeout_s = int(os.environ.get("NANOBOT_STREAM_IDLE_TIMEOUT_S", "90"))
|
idle_timeout_s = int(os.environ.get("NANOBOT_STREAM_IDLE_TIMEOUT_S", "90"))
|
||||||
content_parts: list[str] = []
|
content_parts: list[str] = []
|
||||||
reasoning_parts: list[str] = []
|
reasoning_parts: list[str] = []
|
||||||
|
|||||||
@@ -68,6 +68,10 @@ def _make_provider_core(
|
|||||||
from nanobot.providers.github_copilot_provider import GitHubCopilotProvider
|
from nanobot.providers.github_copilot_provider import GitHubCopilotProvider
|
||||||
|
|
||||||
provider = GitHubCopilotProvider(default_model=model)
|
provider = GitHubCopilotProvider(default_model=model)
|
||||||
|
elif backend == "xai_oauth":
|
||||||
|
from nanobot.providers.xai_oauth_provider import XaiOAuthProvider
|
||||||
|
|
||||||
|
provider = XaiOAuthProvider(default_model=model, config=p)
|
||||||
elif backend == "anthropic":
|
elif backend == "anthropic":
|
||||||
from nanobot.providers.anthropic_provider import AnthropicProvider
|
from nanobot.providers.anthropic_provider import AnthropicProvider
|
||||||
|
|
||||||
|
|||||||
@@ -207,8 +207,9 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
|||||||
|
|
||||||
async def _refresh_client_api_key(self) -> str:
|
async def _refresh_client_api_key(self) -> str:
|
||||||
token = await self._get_copilot_access_token()
|
token = await self._get_copilot_access_token()
|
||||||
|
client = await self._ensure_client()
|
||||||
self.api_key = token
|
self.api_key = token
|
||||||
self._client.api_key = token
|
client.api_key = token
|
||||||
return token
|
return token
|
||||||
|
|
||||||
async def chat(
|
async def chat(
|
||||||
@@ -243,6 +244,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
|||||||
tool_choice: str | dict[str, object] | None = None,
|
tool_choice: str | dict[str, object] | None = None,
|
||||||
on_content_delta: Callable[[str], None] | None = None,
|
on_content_delta: Callable[[str], None] | None = None,
|
||||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, object]], Awaitable[None]] | None = None,
|
||||||
):
|
):
|
||||||
await self._refresh_client_api_key()
|
await self._refresh_client_api_key()
|
||||||
return await super().chat_stream(
|
return await super().chat_stream(
|
||||||
@@ -255,4 +257,5 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
|||||||
tool_choice=tool_choice,
|
tool_choice=tool_choice,
|
||||||
on_content_delta=on_content_delta,
|
on_content_delta=on_content_delta,
|
||||||
on_thinking_delta=on_thinking_delta,
|
on_thinking_delta=on_thinking_delta,
|
||||||
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -3,11 +3,14 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import base64
|
import base64
|
||||||
|
import binascii
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.providers.registry import find_by_name
|
from nanobot.providers.registry import find_by_name
|
||||||
from nanobot.utils.helpers import detect_image_mime
|
from nanobot.utils.helpers import detect_image_mime
|
||||||
@@ -26,6 +29,8 @@ _AIHUBMIX_ASPECT_RATIO_SIZES = {
|
|||||||
"4:3": "1536x1024",
|
"4:3": "1536x1024",
|
||||||
"16:9": "1536x1024",
|
"16:9": "1536x1024",
|
||||||
}
|
}
|
||||||
|
_GEMINI_DEFAULT_TIMEOUT_S = 120.0
|
||||||
|
_GEMINI_IMAGEN_ASPECT_RATIOS = {"1:1", "9:16", "16:9", "3:4", "4:3"}
|
||||||
|
|
||||||
|
|
||||||
class ImageGenerationError(RuntimeError):
|
class ImageGenerationError(RuntimeError):
|
||||||
@@ -41,28 +46,38 @@ class GeneratedImageResponse:
|
|||||||
raw: dict[str, Any]
|
raw: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
def _provider_base_url(provider: str, api_base: str | None, fallback: str) -> str:
|
def _read_image_b64(path: str | Path) -> tuple[str, str]:
|
||||||
if api_base:
|
"""Return ``(mime, base64)`` for the image at ``path``."""
|
||||||
return api_base.rstrip("/")
|
|
||||||
spec = find_by_name(provider)
|
|
||||||
if spec and spec.default_api_base:
|
|
||||||
return spec.default_api_base.rstrip("/")
|
|
||||||
return fallback
|
|
||||||
|
|
||||||
|
|
||||||
def image_path_to_data_url(path: str | Path) -> str:
|
|
||||||
"""Convert a local image path to an image data URL."""
|
|
||||||
p = Path(path).expanduser()
|
p = Path(path).expanduser()
|
||||||
raw = p.read_bytes()
|
raw = p.read_bytes()
|
||||||
mime = detect_image_mime(raw)
|
mime = detect_image_mime(raw)
|
||||||
if mime is None:
|
if mime is None:
|
||||||
raise ImageGenerationError(f"unsupported reference image: {p}")
|
raise ImageGenerationError(f"unsupported reference image: {p}")
|
||||||
encoded = base64.b64encode(raw).decode("ascii")
|
return mime, base64.b64encode(raw).decode("ascii")
|
||||||
|
|
||||||
|
|
||||||
|
def image_path_to_data_url(path: str | Path) -> str:
|
||||||
|
"""Convert a local image path to an image data URL."""
|
||||||
|
mime, encoded = _read_image_b64(path)
|
||||||
return f"data:{mime};base64,{encoded}"
|
return f"data:{mime};base64,{encoded}"
|
||||||
|
|
||||||
|
|
||||||
def _b64_png_data_url(value: str) -> str:
|
def image_path_to_inline_data(path: str | Path) -> dict[str, str]:
|
||||||
return f"data:image/png;base64,{value}"
|
"""Convert a local image path to a Gemini ``inlineData`` payload dict."""
|
||||||
|
mime, encoded = _read_image_b64(path)
|
||||||
|
return {"mimeType": mime, "data": encoded}
|
||||||
|
|
||||||
|
|
||||||
|
def _b64_image_data_url(value: str) -> str:
|
||||||
|
encoded = "".join(value.split())
|
||||||
|
try:
|
||||||
|
raw = base64.b64decode(encoded, validate=True)
|
||||||
|
except binascii.Error as exc:
|
||||||
|
raise ImageGenerationError("generated image payload was not valid base64") from exc
|
||||||
|
mime = detect_image_mime(raw)
|
||||||
|
if mime is None:
|
||||||
|
raise ImageGenerationError("generated image payload was not a supported image")
|
||||||
|
return f"data:{mime};base64,{encoded}"
|
||||||
|
|
||||||
|
|
||||||
def _aihubmix_size(aspect_ratio: str | None, image_size: str | None) -> str:
|
def _aihubmix_size(aspect_ratio: str | None, image_size: str | None) -> str:
|
||||||
@@ -106,8 +121,49 @@ async def _download_image_data_url(
|
|||||||
return f"data:{mime};base64,{encoded}"
|
return f"data:{mime};base64,{encoded}"
|
||||||
|
|
||||||
|
|
||||||
class OpenRouterImageGenerationClient:
|
# ---------------------------------------------------------------------------
|
||||||
"""Small async client for OpenRouter Chat Completions image generation."""
|
# Registry
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
_IMAGE_GEN_PROVIDERS: dict[str, type[ImageGenerationProvider]] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def register_image_gen_provider(cls: type[ImageGenerationProvider]) -> None:
|
||||||
|
name = cls.provider_name
|
||||||
|
if not name:
|
||||||
|
raise ValueError(f"{cls.__name__} must set provider_name")
|
||||||
|
_IMAGE_GEN_PROVIDERS[name] = cls
|
||||||
|
|
||||||
|
|
||||||
|
def get_image_gen_provider(name: str) -> type[ImageGenerationProvider] | None:
|
||||||
|
return _IMAGE_GEN_PROVIDERS.get(name)
|
||||||
|
|
||||||
|
|
||||||
|
def image_gen_provider_names() -> tuple[str, ...]:
|
||||||
|
"""Return registered image generation provider names in registry order."""
|
||||||
|
return tuple(_IMAGE_GEN_PROVIDERS)
|
||||||
|
|
||||||
|
|
||||||
|
def image_gen_provider_configs(config: Any) -> dict[str, Any]:
|
||||||
|
providers_cfg = config.providers
|
||||||
|
return {
|
||||||
|
name: pc
|
||||||
|
for name in _IMAGE_GEN_PROVIDERS
|
||||||
|
if (pc := getattr(providers_cfg, name, None)) is not None
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Base class
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class ImageGenerationProvider(ABC):
|
||||||
|
"""Base class for image generation provider clients."""
|
||||||
|
|
||||||
|
provider_name: str = ""
|
||||||
|
missing_key_message: str = ""
|
||||||
|
default_timeout: float = _DEFAULT_TIMEOUT_S
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -116,20 +172,71 @@ class OpenRouterImageGenerationClient:
|
|||||||
api_base: str | None = None,
|
api_base: str | None = None,
|
||||||
extra_headers: dict[str, str] | None = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
extra_body: dict[str, Any] | None = None,
|
extra_body: dict[str, Any] | None = None,
|
||||||
timeout: float = _DEFAULT_TIMEOUT_S,
|
timeout: float | None = None,
|
||||||
client: httpx.AsyncClient | None = None,
|
client: httpx.AsyncClient | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.api_key = api_key
|
self.api_key = api_key
|
||||||
self.api_base = _provider_base_url(
|
self.api_base = self._resolve_base_url(api_base)
|
||||||
"openrouter",
|
|
||||||
api_base,
|
|
||||||
"https://openrouter.ai/api/v1",
|
|
||||||
)
|
|
||||||
self.extra_headers = extra_headers or {}
|
self.extra_headers = extra_headers or {}
|
||||||
self.extra_body = extra_body or {}
|
self.extra_body = extra_body or {}
|
||||||
self.timeout = timeout
|
self.timeout = timeout if timeout is not None else self.default_timeout
|
||||||
self._client = client
|
self._client = client
|
||||||
|
|
||||||
|
def _resolve_base_url(self, api_base: str | None) -> str:
|
||||||
|
if api_base:
|
||||||
|
return api_base.rstrip("/")
|
||||||
|
spec = find_by_name(self.provider_name)
|
||||||
|
if spec and spec.default_api_base:
|
||||||
|
return spec.default_api_base.rstrip("/")
|
||||||
|
return self._default_base_url()
|
||||||
|
|
||||||
|
def _default_base_url(self) -> str:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
reference_images: list[str] | None = None,
|
||||||
|
aspect_ratio: str | None = None,
|
||||||
|
image_size: str | None = None,
|
||||||
|
) -> GeneratedImageResponse: ...
|
||||||
|
|
||||||
|
def _require_images(self, images: list[str], data: dict[str, Any]) -> None:
|
||||||
|
if images:
|
||||||
|
return
|
||||||
|
provider_error = data.get("error") if isinstance(data, dict) else None
|
||||||
|
label = self.provider_name
|
||||||
|
if provider_error:
|
||||||
|
raise ImageGenerationError(f"{label} returned no images: {provider_error}")
|
||||||
|
raise ImageGenerationError(f"{label} returned no images for this request")
|
||||||
|
|
||||||
|
async def _http_post(
|
||||||
|
self,
|
||||||
|
url: str,
|
||||||
|
*,
|
||||||
|
headers: dict[str, str],
|
||||||
|
body: dict[str, Any],
|
||||||
|
) -> httpx.Response:
|
||||||
|
if self._client is not None:
|
||||||
|
return await self._client.post(url, headers=headers, json=body)
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout) as c:
|
||||||
|
return await c.post(url, headers=headers, json=body)
|
||||||
|
|
||||||
|
|
||||||
|
class OpenRouterImageGenerationClient(ImageGenerationProvider):
|
||||||
|
"""Small async client for OpenRouter Chat Completions image generation."""
|
||||||
|
|
||||||
|
provider_name = "openrouter"
|
||||||
|
missing_key_message = (
|
||||||
|
"OpenRouter API key is not configured. Set providers.openrouter.apiKey."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _default_base_url(self) -> str:
|
||||||
|
return "https://openrouter.ai/api/v1"
|
||||||
|
|
||||||
async def generate(
|
async def generate(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -140,9 +247,7 @@ class OpenRouterImageGenerationClient:
|
|||||||
image_size: str | None = None,
|
image_size: str | None = None,
|
||||||
) -> GeneratedImageResponse:
|
) -> GeneratedImageResponse:
|
||||||
if not self.api_key:
|
if not self.api_key:
|
||||||
raise ImageGenerationError(
|
raise ImageGenerationError(self.missing_key_message)
|
||||||
"OpenRouter API key is not configured. Set providers.openrouter.apiKey."
|
|
||||||
)
|
|
||||||
|
|
||||||
content: str | list[dict[str, Any]]
|
content: str | list[dict[str, Any]]
|
||||||
references = list(reference_images or [])
|
references = list(reference_images or [])
|
||||||
@@ -178,12 +283,7 @@ class OpenRouterImageGenerationClient:
|
|||||||
**self.extra_headers,
|
**self.extra_headers,
|
||||||
}
|
}
|
||||||
url = f"{self.api_base}/chat/completions"
|
url = f"{self.api_base}/chat/completions"
|
||||||
|
response = await self._http_post(url, headers=headers, body=body)
|
||||||
if self._client is not None:
|
|
||||||
response = await self._client.post(url, headers=headers, json=body)
|
|
||||||
else:
|
|
||||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
|
||||||
response = await client.post(url, headers=headers, json=body)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
@@ -208,11 +308,7 @@ class OpenRouterImageGenerationClient:
|
|||||||
if isinstance(url_value, str) and url_value.startswith("data:image/"):
|
if isinstance(url_value, str) and url_value.startswith("data:image/"):
|
||||||
images.append(url_value)
|
images.append(url_value)
|
||||||
|
|
||||||
if not images:
|
self._require_images(images, data)
|
||||||
provider_error = data.get("error") if isinstance(data, dict) else None
|
|
||||||
if provider_error:
|
|
||||||
raise ImageGenerationError(f"OpenRouter returned no images: {provider_error}")
|
|
||||||
raise ImageGenerationError("OpenRouter returned no images for this request")
|
|
||||||
|
|
||||||
return GeneratedImageResponse(
|
return GeneratedImageResponse(
|
||||||
images=images,
|
images=images,
|
||||||
@@ -221,29 +317,17 @@ class OpenRouterImageGenerationClient:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class AIHubMixImageGenerationClient:
|
class AIHubMixImageGenerationClient(ImageGenerationProvider):
|
||||||
"""Small async client for AIHubMix unified image generation."""
|
"""Small async client for AIHubMix unified image generation."""
|
||||||
|
|
||||||
def __init__(
|
provider_name = "aihubmix"
|
||||||
self,
|
missing_key_message = (
|
||||||
*,
|
"AIHubMix API key is not configured. Set providers.aihubmix.apiKey."
|
||||||
api_key: str | None,
|
)
|
||||||
api_base: str | None = None,
|
default_timeout = _AIHUBMIX_TIMEOUT_S
|
||||||
extra_headers: dict[str, str] | None = None,
|
|
||||||
extra_body: dict[str, Any] | None = None,
|
def _default_base_url(self) -> str:
|
||||||
timeout: float = _AIHUBMIX_TIMEOUT_S,
|
return "https://aihubmix.com/v1"
|
||||||
client: httpx.AsyncClient | None = None,
|
|
||||||
) -> None:
|
|
||||||
self.api_key = api_key
|
|
||||||
self.api_base = _provider_base_url(
|
|
||||||
"aihubmix",
|
|
||||||
api_base,
|
|
||||||
"https://aihubmix.com/v1",
|
|
||||||
)
|
|
||||||
self.extra_headers = extra_headers or {}
|
|
||||||
self.extra_body = extra_body or {}
|
|
||||||
self.timeout = timeout
|
|
||||||
self._client = client
|
|
||||||
|
|
||||||
async def generate(
|
async def generate(
|
||||||
self,
|
self,
|
||||||
@@ -255,9 +339,7 @@ class AIHubMixImageGenerationClient:
|
|||||||
image_size: str | None = None,
|
image_size: str | None = None,
|
||||||
) -> GeneratedImageResponse:
|
) -> GeneratedImageResponse:
|
||||||
if not self.api_key:
|
if not self.api_key:
|
||||||
raise ImageGenerationError(
|
raise ImageGenerationError(self.missing_key_message)
|
||||||
"AIHubMix API key is not configured. Set providers.aihubmix.apiKey."
|
|
||||||
)
|
|
||||||
|
|
||||||
refs = list(reference_images or [])
|
refs = list(reference_images or [])
|
||||||
headers = {
|
headers = {
|
||||||
@@ -266,16 +348,8 @@ class AIHubMixImageGenerationClient:
|
|||||||
}
|
}
|
||||||
size = _aihubmix_size(aspect_ratio, image_size)
|
size = _aihubmix_size(aspect_ratio, image_size)
|
||||||
|
|
||||||
if self._client is not None:
|
client = self._client or httpx.AsyncClient(timeout=self.timeout)
|
||||||
return await self._generate_with_client(
|
try:
|
||||||
self._client,
|
|
||||||
prompt=prompt,
|
|
||||||
model=model,
|
|
||||||
reference_images=refs,
|
|
||||||
size=size,
|
|
||||||
headers=headers,
|
|
||||||
)
|
|
||||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
|
||||||
return await self._generate_with_client(
|
return await self._generate_with_client(
|
||||||
client,
|
client,
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
@@ -284,6 +358,9 @@ class AIHubMixImageGenerationClient:
|
|||||||
size=size,
|
size=size,
|
||||||
headers=headers,
|
headers=headers,
|
||||||
)
|
)
|
||||||
|
finally:
|
||||||
|
if self._client is None:
|
||||||
|
await client.aclose()
|
||||||
|
|
||||||
async def _generate_with_client(
|
async def _generate_with_client(
|
||||||
self,
|
self,
|
||||||
@@ -332,15 +409,182 @@ class AIHubMixImageGenerationClient:
|
|||||||
payload = response.json()
|
payload = response.json()
|
||||||
images = await _aihubmix_images_from_payload(client, payload)
|
images = await _aihubmix_images_from_payload(client, payload)
|
||||||
|
|
||||||
if not images:
|
self._require_images(images, payload)
|
||||||
provider_error = payload.get("error") if isinstance(payload, dict) else None
|
|
||||||
if provider_error:
|
|
||||||
raise ImageGenerationError(f"AIHubMix returned no images: {provider_error}")
|
|
||||||
raise ImageGenerationError("AIHubMix returned no images for this request")
|
|
||||||
|
|
||||||
return GeneratedImageResponse(images=images, content="", raw=payload)
|
return GeneratedImageResponse(images=images, content="", raw=payload)
|
||||||
|
|
||||||
|
|
||||||
|
def _http_error_detail(response: httpx.Response) -> str:
|
||||||
|
"""Extract a readable error message from an HTTP error response."""
|
||||||
|
try:
|
||||||
|
data = response.json()
|
||||||
|
if isinstance(data, dict):
|
||||||
|
err = data.get("error")
|
||||||
|
if isinstance(err, dict):
|
||||||
|
return err.get("message") or str(err)
|
||||||
|
if err:
|
||||||
|
return str(err)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return response.text[:500] or "<empty response body>"
|
||||||
|
|
||||||
|
|
||||||
|
class GeminiImageGenerationClient(ImageGenerationProvider):
|
||||||
|
"""Async client for Gemini/Imagen image generation via the Generative Language API."""
|
||||||
|
|
||||||
|
provider_name = "gemini"
|
||||||
|
missing_key_message = (
|
||||||
|
"Gemini API key is not configured. Set providers.gemini.apiKey."
|
||||||
|
)
|
||||||
|
default_timeout = _GEMINI_DEFAULT_TIMEOUT_S
|
||||||
|
|
||||||
|
def _default_base_url(self) -> str:
|
||||||
|
return "https://generativelanguage.googleapis.com/v1beta"
|
||||||
|
|
||||||
|
def _resolve_base_url(self, api_base: str | None) -> str:
|
||||||
|
# The Gemini provider's registry default_api_base is the OpenAI-compat
|
||||||
|
# shim (.../v1beta/openai/), which has no image endpoints.
|
||||||
|
# Skip the registry lookup and use the native API base directly.
|
||||||
|
if api_base:
|
||||||
|
return api_base.rstrip("/")
|
||||||
|
return self._default_base_url()
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
reference_images: list[str] | None = None,
|
||||||
|
aspect_ratio: str | None = None,
|
||||||
|
image_size: str | None = None,
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
if not self.api_key:
|
||||||
|
raise ImageGenerationError(self.missing_key_message)
|
||||||
|
if "imagen" in model.lower():
|
||||||
|
if reference_images:
|
||||||
|
logger.warning(
|
||||||
|
"Imagen models do not support reference images; "
|
||||||
|
"ignoring {} reference image(s) for {}",
|
||||||
|
len(reference_images),
|
||||||
|
model,
|
||||||
|
)
|
||||||
|
return await self._generate_imagen(
|
||||||
|
prompt=prompt, model=model, aspect_ratio=aspect_ratio
|
||||||
|
)
|
||||||
|
return await self._generate_gemini_flash(
|
||||||
|
prompt=prompt, model=model, reference_images=reference_images or []
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _generate_imagen(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
aspect_ratio: str | None,
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
parameters: dict[str, Any] = {"sampleCount": 1}
|
||||||
|
if aspect_ratio in _GEMINI_IMAGEN_ASPECT_RATIOS:
|
||||||
|
parameters["aspectRatio"] = aspect_ratio
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"instances": [{"prompt": prompt}],
|
||||||
|
"parameters": parameters,
|
||||||
|
}
|
||||||
|
body.update(self.extra_body)
|
||||||
|
|
||||||
|
url = f"{self.api_base}/models/{model}:predict"
|
||||||
|
headers = {
|
||||||
|
"x-goog-api-key": self.api_key or "",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
**self.extra_headers,
|
||||||
|
}
|
||||||
|
response = await self._http_post(url, headers=headers, body=body)
|
||||||
|
|
||||||
|
try:
|
||||||
|
response.raise_for_status()
|
||||||
|
except httpx.HTTPStatusError as exc:
|
||||||
|
detail = _http_error_detail(response)
|
||||||
|
logger.error("Gemini Imagen generation failed (HTTP {}): {}", response.status_code, detail)
|
||||||
|
raise ImageGenerationError(
|
||||||
|
f"Gemini Imagen generation failed (HTTP {response.status_code}): {detail}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
data = response.json()
|
||||||
|
images: list[str] = []
|
||||||
|
for prediction in data.get("predictions") or []:
|
||||||
|
if not isinstance(prediction, dict):
|
||||||
|
continue
|
||||||
|
b64 = prediction.get("bytesBase64Encoded")
|
||||||
|
mime = prediction.get("mimeType", "image/png")
|
||||||
|
if isinstance(b64, str) and b64:
|
||||||
|
images.append(f"data:{mime};base64,{b64}")
|
||||||
|
|
||||||
|
self._require_images(images, data)
|
||||||
|
|
||||||
|
return GeneratedImageResponse(images=images, content="", raw=data)
|
||||||
|
|
||||||
|
async def _generate_gemini_flash(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
reference_images: list[str],
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
parts: list[dict[str, Any]] = [
|
||||||
|
{"inlineData": image_path_to_inline_data(path)} for path in reference_images
|
||||||
|
]
|
||||||
|
parts.append({"text": prompt})
|
||||||
|
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"contents": [{"role": "user", "parts": parts}],
|
||||||
|
"generationConfig": {"responseModalities": ["TEXT", "IMAGE"]},
|
||||||
|
}
|
||||||
|
body.update(self.extra_body)
|
||||||
|
|
||||||
|
url = f"{self.api_base}/models/{model}:generateContent"
|
||||||
|
headers = {
|
||||||
|
"x-goog-api-key": self.api_key or "",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
**self.extra_headers,
|
||||||
|
}
|
||||||
|
response = await self._http_post(url, headers=headers, body=body)
|
||||||
|
|
||||||
|
try:
|
||||||
|
response.raise_for_status()
|
||||||
|
except httpx.HTTPStatusError as exc:
|
||||||
|
detail = _http_error_detail(response)
|
||||||
|
logger.error("Gemini image generation failed (HTTP {}): {}", response.status_code, detail)
|
||||||
|
raise ImageGenerationError(
|
||||||
|
f"Gemini image generation failed (HTTP {response.status_code}): {detail}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
data = response.json()
|
||||||
|
images: list[str] = []
|
||||||
|
text_parts: list[str] = []
|
||||||
|
for candidate in data.get("candidates") or []:
|
||||||
|
if not isinstance(candidate, dict):
|
||||||
|
continue
|
||||||
|
content = candidate.get("content") or {}
|
||||||
|
for part in content.get("parts") or []:
|
||||||
|
if not isinstance(part, dict):
|
||||||
|
continue
|
||||||
|
if "text" in part:
|
||||||
|
text_parts.append(part["text"])
|
||||||
|
inline = part.get("inlineData")
|
||||||
|
if isinstance(inline, dict):
|
||||||
|
mime = inline.get("mimeType", "image/png")
|
||||||
|
b64 = inline.get("data", "")
|
||||||
|
if b64:
|
||||||
|
images.append(f"data:{mime};base64,{b64}")
|
||||||
|
|
||||||
|
self._require_images(images, data)
|
||||||
|
|
||||||
|
return GeneratedImageResponse(
|
||||||
|
images=images,
|
||||||
|
content="\n".join(t for t in text_parts if t).strip(),
|
||||||
|
raw=data,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def _aihubmix_images_from_payload(
|
async def _aihubmix_images_from_payload(
|
||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
payload: dict[str, Any],
|
payload: dict[str, Any],
|
||||||
@@ -368,13 +612,13 @@ async def _aihubmix_images_from_payload(
|
|||||||
|
|
||||||
b64_json = value.get("b64_json")
|
b64_json = value.get("b64_json")
|
||||||
if isinstance(b64_json, str) and b64_json:
|
if isinstance(b64_json, str) and b64_json:
|
||||||
images.append(_b64_png_data_url(b64_json))
|
images.append(_b64_image_data_url(b64_json))
|
||||||
elif b64_json is not None:
|
elif b64_json is not None:
|
||||||
await collect(b64_json)
|
await collect(b64_json)
|
||||||
|
|
||||||
bytes_base64 = value.get("bytesBase64") or value.get("bytes_base64") or value.get("base64")
|
bytes_base64 = value.get("bytesBase64") or value.get("bytes_base64") or value.get("base64")
|
||||||
if isinstance(bytes_base64, str) and bytes_base64:
|
if isinstance(bytes_base64, str) and bytes_base64:
|
||||||
images.append(_b64_png_data_url(bytes_base64))
|
images.append(_b64_image_data_url(bytes_base64))
|
||||||
|
|
||||||
image_url = value.get("image_url") or value.get("imageUrl")
|
image_url = value.get("image_url") or value.get("imageUrl")
|
||||||
if isinstance(image_url, dict):
|
if isinstance(image_url, dict):
|
||||||
@@ -393,3 +637,254 @@ async def _aihubmix_images_from_payload(
|
|||||||
for candidate in candidates:
|
for candidate in candidates:
|
||||||
await collect(candidate)
|
await collect(candidate)
|
||||||
return images
|
return images
|
||||||
|
|
||||||
|
|
||||||
|
_MINIMAX_TIMEOUT_S = 300.0
|
||||||
|
|
||||||
|
_MINIMAX_ASPECT_RATIO_SIZES = {
|
||||||
|
"1:1": "1:1",
|
||||||
|
"16:9": "16:9",
|
||||||
|
"4:3": "4:3",
|
||||||
|
"3:2": "3:2",
|
||||||
|
"2:3": "2:3",
|
||||||
|
"3:4": "3:4",
|
||||||
|
"9:16": "9:16",
|
||||||
|
"21:9": "21:9",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class MiniMaxImageGenerationClient(ImageGenerationProvider):
|
||||||
|
"""Async client for MiniMax image generation API."""
|
||||||
|
|
||||||
|
provider_name = "minimax"
|
||||||
|
missing_key_message = (
|
||||||
|
"MiniMax API key is not configured. Set providers.minimax.apiKey."
|
||||||
|
)
|
||||||
|
default_timeout = _MINIMAX_TIMEOUT_S
|
||||||
|
|
||||||
|
def _default_base_url(self) -> str:
|
||||||
|
return "https://api.minimaxi.com/v1"
|
||||||
|
|
||||||
|
def _resolve_aspect_ratio(self, aspect_ratio: str | None) -> str:
|
||||||
|
if aspect_ratio and aspect_ratio in _MINIMAX_ASPECT_RATIO_SIZES:
|
||||||
|
return _MINIMAX_ASPECT_RATIO_SIZES[aspect_ratio]
|
||||||
|
return "1:1"
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
reference_images: list[str] | None = None,
|
||||||
|
aspect_ratio: str | None = None,
|
||||||
|
image_size: str | None = None,
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
if not self.api_key:
|
||||||
|
raise ImageGenerationError(self.missing_key_message)
|
||||||
|
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
**self.extra_headers,
|
||||||
|
}
|
||||||
|
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"model": model,
|
||||||
|
"prompt": prompt,
|
||||||
|
"response_format": "base64",
|
||||||
|
}
|
||||||
|
|
||||||
|
resolved_ratio = self._resolve_aspect_ratio(aspect_ratio)
|
||||||
|
body["aspect_ratio"] = resolved_ratio
|
||||||
|
|
||||||
|
refs = list(reference_images or [])
|
||||||
|
if refs:
|
||||||
|
image_refs = [image_path_to_data_url(path) for path in refs]
|
||||||
|
body["subject_reference"] = [
|
||||||
|
{"type": "character", "image_file": ref} for ref in image_refs
|
||||||
|
]
|
||||||
|
|
||||||
|
body.update(self.extra_body)
|
||||||
|
|
||||||
|
client = self._client or httpx.AsyncClient(timeout=self.timeout)
|
||||||
|
try:
|
||||||
|
return await self._generate_with_client(client, body, headers)
|
||||||
|
finally:
|
||||||
|
if self._client is None:
|
||||||
|
await client.aclose()
|
||||||
|
|
||||||
|
async def _generate_with_client(
|
||||||
|
self,
|
||||||
|
client: httpx.AsyncClient,
|
||||||
|
body: dict[str, Any],
|
||||||
|
headers: dict[str, str],
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
url = f"{self.api_base}/image_generation"
|
||||||
|
try:
|
||||||
|
response = await client.post(url, headers=headers, json=body)
|
||||||
|
except httpx.TimeoutException as exc:
|
||||||
|
raise ImageGenerationError("MiniMax image generation timed out") from exc
|
||||||
|
except httpx.RequestError as exc:
|
||||||
|
raise ImageGenerationError(f"MiniMax image generation request failed: {exc}") from exc
|
||||||
|
|
||||||
|
try:
|
||||||
|
response.raise_for_status()
|
||||||
|
except httpx.HTTPStatusError as exc:
|
||||||
|
detail = response.text[:500]
|
||||||
|
raise ImageGenerationError(f"MiniMax image generation failed: {detail}") from exc
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
images = _minimax_images_from_payload(payload)
|
||||||
|
|
||||||
|
self._require_images(images, payload)
|
||||||
|
|
||||||
|
return GeneratedImageResponse(images=images, content="", raw=payload)
|
||||||
|
|
||||||
|
|
||||||
|
def _minimax_images_from_payload(payload: dict[str, Any]) -> list[str]:
|
||||||
|
"""Extract base64 images from MiniMax API response.
|
||||||
|
|
||||||
|
MiniMax returns images in ``data.image_base64`` (list of base64 strings).
|
||||||
|
"""
|
||||||
|
images: list[str] = []
|
||||||
|
data = payload.get("data")
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
return images
|
||||||
|
for b64 in data.get("image_base64") or []:
|
||||||
|
if isinstance(b64, str) and b64:
|
||||||
|
images.append(_b64_image_data_url(b64))
|
||||||
|
return images
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# StepFun (阶跃星辰) image generation
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
_STEPFUN_ASPECT_RATIO_SIZES = {
|
||||||
|
"1:1": "1024x1024",
|
||||||
|
"16:9": "1280x800",
|
||||||
|
"9:16": "800x1280",
|
||||||
|
"3:4": "768x1360",
|
||||||
|
"4:3": "1360x768",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class StepFunImageGenerationClient(ImageGenerationProvider):
|
||||||
|
"""Async client for StepFun (阶跃星辰) image generation.
|
||||||
|
|
||||||
|
Supports:
|
||||||
|
- Text-to-image via step-image-edit-2 (default model)
|
||||||
|
- Reference-image-guided generation via style_reference (step-1x-medium)
|
||||||
|
"""
|
||||||
|
|
||||||
|
provider_name = "stepfun"
|
||||||
|
missing_key_message = (
|
||||||
|
"StepFun API key is not configured. Set providers.stepfun.apiKey."
|
||||||
|
)
|
||||||
|
default_timeout = 120.0
|
||||||
|
|
||||||
|
def _default_base_url(self) -> str:
|
||||||
|
return "https://api.stepfun.com/v1"
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
reference_images: list[str] | None = None,
|
||||||
|
aspect_ratio: str | None = None,
|
||||||
|
image_size: str | None = None,
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
if not self.api_key:
|
||||||
|
raise ImageGenerationError(self.missing_key_message)
|
||||||
|
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
**self.extra_headers,
|
||||||
|
}
|
||||||
|
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"model": model,
|
||||||
|
"prompt": prompt,
|
||||||
|
"response_format": "b64_json",
|
||||||
|
"n": 1,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Map aspect ratio / image_size to StepFun size string
|
||||||
|
size = _stepfun_size(aspect_ratio, image_size)
|
||||||
|
if size:
|
||||||
|
body["size"] = size
|
||||||
|
|
||||||
|
# step-1x-medium supports style_reference for reference-image-guided generation
|
||||||
|
refs = list(reference_images or [])
|
||||||
|
if refs and "1x" in model:
|
||||||
|
body["style_reference"] = {
|
||||||
|
"source_url": image_path_to_data_url(refs[0]),
|
||||||
|
}
|
||||||
|
|
||||||
|
body.update(self.extra_body)
|
||||||
|
|
||||||
|
response = await self._http_post(
|
||||||
|
f"{self.api_base}/images/generations",
|
||||||
|
headers=headers,
|
||||||
|
body=body,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
response.raise_for_status()
|
||||||
|
except httpx.HTTPStatusError as exc:
|
||||||
|
detail = response.text[:500]
|
||||||
|
raise ImageGenerationError(
|
||||||
|
f"StepFun image generation failed: {detail}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
images = _stepfun_images_from_payload(payload)
|
||||||
|
|
||||||
|
self._require_images(images, payload)
|
||||||
|
|
||||||
|
return GeneratedImageResponse(images=images, content="", raw=payload)
|
||||||
|
|
||||||
|
|
||||||
|
def _stepfun_size(
|
||||||
|
aspect_ratio: str | None,
|
||||||
|
image_size: str | None,
|
||||||
|
) -> str:
|
||||||
|
"""Resolve aspect ratio / image_size to StepFun size string.
|
||||||
|
|
||||||
|
StepFun expects ``WIDTHxHEIGHT`` (note: width x height, not the more
|
||||||
|
common ``HxW`` order used by other providers). The accepted sizes are
|
||||||
|
``1024x1024``, ``768x1360``, ``896x1184``, ``1360x768``, ``1184x896``.
|
||||||
|
"""
|
||||||
|
if image_size and "x" in image_size.lower():
|
||||||
|
return image_size
|
||||||
|
if aspect_ratio and aspect_ratio in _STEPFUN_ASPECT_RATIO_SIZES:
|
||||||
|
return _STEPFUN_ASPECT_RATIO_SIZES[aspect_ratio]
|
||||||
|
return "1024x1024"
|
||||||
|
|
||||||
|
|
||||||
|
def _stepfun_images_from_payload(payload: dict[str, Any]) -> list[str]:
|
||||||
|
"""Extract base64 images from StepFun API response.
|
||||||
|
|
||||||
|
StepFun returns images in ``data[].b64_json`` (base64 strings).
|
||||||
|
"""
|
||||||
|
images: list[str] = []
|
||||||
|
for item in payload.get("data") or []:
|
||||||
|
if not isinstance(item, dict):
|
||||||
|
continue
|
||||||
|
b64 = item.get("b64_json")
|
||||||
|
if isinstance(b64, str) and b64:
|
||||||
|
images.append(_b64_image_data_url(b64))
|
||||||
|
return images
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Provider registration
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
register_image_gen_provider(OpenRouterImageGenerationClient)
|
||||||
|
register_image_gen_provider(AIHubMixImageGenerationClient)
|
||||||
|
register_image_gen_provider(GeminiImageGenerationClient)
|
||||||
|
register_image_gen_provider(MiniMaxImageGenerationClient)
|
||||||
|
register_image_gen_provider(StepFunImageGenerationClient)
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ class OpenAICodexProvider(LLMProvider):
|
|||||||
reasoning_effort: str | None,
|
reasoning_effort: str | None,
|
||||||
tool_choice: str | dict[str, Any] | None,
|
tool_choice: str | dict[str, Any] | None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
"""Shared request logic for both chat() and chat_stream()."""
|
"""Shared request logic for both chat() and chat_stream()."""
|
||||||
model = model or self.default_model
|
model = model or self.default_model
|
||||||
@@ -70,6 +71,7 @@ class OpenAICodexProvider(LLMProvider):
|
|||||||
content, tool_calls, finish_reason = await _request_codex(
|
content, tool_calls, finish_reason = await _request_codex(
|
||||||
DEFAULT_CODEX_URL, headers, body, verify=True,
|
DEFAULT_CODEX_URL, headers, body, verify=True,
|
||||||
on_content_delta=on_content_delta,
|
on_content_delta=on_content_delta,
|
||||||
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
if "CERTIFICATE_VERIFY_FAILED" not in str(e):
|
if "CERTIFICATE_VERIFY_FAILED" not in str(e):
|
||||||
@@ -78,6 +80,7 @@ class OpenAICodexProvider(LLMProvider):
|
|||||||
content, tool_calls, finish_reason = await _request_codex(
|
content, tool_calls, finish_reason = await _request_codex(
|
||||||
DEFAULT_CODEX_URL, headers, body, verify=False,
|
DEFAULT_CODEX_URL, headers, body, verify=False,
|
||||||
on_content_delta=on_content_delta,
|
on_content_delta=on_content_delta,
|
||||||
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
)
|
)
|
||||||
return LLMResponse(content=content, tool_calls=tool_calls, finish_reason=finish_reason)
|
return LLMResponse(content=content, tool_calls=tool_calls, finish_reason=finish_reason)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -100,9 +103,18 @@ class OpenAICodexProvider(LLMProvider):
|
|||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
_ = on_thinking_delta
|
_ = on_thinking_delta
|
||||||
return await self._call_codex(messages, tools, model, reasoning_effort, tool_choice, on_content_delta)
|
return await self._call_codex(
|
||||||
|
messages,
|
||||||
|
tools,
|
||||||
|
model,
|
||||||
|
reasoning_effort,
|
||||||
|
tool_choice,
|
||||||
|
on_content_delta,
|
||||||
|
on_tool_call_delta,
|
||||||
|
)
|
||||||
|
|
||||||
def get_default_model(self) -> str:
|
def get_default_model(self) -> str:
|
||||||
return self.default_model
|
return self.default_model
|
||||||
@@ -138,6 +150,7 @@ async def _request_codex(
|
|||||||
body: dict[str, Any],
|
body: dict[str, Any],
|
||||||
verify: bool,
|
verify: bool,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> tuple[str, list[ToolCallRequest], str]:
|
) -> tuple[str, list[ToolCallRequest], str]:
|
||||||
async with httpx.AsyncClient(timeout=60.0, verify=verify) as client:
|
async with httpx.AsyncClient(timeout=60.0, verify=verify) as client:
|
||||||
async with client.stream("POST", url, headers=headers, json=body) as response:
|
async with client.stream("POST", url, headers=headers, json=body) as response:
|
||||||
@@ -148,7 +161,7 @@ async def _request_codex(
|
|||||||
_friendly_error(response.status_code, text.decode("utf-8", "ignore")),
|
_friendly_error(response.status_code, text.decode("utf-8", "ignore")),
|
||||||
retry_after=retry_after,
|
retry_after=retry_after,
|
||||||
)
|
)
|
||||||
return await consume_sse(response, on_content_delta)
|
return await consume_sse(response, on_content_delta, on_tool_call_delta)
|
||||||
|
|
||||||
|
|
||||||
def _prompt_cache_key(messages: list[dict[str, Any]]) -> str:
|
def _prompt_cache_key(messages: list[dict[str, Any]]) -> str:
|
||||||
|
|||||||
@@ -16,20 +16,9 @@ from ipaddress import ip_address
|
|||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
import httpx
|
|
||||||
import json_repair
|
import json_repair
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
if os.environ.get("LANGFUSE_SECRET_KEY") and importlib.util.find_spec("langfuse"):
|
|
||||||
from langfuse.openai import AsyncOpenAI
|
|
||||||
else:
|
|
||||||
if os.environ.get("LANGFUSE_SECRET_KEY"):
|
|
||||||
logger.warning(
|
|
||||||
"LANGFUSE_SECRET_KEY is set but langfuse is not installed; "
|
|
||||||
"install with `pip install langfuse` to enable tracing"
|
|
||||||
)
|
|
||||||
from openai import AsyncOpenAI
|
|
||||||
|
|
||||||
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
||||||
from nanobot.providers.openai_responses import (
|
from nanobot.providers.openai_responses import (
|
||||||
consume_sdk_stream,
|
consume_sdk_stream,
|
||||||
@@ -39,8 +28,15 @@ from nanobot.providers.openai_responses import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from openai import AsyncOpenAI as AsyncOpenAIType
|
||||||
|
|
||||||
from nanobot.providers.registry import ProviderSpec
|
from nanobot.providers.registry import ProviderSpec
|
||||||
|
|
||||||
|
# Module-level placeholder — set lazily by _ensure_client on first real
|
||||||
|
# use, or replaced by tests via ``patch(...)``. Kept as a plain name so
|
||||||
|
# that ``unittest.mock.patch`` can find and replace it.
|
||||||
|
AsyncOpenAI: Any = None
|
||||||
|
|
||||||
_ALLOWED_MSG_KEYS = frozenset({
|
_ALLOWED_MSG_KEYS = frozenset({
|
||||||
"role", "content", "tool_calls", "tool_call_id", "name",
|
"role", "content", "tool_calls", "tool_call_id", "name",
|
||||||
"reasoning_content", "extra_content",
|
"reasoning_content", "extra_content",
|
||||||
@@ -302,43 +298,76 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
|
|
||||||
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
|
||||||
default_headers = {"x-session-affinity": uuid.uuid4().hex}
|
self._default_headers = {"x-session-affinity": uuid.uuid4().hex}
|
||||||
if _uses_openrouter_attribution(spec, effective_base):
|
if _uses_openrouter_attribution(spec, effective_base):
|
||||||
default_headers.update(_DEFAULT_OPENROUTER_HEADERS)
|
self._default_headers.update(_DEFAULT_OPENROUTER_HEADERS)
|
||||||
if extra_headers:
|
if extra_headers:
|
||||||
default_headers.update(extra_headers)
|
self._default_headers.update(extra_headers)
|
||||||
|
self._api_key_for_client = api_key or "no-key"
|
||||||
|
self._is_local = _is_local_endpoint(spec, effective_base)
|
||||||
|
|
||||||
# Local model servers (Ollama, llama.cpp, vLLM) often close idle
|
# Lazy-init: the OpenAI client and its httpx transport are expensive
|
||||||
# HTTP connections before the client-side keepalive expires. When
|
# to create (~700 ms on Windows). Defer until first use.
|
||||||
# two LLM calls happen seconds apart (e.g. heartbeat _decide then
|
self._client: AsyncOpenAIType | None = None
|
||||||
# process_direct), the second call may grab a now-dead pooled
|
self._client_lock = asyncio.Lock()
|
||||||
# connection, causing a transient APIConnectionError on every first
|
|
||||||
# attempt. Disabling keepalive for local endpoints avoids this by
|
|
||||||
# opening a fresh connection for each request, which is cheap on a
|
|
||||||
# LAN. Cloud providers benefit from keepalive, so we leave the
|
|
||||||
# default pool settings for them.
|
|
||||||
timeout_s = _openai_compat_timeout_s()
|
|
||||||
http_client: httpx.AsyncClient | None = None
|
|
||||||
if _is_local_endpoint(spec, effective_base):
|
|
||||||
http_client = httpx.AsyncClient(
|
|
||||||
limits=httpx.Limits(keepalive_expiry=0),
|
|
||||||
timeout=timeout_s,
|
|
||||||
)
|
|
||||||
|
|
||||||
self._client = AsyncOpenAI(
|
|
||||||
api_key=api_key or "no-key",
|
|
||||||
base_url=effective_base,
|
|
||||||
default_headers=default_headers,
|
|
||||||
max_retries=0,
|
|
||||||
timeout=timeout_s,
|
|
||||||
http_client=http_client,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Responses API circuit breaker: skip after repeated failures,
|
# Responses API circuit breaker: skip after repeated failures,
|
||||||
# probe again after _RESPONSES_PROBE_INTERVAL_S seconds.
|
# probe again after _RESPONSES_PROBE_INTERVAL_S seconds.
|
||||||
self._responses_failures: dict[str, int] = {}
|
self._responses_failures: dict[str, int] = {}
|
||||||
self._responses_tripped_at: dict[str, float] = {}
|
self._responses_tripped_at: dict[str, float] = {}
|
||||||
|
|
||||||
|
def _build_client(self) -> None:
|
||||||
|
"""Create the OpenAI client using the current module-level AsyncOpenAI."""
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
timeout_s = _openai_compat_timeout_s()
|
||||||
|
http_client: httpx.AsyncClient | None = None
|
||||||
|
if self._is_local:
|
||||||
|
# Local model servers (Ollama, llama.cpp, vLLM) often close idle
|
||||||
|
# HTTP connections before the client-side keepalive expires. When
|
||||||
|
# two LLM calls happen seconds apart (e.g. heartbeat _decide then
|
||||||
|
# process_direct), the second call may grab a now-dead pooled
|
||||||
|
# connection, causing a transient APIConnectionError on every first
|
||||||
|
# attempt. Disabling keepalive for local endpoints avoids this by
|
||||||
|
# opening a fresh connection for each request, which is cheap on a
|
||||||
|
# LAN. Cloud providers benefit from keepalive, so we leave the
|
||||||
|
# default pool settings for them.
|
||||||
|
http_client = httpx.AsyncClient(
|
||||||
|
limits=httpx.Limits(keepalive_expiry=0),
|
||||||
|
timeout=timeout_s,
|
||||||
|
)
|
||||||
|
self._client = AsyncOpenAI(
|
||||||
|
api_key=self._api_key_for_client,
|
||||||
|
base_url=self._effective_base,
|
||||||
|
default_headers=self._default_headers,
|
||||||
|
max_retries=0,
|
||||||
|
timeout=timeout_s,
|
||||||
|
http_client=http_client,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _ensure_client(self):
|
||||||
|
"""Return the shared OpenAI client, creating it on first call."""
|
||||||
|
if self._client is not None:
|
||||||
|
return self._client
|
||||||
|
async with self._client_lock:
|
||||||
|
if self._client is not None:
|
||||||
|
return self._client
|
||||||
|
global AsyncOpenAI
|
||||||
|
if AsyncOpenAI is None:
|
||||||
|
if os.environ.get("LANGFUSE_SECRET_KEY") and importlib.util.find_spec("langfuse"):
|
||||||
|
from langfuse.openai import AsyncOpenAI as _AsyncOpenAI
|
||||||
|
else:
|
||||||
|
if os.environ.get("LANGFUSE_SECRET_KEY"):
|
||||||
|
logger.warning(
|
||||||
|
"LANGFUSE_SECRET_KEY is set but langfuse is not installed; "
|
||||||
|
"install with `pip install langfuse` to enable tracing"
|
||||||
|
)
|
||||||
|
from openai import AsyncOpenAI as _AsyncOpenAI
|
||||||
|
AsyncOpenAI = _AsyncOpenAI
|
||||||
|
|
||||||
|
self._build_client()
|
||||||
|
return self._client
|
||||||
|
|
||||||
def _setup_env(self, api_key: str, api_base: str | None) -> None:
|
def _setup_env(self, api_key: str, api_base: str | None) -> None:
|
||||||
"""Set environment variables based on provider spec."""
|
"""Set environment variables based on provider spec."""
|
||||||
spec = self._spec
|
spec = self._spec
|
||||||
@@ -999,6 +1028,21 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
if fn_prov:
|
if fn_prov:
|
||||||
buf["fn_prov"] = fn_prov
|
buf["fn_prov"] = fn_prov
|
||||||
|
|
||||||
|
def _accum_legacy_function_call(function_call: Any) -> None:
|
||||||
|
"""Accumulate legacy ``delta.function_call`` streaming chunks."""
|
||||||
|
if not function_call:
|
||||||
|
return
|
||||||
|
buf = tc_bufs.setdefault(0, {
|
||||||
|
"id": "", "name": "", "arguments": "",
|
||||||
|
"extra_content": None, "prov": None, "fn_prov": None,
|
||||||
|
})
|
||||||
|
fn_name = _get(function_call, "name")
|
||||||
|
if fn_name:
|
||||||
|
buf["name"] = str(fn_name)
|
||||||
|
fn_args = _get(function_call, "arguments")
|
||||||
|
if fn_args:
|
||||||
|
buf["arguments"] += str(fn_args)
|
||||||
|
|
||||||
for chunk in chunks:
|
for chunk in chunks:
|
||||||
if isinstance(chunk, str):
|
if isinstance(chunk, str):
|
||||||
content_parts.append(chunk)
|
content_parts.append(chunk)
|
||||||
@@ -1029,6 +1073,7 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
reasoning_parts.append(text)
|
reasoning_parts.append(text)
|
||||||
for idx, tc in enumerate(delta.get("tool_calls") or []):
|
for idx, tc in enumerate(delta.get("tool_calls") or []):
|
||||||
_accum_tc(tc, idx)
|
_accum_tc(tc, idx)
|
||||||
|
_accum_legacy_function_call(delta.get("function_call"))
|
||||||
usage = cls._extract_usage(chunk_map) or usage
|
usage = cls._extract_usage(chunk_map) or usage
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -1047,8 +1092,10 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
reasoning = getattr(delta, "reasoning", None)
|
reasoning = getattr(delta, "reasoning", None)
|
||||||
if reasoning:
|
if reasoning:
|
||||||
reasoning_parts.append(reasoning)
|
reasoning_parts.append(reasoning)
|
||||||
for tc in (delta.tool_calls or []) if delta else []:
|
for tc in (getattr(delta, "tool_calls", None) or []) if delta else []:
|
||||||
_accum_tc(tc, getattr(tc, "index", 0))
|
_accum_tc(tc, getattr(tc, "index", 0))
|
||||||
|
if delta:
|
||||||
|
_accum_legacy_function_call(getattr(delta, "function_call", None))
|
||||||
|
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
content="".join(content_parts) or None,
|
content="".join(content_parts) or None,
|
||||||
@@ -1164,6 +1211,7 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
reasoning_effort: str | None = None,
|
reasoning_effort: str | None = None,
|
||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
|
await self._ensure_client()
|
||||||
try:
|
try:
|
||||||
if self._should_use_responses_api(model, reasoning_effort):
|
if self._should_use_responses_api(model, reasoning_effort):
|
||||||
try:
|
try:
|
||||||
@@ -1203,7 +1251,9 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
|
await self._ensure_client()
|
||||||
idle_timeout_s = int(os.environ.get("NANOBOT_STREAM_IDLE_TIMEOUT_S", "90"))
|
idle_timeout_s = int(os.environ.get("NANOBOT_STREAM_IDLE_TIMEOUT_S", "90"))
|
||||||
try:
|
try:
|
||||||
if self._should_use_responses_api(model, reasoning_effort):
|
if self._should_use_responses_api(model, reasoning_effort):
|
||||||
@@ -1226,9 +1276,16 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
except StopAsyncIteration:
|
except StopAsyncIteration:
|
||||||
break
|
break
|
||||||
|
|
||||||
content, tool_calls, finish_reason, usage, reasoning_content = await consume_sdk_stream(
|
(
|
||||||
|
content,
|
||||||
|
tool_calls,
|
||||||
|
finish_reason,
|
||||||
|
usage,
|
||||||
|
reasoning_content,
|
||||||
|
) = await consume_sdk_stream(
|
||||||
_timed_stream(),
|
_timed_stream(),
|
||||||
on_content_delta,
|
on_content_delta,
|
||||||
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
)
|
)
|
||||||
self._record_responses_success(model, reasoning_effort)
|
self._record_responses_success(model, reasoning_effort)
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
@@ -1252,6 +1309,12 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
messages, tools, model, max_tokens, temperature,
|
messages, tools, model, max_tokens, temperature,
|
||||||
reasoning_effort, tool_choice,
|
reasoning_effort, tool_choice,
|
||||||
)
|
)
|
||||||
|
if self._spec and self._spec.name == "zhipu" and tools and on_tool_call_delta:
|
||||||
|
# Z.AI/GLM keeps streaming tool-call arguments behind an
|
||||||
|
# explicit provider flag. Pass it through the OpenAI SDK's
|
||||||
|
# extra_body escape hatch so the usual delta.tool_calls path
|
||||||
|
# can surface live file-edit progress.
|
||||||
|
kwargs.setdefault("extra_body", {})["tool_stream"] = True
|
||||||
kwargs["stream"] = True
|
kwargs["stream"] = True
|
||||||
kwargs["stream_options"] = {"include_usage": True}
|
kwargs["stream_options"] = {"include_usage": True}
|
||||||
stream = await self._client.chat.completions.create(**kwargs)
|
stream = await self._client.chat.completions.create(**kwargs)
|
||||||
@@ -1279,6 +1342,28 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
r_text = self._extract_text_content(reasoning)
|
r_text = self._extract_text_content(reasoning)
|
||||||
if r_text:
|
if r_text:
|
||||||
await on_thinking_delta(r_text)
|
await on_thinking_delta(r_text)
|
||||||
|
if on_tool_call_delta:
|
||||||
|
for idx, tool_delta in enumerate(
|
||||||
|
getattr(delta_obj, "tool_calls", None) or []
|
||||||
|
):
|
||||||
|
fn = _get(tool_delta, "function")
|
||||||
|
tool_index = _get(tool_delta, "index")
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": tool_index if tool_index is not None else idx,
|
||||||
|
"call_id": str(_get(tool_delta, "id") or ""),
|
||||||
|
"name": str(_get(fn, "name") or "") if fn is not None else "",
|
||||||
|
"arguments_delta": (
|
||||||
|
str(_get(fn, "arguments") or "") if fn is not None else ""
|
||||||
|
),
|
||||||
|
})
|
||||||
|
function_call = getattr(delta_obj, "function_call", None)
|
||||||
|
if function_call:
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": 0,
|
||||||
|
"call_id": "",
|
||||||
|
"name": str(_get(function_call, "name") or ""),
|
||||||
|
"arguments_delta": str(_get(function_call, "arguments") or ""),
|
||||||
|
})
|
||||||
return self._parse_chunks(chunks)
|
return self._parse_chunks(chunks)
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
|
|||||||
@@ -62,6 +62,7 @@ async def iter_sse(response: httpx.Response) -> AsyncGenerator[dict[str, Any], N
|
|||||||
async def consume_sse(
|
async def consume_sse(
|
||||||
response: httpx.Response,
|
response: httpx.Response,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> tuple[str, list[ToolCallRequest], str]:
|
) -> tuple[str, list[ToolCallRequest], str]:
|
||||||
"""Consume a Responses API SSE stream into ``(content, tool_calls, finish_reason)``."""
|
"""Consume a Responses API SSE stream into ``(content, tool_calls, finish_reason)``."""
|
||||||
content = ""
|
content = ""
|
||||||
@@ -82,6 +83,12 @@ async def consume_sse(
|
|||||||
"name": item.get("name"),
|
"name": item.get("name"),
|
||||||
"arguments": item.get("arguments") or "",
|
"arguments": item.get("arguments") or "",
|
||||||
}
|
}
|
||||||
|
if on_tool_call_delta:
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"call_id": str(call_id),
|
||||||
|
"name": str(item.get("name") or ""),
|
||||||
|
"arguments_delta": "",
|
||||||
|
})
|
||||||
elif event_type == "response.output_text.delta":
|
elif event_type == "response.output_text.delta":
|
||||||
delta_text = event.get("delta") or ""
|
delta_text = event.get("delta") or ""
|
||||||
content += delta_text
|
content += delta_text
|
||||||
@@ -90,7 +97,14 @@ async def consume_sse(
|
|||||||
elif event_type == "response.function_call_arguments.delta":
|
elif event_type == "response.function_call_arguments.delta":
|
||||||
call_id = event.get("call_id")
|
call_id = event.get("call_id")
|
||||||
if call_id and call_id in tool_call_buffers:
|
if call_id and call_id in tool_call_buffers:
|
||||||
tool_call_buffers[call_id]["arguments"] += event.get("delta") or ""
|
delta = event.get("delta") or ""
|
||||||
|
tool_call_buffers[call_id]["arguments"] += delta
|
||||||
|
if on_tool_call_delta and delta:
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"call_id": str(call_id),
|
||||||
|
"name": str(tool_call_buffers[call_id].get("name") or ""),
|
||||||
|
"arguments_delta": str(delta),
|
||||||
|
})
|
||||||
elif event_type == "response.function_call_arguments.done":
|
elif event_type == "response.function_call_arguments.done":
|
||||||
call_id = event.get("call_id")
|
call_id = event.get("call_id")
|
||||||
if call_id and call_id in tool_call_buffers:
|
if call_id and call_id in tool_call_buffers:
|
||||||
@@ -210,6 +224,7 @@ def parse_response_output(response: Any) -> LLMResponse:
|
|||||||
async def consume_sdk_stream(
|
async def consume_sdk_stream(
|
||||||
stream: Any,
|
stream: Any,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
|
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
|
||||||
"""Consume an SDK async stream from ``client.responses.create(stream=True)``."""
|
"""Consume an SDK async stream from ``client.responses.create(stream=True)``."""
|
||||||
content = ""
|
content = ""
|
||||||
@@ -232,6 +247,12 @@ async def consume_sdk_stream(
|
|||||||
"name": getattr(item, "name", None),
|
"name": getattr(item, "name", None),
|
||||||
"arguments": getattr(item, "arguments", None) or "",
|
"arguments": getattr(item, "arguments", None) or "",
|
||||||
}
|
}
|
||||||
|
if on_tool_call_delta:
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"call_id": str(call_id),
|
||||||
|
"name": str(getattr(item, "name", None) or ""),
|
||||||
|
"arguments_delta": "",
|
||||||
|
})
|
||||||
elif event_type == "response.output_text.delta":
|
elif event_type == "response.output_text.delta":
|
||||||
delta_text = getattr(event, "delta", "") or ""
|
delta_text = getattr(event, "delta", "") or ""
|
||||||
content += delta_text
|
content += delta_text
|
||||||
@@ -240,7 +261,14 @@ async def consume_sdk_stream(
|
|||||||
elif event_type == "response.function_call_arguments.delta":
|
elif event_type == "response.function_call_arguments.delta":
|
||||||
call_id = getattr(event, "call_id", None)
|
call_id = getattr(event, "call_id", None)
|
||||||
if call_id and call_id in tool_call_buffers:
|
if call_id and call_id in tool_call_buffers:
|
||||||
tool_call_buffers[call_id]["arguments"] += getattr(event, "delta", "") or ""
|
delta = getattr(event, "delta", "") or ""
|
||||||
|
tool_call_buffers[call_id]["arguments"] += delta
|
||||||
|
if on_tool_call_delta and delta:
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"call_id": str(call_id),
|
||||||
|
"name": str(tool_call_buffers[call_id].get("name") or ""),
|
||||||
|
"arguments_delta": str(delta),
|
||||||
|
})
|
||||||
elif event_type == "response.function_call_arguments.done":
|
elif event_type == "response.function_call_arguments.done":
|
||||||
call_id = getattr(event, "call_id", None)
|
call_id = getattr(event, "call_id", None)
|
||||||
if call_id and call_id in tool_call_buffers:
|
if call_id and call_id in tool_call_buffers:
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ class ProviderSpec:
|
|||||||
display_name: str = "" # shown in `nanobot status`
|
display_name: str = "" # shown in `nanobot status`
|
||||||
|
|
||||||
# which provider implementation to use
|
# which provider implementation to use
|
||||||
# "openai_compat" | "anthropic" | "azure_openai" | "openai_codex" | "github_copilot" | "bedrock"
|
# "openai_compat" | "anthropic" | "azure_openai" | "openai_codex" | "github_copilot" | "xai_oauth" | "bedrock"
|
||||||
backend: str = "openai_compat"
|
backend: str = "openai_compat"
|
||||||
|
|
||||||
# extra env vars, e.g. (("ZHIPUAI_API_KEY", "{api_key}"),)
|
# extra env vars, e.g. (("ZHIPUAI_API_KEY", "{api_key}"),)
|
||||||
@@ -155,6 +155,18 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
|
|||||||
detect_by_base_keyword="huggingface",
|
detect_by_base_keyword="huggingface",
|
||||||
default_api_base="https://router.huggingface.co/v1",
|
default_api_base="https://router.huggingface.co/v1",
|
||||||
),
|
),
|
||||||
|
# Skywork API platform (APIFree): OpenAI-compatible MaaS gateway.
|
||||||
|
ProviderSpec(
|
||||||
|
name="skywork",
|
||||||
|
keywords=("skywork", "skyclaw", "apifree"),
|
||||||
|
env_key="SKYWORK_API_KEY",
|
||||||
|
display_name="Skywork",
|
||||||
|
backend="openai_compat",
|
||||||
|
env_extras=(("APIFREE_API_KEY", "{api_key}"),),
|
||||||
|
is_gateway=True,
|
||||||
|
detect_by_base_keyword="apifree.ai",
|
||||||
|
default_api_base="https://api.apifree.ai/agent/v1",
|
||||||
|
),
|
||||||
# AiHubMix: global gateway, OpenAI-compatible interface.
|
# AiHubMix: global gateway, OpenAI-compatible interface.
|
||||||
# strip_model_prefix=True: doesn't understand "anthropic/claude-3",
|
# strip_model_prefix=True: doesn't understand "anthropic/claude-3",
|
||||||
# strips to bare "claude-3".
|
# strips to bare "claude-3".
|
||||||
@@ -279,6 +291,18 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
|
|||||||
is_oauth=True,
|
is_oauth=True,
|
||||||
supports_max_completion_tokens=True,
|
supports_max_completion_tokens=True,
|
||||||
),
|
),
|
||||||
|
# xAI Grok OAuth: SuperGrok subscription-backed Responses API provider
|
||||||
|
ProviderSpec(
|
||||||
|
name="xai_oauth",
|
||||||
|
keywords=("xai-oauth", "grok-oauth", "x-ai-oauth", "xai-grok-oauth"),
|
||||||
|
env_key="",
|
||||||
|
display_name="xAI Grok OAuth",
|
||||||
|
backend="xai_oauth",
|
||||||
|
default_api_base="https://api.x.ai/v1",
|
||||||
|
strip_model_prefix=True,
|
||||||
|
is_oauth=True,
|
||||||
|
supports_max_completion_tokens=True,
|
||||||
|
),
|
||||||
# DeepSeek: OpenAI-compatible at api.deepseek.com
|
# DeepSeek: OpenAI-compatible at api.deepseek.com
|
||||||
ProviderSpec(
|
ProviderSpec(
|
||||||
name="deepseek",
|
name="deepseek",
|
||||||
@@ -390,13 +414,23 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
|
|||||||
backend="openai_compat",
|
backend="openai_compat",
|
||||||
default_api_base="https://api.longcat.chat/openai/v1",
|
default_api_base="https://api.longcat.chat/openai/v1",
|
||||||
),
|
),
|
||||||
|
# Ant Ling: OpenAI-compatible API for Ling/Ring model families.
|
||||||
|
ProviderSpec(
|
||||||
|
name="ant_ling",
|
||||||
|
keywords=("ant_ling", "ant-ling", "ling-", "ring-"),
|
||||||
|
env_key="ANT_LING_API_KEY",
|
||||||
|
display_name="Ant Ling",
|
||||||
|
backend="openai_compat",
|
||||||
|
detect_by_base_keyword="ant-ling.com",
|
||||||
|
default_api_base="https://api.ant-ling.com/v1",
|
||||||
|
),
|
||||||
# === Local deployment (matched by config key, NOT by api_base) =========
|
# === Local deployment (matched by config key, NOT by api_base) =========
|
||||||
# vLLM / any OpenAI-compatible local server
|
# vLLM / any OpenAI-compatible local server
|
||||||
ProviderSpec(
|
ProviderSpec(
|
||||||
name="vllm",
|
name="vllm",
|
||||||
keywords=("vllm",),
|
keywords=("vllm",),
|
||||||
env_key="HOSTED_VLLM_API_KEY",
|
env_key="HOSTED_VLLM_API_KEY",
|
||||||
display_name="vLLM/Local",
|
display_name="vLLM",
|
||||||
backend="openai_compat",
|
backend="openai_compat",
|
||||||
is_local=True,
|
is_local=True,
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -0,0 +1,768 @@
|
|||||||
|
"""xAI Grok OAuth credential flow and Responses provider."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import base64
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import secrets
|
||||||
|
import time
|
||||||
|
import webbrowser
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from contextlib import suppress
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from hashlib import sha256
|
||||||
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||||
|
from pathlib import Path
|
||||||
|
from threading import Event, Thread
|
||||||
|
from typing import Any
|
||||||
|
from urllib.parse import parse_qs, urlencode, urlparse
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from filelock import FileLock
|
||||||
|
|
||||||
|
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
||||||
|
from nanobot.providers.openai_responses import consume_sse, convert_messages, convert_tools
|
||||||
|
|
||||||
|
DEFAULT_XAI_API_BASE = "https://api.x.ai/v1"
|
||||||
|
DEFAULT_XAI_AUTH_ISSUER = "https://auth.x.ai"
|
||||||
|
DEFAULT_XAI_DISCOVERY_URL = f"{DEFAULT_XAI_AUTH_ISSUER}/.well-known/openid-configuration"
|
||||||
|
DEFAULT_XAI_REDIRECT_URI = "http://127.0.0.1:56121/callback"
|
||||||
|
DEFAULT_XAI_CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828"
|
||||||
|
DEFAULT_XAI_SCOPE = "openid profile email offline_access grok-cli:access api:access"
|
||||||
|
|
||||||
|
_SERVICE_NAME = "nanobot.xai_oauth"
|
||||||
|
_SECRET_USERNAME = "default"
|
||||||
|
_TOKEN_SKEW_SECONDS = 60
|
||||||
|
_LOGIN_TIMEOUT_SECONDS = 300
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class XaiOAuthEndpoints:
|
||||||
|
authorization_endpoint: str
|
||||||
|
token_endpoint: str
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class XaiOAuthCredential:
|
||||||
|
access_token: str
|
||||||
|
refresh_token: str = ""
|
||||||
|
expires_at: float | None = None
|
||||||
|
account_id: str | None = None
|
||||||
|
token_type: str = "Bearer"
|
||||||
|
api_base: str = DEFAULT_XAI_API_BASE
|
||||||
|
storage: str = "unknown"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_expiring(self) -> bool:
|
||||||
|
return self.expires_at is not None and self.expires_at <= time.time() + _TOKEN_SKEW_SECONDS
|
||||||
|
|
||||||
|
|
||||||
|
def _nanobot_home() -> Path:
|
||||||
|
override = os.environ.get("NANOBOT_HOME")
|
||||||
|
if override:
|
||||||
|
return Path(override).expanduser()
|
||||||
|
from nanobot.config.loader import get_config_path
|
||||||
|
|
||||||
|
return get_config_path().parent
|
||||||
|
|
||||||
|
|
||||||
|
def _auth_dir() -> Path:
|
||||||
|
return _nanobot_home() / "auth"
|
||||||
|
|
||||||
|
|
||||||
|
def get_xai_oauth_metadata_path() -> Path:
|
||||||
|
"""Return the non-secret xAI OAuth metadata path."""
|
||||||
|
return _auth_dir() / "xai-oauth.json"
|
||||||
|
|
||||||
|
|
||||||
|
def _lock_path() -> Path:
|
||||||
|
return get_xai_oauth_metadata_path().with_suffix(".lock")
|
||||||
|
|
||||||
|
|
||||||
|
def _write_private_json(path: Path, payload: dict[str, Any]) -> None:
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
with suppress(OSError):
|
||||||
|
path.parent.chmod(0o700)
|
||||||
|
tmp = path.with_suffix(path.suffix + ".tmp")
|
||||||
|
tmp.write_text(json.dumps(payload, indent=2, ensure_ascii=False), encoding="utf-8")
|
||||||
|
with suppress(OSError):
|
||||||
|
tmp.chmod(0o600)
|
||||||
|
tmp.replace(path)
|
||||||
|
with suppress(OSError):
|
||||||
|
path.chmod(0o600)
|
||||||
|
|
||||||
|
|
||||||
|
def _read_json(path: Path) -> dict[str, Any]:
|
||||||
|
return json.loads(path.read_text(encoding="utf-8"))
|
||||||
|
|
||||||
|
|
||||||
|
def _keyring_set(tokens: dict[str, Any]) -> bool:
|
||||||
|
try:
|
||||||
|
import keyring # type: ignore[import-not-found]
|
||||||
|
|
||||||
|
keyring.set_password(_SERVICE_NAME, _SECRET_USERNAME, json.dumps(tokens))
|
||||||
|
return True
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _keyring_get() -> dict[str, Any] | None:
|
||||||
|
try:
|
||||||
|
import keyring # type: ignore[import-not-found]
|
||||||
|
|
||||||
|
raw = keyring.get_password(_SERVICE_NAME, _SECRET_USERNAME)
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
if not raw:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
payload = json.loads(raw)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return None
|
||||||
|
return payload if isinstance(payload, dict) else None
|
||||||
|
|
||||||
|
|
||||||
|
def _keyring_delete() -> None:
|
||||||
|
try:
|
||||||
|
import keyring # type: ignore[import-not-found]
|
||||||
|
|
||||||
|
keyring.delete_password(_SERVICE_NAME, _SECRET_USERNAME)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def _token_payload(credential: XaiOAuthCredential) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"access_token": credential.access_token,
|
||||||
|
"refresh_token": credential.refresh_token,
|
||||||
|
"expires_at": credential.expires_at,
|
||||||
|
"token_type": credential.token_type,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def save_xai_oauth_credential(credential: XaiOAuthCredential) -> XaiOAuthCredential:
|
||||||
|
"""Persist xAI OAuth tokens, preferring OS keychain storage."""
|
||||||
|
with FileLock(str(_lock_path())):
|
||||||
|
tokens = _token_payload(credential)
|
||||||
|
metadata: dict[str, Any] = {
|
||||||
|
"provider": "xai_oauth",
|
||||||
|
"api_base": credential.api_base,
|
||||||
|
"account_id": credential.account_id,
|
||||||
|
"expires_at": credential.expires_at,
|
||||||
|
"updated_at": int(time.time()),
|
||||||
|
}
|
||||||
|
if _keyring_set(tokens):
|
||||||
|
metadata["storage"] = "keyring"
|
||||||
|
else:
|
||||||
|
metadata["storage"] = "file"
|
||||||
|
metadata["tokens"] = tokens
|
||||||
|
_write_private_json(get_xai_oauth_metadata_path(), metadata)
|
||||||
|
return XaiOAuthCredential(
|
||||||
|
access_token=credential.access_token,
|
||||||
|
refresh_token=credential.refresh_token,
|
||||||
|
expires_at=credential.expires_at,
|
||||||
|
account_id=credential.account_id,
|
||||||
|
token_type=credential.token_type,
|
||||||
|
api_base=credential.api_base,
|
||||||
|
storage=str(metadata["storage"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def load_xai_oauth_credential() -> XaiOAuthCredential | None:
|
||||||
|
"""Load xAI OAuth credentials from keyring or the private file fallback."""
|
||||||
|
path = get_xai_oauth_metadata_path()
|
||||||
|
if not path.exists():
|
||||||
|
return None
|
||||||
|
with FileLock(str(_lock_path())):
|
||||||
|
try:
|
||||||
|
metadata = _read_json(path)
|
||||||
|
except (OSError, json.JSONDecodeError):
|
||||||
|
return None
|
||||||
|
|
||||||
|
storage = str(metadata.get("storage") or "file")
|
||||||
|
tokens = _keyring_get() if storage == "keyring" else metadata.get("tokens")
|
||||||
|
if not isinstance(tokens, dict):
|
||||||
|
return None
|
||||||
|
access_token = str(tokens.get("access_token") or "")
|
||||||
|
if not access_token:
|
||||||
|
return None
|
||||||
|
|
||||||
|
return XaiOAuthCredential(
|
||||||
|
access_token=access_token,
|
||||||
|
refresh_token=str(tokens.get("refresh_token") or ""),
|
||||||
|
expires_at=_as_float(tokens.get("expires_at") or metadata.get("expires_at")),
|
||||||
|
account_id=_as_str(metadata.get("account_id")),
|
||||||
|
token_type=str(tokens.get("token_type") or "Bearer"),
|
||||||
|
api_base=str(metadata.get("api_base") or DEFAULT_XAI_API_BASE),
|
||||||
|
storage=storage,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def delete_xai_oauth_credentials() -> list[Path]:
|
||||||
|
"""Delete persisted xAI OAuth credentials and return removed local paths."""
|
||||||
|
removed: list[Path] = []
|
||||||
|
path = get_xai_oauth_metadata_path()
|
||||||
|
lock_path = _lock_path()
|
||||||
|
with FileLock(str(lock_path)):
|
||||||
|
_keyring_delete()
|
||||||
|
try:
|
||||||
|
path.unlink()
|
||||||
|
removed.append(path)
|
||||||
|
except FileNotFoundError:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
lock_path.unlink()
|
||||||
|
except FileNotFoundError:
|
||||||
|
pass
|
||||||
|
return removed
|
||||||
|
|
||||||
|
|
||||||
|
def get_xai_oauth_login_status() -> XaiOAuthCredential | None:
|
||||||
|
return load_xai_oauth_credential()
|
||||||
|
|
||||||
|
|
||||||
|
def pkce_challenge(verifier: str) -> str:
|
||||||
|
digest = sha256(verifier.encode("ascii")).digest()
|
||||||
|
return base64.urlsafe_b64encode(digest).decode("ascii").rstrip("=")
|
||||||
|
|
||||||
|
|
||||||
|
def _new_pkce_verifier() -> str:
|
||||||
|
return base64.urlsafe_b64encode(secrets.token_bytes(48)).decode("ascii").rstrip("=")
|
||||||
|
|
||||||
|
|
||||||
|
def build_xai_authorization_url(
|
||||||
|
endpoints: XaiOAuthEndpoints,
|
||||||
|
*,
|
||||||
|
verifier: str,
|
||||||
|
state: str,
|
||||||
|
nonce: str | None = None,
|
||||||
|
redirect_uri: str = DEFAULT_XAI_REDIRECT_URI,
|
||||||
|
) -> str:
|
||||||
|
params = {
|
||||||
|
"response_type": "code",
|
||||||
|
"client_id": DEFAULT_XAI_CLIENT_ID,
|
||||||
|
"redirect_uri": redirect_uri,
|
||||||
|
"scope": DEFAULT_XAI_SCOPE,
|
||||||
|
"code_challenge": pkce_challenge(verifier),
|
||||||
|
"code_challenge_method": "S256",
|
||||||
|
"state": state,
|
||||||
|
"nonce": nonce or secrets.token_urlsafe(16),
|
||||||
|
"plan": "generic",
|
||||||
|
"referrer": "nanobot",
|
||||||
|
}
|
||||||
|
return f"{endpoints.authorization_endpoint}?{urlencode(params)}"
|
||||||
|
|
||||||
|
|
||||||
|
def discover_xai_oauth_endpoints() -> XaiOAuthEndpoints:
|
||||||
|
try:
|
||||||
|
with httpx.Client(timeout=20.0, follow_redirects=True, trust_env=True) as client:
|
||||||
|
response = client.get(DEFAULT_XAI_DISCOVERY_URL)
|
||||||
|
response.raise_for_status()
|
||||||
|
payload = response.json()
|
||||||
|
except Exception:
|
||||||
|
payload = {}
|
||||||
|
|
||||||
|
endpoints = XaiOAuthEndpoints(
|
||||||
|
authorization_endpoint=str(
|
||||||
|
payload.get("authorization_endpoint")
|
||||||
|
or f"{DEFAULT_XAI_AUTH_ISSUER}/authorize"
|
||||||
|
),
|
||||||
|
token_endpoint=str(
|
||||||
|
payload.get("token_endpoint")
|
||||||
|
or f"{DEFAULT_XAI_AUTH_ISSUER}/oauth/token"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
_validate_xai_endpoint(endpoints.authorization_endpoint, "authorization_endpoint")
|
||||||
|
_validate_xai_endpoint(endpoints.token_endpoint, "token_endpoint")
|
||||||
|
return endpoints
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_xai_endpoint(url: str, label: str) -> None:
|
||||||
|
parsed = urlparse(url)
|
||||||
|
host = parsed.hostname or ""
|
||||||
|
if parsed.scheme != "https" or not (host == "x.ai" or host.endswith(".x.ai")):
|
||||||
|
raise RuntimeError(f"Refusing non-xAI OAuth {label}: {url}")
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_callback_value(raw: str) -> tuple[str, str | None]:
|
||||||
|
raw = raw.strip()
|
||||||
|
parsed = urlparse(raw)
|
||||||
|
if parsed.scheme and parsed.netloc:
|
||||||
|
params = parse_qs(parsed.query)
|
||||||
|
code = (params.get("code") or [""])[0]
|
||||||
|
state = (params.get("state") or [None])[0]
|
||||||
|
if not code:
|
||||||
|
raise RuntimeError("OAuth callback URL did not contain a code.")
|
||||||
|
return code, state
|
||||||
|
if raw.startswith("?") or "=" in raw:
|
||||||
|
params = parse_qs(raw.lstrip("?"))
|
||||||
|
code = (params.get("code") or [""])[0]
|
||||||
|
state = (params.get("state") or [None])[0]
|
||||||
|
if not code:
|
||||||
|
raise RuntimeError("OAuth callback query did not contain a code.")
|
||||||
|
return code, state
|
||||||
|
if raw:
|
||||||
|
return raw, None
|
||||||
|
raise RuntimeError("No OAuth code provided.")
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_jwt_payload(token: str) -> dict[str, Any]:
|
||||||
|
parts = token.split(".")
|
||||||
|
if len(parts) < 2:
|
||||||
|
return {}
|
||||||
|
data = parts[1] + "=" * (-len(parts[1]) % 4)
|
||||||
|
try:
|
||||||
|
decoded = base64.urlsafe_b64decode(data.encode("ascii"))
|
||||||
|
payload = json.loads(decoded)
|
||||||
|
except Exception:
|
||||||
|
return {}
|
||||||
|
return payload if isinstance(payload, dict) else {}
|
||||||
|
|
||||||
|
|
||||||
|
def _credential_from_token_response(payload: dict[str, Any], previous: XaiOAuthCredential | None = None) -> XaiOAuthCredential:
|
||||||
|
access_token = str(payload.get("access_token") or "")
|
||||||
|
if not access_token:
|
||||||
|
raise RuntimeError("xAI token response did not include an access token.")
|
||||||
|
|
||||||
|
claims = _decode_jwt_payload(access_token)
|
||||||
|
id_claims = _decode_jwt_payload(str(payload.get("id_token") or ""))
|
||||||
|
expires_at = _as_float(payload.get("expires_at"))
|
||||||
|
if expires_at is None:
|
||||||
|
expires_in = _as_float(payload.get("expires_in"))
|
||||||
|
expires_at = time.time() + expires_in if expires_in else _as_float(claims.get("exp"))
|
||||||
|
|
||||||
|
account_id = (
|
||||||
|
_as_str(id_claims.get("email"))
|
||||||
|
or _as_str(id_claims.get("preferred_username"))
|
||||||
|
or _as_str(id_claims.get("sub"))
|
||||||
|
or _as_str(claims.get("sub"))
|
||||||
|
or (previous.account_id if previous else None)
|
||||||
|
)
|
||||||
|
refresh_token = str(payload.get("refresh_token") or (previous.refresh_token if previous else ""))
|
||||||
|
|
||||||
|
return XaiOAuthCredential(
|
||||||
|
access_token=access_token,
|
||||||
|
refresh_token=refresh_token,
|
||||||
|
expires_at=expires_at,
|
||||||
|
account_id=account_id,
|
||||||
|
token_type=str(payload.get("token_type") or (previous.token_type if previous else "Bearer")),
|
||||||
|
api_base=previous.api_base if previous else DEFAULT_XAI_API_BASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def exchange_xai_oauth_code(
|
||||||
|
code: str,
|
||||||
|
*,
|
||||||
|
verifier: str,
|
||||||
|
endpoints: XaiOAuthEndpoints | None = None,
|
||||||
|
redirect_uri: str = DEFAULT_XAI_REDIRECT_URI,
|
||||||
|
) -> XaiOAuthCredential:
|
||||||
|
endpoints = endpoints or discover_xai_oauth_endpoints()
|
||||||
|
challenge = pkce_challenge(verifier)
|
||||||
|
with httpx.Client(timeout=30.0, follow_redirects=True, trust_env=True) as client:
|
||||||
|
response = client.post(
|
||||||
|
endpoints.token_endpoint,
|
||||||
|
headers={"Accept": "application/json"},
|
||||||
|
data={
|
||||||
|
"grant_type": "authorization_code",
|
||||||
|
"client_id": DEFAULT_XAI_CLIENT_ID,
|
||||||
|
"code": code,
|
||||||
|
"redirect_uri": redirect_uri,
|
||||||
|
"code_verifier": verifier,
|
||||||
|
"code_challenge": challenge,
|
||||||
|
"code_challenge_method": "S256",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if response.status_code >= 400:
|
||||||
|
raise RuntimeError(f"xAI token exchange failed: HTTP {response.status_code}: {response.text[:500]}")
|
||||||
|
return _credential_from_token_response(response.json())
|
||||||
|
|
||||||
|
|
||||||
|
def refresh_xai_oauth_credential(credential: XaiOAuthCredential | None = None) -> XaiOAuthCredential:
|
||||||
|
credential = credential or load_xai_oauth_credential()
|
||||||
|
if not credential or not credential.refresh_token:
|
||||||
|
raise RuntimeError("xAI Grok OAuth is not logged in. Run: nanobot provider login xai-oauth")
|
||||||
|
|
||||||
|
endpoints = discover_xai_oauth_endpoints()
|
||||||
|
with httpx.Client(timeout=30.0, follow_redirects=True, trust_env=True) as client:
|
||||||
|
response = client.post(
|
||||||
|
endpoints.token_endpoint,
|
||||||
|
headers={"Accept": "application/json"},
|
||||||
|
data={
|
||||||
|
"grant_type": "refresh_token",
|
||||||
|
"client_id": DEFAULT_XAI_CLIENT_ID,
|
||||||
|
"refresh_token": credential.refresh_token,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if response.status_code >= 400:
|
||||||
|
raise RuntimeError(f"xAI token refresh failed: HTTP {response.status_code}: {response.text[:500]}")
|
||||||
|
return save_xai_oauth_credential(_credential_from_token_response(response.json(), previous=credential))
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_xai_oauth_credential(*, force_refresh: bool = False) -> XaiOAuthCredential:
|
||||||
|
credential = load_xai_oauth_credential()
|
||||||
|
if not credential:
|
||||||
|
raise RuntimeError("xAI Grok OAuth is not logged in. Run: nanobot provider login xai-oauth")
|
||||||
|
if force_refresh or credential.is_expiring:
|
||||||
|
credential = refresh_xai_oauth_credential(credential)
|
||||||
|
return credential
|
||||||
|
|
||||||
|
|
||||||
|
def login_xai_oauth_interactive(
|
||||||
|
print_fn: Callable[[str], None] | None = None,
|
||||||
|
prompt_fn: Callable[[str], str] | None = None,
|
||||||
|
open_browser: bool = True,
|
||||||
|
manual_paste: bool = False,
|
||||||
|
timeout_seconds: int = _LOGIN_TIMEOUT_SECONDS,
|
||||||
|
) -> XaiOAuthCredential:
|
||||||
|
"""Run browser PKCE login and persist xAI OAuth credentials."""
|
||||||
|
printer = print_fn or print
|
||||||
|
prompt = prompt_fn or input
|
||||||
|
endpoints = discover_xai_oauth_endpoints()
|
||||||
|
verifier = _new_pkce_verifier()
|
||||||
|
state = secrets.token_urlsafe(24)
|
||||||
|
nonce = secrets.token_urlsafe(24)
|
||||||
|
authorize_url = build_xai_authorization_url(
|
||||||
|
endpoints,
|
||||||
|
verifier=verifier,
|
||||||
|
state=state,
|
||||||
|
nonce=nonce,
|
||||||
|
)
|
||||||
|
|
||||||
|
callback = _LoopbackCallback()
|
||||||
|
server_started = False if manual_paste else callback.start()
|
||||||
|
printer(f"Open: {authorize_url}")
|
||||||
|
if open_browser:
|
||||||
|
with suppress(Exception):
|
||||||
|
webbrowser.open(authorize_url)
|
||||||
|
|
||||||
|
result: dict[str, str] | None = None
|
||||||
|
if manual_paste:
|
||||||
|
printer("Paste the callback URL or xAI fallback code after authorization.")
|
||||||
|
elif server_started:
|
||||||
|
try:
|
||||||
|
result = callback.wait(timeout_seconds)
|
||||||
|
finally:
|
||||||
|
callback.stop()
|
||||||
|
else:
|
||||||
|
printer("Loopback port 56121 is unavailable; paste the callback URL or xAI fallback code.")
|
||||||
|
|
||||||
|
if result:
|
||||||
|
code = result.get("code") or ""
|
||||||
|
returned_state = result.get("state")
|
||||||
|
else:
|
||||||
|
pasted = prompt("Paste callback URL or fallback code")
|
||||||
|
code, returned_state = _parse_callback_value(pasted)
|
||||||
|
|
||||||
|
if not code:
|
||||||
|
raise RuntimeError("OAuth login did not return a code.")
|
||||||
|
if returned_state and returned_state != state:
|
||||||
|
raise RuntimeError("OAuth state mismatch. Please retry login.")
|
||||||
|
|
||||||
|
credential = exchange_xai_oauth_code(code, verifier=verifier, endpoints=endpoints)
|
||||||
|
return save_xai_oauth_credential(credential)
|
||||||
|
|
||||||
|
|
||||||
|
class _LoopbackCallback:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._event = Event()
|
||||||
|
self._result: dict[str, str] = {}
|
||||||
|
self._server: ThreadingHTTPServer | None = None
|
||||||
|
self._thread: Thread | None = None
|
||||||
|
|
||||||
|
def start(self) -> bool:
|
||||||
|
owner = self
|
||||||
|
|
||||||
|
class Handler(BaseHTTPRequestHandler):
|
||||||
|
def do_GET(self) -> None: # noqa: N802 - stdlib callback name
|
||||||
|
parsed = urlparse(self.path)
|
||||||
|
params = parse_qs(parsed.query)
|
||||||
|
code = (params.get("code") or [""])[0]
|
||||||
|
state = (params.get("state") or [""])[0]
|
||||||
|
if parsed.path != "/callback" or not code:
|
||||||
|
self.send_response(404)
|
||||||
|
self.end_headers()
|
||||||
|
return
|
||||||
|
owner._result = {"code": code, "state": state}
|
||||||
|
owner._event.set()
|
||||||
|
self.send_response(200)
|
||||||
|
self.send_header("Content-Type", "text/html; charset=utf-8")
|
||||||
|
self.end_headers()
|
||||||
|
self.wfile.write(b"<html><body>nanobot xAI OAuth complete. You may close this tab.</body></html>")
|
||||||
|
|
||||||
|
def log_message(self, format: str, *args: Any) -> None: # noqa: A002
|
||||||
|
return
|
||||||
|
|
||||||
|
class Server(ThreadingHTTPServer):
|
||||||
|
allow_reuse_address = True
|
||||||
|
daemon_threads = True
|
||||||
|
|
||||||
|
try:
|
||||||
|
self._server = Server(("127.0.0.1", 56121), Handler)
|
||||||
|
except OSError:
|
||||||
|
return False
|
||||||
|
self._thread = Thread(target=self._server.serve_forever, daemon=True)
|
||||||
|
self._thread.start()
|
||||||
|
return True
|
||||||
|
|
||||||
|
def wait(self, timeout_seconds: int) -> dict[str, str] | None:
|
||||||
|
if self._event.wait(timeout_seconds):
|
||||||
|
return dict(self._result)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def stop(self) -> None:
|
||||||
|
if self._server:
|
||||||
|
self._server.shutdown()
|
||||||
|
self._server.server_close()
|
||||||
|
if self._thread:
|
||||||
|
self._thread.join(timeout=1)
|
||||||
|
|
||||||
|
|
||||||
|
def _as_float(value: Any) -> float | None:
|
||||||
|
try:
|
||||||
|
return float(value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _as_str(value: Any) -> str | None:
|
||||||
|
return value if isinstance(value, str) and value else None
|
||||||
|
|
||||||
|
|
||||||
|
DEFAULT_XAI_MODEL = "xai-oauth/grok-4.3"
|
||||||
|
|
||||||
|
|
||||||
|
class XaiOAuthProvider(LLMProvider):
|
||||||
|
"""Use a SuperGrok OAuth session to call xAI's Responses API."""
|
||||||
|
|
||||||
|
supports_progress_deltas = True
|
||||||
|
|
||||||
|
def __init__(self, default_model: str = DEFAULT_XAI_MODEL, config: Any | None = None):
|
||||||
|
super().__init__(api_key=None, api_base=DEFAULT_XAI_API_BASE)
|
||||||
|
self.default_model = default_model
|
||||||
|
self.config = config
|
||||||
|
|
||||||
|
async def _call_xai(
|
||||||
|
self,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
tools: list[dict[str, Any]] | None,
|
||||||
|
model: str | None,
|
||||||
|
max_tokens: int,
|
||||||
|
temperature: float,
|
||||||
|
reasoning_effort: str | None,
|
||||||
|
tool_choice: str | dict[str, Any] | None,
|
||||||
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
|
) -> LLMResponse:
|
||||||
|
body = _build_xai_responses_body(
|
||||||
|
messages=messages,
|
||||||
|
tools=tools,
|
||||||
|
model=model or self.default_model,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
reasoning_effort=reasoning_effort,
|
||||||
|
tool_choice=tool_choice,
|
||||||
|
hosted_x_search=getattr(self.config, "x_search", None),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
credential = await asyncio.to_thread(resolve_xai_oauth_credential)
|
||||||
|
try:
|
||||||
|
content, tool_calls, finish_reason = await _request_xai(
|
||||||
|
credential,
|
||||||
|
body,
|
||||||
|
on_content_delta=on_content_delta,
|
||||||
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
|
)
|
||||||
|
except _XaiHTTPError as exc:
|
||||||
|
if exc.status_code != 401:
|
||||||
|
raise
|
||||||
|
credential = await asyncio.to_thread(resolve_xai_oauth_credential, force_refresh=True)
|
||||||
|
content, tool_calls, finish_reason = await _request_xai(
|
||||||
|
credential,
|
||||||
|
body,
|
||||||
|
on_content_delta=on_content_delta,
|
||||||
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
|
)
|
||||||
|
return LLMResponse(content=content, tool_calls=tool_calls, finish_reason=finish_reason)
|
||||||
|
except Exception as exc:
|
||||||
|
msg = f"Error calling xAI Grok OAuth: {exc}"
|
||||||
|
retry_after = getattr(exc, "retry_after", None) or self._extract_retry_after(msg)
|
||||||
|
return LLMResponse(content=msg, finish_reason="error", retry_after=retry_after)
|
||||||
|
|
||||||
|
async def chat(
|
||||||
|
self,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
tools: list[dict[str, Any]] | None = None,
|
||||||
|
model: str | None = None,
|
||||||
|
max_tokens: int = 4096,
|
||||||
|
temperature: float = 0.7,
|
||||||
|
reasoning_effort: str | None = None,
|
||||||
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
|
) -> LLMResponse:
|
||||||
|
return await self._call_xai(
|
||||||
|
messages,
|
||||||
|
tools,
|
||||||
|
model,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
reasoning_effort,
|
||||||
|
tool_choice,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def chat_stream(
|
||||||
|
self,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
tools: list[dict[str, Any]] | None = None,
|
||||||
|
model: str | None = None,
|
||||||
|
max_tokens: int = 4096,
|
||||||
|
temperature: float = 0.7,
|
||||||
|
reasoning_effort: str | None = None,
|
||||||
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
|
) -> LLMResponse:
|
||||||
|
_ = on_thinking_delta
|
||||||
|
return await self._call_xai(
|
||||||
|
messages,
|
||||||
|
tools,
|
||||||
|
model,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
reasoning_effort,
|
||||||
|
tool_choice,
|
||||||
|
on_content_delta,
|
||||||
|
on_tool_call_delta,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_default_model(self) -> str:
|
||||||
|
return self.default_model
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_model_prefix(model: str) -> str:
|
||||||
|
for prefix in ("xai-oauth/", "xai_oauth/", "grok-oauth/", "grok_oauth/"):
|
||||||
|
if model.startswith(prefix):
|
||||||
|
return model.split("/", 1)[1]
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def _build_xai_responses_body(
|
||||||
|
*,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
tools: list[dict[str, Any]] | None,
|
||||||
|
model: str,
|
||||||
|
max_tokens: int,
|
||||||
|
temperature: float,
|
||||||
|
reasoning_effort: str | None,
|
||||||
|
tool_choice: str | dict[str, Any] | None,
|
||||||
|
hosted_x_search: Any | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
system_prompt, input_items = convert_messages(LLMProvider._sanitize_empty_content(messages))
|
||||||
|
if system_prompt:
|
||||||
|
input_items = [
|
||||||
|
{"role": "system", "content": [{"type": "input_text", "text": system_prompt}]},
|
||||||
|
*input_items,
|
||||||
|
]
|
||||||
|
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"model": _strip_model_prefix(model),
|
||||||
|
"store": False,
|
||||||
|
"stream": True,
|
||||||
|
"input": input_items,
|
||||||
|
"tool_choice": tool_choice or "auto",
|
||||||
|
"parallel_tool_calls": True,
|
||||||
|
}
|
||||||
|
if max_tokens:
|
||||||
|
body["max_output_tokens"] = max_tokens
|
||||||
|
if temperature is not None:
|
||||||
|
body["temperature"] = temperature
|
||||||
|
if reasoning_effort and reasoning_effort.lower() != "none":
|
||||||
|
body["reasoning"] = {"effort": reasoning_effort}
|
||||||
|
converted_tools = convert_tools(tools) if tools else []
|
||||||
|
hosted_tool = _build_xai_hosted_x_search_tool(hosted_x_search)
|
||||||
|
if hosted_tool:
|
||||||
|
converted_tools.append(hosted_tool)
|
||||||
|
if converted_tools:
|
||||||
|
body["tools"] = converted_tools
|
||||||
|
return body
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_x_handles(handles: list[str] | None) -> list[str] | None:
|
||||||
|
if not handles:
|
||||||
|
return None
|
||||||
|
cleaned = [str(handle).strip().lstrip("@") for handle in handles if str(handle).strip()]
|
||||||
|
return cleaned[:10] or None
|
||||||
|
|
||||||
|
|
||||||
|
def _build_xai_hosted_x_search_tool(config: Any | None) -> dict[str, Any] | None:
|
||||||
|
if not config or not getattr(config, "enable", False):
|
||||||
|
return None
|
||||||
|
|
||||||
|
allowed = _clean_x_handles(getattr(config, "allowed_x_handles", None))
|
||||||
|
excluded = _clean_x_handles(getattr(config, "excluded_x_handles", None))
|
||||||
|
if allowed and excluded:
|
||||||
|
raise ValueError("providers.xai_oauth.x_search cannot set both allowed_x_handles and excluded_x_handles")
|
||||||
|
|
||||||
|
tool: dict[str, Any] = {"type": "x_search"}
|
||||||
|
if allowed:
|
||||||
|
tool["allowed_x_handles"] = allowed
|
||||||
|
if excluded:
|
||||||
|
tool["excluded_x_handles"] = excluded
|
||||||
|
if getattr(config, "from_date", None):
|
||||||
|
tool["from_date"] = config.from_date
|
||||||
|
if getattr(config, "to_date", None):
|
||||||
|
tool["to_date"] = config.to_date
|
||||||
|
if getattr(config, "enable_image_understanding", False):
|
||||||
|
tool["enable_image_understanding"] = True
|
||||||
|
if getattr(config, "enable_video_understanding", False):
|
||||||
|
tool["enable_video_understanding"] = True
|
||||||
|
return tool
|
||||||
|
|
||||||
|
|
||||||
|
class _XaiHTTPError(RuntimeError):
|
||||||
|
def __init__(self, message: str, *, status_code: int, retry_after: float | None = None):
|
||||||
|
super().__init__(message)
|
||||||
|
self.status_code = status_code
|
||||||
|
self.retry_after = retry_after
|
||||||
|
|
||||||
|
|
||||||
|
async def _request_xai(
|
||||||
|
credential: XaiOAuthCredential,
|
||||||
|
body: dict[str, Any],
|
||||||
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
|
) -> tuple[str, list[ToolCallRequest], str]:
|
||||||
|
url = credential.api_base.rstrip("/") + "/responses"
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {credential.access_token}",
|
||||||
|
"Accept": "text/event-stream",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
"User-Agent": "nanobot (python)",
|
||||||
|
}
|
||||||
|
timeout = httpx.Timeout(120.0, connect=20.0)
|
||||||
|
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True, trust_env=True) as client:
|
||||||
|
async with client.stream("POST", url, headers=headers, json=body) as response:
|
||||||
|
if response.status_code != 200:
|
||||||
|
raw = await response.aread()
|
||||||
|
retry_after = LLMProvider._extract_retry_after_from_headers(response.headers)
|
||||||
|
raise _XaiHTTPError(
|
||||||
|
_friendly_error(response.status_code, raw.decode("utf-8", "ignore")),
|
||||||
|
status_code=response.status_code,
|
||||||
|
retry_after=retry_after,
|
||||||
|
)
|
||||||
|
return await consume_sse(response, on_content_delta, on_tool_call_delta)
|
||||||
|
|
||||||
|
|
||||||
|
def _friendly_error(status_code: int, raw: str) -> str:
|
||||||
|
if status_code == 401:
|
||||||
|
return "xAI OAuth session expired or was revoked. Run: nanobot provider login xai-oauth"
|
||||||
|
if status_code == 403:
|
||||||
|
return (
|
||||||
|
"xAI accepted the OAuth token, but this account is not entitled for the requested "
|
||||||
|
"Grok API capability yet. Check the active Grok subscription and selected model."
|
||||||
|
)
|
||||||
|
if status_code == 429:
|
||||||
|
return "xAI Grok subscription quota or rate limit was reached. Please try again later."
|
||||||
|
return f"HTTP {status_code}: {raw[:500]}"
|
||||||
@@ -8,7 +8,7 @@ from contextlib import suppress
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Literal
|
from typing import Any
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -581,36 +581,6 @@ class SessionManager:
|
|||||||
return self._session_payload(repaired)
|
return self._session_payload(repaired)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def get_or_create_task_session(
|
|
||||||
self,
|
|
||||||
base_key: str,
|
|
||||||
task_id: str,
|
|
||||||
role: Literal["manager", "worker"] = "worker",
|
|
||||||
) -> Session:
|
|
||||||
"""Get or create an isolated session for a specific task.
|
|
||||||
|
|
||||||
Key format: task:{base_key}:{task_id}:{role}
|
|
||||||
Example: task:slack:C123:root_qml:manager
|
|
||||||
"""
|
|
||||||
task_key = f"task:{base_key}:{task_id}:{role}"
|
|
||||||
return self.get_or_create(task_key)
|
|
||||||
|
|
||||||
def list_task_sessions(self, base_key: str) -> list[Session]:
|
|
||||||
"""List all task-scoped sessions for a given base key."""
|
|
||||||
prefix = f"task:{base_key}:"
|
|
||||||
return [
|
|
||||||
session for key, session in self._cache.items()
|
|
||||||
if key.startswith(prefix)
|
|
||||||
]
|
|
||||||
|
|
||||||
def finalize_task_session(self, task_id: str) -> None:
|
|
||||||
"""Mark a task session as finalized (read-only) by setting metadata."""
|
|
||||||
prefix = f"task:"
|
|
||||||
for key, session in list(self._cache.items()):
|
|
||||||
if f":{task_id}:" in key and key.startswith(prefix):
|
|
||||||
session.metadata["finalized"] = True
|
|
||||||
self.save(session)
|
|
||||||
|
|
||||||
def list_sessions(self) -> list[dict[str, Any]]:
|
def list_sessions(self) -> list[dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
List all sessions.
|
List all sessions.
|
||||||
|
|||||||
@@ -0,0 +1,347 @@
|
|||||||
|
"""Session turn helpers for WebUI-capable WebSocket sessions.
|
||||||
|
|
||||||
|
AgentLoop uses these without importing a concrete channel plugin; only
|
||||||
|
``channel == "websocket"`` messages are affected.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
import time
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.providers.base import LLMProvider
|
||||||
|
from nanobot.session.goal_state import goal_state_ws_blob
|
||||||
|
from nanobot.session.manager import Session, SessionManager
|
||||||
|
from nanobot.utils.helpers import truncate_text
|
||||||
|
from nanobot.utils.llm_runtime import LLMRuntime
|
||||||
|
|
||||||
|
WEBUI_SESSION_METADATA_KEY = "webui"
|
||||||
|
WEBUI_TITLE_METADATA_KEY = "title"
|
||||||
|
WEBUI_TITLE_USER_EDITED_METADATA_KEY = "title_user_edited"
|
||||||
|
TITLE_MAX_CHARS = 60
|
||||||
|
TITLE_GENERATION_MAX_TOKENS = 96
|
||||||
|
TITLE_GENERATION_REASONING_EFFORT = "none"
|
||||||
|
|
||||||
|
# Wall-clock turn start per ``chat_id`` (websocket only). Survives browser refresh while the
|
||||||
|
# gateway process stays up; cleared on idle/stop and implicitly dropped on restart.
|
||||||
|
_WEBSOCKET_TURN_WALL_STARTED_AT: dict[str, float] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def mark_webui_session(session: Session, metadata: dict[str, Any]) -> bool:
|
||||||
|
"""Persist a WebUI marker only when the inbound websocket frame opted in."""
|
||||||
|
if metadata.get(WEBUI_SESSION_METADATA_KEY) is not True:
|
||||||
|
return False
|
||||||
|
session.metadata[WEBUI_SESSION_METADATA_KEY] = True
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def clean_generated_title(raw: str | None) -> str:
|
||||||
|
text = (raw or "").strip()
|
||||||
|
if not text:
|
||||||
|
return ""
|
||||||
|
text = re.sub(r"^\s*(title|标题)\s*[::]\s*", "", text, flags=re.IGNORECASE)
|
||||||
|
text = text.strip().strip("\"'`“”‘’")
|
||||||
|
text = re.sub(r"\s+", " ", text).strip()
|
||||||
|
text = text.rstrip("。.!!??,,;;:")
|
||||||
|
if len(text) > TITLE_MAX_CHARS:
|
||||||
|
text = text[: TITLE_MAX_CHARS - 1].rstrip() + "…"
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def _title_inputs(session: Session) -> tuple[str, str]:
|
||||||
|
user_text = ""
|
||||||
|
assistant_text = ""
|
||||||
|
for message in session.messages:
|
||||||
|
if message.get("_command") is True:
|
||||||
|
continue
|
||||||
|
role = message.get("role")
|
||||||
|
content = message.get("content")
|
||||||
|
if not isinstance(content, str) or not content.strip():
|
||||||
|
continue
|
||||||
|
if role == "user" and not user_text:
|
||||||
|
user_text = content.strip()
|
||||||
|
elif role == "assistant" and not assistant_text:
|
||||||
|
assistant_text = content.strip()
|
||||||
|
if user_text and assistant_text:
|
||||||
|
break
|
||||||
|
return user_text, assistant_text
|
||||||
|
|
||||||
|
|
||||||
|
async def maybe_generate_webui_title(
|
||||||
|
*,
|
||||||
|
sessions: SessionManager,
|
||||||
|
session_key: str,
|
||||||
|
provider: LLMProvider,
|
||||||
|
model: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Generate and persist a short title for WebUI-owned sessions only."""
|
||||||
|
session = sessions.get_or_create(session_key)
|
||||||
|
if session.metadata.get(WEBUI_SESSION_METADATA_KEY) is not True:
|
||||||
|
return False
|
||||||
|
if session.metadata.get(WEBUI_TITLE_USER_EDITED_METADATA_KEY) is True:
|
||||||
|
return False
|
||||||
|
current_title = session.metadata.get(WEBUI_TITLE_METADATA_KEY)
|
||||||
|
if isinstance(current_title, str) and current_title.strip():
|
||||||
|
return False
|
||||||
|
|
||||||
|
user_text, assistant_text = _title_inputs(session)
|
||||||
|
if not user_text:
|
||||||
|
return False
|
||||||
|
|
||||||
|
prompt = (
|
||||||
|
"Generate a concise title for this chat.\n"
|
||||||
|
"Rules:\n"
|
||||||
|
"- Use the same language as the user when practical.\n"
|
||||||
|
"- 3 to 8 words.\n"
|
||||||
|
"- No quotes.\n"
|
||||||
|
"- No punctuation at the end.\n"
|
||||||
|
"- Return only the title.\n\n"
|
||||||
|
f"User: {truncate_text(user_text, 1_000)}"
|
||||||
|
)
|
||||||
|
if assistant_text:
|
||||||
|
prompt += f"\nAssistant: {truncate_text(assistant_text, 1_000)}"
|
||||||
|
|
||||||
|
try:
|
||||||
|
response = await provider.chat_with_retry(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"role": "system",
|
||||||
|
"content": (
|
||||||
|
"You write short, neutral chat titles. "
|
||||||
|
"Return only the title text."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
{"role": "user", "content": prompt},
|
||||||
|
],
|
||||||
|
tools=None,
|
||||||
|
model=model,
|
||||||
|
max_tokens=TITLE_GENERATION_MAX_TOKENS,
|
||||||
|
temperature=0.2,
|
||||||
|
reasoning_effort=TITLE_GENERATION_REASONING_EFFORT,
|
||||||
|
retry_mode="standard",
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.debug("Failed to generate webui session title for {}", session_key, exc_info=True)
|
||||||
|
return False
|
||||||
|
|
||||||
|
title = clean_generated_title(response.content)
|
||||||
|
if not title or title.lower().startswith("error"):
|
||||||
|
logger.debug(
|
||||||
|
"WebUI title generation returned no usable title for {} (finish_reason={})",
|
||||||
|
session_key,
|
||||||
|
response.finish_reason,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
session.metadata[WEBUI_TITLE_METADATA_KEY] = title
|
||||||
|
sessions.save(session)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
async def maybe_generate_webui_title_after_turn(
|
||||||
|
*,
|
||||||
|
channel: str,
|
||||||
|
metadata: dict[str, Any],
|
||||||
|
sessions: SessionManager,
|
||||||
|
session_key: str,
|
||||||
|
provider: LLMProvider,
|
||||||
|
model: str,
|
||||||
|
) -> bool:
|
||||||
|
if channel != "websocket" or metadata.get(WEBUI_SESSION_METADATA_KEY) is not True:
|
||||||
|
return False
|
||||||
|
return await maybe_generate_webui_title(
|
||||||
|
sessions=sessions,
|
||||||
|
session_key=session_key,
|
||||||
|
provider=provider,
|
||||||
|
model=model,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def websocket_turn_wall_started_at(chat_id: str) -> float | None:
|
||||||
|
"""Return ``time.time()`` when the active user turn began, if still running."""
|
||||||
|
return _WEBSOCKET_TURN_WALL_STARTED_AT.get(chat_id)
|
||||||
|
|
||||||
|
|
||||||
|
async def publish_turn_run_status(bus: MessageBus, msg: InboundMessage, status: str) -> None:
|
||||||
|
"""Notify WebSocket clients while a user turn is executing (timing strip)."""
|
||||||
|
if msg.channel != "websocket":
|
||||||
|
return
|
||||||
|
cid = str(msg.chat_id)
|
||||||
|
meta: dict[str, Any] = {
|
||||||
|
**dict(msg.metadata or {}),
|
||||||
|
"_goal_status": True,
|
||||||
|
"goal_status": status,
|
||||||
|
}
|
||||||
|
if status == "running":
|
||||||
|
t0 = time.time()
|
||||||
|
meta["started_at"] = t0
|
||||||
|
_WEBSOCKET_TURN_WALL_STARTED_AT[cid] = t0
|
||||||
|
else:
|
||||||
|
_WEBSOCKET_TURN_WALL_STARTED_AT.pop(cid, None)
|
||||||
|
await bus.publish_outbound(
|
||||||
|
OutboundMessage(
|
||||||
|
channel=msg.channel,
|
||||||
|
chat_id=cid,
|
||||||
|
content="",
|
||||||
|
metadata=meta,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_bus_progress_callback(
|
||||||
|
bus: MessageBus,
|
||||||
|
msg: InboundMessage,
|
||||||
|
) -> Callable[..., Awaitable[None]]:
|
||||||
|
"""Return the bus progress callback for agent runtime events."""
|
||||||
|
|
||||||
|
async def _publish_progress(
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
tool_hint: bool = False,
|
||||||
|
tool_events: list[dict[str, Any]] | None = None,
|
||||||
|
file_edit_events: list[dict[str, Any]] | None = None,
|
||||||
|
reasoning: bool = False,
|
||||||
|
reasoning_end: bool = False,
|
||||||
|
) -> None:
|
||||||
|
meta = dict(msg.metadata or {})
|
||||||
|
meta["_progress"] = True
|
||||||
|
meta["_tool_hint"] = tool_hint
|
||||||
|
if reasoning:
|
||||||
|
meta["_reasoning_delta"] = True
|
||||||
|
if reasoning_end:
|
||||||
|
meta["_reasoning_end"] = True
|
||||||
|
if tool_events:
|
||||||
|
meta["_tool_events"] = tool_events
|
||||||
|
if file_edit_events:
|
||||||
|
meta["_file_edit_events"] = file_edit_events
|
||||||
|
await bus.publish_outbound(
|
||||||
|
OutboundMessage(
|
||||||
|
channel=msg.channel,
|
||||||
|
chat_id=msg.chat_id,
|
||||||
|
content=content,
|
||||||
|
metadata=meta,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if msg.channel == "websocket":
|
||||||
|
async def _websocket_progress(
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
tool_hint: bool = False,
|
||||||
|
tool_events: list[dict[str, Any]] | None = None,
|
||||||
|
file_edit_events: list[dict[str, Any]] | None = None,
|
||||||
|
reasoning: bool = False,
|
||||||
|
reasoning_end: bool = False,
|
||||||
|
) -> None:
|
||||||
|
await _publish_progress(
|
||||||
|
content,
|
||||||
|
tool_hint=tool_hint,
|
||||||
|
tool_events=tool_events,
|
||||||
|
file_edit_events=file_edit_events,
|
||||||
|
reasoning=reasoning,
|
||||||
|
reasoning_end=reasoning_end,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _websocket_progress
|
||||||
|
|
||||||
|
async def _bus_progress(
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
tool_hint: bool = False,
|
||||||
|
tool_events: list[dict[str, Any]] | None = None,
|
||||||
|
reasoning: bool = False,
|
||||||
|
reasoning_end: bool = False,
|
||||||
|
) -> None:
|
||||||
|
await _publish_progress(
|
||||||
|
content,
|
||||||
|
tool_hint=tool_hint,
|
||||||
|
tool_events=tool_events,
|
||||||
|
reasoning=reasoning,
|
||||||
|
reasoning_end=reasoning_end,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _bus_progress
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class WebuiTurnCoordinator:
|
||||||
|
"""Own the WebUI/WebSocket wire details that hang off AgentLoop turns."""
|
||||||
|
|
||||||
|
bus: MessageBus
|
||||||
|
sessions: SessionManager
|
||||||
|
schedule_background: Callable[[Awaitable[None]], None]
|
||||||
|
_title_contexts: dict[str, LLMRuntime] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def capture_title_context(
|
||||||
|
self,
|
||||||
|
session_key: str,
|
||||||
|
msg: InboundMessage,
|
||||||
|
llm: LLMRuntime,
|
||||||
|
) -> None:
|
||||||
|
if msg.channel == "websocket" and msg.metadata.get("webui") is True:
|
||||||
|
self._title_contexts[session_key] = llm
|
||||||
|
|
||||||
|
def discard(self, session_key: str) -> None:
|
||||||
|
self._title_contexts.pop(session_key, None)
|
||||||
|
|
||||||
|
async def publish_run_status(self, msg: InboundMessage, status: str) -> None:
|
||||||
|
await publish_turn_run_status(self.bus, msg, status)
|
||||||
|
|
||||||
|
async def handle_turn_end(
|
||||||
|
self,
|
||||||
|
msg: InboundMessage,
|
||||||
|
*,
|
||||||
|
session_key: str,
|
||||||
|
latency_ms: int | None,
|
||||||
|
) -> None:
|
||||||
|
if msg.channel != "websocket":
|
||||||
|
return
|
||||||
|
|
||||||
|
turn_metadata: dict[str, Any] = {**msg.metadata, "_turn_end": True}
|
||||||
|
if latency_ms is not None:
|
||||||
|
turn_metadata["latency_ms"] = int(latency_ms)
|
||||||
|
session = self.sessions.get_or_create(session_key)
|
||||||
|
turn_metadata["goal_state"] = goal_state_ws_blob(session.metadata)
|
||||||
|
await self.bus.publish_outbound(OutboundMessage(
|
||||||
|
channel=msg.channel,
|
||||||
|
chat_id=msg.chat_id,
|
||||||
|
content="",
|
||||||
|
metadata=turn_metadata,
|
||||||
|
))
|
||||||
|
self._schedule_title_update(msg, session_key=session_key)
|
||||||
|
|
||||||
|
def _schedule_title_update(self, msg: InboundMessage, *, session_key: str) -> None:
|
||||||
|
title_context = self._title_contexts.pop(session_key, None)
|
||||||
|
if msg.metadata.get("webui") is not True or title_context is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
async def _generate_title_and_notify(
|
||||||
|
title_llm: LLMRuntime = title_context,
|
||||||
|
) -> None:
|
||||||
|
generated = await maybe_generate_webui_title_after_turn(
|
||||||
|
channel=msg.channel,
|
||||||
|
metadata=msg.metadata,
|
||||||
|
sessions=self.sessions,
|
||||||
|
session_key=session_key,
|
||||||
|
provider=title_llm.provider,
|
||||||
|
model=title_llm.model,
|
||||||
|
)
|
||||||
|
if generated:
|
||||||
|
await self.bus.publish_outbound(OutboundMessage(
|
||||||
|
channel=msg.channel,
|
||||||
|
chat_id=msg.chat_id,
|
||||||
|
content="",
|
||||||
|
metadata={
|
||||||
|
**msg.metadata,
|
||||||
|
"_session_updated": True,
|
||||||
|
"_session_update_scope": "metadata",
|
||||||
|
},
|
||||||
|
))
|
||||||
|
|
||||||
|
self.schedule_background(_generate_title_and_notify())
|
||||||
@@ -1,64 +0,0 @@
|
|||||||
---
|
|
||||||
name: create-instance
|
|
||||||
description: "Create a new nanobot instance with separate config and workspace. Use when the user wants to set up a new bot, create a new instance for a different channel, persona, or purpose. Triggers on: create instance, new bot, set up bot, add bot, create telegram/discord/feishu/slack/wechat/wecom/dingtalk/qq/email/matrix/msteams/whatsapp bot, multi-instance setup, inter-agent communication."
|
|
||||||
---
|
|
||||||
|
|
||||||
# Create Instance
|
|
||||||
|
|
||||||
Set up a new nanobot instance with its own config and workspace.
|
|
||||||
|
|
||||||
## Steps
|
|
||||||
|
|
||||||
1. **Collect information** (ask one at a time if not already provided):
|
|
||||||
- **Instance name** (required): short identifier, e.g. `telegram-bot`, `work-slack`
|
|
||||||
- **Channel type** (required): see table below
|
|
||||||
- **Model** (optional): LLM model, defaults to current instance
|
|
||||||
|
|
||||||
2. **Do NOT collect secrets** in the chat (API keys, bot tokens). API keys are automatically inherited from the current instance via `--inherit-config`. Channel-specific tokens must be filled in manually after creation.
|
|
||||||
|
|
||||||
3. **Run the creation script**:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python <skill-dir>/scripts/create_instance.py --name <name> --channel <channel> --inherit-config <current-config>
|
|
||||||
```
|
|
||||||
|
|
||||||
- `<skill-dir>` — the directory containing this SKILL.md
|
|
||||||
- `<current-config>` — current instance's config path, typically `~/.nanobot/config.json`
|
|
||||||
- Optional: `--model <model>`, `--config-dir <path>`
|
|
||||||
|
|
||||||
**Exec tool constraints:**
|
|
||||||
- Use forward-slash paths (works on all platforms)
|
|
||||||
- Do not wrap paths in quotes
|
|
||||||
- Do not use `cd`; pass the full script path directly
|
|
||||||
|
|
||||||
4. **Report results** to the user:
|
|
||||||
- Config and workspace paths (script outputs them)
|
|
||||||
- Required fields to fill in (script lists them)
|
|
||||||
- Start command: `nanobot gateway --config <config-path>`
|
|
||||||
|
|
||||||
## Available Channels
|
|
||||||
|
|
||||||
| Channel | Key | Required Fields |
|
|
||||||
|---------|-----|-----------------|
|
|
||||||
| Telegram | `telegram` | token |
|
|
||||||
| Discord | `discord` | token |
|
|
||||||
| Feishu / Lark | `feishu` | app_id, app_secret |
|
|
||||||
| DingTalk | `dingtalk` | client_id, client_secret |
|
|
||||||
| Slack | `slack` | bot_token, app_token |
|
|
||||||
| WeCom | `wecom` | bot_id, secret |
|
|
||||||
| WeChat OA | `weixin` | token |
|
|
||||||
| WhatsApp | `whatsapp` | bridge_token |
|
|
||||||
| QQ | `qq` | app_id, secret |
|
|
||||||
| Email | `email` | imap_host, imap_username, imap_password, smtp_host, smtp_username, smtp_password, from_address |
|
|
||||||
| Matrix | `matrix` | user_id, password or access_token |
|
|
||||||
| MS Teams | `msteams` | app_id, app_password, tenant_id |
|
|
||||||
| MoChat | `mochat` | claw_token |
|
|
||||||
| WebSocket | `websocket` | token |
|
|
||||||
|
|
||||||
For detailed channel configuration including optional fields, see `references/channels.md`.
|
|
||||||
|
|
||||||
## Troubleshooting
|
|
||||||
|
|
||||||
- **"Unknown channel"**: Channel name must match the Key column exactly. Run the script without arguments to see usage.
|
|
||||||
- **"Config already exists"**: Use a different `--name` or `--config-dir` to create in a new location.
|
|
||||||
- **Port conflicts**: The script auto-assigns free ports for gateway and API if defaults are in use.
|
|
||||||
@@ -1,195 +0,0 @@
|
|||||||
# Channel Configuration Reference
|
|
||||||
|
|
||||||
Detailed configuration for each supported channel.
|
|
||||||
|
|
||||||
## Field Types
|
|
||||||
|
|
||||||
- **Required**: defaults to empty string `""`, must be filled in before the instance can start
|
|
||||||
- **Optional**: has a sensible default, can be customized
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## telegram
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `token` — Bot token from @BotFather
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `proxy` — HTTP proxy URL
|
|
||||||
- `group_policy` — `"open"` (all messages) or `"mention"` (default, only when @mentioned)
|
|
||||||
- `streaming` — Enable streaming responses (default: true)
|
|
||||||
- `reply_to_message` — Reply to the triggering message (default: false)
|
|
||||||
- `react_emoji` — Emoji for "thinking" reaction (default: `"eyes"`)
|
|
||||||
- `inline_keyboards` — Enable inline keyboard buttons (default: false)
|
|
||||||
|
|
||||||
## discord
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `token` — Bot token from Discord Developer Portal
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `allow_channels` — Restrict to specific channel IDs
|
|
||||||
- `group_policy` — `"mention"` (default) or `"open"`
|
|
||||||
- `streaming` — Enable streaming (default: true)
|
|
||||||
- `proxy` — HTTP proxy URL
|
|
||||||
- `intents` — Discord gateway intents (default: 37377)
|
|
||||||
- `read_receipt_emoji` — Emoji for read receipt
|
|
||||||
- `working_emoji` — Emoji for "working" indicator
|
|
||||||
|
|
||||||
## feishu
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `app_id` — Feishu app ID
|
|
||||||
- `app_secret` — Feishu app secret
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `encrypt_key` — Event encryption key
|
|
||||||
- `verification_token` — Event verification token
|
|
||||||
- `domain` — `"feishu"` (default) or `"lark"`
|
|
||||||
- `group_policy` — `"mention"` (default) or `"open"`
|
|
||||||
- `streaming` — Enable streaming (default: true)
|
|
||||||
|
|
||||||
## dingtalk
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `client_id` — DingTalk app client ID
|
|
||||||
- `client_secret` — DingTalk app client secret
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `allow_from` — Allowed user IDs
|
|
||||||
|
|
||||||
## slack
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `bot_token` — Bot OAuth token (`xoxb-...`)
|
|
||||||
- `app_token` — App-level token (`xapp-...`)
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `mode` — `"socket"` (default, Socket Mode) or `"webhook"`
|
|
||||||
- `reply_in_thread` — Reply in thread (default: true)
|
|
||||||
- `react_emoji` — "thinking" emoji (default: `"eyes"`)
|
|
||||||
- `done_emoji` — "done" emoji (default: `"white_check_mark"`)
|
|
||||||
- `group_policy` — `"mention"` (default) or `"open"`
|
|
||||||
- `dm.enabled` — Enable DM support
|
|
||||||
- `dm.policy` — DM policy
|
|
||||||
- `dm.allow_from` — Allowed DM users
|
|
||||||
|
|
||||||
## wecom
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `bot_id` — WeCom bot ID
|
|
||||||
- `secret` — WeCom bot secret
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `allow_from` — Allowed users
|
|
||||||
- `welcome_message` — Welcome message for new chats
|
|
||||||
|
|
||||||
## weixin
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `token` — WeChat Official Account token
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `base_url` — API base URL
|
|
||||||
- `cdn_base_url` — CDN base URL
|
|
||||||
- `state_dir` — State persistence directory
|
|
||||||
- `poll_timeout` — Long polling timeout
|
|
||||||
|
|
||||||
## whatsapp
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `bridge_token` — WhatsApp bridge token (auto-generated if absent)
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `bridge_url` — Bridge WebSocket URL (default: `"ws://localhost:3001"`)
|
|
||||||
- `group_policy` — `"open"` (default) or `"mention"`
|
|
||||||
|
|
||||||
## qq
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `app_id` — QQ bot app ID
|
|
||||||
- `secret` — QQ bot secret
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `msg_format` — `"plain"` or `"markdown"`
|
|
||||||
- `ack_message` — Acknowledgment message text
|
|
||||||
- `media_dir` — Media file directory
|
|
||||||
|
|
||||||
## email
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `imap_host` — IMAP server hostname
|
|
||||||
- `imap_username` — IMAP login username
|
|
||||||
- `imap_password` — IMAP login password
|
|
||||||
- `smtp_host` — SMTP server hostname
|
|
||||||
- `smtp_username` — SMTP login username
|
|
||||||
- `smtp_password` — SMTP login password
|
|
||||||
- `from_address` — Sender email address
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `imap_port` — IMAP port (default: 993)
|
|
||||||
- `smtp_port` — SMTP port (default: 587)
|
|
||||||
- `imap_use_ssl` — Use SSL for IMAP (default: true)
|
|
||||||
- `smtp_use_tls` — Use TLS for SMTP (default: true)
|
|
||||||
- `poll_interval_seconds` — Polling interval (default: 30)
|
|
||||||
- `mark_seen` — Mark emails as read (default: true)
|
|
||||||
- `max_body_chars` — Max email body length (default: 12000)
|
|
||||||
- `subject_prefix` — Reply subject prefix (default: `"Re: "`)
|
|
||||||
- `verify_dkim` — Verify DKIM signatures (default: true)
|
|
||||||
- `verify_spf` — Verify SPF records (default: true)
|
|
||||||
- `allowed_attachment_types` — Allowed file extensions
|
|
||||||
- `max_attachment_size` — Max attachment size in bytes
|
|
||||||
- `consent_granted` — Must be set to `true` for the channel to start (default: false)
|
|
||||||
- `auto_reply_enabled` — Enable auto-reply (default: true)
|
|
||||||
|
|
||||||
## matrix
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `user_id` — Matrix user ID (e.g. `@bot:matrix.org`)
|
|
||||||
- `password` or `access_token` — Login password OR access token
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `homeserver` — Homeserver URL (default: `"https://matrix.org"`)
|
|
||||||
- `device_id` — Device ID
|
|
||||||
- `e2eeEnabled` — Enable end-to-end encryption (default: true)
|
|
||||||
- `group_policy` — `"open"`, `"mention"`, or `"allowlist"`
|
|
||||||
- `streaming` — Enable streaming (default: false)
|
|
||||||
- `max_media_bytes` — Max media file size (default: 20MB)
|
|
||||||
|
|
||||||
## msteams
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `app_id` — Azure AD app ID
|
|
||||||
- `app_password` — Azure AD app password/secret
|
|
||||||
- `tenant_id` — Azure AD tenant ID
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `host` — Listen host (default: `"0.0.0.0"`)
|
|
||||||
- `port` — Listen port (default: 3978)
|
|
||||||
- `reply_in_thread` — Reply in thread (default: true)
|
|
||||||
- `validate_inbound_auth` — Validate incoming auth (default: true)
|
|
||||||
|
|
||||||
## mochat
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `claw_token` — MoChat Claw token
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `base_url` — API base URL
|
|
||||||
- `socket_url` — WebSocket URL
|
|
||||||
- `refresh_interval_ms` — Refresh interval in ms
|
|
||||||
- `watch_timeout_ms` — Watch timeout in ms
|
|
||||||
|
|
||||||
## websocket
|
|
||||||
|
|
||||||
Built-in WebSocket channel for programmatic access.
|
|
||||||
|
|
||||||
**Required:**
|
|
||||||
- `token` — Authentication token (enabled by default; set `websocket_requires_token: false` to disable)
|
|
||||||
|
|
||||||
**Notable optional:**
|
|
||||||
- `host` — Listen host (default: `"127.0.0.1"`)
|
|
||||||
- `port` — Listen port (default: 8765)
|
|
||||||
- `allow_from` — Allowed origins (default: `["*"]`)
|
|
||||||
- `streaming` — Enable streaming (default: true)
|
|
||||||
|
|
||||||
@@ -1,252 +0,0 @@
|
|||||||
#!/usr/bin/env python3
|
|
||||||
"""Create a new nanobot instance with a dedicated config and workspace.
|
|
||||||
|
|
||||||
Usage:
|
|
||||||
create_instance.py --name <name> --channel <channel> [--model <model>] [--config-dir <dir>]
|
|
||||||
|
|
||||||
Examples:
|
|
||||||
create_instance.py --name telegram-bot --channel telegram
|
|
||||||
create_instance.py --name discord-bot --channel discord --model deepseek/deepseek-chat
|
|
||||||
create_instance.py --name my-bot --channel telegram --config-dir ~/.nanobot-custom
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import json
|
|
||||||
import re
|
|
||||||
import socket
|
|
||||||
import sys
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_name(name: str) -> str:
|
|
||||||
"""Normalize and validate instance name."""
|
|
||||||
name = name.strip().lower()
|
|
||||||
name = re.sub(r"[^a-z0-9-]", "-", name)
|
|
||||||
name = re.sub(r"-{2,}", "-", name)
|
|
||||||
name = name.strip("-")
|
|
||||||
if not name:
|
|
||||||
print("[ERROR] Instance name must contain at least one letter or digit.", file=sys.stderr)
|
|
||||||
sys.exit(1)
|
|
||||||
if len(name) > 64:
|
|
||||||
print(f"[ERROR] Instance name too long ({len(name)} chars, max 64).", file=sys.stderr)
|
|
||||||
sys.exit(1)
|
|
||||||
return name
|
|
||||||
|
|
||||||
|
|
||||||
def _get_available_channels() -> list[str]:
|
|
||||||
"""Get list of available channel names without importing channel classes."""
|
|
||||||
from nanobot.channels.registry import discover_channel_names
|
|
||||||
|
|
||||||
return discover_channel_names()
|
|
||||||
|
|
||||||
|
|
||||||
def _run_onboard(config_path: Path, workspace: Path) -> None:
|
|
||||||
"""Create skeleton config + workspace using nanobot's programmatic API."""
|
|
||||||
from nanobot.cli.commands import _onboard_plugins
|
|
||||||
from nanobot.config.loader import save_config, set_config_path
|
|
||||||
from nanobot.config.paths import get_workspace_path
|
|
||||||
from nanobot.config.schema import Config
|
|
||||||
from nanobot.utils.helpers import sync_workspace_templates
|
|
||||||
|
|
||||||
config = Config()
|
|
||||||
config.agents.defaults.workspace = str(workspace)
|
|
||||||
set_config_path(config_path)
|
|
||||||
save_config(config, config_path)
|
|
||||||
_onboard_plugins(config_path)
|
|
||||||
|
|
||||||
workspace_path = get_workspace_path(config.workspace_path)
|
|
||||||
if not workspace_path.exists():
|
|
||||||
workspace_path.mkdir(parents=True, exist_ok=True)
|
|
||||||
sync_workspace_templates(workspace_path)
|
|
||||||
|
|
||||||
|
|
||||||
def _patch_config(
|
|
||||||
config_path: Path,
|
|
||||||
*,
|
|
||||||
channel: str,
|
|
||||||
workspace: Path,
|
|
||||||
model: str | None,
|
|
||||||
name: str | None = None,
|
|
||||||
inherit_config_path: Path | None = None,
|
|
||||||
) -> dict:
|
|
||||||
"""Patch the generated config: enable channel, set workspace, optionally set model."""
|
|
||||||
data = json.loads(config_path.read_text(encoding="utf-8"))
|
|
||||||
|
|
||||||
# Inherit providers and model from current instance
|
|
||||||
if inherit_config_path and inherit_config_path.exists():
|
|
||||||
try:
|
|
||||||
src = json.loads(inherit_config_path.read_text(encoding="utf-8"))
|
|
||||||
|
|
||||||
# Inherit providers (API keys, api_base, etc.)
|
|
||||||
src_providers = src.get("providers", {})
|
|
||||||
if src_providers:
|
|
||||||
data.setdefault("providers", {})
|
|
||||||
for key, val in src_providers.items():
|
|
||||||
if isinstance(val, dict) and val.get("apiKey"):
|
|
||||||
data["providers"][key] = val
|
|
||||||
|
|
||||||
# Inherit model if not explicitly overridden
|
|
||||||
if not model:
|
|
||||||
parent_model = src.get("agents", {}).get("defaults", {}).get("model")
|
|
||||||
if parent_model:
|
|
||||||
model = parent_model
|
|
||||||
|
|
||||||
except Exception as exc:
|
|
||||||
print(f"[WARN] Could not inherit from {inherit_config_path}: {exc}", file=sys.stderr)
|
|
||||||
|
|
||||||
# Set workspace and model
|
|
||||||
data.setdefault("agents", {}).setdefault("defaults", {})
|
|
||||||
data["agents"]["defaults"]["workspace"] = str(workspace)
|
|
||||||
if model:
|
|
||||||
data["agents"]["defaults"]["model"] = model
|
|
||||||
|
|
||||||
# Enable the target channel
|
|
||||||
channels = data.setdefault("channels", {})
|
|
||||||
if channel in channels and isinstance(channels[channel], dict):
|
|
||||||
channels[channel]["enabled"] = True
|
|
||||||
else:
|
|
||||||
channels[channel] = {"enabled": True}
|
|
||||||
|
|
||||||
# Auto-assign ports if defaults are already in use
|
|
||||||
_assign_free_ports(data)
|
|
||||||
|
|
||||||
# Validate with Pydantic, then save
|
|
||||||
from nanobot.config.schema import Config
|
|
||||||
|
|
||||||
Config.model_validate(data)
|
|
||||||
config_path.write_text(json.dumps(data, indent=2, ensure_ascii=False), encoding="utf-8")
|
|
||||||
return data
|
|
||||||
|
|
||||||
|
|
||||||
def _is_port_in_use(port: int, host: str = "127.0.0.1") -> bool:
|
|
||||||
"""Check if a port is already in use."""
|
|
||||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
|
||||||
try:
|
|
||||||
s.bind((host, port))
|
|
||||||
return False
|
|
||||||
except OSError:
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
def _find_free_port(start: int, host: str = "127.0.0.1", max_tries: int = 100) -> int:
|
|
||||||
"""Find the first free port starting from `start`."""
|
|
||||||
for port in range(start, start + max_tries):
|
|
||||||
if not _is_port_in_use(port, host):
|
|
||||||
return port
|
|
||||||
# OS-level fallback: ask the kernel for an ephemeral port
|
|
||||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
|
||||||
s.bind((host, 0))
|
|
||||||
return s.getsockname()[1]
|
|
||||||
|
|
||||||
|
|
||||||
def _assign_free_ports(data: dict) -> None:
|
|
||||||
"""If default gateway or API ports are in use, assign free ones."""
|
|
||||||
from nanobot.config.schema import ApiConfig, GatewayConfig
|
|
||||||
|
|
||||||
defaults = [
|
|
||||||
("gateway", GatewayConfig()),
|
|
||||||
("api", ApiConfig()),
|
|
||||||
]
|
|
||||||
for key, default_cfg in defaults:
|
|
||||||
section = data.setdefault(key, {})
|
|
||||||
port = section.get("port", default_cfg.port)
|
|
||||||
host = section.get("host", default_cfg.host)
|
|
||||||
if _is_port_in_use(port, host):
|
|
||||||
section["port"] = _find_free_port(port + 1, host)
|
|
||||||
|
|
||||||
|
|
||||||
def _get_channel_required_fields(channel: str) -> list[str]:
|
|
||||||
"""Inspect a channel's default config and list fields that are empty strings."""
|
|
||||||
try:
|
|
||||||
from nanobot.channels.registry import load_channel_class
|
|
||||||
|
|
||||||
cls = load_channel_class(channel)
|
|
||||||
default = cls.default_config()
|
|
||||||
return sorted(k for k, v in default.items() if isinstance(v, str) and v == "" and k != "enabled")
|
|
||||||
except Exception as exc:
|
|
||||||
print(f"[WARN] Could not inspect channel '{channel}' defaults: {exc}", file=sys.stderr)
|
|
||||||
return []
|
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
|
||||||
parser = argparse.ArgumentParser(
|
|
||||||
description="Create a new nanobot instance.",
|
|
||||||
)
|
|
||||||
parser.add_argument("--name", required=True, help="Instance name (e.g. telegram-bot)")
|
|
||||||
parser.add_argument("--channel", required=True, help="Channel type (e.g. telegram, discord)")
|
|
||||||
parser.add_argument("--model", default=None, help="LLM model (default: same as current instance)")
|
|
||||||
parser.add_argument(
|
|
||||||
"--config-dir",
|
|
||||||
default=None,
|
|
||||||
help="Config directory (default: ~/.nanobot-{name})",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--inherit-config",
|
|
||||||
default=None,
|
|
||||||
help="Path to current instance's config.json to copy API keys from",
|
|
||||||
)
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
# Validate name
|
|
||||||
name = _validate_name(args.name)
|
|
||||||
|
|
||||||
# Validate channel
|
|
||||||
available = _get_available_channels()
|
|
||||||
if args.channel not in available:
|
|
||||||
print(f"[ERROR] Unknown channel: {args.channel}", file=sys.stderr)
|
|
||||||
print(f"Available channels: {', '.join(sorted(available))}", file=sys.stderr)
|
|
||||||
sys.exit(1)
|
|
||||||
|
|
||||||
# Resolve paths
|
|
||||||
home = Path.home()
|
|
||||||
config_dir = Path(args.config_dir).expanduser().resolve() if args.config_dir else home / f".nanobot-{name}"
|
|
||||||
config_path = config_dir / "config.json"
|
|
||||||
workspace = config_dir / "workspace"
|
|
||||||
|
|
||||||
# Check for duplicate
|
|
||||||
if config_path.exists():
|
|
||||||
print(f"[ERROR] Config already exists at {config_path}", file=sys.stderr)
|
|
||||||
print("Delete it first or use a different --config-dir.", file=sys.stderr)
|
|
||||||
sys.exit(1)
|
|
||||||
|
|
||||||
print(f"Creating instance '{name}'...")
|
|
||||||
print(f" Config dir: {config_dir}")
|
|
||||||
print(f" Workspace: {workspace}")
|
|
||||||
print(f" Channel: {args.channel}")
|
|
||||||
if args.model:
|
|
||||||
print(f" Model: {args.model}")
|
|
||||||
|
|
||||||
# Run onboard
|
|
||||||
_run_onboard(config_path, workspace)
|
|
||||||
|
|
||||||
# Patch config
|
|
||||||
inherit_path = Path(args.inherit_config).expanduser().resolve() if args.inherit_config else None
|
|
||||||
_patch_config(
|
|
||||||
config_path,
|
|
||||||
channel=args.channel,
|
|
||||||
workspace=workspace,
|
|
||||||
model=args.model,
|
|
||||||
name=name,
|
|
||||||
inherit_config_path=inherit_path,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Report
|
|
||||||
print(f"\n[OK] Instance '{name}' created successfully.")
|
|
||||||
print(f" Config: {config_path}")
|
|
||||||
print(f" Workspace: {workspace}")
|
|
||||||
|
|
||||||
# List fields the user needs to fill in
|
|
||||||
required_fields = _get_channel_required_fields(args.channel)
|
|
||||||
if required_fields:
|
|
||||||
print(f"\n[IMPORTANT] Edit {config_path} and fill in these fields:")
|
|
||||||
for field in required_fields:
|
|
||||||
print(f" - channels.{args.channel}.{field}")
|
|
||||||
|
|
||||||
print(f"\nTo start the instance:")
|
|
||||||
print(f" nanobot gateway --config {config_path}")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -15,7 +15,7 @@ If the `generate_image` tool is not available in the current tool list, tell the
|
|||||||
- Image editing: pass the saved artifact path or user image path in `reference_images`.
|
- Image editing: pass the saved artifact path or user image path in `reference_images`.
|
||||||
- Iterative edits in the same conversation: prefer the most recent generated image artifact if the user says things like "make it brighter", "change the background", or "try another version".
|
- Iterative edits in the same conversation: prefer the most recent generated image artifact if the user says things like "make it brighter", "change the background", or "try another version".
|
||||||
- Ambiguous edits: ask a short clarifying question if multiple recent images could be the target.
|
- Ambiguous edits: ask a short clarifying question if multiple recent images could be the target.
|
||||||
- In the current chat, do not call `message` just to announce or resend generated images. The runtime attaches images from `generate_image` to the final assistant reply automatically.
|
- After generating images, call the `message` tool with the artifact paths in the `media` parameter to deliver them to the user.
|
||||||
|
|
||||||
## Prompt Rules
|
## Prompt Rules
|
||||||
|
|
||||||
@@ -42,52 +42,6 @@ For follow-up edits, pass the prior artifact `path` to `reference_images`. If th
|
|||||||
|
|
||||||
Do not include internal replay markers such as `[Message Time: ...]`, `[image: /local/path]`, `generate_image(...)`, or `message(...)` in user-facing replies.
|
Do not include internal replay markers such as `[Message Time: ...]`, `[image: /local/path]`, `generate_image(...)`, or `message(...)` in user-facing replies.
|
||||||
|
|
||||||
## Provider Notes
|
|
||||||
|
|
||||||
Do not ask users to paste API keys into chat. If configuration is needed, describe the fields; LLM provider and BYOK changes are hot-reloaded for new turns.
|
|
||||||
|
|
||||||
For OpenRouter, the image tool expects:
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"providers": {
|
|
||||||
"openrouter": {
|
|
||||||
"apiKey": "sk-or-..."
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"tools": {
|
|
||||||
"imageGeneration": {
|
|
||||||
"enabled": true,
|
|
||||||
"provider": "openrouter",
|
|
||||||
"model": "openai/gpt-5.4-image-2"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
For AIHubMix, the image tool expects:
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"providers": {
|
|
||||||
"aihubmix": {
|
|
||||||
"apiKey": "sk-..."
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"tools": {
|
|
||||||
"imageGeneration": {
|
|
||||||
"enabled": true,
|
|
||||||
"provider": "aihubmix",
|
|
||||||
"model": "gpt-image-2-free"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
AIHubMix `gpt-image-2-free` uses AIHubMix's unified predictions endpoint internally (`/v1/models/openai/gpt-image-2-free/predictions`), not the OpenAI Images `/v1/images/generations` endpoint. If it fails with "Incorrect model ID", do not assume the key lacks permission until the provider config, model name, and gateway restart have been checked.
|
|
||||||
|
|
||||||
`providers.aihubmix.extraBody` can be used for provider-specific options. For example, `"extraBody": {"quality": "low"}` is optional but can make `gpt-image-2-free` faster and less likely to time out.
|
|
||||||
|
|
||||||
## Examples
|
## Examples
|
||||||
|
|
||||||
Generate a new image:
|
Generate a new image:
|
||||||
|
|||||||
@@ -30,5 +30,5 @@ Output is rendered in a terminal. Avoid markdown headings and tables. Use plain
|
|||||||
|
|
||||||
Reply directly with text for the current conversation. Do not use the 'message' tool for normal replies in the current chat.
|
Reply directly with text for the current conversation. Do not use the 'message' tool for normal replies in the current chat.
|
||||||
When you need to call tools before answering, do not include the final user-visible answer in the same assistant message as the tool calls. Wait for the tool results, then answer once.
|
When you need to call tools before answering, do not include the final user-visible answer in the same assistant message as the tool calls. Wait for the tool results, then answer once.
|
||||||
Use the 'message' tool only for proactive sends, cross-channel delivery, or explicitly sending existing local files as attachments. When a tool such as 'generate_image' creates user-visible media, the runtime attaches those artifacts to the final assistant reply automatically, so do not call 'message' just to announce or resend them.
|
Use the 'message' tool only for proactive sends, cross-channel delivery, or explicitly sending existing local files as attachments. When 'generate_image' creates images, call 'message' with the artifact paths in the 'media' parameter to deliver them to the user.
|
||||||
To send an existing local file that was not automatically attached by another tool, call 'message' with the 'media' parameter. Do NOT use read_file to "send" a file — reading a file only shows its content to you, it does NOT deliver the file to the user. Example: message(content="Here is the document", channel="telegram", chat_id="...", media=["/path/to/file.pdf"])
|
To send an existing local file that was not automatically attached by another tool, call 'message' with the 'media' parameter. Do NOT use read_file to "send" a file — reading a file only shows its content to you, it does NOT deliver the file to the user. Example: message(content="Here is the document", channel="telegram", chat_id="...", media=["/path/to/file.pdf"])
|
||||||
|
|||||||
@@ -1,6 +1,42 @@
|
|||||||
"""Utility functions for nanobot."""
|
"""Utility functions for nanobot."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import sys
|
||||||
|
from importlib import import_module
|
||||||
|
from types import ModuleType
|
||||||
|
|
||||||
from nanobot.utils.helpers import ensure_dir
|
from nanobot.utils.helpers import ensure_dir
|
||||||
from nanobot.utils.path import abbreviate_path
|
from nanobot.utils.path import abbreviate_path
|
||||||
|
|
||||||
__all__ = ["ensure_dir", "abbreviate_path"]
|
__all__ = ["ensure_dir", "abbreviate_path"]
|
||||||
|
|
||||||
|
|
||||||
|
class _LazyModuleAlias(ModuleType):
|
||||||
|
def __init__(self, name: str, target: str) -> None:
|
||||||
|
super().__init__(name)
|
||||||
|
self.__dict__["_target"] = target
|
||||||
|
|
||||||
|
def _load(self) -> ModuleType:
|
||||||
|
module = import_module(self.__dict__["_target"])
|
||||||
|
sys.modules[self.__name__] = module
|
||||||
|
return module
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> object:
|
||||||
|
return getattr(self._load(), name)
|
||||||
|
|
||||||
|
def __dir__(self) -> list[str]:
|
||||||
|
return sorted(set(super().__dir__()) | set(dir(self._load())))
|
||||||
|
|
||||||
|
|
||||||
|
_LEGACY_MODULE_ALIASES = {
|
||||||
|
"webui_thread_disk": "nanobot.webui.thread_disk",
|
||||||
|
"webui_transcript": "nanobot.webui.transcript",
|
||||||
|
"webui_turn_helpers": "nanobot.session.webui_turns",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _legacy_name, _target_name in _LEGACY_MODULE_ALIASES.items():
|
||||||
|
sys.modules.setdefault(
|
||||||
|
f"{__name__}.{_legacy_name}",
|
||||||
|
_LazyModuleAlias(f"{__name__}.{_legacy_name}", _target_name),
|
||||||
|
)
|
||||||
|
|||||||
@@ -21,8 +21,6 @@ _MIME_EXTENSIONS = {
|
|||||||
"image/webp": ".webp",
|
"image/webp": ".webp",
|
||||||
"image/gif": ".gif",
|
"image/gif": ".gif",
|
||||||
}
|
}
|
||||||
_GENERATE_IMAGE_TOOL_NAME = "generate_image"
|
|
||||||
|
|
||||||
|
|
||||||
class ArtifactError(ValueError):
|
class ArtifactError(ValueError):
|
||||||
"""Raised when an artifact cannot be safely decoded or stored."""
|
"""Raised when an artifact cannot be safely decoded or stored."""
|
||||||
@@ -115,48 +113,10 @@ def generated_image_tool_result(artifacts: list[dict[str, Any]]) -> str:
|
|||||||
"artifacts": artifacts,
|
"artifacts": artifacts,
|
||||||
"next_step": (
|
"next_step": (
|
||||||
"Use these artifact paths as reference_images for follow-up edits. "
|
"Use these artifact paths as reference_images for follow-up edits. "
|
||||||
"For the current chat, reply naturally; the runtime attaches generated images automatically. "
|
"Call the message tool with the artifact paths in the media parameter "
|
||||||
"Do not call message just to announce or resend them. Keep raw paths internal unless the user asks for debug details."
|
"to deliver the images to the user. Keep raw paths internal unless the "
|
||||||
|
"user asks for debug details."
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
ensure_ascii=False,
|
ensure_ascii=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _extract_text_payload(content: Any) -> str | None:
|
|
||||||
if isinstance(content, str):
|
|
||||||
return content
|
|
||||||
if isinstance(content, list):
|
|
||||||
parts: list[str] = []
|
|
||||||
for block in content:
|
|
||||||
if isinstance(block, dict) and isinstance(block.get("text"), str):
|
|
||||||
parts.append(block["text"])
|
|
||||||
return "\n".join(parts) if parts else None
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def generated_image_paths_from_messages(messages: list[dict[str, Any]]) -> list[str]:
|
|
||||||
"""Collect generated image artifact paths from generate_image tool results."""
|
|
||||||
paths: list[str] = []
|
|
||||||
seen: set[str] = set()
|
|
||||||
for message in messages:
|
|
||||||
if message.get("role") != "tool" or message.get("name") != _GENERATE_IMAGE_TOOL_NAME:
|
|
||||||
continue
|
|
||||||
payload = _extract_text_payload(message.get("content"))
|
|
||||||
if not payload:
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
data = json.loads(payload)
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
continue
|
|
||||||
artifacts = data.get("artifacts") if isinstance(data, dict) else None
|
|
||||||
if not isinstance(artifacts, list):
|
|
||||||
continue
|
|
||||||
for artifact in artifacts:
|
|
||||||
if not isinstance(artifact, dict):
|
|
||||||
continue
|
|
||||||
path = artifact.get("path")
|
|
||||||
if isinstance(path, str) and path and path not in seen:
|
|
||||||
paths.append(path)
|
|
||||||
seen.add(path)
|
|
||||||
return paths
|
|
||||||
|
|||||||
@@ -0,0 +1,780 @@
|
|||||||
|
"""File-edit activity helpers for WebUI progress events."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import difflib
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Awaitable, Callable
|
||||||
|
|
||||||
|
|
||||||
|
TRACKED_FILE_EDIT_TOOLS = frozenset({"write_file", "edit_file", "notebook_edit"})
|
||||||
|
_MAX_SNAPSHOT_BYTES = 2 * 1024 * 1024
|
||||||
|
_LIVE_EMIT_INTERVAL_S = 0.18
|
||||||
|
_LIVE_EMIT_LINE_STEP = 24
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class FileSnapshot:
|
||||||
|
path: Path
|
||||||
|
exists: bool
|
||||||
|
text: str | None
|
||||||
|
unreadable: bool = False
|
||||||
|
binary: bool = False
|
||||||
|
oversized: bool = False
|
||||||
|
|
||||||
|
@property
|
||||||
|
def countable(self) -> bool:
|
||||||
|
return (
|
||||||
|
self.text is not None
|
||||||
|
and not self.binary
|
||||||
|
and not self.oversized
|
||||||
|
and not self.unreadable
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class FileEditTracker:
|
||||||
|
call_id: str
|
||||||
|
tool: str
|
||||||
|
path: Path
|
||||||
|
display_path: str
|
||||||
|
before: FileSnapshot
|
||||||
|
|
||||||
|
|
||||||
|
def is_file_edit_tool(tool_name: str | None) -> bool:
|
||||||
|
return bool(tool_name) and tool_name in TRACKED_FILE_EDIT_TOOLS
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_file_edit_path(
|
||||||
|
tool: Any,
|
||||||
|
workspace: Path | None,
|
||||||
|
params: dict[str, Any] | None,
|
||||||
|
) -> Path | None:
|
||||||
|
"""Resolve the target file path after tool argument preparation."""
|
||||||
|
if not isinstance(params, dict):
|
||||||
|
return None
|
||||||
|
raw_path = params.get("path")
|
||||||
|
if not isinstance(raw_path, str) or not raw_path.strip():
|
||||||
|
return None
|
||||||
|
resolver = getattr(tool, "_resolve", None)
|
||||||
|
if callable(resolver):
|
||||||
|
try:
|
||||||
|
resolved = resolver(raw_path)
|
||||||
|
if isinstance(resolved, Path):
|
||||||
|
return resolved
|
||||||
|
if resolved:
|
||||||
|
return Path(resolved)
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
if workspace is None:
|
||||||
|
return Path(raw_path).expanduser().resolve()
|
||||||
|
return (workspace / raw_path).expanduser().resolve()
|
||||||
|
|
||||||
|
|
||||||
|
def display_file_edit_path(path: Path, workspace: Path | None) -> str:
|
||||||
|
if workspace is not None:
|
||||||
|
try:
|
||||||
|
return path.resolve().relative_to(workspace.resolve()).as_posix()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return path.as_posix()
|
||||||
|
|
||||||
|
|
||||||
|
def read_file_snapshot(path: Path, *, max_bytes: int = _MAX_SNAPSHOT_BYTES) -> FileSnapshot:
|
||||||
|
try:
|
||||||
|
if not path.exists() or not path.is_file():
|
||||||
|
return FileSnapshot(path=path, exists=False, text="")
|
||||||
|
size = path.stat().st_size
|
||||||
|
if size > max_bytes:
|
||||||
|
return FileSnapshot(path=path, exists=True, text=None, oversized=True)
|
||||||
|
raw = path.read_bytes()
|
||||||
|
except OSError:
|
||||||
|
return FileSnapshot(path=path, exists=path.exists(), text=None, unreadable=True)
|
||||||
|
if b"\x00" in raw:
|
||||||
|
return FileSnapshot(path=path, exists=True, text=None, binary=True)
|
||||||
|
try:
|
||||||
|
text = raw.decode("utf-8")
|
||||||
|
except UnicodeDecodeError:
|
||||||
|
return FileSnapshot(path=path, exists=True, text=None, binary=True)
|
||||||
|
return FileSnapshot(path=path, exists=True, text=text.replace("\r\n", "\n"))
|
||||||
|
|
||||||
|
|
||||||
|
def line_diff_stats(before: str | None, after: str | None) -> tuple[int, int]:
|
||||||
|
"""Return ``(added, deleted)`` for a UTF-8 text line-level diff."""
|
||||||
|
if before is None or after is None:
|
||||||
|
return 0, 0
|
||||||
|
if before == "":
|
||||||
|
return _text_line_count(after), 0
|
||||||
|
before_lines = before.replace("\r\n", "\n").splitlines()
|
||||||
|
after_lines = after.replace("\r\n", "\n").splitlines()
|
||||||
|
added = 0
|
||||||
|
deleted = 0
|
||||||
|
matcher = difflib.SequenceMatcher(a=before_lines, b=after_lines, autojunk=False)
|
||||||
|
for tag, i1, i2, j1, j2 in matcher.get_opcodes():
|
||||||
|
if tag == "equal":
|
||||||
|
continue
|
||||||
|
if tag in ("replace", "delete"):
|
||||||
|
deleted += i2 - i1
|
||||||
|
if tag in ("replace", "insert"):
|
||||||
|
added += j2 - j1
|
||||||
|
return added, deleted
|
||||||
|
|
||||||
|
|
||||||
|
def _text_line_count(text: str) -> int:
|
||||||
|
if not text:
|
||||||
|
return 0
|
||||||
|
line_count = 0
|
||||||
|
last_was_newline = False
|
||||||
|
last_was_cr = False
|
||||||
|
for ch in text:
|
||||||
|
if ch == "\r":
|
||||||
|
line_count += 1
|
||||||
|
last_was_newline = True
|
||||||
|
last_was_cr = True
|
||||||
|
elif ch == "\n":
|
||||||
|
if not last_was_cr:
|
||||||
|
line_count += 1
|
||||||
|
last_was_newline = True
|
||||||
|
last_was_cr = False
|
||||||
|
else:
|
||||||
|
last_was_newline = False
|
||||||
|
last_was_cr = False
|
||||||
|
return line_count if last_was_newline else line_count + 1
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_file_edit_tracker(
|
||||||
|
*,
|
||||||
|
call_id: str,
|
||||||
|
tool_name: str,
|
||||||
|
tool: Any,
|
||||||
|
workspace: Path | None,
|
||||||
|
params: dict[str, Any] | None,
|
||||||
|
) -> FileEditTracker | None:
|
||||||
|
if not is_file_edit_tool(tool_name):
|
||||||
|
return None
|
||||||
|
path = resolve_file_edit_path(tool, workspace, params)
|
||||||
|
if path is None:
|
||||||
|
return None
|
||||||
|
before = read_file_snapshot(path)
|
||||||
|
return FileEditTracker(
|
||||||
|
call_id=str(call_id or ""),
|
||||||
|
tool=tool_name,
|
||||||
|
path=path,
|
||||||
|
display_path=display_file_edit_path(path, workspace),
|
||||||
|
before=before,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_file_edit_start_event(
|
||||||
|
tracker: FileEditTracker,
|
||||||
|
params: dict[str, Any] | None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
predicted_after = _predict_after_text(tracker.tool, params or {}, tracker.before)
|
||||||
|
if tracker.before.countable and predicted_after is not None:
|
||||||
|
added, deleted = line_diff_stats(tracker.before.text, predicted_after)
|
||||||
|
else:
|
||||||
|
added, deleted = 0, 0
|
||||||
|
return _event_payload(
|
||||||
|
tracker,
|
||||||
|
phase="start",
|
||||||
|
status="editing",
|
||||||
|
added=added,
|
||||||
|
deleted=deleted,
|
||||||
|
approximate=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_file_edit_end_event(
|
||||||
|
tracker: FileEditTracker,
|
||||||
|
params: dict[str, Any] | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
after = read_file_snapshot(tracker.path)
|
||||||
|
counted = False
|
||||||
|
if tracker.before.countable and after.countable:
|
||||||
|
added, deleted = line_diff_stats(tracker.before.text, after.text)
|
||||||
|
counted = True
|
||||||
|
else:
|
||||||
|
predicted_after = _predict_after_text(tracker.tool, params or {}, tracker.before)
|
||||||
|
if tracker.before.countable and predicted_after is not None:
|
||||||
|
added, deleted = line_diff_stats(tracker.before.text, predicted_after)
|
||||||
|
counted = True
|
||||||
|
else:
|
||||||
|
added, deleted = 0, 0
|
||||||
|
return _event_payload(
|
||||||
|
tracker,
|
||||||
|
phase="end",
|
||||||
|
status="done",
|
||||||
|
added=added,
|
||||||
|
deleted=deleted,
|
||||||
|
approximate=False,
|
||||||
|
binary=(after.binary or after.oversized or after.unreadable) and not counted,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_file_edit_error_event(
|
||||||
|
tracker: FileEditTracker,
|
||||||
|
error: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
payload = _event_payload(
|
||||||
|
tracker,
|
||||||
|
phase="error",
|
||||||
|
status="error",
|
||||||
|
added=0,
|
||||||
|
deleted=0,
|
||||||
|
approximate=False,
|
||||||
|
)
|
||||||
|
if error:
|
||||||
|
payload["error"] = error.strip()[:240]
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
def build_file_edit_live_event(
|
||||||
|
tracker: FileEditTracker,
|
||||||
|
*,
|
||||||
|
added: int,
|
||||||
|
deleted: int = 0,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Build an approximate in-progress event while tool-call arguments stream."""
|
||||||
|
return _event_payload(
|
||||||
|
tracker,
|
||||||
|
phase="start",
|
||||||
|
status="editing",
|
||||||
|
added=added,
|
||||||
|
deleted=deleted,
|
||||||
|
approximate=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_file_edit_pending_event(
|
||||||
|
*,
|
||||||
|
call_id: str,
|
||||||
|
tool_name: str,
|
||||||
|
added: int = 0,
|
||||||
|
deleted: int = 0,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Build an early placeholder before the streamed JSON path is available."""
|
||||||
|
return {
|
||||||
|
"version": 1,
|
||||||
|
"call_id": str(call_id or ""),
|
||||||
|
"tool": tool_name,
|
||||||
|
"path": "",
|
||||||
|
"phase": "start",
|
||||||
|
"added": max(0, int(added)),
|
||||||
|
"deleted": max(0, int(deleted)),
|
||||||
|
"approximate": True,
|
||||||
|
"status": "editing",
|
||||||
|
"pending": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class StreamingFileEditTracker:
|
||||||
|
"""Track file-edit tool arguments while the model is still streaming them.
|
||||||
|
|
||||||
|
Tool execution events only begin after the provider has completed the full
|
||||||
|
function call. For large ``write_file`` calls, the long wait is usually the
|
||||||
|
model producing the JSON ``content`` argument. Large ``edit_file`` calls
|
||||||
|
can have the same wait while ``old_text`` / ``new_text`` stream in. This
|
||||||
|
tracker converts those argument deltas into approximate WebUI file-edit
|
||||||
|
events before the final exact diff is available.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
workspace: Path | None,
|
||||||
|
tools: Any,
|
||||||
|
emit: Callable[[list[dict[str, Any]]], Awaitable[None]],
|
||||||
|
) -> None:
|
||||||
|
self._workspace = workspace
|
||||||
|
self._tools = tools
|
||||||
|
self._emit = emit
|
||||||
|
self._states: dict[str, _StreamingFileEditState] = {}
|
||||||
|
|
||||||
|
async def update(self, payload: dict[str, Any]) -> None:
|
||||||
|
key = _stream_key(payload)
|
||||||
|
if not key:
|
||||||
|
return
|
||||||
|
state = self._states.get(key)
|
||||||
|
if state is None:
|
||||||
|
state = _StreamingFileEditState(key=key)
|
||||||
|
self._states[key] = state
|
||||||
|
|
||||||
|
state.apply_delta(payload)
|
||||||
|
if state.name not in {"write_file", "edit_file"}:
|
||||||
|
return
|
||||||
|
if state.path is None:
|
||||||
|
state.path = _extract_complete_json_string(state.arguments, "path")
|
||||||
|
if state.path is None:
|
||||||
|
added, deleted = state.live_diff_counts()
|
||||||
|
now = time.monotonic()
|
||||||
|
if state.should_emit_pending(added, deleted, now):
|
||||||
|
state.mark_pending_emitted(added, deleted, now)
|
||||||
|
await self._emit([build_file_edit_pending_event(
|
||||||
|
call_id=state.call_id or state.key,
|
||||||
|
tool_name=state.name,
|
||||||
|
added=added,
|
||||||
|
deleted=deleted,
|
||||||
|
)])
|
||||||
|
return
|
||||||
|
if state.tracker is None:
|
||||||
|
tool = self._tools.get(state.name) if hasattr(self._tools, "get") else None
|
||||||
|
state.tracker = prepare_file_edit_tracker(
|
||||||
|
call_id=state.call_id or state.key,
|
||||||
|
tool_name=state.name,
|
||||||
|
tool=tool,
|
||||||
|
workspace=self._workspace,
|
||||||
|
params={"path": state.path},
|
||||||
|
)
|
||||||
|
if state.tracker is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
added, deleted = state.live_diff_counts()
|
||||||
|
now = time.monotonic()
|
||||||
|
if not state.should_emit(added, deleted, now):
|
||||||
|
return
|
||||||
|
state.mark_emitted(added, deleted, now)
|
||||||
|
await self._emit([build_file_edit_live_event(
|
||||||
|
state.tracker,
|
||||||
|
added=added,
|
||||||
|
deleted=deleted,
|
||||||
|
)])
|
||||||
|
|
||||||
|
async def flush(self) -> None:
|
||||||
|
events: list[dict[str, Any]] = []
|
||||||
|
now = time.monotonic()
|
||||||
|
for state in self._states.values():
|
||||||
|
if state.tracker is None:
|
||||||
|
continue
|
||||||
|
added, deleted = state.live_diff_counts()
|
||||||
|
if (
|
||||||
|
state.last_emitted_added == added
|
||||||
|
and state.last_emitted_deleted == deleted
|
||||||
|
and state.emitted_once
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
state.mark_emitted(added, deleted, now)
|
||||||
|
events.append(build_file_edit_live_event(
|
||||||
|
state.tracker,
|
||||||
|
added=added,
|
||||||
|
deleted=deleted,
|
||||||
|
))
|
||||||
|
if events:
|
||||||
|
await self._emit(events)
|
||||||
|
|
||||||
|
def apply_final_call_ids(self, final_tool_calls: list[Any]) -> None:
|
||||||
|
"""Keep final start/end events keyed to any earlier streamed placeholder."""
|
||||||
|
for tool_call in final_tool_calls:
|
||||||
|
canonical = self.canonical_call_id_for(tool_call)
|
||||||
|
if canonical:
|
||||||
|
try:
|
||||||
|
tool_call.id = canonical
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def canonical_call_id_for(self, tool_call: Any) -> str | None:
|
||||||
|
for state in self._states.values():
|
||||||
|
if state.matches_final_tool_call(tool_call):
|
||||||
|
return state.call_id or (state.tracker.call_id if state.tracker else None) or state.key
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def error_unmatched(
|
||||||
|
self,
|
||||||
|
final_tool_calls: list[Any],
|
||||||
|
error: str,
|
||||||
|
) -> None:
|
||||||
|
"""Mark streamed edits as failed when no final tool call will run."""
|
||||||
|
events: list[dict[str, Any]] = []
|
||||||
|
for state in self._states.values():
|
||||||
|
if state.tracker is None:
|
||||||
|
continue
|
||||||
|
if any(state.matches_final_tool_call(tool_call) for tool_call in final_tool_calls):
|
||||||
|
continue
|
||||||
|
events.append(build_file_edit_error_event(state.tracker, error))
|
||||||
|
if events:
|
||||||
|
await self._emit(events)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class _StreamingJsonStringField:
|
||||||
|
key: str
|
||||||
|
scan_pos: int | None = None
|
||||||
|
closed: bool = False
|
||||||
|
escape: bool = False
|
||||||
|
unicode_remaining: int = 0
|
||||||
|
unicode_buffer: str = ""
|
||||||
|
newline_count: int = 0
|
||||||
|
has_chars: bool = False
|
||||||
|
last_char_newline: bool = False
|
||||||
|
last_char_cr: bool = False
|
||||||
|
|
||||||
|
@property
|
||||||
|
def line_count(self) -> int:
|
||||||
|
if not self.has_chars:
|
||||||
|
return 0
|
||||||
|
return self.newline_count + (0 if self.last_char_newline else 1)
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
self.scan_pos = None
|
||||||
|
self.closed = False
|
||||||
|
self.escape = False
|
||||||
|
self.unicode_remaining = 0
|
||||||
|
self.unicode_buffer = ""
|
||||||
|
self.newline_count = 0
|
||||||
|
self.has_chars = False
|
||||||
|
self.last_char_newline = False
|
||||||
|
self.last_char_cr = False
|
||||||
|
|
||||||
|
def scan(self, source: str) -> None:
|
||||||
|
if self.closed:
|
||||||
|
return
|
||||||
|
if self.scan_pos is None:
|
||||||
|
match = re.search(rf'"{re.escape(self.key)}"\s*:\s*"', source)
|
||||||
|
if match is None:
|
||||||
|
return
|
||||||
|
self.scan_pos = match.end()
|
||||||
|
i = self.scan_pos
|
||||||
|
while i < len(source):
|
||||||
|
ch = source[i]
|
||||||
|
if self.unicode_remaining > 0:
|
||||||
|
self.unicode_buffer += ch
|
||||||
|
self.unicode_remaining -= 1
|
||||||
|
if self.unicode_remaining == 0:
|
||||||
|
try:
|
||||||
|
decoded = chr(int(self.unicode_buffer, 16))
|
||||||
|
except ValueError:
|
||||||
|
decoded = "x"
|
||||||
|
self.unicode_buffer = ""
|
||||||
|
self._mark_char(decoded)
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
if self.escape:
|
||||||
|
self.escape = False
|
||||||
|
if ch == "u":
|
||||||
|
self.unicode_remaining = 4
|
||||||
|
self.unicode_buffer = ""
|
||||||
|
elif ch == "n":
|
||||||
|
self._mark_char("\n")
|
||||||
|
elif ch == "r":
|
||||||
|
self._mark_char("\r")
|
||||||
|
else:
|
||||||
|
self._mark_char(ch)
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
if ch == "\\":
|
||||||
|
self.escape = True
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
if ch == '"':
|
||||||
|
self.closed = True
|
||||||
|
i += 1
|
||||||
|
break
|
||||||
|
self._mark_char(ch)
|
||||||
|
i += 1
|
||||||
|
self.scan_pos = i
|
||||||
|
|
||||||
|
def _mark_char(self, ch: str) -> None:
|
||||||
|
self.has_chars = True
|
||||||
|
if ch == "\r":
|
||||||
|
self.newline_count += 1
|
||||||
|
self.last_char_newline = True
|
||||||
|
self.last_char_cr = True
|
||||||
|
elif ch == "\n":
|
||||||
|
if not self.last_char_cr:
|
||||||
|
self.newline_count += 1
|
||||||
|
self.last_char_newline = True
|
||||||
|
self.last_char_cr = False
|
||||||
|
else:
|
||||||
|
self.last_char_newline = False
|
||||||
|
self.last_char_cr = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class _StreamingFileEditState:
|
||||||
|
key: str
|
||||||
|
call_id: str = ""
|
||||||
|
name: str = ""
|
||||||
|
arguments: str = ""
|
||||||
|
path: str | None = None
|
||||||
|
tracker: FileEditTracker | None = None
|
||||||
|
content: _StreamingJsonStringField = field(
|
||||||
|
default_factory=lambda: _StreamingJsonStringField("content")
|
||||||
|
)
|
||||||
|
old_text: _StreamingJsonStringField = field(
|
||||||
|
default_factory=lambda: _StreamingJsonStringField("old_text")
|
||||||
|
)
|
||||||
|
new_text: _StreamingJsonStringField = field(
|
||||||
|
default_factory=lambda: _StreamingJsonStringField("new_text")
|
||||||
|
)
|
||||||
|
emitted_once: bool = False
|
||||||
|
last_emitted_added: int = -1
|
||||||
|
last_emitted_deleted: int = -1
|
||||||
|
last_emit_at: float = 0.0
|
||||||
|
pending_emitted: bool = False
|
||||||
|
last_pending_added: int = -1
|
||||||
|
last_pending_deleted: int = -1
|
||||||
|
last_pending_at: float = 0.0
|
||||||
|
|
||||||
|
def apply_delta(self, payload: dict[str, Any]) -> None:
|
||||||
|
call_id = payload.get("call_id")
|
||||||
|
if isinstance(call_id, str) and call_id:
|
||||||
|
self.call_id = call_id
|
||||||
|
name = payload.get("name")
|
||||||
|
if isinstance(name, str) and name:
|
||||||
|
self.name = name
|
||||||
|
args = payload.get("arguments")
|
||||||
|
if isinstance(args, str):
|
||||||
|
self.arguments = args
|
||||||
|
self.content.reset()
|
||||||
|
self.old_text.reset()
|
||||||
|
self.new_text.reset()
|
||||||
|
return
|
||||||
|
delta = payload.get("arguments_delta")
|
||||||
|
if isinstance(delta, str) and delta:
|
||||||
|
self.arguments += delta
|
||||||
|
|
||||||
|
def live_diff_counts(self) -> tuple[int, int]:
|
||||||
|
if self.name == "write_file":
|
||||||
|
self.content.scan(self.arguments)
|
||||||
|
return self.content.line_count, 0
|
||||||
|
if self.name == "edit_file":
|
||||||
|
self.old_text.scan(self.arguments)
|
||||||
|
self.new_text.scan(self.arguments)
|
||||||
|
return self.new_text.line_count, self.old_text.line_count
|
||||||
|
return 0, 0
|
||||||
|
|
||||||
|
def should_emit(self, added: int, deleted: int, now: float) -> bool:
|
||||||
|
if not self.emitted_once:
|
||||||
|
return True
|
||||||
|
if added == self.last_emitted_added and deleted == self.last_emitted_deleted:
|
||||||
|
return False
|
||||||
|
if max(
|
||||||
|
abs(added - self.last_emitted_added),
|
||||||
|
abs(deleted - self.last_emitted_deleted),
|
||||||
|
) >= _LIVE_EMIT_LINE_STEP:
|
||||||
|
return True
|
||||||
|
return now - self.last_emit_at >= _LIVE_EMIT_INTERVAL_S
|
||||||
|
|
||||||
|
def mark_emitted(self, added: int, deleted: int, now: float) -> None:
|
||||||
|
self.emitted_once = True
|
||||||
|
self.last_emitted_added = added
|
||||||
|
self.last_emitted_deleted = deleted
|
||||||
|
self.last_emit_at = now
|
||||||
|
|
||||||
|
def should_emit_pending(self, added: int, deleted: int, now: float) -> bool:
|
||||||
|
if not self.pending_emitted:
|
||||||
|
return True
|
||||||
|
if added == self.last_pending_added and deleted == self.last_pending_deleted:
|
||||||
|
return False
|
||||||
|
if max(
|
||||||
|
abs(added - self.last_pending_added),
|
||||||
|
abs(deleted - self.last_pending_deleted),
|
||||||
|
) >= _LIVE_EMIT_LINE_STEP:
|
||||||
|
return True
|
||||||
|
return now - self.last_pending_at >= _LIVE_EMIT_INTERVAL_S
|
||||||
|
|
||||||
|
def mark_pending_emitted(self, added: int, deleted: int, now: float) -> None:
|
||||||
|
self.pending_emitted = True
|
||||||
|
self.last_pending_added = added
|
||||||
|
self.last_pending_deleted = deleted
|
||||||
|
self.last_pending_at = now
|
||||||
|
|
||||||
|
def matches_final_tool_call(self, tool_call: Any) -> bool:
|
||||||
|
call_id = getattr(tool_call, "id", None)
|
||||||
|
canonical = self.call_id or (self.tracker.call_id if self.tracker else "")
|
||||||
|
if isinstance(call_id, str) and call_id and canonical and call_id == canonical:
|
||||||
|
return True
|
||||||
|
name = getattr(tool_call, "name", None)
|
||||||
|
if name != self.name:
|
||||||
|
return False
|
||||||
|
arguments = getattr(tool_call, "arguments", None)
|
||||||
|
if not isinstance(arguments, dict):
|
||||||
|
return False
|
||||||
|
path = arguments.get("path")
|
||||||
|
if self.path is None and isinstance(path, str) and path:
|
||||||
|
self.path = path
|
||||||
|
return True
|
||||||
|
return isinstance(path, str) and path == self.path
|
||||||
|
|
||||||
|
|
||||||
|
def _stream_key(payload: dict[str, Any]) -> str:
|
||||||
|
index = payload.get("index")
|
||||||
|
if isinstance(index, int):
|
||||||
|
return f"idx:{index}"
|
||||||
|
if isinstance(index, str) and index:
|
||||||
|
return f"idx:{index}"
|
||||||
|
call_id = payload.get("call_id")
|
||||||
|
if isinstance(call_id, str) and call_id:
|
||||||
|
return f"id:{call_id}"
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_complete_json_string(source: str, key: str) -> str | None:
|
||||||
|
match = re.search(rf'"{re.escape(key)}"\s*:\s*"', source)
|
||||||
|
if match is None:
|
||||||
|
return None
|
||||||
|
out: list[str] = []
|
||||||
|
i = match.end()
|
||||||
|
escape = False
|
||||||
|
while i < len(source):
|
||||||
|
ch = source[i]
|
||||||
|
if escape:
|
||||||
|
escape = False
|
||||||
|
if ch == "n":
|
||||||
|
out.append("\n")
|
||||||
|
elif ch == "r":
|
||||||
|
out.append("\r")
|
||||||
|
elif ch == "t":
|
||||||
|
out.append("\t")
|
||||||
|
elif ch == "u":
|
||||||
|
digits = source[i + 1:i + 5]
|
||||||
|
if len(digits) < 4:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
out.append(chr(int(digits, 16)))
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
i += 4
|
||||||
|
else:
|
||||||
|
out.append(ch)
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
if ch == "\\":
|
||||||
|
escape = True
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
if ch == '"':
|
||||||
|
return "".join(out)
|
||||||
|
out.append(ch)
|
||||||
|
i += 1
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _event_payload(
|
||||||
|
tracker: FileEditTracker,
|
||||||
|
*,
|
||||||
|
phase: str,
|
||||||
|
status: str,
|
||||||
|
added: int,
|
||||||
|
deleted: int,
|
||||||
|
approximate: bool,
|
||||||
|
binary: bool = False,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"version": 1,
|
||||||
|
"call_id": tracker.call_id,
|
||||||
|
"tool": tracker.tool,
|
||||||
|
"path": tracker.display_path,
|
||||||
|
"absolute_path": tracker.path.as_posix(),
|
||||||
|
"phase": phase,
|
||||||
|
"added": max(0, int(added)),
|
||||||
|
"deleted": max(0, int(deleted)),
|
||||||
|
"approximate": bool(approximate),
|
||||||
|
"status": status,
|
||||||
|
}
|
||||||
|
if binary:
|
||||||
|
payload["binary"] = True
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
def _predict_after_text(
|
||||||
|
tool_name: str,
|
||||||
|
params: dict[str, Any],
|
||||||
|
before: FileSnapshot,
|
||||||
|
) -> str | None:
|
||||||
|
if not before.countable:
|
||||||
|
return None
|
||||||
|
before_text = before.text or ""
|
||||||
|
if tool_name == "write_file":
|
||||||
|
content = params.get("content")
|
||||||
|
return content if isinstance(content, str) else ""
|
||||||
|
if tool_name == "edit_file":
|
||||||
|
old_text = params.get("old_text")
|
||||||
|
new_text = params.get("new_text")
|
||||||
|
if not isinstance(old_text, str) or not isinstance(new_text, str):
|
||||||
|
return None
|
||||||
|
replace_all = bool(params.get("replace_all"))
|
||||||
|
if old_text == "":
|
||||||
|
return new_text if not before.exists else before_text
|
||||||
|
if old_text in before_text:
|
||||||
|
if replace_all:
|
||||||
|
return before_text.replace(old_text, new_text)
|
||||||
|
return before_text.replace(old_text, new_text, 1)
|
||||||
|
return None
|
||||||
|
if tool_name == "notebook_edit":
|
||||||
|
return _predict_notebook_after_text(params, before_text)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _predict_notebook_after_text(params: dict[str, Any], before_text: str) -> str | None:
|
||||||
|
try:
|
||||||
|
nb = json.loads(before_text) if before_text.strip() else _empty_notebook()
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
cells = nb.get("cells")
|
||||||
|
if not isinstance(cells, list):
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
cell_index = int(params.get("cell_index", 0))
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return None
|
||||||
|
new_source = params.get("new_source")
|
||||||
|
source = new_source if isinstance(new_source, str) else ""
|
||||||
|
cell_type = (
|
||||||
|
params.get("cell_type") if params.get("cell_type") in ("code", "markdown") else "code"
|
||||||
|
)
|
||||||
|
mode = (
|
||||||
|
params.get("edit_mode")
|
||||||
|
if params.get("edit_mode") in ("replace", "insert", "delete")
|
||||||
|
else "replace"
|
||||||
|
)
|
||||||
|
if mode == "delete":
|
||||||
|
if 0 <= cell_index < len(cells):
|
||||||
|
cells.pop(cell_index)
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
elif mode == "insert":
|
||||||
|
insert_at = min(max(cell_index + 1, 0), len(cells))
|
||||||
|
cells.insert(insert_at, _new_notebook_cell(source, str(cell_type)))
|
||||||
|
else:
|
||||||
|
if not (0 <= cell_index < len(cells)):
|
||||||
|
return None
|
||||||
|
cell = cells[cell_index]
|
||||||
|
if not isinstance(cell, dict):
|
||||||
|
return None
|
||||||
|
cell["source"] = source
|
||||||
|
cell["cell_type"] = cell_type
|
||||||
|
if cell_type == "code":
|
||||||
|
cell.setdefault("outputs", [])
|
||||||
|
cell.setdefault("execution_count", None)
|
||||||
|
else:
|
||||||
|
cell.pop("outputs", None)
|
||||||
|
cell.pop("execution_count", None)
|
||||||
|
nb["cells"] = cells
|
||||||
|
try:
|
||||||
|
return json.dumps(nb, indent=1, ensure_ascii=False)
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _empty_notebook() -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"nbformat": 4,
|
||||||
|
"nbformat_minor": 5,
|
||||||
|
"metadata": {
|
||||||
|
"kernelspec": {"display_name": "Python 3", "language": "python", "name": "python3"},
|
||||||
|
"language_info": {"name": "python"},
|
||||||
|
},
|
||||||
|
"cells": [],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _new_notebook_cell(source: str, cell_type: str) -> dict[str, Any]:
|
||||||
|
cell: dict[str, Any] = {"cell_type": cell_type, "source": source, "metadata": {}}
|
||||||
|
if cell_type == "code":
|
||||||
|
cell["outputs"] = []
|
||||||
|
cell["execution_count"] = None
|
||||||
|
return cell
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
"""Small helpers for passing the active LLM provider/model together."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
from nanobot.providers.base import LLMProvider
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class LLMRuntime:
|
||||||
|
provider: LLMProvider
|
||||||
|
model: str
|
||||||
|
|
||||||
|
|
||||||
|
LLMRuntimeResolver = Callable[[], LLMRuntime]
|
||||||
|
|
||||||
|
|
||||||
|
def static_llm_runtime(provider: LLMProvider, model: str) -> LLMRuntimeResolver:
|
||||||
|
runtime = LLMRuntime(provider=provider, model=model)
|
||||||
|
return lambda: runtime
|
||||||
@@ -10,13 +10,21 @@ from nanobot.agent.hook import AgentHookContext
|
|||||||
|
|
||||||
|
|
||||||
def on_progress_accepts_tool_events(cb: Callable[..., Any]) -> bool:
|
def on_progress_accepts_tool_events(cb: Callable[..., Any]) -> bool:
|
||||||
|
return _on_progress_accepts(cb, "tool_events")
|
||||||
|
|
||||||
|
|
||||||
|
def on_progress_accepts_file_edit_events(cb: Callable[..., Any]) -> bool:
|
||||||
|
return _on_progress_accepts(cb, "file_edit_events")
|
||||||
|
|
||||||
|
|
||||||
|
def _on_progress_accepts(cb: Callable[..., Any], name: str) -> bool:
|
||||||
try:
|
try:
|
||||||
sig = inspect.signature(cb)
|
sig = inspect.signature(cb)
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return False
|
return False
|
||||||
if any(p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()):
|
if any(p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()):
|
||||||
return True
|
return True
|
||||||
return "tool_events" in sig.parameters
|
return name in sig.parameters
|
||||||
|
|
||||||
|
|
||||||
async def invoke_on_progress(
|
async def invoke_on_progress(
|
||||||
@@ -32,6 +40,15 @@ async def invoke_on_progress(
|
|||||||
await on_progress(content, tool_hint=tool_hint)
|
await on_progress(content, tool_hint=tool_hint)
|
||||||
|
|
||||||
|
|
||||||
|
async def invoke_file_edit_progress(
|
||||||
|
on_progress: Callable[..., Awaitable[None]],
|
||||||
|
file_edit_events: list[dict[str, Any]],
|
||||||
|
) -> None:
|
||||||
|
if not file_edit_events or not on_progress_accepts_file_edit_events(on_progress):
|
||||||
|
return
|
||||||
|
await on_progress("", file_edit_events=file_edit_events)
|
||||||
|
|
||||||
|
|
||||||
def build_tool_event_start_payload(tool_call: Any) -> dict[str, Any]:
|
def build_tool_event_start_payload(tool_call: Any) -> dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"version": 1,
|
"version": 1,
|
||||||
|
|||||||
@@ -1,74 +0,0 @@
|
|||||||
"""Session replay: ensure assistant ``media`` paths are under the media root.
|
|
||||||
|
|
||||||
WebUI history signing (``/api/.../messages``) only works for files inside
|
|
||||||
``get_media_dir``. Tool-driven attachments may live in the workspace; stage
|
|
||||||
copies into the websocket media bucket before persisting message JSON.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import shutil
|
|
||||||
import uuid
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
from nanobot.config.paths import get_media_dir
|
|
||||||
from nanobot.utils.helpers import safe_filename
|
|
||||||
|
|
||||||
|
|
||||||
def stage_media_paths_for_session_replay(paths: list[str]) -> list[str]:
|
|
||||||
"""Keep local files only; copy anything outside the media root into ``media/websocket``."""
|
|
||||||
root = get_media_dir().resolve()
|
|
||||||
out: list[str] = []
|
|
||||||
seen: set[str] = set()
|
|
||||||
for raw in paths:
|
|
||||||
if not isinstance(raw, str) or not raw.strip():
|
|
||||||
continue
|
|
||||||
if raw.startswith(("http://", "https://")):
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
p = Path(raw).expanduser().resolve()
|
|
||||||
except OSError:
|
|
||||||
continue
|
|
||||||
if not p.is_file():
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
p.relative_to(root)
|
|
||||||
key = str(p)
|
|
||||||
except ValueError:
|
|
||||||
try:
|
|
||||||
media_dir = get_media_dir("websocket")
|
|
||||||
staged = media_dir / f"{uuid.uuid4().hex[:12]}-{safe_filename(p.name) or 'attachment'}"
|
|
||||||
shutil.copyfile(p, staged)
|
|
||||||
key = str(staged.resolve())
|
|
||||||
except OSError as exc:
|
|
||||||
logger.warning("failed to stage session media from {}: {}", raw, exc)
|
|
||||||
continue
|
|
||||||
if key not in seen:
|
|
||||||
out.append(key)
|
|
||||||
seen.add(key)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def merge_turn_media_into_last_assistant(
|
|
||||||
all_messages: list[dict[str, Any]],
|
|
||||||
generated_image_paths: list[str],
|
|
||||||
extra_attachment_paths: list[str],
|
|
||||||
) -> None:
|
|
||||||
"""Attach staged paths to the last assistant row in *all_messages* (in-place)."""
|
|
||||||
merged = list(
|
|
||||||
dict.fromkeys(
|
|
||||||
[
|
|
||||||
*stage_media_paths_for_session_replay(generated_image_paths),
|
|
||||||
*stage_media_paths_for_session_replay(extra_attachment_paths),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
last = all_messages[-1] if all_messages else None
|
|
||||||
if not merged or not last or last.get("role") != "assistant":
|
|
||||||
return
|
|
||||||
existing = last.get("media")
|
|
||||||
base = existing if isinstance(existing, list) else []
|
|
||||||
last["media"] = list(dict.fromkeys([*base, *merged]))
|
|
||||||
@@ -1,138 +0,0 @@
|
|||||||
"""Helpers for WebUI chat title generation."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import re
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
from nanobot.providers.base import LLMProvider
|
|
||||||
from nanobot.session.manager import Session, SessionManager
|
|
||||||
from nanobot.utils.helpers import truncate_text
|
|
||||||
|
|
||||||
WEBUI_SESSION_METADATA_KEY = "webui"
|
|
||||||
WEBUI_TITLE_METADATA_KEY = "title"
|
|
||||||
WEBUI_TITLE_USER_EDITED_METADATA_KEY = "title_user_edited"
|
|
||||||
TITLE_MAX_CHARS = 60
|
|
||||||
|
|
||||||
|
|
||||||
def mark_webui_session(session: Session, metadata: dict[str, Any]) -> bool:
|
|
||||||
"""Persist a WebUI marker only when the inbound websocket frame opted in."""
|
|
||||||
if metadata.get(WEBUI_SESSION_METADATA_KEY) is not True:
|
|
||||||
return False
|
|
||||||
session.metadata[WEBUI_SESSION_METADATA_KEY] = True
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
def clean_generated_title(raw: str | None) -> str:
|
|
||||||
text = (raw or "").strip()
|
|
||||||
if not text:
|
|
||||||
return ""
|
|
||||||
text = re.sub(r"^\s*(title|标题)\s*[::]\s*", "", text, flags=re.IGNORECASE)
|
|
||||||
text = text.strip().strip("\"'`“”‘’")
|
|
||||||
text = re.sub(r"\s+", " ", text).strip()
|
|
||||||
text = text.rstrip("。.!!??,,;;:")
|
|
||||||
if len(text) > TITLE_MAX_CHARS:
|
|
||||||
text = text[: TITLE_MAX_CHARS - 1].rstrip() + "…"
|
|
||||||
return text
|
|
||||||
|
|
||||||
|
|
||||||
def _title_inputs(session: Session) -> tuple[str, str]:
|
|
||||||
user_text = ""
|
|
||||||
assistant_text = ""
|
|
||||||
for message in session.messages:
|
|
||||||
role = message.get("role")
|
|
||||||
content = message.get("content")
|
|
||||||
if not isinstance(content, str) or not content.strip():
|
|
||||||
continue
|
|
||||||
if role == "user" and not user_text:
|
|
||||||
user_text = content.strip()
|
|
||||||
elif role == "assistant" and not assistant_text:
|
|
||||||
assistant_text = content.strip()
|
|
||||||
if user_text and assistant_text:
|
|
||||||
break
|
|
||||||
return user_text, assistant_text
|
|
||||||
|
|
||||||
|
|
||||||
async def maybe_generate_webui_title(
|
|
||||||
*,
|
|
||||||
sessions: SessionManager,
|
|
||||||
session_key: str,
|
|
||||||
provider: LLMProvider,
|
|
||||||
model: str,
|
|
||||||
) -> bool:
|
|
||||||
"""Generate and persist a short title for WebUI-owned sessions only."""
|
|
||||||
session = sessions.get_or_create(session_key)
|
|
||||||
if session.metadata.get(WEBUI_SESSION_METADATA_KEY) is not True:
|
|
||||||
return False
|
|
||||||
if session.metadata.get(WEBUI_TITLE_USER_EDITED_METADATA_KEY) is True:
|
|
||||||
return False
|
|
||||||
current_title = session.metadata.get(WEBUI_TITLE_METADATA_KEY)
|
|
||||||
if isinstance(current_title, str) and current_title.strip():
|
|
||||||
return False
|
|
||||||
|
|
||||||
user_text, assistant_text = _title_inputs(session)
|
|
||||||
if not user_text:
|
|
||||||
return False
|
|
||||||
|
|
||||||
prompt = (
|
|
||||||
"Generate a concise title for this chat.\n"
|
|
||||||
"Rules:\n"
|
|
||||||
"- Use the same language as the user when practical.\n"
|
|
||||||
"- 3 to 8 words.\n"
|
|
||||||
"- No quotes.\n"
|
|
||||||
"- No punctuation at the end.\n"
|
|
||||||
"- Return only the title.\n\n"
|
|
||||||
f"User: {truncate_text(user_text, 1_000)}"
|
|
||||||
)
|
|
||||||
if assistant_text:
|
|
||||||
prompt += f"\nAssistant: {truncate_text(assistant_text, 1_000)}"
|
|
||||||
|
|
||||||
try:
|
|
||||||
response = await provider.chat_with_retry(
|
|
||||||
[
|
|
||||||
{
|
|
||||||
"role": "system",
|
|
||||||
"content": (
|
|
||||||
"You write short, neutral chat titles. "
|
|
||||||
"Return only the title text."
|
|
||||||
),
|
|
||||||
},
|
|
||||||
{"role": "user", "content": prompt},
|
|
||||||
],
|
|
||||||
tools=None,
|
|
||||||
model=model,
|
|
||||||
max_tokens=32,
|
|
||||||
temperature=0.2,
|
|
||||||
retry_mode="standard",
|
|
||||||
)
|
|
||||||
except Exception:
|
|
||||||
logger.debug("Failed to generate webui session title for {}", session_key, exc_info=True)
|
|
||||||
return False
|
|
||||||
|
|
||||||
title = clean_generated_title(response.content)
|
|
||||||
if not title or title.lower().startswith("error"):
|
|
||||||
return False
|
|
||||||
session.metadata[WEBUI_TITLE_METADATA_KEY] = title
|
|
||||||
sessions.save(session)
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
async def maybe_generate_webui_title_after_turn(
|
|
||||||
*,
|
|
||||||
channel: str,
|
|
||||||
metadata: dict[str, Any],
|
|
||||||
sessions: SessionManager,
|
|
||||||
session_key: str,
|
|
||||||
provider: LLMProvider,
|
|
||||||
model: str,
|
|
||||||
) -> bool:
|
|
||||||
if channel != "websocket" or metadata.get(WEBUI_SESSION_METADATA_KEY) is not True:
|
|
||||||
return False
|
|
||||||
return await maybe_generate_webui_title(
|
|
||||||
sessions=sessions,
|
|
||||||
session_key=session_key,
|
|
||||||
provider=provider,
|
|
||||||
model=model,
|
|
||||||
)
|
|
||||||
@@ -1,48 +0,0 @@
|
|||||||
"""Outbound helpers for the WebSocket/WebUI wire contract.
|
|
||||||
|
|
||||||
AgentLoop uses these without importing a concrete channel plugin; only
|
|
||||||
``channel == "websocket"`` messages are affected.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import time
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
|
||||||
from nanobot.bus.queue import MessageBus
|
|
||||||
|
|
||||||
# Wall-clock turn start per ``chat_id`` (websocket only). Survives browser refresh while the
|
|
||||||
# gateway process stays up; cleared on idle/stop and implicitly dropped on restart.
|
|
||||||
_WEBSOCKET_TURN_WALL_STARTED_AT: dict[str, float] = {}
|
|
||||||
|
|
||||||
|
|
||||||
def websocket_turn_wall_started_at(chat_id: str) -> float | None:
|
|
||||||
"""Return ``time.time()`` when the active user turn began, if still running."""
|
|
||||||
return _WEBSOCKET_TURN_WALL_STARTED_AT.get(chat_id)
|
|
||||||
|
|
||||||
|
|
||||||
async def publish_turn_run_status(bus: MessageBus, msg: InboundMessage, status: str) -> None:
|
|
||||||
"""Notify WebSocket clients while a user turn is executing (timing strip)."""
|
|
||||||
if msg.channel != "websocket":
|
|
||||||
return
|
|
||||||
cid = str(msg.chat_id)
|
|
||||||
meta: dict[str, Any] = {
|
|
||||||
**dict(msg.metadata or {}),
|
|
||||||
"_goal_status": True,
|
|
||||||
"goal_status": status,
|
|
||||||
}
|
|
||||||
if status == "running":
|
|
||||||
t0 = time.time()
|
|
||||||
meta["started_at"] = t0
|
|
||||||
_WEBSOCKET_TURN_WALL_STARTED_AT[cid] = t0
|
|
||||||
else:
|
|
||||||
_WEBSOCKET_TURN_WALL_STARTED_AT.pop(cid, None)
|
|
||||||
await bus.publish_outbound(
|
|
||||||
OutboundMessage(
|
|
||||||
channel=msg.channel,
|
|
||||||
chat_id=cid,
|
|
||||||
content="",
|
|
||||||
metadata=meta,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
@@ -1,6 +1,8 @@
|
|||||||
"""Embedded web UI assets.
|
"""Embedded web UI assets.
|
||||||
|
|
||||||
The ``dist/`` subdirectory is populated by ``cd webui && bun run build`` and
|
The ``dist/`` subdirectory holds the production WebUI bundle served by the
|
||||||
is shipped in the wheel; it stays empty in source checkouts until that command
|
gateway. It is shipped inside the published wheel and is rebuilt automatically
|
||||||
has been run.
|
by the ``webui-build`` Hatch hook during ``python -m build``. In an editable
|
||||||
|
source checkout it stays empty until you run ``cd webui && bun run build``
|
||||||
|
(or use the Vite dev server at ``cd webui && bun run dev``).
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -0,0 +1,2 @@
|
|||||||
|
"""Backend helpers for the bundled WebUI surface."""
|
||||||
|
|
||||||
@@ -0,0 +1,609 @@
|
|||||||
|
"""Settings REST helpers for the WebUI HTTP surface.
|
||||||
|
|
||||||
|
The WebSocket channel owns transport/authentication. This module owns the
|
||||||
|
settings payload shape and the allowlisted config mutations exposed to WebUI.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from zoneinfo import ZoneInfo
|
||||||
|
|
||||||
|
from nanobot.config.loader import get_config_path, load_config, save_config
|
||||||
|
from nanobot.providers.image_generation import (
|
||||||
|
get_image_gen_provider,
|
||||||
|
image_gen_provider_names,
|
||||||
|
)
|
||||||
|
from nanobot.providers.registry import PROVIDERS, find_by_name
|
||||||
|
|
||||||
|
QueryParams = dict[str, list[str]]
|
||||||
|
|
||||||
|
_WEB_SEARCH_PROVIDER_OPTIONS: tuple[dict[str, str], ...] = (
|
||||||
|
{"name": "duckduckgo", "label": "DuckDuckGo", "credential": "none"},
|
||||||
|
{"name": "brave", "label": "Brave Search", "credential": "api_key"},
|
||||||
|
{"name": "tavily", "label": "Tavily", "credential": "api_key"},
|
||||||
|
{"name": "searxng", "label": "SearXNG", "credential": "base_url"},
|
||||||
|
{"name": "jina", "label": "Jina", "credential": "api_key"},
|
||||||
|
{"name": "kagi", "label": "Kagi", "credential": "api_key"},
|
||||||
|
{"name": "olostep", "label": "Olostep", "credential": "api_key"},
|
||||||
|
)
|
||||||
|
_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",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
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 _mask_secret_hint(secret: str | None) -> str | None:
|
||||||
|
if not secret:
|
||||||
|
return None
|
||||||
|
if len(secret) <= 8:
|
||||||
|
return "••••"
|
||||||
|
return f"{secret[:4]}••••{secret[-4:]}"
|
||||||
|
|
||||||
|
|
||||||
|
def _provider_requires_api_key(spec: Any) -> bool:
|
||||||
|
if spec.backend == "azure_openai":
|
||||||
|
return True
|
||||||
|
if spec.is_local or spec.is_direct:
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def _provider_configured_for_settings(spec: Any, provider_config: Any) -> bool:
|
||||||
|
if _provider_requires_api_key(spec):
|
||||||
|
return bool(provider_config.api_key)
|
||||||
|
return bool(
|
||||||
|
provider_config.api_key
|
||||||
|
or provider_config.api_base
|
||||||
|
or getattr(provider_config, "region", None)
|
||||||
|
or getattr(provider_config, "profile", None)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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"}
|
||||||
|
|
||||||
|
|
||||||
|
def _image_generation_provider_rows(config: Any) -> list[dict[str, Any]]:
|
||||||
|
rows: list[dict[str, Any]] = []
|
||||||
|
for name in image_gen_provider_names():
|
||||||
|
spec = find_by_name(name)
|
||||||
|
provider_config = getattr(config.providers, name, None)
|
||||||
|
configured = (
|
||||||
|
_provider_configured_for_settings(spec, provider_config)
|
||||||
|
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,
|
||||||
|
"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 settings_payload(*, requires_restart: bool = False) -> dict[str, Any]:
|
||||||
|
config = load_config()
|
||||||
|
defaults = config.agents.defaults
|
||||||
|
active_preset_name = defaults.model_preset or "default"
|
||||||
|
try:
|
||||||
|
effective_preset = config.resolve_preset()
|
||||||
|
except Exception:
|
||||||
|
effective_preset = config.resolve_default_preset()
|
||||||
|
active_preset_name = "default"
|
||||||
|
|
||||||
|
provider_name = (
|
||||||
|
config.get_provider_name(effective_preset.model, preset=effective_preset)
|
||||||
|
or effective_preset.provider
|
||||||
|
)
|
||||||
|
provider = config.get_provider(effective_preset.model, preset=effective_preset)
|
||||||
|
selected_provider = provider_name
|
||||||
|
if effective_preset.provider != "auto":
|
||||||
|
spec = find_by_name(effective_preset.provider)
|
||||||
|
selected_provider = spec.name if spec else provider_name
|
||||||
|
|
||||||
|
providers = []
|
||||||
|
for spec in PROVIDERS:
|
||||||
|
provider_config = getattr(config.providers, spec.name, None)
|
||||||
|
if provider_config is None or spec.is_oauth:
|
||||||
|
continue
|
||||||
|
providers.append(
|
||||||
|
{
|
||||||
|
"name": spec.name,
|
||||||
|
"label": spec.label,
|
||||||
|
"configured": _provider_configured_for_settings(spec, provider_config),
|
||||||
|
"api_key_required": _provider_requires_api_key(spec),
|
||||||
|
"api_key_hint": _mask_secret_hint(provider_config.api_key),
|
||||||
|
"api_base": provider_config.api_base,
|
||||||
|
"default_api_base": spec.default_api_base or None,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
search_config = config.tools.web.search
|
||||||
|
image_config = config.tools.image_generation
|
||||||
|
search_provider = (
|
||||||
|
search_config.provider
|
||||||
|
if search_config.provider in _WEB_SEARCH_PROVIDER_BY_NAME
|
||||||
|
else "duckduckgo"
|
||||||
|
)
|
||||||
|
image_providers = _image_generation_provider_rows(config)
|
||||||
|
selected_image_provider = next(
|
||||||
|
(
|
||||||
|
provider
|
||||||
|
for provider in image_providers
|
||||||
|
if provider["name"] == image_config.provider
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
model_presets = [
|
||||||
|
{
|
||||||
|
"name": "default",
|
||||||
|
"label": "Default",
|
||||||
|
"active": active_preset_name == "default",
|
||||||
|
"is_default": True,
|
||||||
|
"model": defaults.model,
|
||||||
|
"provider": defaults.provider,
|
||||||
|
"max_tokens": defaults.max_tokens,
|
||||||
|
"context_window_tokens": defaults.context_window_tokens,
|
||||||
|
"temperature": defaults.temperature,
|
||||||
|
"reasoning_effort": defaults.reasoning_effort,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
for name, preset in config.model_presets.items():
|
||||||
|
model_presets.append(
|
||||||
|
{
|
||||||
|
"name": name,
|
||||||
|
"label": name,
|
||||||
|
"active": active_preset_name == name,
|
||||||
|
"is_default": False,
|
||||||
|
"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,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
exec_config = config.tools.exec
|
||||||
|
return {
|
||||||
|
"agent": {
|
||||||
|
"model": effective_preset.model,
|
||||||
|
"provider": selected_provider,
|
||||||
|
"resolved_provider": provider_name,
|
||||||
|
"has_api_key": bool(provider and provider.api_key),
|
||||||
|
"model_preset": active_preset_name,
|
||||||
|
"max_tokens": effective_preset.max_tokens,
|
||||||
|
"context_window_tokens": effective_preset.context_window_tokens,
|
||||||
|
"temperature": effective_preset.temperature,
|
||||||
|
"reasoning_effort": effective_preset.reasoning_effort,
|
||||||
|
"timezone": defaults.timezone,
|
||||||
|
"bot_name": defaults.bot_name,
|
||||||
|
"bot_icon": defaults.bot_icon,
|
||||||
|
"tool_hint_max_length": defaults.tool_hint_max_length,
|
||||||
|
},
|
||||||
|
"model_presets": model_presets,
|
||||||
|
"providers": providers,
|
||||||
|
"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,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"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,
|
||||||
|
},
|
||||||
|
"runtime": {
|
||||||
|
"config_path": str(get_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(),
|
||||||
|
"max_batch_size": defaults.dream.max_batch_size,
|
||||||
|
"max_iterations": defaults.dream.max_iterations,
|
||||||
|
"annotate_line_ages": defaults.dream.annotate_line_ages,
|
||||||
|
},
|
||||||
|
"unified_session": defaults.unified_session,
|
||||||
|
},
|
||||||
|
"advanced": {
|
||||||
|
"restrict_to_workspace": config.tools.restrict_to_workspace,
|
||||||
|
"ssrf_whitelist_count": len(config.tools.ssrf_whitelist),
|
||||||
|
"mcp_server_count": len(config.tools.mcp_servers),
|
||||||
|
"exec_enabled": exec_config.enable,
|
||||||
|
"exec_sandbox": exec_config.sandbox or None,
|
||||||
|
"exec_path_append_set": bool(exec_config.path_append),
|
||||||
|
},
|
||||||
|
"requires_restart": requires_restart,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def update_agent_settings(query: QueryParams) -> dict[str, Any]:
|
||||||
|
config = load_config()
|
||||||
|
defaults = config.agents.defaults
|
||||||
|
changed = False
|
||||||
|
restart_required = False
|
||||||
|
|
||||||
|
if "model_preset" in query or "modelPreset" in query:
|
||||||
|
preset = (_query_first_alias(query, "model_preset", "modelPreset") or "").strip()
|
||||||
|
preset_value = None if not preset or preset == "default" else preset
|
||||||
|
if preset_value is not None and preset_value not in config.model_presets:
|
||||||
|
raise WebUISettingsError("unknown model preset")
|
||||||
|
if defaults.model_preset != preset_value:
|
||||||
|
defaults.model_preset = preset_value
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
model = _query_first(query, "model")
|
||||||
|
if model is not None:
|
||||||
|
model = model.strip()
|
||||||
|
if not model:
|
||||||
|
raise WebUISettingsError("model is required")
|
||||||
|
if defaults.model != model:
|
||||||
|
defaults.model = model
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
provider = _query_first(query, "provider")
|
||||||
|
if provider is not None:
|
||||||
|
provider = provider.strip()
|
||||||
|
if not provider:
|
||||||
|
raise WebUISettingsError("provider is required")
|
||||||
|
spec = find_by_name(provider)
|
||||||
|
if spec is None:
|
||||||
|
raise WebUISettingsError("unknown provider")
|
||||||
|
provider_config = getattr(config.providers, provider, None)
|
||||||
|
if (
|
||||||
|
provider_config is None
|
||||||
|
or not _provider_configured_for_settings(spec, provider_config)
|
||||||
|
):
|
||||||
|
raise WebUISettingsError("provider is not configured")
|
||||||
|
if defaults.provider != provider:
|
||||||
|
defaults.provider = provider
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
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
|
||||||
|
if defaults.timezone != timezone:
|
||||||
|
defaults.timezone = timezone
|
||||||
|
changed = True
|
||||||
|
restart_required = True
|
||||||
|
|
||||||
|
bot_name = _query_first_alias(query, "bot_name", "botName")
|
||||||
|
if bot_name is not None:
|
||||||
|
bot_name = bot_name.strip()
|
||||||
|
if not bot_name:
|
||||||
|
raise WebUISettingsError("bot_name is required")
|
||||||
|
if defaults.bot_name != bot_name:
|
||||||
|
defaults.bot_name = bot_name
|
||||||
|
changed = True
|
||||||
|
restart_required = True
|
||||||
|
|
||||||
|
bot_icon = _query_first_alias(query, "bot_icon", "botIcon")
|
||||||
|
if bot_icon is not None:
|
||||||
|
bot_icon = bot_icon.strip()
|
||||||
|
if defaults.bot_icon != bot_icon:
|
||||||
|
defaults.bot_icon = bot_icon
|
||||||
|
changed = True
|
||||||
|
restart_required = True
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
if changed:
|
||||||
|
save_config(config)
|
||||||
|
return settings_payload(requires_restart=restart_required)
|
||||||
|
|
||||||
|
|
||||||
|
def update_provider_settings(query: QueryParams) -> dict[str, Any]:
|
||||||
|
provider_name = (_query_first(query, "provider") or "").strip()
|
||||||
|
if not provider_name:
|
||||||
|
raise WebUISettingsError("provider is required")
|
||||||
|
spec = find_by_name(provider_name)
|
||||||
|
if spec is None or spec.is_oauth:
|
||||||
|
raise WebUISettingsError("unknown provider")
|
||||||
|
|
||||||
|
config = load_config()
|
||||||
|
provider_config = getattr(config.providers, spec.name, None)
|
||||||
|
if provider_config is None:
|
||||||
|
raise WebUISettingsError("unknown provider")
|
||||||
|
|
||||||
|
changed = False
|
||||||
|
if "api_key" in query or "apiKey" in query:
|
||||||
|
api_key = _query_first_alias(query, "api_key", "apiKey")
|
||||||
|
api_key = (api_key or "").strip() or None
|
||||||
|
if provider_config.api_key != api_key:
|
||||||
|
provider_config.api_key = api_key
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
if "api_base" in query or "apiBase" in query:
|
||||||
|
api_base = _query_first_alias(query, "api_base", "apiBase")
|
||||||
|
api_base = (api_base or "").strip() or None
|
||||||
|
if provider_config.api_base != api_base:
|
||||||
|
provider_config.api_base = api_base
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
if changed:
|
||||||
|
save_config(config)
|
||||||
|
image_config = config.tools.image_generation
|
||||||
|
restart_required = (
|
||||||
|
changed
|
||||||
|
and image_config.enabled
|
||||||
|
and image_config.provider == spec.name
|
||||||
|
and get_image_gen_provider(spec.name) is not None
|
||||||
|
)
|
||||||
|
return settings_payload(requires_restart=restart_required)
|
||||||
|
|
||||||
|
|
||||||
|
def update_web_search_settings(query: QueryParams) -> dict[str, Any]:
|
||||||
|
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")
|
||||||
|
|
||||||
|
config = load_config()
|
||||||
|
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", "")
|
||||||
|
else:
|
||||||
|
api_key = _query_first_alias(query, "api_key", "apiKey")
|
||||||
|
api_key = api_key.strip() if api_key is not None else None
|
||||||
|
if not api_key and previous_provider == provider_name and search_config.api_key:
|
||||||
|
api_key = search_config.api_key
|
||||||
|
if not api_key:
|
||||||
|
raise WebUISettingsError("api_key is required")
|
||||||
|
set_search_value("api_key", api_key)
|
||||||
|
set_search_value("base_url", "")
|
||||||
|
|
||||||
|
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:
|
||||||
|
normalized = use_jina_reader.strip().lower()
|
||||||
|
if normalized not in {"1", "0", "true", "false", "yes", "no"}:
|
||||||
|
raise WebUISettingsError("use_jina_reader must be boolean")
|
||||||
|
previous_jina_reader = web_config.fetch.use_jina_reader
|
||||||
|
set_fetch_value("use_jina_reader", normalized in {"1", "true", "yes"})
|
||||||
|
if web_config.fetch.use_jina_reader != previous_jina_reader:
|
||||||
|
restart_required = True
|
||||||
|
|
||||||
|
if changed:
|
||||||
|
save_config(config)
|
||||||
|
return settings_payload(requires_restart=restart_required)
|
||||||
|
|
||||||
|
|
||||||
|
def update_image_generation_settings(query: QueryParams) -> dict[str, Any]:
|
||||||
|
config = load_config()
|
||||||
|
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)
|
||||||
|
if provider["name"] == image_config.provider
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if not selected_provider or not selected_provider["configured"]:
|
||||||
|
raise WebUISettingsError("image generation provider is not configured")
|
||||||
|
|
||||||
|
if changed:
|
||||||
|
save_config(config)
|
||||||
|
return settings_payload(requires_restart=changed)
|
||||||
@@ -0,0 +1,193 @@
|
|||||||
|
"""Persisted WebUI sidebar workspace state.
|
||||||
|
|
||||||
|
This state is UI-only metadata, scoped to the active nanobot instance data
|
||||||
|
directory (the directory containing the current config.json). It deliberately
|
||||||
|
does not modify agent sessions.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from nanobot.config.paths import get_webui_dir
|
||||||
|
|
||||||
|
WEBUI_SIDEBAR_STATE_SCHEMA_VERSION = 1
|
||||||
|
_MAX_STATE_FILE_BYTES = 256 * 1024
|
||||||
|
_MAX_LIST_ITEMS = 2_000
|
||||||
|
_MAX_MAP_ITEMS = 2_000
|
||||||
|
_MAX_KEY_LEN = 512
|
||||||
|
_MAX_TITLE_LEN = 160
|
||||||
|
_MAX_TAG_LEN = 40
|
||||||
|
_ALLOWED_DENSITIES = {"comfortable", "compact"}
|
||||||
|
_ALLOWED_SORTS = {"updated_desc", "created_desc", "title_asc"}
|
||||||
|
|
||||||
|
|
||||||
|
def webui_sidebar_state_path() -> Path:
|
||||||
|
return get_webui_dir() / "sidebar-state.json"
|
||||||
|
|
||||||
|
|
||||||
|
def default_webui_sidebar_state() -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"schema_version": WEBUI_SIDEBAR_STATE_SCHEMA_VERSION,
|
||||||
|
"pinned_keys": [],
|
||||||
|
"archived_keys": [],
|
||||||
|
"title_overrides": {},
|
||||||
|
"tags_by_key": {},
|
||||||
|
"collapsed_groups": {},
|
||||||
|
"view": {
|
||||||
|
"density": "comfortable",
|
||||||
|
"show_previews": False,
|
||||||
|
"show_timestamps": False,
|
||||||
|
"show_archived": False,
|
||||||
|
"sort": "updated_desc",
|
||||||
|
},
|
||||||
|
"updated_at": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_string(value: Any, *, max_len: int = _MAX_KEY_LEN) -> str | None:
|
||||||
|
if not isinstance(value, str):
|
||||||
|
return None
|
||||||
|
cleaned = value.strip()
|
||||||
|
if not cleaned:
|
||||||
|
return None
|
||||||
|
return cleaned[:max_len]
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_string_list(value: Any, *, max_len: int = _MAX_KEY_LEN) -> list[str]:
|
||||||
|
if not isinstance(value, list):
|
||||||
|
return []
|
||||||
|
out: list[str] = []
|
||||||
|
seen: set[str] = set()
|
||||||
|
for item in value[:_MAX_LIST_ITEMS]:
|
||||||
|
cleaned = _clean_string(item, max_len=max_len)
|
||||||
|
if cleaned is None or cleaned in seen:
|
||||||
|
continue
|
||||||
|
seen.add(cleaned)
|
||||||
|
out.append(cleaned)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_bool_map(value: Any) -> dict[str, bool]:
|
||||||
|
if not isinstance(value, dict):
|
||||||
|
return {}
|
||||||
|
out: dict[str, bool] = {}
|
||||||
|
for key, raw in list(value.items())[:_MAX_MAP_ITEMS]:
|
||||||
|
cleaned_key = _clean_string(key)
|
||||||
|
if cleaned_key is None:
|
||||||
|
continue
|
||||||
|
out[cleaned_key] = bool(raw)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_title_overrides(value: Any) -> dict[str, str]:
|
||||||
|
if not isinstance(value, dict):
|
||||||
|
return {}
|
||||||
|
out: dict[str, str] = {}
|
||||||
|
for key, raw_title in list(value.items())[:_MAX_MAP_ITEMS]:
|
||||||
|
cleaned_key = _clean_string(key)
|
||||||
|
cleaned_title = _clean_string(raw_title, max_len=_MAX_TITLE_LEN)
|
||||||
|
if cleaned_key is None or cleaned_title is None:
|
||||||
|
continue
|
||||||
|
out[cleaned_key] = cleaned_title
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_tags_by_key(value: Any) -> dict[str, list[str]]:
|
||||||
|
if not isinstance(value, dict):
|
||||||
|
return {}
|
||||||
|
out: dict[str, list[str]] = {}
|
||||||
|
for key, raw_tags in list(value.items())[:_MAX_MAP_ITEMS]:
|
||||||
|
cleaned_key = _clean_string(key)
|
||||||
|
if cleaned_key is None:
|
||||||
|
continue
|
||||||
|
tags = _clean_string_list(raw_tags, max_len=_MAX_TAG_LEN)[:12]
|
||||||
|
if tags:
|
||||||
|
out[cleaned_key] = tags
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_view(value: Any) -> dict[str, Any]:
|
||||||
|
default = default_webui_sidebar_state()["view"]
|
||||||
|
if not isinstance(value, dict):
|
||||||
|
return dict(default)
|
||||||
|
density = value.get("density")
|
||||||
|
sort = value.get("sort")
|
||||||
|
return {
|
||||||
|
"density": density if density in _ALLOWED_DENSITIES else default["density"],
|
||||||
|
"show_previews": bool(value.get("show_previews", default["show_previews"])),
|
||||||
|
"show_timestamps": bool(value.get("show_timestamps", default["show_timestamps"])),
|
||||||
|
"show_archived": bool(value.get("show_archived", default["show_archived"])),
|
||||||
|
"sort": sort if sort in _ALLOWED_SORTS else default["sort"],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_webui_sidebar_state(raw: Any) -> dict[str, Any]:
|
||||||
|
"""Return a schema-v1 sidebar state from any older/partial input."""
|
||||||
|
if not isinstance(raw, dict):
|
||||||
|
raw = {}
|
||||||
|
state = default_webui_sidebar_state()
|
||||||
|
state["pinned_keys"] = _clean_string_list(raw.get("pinned_keys"))
|
||||||
|
state["archived_keys"] = _clean_string_list(raw.get("archived_keys"))
|
||||||
|
state["title_overrides"] = _clean_title_overrides(raw.get("title_overrides"))
|
||||||
|
state["tags_by_key"] = _clean_tags_by_key(raw.get("tags_by_key"))
|
||||||
|
state["collapsed_groups"] = _clean_bool_map(raw.get("collapsed_groups"))
|
||||||
|
state["view"] = _clean_view(raw.get("view"))
|
||||||
|
updated_at = raw.get("updated_at")
|
||||||
|
state["updated_at"] = updated_at if isinstance(updated_at, str) else None
|
||||||
|
return state
|
||||||
|
|
||||||
|
|
||||||
|
def read_webui_sidebar_state() -> dict[str, Any]:
|
||||||
|
path = webui_sidebar_state_path()
|
||||||
|
if not path.is_file():
|
||||||
|
return default_webui_sidebar_state()
|
||||||
|
try:
|
||||||
|
if path.stat().st_size > _MAX_STATE_FILE_BYTES:
|
||||||
|
logger.warning("webui sidebar state too large, ignoring: {}", path)
|
||||||
|
return default_webui_sidebar_state()
|
||||||
|
with open(path, encoding="utf-8") as f:
|
||||||
|
raw = json.load(f)
|
||||||
|
except (OSError, json.JSONDecodeError) as e:
|
||||||
|
logger.warning("read webui sidebar state failed {}: {}", path, e)
|
||||||
|
return default_webui_sidebar_state()
|
||||||
|
return normalize_webui_sidebar_state(raw)
|
||||||
|
|
||||||
|
|
||||||
|
def write_webui_sidebar_state(raw: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
state = normalize_webui_sidebar_state(raw)
|
||||||
|
state["updated_at"] = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
|
||||||
|
encoded = json.dumps(
|
||||||
|
state,
|
||||||
|
ensure_ascii=False,
|
||||||
|
indent=2,
|
||||||
|
sort_keys=True,
|
||||||
|
).encode("utf-8")
|
||||||
|
if len(encoded) > _MAX_STATE_FILE_BYTES:
|
||||||
|
raise ValueError("sidebar state is too large")
|
||||||
|
|
||||||
|
path = webui_sidebar_state_path()
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
tmp = path.with_suffix(".json.tmp")
|
||||||
|
with open(tmp, "wb") as f:
|
||||||
|
f.write(encoded)
|
||||||
|
f.write(b"\n")
|
||||||
|
f.flush()
|
||||||
|
os.fsync(f.fileno())
|
||||||
|
os.replace(tmp, path)
|
||||||
|
try:
|
||||||
|
dir_fd = os.open(path.parent, os.O_RDONLY)
|
||||||
|
except OSError:
|
||||||
|
return state
|
||||||
|
try:
|
||||||
|
os.fsync(dir_fd)
|
||||||
|
finally:
|
||||||
|
os.close(dir_fd)
|
||||||
|
return state
|
||||||
|
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Legacy WebUI JSON snapshot path helpers (JSON file); transcripts use webui_transcript."""
|
"""Legacy WebUI JSON snapshot path helpers (JSON file); transcripts use transcript."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -8,7 +8,7 @@ from loguru import logger
|
|||||||
|
|
||||||
from nanobot.config.paths import get_webui_dir
|
from nanobot.config.paths import get_webui_dir
|
||||||
from nanobot.session.manager import SessionManager
|
from nanobot.session.manager import SessionManager
|
||||||
from nanobot.utils.webui_transcript import delete_webui_transcript
|
from nanobot.webui.transcript import delete_webui_transcript
|
||||||
|
|
||||||
|
|
||||||
def webui_thread_file_path(session_key: str) -> Path:
|
def webui_thread_file_path(session_key: str) -> Path:
|
||||||
@@ -99,17 +99,39 @@ def tool_trace_lines_from_events(events: Any) -> list[str]:
|
|||||||
if not isinstance(events, list):
|
if not isinstance(events, list):
|
||||||
return []
|
return []
|
||||||
lines: list[str] = []
|
lines: list[str] = []
|
||||||
|
seen: set[str] = set()
|
||||||
for event in events:
|
for event in events:
|
||||||
if not event or not isinstance(event, dict):
|
if not event or not isinstance(event, dict):
|
||||||
continue
|
continue
|
||||||
if event.get("phase") != "start":
|
if event.get("phase") not in {"start", "end", "error"}:
|
||||||
continue
|
continue
|
||||||
|
call_id = event.get("call_id")
|
||||||
|
if isinstance(call_id, str) and call_id:
|
||||||
|
if call_id in seen:
|
||||||
|
continue
|
||||||
|
seen.add(call_id)
|
||||||
t = _format_tool_call_trace(event)
|
t = _format_tool_call_trace(event)
|
||||||
if t:
|
if t:
|
||||||
lines.append(t)
|
lines.append(t)
|
||||||
return lines
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
def _merge_unique_tool_trace_lines(
|
||||||
|
previous_traces: list[str],
|
||||||
|
lines: list[str],
|
||||||
|
) -> tuple[list[str], bool]:
|
||||||
|
seen_lines = set(previous_traces)
|
||||||
|
traces = list(previous_traces)
|
||||||
|
added = False
|
||||||
|
for line in lines:
|
||||||
|
if line in seen_lines:
|
||||||
|
continue
|
||||||
|
seen_lines.add(line)
|
||||||
|
traces.append(line)
|
||||||
|
added = True
|
||||||
|
return traces, added
|
||||||
|
|
||||||
|
|
||||||
def replay_transcript_to_ui_messages(
|
def replay_transcript_to_ui_messages(
|
||||||
lines: list[dict[str, Any]],
|
lines: list[dict[str, Any]],
|
||||||
*,
|
*,
|
||||||
@@ -125,11 +147,36 @@ def replay_transcript_to_ui_messages(
|
|||||||
buffer_message_id: str | None = None
|
buffer_message_id: str | None = None
|
||||||
buffer_parts: list[str] = []
|
buffer_parts: list[str] = []
|
||||||
suppress_until_turn_end = False
|
suppress_until_turn_end = False
|
||||||
|
active_activity_segment_id: str | None = None
|
||||||
|
active_file_edit_segment_id: str | None = None
|
||||||
|
activity_segment_counter = 0
|
||||||
_ts_base = int(time.time() * 1000)
|
_ts_base = int(time.time() * 1000)
|
||||||
|
|
||||||
def _new_id(prefix: str, idx: int) -> str:
|
def _new_id(prefix: str, idx: int) -> str:
|
||||||
return f"{prefix}-{idx}-{uuid.uuid4().hex[:8]}"
|
return f"{prefix}-{idx}-{uuid.uuid4().hex[:8]}"
|
||||||
|
|
||||||
|
def _new_activity_segment(*, activate: bool = True) -> str:
|
||||||
|
nonlocal active_activity_segment_id, activity_segment_counter
|
||||||
|
activity_segment_counter += 1
|
||||||
|
segment_id = f"activity-{activity_segment_counter}"
|
||||||
|
if activate:
|
||||||
|
active_activity_segment_id = segment_id
|
||||||
|
return segment_id
|
||||||
|
|
||||||
|
def _ensure_activity_segment() -> str:
|
||||||
|
return active_activity_segment_id or _new_activity_segment()
|
||||||
|
|
||||||
|
def close_activity_for_answer() -> None:
|
||||||
|
nonlocal active_activity_segment_id, active_file_edit_segment_id
|
||||||
|
active_activity_segment_id = None
|
||||||
|
active_file_edit_segment_id = None
|
||||||
|
|
||||||
|
def close_file_edit_phase_before_activity() -> None:
|
||||||
|
nonlocal active_activity_segment_id, active_file_edit_segment_id
|
||||||
|
if active_file_edit_segment_id:
|
||||||
|
active_activity_segment_id = None
|
||||||
|
active_file_edit_segment_id = None
|
||||||
|
|
||||||
def attach_reasoning_chunk(prev: list[dict[str, Any]], chunk: str, idx: int) -> None:
|
def attach_reasoning_chunk(prev: list[dict[str, Any]], chunk: str, idx: int) -> None:
|
||||||
for i in range(len(prev) - 1, -1, -1):
|
for i in range(len(prev) - 1, -1, -1):
|
||||||
candidate = prev[i]
|
candidate = prev[i]
|
||||||
@@ -151,12 +198,19 @@ def replay_transcript_to_ui_messages(
|
|||||||
**candidate,
|
**candidate,
|
||||||
"reasoning": (str(candidate.get("reasoning") or "")) + chunk,
|
"reasoning": (str(candidate.get("reasoning") or "")) + chunk,
|
||||||
"reasoningStreaming": True,
|
"reasoningStreaming": True,
|
||||||
|
"activitySegmentId": candidate.get("activitySegmentId") or _ensure_activity_segment(),
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
if not has_answer and candidate.get("isStreaming"):
|
if not has_answer and candidate.get("isStreaming"):
|
||||||
prev[i] = {**candidate, "reasoning": chunk, "reasoningStreaming": True}
|
prev[i] = {
|
||||||
|
**candidate,
|
||||||
|
"reasoning": chunk,
|
||||||
|
"reasoningStreaming": True,
|
||||||
|
"activitySegmentId": candidate.get("activitySegmentId") or _ensure_activity_segment(),
|
||||||
|
}
|
||||||
return
|
return
|
||||||
break
|
break
|
||||||
|
segment = _ensure_activity_segment()
|
||||||
prev.append(
|
prev.append(
|
||||||
{
|
{
|
||||||
"id": _new_id("as", idx),
|
"id": _new_id("as", idx),
|
||||||
@@ -165,6 +219,7 @@ def replay_transcript_to_ui_messages(
|
|||||||
"isStreaming": True,
|
"isStreaming": True,
|
||||||
"reasoning": chunk,
|
"reasoning": chunk,
|
||||||
"reasoningStreaming": True,
|
"reasoningStreaming": True,
|
||||||
|
"activitySegmentId": segment,
|
||||||
"createdAt": _ts_base + idx,
|
"createdAt": _ts_base + idx,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
@@ -221,6 +276,7 @@ def replay_transcript_to_ui_messages(
|
|||||||
return
|
return
|
||||||
|
|
||||||
def absorb_complete(extra: dict[str, Any], idx: int) -> None:
|
def absorb_complete(extra: dict[str, Any], idx: int) -> None:
|
||||||
|
nonlocal active_activity_segment_id, active_file_edit_segment_id
|
||||||
last = messages[-1] if messages else None
|
last = messages[-1] if messages else None
|
||||||
if last and is_reasoning_only_placeholder(last):
|
if last and is_reasoning_only_placeholder(last):
|
||||||
messages[-1] = {
|
messages[-1] = {
|
||||||
@@ -238,10 +294,98 @@ def replay_transcript_to_ui_messages(
|
|||||||
**extra,
|
**extra,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
active_activity_segment_id = None
|
||||||
|
active_file_edit_segment_id = None
|
||||||
|
|
||||||
|
def _file_edit_key(edit: dict[str, Any]) -> str:
|
||||||
|
call_id = str(edit.get("call_id") or "")
|
||||||
|
tool = str(edit.get("tool") or "")
|
||||||
|
if call_id:
|
||||||
|
return f"{call_id}|{tool}"
|
||||||
|
return f"{tool}|{edit.get('path') or ''}"
|
||||||
|
|
||||||
|
def find_file_edit_trace_index(
|
||||||
|
segment: str | None,
|
||||||
|
edits: list[dict[str, Any]],
|
||||||
|
) -> int | None:
|
||||||
|
incoming_keys = {_file_edit_key(edit) for edit in edits if isinstance(edit, dict)}
|
||||||
|
for i in range(len(messages) - 1, -1, -1):
|
||||||
|
candidate = messages[i]
|
||||||
|
if candidate.get("role") == "user":
|
||||||
|
break
|
||||||
|
if candidate.get("kind") != "trace" or not candidate.get("fileEdits"):
|
||||||
|
continue
|
||||||
|
if segment and candidate.get("activitySegmentId") == segment:
|
||||||
|
return i
|
||||||
|
existing_edits = candidate.get("fileEdits")
|
||||||
|
if not isinstance(existing_edits, list):
|
||||||
|
continue
|
||||||
|
for existing in existing_edits:
|
||||||
|
if isinstance(existing, dict) and _file_edit_key(existing) in incoming_keys:
|
||||||
|
return i
|
||||||
|
return None
|
||||||
|
|
||||||
|
def upsert_file_edits(edits: list[dict[str, Any]], idx: int) -> None:
|
||||||
|
nonlocal active_file_edit_segment_id
|
||||||
|
if not edits:
|
||||||
|
return
|
||||||
|
segment = active_file_edit_segment_id
|
||||||
|
target_index = find_file_edit_trace_index(segment, edits)
|
||||||
|
if target_index is not None:
|
||||||
|
last = messages[target_index]
|
||||||
|
segment = str(last.get("activitySegmentId") or segment or _new_activity_segment(activate=False))
|
||||||
|
active_file_edit_segment_id = segment
|
||||||
|
else:
|
||||||
|
if not segment:
|
||||||
|
segment = _new_activity_segment(activate=False)
|
||||||
|
active_file_edit_segment_id = segment
|
||||||
|
messages.append(
|
||||||
|
{
|
||||||
|
"id": _new_id("tr", idx),
|
||||||
|
"role": "tool",
|
||||||
|
"kind": "trace",
|
||||||
|
"content": "",
|
||||||
|
"traces": [],
|
||||||
|
"fileEdits": [],
|
||||||
|
"activitySegmentId": segment,
|
||||||
|
"createdAt": _ts_base + idx,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
target_index = len(messages) - 1
|
||||||
|
last = messages[target_index]
|
||||||
|
if not segment:
|
||||||
|
segment = _new_activity_segment(activate=False)
|
||||||
|
active_file_edit_segment_id = segment
|
||||||
|
existing = list(last.get("fileEdits") or [])
|
||||||
|
index_by_key = {
|
||||||
|
_file_edit_key(edit): pos
|
||||||
|
for pos, edit in enumerate(existing)
|
||||||
|
if isinstance(edit, dict)
|
||||||
|
}
|
||||||
|
for edit in edits:
|
||||||
|
if not isinstance(edit, dict):
|
||||||
|
continue
|
||||||
|
key = _file_edit_key(edit)
|
||||||
|
if key in index_by_key:
|
||||||
|
pos = index_by_key[key]
|
||||||
|
merged = {**existing[pos], **edit}
|
||||||
|
if edit.get("path") and not edit.get("pending"):
|
||||||
|
merged.pop("pending", None)
|
||||||
|
existing[pos] = merged
|
||||||
|
else:
|
||||||
|
index_by_key[key] = len(existing)
|
||||||
|
existing.append(dict(edit))
|
||||||
|
messages[target_index] = {
|
||||||
|
**last,
|
||||||
|
"fileEdits": existing,
|
||||||
|
"activitySegmentId": last.get("activitySegmentId") or segment,
|
||||||
|
}
|
||||||
|
|
||||||
for idx, rec in enumerate(lines):
|
for idx, rec in enumerate(lines):
|
||||||
ev = rec.get("event")
|
ev = rec.get("event")
|
||||||
if ev == "user":
|
if ev == "user":
|
||||||
|
active_activity_segment_id = None
|
||||||
|
active_file_edit_segment_id = None
|
||||||
text = rec.get("text")
|
text = rec.get("text")
|
||||||
text_s = text if isinstance(text, str) else ""
|
text_s = text if isinstance(text, str) else ""
|
||||||
media_paths = rec.get("media_paths")
|
media_paths = rec.get("media_paths")
|
||||||
@@ -264,12 +408,19 @@ def replay_transcript_to_ui_messages(
|
|||||||
messages.append(row)
|
messages.append(row)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
if ev == "file_edit":
|
||||||
|
raw_edits = rec.get("edits")
|
||||||
|
if isinstance(raw_edits, list):
|
||||||
|
upsert_file_edits([e for e in raw_edits if isinstance(e, dict)], idx)
|
||||||
|
continue
|
||||||
|
|
||||||
if ev == "delta":
|
if ev == "delta":
|
||||||
if suppress_until_turn_end:
|
if suppress_until_turn_end:
|
||||||
continue
|
continue
|
||||||
chunk = rec.get("text")
|
chunk = rec.get("text")
|
||||||
if not isinstance(chunk, str):
|
if not isinstance(chunk, str):
|
||||||
continue
|
continue
|
||||||
|
close_activity_for_answer()
|
||||||
adopted = find_active_placeholder(messages) if buffer_message_id is None else None
|
adopted = find_active_placeholder(messages) if buffer_message_id is None else None
|
||||||
if buffer_message_id is None:
|
if buffer_message_id is None:
|
||||||
if adopted:
|
if adopted:
|
||||||
@@ -308,6 +459,7 @@ def replay_transcript_to_ui_messages(
|
|||||||
chunk = rec.get("text")
|
chunk = rec.get("text")
|
||||||
if not isinstance(chunk, str) or not chunk:
|
if not isinstance(chunk, str) or not chunk:
|
||||||
continue
|
continue
|
||||||
|
close_file_edit_phase_before_activity()
|
||||||
attach_reasoning_chunk(messages, chunk, idx)
|
attach_reasoning_chunk(messages, chunk, idx)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -329,6 +481,7 @@ def replay_transcript_to_ui_messages(
|
|||||||
line = rec.get("text")
|
line = rec.get("text")
|
||||||
if not isinstance(line, str) or not line:
|
if not isinstance(line, str) or not line:
|
||||||
continue
|
continue
|
||||||
|
close_file_edit_phase_before_activity()
|
||||||
attach_reasoning_chunk(messages, line, idx)
|
attach_reasoning_chunk(messages, line, idx)
|
||||||
close_reasoning(messages)
|
close_reasoning(messages)
|
||||||
continue
|
continue
|
||||||
@@ -338,15 +491,28 @@ def replay_transcript_to_ui_messages(
|
|||||||
trace_lines = structured if structured else ([text] if isinstance(text, str) and text else [])
|
trace_lines = structured if structured else ([text] if isinstance(text, str) and text else [])
|
||||||
if not trace_lines:
|
if not trace_lines:
|
||||||
continue
|
continue
|
||||||
|
segment = _ensure_activity_segment()
|
||||||
last = messages[-1] if messages else None
|
last = messages[-1] if messages else None
|
||||||
if last and last.get("kind") == "trace" and not last.get("isStreaming"):
|
if (
|
||||||
|
last
|
||||||
|
and last.get("kind") == "trace"
|
||||||
|
and not last.get("isStreaming")
|
||||||
|
and (last.get("activitySegmentId") in (None, segment))
|
||||||
|
):
|
||||||
prev_traces = list(last.get("traces") or [last.get("content")])
|
prev_traces = list(last.get("traces") or [last.get("content")])
|
||||||
merged_traces = prev_traces + trace_lines
|
if structured:
|
||||||
messages[-1] = {
|
merged_traces, added = _merge_unique_tool_trace_lines(prev_traces, structured)
|
||||||
|
if not added:
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
merged_traces = prev_traces + trace_lines
|
||||||
|
merged = {
|
||||||
**last,
|
**last,
|
||||||
"traces": merged_traces,
|
"traces": merged_traces,
|
||||||
"content": trace_lines[-1],
|
"content": merged_traces[-1],
|
||||||
|
"activitySegmentId": last.get("activitySegmentId") or segment,
|
||||||
}
|
}
|
||||||
|
messages[-1] = merged
|
||||||
else:
|
else:
|
||||||
messages.append(
|
messages.append(
|
||||||
{
|
{
|
||||||
@@ -355,6 +521,7 @@ def replay_transcript_to_ui_messages(
|
|||||||
"kind": "trace",
|
"kind": "trace",
|
||||||
"content": trace_lines[-1],
|
"content": trace_lines[-1],
|
||||||
"traces": trace_lines,
|
"traces": trace_lines,
|
||||||
|
"activitySegmentId": segment,
|
||||||
"createdAt": _ts_base + idx,
|
"createdAt": _ts_base + idx,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
@@ -389,6 +556,8 @@ def replay_transcript_to_ui_messages(
|
|||||||
|
|
||||||
if ev == "turn_end":
|
if ev == "turn_end":
|
||||||
suppress_until_turn_end = False
|
suppress_until_turn_end = False
|
||||||
|
active_activity_segment_id = None
|
||||||
|
active_file_edit_segment_id = None
|
||||||
for i, m in enumerate(messages):
|
for i, m in enumerate(messages):
|
||||||
if m.get("isStreaming"):
|
if m.get("isStreaming"):
|
||||||
messages[i] = {**m, "isStreaming": False}
|
messages[i] = {**m, "isStreaming": False}
|
||||||
+14
-1
@@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "nanobot-ai"
|
name = "nanobot-ai"
|
||||||
version = "0.1.5.post3"
|
version = "0.2.0"
|
||||||
description = "A lightweight personal AI assistant framework"
|
description = "A lightweight personal AI assistant framework"
|
||||||
readme = { file = "README.md", content-type = "text/markdown" }
|
readme = { file = "README.md", content-type = "text/markdown" }
|
||||||
requires-python = ">=3.11"
|
requires-python = ">=3.11"
|
||||||
@@ -61,6 +61,7 @@ dependencies = [
|
|||||||
"openpyxl>=3.1.0,<4.0.0",
|
"openpyxl>=3.1.0,<4.0.0",
|
||||||
"python-pptx>=1.0.0,<2.0.0",
|
"python-pptx>=1.0.0,<2.0.0",
|
||||||
"filelock>=3.25.2",
|
"filelock>=3.25.2",
|
||||||
|
"keyring>=25.0.0,<26.0.0",
|
||||||
"boto3>=1.43.0",
|
"boto3>=1.43.0",
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -121,12 +122,22 @@ build-backend = "hatchling.build"
|
|||||||
[tool.hatch.metadata]
|
[tool.hatch.metadata]
|
||||||
allow-direct-references = true
|
allow-direct-references = true
|
||||||
|
|
||||||
|
[tool.hatch.build.hooks.custom]
|
||||||
|
# Implementation lives in the conventional `hatch_build.py` at the repo root.
|
||||||
|
|
||||||
[tool.hatch.build]
|
[tool.hatch.build]
|
||||||
include = [
|
include = [
|
||||||
"nanobot/**/*.py",
|
"nanobot/**/*.py",
|
||||||
"nanobot/templates/**/*.md",
|
"nanobot/templates/**/*.md",
|
||||||
"nanobot/skills/**/*.md",
|
"nanobot/skills/**/*.md",
|
||||||
"nanobot/skills/**/*.sh",
|
"nanobot/skills/**/*.sh",
|
||||||
|
"nanobot/web/dist/**/*",
|
||||||
|
]
|
||||||
|
# nanobot/web/dist/ is produced by `cd webui && bun run build` and is
|
||||||
|
# git-ignored. List it as an artifact so hatch ships it in both wheel and
|
||||||
|
# sdist even though VCS does not track it.
|
||||||
|
artifacts = [
|
||||||
|
"nanobot/web/dist/**/*",
|
||||||
]
|
]
|
||||||
|
|
||||||
[tool.hatch.build.targets.wheel]
|
[tool.hatch.build.targets.wheel]
|
||||||
@@ -141,7 +152,9 @@ packages = ["nanobot"]
|
|||||||
[tool.hatch.build.targets.sdist]
|
[tool.hatch.build.targets.sdist]
|
||||||
include = [
|
include = [
|
||||||
"nanobot/",
|
"nanobot/",
|
||||||
|
"nanobot/web/dist/",
|
||||||
"bridge/",
|
"bridge/",
|
||||||
|
"hatch_build.py",
|
||||||
"README.md",
|
"README.md",
|
||||||
"LICENSE",
|
"LICENSE",
|
||||||
"THIRD_PARTY_NOTICES.md",
|
"THIRD_PARTY_NOTICES.md",
|
||||||
|
|||||||
+140
-163
@@ -45,6 +45,73 @@ def _add_turns(session, turns: int, *, prefix: str = "msg") -> None:
|
|||||||
session.add_message("assistant", f"{prefix} assistant {i}")
|
session.add_message("assistant", f"{prefix} assistant {i}")
|
||||||
|
|
||||||
|
|
||||||
|
def _make_fake_compact(
|
||||||
|
loop: AgentLoop,
|
||||||
|
*,
|
||||||
|
summary: str = "Summary.",
|
||||||
|
on_archive=None,
|
||||||
|
track_archived: list | None = None,
|
||||||
|
track_count: bool = False,
|
||||||
|
):
|
||||||
|
"""Return a fake compact_idle_session that mirrors the real method's session mutation."""
|
||||||
|
from nanobot.session.manager import Session as _Session
|
||||||
|
|
||||||
|
state = {"count": 0}
|
||||||
|
|
||||||
|
async def _fake_compact(key: str, max_suffix: int = 8) -> str:
|
||||||
|
state["count"] += 1
|
||||||
|
session = loop.sessions.get_or_create(key)
|
||||||
|
|
||||||
|
tail = list(session.messages[session.last_consolidated:])
|
||||||
|
if not tail:
|
||||||
|
session.updated_at = datetime.now()
|
||||||
|
loop.sessions.save(session)
|
||||||
|
return ""
|
||||||
|
|
||||||
|
probe = _Session(
|
||||||
|
key=session.key,
|
||||||
|
messages=tail.copy(),
|
||||||
|
created_at=session.created_at,
|
||||||
|
updated_at=session.updated_at,
|
||||||
|
metadata={},
|
||||||
|
last_consolidated=0,
|
||||||
|
)
|
||||||
|
probe.retain_recent_legal_suffix(max_suffix)
|
||||||
|
kept = probe.messages
|
||||||
|
cut = len(tail) - len(kept)
|
||||||
|
archive_msgs = tail[:cut]
|
||||||
|
|
||||||
|
if not archive_msgs and not kept:
|
||||||
|
session.updated_at = datetime.now()
|
||||||
|
loop.sessions.save(session)
|
||||||
|
return ""
|
||||||
|
|
||||||
|
last_active = session.updated_at
|
||||||
|
s = summary
|
||||||
|
if archive_msgs:
|
||||||
|
if on_archive:
|
||||||
|
result = on_archive(archive_msgs)
|
||||||
|
s = result if isinstance(result, str) else summary
|
||||||
|
if track_archived is not None:
|
||||||
|
track_archived.extend(archive_msgs)
|
||||||
|
|
||||||
|
if s and s != "(nothing)":
|
||||||
|
session.metadata["_last_summary"] = {
|
||||||
|
"text": s,
|
||||||
|
"last_active": last_active.isoformat(),
|
||||||
|
}
|
||||||
|
|
||||||
|
session.messages = kept
|
||||||
|
session.last_consolidated = 0
|
||||||
|
session.updated_at = datetime.now()
|
||||||
|
loop.sessions.save(session)
|
||||||
|
return s
|
||||||
|
|
||||||
|
# Attach state for count access
|
||||||
|
_fake_compact.state = state # type: ignore[attr-defined]
|
||||||
|
return _fake_compact
|
||||||
|
|
||||||
|
|
||||||
class TestSessionTTLConfig:
|
class TestSessionTTLConfig:
|
||||||
"""Test session TTL configuration."""
|
"""Test session TTL configuration."""
|
||||||
|
|
||||||
@@ -201,10 +268,7 @@ class TestAutoCompact:
|
|||||||
s2.add_message("user", "recent")
|
s2.add_message("user", "recent")
|
||||||
loop.sessions.save(s2)
|
loop.sessions.save(s2)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
loop.auto_compact.check_expired(loop._schedule_background)
|
loop.auto_compact.check_expired(loop._schedule_background)
|
||||||
await asyncio.sleep(0.1)
|
await asyncio.sleep(0.1)
|
||||||
|
|
||||||
@@ -222,12 +286,9 @@ class TestAutoCompact:
|
|||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archived_messages = []
|
archived_messages = []
|
||||||
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
async def _fake_archive(messages):
|
loop, track_archived=archived_messages,
|
||||||
archived_messages.extend(messages)
|
)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
@@ -246,10 +307,9 @@ class TestAutoCompact:
|
|||||||
_add_turns(session, 6, prefix="hello")
|
_add_turns(session, 6, prefix="hello")
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
return "User said hello."
|
loop, summary="User said hello.",
|
||||||
|
)
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
@@ -262,23 +322,16 @@ class TestAutoCompact:
|
|||||||
|
|
||||||
@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):
|
||||||
"""_archive on empty session should not archive."""
|
"""_archive on empty session should not store a summary."""
|
||||||
loop = _make_loop(tmp_path, session_ttl_minutes=15)
|
loop = _make_loop(tmp_path, session_ttl_minutes=15)
|
||||||
|
|
||||||
archive_called = False
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
|
||||||
nonlocal archive_called
|
|
||||||
archive_called = True
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
assert not archive_called
|
|
||||||
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
|
||||||
await loop.close_mcp()
|
await loop.close_mcp()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -290,18 +343,14 @@ class TestAutoCompact:
|
|||||||
session.last_consolidated = 18
|
session.last_consolidated = 18
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archived_count = 0
|
archived_messages = []
|
||||||
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
async def _fake_archive(messages):
|
loop, track_archived=archived_messages,
|
||||||
nonlocal archived_count
|
)
|
||||||
archived_count = len(messages)
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
assert archived_count == 2
|
assert len(archived_messages) == 2
|
||||||
await loop.close_mcp()
|
await loop.close_mcp()
|
||||||
|
|
||||||
|
|
||||||
@@ -334,12 +383,9 @@ class TestAutoCompactIdleDetection:
|
|||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archived_messages = []
|
archived_messages = []
|
||||||
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
async def _fake_archive(messages):
|
loop, track_archived=archived_messages,
|
||||||
archived_messages.extend(messages)
|
)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# Simulate proactive archive completing before message arrives
|
# Simulate proactive archive completing before message arrives
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
@@ -402,10 +448,7 @@ class TestAutoCompactIdleDetection:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||||
response = await loop._process_message(msg)
|
response = await loop._process_message(msg)
|
||||||
@@ -466,10 +509,7 @@ class TestAutoCompactSystemMessages:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# Simulate proactive archive completing before system message arrives
|
# Simulate proactive archive completing before system message arrives
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
@@ -547,12 +587,9 @@ class TestAutoCompactEdgeCases:
|
|||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archived_messages = []
|
archived_messages = []
|
||||||
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
async def _fake_archive(messages):
|
loop, track_archived=archived_messages,
|
||||||
archived_messages.extend(messages)
|
)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# Simulate proactive archive completing before message arrives
|
# Simulate proactive archive completing before message arrives
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
@@ -644,10 +681,7 @@ class TestAutoCompactIntegration:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# Simulate proactive archive completing before message arrives
|
# Simulate proactive archive completing before message arrives
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
@@ -704,12 +738,9 @@ class TestProactiveAutoCompact:
|
|||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archived_messages = []
|
archived_messages = []
|
||||||
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
async def _fake_archive(messages):
|
loop, summary="User chatted about old things.", track_archived=archived_messages,
|
||||||
archived_messages.extend(messages)
|
)
|
||||||
return "User chatted about old things."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
|
|
||||||
@@ -748,14 +779,14 @@ class TestProactiveAutoCompact:
|
|||||||
started = asyncio.Event()
|
started = asyncio.Event()
|
||||||
block_forever = asyncio.Event()
|
block_forever = asyncio.Event()
|
||||||
|
|
||||||
async def _slow_archive(messages):
|
async def _slow_compact(key, max_suffix=8):
|
||||||
nonlocal archive_count
|
nonlocal archive_count
|
||||||
archive_count += 1
|
archive_count += 1
|
||||||
started.set()
|
started.set()
|
||||||
await block_forever.wait()
|
await block_forever.wait()
|
||||||
return "Summary."
|
return "Summary."
|
||||||
|
|
||||||
loop.consolidator.archive = _slow_archive
|
loop.consolidator.compact_idle_session = _slow_compact
|
||||||
|
|
||||||
# First call starts archiving via callback
|
# First call starts archiving via callback
|
||||||
loop.auto_compact.check_expired(loop._schedule_background)
|
loop.auto_compact.check_expired(loop._schedule_background)
|
||||||
@@ -781,10 +812,10 @@ class TestProactiveAutoCompact:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _failing_archive(messages):
|
async def _failing_compact(key, max_suffix=8):
|
||||||
raise RuntimeError("LLM down")
|
raise RuntimeError("LLM down")
|
||||||
|
|
||||||
loop.consolidator.archive = _failing_archive
|
loop.consolidator.compact_idle_session = _failing_compact
|
||||||
|
|
||||||
# Should not raise
|
# Should not raise
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
@@ -795,24 +826,18 @@ class TestProactiveAutoCompact:
|
|||||||
|
|
||||||
@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):
|
||||||
"""Proactive archive should not call LLM for sessions with no un-consolidated messages."""
|
"""Proactive archive should not produce a summary for sessions with no messages."""
|
||||||
loop = _make_loop(tmp_path, session_ttl_minutes=15)
|
loop = _make_loop(tmp_path, session_ttl_minutes=15)
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archive_called = False
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
|
||||||
nonlocal archive_called
|
|
||||||
archive_called = True
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
|
|
||||||
assert not archive_called
|
# Empty session should not produce a summary
|
||||||
|
assert "cli:test" not in loop.auto_compact._summaries
|
||||||
await loop.close_mcp()
|
await loop.close_mcp()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -824,18 +849,12 @@ class TestProactiveAutoCompact:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archive_count = 0
|
_fake_compact = _make_fake_compact(loop)
|
||||||
|
loop.consolidator.compact_idle_session = _fake_compact
|
||||||
async def _fake_archive(messages):
|
|
||||||
nonlocal archive_count
|
|
||||||
archive_count += 1
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# Simulate an active agent task for this session
|
# Simulate an active agent task for this session
|
||||||
await self._run_check_expired(loop, active_session_keys={"cli:test"})
|
await self._run_check_expired(loop, active_session_keys={"cli:test"})
|
||||||
assert archive_count == 0
|
assert _fake_compact.state["count"] == 0
|
||||||
|
|
||||||
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
|
||||||
@@ -851,22 +870,16 @@ class TestProactiveAutoCompact:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archive_count = 0
|
_fake_compact = _make_fake_compact(loop)
|
||||||
|
loop.consolidator.compact_idle_session = _fake_compact
|
||||||
async def _fake_archive(messages):
|
|
||||||
nonlocal archive_count
|
|
||||||
archive_count += 1
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# First tick: active task, skip
|
# First tick: active task, skip
|
||||||
await self._run_check_expired(loop, active_session_keys={"cli:test"})
|
await self._run_check_expired(loop, active_session_keys={"cli:test"})
|
||||||
assert archive_count == 0
|
assert _fake_compact.state["count"] == 0
|
||||||
|
|
||||||
# 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 archive_count == 1
|
assert _fake_compact.state["count"] == 1
|
||||||
await loop.close_mcp()
|
await loop.close_mcp()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -888,18 +901,12 @@ class TestProactiveAutoCompact:
|
|||||||
s3.add_message("user", "recent")
|
s3.add_message("user", "recent")
|
||||||
loop.sessions.save(s3)
|
loop.sessions.save(s3)
|
||||||
|
|
||||||
archive_count = 0
|
_fake_compact = _make_fake_compact(loop)
|
||||||
|
loop.consolidator.compact_idle_session = _fake_compact
|
||||||
async def _fake_archive(messages):
|
|
||||||
nonlocal archive_count
|
|
||||||
archive_count += 1
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await self._run_check_expired(loop, active_session_keys={"cli:expired_active"})
|
await self._run_check_expired(loop, active_session_keys={"cli:expired_active"})
|
||||||
|
|
||||||
assert archive_count == 1
|
assert _fake_compact.state["count"] == 1
|
||||||
s1_after = loop.sessions.get_or_create("cli:expired_idle")
|
s1_after = loop.sessions.get_or_create("cli:expired_idle")
|
||||||
assert len(s1_after.messages) == loop.auto_compact._RECENT_SUFFIX_MESSAGES
|
assert len(s1_after.messages) == loop.auto_compact._RECENT_SUFFIX_MESSAGES
|
||||||
s2_after = loop.sessions.get_or_create("cli:expired_active")
|
s2_after = loop.sessions.get_or_create("cli:expired_active")
|
||||||
@@ -917,22 +924,16 @@ class TestProactiveAutoCompact:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archive_count = 0
|
_fake_compact = _make_fake_compact(loop)
|
||||||
|
loop.consolidator.compact_idle_session = _fake_compact
|
||||||
async def _fake_archive(messages):
|
|
||||||
nonlocal archive_count
|
|
||||||
archive_count += 1
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# First tick: archives the session
|
# First tick: archives the session
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
assert archive_count == 1
|
assert _fake_compact.state["count"] == 1
|
||||||
|
|
||||||
# Second tick: should NOT re-schedule (updated_at is fresh after clear)
|
# Second tick: should NOT re-schedule (updated_at is fresh after clear)
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
assert archive_count == 1 # Still 1, not re-scheduled
|
assert _fake_compact.state["count"] == 1 # Still 1, not re-scheduled
|
||||||
await loop.close_mcp()
|
await loop.close_mcp()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -943,22 +944,15 @@ class TestProactiveAutoCompact:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archive_count = 0
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
|
||||||
nonlocal archive_count
|
|
||||||
archive_count += 1
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# First tick: skips (no messages), refreshes updated_at
|
# First tick: skips (no messages), refreshes updated_at
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
assert archive_count == 0
|
assert "cli:test" not in loop.auto_compact._summaries
|
||||||
|
|
||||||
# Second tick: should NOT re-schedule because updated_at is fresh
|
# Second tick: should NOT re-schedule because updated_at is fresh
|
||||||
await self._run_check_expired(loop)
|
await self._run_check_expired(loop)
|
||||||
assert archive_count == 0
|
assert "cli:test" not in loop.auto_compact._summaries
|
||||||
await loop.close_mcp()
|
await loop.close_mcp()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -970,18 +964,12 @@ class TestProactiveAutoCompact:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archive_count = 0
|
_fake_compact = _make_fake_compact(loop)
|
||||||
|
loop.consolidator.compact_idle_session = _fake_compact
|
||||||
async def _fake_archive(messages):
|
|
||||||
nonlocal archive_count
|
|
||||||
archive_count += 1
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# First compact cycle
|
# First compact cycle
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
assert archive_count == 1
|
assert _fake_compact.state["count"] == 1
|
||||||
|
|
||||||
# User returns, sends new messages
|
# User returns, sends new messages
|
||||||
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="second topic")
|
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="second topic")
|
||||||
@@ -995,7 +983,7 @@ class TestProactiveAutoCompact:
|
|||||||
|
|
||||||
# Second compact cycle should succeed
|
# Second compact cycle should succeed
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
assert archive_count == 2
|
assert _fake_compact.state["count"] == 2
|
||||||
await loop.close_mcp()
|
await loop.close_mcp()
|
||||||
|
|
||||||
|
|
||||||
@@ -1011,10 +999,9 @@ class TestSummaryPersistence:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
return "User said hello."
|
loop, summary="User said hello.",
|
||||||
|
)
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
@@ -1036,10 +1023,9 @@ class TestSummaryPersistence:
|
|||||||
session.updated_at = last_active
|
session.updated_at = last_active
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
return "User said hello."
|
loop, summary="User said hello.",
|
||||||
|
)
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
# Archive
|
# Archive
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
@@ -1069,10 +1055,7 @@ class TestSummaryPersistence:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
@@ -1100,10 +1083,7 @@ class TestSummaryPersistence:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
@@ -1129,10 +1109,9 @@ class TestSummaryPersistence:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
return "First summary."
|
loop, summary="First summary.",
|
||||||
|
)
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
# Consume the first summary via hot path
|
# Consume the first summary via hot path
|
||||||
@@ -1148,10 +1127,9 @@ class TestSummaryPersistence:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive2(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
return "Second summary."
|
loop, summary="Second summary.",
|
||||||
|
)
|
||||||
loop.consolidator.archive = _fake_archive2
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
# The second archive writes a new summary
|
# The second archive writes a new summary
|
||||||
@@ -1173,10 +1151,9 @@ class TestSummaryPersistence:
|
|||||||
session.updated_at = datetime.now() - timedelta(minutes=20)
|
session.updated_at = datetime.now() - timedelta(minutes=20)
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def _fake_archive(messages):
|
loop.consolidator.compact_idle_session = _make_fake_compact(
|
||||||
return "Old summary."
|
loop, summary="Old summary.",
|
||||||
|
)
|
||||||
loop.consolidator.archive = _fake_archive
|
|
||||||
await loop.auto_compact._archive("cli:test")
|
await loop.auto_compact._archive("cli:test")
|
||||||
|
|
||||||
# Verify summary exists before /new
|
# Verify summary exists before /new
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ def _make_autocompact(
|
|||||||
sessions = MagicMock(spec=SessionManager)
|
sessions = MagicMock(spec=SessionManager)
|
||||||
if consolidator is None:
|
if consolidator is None:
|
||||||
consolidator = MagicMock()
|
consolidator = MagicMock()
|
||||||
consolidator.archive = AsyncMock(return_value="Summary.")
|
consolidator.compact_idle_session = AsyncMock(return_value="Summary.")
|
||||||
return AutoCompact(
|
return AutoCompact(
|
||||||
sessions=sessions,
|
sessions=sessions,
|
||||||
consolidator=consolidator,
|
consolidator=consolidator,
|
||||||
@@ -178,62 +178,6 @@ class TestFormatSummary:
|
|||||||
assert result.startswith("Previous conversation summary (last active ")
|
assert result.startswith("Previous conversation summary (last active ")
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# _split_unconsolidated
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
class TestSplitUnconsolidated:
|
|
||||||
"""Test AutoCompact._split_unconsolidated splitting logic."""
|
|
||||||
|
|
||||||
def test_empty_session_returns_both_empty(self):
|
|
||||||
"""Empty session should return ([], [])."""
|
|
||||||
ac = _make_autocompact()
|
|
||||||
session = _make_session(messages=[])
|
|
||||||
archive, kept = ac._split_unconsolidated(session)
|
|
||||||
assert archive == []
|
|
||||||
assert kept == []
|
|
||||||
|
|
||||||
def test_all_messages_archivable_when_more_than_suffix(self):
|
|
||||||
"""Session with many messages should archive a prefix and keep suffix."""
|
|
||||||
ac = _make_autocompact()
|
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
|
||||||
session = _make_session(messages=msgs)
|
|
||||||
archive, kept = ac._split_unconsolidated(session)
|
|
||||||
assert len(archive) > 0
|
|
||||||
assert len(kept) <= AutoCompact._RECENT_SUFFIX_MESSAGES
|
|
||||||
|
|
||||||
def test_fewer_messages_than_suffix_returns_empty_archive(self):
|
|
||||||
"""Session with fewer messages than suffix should have empty archive."""
|
|
||||||
ac = _make_autocompact()
|
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(3)]
|
|
||||||
session = _make_session(messages=msgs)
|
|
||||||
archive, kept = ac._split_unconsolidated(session)
|
|
||||||
assert archive == []
|
|
||||||
assert len(kept) == len(msgs)
|
|
||||||
|
|
||||||
def test_respects_last_consolidated_offset(self):
|
|
||||||
"""Only messages after last_consolidated should be considered."""
|
|
||||||
ac = _make_autocompact()
|
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
|
||||||
# First 10 are already consolidated
|
|
||||||
session = _make_session(messages=msgs, last_consolidated=10)
|
|
||||||
archive, kept = ac._split_unconsolidated(session)
|
|
||||||
# Only the tail of 10 messages is considered for splitting
|
|
||||||
assert all(m["content"] in [f"u{i}" for i in range(10, 20)] for m in kept)
|
|
||||||
assert all(m["content"] in [f"u{i}" for i in range(10, 20)] for m in archive)
|
|
||||||
|
|
||||||
def test_retain_recent_legal_suffix_keeps_last_n(self):
|
|
||||||
"""The kept suffix should be at most _RECENT_SUFFIX_MESSAGES long."""
|
|
||||||
ac = _make_autocompact()
|
|
||||||
# 20 user messages = 20 messages total, all after last_consolidated=0
|
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
|
||||||
session = _make_session(messages=msgs)
|
|
||||||
archive, kept = ac._split_unconsolidated(session)
|
|
||||||
assert len(kept) <= AutoCompact._RECENT_SUFFIX_MESSAGES
|
|
||||||
assert len(archive) == len(msgs) - len(kept)
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# check_expired
|
# check_expired
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -313,126 +257,71 @@ class TestCheckExpired:
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
class TestArchive:
|
class TestArchiveDelegates:
|
||||||
"""Test AutoCompact._archive async method."""
|
"""_archive should delegate all session mutation to Consolidator."""
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_empty_session_updates_timestamp_no_archive_call(self):
|
async def test_calls_compact_idle_session(self):
|
||||||
"""Empty session should refresh updated_at and not call consolidator.archive."""
|
|
||||||
ac = _make_autocompact()
|
ac = _make_autocompact()
|
||||||
mock_sm = MagicMock(spec=SessionManager)
|
mock_sm = MagicMock(spec=SessionManager)
|
||||||
empty_session = _make_session(messages=[])
|
|
||||||
mock_sm.get_or_create.return_value = empty_session
|
|
||||||
ac.sessions = mock_sm
|
ac.sessions = mock_sm
|
||||||
ac.consolidator.archive = AsyncMock(return_value="Summary.")
|
ac.consolidator.compact_idle_session = AsyncMock(return_value="Summary.")
|
||||||
|
|
||||||
await ac._archive("cli:test")
|
await ac._archive("cli:test")
|
||||||
|
|
||||||
ac.consolidator.archive.assert_not_called()
|
ac.consolidator.compact_idle_session.assert_awaited_once_with(
|
||||||
mock_sm.save.assert_called_once_with(empty_session)
|
"cli:test", ac._RECENT_SUFFIX_MESSAGES,
|
||||||
# updated_at was refreshed
|
)
|
||||||
assert empty_session.updated_at > datetime.now() - timedelta(seconds=5)
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_archive_returns_empty_string_no_summary_stored(self):
|
async def test_populates_summaries_from_metadata(self):
|
||||||
"""If archive returns empty string, no summary should be stored."""
|
|
||||||
ac = _make_autocompact()
|
ac = _make_autocompact()
|
||||||
mock_sm = MagicMock(spec=SessionManager)
|
mock_sm = MagicMock(spec=SessionManager)
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
session = _make_session(
|
||||||
session = _make_session(messages=msgs)
|
metadata={"_last_summary": {"text": "Hello.", "last_active": "2026-05-13T10:00:00"}}
|
||||||
|
)
|
||||||
mock_sm.get_or_create.return_value = session
|
mock_sm.get_or_create.return_value = session
|
||||||
ac.sessions = mock_sm
|
ac.sessions = mock_sm
|
||||||
ac.consolidator.archive = AsyncMock(return_value="")
|
ac.consolidator.compact_idle_session = AsyncMock(return_value="Hello.")
|
||||||
|
|
||||||
await ac._archive("cli:test")
|
await ac._archive("cli:test")
|
||||||
|
|
||||||
assert "cli:test" not in ac._summaries
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_archive_returns_nothing_no_summary_stored(self):
|
|
||||||
"""If archive returns '(nothing)', no summary should be stored."""
|
|
||||||
ac = _make_autocompact()
|
|
||||||
mock_sm = MagicMock(spec=SessionManager)
|
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
|
||||||
session = _make_session(messages=msgs)
|
|
||||||
mock_sm.get_or_create.return_value = session
|
|
||||||
ac.sessions = mock_sm
|
|
||||||
ac.consolidator.archive = AsyncMock(return_value="(nothing)")
|
|
||||||
|
|
||||||
await ac._archive("cli:test")
|
|
||||||
|
|
||||||
assert "cli:test" not in ac._summaries
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_archive_exception_caught_key_removed_from_archiving(self):
|
|
||||||
"""If archive raises, exception is caught and key removed from _archiving."""
|
|
||||||
ac = _make_autocompact()
|
|
||||||
mock_sm = MagicMock(spec=SessionManager)
|
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
|
||||||
session = _make_session(messages=msgs)
|
|
||||||
mock_sm.get_or_create.return_value = session
|
|
||||||
ac.sessions = mock_sm
|
|
||||||
ac.consolidator.archive = AsyncMock(side_effect=RuntimeError("LLM down"))
|
|
||||||
|
|
||||||
# Should not raise
|
|
||||||
await ac._archive("cli:test")
|
|
||||||
|
|
||||||
assert "cli:test" not in ac._archiving
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_successful_archive_stores_summary_in_summaries_and_metadata(self):
|
|
||||||
"""Successful archive should store summary in _summaries dict and metadata."""
|
|
||||||
ac = _make_autocompact()
|
|
||||||
mock_sm = MagicMock(spec=SessionManager)
|
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
|
||||||
last_active = datetime(2026, 5, 13, 10, 0, 0)
|
|
||||||
session = _make_session(messages=msgs, updated_at=last_active)
|
|
||||||
mock_sm.get_or_create.return_value = session
|
|
||||||
ac.sessions = mock_sm
|
|
||||||
ac.consolidator.archive = AsyncMock(return_value="User discussed AI.")
|
|
||||||
|
|
||||||
await ac._archive("cli:test")
|
|
||||||
|
|
||||||
# _summaries
|
|
||||||
entry = ac._summaries.get("cli:test")
|
entry = ac._summaries.get("cli:test")
|
||||||
assert entry is not None
|
assert entry is not None
|
||||||
assert entry[0] == "User discussed AI."
|
assert entry[0] == "Hello."
|
||||||
assert entry[1] == last_active
|
|
||||||
# metadata
|
|
||||||
meta = session.metadata.get("_last_summary")
|
|
||||||
assert meta is not None
|
|
||||||
assert meta["text"] == "User discussed AI."
|
|
||||||
assert "last_active" in meta
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_finally_block_always_removes_from_archiving(self):
|
async def test_no_summary_when_compact_returns_empty(self):
|
||||||
"""Finally block should always remove key from _archiving, even on error."""
|
|
||||||
ac = _make_autocompact()
|
ac = _make_autocompact()
|
||||||
mock_sm = MagicMock(spec=SessionManager)
|
mock_sm = MagicMock(spec=SessionManager)
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
|
||||||
session = _make_session(messages=msgs)
|
|
||||||
mock_sm.get_or_create.return_value = session
|
|
||||||
ac.sessions = mock_sm
|
ac.sessions = mock_sm
|
||||||
ac.consolidator.archive = AsyncMock(side_effect=RuntimeError("fail"))
|
ac.consolidator.compact_idle_session = AsyncMock(return_value="")
|
||||||
|
|
||||||
# Pre-add key to archiving to verify it gets removed
|
|
||||||
ac._archiving.add("cli:test")
|
|
||||||
await ac._archive("cli:test")
|
await ac._archive("cli:test")
|
||||||
assert "cli:test" not in ac._archiving
|
|
||||||
|
assert "cli:test" not in ac._summaries
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_finally_removes_from_archiving_on_success(self):
|
async def test_no_summary_when_compact_returns_nothing(self):
|
||||||
"""Finally block should remove key from _archiving on success too."""
|
|
||||||
ac = _make_autocompact()
|
ac = _make_autocompact()
|
||||||
mock_sm = MagicMock(spec=SessionManager)
|
mock_sm = MagicMock(spec=SessionManager)
|
||||||
msgs = [{"role": "user", "content": f"u{i}"} for i in range(20)]
|
|
||||||
session = _make_session(messages=msgs)
|
|
||||||
mock_sm.get_or_create.return_value = session
|
|
||||||
ac.sessions = mock_sm
|
ac.sessions = mock_sm
|
||||||
ac.consolidator.archive = AsyncMock(return_value="Summary.")
|
ac.consolidator.compact_idle_session = AsyncMock(return_value="(nothing)")
|
||||||
|
|
||||||
|
await ac._archive("cli:test")
|
||||||
|
|
||||||
|
assert "cli:test" not in ac._summaries
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_exception_still_removes_from_archiving(self):
|
||||||
|
ac = _make_autocompact()
|
||||||
|
mock_sm = MagicMock(spec=SessionManager)
|
||||||
|
ac.sessions = mock_sm
|
||||||
|
ac.consolidator.compact_idle_session = AsyncMock(side_effect=RuntimeError("fail"))
|
||||||
|
|
||||||
ac._archiving.add("cli:test")
|
ac._archiving.add("cli:test")
|
||||||
await ac._archive("cli:test")
|
await ac._archive("cli:test")
|
||||||
|
|
||||||
assert "cli:test" not in ac._archiving
|
assert "cli:test" not in ac._archiving
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -28,6 +28,12 @@ def mock_provider():
|
|||||||
def consolidator(store, mock_provider):
|
def consolidator(store, mock_provider):
|
||||||
sessions = MagicMock()
|
sessions = MagicMock()
|
||||||
sessions.save = MagicMock()
|
sessions.save = MagicMock()
|
||||||
|
# When maybe_consolidate_by_tokens refreshes the session reference via
|
||||||
|
# get_or_create(session.key), it should get back the same object the test
|
||||||
|
# passed in. Store sessions by key so the lookup is transparent.
|
||||||
|
_session_cache: dict[str, MagicMock] = {}
|
||||||
|
sessions.get_or_create = MagicMock(side_effect=lambda key: _session_cache.get(key, MagicMock()))
|
||||||
|
sessions._session_cache = _session_cache
|
||||||
return Consolidator(
|
return Consolidator(
|
||||||
store=store,
|
store=store,
|
||||||
provider=mock_provider,
|
provider=mock_provider,
|
||||||
@@ -117,6 +123,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
session.last_consolidated = 0
|
session.last_consolidated = 0
|
||||||
session.messages = [{"role": "user", "content": "hi"}]
|
session.messages = [{"role": "user", "content": "hi"}]
|
||||||
session.key = "test:key"
|
session.key = "test:key"
|
||||||
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(100, "tiktoken"))
|
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(100, "tiktoken"))
|
||||||
consolidator.archive = AsyncMock(return_value=True)
|
consolidator.archive = AsyncMock(return_value=True)
|
||||||
await consolidator.maybe_consolidate_by_tokens(session)
|
await consolidator.maybe_consolidate_by_tokens(session)
|
||||||
@@ -152,6 +159,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
session.add_message("user", f"u{i}")
|
session.add_message("user", f"u{i}")
|
||||||
session.add_message("assistant", f"a{i}")
|
session.add_message("assistant", f"a{i}")
|
||||||
|
|
||||||
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(100, "tiktoken"))
|
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(100, "tiktoken"))
|
||||||
consolidator.archive = AsyncMock(return_value="old conversation summary")
|
consolidator.archive = AsyncMock(return_value="old conversation summary")
|
||||||
|
|
||||||
@@ -184,6 +192,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
session.add_message("tool", "tool result", tool_call_id="call-1", name="x")
|
session.add_message("tool", "tool result", tool_call_id="call-1", name="x")
|
||||||
session.add_message("assistant", "final answer")
|
session.add_message("assistant", "final answer")
|
||||||
|
|
||||||
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(100, "tiktoken"))
|
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(100, "tiktoken"))
|
||||||
consolidator.archive = AsyncMock(return_value="tool turn summary")
|
consolidator.archive = AsyncMock(return_value="tool turn summary")
|
||||||
|
|
||||||
@@ -210,6 +219,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
}
|
}
|
||||||
for i in range(70)
|
for i in range(70)
|
||||||
]
|
]
|
||||||
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(
|
consolidator.estimate_session_prompt_tokens = MagicMock(
|
||||||
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
||||||
)
|
)
|
||||||
@@ -238,6 +248,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
for i in range(70)
|
for i in range(70)
|
||||||
]
|
]
|
||||||
session.metadata = {}
|
session.metadata = {}
|
||||||
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(
|
consolidator.estimate_session_prompt_tokens = MagicMock(
|
||||||
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
||||||
)
|
)
|
||||||
@@ -263,6 +274,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
for i in range(70)
|
for i in range(70)
|
||||||
]
|
]
|
||||||
session.metadata = {}
|
session.metadata = {}
|
||||||
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
# Keep estimates high so the loop would otherwise run multiple rounds.
|
# Keep estimates high so the loop would otherwise run multiple rounds.
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(
|
consolidator.estimate_session_prompt_tokens = MagicMock(
|
||||||
return_value=(1200, "tiktoken")
|
return_value=(1200, "tiktoken")
|
||||||
@@ -287,6 +299,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
}
|
}
|
||||||
for i in range(70)
|
for i in range(70)
|
||||||
]
|
]
|
||||||
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(
|
consolidator.estimate_session_prompt_tokens = MagicMock(
|
||||||
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
||||||
)
|
)
|
||||||
@@ -299,6 +312,260 @@ class TestConsolidatorTokenBudget:
|
|||||||
assert session.last_consolidated == 61
|
assert session.last_consolidated == 61
|
||||||
|
|
||||||
|
|
||||||
|
class TestCompactIdleSession:
|
||||||
|
"""Tests for Consolidator.compact_idle_session — lock-protected idle truncation."""
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def real_consolidator(self, store, mock_provider):
|
||||||
|
"""Create a Consolidator with a real SessionManager (not a mock)."""
|
||||||
|
from nanobot.session.manager import SessionManager
|
||||||
|
|
||||||
|
sessions = SessionManager(store.workspace)
|
||||||
|
return Consolidator(
|
||||||
|
store=store,
|
||||||
|
provider=mock_provider,
|
||||||
|
model="test-model",
|
||||||
|
sessions=sessions,
|
||||||
|
context_window_tokens=1000,
|
||||||
|
build_messages=MagicMock(return_value=[]),
|
||||||
|
get_tool_definitions=MagicMock(return_value=[]),
|
||||||
|
max_completion_tokens=100,
|
||||||
|
)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_archives_prefix_keeps_suffix(self, real_consolidator, mock_provider):
|
||||||
|
"""20 user/assistant turns → compact with max_suffix=8 → messages ≤ 8,
|
||||||
|
last_consolidated=0, _last_summary stored."""
|
||||||
|
mock_provider.chat_with_retry.return_value = MagicMock(
|
||||||
|
content="Summary of old conversation.", finish_reason="stop"
|
||||||
|
)
|
||||||
|
sessions = real_consolidator.sessions
|
||||||
|
session = sessions.get_or_create("cli:test")
|
||||||
|
for i in range(20):
|
||||||
|
session.add_message("user", f"user msg {i}")
|
||||||
|
session.add_message("assistant", f"assistant msg {i}")
|
||||||
|
sessions.save(session)
|
||||||
|
|
||||||
|
result = await real_consolidator.compact_idle_session("cli:test", max_suffix=8)
|
||||||
|
assert result == "Summary of old conversation."
|
||||||
|
|
||||||
|
reloaded = sessions.get_or_create("cli:test")
|
||||||
|
assert len(reloaded.messages) <= 8
|
||||||
|
assert reloaded.last_consolidated == 0
|
||||||
|
meta = reloaded.metadata.get("_last_summary")
|
||||||
|
assert meta is not None
|
||||||
|
assert meta["text"] == "Summary of old conversation."
|
||||||
|
assert "last_active" in meta
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_empty_session_refreshes_timestamp(self, real_consolidator):
|
||||||
|
"""Empty session with old updated_at → refreshed after call, returns ''."""
|
||||||
|
from datetime import datetime, timedelta
|
||||||
|
|
||||||
|
sessions = real_consolidator.sessions
|
||||||
|
session = sessions.get_or_create("cli:empty")
|
||||||
|
old_ts = datetime.now() - timedelta(hours=2)
|
||||||
|
session.updated_at = old_ts
|
||||||
|
sessions.save(session)
|
||||||
|
|
||||||
|
result = await real_consolidator.compact_idle_session("cli:empty")
|
||||||
|
assert result == ""
|
||||||
|
|
||||||
|
reloaded = sessions.get_or_create("cli:empty")
|
||||||
|
assert reloaded.updated_at > old_ts
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_nothing_summary_not_stored(self, real_consolidator, mock_provider):
|
||||||
|
"""LLM returns '(nothing)' → _last_summary NOT in metadata."""
|
||||||
|
mock_provider.chat_with_retry.return_value = MagicMock(
|
||||||
|
content="(nothing)", finish_reason="stop"
|
||||||
|
)
|
||||||
|
sessions = real_consolidator.sessions
|
||||||
|
session = sessions.get_or_create("cli:nothing")
|
||||||
|
for i in range(10):
|
||||||
|
session.add_message("user", f"u{i}")
|
||||||
|
session.add_message("assistant", f"a{i}")
|
||||||
|
sessions.save(session)
|
||||||
|
|
||||||
|
result = await real_consolidator.compact_idle_session("cli:nothing", max_suffix=4)
|
||||||
|
assert result == "(nothing)"
|
||||||
|
|
||||||
|
reloaded = sessions.get_or_create("cli:nothing")
|
||||||
|
assert "_last_summary" not in reloaded.metadata
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_llm_failure_still_truncates(self, real_consolidator, mock_provider, store):
|
||||||
|
"""LLM raises RuntimeError → raw_archive fires, session still truncated, returns None."""
|
||||||
|
mock_provider.chat_with_retry.side_effect = RuntimeError("LLM unavailable")
|
||||||
|
sessions = real_consolidator.sessions
|
||||||
|
session = sessions.get_or_create("cli:fail")
|
||||||
|
for i in range(10):
|
||||||
|
session.add_message("user", f"u{i}")
|
||||||
|
session.add_message("assistant", f"a{i}")
|
||||||
|
sessions.save(session)
|
||||||
|
|
||||||
|
result = await real_consolidator.compact_idle_session("cli:fail", max_suffix=4)
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
# raw_archive should have been called (history.jsonl gets an entry)
|
||||||
|
entries = store.read_unprocessed_history(since_cursor=0)
|
||||||
|
assert any("[RAW]" in e["content"] for e in entries)
|
||||||
|
|
||||||
|
# Session should still be truncated
|
||||||
|
reloaded = sessions.get_or_create("cli:fail")
|
||||||
|
assert len(reloaded.messages) <= 4
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_respects_last_consolidated(self, real_consolidator, mock_provider):
|
||||||
|
"""30 turns with last_consolidated=50 → only unconsolidated tail considered."""
|
||||||
|
mock_provider.chat_with_retry.return_value = MagicMock(
|
||||||
|
content="Tail summary.", finish_reason="stop"
|
||||||
|
)
|
||||||
|
sessions = real_consolidator.sessions
|
||||||
|
session = sessions.get_or_create("cli:offset")
|
||||||
|
for i in range(30):
|
||||||
|
session.add_message("user", f"u{i}")
|
||||||
|
session.add_message("assistant", f"a{i}")
|
||||||
|
session.last_consolidated = 50 # Only 10 messages unconsolidated
|
||||||
|
sessions.save(session)
|
||||||
|
|
||||||
|
result = await real_consolidator.compact_idle_session("cli:offset", max_suffix=4)
|
||||||
|
assert result == "Tail summary."
|
||||||
|
|
||||||
|
# Verify only the unconsolidated tail was processed:
|
||||||
|
# 10 unconsolidated messages (50-59), keep suffix of 4 → archive 6
|
||||||
|
archived_call = mock_provider.chat_with_retry.call_args
|
||||||
|
user_content = archived_call.kwargs["messages"][1]["content"]
|
||||||
|
# Should contain only tail messages, not early ones
|
||||||
|
assert "u0" not in user_content
|
||||||
|
assert "u25" in user_content or "a25" in user_content
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_acquires_consolidation_lock(self, real_consolidator, mock_provider):
|
||||||
|
"""Verify lock is held during execution."""
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
# Use a slow LLM response to ensure the lock is held while we check
|
||||||
|
started = asyncio.Event()
|
||||||
|
|
||||||
|
async def slow_chat(**kwargs):
|
||||||
|
started.set()
|
||||||
|
await asyncio.sleep(0.1)
|
||||||
|
return MagicMock(content="Summary.", finish_reason="stop")
|
||||||
|
|
||||||
|
mock_provider.chat_with_retry = slow_chat
|
||||||
|
|
||||||
|
sessions = real_consolidator.sessions
|
||||||
|
session = sessions.get_or_create("cli:lock")
|
||||||
|
for i in range(10):
|
||||||
|
session.add_message("user", f"u{i}")
|
||||||
|
session.add_message("assistant", f"a{i}")
|
||||||
|
sessions.save(session)
|
||||||
|
|
||||||
|
lock = real_consolidator.get_lock("cli:lock")
|
||||||
|
assert not lock.locked()
|
||||||
|
|
||||||
|
task = asyncio.ensure_future(
|
||||||
|
real_consolidator.compact_idle_session("cli:lock", max_suffix=4)
|
||||||
|
)
|
||||||
|
await started.wait()
|
||||||
|
assert lock.locked()
|
||||||
|
await task
|
||||||
|
assert not lock.locked()
|
||||||
|
|
||||||
|
|
||||||
|
class TestConsolidatorSessionRefresh:
|
||||||
|
"""Background consolidation must detect stale session references."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reloads_before_empty_session_guard(self, tmp_path):
|
||||||
|
"""A stale empty reference must not skip a non-empty cached session."""
|
||||||
|
from nanobot.agent.memory import Consolidator, MemoryStore
|
||||||
|
from nanobot.session.manager import Session, SessionManager
|
||||||
|
|
||||||
|
store = MemoryStore(tmp_path)
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.chat_with_retry = AsyncMock(
|
||||||
|
return_value=MagicMock(content="summary", finish_reason="stop")
|
||||||
|
)
|
||||||
|
provider.generation.max_tokens = 4096
|
||||||
|
provider.estimate_prompt_tokens = MagicMock(return_value=(10, "test"))
|
||||||
|
sessions = SessionManager(tmp_path)
|
||||||
|
consolidator = Consolidator(
|
||||||
|
store=store,
|
||||||
|
provider=provider,
|
||||||
|
model="test-model",
|
||||||
|
sessions=sessions,
|
||||||
|
context_window_tokens=128_000,
|
||||||
|
build_messages=MagicMock(return_value=[]),
|
||||||
|
get_tool_definitions=MagicMock(return_value=[]),
|
||||||
|
)
|
||||||
|
|
||||||
|
fresh = sessions.get_or_create("cli:test")
|
||||||
|
fresh.add_message("user", "fresh message")
|
||||||
|
sessions.save(fresh)
|
||||||
|
stale_empty = Session(key="cli:test")
|
||||||
|
|
||||||
|
seen: dict[str, Session] = {}
|
||||||
|
|
||||||
|
def estimate(session: Session):
|
||||||
|
seen["session"] = session
|
||||||
|
return 10, "test"
|
||||||
|
|
||||||
|
consolidator.estimate_session_prompt_tokens = MagicMock(side_effect=estimate)
|
||||||
|
|
||||||
|
await consolidator.maybe_consolidate_by_tokens(stale_empty)
|
||||||
|
|
||||||
|
assert seen["session"] is fresh
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reloads_stale_session_after_compact(self, tmp_path):
|
||||||
|
"""After compact_idle_session replaces the session, a concurrent
|
||||||
|
maybe_consolidate_by_tokens with the old reference should use the
|
||||||
|
fresh session from cache instead of overwriting."""
|
||||||
|
from nanobot.agent.memory import Consolidator, MemoryStore
|
||||||
|
from nanobot.session.manager import SessionManager
|
||||||
|
|
||||||
|
store = MemoryStore(tmp_path)
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.chat_with_retry = AsyncMock(
|
||||||
|
return_value=MagicMock(content="summary", finish_reason="stop")
|
||||||
|
)
|
||||||
|
provider.generation.max_tokens = 4096
|
||||||
|
provider.estimate_prompt_tokens = MagicMock(return_value=(10, "test"))
|
||||||
|
sessions = SessionManager(tmp_path)
|
||||||
|
consolidator = Consolidator(
|
||||||
|
store=store,
|
||||||
|
provider=provider,
|
||||||
|
model="test-model",
|
||||||
|
sessions=sessions,
|
||||||
|
context_window_tokens=128_000,
|
||||||
|
build_messages=MagicMock(return_value=[]),
|
||||||
|
get_tool_definitions=MagicMock(return_value=[]),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Populate session with many messages
|
||||||
|
session = sessions.get_or_create("cli:test")
|
||||||
|
for i in range(20):
|
||||||
|
session.add_message("user", f"u{i}")
|
||||||
|
session.add_message("assistant", f"a{i}")
|
||||||
|
sessions.save(session)
|
||||||
|
|
||||||
|
# Simulate: background consolidation captures old reference
|
||||||
|
old_ref = session
|
||||||
|
|
||||||
|
# AutoCompact runs first and truncates to 8
|
||||||
|
await consolidator.compact_idle_session("cli:test", max_suffix=8)
|
||||||
|
|
||||||
|
# Background consolidation runs with stale reference —
|
||||||
|
# should detect the session was replaced and not undo the compact.
|
||||||
|
await consolidator.maybe_consolidate_by_tokens(old_ref)
|
||||||
|
|
||||||
|
session_after = sessions.get_or_create("cli:test")
|
||||||
|
# Messages should still be truncated (not restored to 40)
|
||||||
|
assert len(session_after.messages) <= 8
|
||||||
|
|
||||||
|
|
||||||
class TestRawArchiveTruncation:
|
class TestRawArchiveTruncation:
|
||||||
"""raw_archive() must cap entry size to avoid bloating history.jsonl."""
|
"""raw_archive() must cap entry size to avoid bloating history.jsonl."""
|
||||||
|
|
||||||
|
|||||||
@@ -314,8 +314,8 @@ def test_system_prompt_keeps_message_tool_out_of_current_chat_replies(tmp_path)
|
|||||||
prompt = builder.build_system_prompt(channel="slack")
|
prompt = builder.build_system_prompt(channel="slack")
|
||||||
|
|
||||||
assert "Do not use the 'message' tool for normal replies in the current chat" in prompt
|
assert "Do not use the 'message' tool for normal replies in the current chat" in prompt
|
||||||
assert "the runtime attaches those artifacts to the final assistant reply automatically" in prompt
|
assert "When 'generate_image' creates images" in prompt
|
||||||
assert "do not call 'message' just to announce or resend them" in prompt
|
assert "call 'message' with the artifact paths in the 'media' parameter" in prompt
|
||||||
assert "Wait for the tool results, then answer once" in prompt
|
assert "Wait for the tool results, then answer once" in prompt
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import pytest
|
|||||||
|
|
||||||
from nanobot.heartbeat.service import HeartbeatService
|
from nanobot.heartbeat.service import HeartbeatService
|
||||||
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
||||||
|
from nanobot.utils.llm_runtime import LLMRuntime
|
||||||
|
|
||||||
|
|
||||||
class DummyProvider(LLMProvider):
|
class DummyProvider(LLMProvider):
|
||||||
@@ -11,9 +12,11 @@ class DummyProvider(LLMProvider):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self._responses = list(responses)
|
self._responses = list(responses)
|
||||||
self.calls = 0
|
self.calls = 0
|
||||||
|
self.models: list[str | None] = []
|
||||||
|
|
||||||
async def chat(self, *args, **kwargs) -> LLMResponse:
|
async def chat(self, *args, **kwargs) -> LLMResponse:
|
||||||
self.calls += 1
|
self.calls += 1
|
||||||
|
self.models.append(kwargs.get("model"))
|
||||||
if self._responses:
|
if self._responses:
|
||||||
return self._responses.pop(0)
|
return self._responses.pop(0)
|
||||||
return LLMResponse(content="", tool_calls=[])
|
return LLMResponse(content="", tool_calls=[])
|
||||||
@@ -215,6 +218,51 @@ async def test_tick_suppresses_when_evaluator_says_no(tmp_path, monkeypatch) ->
|
|||||||
assert notified == []
|
assert notified == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_tick_uses_runtime_provider_and_model(tmp_path, monkeypatch) -> None:
|
||||||
|
"""Preset changes must apply to heartbeat decision and post-run evaluation."""
|
||||||
|
(tmp_path / "HEARTBEAT.md").write_text("- [ ] check runtime model", encoding="utf-8")
|
||||||
|
|
||||||
|
runtime_provider = DummyProvider([
|
||||||
|
LLMResponse(
|
||||||
|
content="",
|
||||||
|
tool_calls=[
|
||||||
|
ToolCallRequest(
|
||||||
|
id="hb_1",
|
||||||
|
name="heartbeat",
|
||||||
|
arguments={"action": "run", "tasks": "check runtime model"},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
),
|
||||||
|
])
|
||||||
|
runtime_model = "openai/gpt-4.1"
|
||||||
|
|
||||||
|
executed: list[str] = []
|
||||||
|
evaluated: list[tuple[LLMProvider, str]] = []
|
||||||
|
|
||||||
|
async def _on_execute(tasks: str) -> str:
|
||||||
|
executed.append(tasks)
|
||||||
|
return "runtime model produced a user-facing update"
|
||||||
|
|
||||||
|
async def _eval_capture(response, tasks, provider, model):
|
||||||
|
evaluated.append((provider, model))
|
||||||
|
return False
|
||||||
|
|
||||||
|
service = HeartbeatService(
|
||||||
|
workspace=tmp_path,
|
||||||
|
llm_runtime=lambda: LLMRuntime(runtime_provider, runtime_model),
|
||||||
|
on_execute=_on_execute,
|
||||||
|
)
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.utils.evaluator.evaluate_response", _eval_capture)
|
||||||
|
|
||||||
|
asyncio.run(service._tick())
|
||||||
|
|
||||||
|
assert runtime_provider.calls == 1
|
||||||
|
assert runtime_provider.models == [runtime_model]
|
||||||
|
assert executed == ["check runtime model"]
|
||||||
|
assert evaluated == [(runtime_provider, runtime_model)]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_decide_retries_transient_error_then_succeeds(tmp_path, monkeypatch) -> None:
|
async def test_decide_retries_transient_error_then_succeeds(tmp_path, monkeypatch) -> None:
|
||||||
provider = DummyProvider([
|
provider = DummyProvider([
|
||||||
@@ -286,4 +334,3 @@ async def test_decide_prompt_includes_current_time(tmp_path) -> None:
|
|||||||
user_msg = captured_messages[1]
|
user_msg = captured_messages[1]
|
||||||
assert user_msg["role"] == "user"
|
assert user_msg["role"] == "user"
|
||||||
assert "Current Time:" in user_msg["content"]
|
assert "Current Time:" in user_msg["content"]
|
||||||
|
|
||||||
|
|||||||
@@ -29,14 +29,15 @@ class FakeImageClient:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_generated_image_media_is_attached_to_final_assistant_message(
|
async def test_outbound_no_longer_carries_generated_media(
|
||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
"""Media delivery is now the LLM's responsibility via the message tool."""
|
||||||
set_config_path(tmp_path / "config.json")
|
set_config_path(tmp_path / "config.json")
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.agent.tools.image_generation.OpenRouterImageGenerationClient",
|
"nanobot.agent.tools.image_generation.get_image_gen_provider",
|
||||||
FakeImageClient,
|
lambda name: FakeImageClient if name == "openrouter" else None,
|
||||||
)
|
)
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
provider.get_default_model.return_value = "test-model"
|
provider.get_default_model.return_value = "test-model"
|
||||||
@@ -81,9 +82,6 @@ async def test_generated_image_media_is_attached_to_final_assistant_message(
|
|||||||
|
|
||||||
assert result is not None
|
assert result is not None
|
||||||
assert result.content == "Done"
|
assert result.content == "Done"
|
||||||
assert len(result.media) == 1
|
# OutboundMessage no longer carries generated media —
|
||||||
assert Path(result.media[0]).is_file()
|
# the LLM sends images via the message tool instead.
|
||||||
|
assert result.media == []
|
||||||
session = loop.sessions.get_or_create("websocket:chat-image")
|
|
||||||
assert session.messages[-1]["role"] == "assistant"
|
|
||||||
assert session.messages[-1]["media"] == result.media
|
|
||||||
|
|||||||
@@ -6,10 +6,15 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
import nanobot.agent.runner as runner_module
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
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.providers.base import LLMResponse, ToolCallRequest
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
from nanobot.utils.progress_events import (
|
||||||
|
invoke_file_edit_progress,
|
||||||
|
on_progress_accepts_file_edit_events,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _make_loop(tmp_path: Path) -> AgentLoop:
|
def _make_loop(tmp_path: Path) -> AgentLoop:
|
||||||
@@ -82,6 +87,143 @@ class TestToolEventProgress:
|
|||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_write_file_emits_file_edit_progress(self, tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
target = tmp_path / "foo.txt"
|
||||||
|
target.write_text("old\n", encoding="utf-8")
|
||||||
|
tool_call = ToolCallRequest(
|
||||||
|
id="call-write",
|
||||||
|
name="write_file",
|
||||||
|
arguments={"path": "foo.txt", "content": "new\nextra\n"},
|
||||||
|
)
|
||||||
|
calls = iter([
|
||||||
|
LLMResponse(content="", tool_calls=[tool_call]),
|
||||||
|
LLMResponse(content="Done", tool_calls=[]),
|
||||||
|
])
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(side_effect=lambda *a, **kw: next(calls))
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
loop.tools.prepare_call = MagicMock(
|
||||||
|
return_value=(None, {"path": "foo.txt", "content": "new\nextra\n"}, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def execute(name: str, params: dict) -> str:
|
||||||
|
target.write_text(params["content"], encoding="utf-8")
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
loop.tools.execute = AsyncMock(side_effect=execute)
|
||||||
|
file_events: list[dict] = []
|
||||||
|
|
||||||
|
async def on_progress(
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
tool_hint: bool = False,
|
||||||
|
tool_events: list[dict] | None = None,
|
||||||
|
file_edit_events: list[dict] | None = None,
|
||||||
|
) -> None:
|
||||||
|
if file_edit_events:
|
||||||
|
file_events.extend(file_edit_events)
|
||||||
|
|
||||||
|
final_content, _, _, _, _ = await loop._run_agent_loop([], on_progress=on_progress)
|
||||||
|
|
||||||
|
assert final_content == "Done"
|
||||||
|
assert [event["phase"] for event in file_events] == ["start", "end"]
|
||||||
|
assert file_events[0] == {
|
||||||
|
"version": 1,
|
||||||
|
"call_id": "call-write",
|
||||||
|
"tool": "write_file",
|
||||||
|
"path": "foo.txt",
|
||||||
|
"absolute_path": (tmp_path / "foo.txt").resolve().as_posix(),
|
||||||
|
"phase": "start",
|
||||||
|
"added": 2,
|
||||||
|
"deleted": 1,
|
||||||
|
"approximate": True,
|
||||||
|
"status": "editing",
|
||||||
|
}
|
||||||
|
assert file_events[1]["status"] == "done"
|
||||||
|
assert file_events[1]["approximate"] is False
|
||||||
|
assert (file_events[1]["added"], file_events[1]["deleted"]) == (2, 1)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_file_edit_snapshot_skipped_when_progress_callback_cannot_emit_file_edits(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
target = tmp_path / "foo.txt"
|
||||||
|
target.write_text("old\n", encoding="utf-8")
|
||||||
|
tool_call = ToolCallRequest(
|
||||||
|
id="call-write",
|
||||||
|
name="write_file",
|
||||||
|
arguments={"path": "foo.txt", "content": "new\n"},
|
||||||
|
)
|
||||||
|
calls = iter([
|
||||||
|
LLMResponse(content="", tool_calls=[tool_call]),
|
||||||
|
LLMResponse(content="Done", tool_calls=[]),
|
||||||
|
])
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(side_effect=lambda *a, **kw: next(calls))
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
loop.tools.prepare_call = MagicMock(
|
||||||
|
return_value=(None, {"path": "foo.txt", "content": "new\n"}, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def execute(name: str, params: dict) -> str:
|
||||||
|
target.write_text(params["content"], encoding="utf-8")
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
loop.tools.execute = AsyncMock(side_effect=execute)
|
||||||
|
prepare_tracker = MagicMock(side_effect=AssertionError("unexpected file snapshot"))
|
||||||
|
monkeypatch.setattr(runner_module, "prepare_file_edit_tracker", prepare_tracker)
|
||||||
|
|
||||||
|
async def on_progress(
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
tool_hint: bool = False,
|
||||||
|
tool_events: list[dict] | None = None,
|
||||||
|
) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
final_content, _, _, _, _ = await loop._run_agent_loop([], on_progress=on_progress)
|
||||||
|
|
||||||
|
assert final_content == "Done"
|
||||||
|
assert target.read_text(encoding="utf-8") == "new\n"
|
||||||
|
prepare_tracker.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_exec_does_not_emit_file_edit_progress(self, tmp_path: Path) -> None:
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
tool_call = ToolCallRequest(
|
||||||
|
id="call-exec",
|
||||||
|
name="exec",
|
||||||
|
arguments={"command": "printf hi > foo.txt"},
|
||||||
|
)
|
||||||
|
calls = iter([
|
||||||
|
LLMResponse(content="", tool_calls=[tool_call]),
|
||||||
|
LLMResponse(content="Done", tool_calls=[]),
|
||||||
|
])
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(side_effect=lambda *a, **kw: next(calls))
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
loop.tools.prepare_call = MagicMock(
|
||||||
|
return_value=(None, {"command": "printf hi > foo.txt"}, None),
|
||||||
|
)
|
||||||
|
loop.tools.execute = AsyncMock(return_value="ok")
|
||||||
|
file_events: list[dict] = []
|
||||||
|
|
||||||
|
async def on_progress(
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
tool_hint: bool = False,
|
||||||
|
tool_events: list[dict] | None = None,
|
||||||
|
file_edit_events: list[dict] | None = None,
|
||||||
|
) -> None:
|
||||||
|
if file_edit_events:
|
||||||
|
file_events.extend(file_edit_events)
|
||||||
|
|
||||||
|
await loop._run_agent_loop([], on_progress=on_progress)
|
||||||
|
|
||||||
|
assert file_events == []
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_bus_progress_forwards_tool_events_to_outbound_metadata(self, tmp_path: Path) -> None:
|
async def test_bus_progress_forwards_tool_events_to_outbound_metadata(self, tmp_path: Path) -> None:
|
||||||
"""When run() handles a bus message, _tool_events lands in OutboundMessage metadata."""
|
"""When run() handles a bus message, _tool_events lands in OutboundMessage metadata."""
|
||||||
@@ -130,6 +272,138 @@ class TestToolEventProgress:
|
|||||||
assert finish["phase"] == "end"
|
assert finish["phase"] == "end"
|
||||||
assert finish["result"] == "file.txt"
|
assert finish["result"] == "file.txt"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_bus_progress_forwards_file_edit_events_for_websocket_only(self, tmp_path: Path) -> None:
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
|
||||||
|
edit_events = [{
|
||||||
|
"call_id": "call-write",
|
||||||
|
"tool": "write_file",
|
||||||
|
"path": "foo.txt",
|
||||||
|
"phase": "start",
|
||||||
|
"added": 1,
|
||||||
|
"deleted": 0,
|
||||||
|
"approximate": True,
|
||||||
|
"status": "editing",
|
||||||
|
}]
|
||||||
|
|
||||||
|
websocket_progress = await loop._build_bus_progress_callback(InboundMessage(
|
||||||
|
channel="websocket",
|
||||||
|
sender_id="u1",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="edit",
|
||||||
|
))
|
||||||
|
assert on_progress_accepts_file_edit_events(websocket_progress) is True
|
||||||
|
await websocket_progress("", file_edit_events=edit_events)
|
||||||
|
outbound = await bus.consume_outbound()
|
||||||
|
assert outbound.metadata["_file_edit_events"] == edit_events
|
||||||
|
|
||||||
|
telegram_progress = await loop._build_bus_progress_callback(InboundMessage(
|
||||||
|
channel="telegram",
|
||||||
|
sender_id="u1",
|
||||||
|
chat_id="chat2",
|
||||||
|
content="edit",
|
||||||
|
))
|
||||||
|
assert on_progress_accepts_file_edit_events(telegram_progress) is False
|
||||||
|
await invoke_file_edit_progress(telegram_progress, edit_events)
|
||||||
|
assert bus.outbound_size == 0
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_goal_turn_keeps_live_file_edit_progress_for_webui(self, tmp_path: Path) -> None:
|
||||||
|
"""The /goal command rewrites the prompt but must not bypass WebUI file-edit progress."""
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
call_count = 0
|
||||||
|
target = tmp_path / "goal.txt"
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(*, on_tool_call_delta=None, **kwargs):
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
if call_count == 1:
|
||||||
|
assert on_tool_call_delta is not None
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": 0,
|
||||||
|
"call_id": "call-goal-write",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": '{"path":"goal.txt","content":"',
|
||||||
|
})
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": 0,
|
||||||
|
"arguments_delta": "one\\ntwo\\nthree\\n",
|
||||||
|
})
|
||||||
|
await on_tool_call_delta({"index": 0, "arguments_delta": '"}'})
|
||||||
|
return LLMResponse(
|
||||||
|
content=None,
|
||||||
|
tool_calls=[
|
||||||
|
ToolCallRequest(
|
||||||
|
id="call-goal-write",
|
||||||
|
name="write_file",
|
||||||
|
arguments={
|
||||||
|
"path": "goal.txt",
|
||||||
|
"content": "one\ntwo\nthree\n",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
usage={},
|
||||||
|
)
|
||||||
|
return LLMResponse(content="Done", tool_calls=[], usage={})
|
||||||
|
|
||||||
|
async def execute(name: str, params: dict) -> str:
|
||||||
|
assert name == "write_file"
|
||||||
|
target.write_text(params["content"], encoding="utf-8")
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
|
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[
|
||||||
|
{"type": "function", "function": {"name": "write_file"}},
|
||||||
|
])
|
||||||
|
loop.tools.prepare_call = MagicMock(
|
||||||
|
return_value=(
|
||||||
|
None,
|
||||||
|
{"path": "goal.txt", "content": "one\ntwo\nthree\n"},
|
||||||
|
None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
loop.tools.execute = AsyncMock(side_effect=execute)
|
||||||
|
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await loop._dispatch(InboundMessage(
|
||||||
|
channel="websocket",
|
||||||
|
sender_id="u1",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="/goal create goal file",
|
||||||
|
metadata={"_wants_stream": True},
|
||||||
|
))
|
||||||
|
|
||||||
|
outbound = []
|
||||||
|
while bus.outbound_size > 0:
|
||||||
|
outbound.append(await bus.consume_outbound())
|
||||||
|
|
||||||
|
edit_events = [
|
||||||
|
event
|
||||||
|
for msg in outbound
|
||||||
|
for event in msg.metadata.get("_file_edit_events", [])
|
||||||
|
]
|
||||||
|
assert any(
|
||||||
|
event["status"] == "editing"
|
||||||
|
and event["approximate"]
|
||||||
|
and event["added"] == 3
|
||||||
|
for event in edit_events
|
||||||
|
)
|
||||||
|
assert any(
|
||||||
|
event["status"] == "done"
|
||||||
|
and not event["approximate"]
|
||||||
|
and event["added"] == 3
|
||||||
|
for event in edit_events
|
||||||
|
)
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_non_streaming_channel_does_not_publish_codex_progress_deltas(
|
async def test_non_streaming_channel_does_not_publish_codex_progress_deltas(
|
||||||
self,
|
self,
|
||||||
@@ -353,8 +627,93 @@ class TestToolEventProgress:
|
|||||||
assert session_updated is not None
|
assert session_updated is not None
|
||||||
|
|
||||||
assert (session_updated.metadata or {}).get("_session_updated") is True
|
assert (session_updated.metadata or {}).get("_session_updated") is True
|
||||||
|
assert (session_updated.metadata or {}).get("_session_update_scope") == "metadata"
|
||||||
assert provider.chat_with_retry.await_count == 2
|
assert provider.chat_with_retry.await_count == 2
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_webui_title_generation_uses_turn_model_snapshot(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="Done", tool_calls=[]))
|
||||||
|
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||||
|
|
||||||
|
captured: dict[str, object] = {}
|
||||||
|
|
||||||
|
async def fake_title_after_turn(**kwargs: object) -> bool:
|
||||||
|
captured.update(kwargs)
|
||||||
|
return False
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.session.webui_turns.maybe_generate_webui_title_after_turn",
|
||||||
|
fake_title_after_turn,
|
||||||
|
)
|
||||||
|
scheduled_title: list[object] = []
|
||||||
|
|
||||||
|
def schedule_background(coro: object) -> None:
|
||||||
|
name = getattr(coro, "__qualname__", "")
|
||||||
|
if "_generate_title_and_notify" in name:
|
||||||
|
scheduled_title.append(coro)
|
||||||
|
elif hasattr(coro, "close"):
|
||||||
|
coro.close()
|
||||||
|
|
||||||
|
loop._schedule_background = schedule_background # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await loop._dispatch(InboundMessage(
|
||||||
|
channel="websocket",
|
||||||
|
sender_id="u1",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="say hello",
|
||||||
|
metadata={"webui": True},
|
||||||
|
))
|
||||||
|
|
||||||
|
assert len(scheduled_title) == 1
|
||||||
|
loop.provider = MagicMock()
|
||||||
|
loop.model = "switched-after-turn"
|
||||||
|
|
||||||
|
await scheduled_title[0] # type: ignore[misc]
|
||||||
|
|
||||||
|
assert captured["provider"] is provider
|
||||||
|
assert captured["model"] == "test-model"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_webui_command_turn_does_not_schedule_title_generation(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="Done", tool_calls=[]))
|
||||||
|
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
|
||||||
|
|
||||||
|
async def fake_title_after_turn(**_kwargs: object) -> bool:
|
||||||
|
raise AssertionError("command-only turns should not generate titles")
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.session.webui_turns.maybe_generate_webui_title_after_turn",
|
||||||
|
fake_title_after_turn,
|
||||||
|
)
|
||||||
|
scheduled: list[object] = []
|
||||||
|
loop._schedule_background = scheduled.append # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await loop._dispatch(InboundMessage(
|
||||||
|
channel="websocket",
|
||||||
|
sender_id="u1",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="/model",
|
||||||
|
metadata={"webui": True},
|
||||||
|
))
|
||||||
|
|
||||||
|
assert scheduled == []
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_non_websocket_dispatch_does_not_publish_turn_end_marker(self, tmp_path: Path) -> None:
|
async def test_non_websocket_dispatch_does_not_publish_turn_end_marker(self, tmp_path: Path) -> None:
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
|
|||||||
@@ -10,12 +10,16 @@ from nanobot.bus.events import InboundMessage
|
|||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.providers.base import LLMResponse
|
from nanobot.providers.base import LLMResponse
|
||||||
from nanobot.session.goal_state import GOAL_STATE_KEY
|
from nanobot.session.goal_state import GOAL_STATE_KEY
|
||||||
from nanobot.session.manager import Session
|
from nanobot.session.manager import Session, SessionManager
|
||||||
from nanobot.utils.webui_titles import (
|
from nanobot.session.webui_turns import (
|
||||||
|
TITLE_GENERATION_MAX_TOKENS,
|
||||||
|
TITLE_GENERATION_REASONING_EFFORT,
|
||||||
WEBUI_SESSION_METADATA_KEY,
|
WEBUI_SESSION_METADATA_KEY,
|
||||||
WEBUI_TITLE_METADATA_KEY,
|
WEBUI_TITLE_METADATA_KEY,
|
||||||
|
WebuiTurnCoordinator,
|
||||||
maybe_generate_webui_title,
|
maybe_generate_webui_title,
|
||||||
)
|
)
|
||||||
|
from nanobot.utils.llm_runtime import LLMRuntime
|
||||||
|
|
||||||
|
|
||||||
def _mk_loop() -> AgentLoop:
|
def _mk_loop() -> AgentLoop:
|
||||||
@@ -33,6 +37,22 @@ def _make_full_loop(tmp_path: Path) -> AgentLoop:
|
|||||||
return AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model")
|
return AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model")
|
||||||
|
|
||||||
|
|
||||||
|
def test_agent_loop_llm_runtime_reflects_current_provider_and_model(tmp_path: Path) -> None:
|
||||||
|
loop = _make_full_loop(tmp_path)
|
||||||
|
runtime = loop.llm_runtime()
|
||||||
|
|
||||||
|
assert runtime.provider is loop.provider
|
||||||
|
assert runtime.model == "test-model"
|
||||||
|
|
||||||
|
next_provider = MagicMock()
|
||||||
|
loop.provider = next_provider
|
||||||
|
loop.model = "next-model"
|
||||||
|
runtime = loop.llm_runtime()
|
||||||
|
|
||||||
|
assert runtime.provider is next_provider
|
||||||
|
assert runtime.model == "next-model"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_generate_webui_title_only_for_marked_webui_sessions(tmp_path: Path) -> None:
|
async def test_generate_webui_title_only_for_marked_webui_sessions(tmp_path: Path) -> None:
|
||||||
loop = _make_full_loop(tmp_path)
|
loop = _make_full_loop(tmp_path)
|
||||||
@@ -55,6 +75,11 @@ async def test_generate_webui_title_only_for_marked_webui_sessions(tmp_path: Pat
|
|||||||
assert generated is True
|
assert generated is True
|
||||||
assert session.metadata[WEBUI_TITLE_METADATA_KEY] == "优化 WebUI 侧边栏"
|
assert session.metadata[WEBUI_TITLE_METADATA_KEY] == "优化 WebUI 侧边栏"
|
||||||
loop.provider.chat_with_retry.assert_awaited_once()
|
loop.provider.chat_with_retry.assert_awaited_once()
|
||||||
|
assert loop.provider.chat_with_retry.await_args.kwargs["max_tokens"] == TITLE_GENERATION_MAX_TOKENS
|
||||||
|
assert (
|
||||||
|
loop.provider.chat_with_retry.await_args.kwargs["reasoning_effort"]
|
||||||
|
== TITLE_GENERATION_REASONING_EFFORT
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -79,6 +104,80 @@ async def test_generate_webui_title_skips_plain_websocket_sessions(tmp_path: Pat
|
|||||||
loop.provider.chat_with_retry.assert_not_awaited()
|
loop.provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_generate_webui_title_ignores_command_only_sessions(tmp_path: Path) -> None:
|
||||||
|
loop = _make_full_loop(tmp_path)
|
||||||
|
session = loop.sessions.get_or_create("websocket:command-title")
|
||||||
|
session.metadata[WEBUI_SESSION_METADATA_KEY] = True
|
||||||
|
session.add_message("user", "/model deep", _command=True)
|
||||||
|
session.add_message(
|
||||||
|
"assistant",
|
||||||
|
"Switched model preset to `deep`.\n- Model: `deepseek-v4-pro`",
|
||||||
|
_command=True,
|
||||||
|
)
|
||||||
|
loop.sessions.save(session)
|
||||||
|
|
||||||
|
generated = await maybe_generate_webui_title(
|
||||||
|
sessions=loop.sessions,
|
||||||
|
session_key="websocket:command-title",
|
||||||
|
provider=loop.provider,
|
||||||
|
model=loop.model,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert generated is False
|
||||||
|
assert WEBUI_TITLE_METADATA_KEY not in session.metadata
|
||||||
|
loop.provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
def test_webui_title_update_uses_captured_llm_runtime(
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
bus = MessageBus()
|
||||||
|
sessions = SessionManager(tmp_path)
|
||||||
|
scheduled: list[object] = []
|
||||||
|
captured: dict[str, object] = {}
|
||||||
|
|
||||||
|
async def fake_title_after_turn(**kwargs: object) -> bool:
|
||||||
|
captured.update(kwargs)
|
||||||
|
return False
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.session.webui_turns.maybe_generate_webui_title_after_turn",
|
||||||
|
fake_title_after_turn,
|
||||||
|
)
|
||||||
|
coordinator = WebuiTurnCoordinator(
|
||||||
|
bus=bus,
|
||||||
|
sessions=sessions,
|
||||||
|
schedule_background=lambda coro: scheduled.append(coro),
|
||||||
|
)
|
||||||
|
provider = MagicMock()
|
||||||
|
msg = InboundMessage(
|
||||||
|
channel="websocket",
|
||||||
|
sender_id="u1",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="say hello",
|
||||||
|
metadata={"webui": True},
|
||||||
|
)
|
||||||
|
|
||||||
|
coordinator.capture_title_context(
|
||||||
|
"websocket:chat1",
|
||||||
|
msg,
|
||||||
|
LLMRuntime(provider, "turn-model"),
|
||||||
|
)
|
||||||
|
asyncio.run(coordinator.handle_turn_end(
|
||||||
|
msg,
|
||||||
|
session_key="websocket:chat1",
|
||||||
|
latency_ms=None,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert len(scheduled) == 1
|
||||||
|
asyncio.run(scheduled[0]) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
assert captured["provider"] is provider
|
||||||
|
assert captured["model"] == "turn-model"
|
||||||
|
|
||||||
|
|
||||||
def test_save_turn_skips_multimodal_user_when_only_runtime_context() -> None:
|
def test_save_turn_skips_multimodal_user_when_only_runtime_context() -> None:
|
||||||
loop = _mk_loop()
|
loop = _mk_loop()
|
||||||
session = Session(key="test:runtime-only")
|
session = Session(key="test:runtime-only")
|
||||||
|
|||||||
@@ -1074,3 +1074,242 @@ class TestConfigurePydanticModelEmptyString:
|
|||||||
result = _configure_pydantic_model(model, "Test")
|
result = _configure_pydantic_model(model, "Test")
|
||||||
assert result is not None
|
assert result is not None
|
||||||
assert result.api_key == ""
|
assert result.api_key == ""
|
||||||
|
|
||||||
|
|
||||||
|
class TestModelPresetWizard:
|
||||||
|
"""Tests for model preset CRUD in the onboard wizard."""
|
||||||
|
|
||||||
|
def test_sync_preset_cache(self):
|
||||||
|
"""_sync_preset_cache should populate the module-level cache."""
|
||||||
|
from nanobot.cli.onboard import _MODEL_PRESET_CACHE, _sync_preset_cache
|
||||||
|
from nanobot.config.schema import ModelPresetConfig
|
||||||
|
|
||||||
|
config = Config()
|
||||||
|
config.model_presets["fast"] = ModelPresetConfig(model="gpt-4.1-mini")
|
||||||
|
config.model_presets["power"] = ModelPresetConfig(model="gpt-4.1")
|
||||||
|
_sync_preset_cache(config)
|
||||||
|
assert _MODEL_PRESET_CACHE == {"fast", "power"}
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
|
||||||
|
def test_model_preset_add(self, monkeypatch):
|
||||||
|
"""_configure_model_presets should add a new preset."""
|
||||||
|
from nanobot.cli.onboard import _MODEL_PRESET_CACHE, _configure_model_presets
|
||||||
|
from nanobot.config.schema import ModelPresetConfig
|
||||||
|
|
||||||
|
config = Config()
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
|
||||||
|
responses = iter([
|
||||||
|
"[+] Add new preset",
|
||||||
|
"my-preset",
|
||||||
|
"<- Back",
|
||||||
|
])
|
||||||
|
|
||||||
|
class FakePrompt:
|
||||||
|
def __init__(self, response):
|
||||||
|
self.response = response
|
||||||
|
|
||||||
|
def ask(self):
|
||||||
|
if isinstance(self.response, BaseException):
|
||||||
|
raise self.response
|
||||||
|
return self.response
|
||||||
|
|
||||||
|
def fake_select(*_args, **_kwargs):
|
||||||
|
return FakePrompt(next(responses))
|
||||||
|
|
||||||
|
def fake_text(*_args, **_kwargs):
|
||||||
|
return FakePrompt(next(responses))
|
||||||
|
|
||||||
|
def fake_configure(*_model, **_kwargs):
|
||||||
|
return ModelPresetConfig(model="gpt-test", temperature=0.5)
|
||||||
|
|
||||||
|
def fake_select_with_back(*_args, **_kwargs):
|
||||||
|
return next(responses)
|
||||||
|
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_select_with_back", fake_select_with_back)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
onboard_wizard, "questionary", SimpleNamespace(select=fake_select, text=fake_text)
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_configure_pydantic_model", fake_configure)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_show_section_header", lambda *a, **kw: None)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "console", SimpleNamespace(clear=lambda: None))
|
||||||
|
|
||||||
|
_configure_model_presets(config)
|
||||||
|
|
||||||
|
assert "my-preset" in config.model_presets
|
||||||
|
assert config.model_presets["my-preset"].model == "gpt-test"
|
||||||
|
assert config.model_presets["my-preset"].temperature == 0.5
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
|
||||||
|
def test_model_preset_delete(self, monkeypatch):
|
||||||
|
"""_configure_model_presets should delete an existing preset."""
|
||||||
|
from nanobot.cli.onboard import _MODEL_PRESET_CACHE, _configure_model_presets
|
||||||
|
from nanobot.config.schema import ModelPresetConfig
|
||||||
|
|
||||||
|
config = Config()
|
||||||
|
config.model_presets["old"] = ModelPresetConfig(model="x")
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
_MODEL_PRESET_CACHE.update({"old", "default"})
|
||||||
|
|
||||||
|
responses = iter([
|
||||||
|
"old (x)",
|
||||||
|
"Delete",
|
||||||
|
True,
|
||||||
|
"<- Back",
|
||||||
|
])
|
||||||
|
|
||||||
|
class FakePrompt:
|
||||||
|
def __init__(self, response):
|
||||||
|
self.response = response
|
||||||
|
|
||||||
|
def ask(self):
|
||||||
|
if isinstance(self.response, BaseException):
|
||||||
|
raise self.response
|
||||||
|
return self.response
|
||||||
|
|
||||||
|
def fake_select(*_args, **_kwargs):
|
||||||
|
return FakePrompt(next(responses))
|
||||||
|
|
||||||
|
def fake_confirm(*_args, **_kwargs):
|
||||||
|
return FakePrompt(next(responses))
|
||||||
|
|
||||||
|
def fake_select_with_back(*_args, **_kwargs):
|
||||||
|
return next(responses)
|
||||||
|
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_select_with_back", fake_select_with_back)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
onboard_wizard, "questionary", SimpleNamespace(select=fake_select, confirm=fake_confirm)
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_show_section_header", lambda *a, **kw: None)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "console", SimpleNamespace(clear=lambda: None))
|
||||||
|
|
||||||
|
_configure_model_presets(config)
|
||||||
|
|
||||||
|
assert "old" not in config.model_presets
|
||||||
|
assert "old" not in _MODEL_PRESET_CACHE
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
|
||||||
|
def test_model_preset_field_handler(self, monkeypatch):
|
||||||
|
"""_handle_model_preset_field should set a preset name from choices."""
|
||||||
|
from nanobot.cli.onboard import _MODEL_PRESET_CACHE, _handle_model_preset_field
|
||||||
|
from nanobot.config.schema import AgentDefaults
|
||||||
|
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
_MODEL_PRESET_CACHE.update({"fast", "power", "default"})
|
||||||
|
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_select_with_back", lambda *a, **kw: "fast")
|
||||||
|
|
||||||
|
defaults = AgentDefaults()
|
||||||
|
_handle_model_preset_field(defaults, "model_preset", "Model Preset", None)
|
||||||
|
assert defaults.model_preset == "fast"
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
|
||||||
|
def test_model_preset_field_handler_clear(self, monkeypatch):
|
||||||
|
"""_handle_model_preset_field should clear preset when (clear/unset) chosen."""
|
||||||
|
from nanobot.cli.onboard import _MODEL_PRESET_CACHE, _handle_model_preset_field
|
||||||
|
from nanobot.config.schema import AgentDefaults
|
||||||
|
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
_MODEL_PRESET_CACHE.add("fast")
|
||||||
|
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_select_with_back", lambda *a, **kw: "(clear/unset)")
|
||||||
|
|
||||||
|
defaults = AgentDefaults(model_preset="fast")
|
||||||
|
_handle_model_preset_field(defaults, "model_preset", "Model Preset", "fast")
|
||||||
|
assert defaults.model_preset is None
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
|
||||||
|
def test_main_menu_dispatch_includes_model_presets(self):
|
||||||
|
"""_configure_model_presets should be importable and callable."""
|
||||||
|
from nanobot.cli.onboard import _configure_model_presets
|
||||||
|
|
||||||
|
assert callable(_configure_model_presets)
|
||||||
|
|
||||||
|
def test_run_onboard_model_presets_edit(self, monkeypatch):
|
||||||
|
"""run_onboard should handle [M] Model Presets correctly."""
|
||||||
|
from nanobot.config.schema import ModelPresetConfig
|
||||||
|
|
||||||
|
initial_config = Config()
|
||||||
|
|
||||||
|
responses = iter([
|
||||||
|
"[M] Model Presets",
|
||||||
|
"[S] Save and Exit",
|
||||||
|
])
|
||||||
|
|
||||||
|
class FakePrompt:
|
||||||
|
def __init__(self, response):
|
||||||
|
self.response = response
|
||||||
|
|
||||||
|
def ask(self):
|
||||||
|
if isinstance(self.response, BaseException):
|
||||||
|
raise self.response
|
||||||
|
return self.response
|
||||||
|
|
||||||
|
def fake_select(*_args, **_kwargs):
|
||||||
|
return FakePrompt(next(responses))
|
||||||
|
|
||||||
|
preset_mutated = {"n": 0}
|
||||||
|
|
||||||
|
def fake_configure_model_presets(config):
|
||||||
|
preset_mutated["n"] += 1
|
||||||
|
config.model_presets["test"] = ModelPresetConfig(model="gpt-test")
|
||||||
|
|
||||||
|
monkeypatch.setattr(onboard_wizard, "questionary", SimpleNamespace(select=fake_select))
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_configure_model_presets", fake_configure_model_presets)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_show_main_menu_header", lambda: None)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_show_section_header", lambda *a, **kw: None)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "console", SimpleNamespace(clear=lambda: None))
|
||||||
|
|
||||||
|
result = run_onboard(initial_config)
|
||||||
|
assert result.should_save is True
|
||||||
|
assert preset_mutated["n"] == 1
|
||||||
|
assert "test" in result.config.model_presets
|
||||||
|
|
||||||
|
def test_fallback_models_field_add(self, monkeypatch):
|
||||||
|
"""_handle_fallback_models_field should add a preset name."""
|
||||||
|
from nanobot.cli.onboard import _MODEL_PRESET_CACHE, _handle_fallback_models_field
|
||||||
|
from nanobot.config.schema import AgentDefaults
|
||||||
|
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
_MODEL_PRESET_CACHE.update({"fast", "default"})
|
||||||
|
|
||||||
|
select_responses = iter(["fast"])
|
||||||
|
questionary_responses = iter(["[+] Add preset", "[Done]"])
|
||||||
|
|
||||||
|
class FakePrompt:
|
||||||
|
def __init__(self, response):
|
||||||
|
self.response = response
|
||||||
|
|
||||||
|
def ask(self):
|
||||||
|
if isinstance(self.response, BaseException):
|
||||||
|
raise self.response
|
||||||
|
return self.response
|
||||||
|
|
||||||
|
def fake_questionary_select(*_args, **_kwargs):
|
||||||
|
return FakePrompt(next(questionary_responses))
|
||||||
|
|
||||||
|
def fake_select_with_back(*_args, **_kwargs):
|
||||||
|
return next(select_responses)
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
onboard_wizard, "questionary",
|
||||||
|
SimpleNamespace(select=fake_questionary_select, press_any_key_to_continue=lambda: FakePrompt(None)),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_select_with_back", fake_select_with_back)
|
||||||
|
monkeypatch.setattr(onboard_wizard, "console", SimpleNamespace(clear=lambda: None, print=lambda *a, **kw: None))
|
||||||
|
|
||||||
|
defaults = AgentDefaults()
|
||||||
|
_handle_fallback_models_field(defaults, "fallback_models", "Fallback Models", [])
|
||||||
|
assert defaults.fallback_models == ["fast"]
|
||||||
|
_MODEL_PRESET_CACHE.clear()
|
||||||
|
|
||||||
|
def test_provider_field_handler(self, monkeypatch):
|
||||||
|
"""_handle_provider_field should set provider from choices."""
|
||||||
|
from nanobot.cli.onboard import _handle_provider_field
|
||||||
|
from nanobot.config.schema import AgentDefaults
|
||||||
|
|
||||||
|
monkeypatch.setattr(onboard_wizard, "_select_with_back", lambda *a, **kw: "anthropic")
|
||||||
|
|
||||||
|
defaults = AgentDefaults()
|
||||||
|
_handle_provider_field(defaults, "provider", "Provider", "auto")
|
||||||
|
assert defaults.provider == "anthropic"
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ import pytest
|
|||||||
|
|
||||||
from nanobot.agent.runner import AgentRunner, AgentRunSpec
|
from nanobot.agent.runner import AgentRunner, AgentRunSpec
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
from nanobot.providers.base import LLMResponse
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
|
||||||
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
|
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
|
||||||
|
|
||||||
@@ -77,3 +77,220 @@ async def test_runner_streams_provider_progress_deltas_by_default():
|
|||||||
assert result.final_content == "hello"
|
assert result.final_content == "hello"
|
||||||
assert [call.args[0] for call in progress_cb.await_args_list] == ["he", "llo"]
|
assert [call.args[0] for call in progress_cb.await_args_list] == ["he", "llo"]
|
||||||
provider.chat_with_retry.assert_not_awaited()
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_streams_live_write_file_activity_from_tool_argument_deltas(tmp_path):
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
call_count = 0
|
||||||
|
progress_events: list[dict] = []
|
||||||
|
|
||||||
|
async def progress_cb(content, *, file_edit_events=None, **kwargs):
|
||||||
|
if file_edit_events:
|
||||||
|
progress_events.extend(file_edit_events)
|
||||||
|
|
||||||
|
class Tools:
|
||||||
|
def get_definitions(self):
|
||||||
|
return [{"type": "function", "function": {"name": "write_file"}}]
|
||||||
|
|
||||||
|
def get(self, name):
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def execute(self, name, params):
|
||||||
|
assert name == "write_file"
|
||||||
|
assert any(event["approximate"] and event["added"] == 24 for event in progress_events)
|
||||||
|
target = tmp_path / params["path"]
|
||||||
|
target.write_text(params["content"], encoding="utf-8")
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(*, on_tool_call_delta=None, **kwargs):
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
if call_count == 1:
|
||||||
|
assert on_tool_call_delta is not None
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": 0,
|
||||||
|
"call_id": "call-write",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": '{"path":"big.txt","content":"',
|
||||||
|
})
|
||||||
|
await on_tool_call_delta({"index": 0, "arguments_delta": "line\\n" * 24})
|
||||||
|
return LLMResponse(
|
||||||
|
content=None,
|
||||||
|
tool_calls=[
|
||||||
|
ToolCallRequest(
|
||||||
|
id="call-write",
|
||||||
|
name="write_file",
|
||||||
|
arguments={"path": "big.txt", "content": "line\n" * 24},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
usage={},
|
||||||
|
)
|
||||||
|
return LLMResponse(content="done", tool_calls=[], usage={})
|
||||||
|
|
||||||
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
result = await runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[{"role": "user", "content": "write a large file"}],
|
||||||
|
tools=Tools(),
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=2,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
|
workspace=tmp_path,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.final_content == "done"
|
||||||
|
assert any(event["approximate"] and event["added"] == 24 for event in progress_events)
|
||||||
|
assert any(
|
||||||
|
not event["approximate"] and event["phase"] == "end" and event["added"] == 24
|
||||||
|
for event in progress_events
|
||||||
|
)
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_streams_live_edit_file_activity_from_tool_argument_deltas(tmp_path):
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
call_count = 0
|
||||||
|
progress_events: list[dict] = []
|
||||||
|
target = tmp_path / "notes.txt"
|
||||||
|
target.write_text("old\nkeep\n", encoding="utf-8")
|
||||||
|
|
||||||
|
async def progress_cb(content, *, file_edit_events=None, **kwargs):
|
||||||
|
if file_edit_events:
|
||||||
|
progress_events.extend(file_edit_events)
|
||||||
|
|
||||||
|
class Tools:
|
||||||
|
def get_definitions(self):
|
||||||
|
return [{"type": "function", "function": {"name": "edit_file"}}]
|
||||||
|
|
||||||
|
def get(self, name):
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def execute(self, name, params):
|
||||||
|
assert name == "edit_file"
|
||||||
|
assert any(
|
||||||
|
event["tool"] == "edit_file"
|
||||||
|
and event["approximate"]
|
||||||
|
and event["added"] == 3
|
||||||
|
and event["deleted"] == 2
|
||||||
|
for event in progress_events
|
||||||
|
)
|
||||||
|
target.write_text(params["new_text"], encoding="utf-8")
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(*, on_tool_call_delta=None, **kwargs):
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
if call_count == 1:
|
||||||
|
assert on_tool_call_delta is not None
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": 0,
|
||||||
|
"call_id": "call-edit",
|
||||||
|
"name": "edit_file",
|
||||||
|
"arguments_delta": (
|
||||||
|
'{"path":"notes.txt","old_text":"old\\nkeep\\n","new_text":"'
|
||||||
|
),
|
||||||
|
})
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": 0,
|
||||||
|
"arguments_delta": "new\\nkeep\\nextra\\n",
|
||||||
|
})
|
||||||
|
await on_tool_call_delta({"index": 0, "arguments_delta": '"}'})
|
||||||
|
return LLMResponse(
|
||||||
|
content=None,
|
||||||
|
tool_calls=[
|
||||||
|
ToolCallRequest(
|
||||||
|
id="call-edit",
|
||||||
|
name="edit_file",
|
||||||
|
arguments={
|
||||||
|
"path": "notes.txt",
|
||||||
|
"old_text": "old\nkeep\n",
|
||||||
|
"new_text": "new\nkeep\nextra\n",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
usage={},
|
||||||
|
)
|
||||||
|
return LLMResponse(content="done", tool_calls=[], usage={})
|
||||||
|
|
||||||
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
result = await runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[{"role": "user", "content": "edit a file"}],
|
||||||
|
tools=Tools(),
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=2,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
|
workspace=tmp_path,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.final_content == "done"
|
||||||
|
assert any(
|
||||||
|
event["tool"] == "edit_file"
|
||||||
|
and event["approximate"]
|
||||||
|
and event["added"] == 3
|
||||||
|
and event["deleted"] == 2
|
||||||
|
for event in progress_events
|
||||||
|
)
|
||||||
|
assert any(
|
||||||
|
event["tool"] == "edit_file"
|
||||||
|
and not event["approximate"]
|
||||||
|
and event["phase"] == "end"
|
||||||
|
and event["added"] == 2
|
||||||
|
and event["deleted"] == 1
|
||||||
|
for event in progress_events
|
||||||
|
)
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_marks_unfinished_live_write_file_activity_failed(tmp_path):
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
progress_events: list[dict] = []
|
||||||
|
|
||||||
|
async def progress_cb(content, *, file_edit_events=None, **kwargs):
|
||||||
|
if file_edit_events:
|
||||||
|
progress_events.extend(file_edit_events)
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(*, on_tool_call_delta=None, **kwargs):
|
||||||
|
assert on_tool_call_delta is not None
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": 0,
|
||||||
|
"call_id": "call-write",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": '{"path":"aborted.txt","content":"partial\\n',
|
||||||
|
})
|
||||||
|
return LLMResponse(content="stopped", tool_calls=[], finish_reason="stop", usage={})
|
||||||
|
|
||||||
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
|
tools = MagicMock()
|
||||||
|
tools.get_definitions.return_value = [{"type": "function", "function": {"name": "write_file"}}]
|
||||||
|
tools.get.return_value = None
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
result = await runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[{"role": "user", "content": "write a large file"}],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=1,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
|
workspace=tmp_path,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.final_content == "stopped"
|
||||||
|
assert progress_events[-1]["path"] == "aborted.txt"
|
||||||
|
assert progress_events[-1]["phase"] == "error"
|
||||||
|
assert progress_events[-1]["status"] == "error"
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|||||||
@@ -47,3 +47,28 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
|
|||||||
assert loop.dream.provider is new_provider
|
assert loop.dream.provider is new_provider
|
||||||
assert loop.dream.model == "new-model"
|
assert loop.dream.model == "new-model"
|
||||||
assert loop.dream._runner.provider is new_provider
|
assert loop.dream._runner.provider is new_provider
|
||||||
|
|
||||||
|
|
||||||
|
def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
|
||||||
|
old_provider = _provider("old-model")
|
||||||
|
new_provider = _provider("new-model", max_tokens=456)
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=MessageBus(),
|
||||||
|
provider=old_provider,
|
||||||
|
workspace=tmp_path,
|
||||||
|
model="old-model",
|
||||||
|
context_window_tokens=1000,
|
||||||
|
provider_snapshot_loader=lambda: ProviderSnapshot(
|
||||||
|
provider=new_provider,
|
||||||
|
model="new-model",
|
||||||
|
context_window_tokens=2000,
|
||||||
|
signature=("new-model",),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
runtime = loop.llm_runtime()
|
||||||
|
|
||||||
|
assert runtime.provider is new_provider
|
||||||
|
assert runtime.model == "new-model"
|
||||||
|
assert loop.provider is new_provider
|
||||||
|
assert loop.runner.provider is new_provider
|
||||||
|
|||||||
@@ -1,34 +0,0 @@
|
|||||||
"""Tests for staging attachment paths into the media bucket for session replay."""
|
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from nanobot.config.loader import set_config_path
|
|
||||||
from nanobot.config.paths import get_media_dir
|
|
||||||
from nanobot.utils.session_attachments import stage_media_paths_for_session_replay
|
|
||||||
|
|
||||||
|
|
||||||
def test_persist_media_stages_workspace_file(tmp_path: Path) -> None:
|
|
||||||
set_config_path(tmp_path / "config.json")
|
|
||||||
outside = tmp_path / "workspace" / "report.md"
|
|
||||||
outside.parent.mkdir(parents=True)
|
|
||||||
outside.write_text("body", encoding="utf-8")
|
|
||||||
|
|
||||||
out = stage_media_paths_for_session_replay([str(outside)])
|
|
||||||
|
|
||||||
assert len(out) == 1
|
|
||||||
staged = Path(out[0])
|
|
||||||
assert staged.is_file()
|
|
||||||
assert staged.read_text(encoding="utf-8") == "body"
|
|
||||||
assert staged.resolve().is_relative_to(get_media_dir().resolve())
|
|
||||||
|
|
||||||
|
|
||||||
def test_persist_media_keeps_files_already_under_media_root(tmp_path: Path) -> None:
|
|
||||||
set_config_path(tmp_path / "config.json")
|
|
||||||
media = get_media_dir("websocket")
|
|
||||||
media.mkdir(parents=True, exist_ok=True)
|
|
||||||
inside = media / "keep-me.txt"
|
|
||||||
inside.write_text("x", encoding="utf-8")
|
|
||||||
|
|
||||||
out = stage_media_paths_for_session_replay([str(inside.resolve())])
|
|
||||||
|
|
||||||
assert out == [str(inside.resolve())]
|
|
||||||
@@ -387,6 +387,7 @@ class TestConsolidationUnaffectedByUnifiedSession:
|
|||||||
|
|
||||||
session = Session(key="unified:default")
|
session = Session(key="unified:default")
|
||||||
session.messages = [{"role": "user", "content": "msg"}]
|
session.messages = [{"role": "user", "content": "msg"}]
|
||||||
|
sessions.get_or_create.return_value = session
|
||||||
|
|
||||||
# Simulate over-budget: estimated > budget
|
# Simulate over-budget: estimated > budget
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(950, "tiktoken"))
|
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(950, "tiktoken"))
|
||||||
|
|||||||
@@ -111,6 +111,23 @@ def test_discover_plugins_loads_entry_points():
|
|||||||
assert result["line"] is _FakePlugin
|
assert result["line"] is _FakePlugin
|
||||||
|
|
||||||
|
|
||||||
|
def test_discover_plugins_skips_names_outside_enabled_set():
|
||||||
|
from nanobot.channels.registry import discover_plugins
|
||||||
|
|
||||||
|
loaded: list[str] = []
|
||||||
|
|
||||||
|
def _load_disabled():
|
||||||
|
loaded.append("disabled")
|
||||||
|
return _FakePlugin
|
||||||
|
|
||||||
|
ep = SimpleNamespace(name="disabled", load=_load_disabled)
|
||||||
|
with patch(_EP_TARGET, return_value=[ep]):
|
||||||
|
result = discover_plugins({"enabled"})
|
||||||
|
|
||||||
|
assert result == {}
|
||||||
|
assert loaded == []
|
||||||
|
|
||||||
|
|
||||||
def test_discover_plugins_handles_load_error():
|
def test_discover_plugins_handles_load_error():
|
||||||
from nanobot.channels.registry import discover_plugins
|
from nanobot.channels.registry import discover_plugins
|
||||||
|
|
||||||
@@ -152,6 +169,25 @@ def test_discover_all_includes_external_plugin():
|
|||||||
assert result["line"] is _FakePlugin
|
assert result["line"] is _FakePlugin
|
||||||
|
|
||||||
|
|
||||||
|
def test_discover_enabled_imports_only_enabled_builtins():
|
||||||
|
from nanobot.channels.registry import discover_enabled
|
||||||
|
|
||||||
|
loaded: list[str] = []
|
||||||
|
|
||||||
|
def _load_channel(name: str):
|
||||||
|
loaded.append(name)
|
||||||
|
return _FakePlugin
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("nanobot.channels.registry.load_channel_class", side_effect=_load_channel),
|
||||||
|
patch(_EP_TARGET, return_value=[]),
|
||||||
|
):
|
||||||
|
result = discover_enabled({"enabled"}, _names=["enabled", "disabled"])
|
||||||
|
|
||||||
|
assert result == {"enabled": _FakePlugin}
|
||||||
|
assert loaded == ["enabled"]
|
||||||
|
|
||||||
|
|
||||||
def test_discover_all_builtin_shadows_plugin():
|
def test_discover_all_builtin_shadows_plugin():
|
||||||
from nanobot.channels.registry import discover_all
|
from nanobot.channels.registry import discover_all
|
||||||
|
|
||||||
@@ -180,7 +216,7 @@ async def test_manager_loads_plugin_from_dict_config():
|
|||||||
)
|
)
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
"nanobot.channels.registry.discover_all",
|
"nanobot.channels.registry.discover_enabled",
|
||||||
return_value={"fakeplugin": _FakePlugin},
|
return_value={"fakeplugin": _FakePlugin},
|
||||||
):
|
):
|
||||||
mgr = ChannelManager.__new__(ChannelManager)
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
@@ -210,7 +246,7 @@ async def test_manager_propagates_groq_transcription_api_base_to_channels():
|
|||||||
)
|
)
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
"nanobot.channels.registry.discover_all",
|
"nanobot.channels.registry.discover_enabled",
|
||||||
return_value={"fakeplugin": _FakePlugin},
|
return_value={"fakeplugin": _FakePlugin},
|
||||||
):
|
):
|
||||||
mgr = ChannelManager.__new__(ChannelManager)
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
@@ -246,7 +282,7 @@ async def test_manager_propagates_openai_transcription_api_base_to_channels():
|
|||||||
)
|
)
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
"nanobot.channels.registry.discover_all",
|
"nanobot.channels.registry.discover_enabled",
|
||||||
return_value={"fakeplugin": _FakePlugin},
|
return_value={"fakeplugin": _FakePlugin},
|
||||||
):
|
):
|
||||||
mgr = ChannelManager.__new__(ChannelManager)
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
@@ -498,10 +534,8 @@ async def test_manager_skips_disabled_plugin():
|
|||||||
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
)
|
)
|
||||||
|
|
||||||
with patch(
|
ep = _make_entry_point("fakeplugin", _FakePlugin)
|
||||||
"nanobot.channels.registry.discover_all",
|
with patch(_EP_TARGET, return_value=[ep]):
|
||||||
return_value={"fakeplugin": _FakePlugin},
|
|
||||||
):
|
|
||||||
mgr = ChannelManager.__new__(ChannelManager)
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
mgr.config = fake_config
|
mgr.config = fake_config
|
||||||
mgr.bus = MessageBus()
|
mgr.bus = MessageBus()
|
||||||
|
|||||||
@@ -29,7 +29,8 @@ from nanobot.channels.websocket import (
|
|||||||
publish_runtime_model_update,
|
publish_runtime_model_update,
|
||||||
)
|
)
|
||||||
from nanobot.config.loader import load_config, save_config
|
from nanobot.config.loader import load_config, save_config
|
||||||
from nanobot.config.schema import Config
|
from nanobot.config.schema import Config, ModelPresetConfig
|
||||||
|
from nanobot.webui.settings_api import settings_payload
|
||||||
|
|
||||||
# -- Shared helpers (aligned with test_websocket_integration.py) ---------------
|
# -- Shared helpers (aligned with test_websocket_integration.py) ---------------
|
||||||
|
|
||||||
@@ -370,6 +371,55 @@ async def test_send_progress_includes_structured_tool_events() -> None:
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_file_edit_progress_uses_file_edit_event() -> None:
|
||||||
|
bus = MagicMock()
|
||||||
|
channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"]}, bus)
|
||||||
|
mock_ws = AsyncMock()
|
||||||
|
channel._attach(mock_ws, "chat-1")
|
||||||
|
|
||||||
|
await channel.send(OutboundMessage(
|
||||||
|
channel="websocket",
|
||||||
|
chat_id="chat-1",
|
||||||
|
content="",
|
||||||
|
metadata={
|
||||||
|
"_progress": True,
|
||||||
|
"_file_edit_events": [
|
||||||
|
{
|
||||||
|
"version": 1,
|
||||||
|
"phase": "start",
|
||||||
|
"call_id": "call-1",
|
||||||
|
"tool": "write_file",
|
||||||
|
"path": "src/app.py",
|
||||||
|
"added": 12,
|
||||||
|
"deleted": 2,
|
||||||
|
"approximate": True,
|
||||||
|
"status": "editing",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
},
|
||||||
|
))
|
||||||
|
|
||||||
|
payload = json.loads(mock_ws.send.await_args.args[0])
|
||||||
|
assert payload == {
|
||||||
|
"event": "file_edit",
|
||||||
|
"chat_id": "chat-1",
|
||||||
|
"edits": [
|
||||||
|
{
|
||||||
|
"version": 1,
|
||||||
|
"phase": "start",
|
||||||
|
"call_id": "call-1",
|
||||||
|
"tool": "write_file",
|
||||||
|
"path": "src/app.py",
|
||||||
|
"added": 12,
|
||||||
|
"deleted": 2,
|
||||||
|
"approximate": True,
|
||||||
|
"status": "editing",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_send_progress_includes_agent_ui_blob() -> None:
|
async def test_send_progress_includes_agent_ui_blob() -> None:
|
||||||
bus = MagicMock()
|
bus = MagicMock()
|
||||||
@@ -707,7 +757,7 @@ async def test_maybe_push_turn_run_wall_clock_skips_when_no_active_turn() -> Non
|
|||||||
channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"]}, bus)
|
channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"]}, bus)
|
||||||
mock_ws = AsyncMock()
|
mock_ws = AsyncMock()
|
||||||
channel._attach(mock_ws, "chat-1")
|
channel._attach(mock_ws, "chat-1")
|
||||||
from nanobot.utils import webui_turn_helpers as wth
|
from nanobot.session import webui_turns as wth
|
||||||
|
|
||||||
wth._WEBSOCKET_TURN_WALL_STARTED_AT.clear()
|
wth._WEBSOCKET_TURN_WALL_STARTED_AT.clear()
|
||||||
await channel._maybe_push_turn_run_wall_clock("chat-1")
|
await channel._maybe_push_turn_run_wall_clock("chat-1")
|
||||||
@@ -720,7 +770,7 @@ async def test_maybe_push_turn_run_wall_clock_replays_running() -> None:
|
|||||||
channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"]}, bus)
|
channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"]}, bus)
|
||||||
mock_ws = AsyncMock()
|
mock_ws = AsyncMock()
|
||||||
channel._attach(mock_ws, "chat-1")
|
channel._attach(mock_ws, "chat-1")
|
||||||
from nanobot.utils import webui_turn_helpers as wth
|
from nanobot.session import webui_turns as wth
|
||||||
|
|
||||||
wth._WEBSOCKET_TURN_WALL_STARTED_AT.clear()
|
wth._WEBSOCKET_TURN_WALL_STARTED_AT.clear()
|
||||||
try:
|
try:
|
||||||
@@ -758,6 +808,25 @@ async def test_send_session_updated_emits_session_updated_event() -> None:
|
|||||||
assert body == {"event": "session_updated", "chat_id": "chat-1"}
|
assert body == {"event": "session_updated", "chat_id": "chat-1"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_session_updated_includes_scope_when_present() -> None:
|
||||||
|
bus = MagicMock()
|
||||||
|
channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"]}, bus)
|
||||||
|
mock_ws = AsyncMock()
|
||||||
|
channel._attach(mock_ws, "chat-1")
|
||||||
|
|
||||||
|
await channel.send(OutboundMessage(
|
||||||
|
channel="websocket",
|
||||||
|
chat_id="chat-1",
|
||||||
|
content="",
|
||||||
|
metadata={"_session_updated": True, "_session_update_scope": "metadata"},
|
||||||
|
))
|
||||||
|
|
||||||
|
mock_ws.send.assert_awaited_once()
|
||||||
|
body = json.loads(mock_ws.send.await_args.args[0])
|
||||||
|
assert body == {"event": "session_updated", "chat_id": "chat-1", "scope": "metadata"}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_send_non_connection_closed_exception_is_raised() -> None:
|
async def test_send_non_connection_closed_exception_is_raised() -> None:
|
||||||
bus = MagicMock()
|
bus = MagicMock()
|
||||||
@@ -923,6 +992,11 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
config = Config()
|
config = Config()
|
||||||
config.agents.defaults.model = "openai/gpt-4o"
|
config.agents.defaults.model = "openai/gpt-4o"
|
||||||
config.providers.openai.api_key = "secret-key"
|
config.providers.openai.api_key = "secret-key"
|
||||||
|
config.model_presets["deep"] = ModelPresetConfig(
|
||||||
|
model="anthropic/claude-opus-4-5",
|
||||||
|
provider="anthropic",
|
||||||
|
reasoning_effort="high",
|
||||||
|
)
|
||||||
config.tools.web.search.provider = "brave"
|
config.tools.web.search.provider = "brave"
|
||||||
config.tools.web.search.api_key = "brave-secret"
|
config.tools.web.search.api_key = "brave-secret"
|
||||||
save_config(config, config_path)
|
save_config(config, config_path)
|
||||||
@@ -943,16 +1017,52 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
body = settings.json()
|
body = settings.json()
|
||||||
assert body["agent"]["model"] == "openai/gpt-4o"
|
assert body["agent"]["model"] == "openai/gpt-4o"
|
||||||
assert body["agent"]["provider"] == "openai"
|
assert body["agent"]["provider"] == "openai"
|
||||||
|
assert body["agent"]["model_preset"] == "default"
|
||||||
|
assert body["agent"]["max_tokens"] == 8192
|
||||||
|
assert body["agent"]["timezone"] == "UTC"
|
||||||
|
assert body["agent"]["tool_hint_max_length"] == 40
|
||||||
|
presets = {preset["name"]: preset for preset in body["model_presets"]}
|
||||||
|
assert presets["default"]["active"] is True
|
||||||
|
assert presets["deep"]["reasoning_effort"] == "high"
|
||||||
providers = {provider["name"]: provider for provider in body["providers"]}
|
providers = {provider["name"]: provider for provider in body["providers"]}
|
||||||
assert providers["openai"]["configured"] is True
|
assert providers["openai"]["configured"] is True
|
||||||
assert providers["openai"]["api_key_hint"] == "secr••••-key"
|
assert providers["openai"]["api_key_hint"] == "secr••••-key"
|
||||||
|
assert providers["azure_openai"]["api_key_required"] is True
|
||||||
assert providers["openrouter"]["configured"] is False
|
assert providers["openrouter"]["configured"] is False
|
||||||
|
assert providers["openrouter"]["api_key_required"] is True
|
||||||
|
assert providers["skywork"]["label"] == "Skywork"
|
||||||
|
assert providers["skywork"]["default_api_base"] == "https://api.apifree.ai/agent/v1"
|
||||||
|
assert providers["ant_ling"]["label"] == "Ant Ling"
|
||||||
|
assert providers["ant_ling"]["default_api_base"] == "https://api.ant-ling.com/v1"
|
||||||
|
assert providers["atomic_chat"]["configured"] is False
|
||||||
|
assert providers["atomic_chat"]["api_key_required"] is False
|
||||||
|
assert providers["atomic_chat"]["default_api_base"] == "http://localhost:1337/v1"
|
||||||
assert body["agent"]["has_api_key"] is True
|
assert body["agent"]["has_api_key"] is True
|
||||||
assert body["web_search"]["provider"] == "brave"
|
assert body["web_search"]["provider"] == "brave"
|
||||||
assert body["web_search"]["api_key_hint"] == "brav••••cret"
|
assert body["web_search"]["api_key_hint"] == "brav••••cret"
|
||||||
|
assert body["web_search"]["max_results"] == 5
|
||||||
|
assert body["web"]["fetch"]["use_jina_reader"] is True
|
||||||
search_providers = {provider["name"]: provider for provider in body["web_search"]["providers"]}
|
search_providers = {provider["name"]: provider for provider in body["web_search"]["providers"]}
|
||||||
assert search_providers["duckduckgo"]["credential"] == "none"
|
assert search_providers["duckduckgo"]["credential"] == "none"
|
||||||
assert search_providers["searxng"]["credential"] == "base_url"
|
assert search_providers["searxng"]["credential"] == "base_url"
|
||||||
|
assert body["image_generation"]["enabled"] is False
|
||||||
|
assert body["image_generation"]["provider"] == "openrouter"
|
||||||
|
assert body["image_generation"]["provider_configured"] is False
|
||||||
|
assert body["image_generation"]["default_aspect_ratio"] == "1:1"
|
||||||
|
image_providers = {
|
||||||
|
provider["name"]: provider
|
||||||
|
for provider in body["image_generation"]["providers"]
|
||||||
|
}
|
||||||
|
assert image_providers["openrouter"]["label"] == "OpenRouter"
|
||||||
|
assert image_providers["openrouter"]["configured"] is False
|
||||||
|
assert image_providers["gemini"]["label"] == "Gemini"
|
||||||
|
assert body["runtime"]["config_path"] == str(config_path)
|
||||||
|
workspace_path = body["runtime"]["workspace_path"].replace("\\", "/")
|
||||||
|
assert workspace_path.endswith("/.nanobot/workspace")
|
||||||
|
assert body["runtime"]["gateway_port"] == 18790
|
||||||
|
assert body["advanced"]["exec_enabled"] is True
|
||||||
|
assert body["advanced"]["mcp_server_count"] == 0
|
||||||
|
assert body["restart_required_sections"] == []
|
||||||
assert "secret-key" not in settings.text
|
assert "secret-key" not in settings.text
|
||||||
assert "brave-secret" not in settings.text
|
assert "brave-secret" not in settings.text
|
||||||
|
|
||||||
@@ -967,38 +1077,137 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
assert provider_body["requires_restart"] is False
|
assert provider_body["requires_restart"] is False
|
||||||
provider_rows = {provider["name"]: provider for provider in provider_body["providers"]}
|
provider_rows = {provider["name"]: provider for provider in provider_body["providers"]}
|
||||||
assert provider_rows["openrouter"]["configured"] is True
|
assert provider_rows["openrouter"]["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
|
||||||
|
|
||||||
|
local_provider_updated = await _http_get(
|
||||||
|
"http://127.0.0.1:"
|
||||||
|
f"{port}/api/settings/provider/update?provider=atomic_chat"
|
||||||
|
"&api_base=http%3A%2F%2Flocalhost%3A1337%2Fv1",
|
||||||
|
headers={"Authorization": "Bearer tok"},
|
||||||
|
)
|
||||||
|
assert local_provider_updated.status_code == 200
|
||||||
|
local_provider_body = local_provider_updated.json()
|
||||||
|
local_provider_rows = {
|
||||||
|
provider["name"]: provider for provider in local_provider_body["providers"]
|
||||||
|
}
|
||||||
|
assert local_provider_rows["atomic_chat"]["configured"] is True
|
||||||
|
assert "localhost:1337" in local_provider_updated.text
|
||||||
|
|
||||||
updated = await _http_get(
|
updated = await _http_get(
|
||||||
"http://127.0.0.1:"
|
"http://127.0.0.1:"
|
||||||
f"{port}/api/settings/update?model=openrouter/test"
|
f"{port}/api/settings/update?model=atomic_chat/test"
|
||||||
"&provider=openrouter",
|
"&provider=atomic_chat&timezone=Asia%2FShanghai"
|
||||||
|
"&bot_name=Nano&bot_icon=N&tool_hint_max_length=120",
|
||||||
headers={"Authorization": "Bearer tok"},
|
headers={"Authorization": "Bearer tok"},
|
||||||
)
|
)
|
||||||
assert updated.status_code == 200
|
assert updated.status_code == 200
|
||||||
assert updated.json()["requires_restart"] is False
|
updated_body = updated.json()
|
||||||
|
assert updated_body["requires_restart"] is True
|
||||||
|
assert updated_body["restart_required_sections"] == ["runtime"]
|
||||||
|
|
||||||
|
preset_updated = await _http_get(
|
||||||
|
"http://127.0.0.1:"
|
||||||
|
f"{port}/api/settings/update?model_preset=deep",
|
||||||
|
headers={"Authorization": "Bearer tok"},
|
||||||
|
)
|
||||||
|
assert preset_updated.status_code == 200
|
||||||
|
assert preset_updated.json()["agent"]["model"] == "anthropic/claude-opus-4-5"
|
||||||
|
|
||||||
|
bad_preset = await _http_get(
|
||||||
|
"http://127.0.0.1:"
|
||||||
|
f"{port}/api/settings/update?model_preset=missing",
|
||||||
|
headers={"Authorization": "Bearer tok"},
|
||||||
|
)
|
||||||
|
assert bad_preset.status_code == 400
|
||||||
|
|
||||||
search_updated = await _http_get(
|
search_updated = await _http_get(
|
||||||
"http://127.0.0.1:"
|
"http://127.0.0.1:"
|
||||||
f"{port}/api/settings/web-search/update?provider=searxng"
|
f"{port}/api/settings/web-search/update?provider=searxng"
|
||||||
"&base_url=https%3A%2F%2Fsearch.example.com",
|
"&base_url=https%3A%2F%2Fsearch.example.com"
|
||||||
|
"&max_results=8&timeout=45&use_jina_reader=false",
|
||||||
headers={"Authorization": "Bearer tok"},
|
headers={"Authorization": "Bearer tok"},
|
||||||
)
|
)
|
||||||
assert search_updated.status_code == 200
|
assert search_updated.status_code == 200
|
||||||
search_body = search_updated.json()
|
search_body = search_updated.json()
|
||||||
assert search_body["requires_restart"] is False
|
assert search_body["requires_restart"] is True
|
||||||
|
assert search_body["restart_required_sections"] == ["runtime", "web"]
|
||||||
assert search_body["web_search"]["provider"] == "searxng"
|
assert search_body["web_search"]["provider"] == "searxng"
|
||||||
assert search_body["web_search"]["api_key_hint"] is None
|
assert search_body["web_search"]["api_key_hint"] is None
|
||||||
assert search_body["web_search"]["base_url"] == "https://search.example.com"
|
assert search_body["web_search"]["base_url"] == "https://search.example.com"
|
||||||
|
assert search_body["web_search"]["max_results"] == 8
|
||||||
|
assert search_body["web"]["fetch"]["use_jina_reader"] is False
|
||||||
|
|
||||||
|
image_updated = await _http_get(
|
||||||
|
"http://127.0.0.1:"
|
||||||
|
f"{port}/api/settings/image-generation/update?enabled=true"
|
||||||
|
"&provider=openrouter&model=openai%2Fgpt-image-1"
|
||||||
|
"&default_aspect_ratio=16%3A9&default_image_size=2K"
|
||||||
|
"&max_images_per_turn=3",
|
||||||
|
headers={"Authorization": "Bearer tok"},
|
||||||
|
)
|
||||||
|
assert image_updated.status_code == 200
|
||||||
|
image_body = image_updated.json()
|
||||||
|
assert image_body["requires_restart"] is True
|
||||||
|
assert image_body["restart_required_sections"] == ["image", "runtime", "web"]
|
||||||
|
assert image_body["image_generation"]["enabled"] is True
|
||||||
|
assert image_body["image_generation"]["model"] == "openai/gpt-image-1"
|
||||||
|
assert image_body["image_generation"]["default_aspect_ratio"] == "16:9"
|
||||||
|
assert image_body["image_generation"]["default_image_size"] == "2K"
|
||||||
|
assert image_body["image_generation"]["max_images_per_turn"] == 3
|
||||||
|
|
||||||
|
image_provider_updated = await _http_get(
|
||||||
|
"http://127.0.0.1:"
|
||||||
|
f"{port}/api/settings/provider/update?provider=openrouter"
|
||||||
|
"&api_key=sk-or-next&api_base=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1",
|
||||||
|
headers={"Authorization": "Bearer tok"},
|
||||||
|
)
|
||||||
|
assert image_provider_updated.status_code == 200
|
||||||
|
assert image_provider_updated.json()["requires_restart"] is True
|
||||||
|
assert image_provider_updated.json()["restart_required_sections"] == [
|
||||||
|
"image",
|
||||||
|
"runtime",
|
||||||
|
"web",
|
||||||
|
]
|
||||||
|
assert "sk-or-next" not in image_provider_updated.text
|
||||||
|
|
||||||
|
bad_web = await _http_get(
|
||||||
|
"http://127.0.0.1:"
|
||||||
|
f"{port}/api/settings/web-search/update?provider=duckduckgo&max_results=99",
|
||||||
|
headers={"Authorization": "Bearer tok"},
|
||||||
|
)
|
||||||
|
assert bad_web.status_code == 400
|
||||||
|
|
||||||
|
bad_image = await _http_get(
|
||||||
|
"http://127.0.0.1:"
|
||||||
|
f"{port}/api/settings/image-generation/update?provider=missing",
|
||||||
|
headers={"Authorization": "Bearer tok"},
|
||||||
|
)
|
||||||
|
assert bad_image.status_code == 400
|
||||||
|
|
||||||
saved = load_config(config_path)
|
saved = load_config(config_path)
|
||||||
assert saved.agents.defaults.model == "openrouter/test"
|
assert saved.agents.defaults.model == "atomic_chat/test"
|
||||||
assert saved.agents.defaults.provider == "openrouter"
|
assert saved.agents.defaults.provider == "atomic_chat"
|
||||||
assert saved.providers.openrouter.api_key == "sk-or-test"
|
assert saved.agents.defaults.model_preset == "deep"
|
||||||
|
assert saved.agents.defaults.timezone == "Asia/Shanghai"
|
||||||
|
assert saved.agents.defaults.bot_name == "Nano"
|
||||||
|
assert saved.agents.defaults.bot_icon == "N"
|
||||||
|
assert saved.agents.defaults.tool_hint_max_length == 120
|
||||||
|
assert saved.providers.openrouter.api_key == "sk-or-next"
|
||||||
assert saved.providers.openrouter.api_base == "https://openrouter.ai/api/v1"
|
assert saved.providers.openrouter.api_base == "https://openrouter.ai/api/v1"
|
||||||
|
assert saved.providers.atomic_chat.api_base == "http://localhost:1337/v1"
|
||||||
assert saved.tools.web.search.provider == "searxng"
|
assert saved.tools.web.search.provider == "searxng"
|
||||||
assert saved.tools.web.search.api_key == ""
|
assert saved.tools.web.search.api_key == ""
|
||||||
assert saved.tools.web.search.base_url == "https://search.example.com"
|
assert saved.tools.web.search.base_url == "https://search.example.com"
|
||||||
|
assert saved.tools.web.search.max_results == 8
|
||||||
|
assert saved.tools.web.search.timeout == 45
|
||||||
|
assert saved.tools.web.fetch.use_jina_reader is False
|
||||||
|
assert saved.tools.image_generation.enabled is True
|
||||||
|
assert saved.tools.image_generation.provider == "openrouter"
|
||||||
|
assert saved.tools.image_generation.model == "openai/gpt-image-1"
|
||||||
|
assert saved.tools.image_generation.default_aspect_ratio == "16:9"
|
||||||
|
assert saved.tools.image_generation.default_image_size == "2K"
|
||||||
|
assert saved.tools.image_generation.max_images_per_turn == 3
|
||||||
finally:
|
finally:
|
||||||
await channel.stop()
|
await channel.stop()
|
||||||
await server_task
|
await server_task
|
||||||
@@ -1043,7 +1252,7 @@ def test_settings_payload_normalizes_camel_case_provider(
|
|||||||
save_config(config, config_path)
|
save_config(config, config_path)
|
||||||
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
||||||
|
|
||||||
body = _ch(bus)._settings_payload()
|
body = settings_payload()
|
||||||
|
|
||||||
assert body["agent"]["provider"] == "minimax_anthropic"
|
assert body["agent"]["provider"] == "minimax_anthropic"
|
||||||
|
|
||||||
@@ -1460,6 +1669,54 @@ def test_parse_envelope_rejects_legacy_and_garbage() -> None:
|
|||||||
assert _parse_envelope('{"type":123}') is None
|
assert _parse_envelope('{"type":123}') is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_sessions_list_includes_active_run_started_at() -> None:
|
||||||
|
from websockets.datastructures import Headers
|
||||||
|
from websockets.http11 import Request
|
||||||
|
|
||||||
|
from nanobot.session import webui_turns as wth
|
||||||
|
|
||||||
|
bus = MagicMock()
|
||||||
|
channel = _ch(bus)
|
||||||
|
channel._api_tokens["tok"] = time.monotonic() + 300.0
|
||||||
|
channel._session_manager = MagicMock()
|
||||||
|
channel._session_manager.list_sessions.return_value = [
|
||||||
|
{
|
||||||
|
"key": "websocket:chat-1",
|
||||||
|
"created_at": "2026-05-19T10:00:00Z",
|
||||||
|
"updated_at": "2026-05-19T10:01:00Z",
|
||||||
|
"title": "Running",
|
||||||
|
"preview": "work",
|
||||||
|
"path": "/private/path",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"key": "cli:chat-2",
|
||||||
|
"created_at": "2026-05-19T10:00:00Z",
|
||||||
|
"updated_at": "2026-05-19T10:01:00Z",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
wth._WEBSOCKET_TURN_WALL_STARTED_AT.clear()
|
||||||
|
try:
|
||||||
|
wth._WEBSOCKET_TURN_WALL_STARTED_AT["chat-1"] = 1_700_000_000.0
|
||||||
|
req = Request("/api/sessions", Headers([("Authorization", "Bearer tok")]))
|
||||||
|
resp = channel._handle_sessions_list(req)
|
||||||
|
finally:
|
||||||
|
wth._WEBSOCKET_TURN_WALL_STARTED_AT.clear()
|
||||||
|
|
||||||
|
assert resp.status_code == 200
|
||||||
|
body = json.loads(resp.body.decode())
|
||||||
|
assert body["sessions"] == [
|
||||||
|
{
|
||||||
|
"key": "websocket:chat-1",
|
||||||
|
"created_at": "2026-05-19T10:00:00Z",
|
||||||
|
"updated_at": "2026-05-19T10:01:00Z",
|
||||||
|
"title": "Running",
|
||||||
|
"preview": "work",
|
||||||
|
"run_started_at": 1_700_000_000.0,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
("value", "expected"),
|
("value", "expected"),
|
||||||
[
|
[
|
||||||
@@ -1486,7 +1743,7 @@ def test_handle_webui_thread_get_returns_json(tmp_path, monkeypatch) -> None:
|
|||||||
from websockets.datastructures import Headers
|
from websockets.datastructures import Headers
|
||||||
from websockets.http11 import Request
|
from websockets.http11 import Request
|
||||||
|
|
||||||
from nanobot.utils.webui_transcript import append_transcript_object
|
from nanobot.webui.transcript import append_transcript_object
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
|
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
|
||||||
key = "websocket:c1"
|
key = "websocket:c1"
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import json
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
@@ -176,13 +177,62 @@ async def test_sessions_list_only_returns_websocket_sessions_by_default(
|
|||||||
await server_task
|
await server_task
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_webui_sidebar_state_routes_are_config_dir_scoped(
|
||||||
|
bus: MagicMock, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
|
||||||
|
sm = _seed_session(tmp_path, key="websocket:sidebar")
|
||||||
|
channel = _ch(bus, session_manager=sm, port=29911)
|
||||||
|
server_task = asyncio.create_task(channel.start())
|
||||||
|
await asyncio.sleep(0.3)
|
||||||
|
try:
|
||||||
|
boot = await _http_get("http://127.0.0.1:29911/webui/bootstrap")
|
||||||
|
token = boot.json()["token"]
|
||||||
|
auth = {"Authorization": f"Bearer {token}"}
|
||||||
|
|
||||||
|
initial = await _http_get(
|
||||||
|
"http://127.0.0.1:29911/api/webui/sidebar-state",
|
||||||
|
headers=auth,
|
||||||
|
)
|
||||||
|
assert initial.status_code == 200
|
||||||
|
assert initial.json()["schema_version"] == 1
|
||||||
|
assert initial.json()["pinned_keys"] == []
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"pinned_keys": ["websocket:sidebar"],
|
||||||
|
"archived_keys": ["websocket:old"],
|
||||||
|
"title_overrides": {"websocket:sidebar": "Pinned work"},
|
||||||
|
"view": {"density": "compact", "show_archived": True},
|
||||||
|
}
|
||||||
|
query = urlencode({"state": json.dumps(payload)})
|
||||||
|
updated = await _http_get(
|
||||||
|
f"http://127.0.0.1:29911/api/webui/sidebar-state/update?{query}",
|
||||||
|
headers=auth,
|
||||||
|
)
|
||||||
|
assert updated.status_code == 200
|
||||||
|
body = updated.json()
|
||||||
|
assert body["pinned_keys"] == ["websocket:sidebar"]
|
||||||
|
assert body["title_overrides"] == {"websocket:sidebar": "Pinned work"}
|
||||||
|
assert body["view"]["density"] == "compact"
|
||||||
|
|
||||||
|
state_path = tmp_path / "webui" / "sidebar-state.json"
|
||||||
|
assert state_path.is_file()
|
||||||
|
assert json.loads(state_path.read_text(encoding="utf-8"))["pinned_keys"] == [
|
||||||
|
"websocket:sidebar"
|
||||||
|
]
|
||||||
|
finally:
|
||||||
|
await channel.stop()
|
||||||
|
await server_task
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_session_delete_removes_file(
|
async def test_session_delete_removes_file(
|
||||||
bus: MagicMock, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
bus: MagicMock, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||||
) -> None:
|
) -> None:
|
||||||
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
|
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
|
||||||
sm = _seed_session(tmp_path, key="websocket:doomed")
|
sm = _seed_session(tmp_path, key="websocket:doomed")
|
||||||
from nanobot.utils.webui_transcript import append_transcript_object
|
from nanobot.webui.transcript import append_transcript_object
|
||||||
|
|
||||||
append_transcript_object("websocket:doomed", {"event": "user", "chat_id": "doomed", "text": "x"})
|
append_transcript_object("websocket:doomed", {"event": "user", "chat_id": "doomed", "text": "x"})
|
||||||
channel = _ch(bus, session_manager=sm, port=29903)
|
channel = _ch(bus, session_manager=sm, port=29903)
|
||||||
|
|||||||
+213
-4
@@ -11,7 +11,7 @@ from typer.testing import CliRunner
|
|||||||
from nanobot.bus.events import OutboundMessage
|
from nanobot.bus.events import OutboundMessage
|
||||||
from nanobot.cli.commands import app
|
from nanobot.cli.commands import app
|
||||||
from nanobot.providers.factory import make_provider
|
from nanobot.providers.factory import make_provider
|
||||||
from nanobot.config.schema import Config
|
from nanobot.config.schema import Config, ModelPresetConfig
|
||||||
from nanobot.cron.types import CronJob, CronPayload
|
from nanobot.cron.types import CronJob, CronPayload
|
||||||
from nanobot.providers.factory import ProviderSnapshot
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
from nanobot.providers.openai_codex_provider import _strip_model_prefix
|
from nanobot.providers.openai_codex_provider import _strip_model_prefix
|
||||||
@@ -226,6 +226,16 @@ def test_config_dump_excludes_oauth_provider_blocks():
|
|||||||
|
|
||||||
assert "openaiCodex" not in providers
|
assert "openaiCodex" not in providers
|
||||||
assert "githubCopilot" not in providers
|
assert "githubCopilot" not in providers
|
||||||
|
assert "xaiOauth" not in providers
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_dump_includes_xai_oauth_when_hosted_search_is_disabled():
|
||||||
|
config = Config()
|
||||||
|
config.providers.xai_oauth.x_search.enable = False
|
||||||
|
|
||||||
|
providers = config.model_dump(by_alias=True)["providers"]
|
||||||
|
|
||||||
|
assert providers["xaiOauth"]["xSearch"]["enable"] is False
|
||||||
|
|
||||||
|
|
||||||
def test_provider_logout_openai_codex_removes_local_oauth_files(tmp_path, monkeypatch):
|
def test_provider_logout_openai_codex_removes_local_oauth_files(tmp_path, monkeypatch):
|
||||||
@@ -280,6 +290,175 @@ def test_provider_logout_github_copilot_succeeds_when_no_local_oauth_file(monkey
|
|||||||
assert "No local OAuth credentials found for GitHub Copilot" in result.stdout
|
assert "No local OAuth credentials found for GitHub Copilot" in result.stdout
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_logout_xai_oauth_removes_local_oauth_files(tmp_path, monkeypatch):
|
||||||
|
token_path = tmp_path / "auth" / "xai-oauth.json"
|
||||||
|
lock_path = token_path.with_suffix(".lock")
|
||||||
|
token_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
token_path.write_text("{}", encoding="utf-8")
|
||||||
|
lock_path.write_text("", encoding="utf-8")
|
||||||
|
monkeypatch.setenv("NANOBOT_HOME", str(tmp_path))
|
||||||
|
monkeypatch.setattr("nanobot.providers.xai_oauth_provider._keyring_delete", lambda: None)
|
||||||
|
|
||||||
|
result = runner.invoke(app, ["provider", "logout", "xai-oauth"])
|
||||||
|
|
||||||
|
assert result.exit_code == 0
|
||||||
|
assert not token_path.exists()
|
||||||
|
assert not lock_path.exists()
|
||||||
|
assert "Logged out from xAI Grok OAuth" in result.stdout
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_logout_xai_oauth_succeeds_when_no_local_oauth_file(monkeypatch, tmp_path):
|
||||||
|
monkeypatch.setenv("NANOBOT_HOME", str(tmp_path))
|
||||||
|
monkeypatch.setattr("nanobot.providers.xai_oauth_provider._keyring_delete", lambda: None)
|
||||||
|
|
||||||
|
result = runner.invoke(app, ["provider", "logout", "xai-oauth"])
|
||||||
|
|
||||||
|
assert result.exit_code == 0
|
||||||
|
assert "No local OAuth credentials found for xAI Grok OAuth" in result.stdout
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_login_xai_oauth_forwards_manual_options(monkeypatch):
|
||||||
|
from nanobot.providers.xai_oauth_provider import XaiOAuthCredential
|
||||||
|
|
||||||
|
captured: dict[str, object] = {}
|
||||||
|
|
||||||
|
def fake_login_xai_oauth_interactive(**kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
return XaiOAuthCredential(access_token="access", account_id="acct", storage="keyring")
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.providers.xai_oauth_provider.login_xai_oauth_interactive",
|
||||||
|
fake_login_xai_oauth_interactive,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = runner.invoke(app, ["provider", "login", "xai-oauth", "--no-browser", "--manual-paste"])
|
||||||
|
|
||||||
|
assert result.exit_code == 0
|
||||||
|
assert captured["open_browser"] is False
|
||||||
|
assert captured["manual_paste"] is True
|
||||||
|
assert "Authenticated with xAI Grok OAuth" in result.stdout
|
||||||
|
assert "nanobot config set agents.defaults.provider xai-oauth" in result.stdout
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_set_updates_default_model_selection(tmp_path):
|
||||||
|
config_path = tmp_path / "config.json"
|
||||||
|
|
||||||
|
result = runner.invoke(app, [
|
||||||
|
"config",
|
||||||
|
"set",
|
||||||
|
"--config",
|
||||||
|
str(config_path),
|
||||||
|
"agents.defaults.model_preset",
|
||||||
|
"null",
|
||||||
|
])
|
||||||
|
assert result.exit_code == 0
|
||||||
|
|
||||||
|
result = runner.invoke(app, [
|
||||||
|
"config",
|
||||||
|
"set",
|
||||||
|
"--config",
|
||||||
|
str(config_path),
|
||||||
|
"agents.defaults.provider",
|
||||||
|
"xai-oauth",
|
||||||
|
])
|
||||||
|
assert result.exit_code == 0
|
||||||
|
|
||||||
|
result = runner.invoke(app, [
|
||||||
|
"config",
|
||||||
|
"set",
|
||||||
|
"--config",
|
||||||
|
str(config_path),
|
||||||
|
"agents.defaults.model",
|
||||||
|
"xai-oauth/grok-4.3",
|
||||||
|
])
|
||||||
|
assert result.exit_code == 0
|
||||||
|
|
||||||
|
data = json.loads(config_path.read_text(encoding="utf-8"))
|
||||||
|
config = Config.model_validate(data)
|
||||||
|
assert config.agents.defaults.model_preset is None
|
||||||
|
assert config.agents.defaults.provider == "xai-oauth"
|
||||||
|
assert config.agents.defaults.model == "xai-oauth/grok-4.3"
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_set_warns_when_model_preset_would_override_selection(tmp_path):
|
||||||
|
config = Config()
|
||||||
|
config.agents.defaults.model_preset = "fast"
|
||||||
|
config.model_presets["fast"] = ModelPresetConfig(
|
||||||
|
provider="openrouter",
|
||||||
|
model="openrouter/openai/gpt-4o-mini",
|
||||||
|
)
|
||||||
|
config_path = tmp_path / "config.json"
|
||||||
|
config_path.write_text(json.dumps(config.model_dump(mode="json", by_alias=True)), encoding="utf-8")
|
||||||
|
|
||||||
|
result = runner.invoke(app, [
|
||||||
|
"config",
|
||||||
|
"set",
|
||||||
|
"--config",
|
||||||
|
str(config_path),
|
||||||
|
"agents.defaults.provider",
|
||||||
|
"xai-oauth",
|
||||||
|
])
|
||||||
|
|
||||||
|
assert result.exit_code == 0
|
||||||
|
assert "model_preset is set and may override this" in result.stdout
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_set_disables_xai_oauth_hosted_search(tmp_path):
|
||||||
|
config_path = tmp_path / "config.json"
|
||||||
|
|
||||||
|
result = runner.invoke(app, [
|
||||||
|
"config",
|
||||||
|
"set",
|
||||||
|
"--config",
|
||||||
|
str(config_path),
|
||||||
|
"providers.xai_oauth.x_search.enable",
|
||||||
|
"false",
|
||||||
|
])
|
||||||
|
|
||||||
|
assert result.exit_code == 0
|
||||||
|
data = json.loads(config_path.read_text(encoding="utf-8"))
|
||||||
|
assert data["providers"]["xaiOauth"]["xSearch"]["enable"] is False
|
||||||
|
assert Config.model_validate(data).providers.xai_oauth.x_search.enable is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_set_rejects_unknown_path(tmp_path):
|
||||||
|
result = runner.invoke(app, [
|
||||||
|
"config",
|
||||||
|
"set",
|
||||||
|
"--config",
|
||||||
|
str(tmp_path / "config.json"),
|
||||||
|
"agents.defaults.not_a_field",
|
||||||
|
"value",
|
||||||
|
])
|
||||||
|
|
||||||
|
assert result.exit_code == 1
|
||||||
|
assert "Could not set config value" in result.stdout
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_login_xai_oauth_does_not_update_config(monkeypatch, tmp_path):
|
||||||
|
from nanobot.providers.xai_oauth_provider import XaiOAuthCredential
|
||||||
|
|
||||||
|
config = Config()
|
||||||
|
config.agents.defaults.provider = "auto"
|
||||||
|
config.agents.defaults.model = "anthropic/claude-opus-4-5"
|
||||||
|
config_path = tmp_path / "config.json"
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.providers.xai_oauth_provider.login_xai_oauth_interactive",
|
||||||
|
lambda **_kwargs: XaiOAuthCredential(access_token="access", account_id="acct", storage="keyring"),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("nanobot.config.loader.get_config_path", lambda: config_path)
|
||||||
|
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
||||||
|
save_config = MagicMock()
|
||||||
|
monkeypatch.setattr("nanobot.config.loader.save_config", save_config)
|
||||||
|
|
||||||
|
result = runner.invoke(app, ["provider", "login", "xai-oauth"])
|
||||||
|
|
||||||
|
assert result.exit_code == 0
|
||||||
|
save_config.assert_not_called()
|
||||||
|
assert "nanobot config set agents.defaults.model xai-oauth/grok-4.3" in result.stdout
|
||||||
|
|
||||||
|
|
||||||
def test_provider_logout_rejects_unknown_provider():
|
def test_provider_logout_rejects_unknown_provider():
|
||||||
result = runner.invoke(app, ["provider", "logout", "not-a-real-provider"])
|
result = runner.invoke(app, ["provider", "logout", "not-a-real-provider"])
|
||||||
|
|
||||||
@@ -398,6 +577,8 @@ def test_find_by_name_accepts_camel_case_and_hyphen_aliases():
|
|||||||
assert find_by_name("volcengineCodingPlan").name == "volcengine_coding_plan"
|
assert find_by_name("volcengineCodingPlan").name == "volcengine_coding_plan"
|
||||||
assert find_by_name("github-copilot") is not None
|
assert find_by_name("github-copilot") is not None
|
||||||
assert find_by_name("github-copilot").name == "github_copilot"
|
assert find_by_name("github-copilot").name == "github_copilot"
|
||||||
|
assert find_by_name("xai-oauth") is not None
|
||||||
|
assert find_by_name("xai-oauth").name == "xai_oauth"
|
||||||
assert find_by_name("longcat") is not None
|
assert find_by_name("longcat") is not None
|
||||||
assert find_by_name("longcat").name == "longcat"
|
assert find_by_name("longcat").name == "longcat"
|
||||||
assert find_by_name("atomic-chat") is not None
|
assert find_by_name("atomic-chat") is not None
|
||||||
@@ -540,6 +721,23 @@ def test_make_provider_uses_github_copilot_backend():
|
|||||||
assert provider.__class__.__name__ == "GitHubCopilotProvider"
|
assert provider.__class__.__name__ == "GitHubCopilotProvider"
|
||||||
|
|
||||||
|
|
||||||
|
def test_make_provider_uses_xai_oauth_backend():
|
||||||
|
config = Config.model_validate(
|
||||||
|
{
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"provider": "xai-oauth",
|
||||||
|
"model": "xai-oauth/grok-4.3",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
provider = make_provider(config)
|
||||||
|
|
||||||
|
assert provider.__class__.__name__ == "XaiOAuthProvider"
|
||||||
|
|
||||||
|
|
||||||
def test_github_copilot_provider_strips_prefixed_model_name():
|
def test_github_copilot_provider_strips_prefixed_model_name():
|
||||||
from nanobot.providers.github_copilot_provider import GitHubCopilotProvider
|
from nanobot.providers.github_copilot_provider import GitHubCopilotProvider
|
||||||
|
|
||||||
@@ -572,6 +770,7 @@ async def test_github_copilot_provider_refreshes_client_api_key_before_chat():
|
|||||||
|
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI", return_value=mock_client):
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI", return_value=mock_client):
|
||||||
provider = GitHubCopilotProvider(default_model="github-copilot/gpt-4")
|
provider = GitHubCopilotProvider(default_model="github-copilot/gpt-4")
|
||||||
|
await provider._ensure_client()
|
||||||
|
|
||||||
provider._get_copilot_access_token = AsyncMock(return_value="copilot-access-token")
|
provider._get_copilot_access_token = AsyncMock(return_value="copilot-access-token")
|
||||||
|
|
||||||
@@ -611,7 +810,8 @@ def test_make_provider_passes_extra_headers_to_custom_provider():
|
|||||||
)
|
)
|
||||||
|
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_async_openai:
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_async_openai:
|
||||||
make_provider(config)
|
provider = make_provider(config)
|
||||||
|
asyncio.run(provider._ensure_client())
|
||||||
|
|
||||||
kwargs = mock_async_openai.call_args.kwargs
|
kwargs = mock_async_openai.call_args.kwargs
|
||||||
assert kwargs["api_key"] == "test-key"
|
assert kwargs["api_key"] == "test-key"
|
||||||
@@ -1170,6 +1370,7 @@ def test_gateway_cron_evaluator_receives_scheduled_reminder_context(
|
|||||||
self.model = "test-model"
|
self.model = "test-model"
|
||||||
self.provider = kwargs.get("provider", object())
|
self.provider = kwargs.get("provider", object())
|
||||||
self.tools = {}
|
self.tools = {}
|
||||||
|
seen["agent"] = self
|
||||||
|
|
||||||
async def process_direct(self, *_args, **_kwargs):
|
async def process_direct(self, *_args, **_kwargs):
|
||||||
return OutboundMessage(
|
return OutboundMessage(
|
||||||
@@ -1218,6 +1419,11 @@ def test_gateway_cron_evaluator_receives_scheduled_reminder_context(
|
|||||||
assert isinstance(cron, _FakeCron)
|
assert isinstance(cron, _FakeCron)
|
||||||
assert cron.on_job is not None
|
assert cron.on_job is not None
|
||||||
|
|
||||||
|
runtime_provider = object()
|
||||||
|
agent = seen["agent"]
|
||||||
|
agent.provider = runtime_provider
|
||||||
|
agent.model = "runtime-model"
|
||||||
|
|
||||||
job = CronJob(
|
job = CronJob(
|
||||||
id="cron-1",
|
id="cron-1",
|
||||||
name="stretch",
|
name="stretch",
|
||||||
@@ -1233,8 +1439,8 @@ def test_gateway_cron_evaluator_receives_scheduled_reminder_context(
|
|||||||
|
|
||||||
assert response == "Time to stretch."
|
assert response == "Time to stretch."
|
||||||
assert seen["response"] == "Time to stretch."
|
assert seen["response"] == "Time to stretch."
|
||||||
assert seen["provider"] is provider
|
assert seen["provider"] is runtime_provider
|
||||||
assert seen["model"] == "test-model"
|
assert seen["model"] == "runtime-model"
|
||||||
assert seen["task_context"] == (
|
assert seen["task_context"] == (
|
||||||
"The scheduled time has arrived. Deliver this reminder to the user now, "
|
"The scheduled time has arrived. Deliver this reminder to the user now, "
|
||||||
"as a brief and natural message in their language. Speak directly to them — "
|
"as a brief and natural message in their language. Speak directly to them — "
|
||||||
@@ -1543,6 +1749,9 @@ def test_gateway_health_endpoint_binds_and_serves_expected_responses(
|
|||||||
self.dream = _FakeDream()
|
self.dream = _FakeDream()
|
||||||
self.sessions = _FakeSessionManager()
|
self.sessions = _FakeSessionManager()
|
||||||
|
|
||||||
|
def llm_runtime(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
async def run(self) -> None:
|
async def run(self) -> None:
|
||||||
await asyncio.Event().wait()
|
await asyncio.Event().wait()
|
||||||
|
|
||||||
|
|||||||
@@ -69,6 +69,72 @@ async def test_reasoning_delta_displayed_when_show_reasoning_enabled():
|
|||||||
assert calls == ["I should search first."]
|
assert calls == ["I should search first."]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reasoning_delta_buffers_until_sentence_boundary():
|
||||||
|
calls: list[str] = []
|
||||||
|
channels_config = SimpleNamespace(
|
||||||
|
send_progress=True, send_tool_hints=False, show_reasoning=True,
|
||||||
|
)
|
||||||
|
reasoning_buffer = commands._ReasoningBuffer()
|
||||||
|
|
||||||
|
with patch("nanobot.cli.commands._print_cli_reasoning", side_effect=lambda t, th, r=None: calls.append(t)):
|
||||||
|
first = await commands._maybe_print_interactive_progress(
|
||||||
|
SimpleNamespace(
|
||||||
|
content="The",
|
||||||
|
metadata={"_progress": True, "_reasoning_delta": True},
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
channels_config,
|
||||||
|
reasoning_buffer=reasoning_buffer,
|
||||||
|
)
|
||||||
|
second = await commands._maybe_print_interactive_progress(
|
||||||
|
SimpleNamespace(
|
||||||
|
content=" user asked.",
|
||||||
|
metadata={"_progress": True, "_reasoning_delta": True},
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
channels_config,
|
||||||
|
reasoning_buffer=reasoning_buffer,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert first is True
|
||||||
|
assert second is True
|
||||||
|
assert calls == ["The user asked."]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reasoning_end_flushes_buffered_delta():
|
||||||
|
calls: list[str] = []
|
||||||
|
channels_config = SimpleNamespace(
|
||||||
|
send_progress=True, send_tool_hints=False, show_reasoning=True,
|
||||||
|
)
|
||||||
|
reasoning_buffer = commands._ReasoningBuffer()
|
||||||
|
|
||||||
|
with patch("nanobot.cli.commands._print_cli_reasoning", side_effect=lambda t, th, r=None: calls.append(t)):
|
||||||
|
delta = await commands._maybe_print_interactive_progress(
|
||||||
|
SimpleNamespace(
|
||||||
|
content="The user asked",
|
||||||
|
metadata={"_progress": True, "_reasoning_delta": True},
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
channels_config,
|
||||||
|
reasoning_buffer=reasoning_buffer,
|
||||||
|
)
|
||||||
|
end = await commands._maybe_print_interactive_progress(
|
||||||
|
SimpleNamespace(
|
||||||
|
content="",
|
||||||
|
metadata={"_progress": True, "_reasoning_end": True},
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
channels_config,
|
||||||
|
reasoning_buffer=reasoning_buffer,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert delta is True
|
||||||
|
assert end is True
|
||||||
|
assert calls == ["The user asked"]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_reasoning_hidden_when_show_reasoning_disabled():
|
async def test_reasoning_hidden_when_show_reasoning_disabled():
|
||||||
"""Reasoning content should be suppressed when show_reasoning is False."""
|
"""Reasoning content should be suppressed when show_reasoning is False."""
|
||||||
|
|||||||
@@ -0,0 +1,73 @@
|
|||||||
|
"""Tests for the Ant Ling provider registration."""
|
||||||
|
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from nanobot.config.schema import Config, ProvidersConfig
|
||||||
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
|
from nanobot.providers.registry import PROVIDERS, find_by_name
|
||||||
|
|
||||||
|
|
||||||
|
def test_ant_ling_config_field_exists() -> None:
|
||||||
|
config = ProvidersConfig()
|
||||||
|
|
||||||
|
assert hasattr(config, "ant_ling")
|
||||||
|
|
||||||
|
|
||||||
|
def test_ant_ling_provider_in_registry() -> None:
|
||||||
|
specs = {spec.name: spec for spec in PROVIDERS}
|
||||||
|
|
||||||
|
assert "ant_ling" in specs
|
||||||
|
ant_ling = specs["ant_ling"]
|
||||||
|
assert ant_ling.backend == "openai_compat"
|
||||||
|
assert ant_ling.env_key == "ANT_LING_API_KEY"
|
||||||
|
assert ant_ling.display_name == "Ant Ling"
|
||||||
|
assert ant_ling.default_api_base == "https://api.ant-ling.com/v1"
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_by_name_accepts_ant_ling_spellings() -> None:
|
||||||
|
spec = find_by_name("ant_ling")
|
||||||
|
|
||||||
|
assert spec is not None
|
||||||
|
assert find_by_name("ant-ling") is spec
|
||||||
|
assert find_by_name("antLing") is spec
|
||||||
|
|
||||||
|
|
||||||
|
def test_ant_ling_model_auto_matches_with_default_api_base() -> None:
|
||||||
|
config = Config.model_validate({
|
||||||
|
"providers": {
|
||||||
|
"antLing": {
|
||||||
|
"apiKey": "ling-key",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "Ling-2.6-flash",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
assert config.get_provider_name("Ling-2.6-flash") == "ant_ling"
|
||||||
|
assert config.get_api_key("Ling-2.6-flash") == "ling-key"
|
||||||
|
assert config.get_api_base("Ling-2.6-flash") == "https://api.ant-ling.com/v1"
|
||||||
|
|
||||||
|
|
||||||
|
def test_ant_ling_preserves_official_model_name() -> None:
|
||||||
|
spec = find_by_name("ant_ling")
|
||||||
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
||||||
|
provider = OpenAICompatProvider(
|
||||||
|
api_key="ling-key",
|
||||||
|
default_model="Ling-2.6-flash",
|
||||||
|
spec=spec,
|
||||||
|
)
|
||||||
|
|
||||||
|
kwargs = provider._build_kwargs(
|
||||||
|
messages=[{"role": "user", "content": "hi"}],
|
||||||
|
tools=None,
|
||||||
|
model="Ling-2.6-flash",
|
||||||
|
max_tokens=1024,
|
||||||
|
temperature=0.7,
|
||||||
|
reasoning_effort=None,
|
||||||
|
tool_choice=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert kwargs["model"] == "Ling-2.6-flash"
|
||||||
@@ -129,6 +129,74 @@ async def test_chat_stream_invokes_on_thinking_delta_for_thinking_delta() -> Non
|
|||||||
assert text_parts == ["X"]
|
assert text_parts == ["X"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_chat_stream_invokes_tool_call_delta_for_input_json_delta() -> None:
|
||||||
|
provider = AnthropicProvider(api_key="sk-test")
|
||||||
|
provider._client = MagicMock()
|
||||||
|
|
||||||
|
chunks = [
|
||||||
|
SimpleNamespace(
|
||||||
|
type="content_block_start",
|
||||||
|
index=1,
|
||||||
|
content_block=SimpleNamespace(
|
||||||
|
type="tool_use",
|
||||||
|
id="toolu_1",
|
||||||
|
name="write_file",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
SimpleNamespace(
|
||||||
|
type="content_block_delta",
|
||||||
|
index=1,
|
||||||
|
delta=SimpleNamespace(
|
||||||
|
type="input_json_delta",
|
||||||
|
partial_json='{"path":"notes.md","content":"',
|
||||||
|
),
|
||||||
|
),
|
||||||
|
SimpleNamespace(
|
||||||
|
type="content_block_delta",
|
||||||
|
index=1,
|
||||||
|
delta=SimpleNamespace(type="input_json_delta", partial_json="line\\n"),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
fake = _FakeAsyncStream(chunks)
|
||||||
|
stream_cm = MagicMock()
|
||||||
|
stream_cm.__aenter__ = AsyncMock(return_value=fake)
|
||||||
|
stream_cm.__aexit__ = AsyncMock(return_value=None)
|
||||||
|
provider._client.messages.stream = MagicMock(return_value=stream_cm)
|
||||||
|
|
||||||
|
deltas: list[dict] = []
|
||||||
|
|
||||||
|
async def on_tool_delta(delta: dict) -> None:
|
||||||
|
deltas.append(delta)
|
||||||
|
|
||||||
|
await provider.chat_stream(
|
||||||
|
messages=[{"role": "user", "content": "write"}],
|
||||||
|
on_tool_call_delta=on_tool_delta,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert deltas == [
|
||||||
|
{
|
||||||
|
"index": 1,
|
||||||
|
"call_id": "toolu_1",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"index": 1,
|
||||||
|
"call_id": "toolu_1",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": '{"path":"notes.md","content":"',
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"index": 1,
|
||||||
|
"call_id": "toolu_1",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": "line\\n",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
fake.get_final_message.assert_awaited_once()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_chat_stream_without_callback_still_finalizes() -> None:
|
async def test_chat_stream_without_callback_still_finalizes() -> None:
|
||||||
provider = AnthropicProvider(api_key="sk-test")
|
provider = AnthropicProvider(api_key="sk-test")
|
||||||
|
|||||||
@@ -65,6 +65,7 @@ async def test_github_copilot_does_not_fall_back_from_responses_error():
|
|||||||
|
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI", return_value=mock_client):
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI", return_value=mock_client):
|
||||||
provider = GitHubCopilotProvider(default_model="github_copilot/gpt-5.4-mini")
|
provider = GitHubCopilotProvider(default_model="github_copilot/gpt-5.4-mini")
|
||||||
|
await provider._ensure_client()
|
||||||
provider._get_copilot_access_token = AsyncMock(return_value="copilot-access-token")
|
provider._get_copilot_access_token = AsyncMock(return_value="copilot-access-token")
|
||||||
|
|
||||||
response = await provider.chat(
|
response = await provider.chat(
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import base64
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -8,9 +9,12 @@ import pytest
|
|||||||
|
|
||||||
from nanobot.providers.image_generation import (
|
from nanobot.providers.image_generation import (
|
||||||
AIHubMixImageGenerationClient,
|
AIHubMixImageGenerationClient,
|
||||||
|
GeminiImageGenerationClient,
|
||||||
GeneratedImageResponse,
|
GeneratedImageResponse,
|
||||||
ImageGenerationError,
|
ImageGenerationError,
|
||||||
|
MiniMaxImageGenerationClient,
|
||||||
OpenRouterImageGenerationClient,
|
OpenRouterImageGenerationClient,
|
||||||
|
StepFunImageGenerationClient,
|
||||||
)
|
)
|
||||||
|
|
||||||
PNG_BYTES = (
|
PNG_BYTES = (
|
||||||
@@ -23,6 +27,7 @@ PNG_DATA_URL = (
|
|||||||
"data:image/png;base64,"
|
"data:image/png;base64,"
|
||||||
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/x8AAwMCAO+/p9sAAAAASUVORK5CYII="
|
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/x8AAwMCAO+/p9sAAAAASUVORK5CYII="
|
||||||
)
|
)
|
||||||
|
JPEG_BYTES = b"\xff\xd8\xff\xe0" + b"0" * 12
|
||||||
|
|
||||||
|
|
||||||
class FakeResponse:
|
class FakeResponse:
|
||||||
@@ -202,3 +207,311 @@ async def test_aihubmix_image_generation_downloads_url_response() -> None:
|
|||||||
|
|
||||||
assert response.images[0].startswith("data:image/png;base64,")
|
assert response.images[0].startswith("data:image/png;base64,")
|
||||||
assert fake.get_calls[0]["url"] == "https://cdn.example/image.png"
|
assert fake.get_calls[0]["url"] == "https://cdn.example/image.png"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_aihubmix_base64_response_uses_detected_mime() -> None:
|
||||||
|
raw_b64 = base64.b64encode(JPEG_BYTES).decode("ascii")
|
||||||
|
fake = FakeClient(FakeResponse({"output": {"b64_json": raw_b64}}))
|
||||||
|
client = AIHubMixImageGenerationClient(
|
||||||
|
api_key="sk-ahm-test",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await client.generate(prompt="draw", model="gpt-image-2-free")
|
||||||
|
|
||||||
|
assert response.images == [f"data:image/jpeg;base64,{raw_b64}"]
|
||||||
|
|
||||||
|
|
||||||
|
RAW_B64 = PNG_DATA_URL.removeprefix("data:image/png;base64,")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_imagen_payload_and_response() -> None:
|
||||||
|
fake = FakeClient(
|
||||||
|
FakeResponse({"predictions": [{"bytesBase64Encoded": RAW_B64, "mimeType": "image/png"}]})
|
||||||
|
)
|
||||||
|
client = GeminiImageGenerationClient(
|
||||||
|
api_key="AIza-test",
|
||||||
|
api_base="https://generativelanguage.googleapis.com/v1beta",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await client.generate(
|
||||||
|
prompt="a sunset",
|
||||||
|
model="imagen-4.0-generate-001",
|
||||||
|
aspect_ratio="16:9",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.images == [PNG_DATA_URL]
|
||||||
|
assert response.content == ""
|
||||||
|
call = fake.calls[0]
|
||||||
|
assert call["url"].endswith(":predict")
|
||||||
|
assert call["headers"]["x-goog-api-key"] == "AIza-test"
|
||||||
|
assert "params" not in call
|
||||||
|
body = call["json"]
|
||||||
|
assert body["instances"] == [{"prompt": "a sunset"}]
|
||||||
|
assert body["parameters"]["sampleCount"] == 1
|
||||||
|
assert body["parameters"]["aspectRatio"] == "16:9"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_imagen_ignores_unsupported_aspect_ratio() -> None:
|
||||||
|
fake = FakeClient(
|
||||||
|
FakeResponse({"predictions": [{"bytesBase64Encoded": RAW_B64, "mimeType": "image/png"}]})
|
||||||
|
)
|
||||||
|
client = GeminiImageGenerationClient(api_key="AIza-test", client=fake) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
await client.generate(prompt="a sunset", model="imagen-4.0-generate-001", aspect_ratio="2:3")
|
||||||
|
|
||||||
|
body = fake.calls[0]["json"]
|
||||||
|
assert "aspectRatio" not in body["parameters"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_flash_payload_and_response() -> None:
|
||||||
|
fake = FakeClient(
|
||||||
|
FakeResponse(
|
||||||
|
{
|
||||||
|
"candidates": [
|
||||||
|
{
|
||||||
|
"content": {
|
||||||
|
"parts": [
|
||||||
|
{"text": "here is your image"},
|
||||||
|
{"inlineData": {"mimeType": "image/png", "data": RAW_B64}},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
client = GeminiImageGenerationClient(
|
||||||
|
api_key="AIza-test",
|
||||||
|
api_base="https://generativelanguage.googleapis.com/v1beta",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await client.generate(
|
||||||
|
prompt="draw a cat",
|
||||||
|
model="gemini-2.0-flash-preview-image-generation",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.images == [PNG_DATA_URL]
|
||||||
|
assert response.content == "here is your image"
|
||||||
|
call = fake.calls[0]
|
||||||
|
assert call["url"].endswith(":generateContent")
|
||||||
|
assert call["headers"]["x-goog-api-key"] == "AIza-test"
|
||||||
|
assert "params" not in call
|
||||||
|
body = call["json"]
|
||||||
|
assert body["generationConfig"]["responseModalities"] == ["TEXT", "IMAGE"]
|
||||||
|
assert body["contents"][0]["parts"][-1] == {"text": "draw a cat"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_flash_reference_images(tmp_path: Path) -> None:
|
||||||
|
ref = tmp_path / "ref.png"
|
||||||
|
ref.write_bytes(PNG_BYTES)
|
||||||
|
fake = FakeClient(
|
||||||
|
FakeResponse(
|
||||||
|
{
|
||||||
|
"candidates": [
|
||||||
|
{
|
||||||
|
"content": {
|
||||||
|
"parts": [{"inlineData": {"mimeType": "image/png", "data": RAW_B64}}]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
client = GeminiImageGenerationClient(api_key="AIza-test", client=fake) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
response = await client.generate(
|
||||||
|
prompt="edit this",
|
||||||
|
model="gemini-2.0-flash-preview-image-generation",
|
||||||
|
reference_images=[str(ref)],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.images == [PNG_DATA_URL]
|
||||||
|
parts = fake.calls[0]["json"]["contents"][0]["parts"]
|
||||||
|
assert parts[0]["inlineData"]["mimeType"] == "image/png"
|
||||||
|
assert parts[0]["inlineData"]["data"].startswith("iVBOR")
|
||||||
|
assert parts[1] == {"text": "edit this"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_requires_api_key() -> None:
|
||||||
|
client = GeminiImageGenerationClient(api_key=None)
|
||||||
|
|
||||||
|
with pytest.raises(ImageGenerationError, match="API key"):
|
||||||
|
await client.generate(prompt="draw", model="imagen-4.0-generate-001")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_no_images_raises() -> None:
|
||||||
|
fake = FakeClient(FakeResponse({"candidates": [{"content": {"parts": [{"text": "sorry"}]}}]}))
|
||||||
|
client = GeminiImageGenerationClient(api_key="AIza-test", client=fake) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
with pytest.raises(ImageGenerationError, match="returned no images"):
|
||||||
|
await client.generate(prompt="draw", model="gemini-2.0-flash-preview-image-generation")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_minimax_payload_and_response_with_reference_image(tmp_path: Path) -> None:
|
||||||
|
ref = tmp_path / "ref.png"
|
||||||
|
ref.write_bytes(PNG_BYTES)
|
||||||
|
fake = FakeClient(FakeResponse({"data": {"image_base64": [RAW_B64]}}))
|
||||||
|
client = MiniMaxImageGenerationClient(
|
||||||
|
api_key="sk-mm-test",
|
||||||
|
api_base="https://api.minimaxi.com/v1/",
|
||||||
|
extra_headers={"X-Test": "1"},
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await client.generate(
|
||||||
|
prompt="draw a character",
|
||||||
|
model="image-01",
|
||||||
|
reference_images=[str(ref)],
|
||||||
|
aspect_ratio="21:9",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.images == [PNG_DATA_URL]
|
||||||
|
call = fake.calls[0]
|
||||||
|
assert call["url"] == "https://api.minimaxi.com/v1/image_generation"
|
||||||
|
assert call["headers"]["Authorization"] == "Bearer sk-mm-test"
|
||||||
|
assert call["headers"]["X-Test"] == "1"
|
||||||
|
body = call["json"]
|
||||||
|
assert body["model"] == "image-01"
|
||||||
|
assert body["prompt"] == "draw a character"
|
||||||
|
assert body["response_format"] == "base64"
|
||||||
|
assert body["aspect_ratio"] == "21:9"
|
||||||
|
assert body["subject_reference"][0]["type"] == "character"
|
||||||
|
assert body["subject_reference"][0]["image_file"].startswith("data:image/png;base64,")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# StepFun (阶跃星辰)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stepfun_payload_and_response_with_aspect_ratio() -> None:
|
||||||
|
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
||||||
|
client = StepFunImageGenerationClient(
|
||||||
|
api_key="sk-sf-test",
|
||||||
|
api_base="https://api.stepfun.com/v1",
|
||||||
|
extra_headers={"X-Test": "1"},
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await client.generate(
|
||||||
|
prompt="a cat on the moon",
|
||||||
|
model="step-image-edit-2",
|
||||||
|
aspect_ratio="16:9",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.images == [PNG_DATA_URL]
|
||||||
|
call = fake.calls[0]
|
||||||
|
assert call["url"] == "https://api.stepfun.com/v1/images/generations"
|
||||||
|
assert call["headers"]["Authorization"] == "Bearer sk-sf-test"
|
||||||
|
assert call["headers"]["X-Test"] == "1"
|
||||||
|
body = call["json"]
|
||||||
|
assert body["model"] == "step-image-edit-2"
|
||||||
|
assert body["prompt"] == "a cat on the moon"
|
||||||
|
assert body["response_format"] == "b64_json"
|
||||||
|
assert body["n"] == 1
|
||||||
|
assert body["size"] == "1280x800"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stepfun_default_size_when_no_aspect_ratio() -> None:
|
||||||
|
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
||||||
|
client = StepFunImageGenerationClient(
|
||||||
|
api_key="sk-sf-test",
|
||||||
|
api_base="https://api.stepfun.com/v1",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
await client.generate(prompt="a dog", model="step-image-edit-2")
|
||||||
|
|
||||||
|
body = fake.calls[0]["json"]
|
||||||
|
assert body["size"] == "1024x1024"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stepfun_uses_explicit_image_size() -> None:
|
||||||
|
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
||||||
|
client = StepFunImageGenerationClient(
|
||||||
|
api_key="sk-sf-test",
|
||||||
|
api_base="https://api.stepfun.com/v1",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
await client.generate(
|
||||||
|
prompt="a bird",
|
||||||
|
model="step-image-edit-2",
|
||||||
|
image_size="1024x1024",
|
||||||
|
)
|
||||||
|
|
||||||
|
body = fake.calls[0]["json"]
|
||||||
|
assert body["size"] == "1024x1024"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stepfun_style_reference_on_1x_model(tmp_path: Path) -> None:
|
||||||
|
"""step-1x-medium supports style_reference for reference-image generation."""
|
||||||
|
ref = tmp_path / "ref.png"
|
||||||
|
ref.write_bytes(PNG_BYTES)
|
||||||
|
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
||||||
|
client = StepFunImageGenerationClient(
|
||||||
|
api_key="sk-sf-test",
|
||||||
|
api_base="https://api.stepfun.com/v1",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
await client.generate(
|
||||||
|
prompt="in this style",
|
||||||
|
model="step-1x-medium",
|
||||||
|
reference_images=[str(ref)],
|
||||||
|
)
|
||||||
|
|
||||||
|
body = fake.calls[0]["json"]
|
||||||
|
assert "style_reference" in body
|
||||||
|
assert body["style_reference"]["source_url"].startswith("data:image/png;base64,")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stepfun_no_style_reference_on_non_1x_model() -> None:
|
||||||
|
"""step-image-edit-2 does not use style_reference; reference images are ignored."""
|
||||||
|
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
||||||
|
client = StepFunImageGenerationClient(
|
||||||
|
api_key="sk-sf-test",
|
||||||
|
api_base="https://api.stepfun.com/v1",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
await client.generate(
|
||||||
|
prompt="a flower",
|
||||||
|
model="step-image-edit-2",
|
||||||
|
reference_images=["/tmp/ref.png"],
|
||||||
|
)
|
||||||
|
|
||||||
|
body = fake.calls[0]["json"]
|
||||||
|
assert "style_reference" not in body
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stepfun_requires_api_key() -> None:
|
||||||
|
client = StepFunImageGenerationClient(api_key=None)
|
||||||
|
|
||||||
|
with pytest.raises(ImageGenerationError, match="API key"):
|
||||||
|
await client.generate(prompt="draw", model="step-image-edit-2")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stepfun_no_images_raises() -> None:
|
||||||
|
fake = FakeClient(FakeResponse({"data": [{"text": "sorry"}]}))
|
||||||
|
client = StepFunImageGenerationClient(api_key="sk-sf-test", client=fake) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
with pytest.raises(ImageGenerationError, match="returned no images"):
|
||||||
|
await client.generate(prompt="draw", model="step-image-edit-2")
|
||||||
|
|||||||
@@ -164,6 +164,130 @@ def _fake_chat_stream_reasoning_chunks():
|
|||||||
return _stream()
|
return _stream()
|
||||||
|
|
||||||
|
|
||||||
|
def _fake_chat_stream_tool_call_chunks():
|
||||||
|
"""Mimic OpenAI-compatible streaming tool-call argument deltas."""
|
||||||
|
|
||||||
|
async def _stream():
|
||||||
|
yield SimpleNamespace(
|
||||||
|
choices=[
|
||||||
|
SimpleNamespace(
|
||||||
|
finish_reason=None,
|
||||||
|
delta=SimpleNamespace(
|
||||||
|
content=None,
|
||||||
|
reasoning_content=None,
|
||||||
|
reasoning=None,
|
||||||
|
tool_calls=[
|
||||||
|
SimpleNamespace(
|
||||||
|
index=0,
|
||||||
|
id="call_write",
|
||||||
|
function=SimpleNamespace(
|
||||||
|
name="write_file",
|
||||||
|
arguments='{"path":"notes.md","content":"',
|
||||||
|
),
|
||||||
|
)
|
||||||
|
],
|
||||||
|
),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
usage=None,
|
||||||
|
)
|
||||||
|
yield SimpleNamespace(
|
||||||
|
choices=[
|
||||||
|
SimpleNamespace(
|
||||||
|
finish_reason=None,
|
||||||
|
delta=SimpleNamespace(
|
||||||
|
content=None,
|
||||||
|
reasoning_content=None,
|
||||||
|
reasoning=None,
|
||||||
|
tool_calls=[
|
||||||
|
SimpleNamespace(
|
||||||
|
index=0,
|
||||||
|
id=None,
|
||||||
|
function=SimpleNamespace(name=None, arguments='line\\n"}'),
|
||||||
|
)
|
||||||
|
],
|
||||||
|
),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
usage=None,
|
||||||
|
)
|
||||||
|
yield SimpleNamespace(
|
||||||
|
choices=[
|
||||||
|
SimpleNamespace(
|
||||||
|
finish_reason="tool_calls",
|
||||||
|
delta=SimpleNamespace(
|
||||||
|
content=None,
|
||||||
|
reasoning_content=None,
|
||||||
|
reasoning=None,
|
||||||
|
tool_calls=None,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
usage=SimpleNamespace(prompt_tokens=10, completion_tokens=5, total_tokens=15),
|
||||||
|
)
|
||||||
|
|
||||||
|
return _stream()
|
||||||
|
|
||||||
|
|
||||||
|
def _fake_chat_stream_legacy_function_call_chunks():
|
||||||
|
"""Mimic older OpenAI-compatible ``delta.function_call`` chunks."""
|
||||||
|
|
||||||
|
async def _stream():
|
||||||
|
yield SimpleNamespace(
|
||||||
|
choices=[
|
||||||
|
SimpleNamespace(
|
||||||
|
finish_reason=None,
|
||||||
|
delta=SimpleNamespace(
|
||||||
|
content=None,
|
||||||
|
reasoning_content=None,
|
||||||
|
reasoning=None,
|
||||||
|
tool_calls=None,
|
||||||
|
function_call=SimpleNamespace(
|
||||||
|
name="write_file",
|
||||||
|
arguments='{"path":"notes.md","content":"',
|
||||||
|
),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
usage=None,
|
||||||
|
)
|
||||||
|
yield SimpleNamespace(
|
||||||
|
choices=[
|
||||||
|
SimpleNamespace(
|
||||||
|
finish_reason=None,
|
||||||
|
delta=SimpleNamespace(
|
||||||
|
content=None,
|
||||||
|
reasoning_content=None,
|
||||||
|
reasoning=None,
|
||||||
|
tool_calls=None,
|
||||||
|
function_call=SimpleNamespace(
|
||||||
|
name=None,
|
||||||
|
arguments='line\\n"}',
|
||||||
|
),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
usage=None,
|
||||||
|
)
|
||||||
|
yield SimpleNamespace(
|
||||||
|
choices=[
|
||||||
|
SimpleNamespace(
|
||||||
|
finish_reason="function_call",
|
||||||
|
delta=SimpleNamespace(
|
||||||
|
content=None,
|
||||||
|
reasoning_content=None,
|
||||||
|
reasoning=None,
|
||||||
|
tool_calls=None,
|
||||||
|
function_call=None,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
usage=SimpleNamespace(prompt_tokens=10, completion_tokens=5, total_tokens=15),
|
||||||
|
)
|
||||||
|
|
||||||
|
return _stream()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_openai_compat_stream_forwards_reasoning_deltas_deepseek_style() -> None:
|
async def test_openai_compat_stream_forwards_reasoning_deltas_deepseek_style() -> None:
|
||||||
"""Regression: DeepSeek-V4 / reasoner expose ``delta.reasoning_content`` during streaming."""
|
"""Regression: DeepSeek-V4 / reasoner expose ``delta.reasoning_content`` during streaming."""
|
||||||
@@ -202,6 +326,98 @@ async def test_openai_compat_stream_forwards_reasoning_deltas_deepseek_style() -
|
|||||||
mock_chat.assert_awaited_once()
|
mock_chat.assert_awaited_once()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("provider_name", "model"),
|
||||||
|
[
|
||||||
|
("openai", "gpt-4o"),
|
||||||
|
("deepseek", "deepseek-chat"),
|
||||||
|
("minimax", "MiniMax-M2.7"),
|
||||||
|
("zhipu", "glm-4.6"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_openai_compat_stream_forwards_tool_call_argument_deltas(
|
||||||
|
provider_name: str,
|
||||||
|
model: str,
|
||||||
|
) -> None:
|
||||||
|
mock_chat = AsyncMock(return_value=_fake_chat_stream_tool_call_chunks())
|
||||||
|
spec = find_by_name(provider_name)
|
||||||
|
deltas: list[dict] = []
|
||||||
|
|
||||||
|
async def on_tool_delta(delta: dict) -> None:
|
||||||
|
deltas.append(delta)
|
||||||
|
|
||||||
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_openai:
|
||||||
|
client_instance = mock_openai.return_value
|
||||||
|
client_instance.chat.completions.create = mock_chat
|
||||||
|
|
||||||
|
provider = OpenAICompatProvider(
|
||||||
|
api_key="sk-test",
|
||||||
|
default_model=model,
|
||||||
|
spec=spec,
|
||||||
|
)
|
||||||
|
result = await provider.chat_stream(
|
||||||
|
messages=[{"role": "user", "content": "write"}],
|
||||||
|
tools=[{"type": "function", "function": {"name": "write_file"}}],
|
||||||
|
model=model,
|
||||||
|
on_tool_call_delta=on_tool_delta,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert deltas == [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"call_id": "call_write",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": '{"path":"notes.md","content":"',
|
||||||
|
},
|
||||||
|
{"index": 0, "call_id": "", "name": "", "arguments_delta": 'line\\n"}'},
|
||||||
|
]
|
||||||
|
assert result.tool_calls[0].name == "write_file"
|
||||||
|
assert result.tool_calls[0].arguments == {"path": "notes.md", "content": "line\n"}
|
||||||
|
kwargs = mock_chat.await_args.kwargs
|
||||||
|
if provider_name == "zhipu":
|
||||||
|
assert kwargs["extra_body"]["tool_stream"] is True
|
||||||
|
else:
|
||||||
|
assert kwargs.get("extra_body", {}).get("tool_stream") is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_openai_compat_stream_forwards_legacy_function_call_argument_deltas() -> None:
|
||||||
|
mock_chat = AsyncMock(return_value=_fake_chat_stream_legacy_function_call_chunks())
|
||||||
|
deltas: list[dict] = []
|
||||||
|
|
||||||
|
async def on_tool_delta(delta: dict) -> None:
|
||||||
|
deltas.append(delta)
|
||||||
|
|
||||||
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_openai:
|
||||||
|
client_instance = mock_openai.return_value
|
||||||
|
client_instance.chat.completions.create = mock_chat
|
||||||
|
|
||||||
|
provider = OpenAICompatProvider(
|
||||||
|
api_key="sk-test",
|
||||||
|
default_model="deepseek-chat",
|
||||||
|
spec=find_by_name("deepseek"),
|
||||||
|
)
|
||||||
|
result = await provider.chat_stream(
|
||||||
|
messages=[{"role": "user", "content": "write"}],
|
||||||
|
tools=[{"type": "function", "function": {"name": "write_file"}}],
|
||||||
|
model="deepseek-chat",
|
||||||
|
on_tool_call_delta=on_tool_delta,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert deltas == [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"call_id": "",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": '{"path":"notes.md","content":"',
|
||||||
|
},
|
||||||
|
{"index": 0, "call_id": "", "name": "", "arguments_delta": 'line\\n"}'},
|
||||||
|
]
|
||||||
|
assert result.tool_calls[0].name == "write_file"
|
||||||
|
assert result.tool_calls[0].arguments == {"path": "notes.md", "content": "line\n"}
|
||||||
|
|
||||||
|
|
||||||
class _FakeResponsesError(Exception):
|
class _FakeResponsesError(Exception):
|
||||||
def __init__(self, status_code: int, text: str):
|
def __init__(self, status_code: int, text: str):
|
||||||
super().__init__(text)
|
super().__init__(text)
|
||||||
@@ -233,27 +449,28 @@ def test_gemma_routes_to_gemini_provider() -> None:
|
|||||||
assert "gemma" in spec.keywords
|
assert "gemma" in spec.keywords
|
||||||
|
|
||||||
|
|
||||||
def test_openrouter_sets_default_attribution_headers() -> None:
|
async def test_openrouter_sets_default_attribution_headers() -> None:
|
||||||
spec = find_by_name("openrouter")
|
spec = find_by_name("openrouter")
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient:
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_cls:
|
||||||
OpenAICompatProvider(
|
provider = OpenAICompatProvider(
|
||||||
api_key="sk-or-test-key",
|
api_key="sk-or-test-key",
|
||||||
api_base="https://openrouter.ai/api/v1",
|
api_base="https://openrouter.ai/api/v1",
|
||||||
default_model="anthropic/claude-sonnet-4-5",
|
default_model="anthropic/claude-sonnet-4-5",
|
||||||
spec=spec,
|
spec=spec,
|
||||||
)
|
)
|
||||||
|
await provider._ensure_client()
|
||||||
|
|
||||||
headers = MockClient.call_args.kwargs["default_headers"]
|
headers = mock_client_cls.call_args.kwargs["default_headers"]
|
||||||
assert headers["HTTP-Referer"] == "https://github.com/HKUDS/nanobot"
|
assert headers["HTTP-Referer"] == "https://github.com/HKUDS/nanobot"
|
||||||
assert headers["X-OpenRouter-Title"] == "nanobot"
|
assert headers["X-OpenRouter-Title"] == "nanobot"
|
||||||
assert headers["X-OpenRouter-Categories"] == "cli-agent,personal-agent"
|
assert headers["X-OpenRouter-Categories"] == "cli-agent,personal-agent"
|
||||||
assert "x-session-affinity" in headers
|
assert "x-session-affinity" in headers
|
||||||
|
|
||||||
|
|
||||||
def test_openrouter_user_headers_override_default_attribution() -> None:
|
async def test_openrouter_user_headers_override_default_attribution() -> None:
|
||||||
spec = find_by_name("openrouter")
|
spec = find_by_name("openrouter")
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient:
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_cls:
|
||||||
OpenAICompatProvider(
|
provider = OpenAICompatProvider(
|
||||||
api_key="sk-or-test-key",
|
api_key="sk-or-test-key",
|
||||||
api_base="https://openrouter.ai/api/v1",
|
api_base="https://openrouter.ai/api/v1",
|
||||||
default_model="anthropic/claude-sonnet-4-5",
|
default_model="anthropic/claude-sonnet-4-5",
|
||||||
@@ -264,8 +481,9 @@ def test_openrouter_user_headers_override_default_attribution() -> None:
|
|||||||
},
|
},
|
||||||
spec=spec,
|
spec=spec,
|
||||||
)
|
)
|
||||||
|
await provider._ensure_client()
|
||||||
|
|
||||||
headers = MockClient.call_args.kwargs["default_headers"]
|
headers = mock_client_cls.call_args.kwargs["default_headers"]
|
||||||
assert headers["HTTP-Referer"] == "https://nanobot.ai"
|
assert headers["HTTP-Referer"] == "https://nanobot.ai"
|
||||||
assert headers["X-OpenRouter-Title"] == "Nanobot Pro"
|
assert headers["X-OpenRouter-Title"] == "Nanobot Pro"
|
||||||
assert headers["X-OpenRouter-Categories"] == "cli-agent,personal-agent"
|
assert headers["X-OpenRouter-Categories"] == "cli-agent,personal-agent"
|
||||||
|
|||||||
@@ -44,9 +44,15 @@ class TestShouldExecuteTools:
|
|||||||
resp = _response("stop")
|
resp = _response("stop")
|
||||||
assert resp.should_execute_tools is True
|
assert resp.should_execute_tools is True
|
||||||
|
|
||||||
|
def test_legacy_function_call_reason_executes(self) -> None:
|
||||||
|
# Older OpenAI-compatible streaming APIs can still use the singular
|
||||||
|
# function_call finish reason while carrying a tool-call-shaped payload.
|
||||||
|
resp = _response("function_call")
|
||||||
|
assert resp.should_execute_tools is True
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"anomalous_reason",
|
"anomalous_reason",
|
||||||
["refusal", "content_filter", "error", "length", "function_call", ""],
|
["refusal", "content_filter", "error", "length", ""],
|
||||||
)
|
)
|
||||||
def test_tool_calls_under_anomalous_reason_blocked(self, anomalous_reason: str) -> None:
|
def test_tool_calls_under_anomalous_reason_blocked(self, anomalous_reason: str) -> None:
|
||||||
# This is the #3220 bug: gateways injecting tool_calls under any of these
|
# This is the #3220 bug: gateways injecting tool_calls under any of these
|
||||||
|
|||||||
@@ -85,17 +85,18 @@ class TestIsLocalEndpoint:
|
|||||||
class TestLocalKeepaliveConfig:
|
class TestLocalKeepaliveConfig:
|
||||||
"""Verify that local endpoints get keepalive_expiry=0."""
|
"""Verify that local endpoints get keepalive_expiry=0."""
|
||||||
|
|
||||||
def test_local_spec_disables_keepalive(self):
|
async def test_local_spec_disables_keepalive(self):
|
||||||
spec = _make_spec(is_local=True)
|
spec = _make_spec(is_local=True)
|
||||||
spec.env_key = ""
|
spec.env_key = ""
|
||||||
spec.default_api_base = "http://localhost:11434/v1"
|
spec.default_api_base = "http://localhost:11434/v1"
|
||||||
provider = OpenAICompatProvider(
|
provider = OpenAICompatProvider(
|
||||||
api_key="test", api_base="http://localhost:11434/v1", spec=spec,
|
api_key="test", api_base="http://localhost:11434/v1", spec=spec,
|
||||||
)
|
)
|
||||||
|
await provider._ensure_client()
|
||||||
pool = provider._client._client._transport._pool
|
pool = provider._client._client._transport._pool
|
||||||
assert pool._keepalive_expiry == 0
|
assert pool._keepalive_expiry == 0
|
||||||
|
|
||||||
def test_lan_ip_disables_keepalive(self):
|
async def test_lan_ip_disables_keepalive(self):
|
||||||
"""A generic 'openai' spec with a LAN IP should still disable keepalive."""
|
"""A generic 'openai' spec with a LAN IP should still disable keepalive."""
|
||||||
spec = _make_spec(is_local=False)
|
spec = _make_spec(is_local=False)
|
||||||
spec.env_key = ""
|
spec.env_key = ""
|
||||||
@@ -103,16 +104,18 @@ class TestLocalKeepaliveConfig:
|
|||||||
provider = OpenAICompatProvider(
|
provider = OpenAICompatProvider(
|
||||||
api_key="test", api_base="http://192.168.8.188:1234/v1", spec=spec,
|
api_key="test", api_base="http://192.168.8.188:1234/v1", spec=spec,
|
||||||
)
|
)
|
||||||
|
await provider._ensure_client()
|
||||||
pool = provider._client._client._transport._pool
|
pool = provider._client._client._transport._pool
|
||||||
assert pool._keepalive_expiry == 0
|
assert pool._keepalive_expiry == 0
|
||||||
|
|
||||||
def test_cloud_keeps_default_keepalive(self):
|
async def test_cloud_keeps_default_keepalive(self):
|
||||||
spec = _make_spec(is_local=False)
|
spec = _make_spec(is_local=False)
|
||||||
spec.env_key = ""
|
spec.env_key = ""
|
||||||
spec.default_api_base = "https://api.openai.com/v1"
|
spec.default_api_base = "https://api.openai.com/v1"
|
||||||
provider = OpenAICompatProvider(
|
provider = OpenAICompatProvider(
|
||||||
api_key="test", api_base=None, spec=spec,
|
api_key="test", api_base=None, spec=spec,
|
||||||
)
|
)
|
||||||
|
await provider._ensure_client()
|
||||||
pool = provider._client._client._transport._pool
|
pool = provider._client._client._transport._pool
|
||||||
# Default httpx keepalive is 5.0s
|
# Default httpx keepalive is 5.0s
|
||||||
assert pool._keepalive_expiry == 5.0
|
assert pool._keepalive_expiry == 5.0
|
||||||
|
|||||||
@@ -16,7 +16,15 @@ async def test_codex_prompt_cache_key_uses_stable_conversation_prefix(monkeypatc
|
|||||||
lambda: SimpleNamespace(account_id="acct", access="token"),
|
lambda: SimpleNamespace(account_id="acct", access="token"),
|
||||||
)
|
)
|
||||||
|
|
||||||
async def fake_request(url, headers, body, verify, on_content_delta=None):
|
async def fake_request(
|
||||||
|
url,
|
||||||
|
headers,
|
||||||
|
body,
|
||||||
|
verify,
|
||||||
|
on_content_delta=None,
|
||||||
|
on_tool_call_delta=None,
|
||||||
|
):
|
||||||
|
_ = on_tool_call_delta
|
||||||
bodies.append(body)
|
bodies.append(body)
|
||||||
return "ok", [], "stop"
|
return "ok", [], "stop"
|
||||||
|
|
||||||
|
|||||||
@@ -8,16 +8,18 @@ def _assert_openai_compat_timeout(timeout) -> None:
|
|||||||
assert timeout == 120.0
|
assert timeout == 120.0
|
||||||
|
|
||||||
|
|
||||||
def test_openai_compat_provider_sets_sdk_timeout() -> None:
|
async def test_openai_compat_provider_defers_sdk_client_until_first_use() -> None:
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_async_openai:
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_async_openai:
|
||||||
OpenAICompatProvider(api_key="test-key", api_base="https://example.com/v1")
|
provider = OpenAICompatProvider(api_key="test-key", api_base="https://example.com/v1")
|
||||||
|
mock_async_openai.assert_not_called()
|
||||||
|
await provider._ensure_client()
|
||||||
|
|
||||||
kwargs = mock_async_openai.call_args.kwargs
|
kwargs = mock_async_openai.call_args.kwargs
|
||||||
_assert_openai_compat_timeout(kwargs["timeout"])
|
_assert_openai_compat_timeout(kwargs["timeout"])
|
||||||
assert kwargs["http_client"] is None
|
assert kwargs["http_client"] is None
|
||||||
|
|
||||||
|
|
||||||
def test_openai_compat_provider_sets_timeout_on_local_http_client() -> None:
|
async def test_openai_compat_provider_sets_timeout_on_local_http_client() -> None:
|
||||||
spec = ProviderSpec(
|
spec = ProviderSpec(
|
||||||
name="local",
|
name="local",
|
||||||
keywords=(),
|
keywords=(),
|
||||||
@@ -29,11 +31,13 @@ def test_openai_compat_provider_sets_timeout_on_local_http_client() -> None:
|
|||||||
with (
|
with (
|
||||||
patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_async_openai,
|
patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_async_openai,
|
||||||
patch(
|
patch(
|
||||||
"nanobot.providers.openai_compat_provider.httpx.AsyncClient",
|
"httpx.AsyncClient",
|
||||||
return_value=sentinel.http_client,
|
return_value=sentinel.http_client,
|
||||||
) as mock_http_client,
|
) as mock_http_client,
|
||||||
):
|
):
|
||||||
OpenAICompatProvider(spec=spec)
|
provider = OpenAICompatProvider(spec=spec)
|
||||||
|
mock_async_openai.assert_not_called()
|
||||||
|
await provider._ensure_client()
|
||||||
|
|
||||||
client_kwargs = mock_http_client.call_args.kwargs
|
client_kwargs = mock_http_client.call_args.kwargs
|
||||||
_assert_openai_compat_timeout(client_kwargs["timeout"])
|
_assert_openai_compat_timeout(client_kwargs["timeout"])
|
||||||
@@ -44,10 +48,11 @@ def test_openai_compat_provider_sets_timeout_on_local_http_client() -> None:
|
|||||||
assert openai_kwargs["http_client"] is sentinel.http_client
|
assert openai_kwargs["http_client"] is sentinel.http_client
|
||||||
|
|
||||||
|
|
||||||
def test_openai_compat_provider_timeout_can_be_overridden_by_env(monkeypatch) -> None:
|
async def test_openai_compat_provider_timeout_can_be_overridden_by_env(monkeypatch) -> None:
|
||||||
monkeypatch.setenv("NANOBOT_OPENAI_COMPAT_TIMEOUT_S", "45")
|
monkeypatch.setenv("NANOBOT_OPENAI_COMPAT_TIMEOUT_S", "45")
|
||||||
|
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_async_openai:
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_async_openai:
|
||||||
OpenAICompatProvider(api_key="test-key", api_base="https://example.com/v1")
|
provider = OpenAICompatProvider(api_key="test-key", api_base="https://example.com/v1")
|
||||||
|
await provider._ensure_client()
|
||||||
|
|
||||||
assert mock_async_openai.call_args.kwargs["timeout"] == 45.0
|
assert mock_async_openai.call_args.kwargs["timeout"] == 45.0
|
||||||
|
|||||||
@@ -453,6 +453,56 @@ class TestConsumeSdkStream:
|
|||||||
assert tool_calls[0].name == "get_weather"
|
assert tool_calls[0].name == "get_weather"
|
||||||
assert tool_calls[0].arguments == {"city": "SF"}
|
assert tool_calls[0].arguments == {"city": "SF"}
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_tool_call_argument_delta_callback(self):
|
||||||
|
item_added = MagicMock(type="function_call", call_id="c1", id="fc1", arguments="")
|
||||||
|
item_added.name = "write_file"
|
||||||
|
ev1 = MagicMock(type="response.output_item.added", item=item_added)
|
||||||
|
ev2 = MagicMock(
|
||||||
|
type="response.function_call_arguments.delta",
|
||||||
|
call_id="c1",
|
||||||
|
delta='{"path":"a.txt","content":"',
|
||||||
|
)
|
||||||
|
ev3 = MagicMock(
|
||||||
|
type="response.function_call_arguments.delta",
|
||||||
|
call_id="c1",
|
||||||
|
delta='hello\\n',
|
||||||
|
)
|
||||||
|
ev4 = MagicMock(
|
||||||
|
type="response.function_call_arguments.done",
|
||||||
|
call_id="c1",
|
||||||
|
arguments='{"path":"a.txt","content":"hello\\n"}',
|
||||||
|
)
|
||||||
|
item_done = MagicMock(
|
||||||
|
type="function_call",
|
||||||
|
call_id="c1",
|
||||||
|
id="fc1",
|
||||||
|
arguments='{"path":"a.txt","content":"hello\\n"}',
|
||||||
|
)
|
||||||
|
item_done.name = "write_file"
|
||||||
|
ev5 = MagicMock(type="response.output_item.done", item=item_done)
|
||||||
|
resp_obj = MagicMock(status="completed", usage=None, output=[])
|
||||||
|
ev6 = MagicMock(type="response.completed", response=resp_obj)
|
||||||
|
deltas: list[dict] = []
|
||||||
|
|
||||||
|
async def cb(delta: dict) -> None:
|
||||||
|
deltas.append(delta)
|
||||||
|
|
||||||
|
async def stream():
|
||||||
|
for e in [ev1, ev2, ev3, ev4, ev5, ev6]:
|
||||||
|
yield e
|
||||||
|
|
||||||
|
await consume_sdk_stream(stream(), on_tool_call_delta=cb)
|
||||||
|
assert deltas == [
|
||||||
|
{"call_id": "c1", "name": "write_file", "arguments_delta": ""},
|
||||||
|
{
|
||||||
|
"call_id": "c1",
|
||||||
|
"name": "write_file",
|
||||||
|
"arguments_delta": '{"path":"a.txt","content":"',
|
||||||
|
},
|
||||||
|
{"call_id": "c1", "name": "write_file", "arguments_delta": "hello\\n"},
|
||||||
|
]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_usage_extracted(self):
|
async def test_usage_extracted(self):
|
||||||
usage_obj = MagicMock(input_tokens=10, output_tokens=5, total_tokens=15)
|
usage_obj = MagicMock(input_tokens=10, output_tokens=5, total_tokens=15)
|
||||||
|
|||||||
@@ -5,9 +5,10 @@ from nanobot.providers.azure_openai_provider import AzureOpenAIProvider
|
|||||||
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
|
|
||||||
|
|
||||||
def test_openai_compat_disables_sdk_retries_by_default() -> None:
|
async def test_openai_compat_disables_sdk_retries_by_default() -> None:
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client:
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client:
|
||||||
OpenAICompatProvider(api_key="sk-test", default_model="gpt-4o")
|
provider = OpenAICompatProvider(api_key="sk-test", default_model="gpt-4o")
|
||||||
|
await provider._ensure_client()
|
||||||
|
|
||||||
kwargs = mock_client.call_args.kwargs
|
kwargs = mock_client.call_args.kwargs
|
||||||
assert kwargs["max_retries"] == 0
|
assert kwargs["max_retries"] == 0
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ def test_importing_providers_package_is_lazy(monkeypatch) -> None:
|
|||||||
monkeypatch.delitem(sys.modules, "nanobot.providers.openai_compat_provider", raising=False)
|
monkeypatch.delitem(sys.modules, "nanobot.providers.openai_compat_provider", raising=False)
|
||||||
monkeypatch.delitem(sys.modules, "nanobot.providers.openai_codex_provider", raising=False)
|
monkeypatch.delitem(sys.modules, "nanobot.providers.openai_codex_provider", raising=False)
|
||||||
monkeypatch.delitem(sys.modules, "nanobot.providers.github_copilot_provider", raising=False)
|
monkeypatch.delitem(sys.modules, "nanobot.providers.github_copilot_provider", raising=False)
|
||||||
|
monkeypatch.delitem(sys.modules, "nanobot.providers.xai_oauth_provider", raising=False)
|
||||||
monkeypatch.delitem(sys.modules, "nanobot.providers.azure_openai_provider", raising=False)
|
monkeypatch.delitem(sys.modules, "nanobot.providers.azure_openai_provider", raising=False)
|
||||||
monkeypatch.delitem(sys.modules, "nanobot.providers.bedrock_provider", raising=False)
|
monkeypatch.delitem(sys.modules, "nanobot.providers.bedrock_provider", raising=False)
|
||||||
|
|
||||||
@@ -21,6 +22,7 @@ def test_importing_providers_package_is_lazy(monkeypatch) -> None:
|
|||||||
assert "nanobot.providers.openai_compat_provider" not in sys.modules
|
assert "nanobot.providers.openai_compat_provider" not in sys.modules
|
||||||
assert "nanobot.providers.openai_codex_provider" not in sys.modules
|
assert "nanobot.providers.openai_codex_provider" not in sys.modules
|
||||||
assert "nanobot.providers.github_copilot_provider" not in sys.modules
|
assert "nanobot.providers.github_copilot_provider" not in sys.modules
|
||||||
|
assert "nanobot.providers.xai_oauth_provider" not in sys.modules
|
||||||
assert "nanobot.providers.azure_openai_provider" not in sys.modules
|
assert "nanobot.providers.azure_openai_provider" not in sys.modules
|
||||||
assert "nanobot.providers.bedrock_provider" not in sys.modules
|
assert "nanobot.providers.bedrock_provider" not in sys.modules
|
||||||
assert providers.__all__ == [
|
assert providers.__all__ == [
|
||||||
@@ -30,6 +32,7 @@ def test_importing_providers_package_is_lazy(monkeypatch) -> None:
|
|||||||
"OpenAICompatProvider",
|
"OpenAICompatProvider",
|
||||||
"OpenAICodexProvider",
|
"OpenAICodexProvider",
|
||||||
"GitHubCopilotProvider",
|
"GitHubCopilotProvider",
|
||||||
|
"XaiOAuthProvider",
|
||||||
"AzureOpenAIProvider",
|
"AzureOpenAIProvider",
|
||||||
"BedrockProvider",
|
"BedrockProvider",
|
||||||
]
|
]
|
||||||
@@ -50,3 +53,9 @@ def test_openai_codex_supports_progress_deltas() -> None:
|
|||||||
from nanobot.providers.openai_codex_provider import OpenAICodexProvider
|
from nanobot.providers.openai_codex_provider import OpenAICodexProvider
|
||||||
|
|
||||||
assert OpenAICodexProvider.supports_progress_deltas is True
|
assert OpenAICodexProvider.supports_progress_deltas is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_xai_oauth_supports_progress_deltas() -> None:
|
||||||
|
from nanobot.providers.xai_oauth_provider import XaiOAuthProvider
|
||||||
|
|
||||||
|
assert XaiOAuthProvider.supports_progress_deltas is True
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
"""Tests for the Skywork provider registration."""
|
||||||
|
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from nanobot.config.schema import Config, ProvidersConfig
|
||||||
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
|
from nanobot.providers.registry import PROVIDERS, find_by_name
|
||||||
|
|
||||||
|
|
||||||
|
def test_skywork_config_field_exists() -> None:
|
||||||
|
config = ProvidersConfig()
|
||||||
|
|
||||||
|
assert hasattr(config, "skywork")
|
||||||
|
|
||||||
|
|
||||||
|
def test_skywork_provider_in_registry() -> None:
|
||||||
|
specs = {spec.name: spec for spec in PROVIDERS}
|
||||||
|
|
||||||
|
assert "skywork" in specs
|
||||||
|
skywork = specs["skywork"]
|
||||||
|
assert skywork.backend == "openai_compat"
|
||||||
|
assert skywork.env_key == "SKYWORK_API_KEY"
|
||||||
|
assert ("APIFREE_API_KEY", "{api_key}") in skywork.env_extras
|
||||||
|
assert skywork.display_name == "Skywork"
|
||||||
|
assert skywork.is_gateway is True
|
||||||
|
assert skywork.detect_by_base_keyword == "apifree.ai"
|
||||||
|
assert skywork.default_api_base == "https://api.apifree.ai/agent/v1"
|
||||||
|
assert skywork.supports_max_completion_tokens is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_by_name_skywork() -> None:
|
||||||
|
spec = find_by_name("skywork")
|
||||||
|
|
||||||
|
assert spec is not None
|
||||||
|
assert spec.name == "skywork"
|
||||||
|
|
||||||
|
|
||||||
|
def test_skywork_model_auto_matches_with_default_api_base() -> None:
|
||||||
|
config = Config.model_validate(
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"skywork": {
|
||||||
|
"apiKey": "sky-key",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"model": "skywork-ai/skyclaw-v1",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert config.get_provider_name("skywork-ai/skyclaw-v1") == "skywork"
|
||||||
|
assert config.get_api_key("skywork-ai/skyclaw-v1") == "sky-key"
|
||||||
|
assert config.get_api_base("skywork-ai/skyclaw-v1") == "https://api.apifree.ai/agent/v1"
|
||||||
|
|
||||||
|
|
||||||
|
def test_skywork_preserves_model_id_and_uses_chat_completion_max_tokens() -> None:
|
||||||
|
spec = find_by_name("skywork")
|
||||||
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
||||||
|
provider = OpenAICompatProvider(
|
||||||
|
api_key="sky-key",
|
||||||
|
default_model="skywork-ai/skyclaw-v1",
|
||||||
|
spec=spec,
|
||||||
|
)
|
||||||
|
|
||||||
|
kwargs = provider._build_kwargs(
|
||||||
|
messages=[{"role": "user", "content": "hi"}],
|
||||||
|
tools=None,
|
||||||
|
model="skywork-ai/skyclaw-v1",
|
||||||
|
max_tokens=1024,
|
||||||
|
temperature=0.7,
|
||||||
|
reasoning_effort=None,
|
||||||
|
tool_choice=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert kwargs["model"] == "skywork-ai/skyclaw-v1"
|
||||||
|
assert kwargs["max_tokens"] == 1024
|
||||||
|
assert "max_completion_tokens" not in kwargs
|
||||||
@@ -0,0 +1,146 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import stat
|
||||||
|
from urllib.parse import parse_qs, urlparse
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
import nanobot.providers.xai_oauth_provider as auth
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_xai_authorization_url_includes_pkce_and_grok_scope() -> None:
|
||||||
|
endpoints = auth.XaiOAuthEndpoints(
|
||||||
|
authorization_endpoint="https://auth.x.ai/authorize",
|
||||||
|
token_endpoint="https://auth.x.ai/oauth/token",
|
||||||
|
)
|
||||||
|
|
||||||
|
url = auth.build_xai_authorization_url(
|
||||||
|
endpoints,
|
||||||
|
verifier="verifier",
|
||||||
|
state="state",
|
||||||
|
nonce="nonce",
|
||||||
|
)
|
||||||
|
|
||||||
|
parsed = urlparse(url)
|
||||||
|
params = parse_qs(parsed.query)
|
||||||
|
assert parsed.scheme == "https"
|
||||||
|
assert parsed.hostname == "auth.x.ai"
|
||||||
|
assert params["client_id"] == [auth.DEFAULT_XAI_CLIENT_ID]
|
||||||
|
assert params["code_challenge"] == [auth.pkce_challenge("verifier")]
|
||||||
|
assert params["code_challenge_method"] == ["S256"]
|
||||||
|
assert params["scope"] == [auth.DEFAULT_XAI_SCOPE]
|
||||||
|
assert params["nonce"] == ["nonce"]
|
||||||
|
assert params["plan"] == ["generic"]
|
||||||
|
assert params["referrer"] == ["nanobot"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_parse_callback_value_accepts_fallback_shapes() -> None:
|
||||||
|
assert auth._parse_callback_value("https://localhost/callback?code=abc&state=state") == ("abc", "state")
|
||||||
|
assert auth._parse_callback_value("?code=abc&state=state") == ("abc", "state")
|
||||||
|
assert auth._parse_callback_value("code=abc&state=state") == ("abc", "state")
|
||||||
|
assert auth._parse_callback_value("fallback-code") == ("fallback-code", None)
|
||||||
|
|
||||||
|
|
||||||
|
def test_file_storage_fallback_is_private_and_round_trips(tmp_path, monkeypatch) -> None:
|
||||||
|
monkeypatch.setenv("NANOBOT_HOME", str(tmp_path))
|
||||||
|
monkeypatch.setattr(auth, "_keyring_set", lambda _tokens: False)
|
||||||
|
monkeypatch.setattr(auth, "_keyring_get", lambda: None)
|
||||||
|
|
||||||
|
saved = auth.save_xai_oauth_credential(
|
||||||
|
auth.XaiOAuthCredential(
|
||||||
|
access_token="access",
|
||||||
|
refresh_token="refresh",
|
||||||
|
expires_at=123.0,
|
||||||
|
account_id="acct",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
path = auth.get_xai_oauth_metadata_path()
|
||||||
|
payload = json.loads(path.read_text(encoding="utf-8"))
|
||||||
|
assert saved.storage == "file"
|
||||||
|
assert payload["storage"] == "file"
|
||||||
|
assert payload["tokens"]["access_token"] == "access"
|
||||||
|
if os.name != "nt":
|
||||||
|
assert stat.S_IMODE(path.stat().st_mode) == 0o600
|
||||||
|
|
||||||
|
loaded = auth.load_xai_oauth_credential()
|
||||||
|
assert loaded is not None
|
||||||
|
assert loaded.access_token == "access"
|
||||||
|
assert loaded.refresh_token == "refresh"
|
||||||
|
assert loaded.account_id == "acct"
|
||||||
|
assert loaded.storage == "file"
|
||||||
|
|
||||||
|
|
||||||
|
def test_keyring_storage_keeps_tokens_out_of_metadata(tmp_path, monkeypatch) -> None:
|
||||||
|
monkeypatch.setenv("NANOBOT_HOME", str(tmp_path))
|
||||||
|
secret: dict[str, object] = {}
|
||||||
|
|
||||||
|
def fake_set(tokens: dict[str, object]) -> bool:
|
||||||
|
secret.update(tokens)
|
||||||
|
return True
|
||||||
|
|
||||||
|
monkeypatch.setattr(auth, "_keyring_set", fake_set)
|
||||||
|
monkeypatch.setattr(auth, "_keyring_get", lambda: dict(secret))
|
||||||
|
|
||||||
|
auth.save_xai_oauth_credential(
|
||||||
|
auth.XaiOAuthCredential(
|
||||||
|
access_token="access",
|
||||||
|
refresh_token="refresh",
|
||||||
|
expires_at=123.0,
|
||||||
|
account_id="acct",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = json.loads(auth.get_xai_oauth_metadata_path().read_text(encoding="utf-8"))
|
||||||
|
assert payload["storage"] == "keyring"
|
||||||
|
assert "tokens" not in payload
|
||||||
|
assert auth.load_xai_oauth_credential().access_token == "access"
|
||||||
|
|
||||||
|
|
||||||
|
def test_exchange_xai_oauth_code_sends_required_code_challenge(monkeypatch) -> None:
|
||||||
|
captured: dict[str, object] = {}
|
||||||
|
|
||||||
|
class FakeResponse:
|
||||||
|
status_code = 200
|
||||||
|
text = ""
|
||||||
|
|
||||||
|
def json(self) -> dict[str, object]:
|
||||||
|
return {"access_token": "access", "refresh_token": "refresh", "expires_in": 3600}
|
||||||
|
|
||||||
|
class FakeClient:
|
||||||
|
def __init__(self, *args, **kwargs) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def __enter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, *args) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def post(self, url: str, headers: dict[str, str], data: dict[str, str]) -> FakeResponse:
|
||||||
|
captured["url"] = url
|
||||||
|
captured["headers"] = headers
|
||||||
|
captured["data"] = data
|
||||||
|
return FakeResponse()
|
||||||
|
|
||||||
|
monkeypatch.setattr(auth.httpx, "Client", FakeClient)
|
||||||
|
endpoints = auth.XaiOAuthEndpoints(
|
||||||
|
authorization_endpoint="https://auth.x.ai/authorize",
|
||||||
|
token_endpoint="https://auth.x.ai/oauth/token",
|
||||||
|
)
|
||||||
|
|
||||||
|
credential = auth.exchange_xai_oauth_code("code", verifier="verifier", endpoints=endpoints)
|
||||||
|
|
||||||
|
assert credential.access_token == "access"
|
||||||
|
assert captured["url"] == "https://auth.x.ai/oauth/token"
|
||||||
|
data = captured["data"]
|
||||||
|
assert data["code_verifier"] == "verifier"
|
||||||
|
assert data["code_challenge"] == auth.pkce_challenge("verifier")
|
||||||
|
assert data["code_challenge_method"] == "S256"
|
||||||
|
|
||||||
|
|
||||||
|
def test_rejects_non_xai_discovery_endpoints() -> None:
|
||||||
|
with pytest.raises(RuntimeError):
|
||||||
|
auth._validate_xai_endpoint("https://example.com/oauth/token", "token_endpoint")
|
||||||
@@ -0,0 +1,141 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
from nanobot.config.schema import XaiOAuthXSearchConfig
|
||||||
|
import nanobot.providers.xai_oauth_provider as xai_oauth_provider
|
||||||
|
from nanobot.providers.xai_oauth_provider import (
|
||||||
|
XaiOAuthCredential,
|
||||||
|
XaiOAuthProvider,
|
||||||
|
_build_xai_responses_body,
|
||||||
|
_strip_model_prefix,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_xai_oauth_strip_prefix_supports_aliases() -> None:
|
||||||
|
assert _strip_model_prefix("xai-oauth/grok-4.3") == "grok-4.3"
|
||||||
|
assert _strip_model_prefix("xai_oauth/grok-4.3") == "grok-4.3"
|
||||||
|
assert _strip_model_prefix("grok-oauth/grok-4.3") == "grok-4.3"
|
||||||
|
assert _strip_model_prefix("grok-4.3") == "grok-4.3"
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_xai_responses_body_keeps_system_prompt_in_input() -> None:
|
||||||
|
body = _build_xai_responses_body(
|
||||||
|
messages=[
|
||||||
|
{"role": "system", "content": "You are nanobot."},
|
||||||
|
{"role": "user", "content": "hi"},
|
||||||
|
],
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "ping",
|
||||||
|
"description": "Ping",
|
||||||
|
"parameters": {"type": "object", "properties": {}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
model="xai-oauth/grok-4.3",
|
||||||
|
max_tokens=32,
|
||||||
|
temperature=0.2,
|
||||||
|
reasoning_effort="high",
|
||||||
|
tool_choice=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert body["model"] == "grok-4.3"
|
||||||
|
assert "instructions" not in body
|
||||||
|
assert body["input"][0] == {
|
||||||
|
"role": "system",
|
||||||
|
"content": [{"type": "input_text", "text": "You are nanobot."}],
|
||||||
|
}
|
||||||
|
assert body["input"][1]["role"] == "user"
|
||||||
|
assert body["max_output_tokens"] == 32
|
||||||
|
assert body["temperature"] == 0.2
|
||||||
|
assert body["reasoning"] == {"effort": "high"}
|
||||||
|
assert body["tools"][0]["name"] == "ping"
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_xai_responses_body_attaches_hosted_x_search_by_default() -> None:
|
||||||
|
body = _build_xai_responses_body(
|
||||||
|
messages=[{"role": "user", "content": "what is happening on X?"}],
|
||||||
|
tools=None,
|
||||||
|
model="xai-oauth/grok-4.3",
|
||||||
|
max_tokens=32,
|
||||||
|
temperature=0.2,
|
||||||
|
reasoning_effort=None,
|
||||||
|
tool_choice=None,
|
||||||
|
hosted_x_search=XaiOAuthXSearchConfig(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert body["tools"] == [{"type": "x_search"}]
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_xai_responses_body_can_customize_hosted_x_search() -> None:
|
||||||
|
body = _build_xai_responses_body(
|
||||||
|
messages=[{"role": "user", "content": "what is happening on X?"}],
|
||||||
|
tools=None,
|
||||||
|
model="xai-oauth/grok-4.3",
|
||||||
|
max_tokens=32,
|
||||||
|
temperature=0.2,
|
||||||
|
reasoning_effort=None,
|
||||||
|
tool_choice=None,
|
||||||
|
hosted_x_search=XaiOAuthXSearchConfig(
|
||||||
|
allowed_x_handles=["@xai", " nanobot "],
|
||||||
|
enable_image_understanding=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert body["tools"] == [
|
||||||
|
{
|
||||||
|
"type": "x_search",
|
||||||
|
"allowed_x_handles": ["xai", "nanobot"],
|
||||||
|
"enable_image_understanding": True,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_xai_responses_body_omits_disabled_hosted_x_search() -> None:
|
||||||
|
body = _build_xai_responses_body(
|
||||||
|
messages=[{"role": "user", "content": "hi"}],
|
||||||
|
tools=None,
|
||||||
|
model="xai-oauth/grok-4.3",
|
||||||
|
max_tokens=32,
|
||||||
|
temperature=0.2,
|
||||||
|
reasoning_effort=None,
|
||||||
|
tool_choice=None,
|
||||||
|
hosted_x_search=XaiOAuthXSearchConfig(enable=False),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "tools" not in body
|
||||||
|
|
||||||
|
|
||||||
|
def test_xai_oauth_provider_refreshes_once_on_401(monkeypatch) -> None:
|
||||||
|
async def run() -> None:
|
||||||
|
response = await provider.chat([{"role": "user", "content": "hi"}])
|
||||||
|
|
||||||
|
assert response.content == "ok"
|
||||||
|
assert response.finish_reason == "stop"
|
||||||
|
assert calls == [("resolve", False), ("resolve", True)]
|
||||||
|
|
||||||
|
provider = XaiOAuthProvider(default_model="xai-oauth/grok-4.3")
|
||||||
|
credentials = [
|
||||||
|
XaiOAuthCredential(access_token="expired"),
|
||||||
|
XaiOAuthCredential(access_token="fresh"),
|
||||||
|
]
|
||||||
|
calls: list[tuple[str, bool]] = []
|
||||||
|
|
||||||
|
def fake_resolve(*, force_refresh: bool = False) -> XaiOAuthCredential:
|
||||||
|
calls.append(("resolve", force_refresh))
|
||||||
|
return credentials.pop(0)
|
||||||
|
|
||||||
|
async def fake_request(credential, body, on_content_delta=None, on_tool_call_delta=None):
|
||||||
|
from nanobot.providers.xai_oauth_provider import _XaiHTTPError
|
||||||
|
|
||||||
|
if credential.access_token == "expired":
|
||||||
|
raise _XaiHTTPError("expired", status_code=401)
|
||||||
|
return "ok", [], "stop"
|
||||||
|
|
||||||
|
monkeypatch.setattr(xai_oauth_provider, "resolve_xai_oauth_credential", fake_resolve)
|
||||||
|
monkeypatch.setattr(xai_oauth_provider, "_request_xai", fake_request)
|
||||||
|
|
||||||
|
asyncio.run(run())
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user