mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-08 21:38:40 +03:00
Compare commits
24
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f44ee6cb27 | ||
|
|
5be3df1d6f | ||
|
|
79c23787f6 | ||
|
|
b647aa5f47 | ||
|
|
a9a8bdcef6 | ||
|
|
2d81cc0ae1 | ||
|
|
aed6b6967c | ||
|
|
3874b3acf4 | ||
|
|
01725bab11 | ||
|
|
a786e3d225 | ||
|
|
626f262121 | ||
|
|
7caf492ae2 | ||
|
|
9aa2ab1657 | ||
|
|
971b774282 | ||
|
|
1377759705 | ||
|
|
d56bafa6d0 | ||
|
|
6ec6c9bb83 | ||
|
|
8a2a5eecdd | ||
|
|
08154b4374 | ||
|
|
880097acd5 | ||
|
|
e02615c93d | ||
|
|
e9259e680e | ||
|
|
a5b85a3d6b | ||
|
|
82c323c2d9 |
@@ -20,7 +20,7 @@ jobs:
|
|||||||
strategy:
|
strategy:
|
||||||
fail-fast: false
|
fail-fast: false
|
||||||
matrix:
|
matrix:
|
||||||
os: ${{ fromJSON('["ubuntu-latest","windows-latest"]') }}
|
os: ${{ github.event_name == 'pull_request' && fromJSON('["ubuntu-latest"]') || fromJSON('["ubuntu-latest","windows-latest"]') }}
|
||||||
# CI concentrates on newer runtimes (3.11/3.12 still supported per pyproject requires-python).
|
# CI concentrates on newer runtimes (3.11/3.12 still supported per pyproject requires-python).
|
||||||
python-version: ${{ fromJSON('["3.13","3.14"]') }}
|
python-version: ${{ fromJSON('["3.13","3.14"]') }}
|
||||||
|
|
||||||
|
|||||||
@@ -97,4 +97,3 @@ logs/
|
|||||||
tmp/
|
tmp/
|
||||||
temp/
|
temp/
|
||||||
*.tmp
|
*.tmp
|
||||||
exp/
|
|
||||||
|
|||||||
@@ -1,18 +1,6 @@
|
|||||||

|

|
||||||
|
|
||||||
<div align="center">
|
<div align="center">
|
||||||
<p>
|
|
||||||
<a href="https://nanobot.wiki/docs/latest/getting-started/nanobot-overview">English</a> |
|
|
||||||
<a href="https://nanobot.wiki/cn/docs/latest/getting-started/nanobot-overview">简体中文</a> |
|
|
||||||
<a href="https://nanobot.wiki/zh-Hant/docs/latest/getting-started/nanobot-overview">繁體中文</a> |
|
|
||||||
<a href="https://nanobot.wiki/es/docs/latest/getting-started/nanobot-overview">Español</a> |
|
|
||||||
<a href="https://nanobot.wiki/fr/docs/latest/getting-started/nanobot-overview">Français</a> |
|
|
||||||
<a href="https://nanobot.wiki/id/docs/latest/getting-started/nanobot-overview">Bahasa Indonesia</a> |
|
|
||||||
<a href="https://nanobot.wiki/ja/docs/latest/getting-started/nanobot-overview">日本語</a> |
|
|
||||||
<a href="https://nanobot.wiki/ko/docs/latest/getting-started/nanobot-overview">한국어</a> |
|
|
||||||
<a href="https://nanobot.wiki/ru/docs/latest/getting-started/nanobot-overview">Русский</a> |
|
|
||||||
<a href="https://nanobot.wiki/vi/docs/latest/getting-started/nanobot-overview">Tiếng Việt</a>
|
|
||||||
</p>
|
|
||||||
<p>
|
<p>
|
||||||
<a href="https://pypi.org/project/nanobot-ai/"><img src="https://img.shields.io/pypi/v/nanobot-ai" alt="PyPI"></a>
|
<a href="https://pypi.org/project/nanobot-ai/"><img src="https://img.shields.io/pypi/v/nanobot-ai" alt="PyPI"></a>
|
||||||
<a href="https://pepy.tech/project/nanobot-ai"><img src="https://static.pepy.tech/badge/nanobot-ai" alt="Downloads"></a>
|
<a href="https://pepy.tech/project/nanobot-ai"><img src="https://static.pepy.tech/badge/nanobot-ai" alt="Downloads"></a>
|
||||||
@@ -73,7 +61,7 @@
|
|||||||
- **2026-04-13** 🛡️ Agent turn hardened — user messages persisted early, auto-compact skips active tasks.
|
- **2026-04-13** 🛡️ Agent turn hardened — user messages persisted early, auto-compact skips active tasks.
|
||||||
- **2026-04-12** 🔒 Lark global domain support, Dream learns discovered skills, shell sandbox tightened.
|
- **2026-04-12** 🔒 Lark global domain support, Dream learns discovered skills, shell sandbox tightened.
|
||||||
- **2026-04-11** ⚡ Context compact shrinks sessions on the fly; Kagi web search; QQ & WeCom full media.
|
- **2026-04-11** ⚡ Context compact shrinks sessions on the fly; Kagi web search; QQ & WeCom full media.
|
||||||
- **2026-04-10** 📓 Multiple MCP servers, Feishu streaming & done-emoji.
|
- **2026-04-10** 📓 Notebook editing tool, multiple MCP servers, Feishu streaming & done-emoji.
|
||||||
- **2026-04-09** 🔌 WebSocket channel, unified cross-channel session, `disabled_skills` config.
|
- **2026-04-09** 🔌 WebSocket channel, unified cross-channel session, `disabled_skills` config.
|
||||||
- **2026-04-08** 📤 API file uploads, OpenAI reasoning auto-routing with Responses fallback.
|
- **2026-04-08** 📤 API file uploads, OpenAI reasoning auto-routing with Responses fallback.
|
||||||
- **2026-04-07** 🧠 Anthropic adaptive thinking, MCP resources & prompts exposed as tools.
|
- **2026-04-07** 🧠 Anthropic adaptive thinking, MCP resources & prompts exposed as tools.
|
||||||
@@ -224,7 +212,6 @@ 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)
|
||||||
|
|
||||||
@@ -342,4 +329,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>
|
||||||
+4
-75
@@ -134,7 +134,6 @@ ANTHROPIC_API_KEY="$(bw get password api/anthropic)" nanobot agent
|
|||||||
| `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) |
|
||||||
@@ -148,13 +147,11 @@ ANTHROPIC_API_KEY="$(bw get password api/anthropic)" nanobot agent
|
|||||||
| `gemini` | LLM (Gemini direct) | [aistudio.google.com](https://aistudio.google.com) |
|
| `gemini` | LLM (Gemini direct) | [aistudio.google.com](https://aistudio.google.com) |
|
||||||
| `aihubmix` | LLM (API gateway, access to all models) | [aihubmix.com](https://aihubmix.com) |
|
| `aihubmix` | LLM (API gateway, access to all models) | [aihubmix.com](https://aihubmix.com) |
|
||||||
| `siliconflow` | LLM (SiliconFlow/硅基流动) | [siliconflow.cn](https://siliconflow.cn) |
|
| `siliconflow` | LLM (SiliconFlow/硅基流动) | [siliconflow.cn](https://siliconflow.cn) |
|
||||||
| `novita` | LLM (Novita AI OpenAI-compatible gateway) | [novita.ai](https://novita.ai) |
|
|
||||||
| `dashscope` | LLM (Qwen) | [dashscope.console.aliyun.com](https://dashscope.console.aliyun.com) |
|
| `dashscope` | LLM (Qwen) | [dashscope.console.aliyun.com](https://dashscope.console.aliyun.com) |
|
||||||
| `moonshot` | LLM (Moonshot/Kimi) | [platform.moonshot.cn](https://platform.moonshot.cn) |
|
| `moonshot` | LLM (Moonshot/Kimi) | [platform.moonshot.cn](https://platform.moonshot.cn) |
|
||||||
| `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/)) | — |
|
||||||
@@ -166,36 +163,6 @@ ANTHROPIC_API_KEY="$(bw get password api/anthropic)" nanobot agent
|
|||||||
| `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>
|
||||||
|
|
||||||
@@ -477,34 +444,6 @@ 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>
|
||||||
|
|
||||||
@@ -573,8 +512,6 @@ 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>
|
||||||
|
|
||||||
@@ -640,19 +577,12 @@ 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`). Use it when you want to run nanobot against a model on your own machine instead of a hosted API provider.
|
[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.
|
||||||
|
|
||||||
**1. Start Atomic Chat**
|
**1. Add to config** (partial — merge into `~/.nanobot/config.json`):
|
||||||
|
|
||||||
- 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
|
||||||
{
|
{
|
||||||
@@ -665,13 +595,13 @@ ollama run llama3.2
|
|||||||
"agents": {
|
"agents": {
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"provider": "atomic_chat",
|
"provider": "atomic_chat",
|
||||||
"model": "qwen3-32b"
|
"model": "your-model-id-from-atomic-chat"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
> **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.
|
> **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.
|
||||||
|
|
||||||
> `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.
|
||||||
|
|
||||||
@@ -752,7 +682,6 @@ 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>
|
||||||
|
|
||||||
|
|||||||
+49
-103
@@ -6,6 +6,8 @@ The feature is disabled by default. Enable it in `~/.nanobot/config.json`, confi
|
|||||||
|
|
||||||
## Quick Setup
|
## Quick Setup
|
||||||
|
|
||||||
|
OpenRouter example:
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"providers": {
|
"providers": {
|
||||||
@@ -17,13 +19,56 @@ The feature is disabled by default. Enable it in `~/.nanobot/config.json`, confi
|
|||||||
"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"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
See [Provider Notes](#provider-notes) for AIHubMix, MiniMax, Gemini, Ollama, and StepFun configuration examples.
|
AIHubMix example:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"aihubmix": {
|
||||||
|
"apiKey": "${AIHUBMIX_API_KEY}"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "aihubmix",
|
||||||
|
"model": "gpt-image-2-free",
|
||||||
|
"defaultAspectRatio": "1:1",
|
||||||
|
"defaultImageSize": "1K"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Gemini example (Imagen 4):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"gemini": {
|
||||||
|
"apiKey": "${GEMINI_API_KEY}"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "gemini",
|
||||||
|
"model": "imagen-4.0-generate-001",
|
||||||
|
"defaultAspectRatio": "1:1"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
For Gemini Flash (which supports reference-image edits) see the [Gemini](#gemini) section below.
|
||||||
|
|
||||||
> [!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.
|
||||||
@@ -46,7 +91,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. Supported values: `openrouter`, `aihubmix`, `minimax`, `gemini`, `ollama`, `stepfun` |
|
| `tools.imageGeneration.provider` | string | `"openrouter"` | Image provider name. Supported values: `openrouter`, `aihubmix`, `gemini` |
|
||||||
| `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` |
|
||||||
@@ -116,28 +161,6 @@ 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
|
### Gemini
|
||||||
|
|
||||||
nanobot supports two Gemini image generation model families via Google's Generative Language API:
|
nanobot supports two Gemini image generation model families via Google's Generative Language API:
|
||||||
@@ -168,83 +191,6 @@ For reference-image edits, use a Gemini Flash image model:
|
|||||||
|
|
||||||
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).
|
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).
|
||||||
|
|
||||||
### Ollama
|
|
||||||
|
|
||||||
Ollama's experimental native image generation API works with local servers and hosted ollama.com models. Local access at `http://localhost:11434/api` does not require an API key; set `providers.ollama.apiKey` only when targeting `https://ollama.com/api`.
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"providers": {
|
|
||||||
"ollama": {
|
|
||||||
"apiBase": "http://localhost:11434/api"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"tools": {
|
|
||||||
"imageGeneration": {
|
|
||||||
"enabled": true,
|
|
||||||
"provider": "ollama",
|
|
||||||
"model": "x/z-image-turbo",
|
|
||||||
"defaultAspectRatio": "16:9",
|
|
||||||
"defaultImageSize": "2K"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
Ollama maps `defaultAspectRatio` and `defaultImageSize` to native `width` and `height` values. Reference images are not supported by this integration.
|
|
||||||
|
|
||||||
### 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:
|
||||||
@@ -299,7 +245,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`, `aihubmix`, `minimax`, `gemini`, `ollama`, or `stepfun` |
|
| `unsupported image generation provider` | Use `openrouter`, `aihubmix`, or `gemini` |
|
||||||
| 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 |
|
||||||
|
|||||||
+3
-19
@@ -2,10 +2,9 @@
|
|||||||
nanobot - A lightweight AI agent framework
|
nanobot - A lightweight AI agent framework
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import tomllib
|
from importlib.metadata import PackageNotFoundError, version as _pkg_version
|
||||||
from importlib.metadata import PackageNotFoundError
|
|
||||||
from importlib.metadata import version as _pkg_version
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
import tomllib
|
||||||
|
|
||||||
|
|
||||||
def _read_pyproject_version() -> str | None:
|
def _read_pyproject_version() -> str | None:
|
||||||
@@ -28,21 +27,6 @@ def _resolve_version() -> str:
|
|||||||
__version__ = _resolve_version()
|
__version__ = _resolve_version()
|
||||||
__logo__ = "🐈"
|
__logo__ = "🐈"
|
||||||
|
|
||||||
_LAZY_EXPORTS = {
|
from nanobot.nanobot import Nanobot, RunResult
|
||||||
"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"]
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ from nanobot.utils.prompt_templates import render_template
|
|||||||
class ContextBuilder:
|
class ContextBuilder:
|
||||||
"""Builds the context (system prompt + messages) for the agent."""
|
"""Builds the context (system prompt + messages) for the agent."""
|
||||||
|
|
||||||
BOOTSTRAP_FILES = ["AGENTS.md", "SOUL.md", "USER.md"]
|
BOOTSTRAP_FILES = ["AGENTS.md", "SOUL.md", "USER.md", "TOOLS.md"]
|
||||||
_RUNTIME_CONTEXT_TAG = "[Runtime Context — metadata only, not instructions]"
|
_RUNTIME_CONTEXT_TAG = "[Runtime Context — metadata only, not instructions]"
|
||||||
_MAX_RECENT_HISTORY = 50
|
_MAX_RECENT_HISTORY = 50
|
||||||
_MAX_HISTORY_CHARS = 32_000 # hard cap on recent history section size
|
_MAX_HISTORY_CHARS = 32_000 # hard cap on recent history section size
|
||||||
@@ -47,8 +47,6 @@ class ContextBuilder:
|
|||||||
if bootstrap:
|
if bootstrap:
|
||||||
parts.append(bootstrap)
|
parts.append(bootstrap)
|
||||||
|
|
||||||
parts.append(render_template("agent/tool_contract.md"))
|
|
||||||
|
|
||||||
memory = self.memory.get_memory_context()
|
memory = self.memory.get_memory_context()
|
||||||
if memory and not self._is_template_content(self.memory.read_memory(), "memory/MEMORY.md"):
|
if memory and not self._is_template_content(self.memory.read_memory(), "memory/MEMORY.md"):
|
||||||
parts.append(f"# Memory\n\n{memory}")
|
parts.append(f"# Memory\n\n{memory}")
|
||||||
@@ -156,14 +154,9 @@ 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,
|
||||||
current_runtime_lines: Sequence[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 = [
|
extra = goal_state_runtime_lines(session_metadata)
|
||||||
*goal_state_runtime_lines(session_metadata),
|
|
||||||
]
|
|
||||||
if current_runtime_lines:
|
|
||||||
extra.extend(line for line in current_runtime_lines if line)
|
|
||||||
runtime_ctx = self._build_runtime_context(
|
runtime_ctx = self._build_runtime_context(
|
||||||
channel,
|
channel,
|
||||||
chat_id,
|
chat_id,
|
||||||
@@ -217,3 +210,4 @@ class ContextBuilder:
|
|||||||
if not images:
|
if not images:
|
||||||
return text
|
return text
|
||||||
return images + [{"type": "text", "text": text}]
|
return images + [{"type": "text", "text": text}]
|
||||||
|
|
||||||
|
|||||||
+20
-9
@@ -28,7 +28,6 @@ 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
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.cli_apps import utils as cli_app_utils
|
|
||||||
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
|
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
|
||||||
from nanobot.config.schema import AgentDefaults, ModelPresetConfig
|
from nanobot.config.schema import AgentDefaults, ModelPresetConfig
|
||||||
from nanobot.providers.base import LLMProvider
|
from nanobot.providers.base import LLMProvider
|
||||||
@@ -37,17 +36,19 @@ from nanobot.session.goal_state import (
|
|||||||
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.session.webui_turns import (
|
from nanobot.utils.artifacts import generated_image_paths_from_messages
|
||||||
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.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_turn_helpers import (
|
||||||
|
WebuiTurnCoordinator,
|
||||||
|
build_bus_progress_callback,
|
||||||
|
mark_webui_session,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nanobot.config.schema import (
|
from nanobot.config.schema import (
|
||||||
@@ -60,6 +61,7 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
UNIFIED_SESSION_KEY = "unified:default"
|
UNIFIED_SESSION_KEY = "unified:default"
|
||||||
|
|
||||||
|
|
||||||
class TurnState(Enum):
|
class TurnState(Enum):
|
||||||
RESTORE = auto()
|
RESTORE = auto()
|
||||||
COMPACT = auto()
|
COMPACT = auto()
|
||||||
@@ -101,6 +103,7 @@ 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
|
||||||
@@ -568,7 +571,7 @@ class AgentLoop:
|
|||||||
media_paths = [p for p in (msg.media or []) if isinstance(p, str) and p]
|
media_paths = [p for p in (msg.media or []) if isinstance(p, str) and p]
|
||||||
has_text = isinstance(msg.content, str) and msg.content.strip()
|
has_text = isinstance(msg.content, str) and msg.content.strip()
|
||||||
if has_text or media_paths:
|
if has_text or media_paths:
|
||||||
extra: dict[str, Any] = ({"media": list(media_paths)} if media_paths else {}) | cli_app_utils.session_extra(msg.metadata)
|
extra: dict[str, Any] = {"media": list(media_paths)} if media_paths else {}
|
||||||
extra.update(kwargs)
|
extra.update(kwargs)
|
||||||
text = msg.content if isinstance(msg.content, str) else ""
|
text = msg.content if isinstance(msg.content, str) else ""
|
||||||
session.add_message("user", text, **extra)
|
session.add_message("user", text, **extra)
|
||||||
@@ -593,7 +596,7 @@ class AgentLoop:
|
|||||||
chat_id=self._runtime_chat_id(msg),
|
chat_id=self._runtime_chat_id(msg),
|
||||||
sender_id=msg.sender_id,
|
sender_id=msg.sender_id,
|
||||||
session_summary=pending_summary,
|
session_summary=pending_summary,
|
||||||
session_metadata=session.metadata, current_runtime_lines=cli_app_utils.runtime_lines(msg, self.context.workspace),
|
session_metadata=session.metadata,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _dispatch_command_inline(
|
async def _dispatch_command_inline(
|
||||||
@@ -1058,7 +1061,7 @@ class AgentLoop:
|
|||||||
current_role=current_role,
|
current_role=current_role,
|
||||||
sender_id=msg.sender_id,
|
sender_id=msg.sender_id,
|
||||||
session_summary=pending,
|
session_summary=pending,
|
||||||
session_metadata=session.metadata, current_runtime_lines=cli_app_utils.runtime_lines(msg, self.context.workspace, skip=is_subagent),
|
session_metadata=session.metadata,
|
||||||
)
|
)
|
||||||
t_wall = time.time()
|
t_wall = time.time()
|
||||||
final_content, _, all_msgs, stop_reason, _ = await self._run_agent_loop(
|
final_content, _, all_msgs, stop_reason, _ = await self._run_agent_loop(
|
||||||
@@ -1191,6 +1194,7 @@ 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,
|
||||||
@@ -1214,6 +1218,7 @@ 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,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1343,6 +1348,11 @@ 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(
|
||||||
@@ -1370,6 +1380,7 @@ 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,
|
||||||
)
|
)
|
||||||
|
|||||||
+13
-55
@@ -19,9 +19,7 @@ from nanobot.utils.file_edit_events import (
|
|||||||
build_file_edit_end_event,
|
build_file_edit_end_event,
|
||||||
build_file_edit_error_event,
|
build_file_edit_error_event,
|
||||||
build_file_edit_start_event,
|
build_file_edit_start_event,
|
||||||
prepare_file_edit_tracker as _prepare_file_edit_tracker,
|
prepare_file_edit_tracker,
|
||||||
prepare_file_edit_trackers,
|
|
||||||
StreamingFileEditTracker,
|
|
||||||
)
|
)
|
||||||
from nanobot.utils.helpers import (
|
from nanobot.utils.helpers import (
|
||||||
IncrementalThinkExtractor,
|
IncrementalThinkExtractor,
|
||||||
@@ -59,14 +57,11 @@ _SNIP_SAFETY_BUFFER = 1024
|
|||||||
_MICROCOMPACT_KEEP_RECENT = 10
|
_MICROCOMPACT_KEEP_RECENT = 10
|
||||||
_MICROCOMPACT_MIN_CHARS = 500
|
_MICROCOMPACT_MIN_CHARS = 500
|
||||||
_COMPACTABLE_TOOLS = frozenset({
|
_COMPACTABLE_TOOLS = frozenset({
|
||||||
"read_file", "exec", "grep", "find_files",
|
"read_file", "exec", "grep",
|
||||||
"web_search", "web_fetch", "list_dir", "list_exec_sessions",
|
"web_search", "web_fetch", "list_dir",
|
||||||
})
|
})
|
||||||
_BACKFILL_CONTENT = "[Tool result unavailable — call was interrupted or lost]"
|
_BACKFILL_CONTENT = "[Tool result unavailable — call was interrupted or lost]"
|
||||||
|
|
||||||
# Backward-compatible module attribute for tests/extensions that monkeypatch
|
|
||||||
# the former single-file tracker hook. Runtime uses prepare_file_edit_trackers.
|
|
||||||
prepare_file_edit_tracker = _prepare_file_edit_tracker
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
@@ -634,24 +629,6 @@ 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:
|
||||||
@@ -669,7 +646,6 @@ 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 = ""
|
||||||
@@ -699,7 +675,6 @@ 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)
|
||||||
@@ -714,14 +689,6 @@ 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(
|
||||||
@@ -861,8 +828,8 @@ class AgentRunner:
|
|||||||
and on_progress_accepts_file_edit_events(spec.progress_callback)
|
and on_progress_accepts_file_edit_events(spec.progress_callback)
|
||||||
)
|
)
|
||||||
progress_callback = spec.progress_callback if emit_file_edit_events else None
|
progress_callback = spec.progress_callback if emit_file_edit_events else None
|
||||||
file_edit_trackers = (
|
file_edit_tracker = (
|
||||||
prepare_file_edit_trackers(
|
prepare_file_edit_tracker(
|
||||||
call_id=tool_call.id,
|
call_id=tool_call.id,
|
||||||
tool_name=tool_call.name,
|
tool_name=tool_call.name,
|
||||||
tool=tool,
|
tool=tool,
|
||||||
@@ -872,13 +839,13 @@ class AgentRunner:
|
|||||||
if progress_callback is not None
|
if progress_callback is not None
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
if file_edit_trackers and progress_callback is not None:
|
if file_edit_tracker is not None and progress_callback is not None:
|
||||||
await invoke_file_edit_progress(
|
await invoke_file_edit_progress(
|
||||||
progress_callback,
|
progress_callback,
|
||||||
[build_file_edit_start_event(
|
[build_file_edit_start_event(
|
||||||
file_edit_tracker,
|
file_edit_tracker,
|
||||||
params if isinstance(params, dict) else None,
|
params if isinstance(params, dict) else None,
|
||||||
) for file_edit_tracker in file_edit_trackers],
|
)],
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
if tool is not None:
|
if tool is not None:
|
||||||
@@ -888,13 +855,10 @@ class AgentRunner:
|
|||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
raise
|
raise
|
||||||
except BaseException as exc:
|
except BaseException as exc:
|
||||||
if file_edit_trackers and progress_callback is not None:
|
if file_edit_tracker is not None and progress_callback is not None:
|
||||||
await invoke_file_edit_progress(
|
await invoke_file_edit_progress(
|
||||||
progress_callback,
|
progress_callback,
|
||||||
[
|
[build_file_edit_error_event(file_edit_tracker, str(exc))],
|
||||||
build_file_edit_error_event(file_edit_tracker, str(exc))
|
|
||||||
for file_edit_tracker in file_edit_trackers
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
event = {
|
event = {
|
||||||
"name": tool_call.name,
|
"name": tool_call.name,
|
||||||
@@ -917,13 +881,10 @@ 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_trackers and progress_callback is not None:
|
if file_edit_tracker is not None and progress_callback is not None:
|
||||||
await invoke_file_edit_progress(
|
await invoke_file_edit_progress(
|
||||||
progress_callback,
|
progress_callback,
|
||||||
[
|
[build_file_edit_error_event(file_edit_tracker, result)],
|
||||||
build_file_edit_error_event(file_edit_tracker, result)
|
|
||||||
for file_edit_tracker in file_edit_trackers
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
event = {
|
event = {
|
||||||
"name": tool_call.name,
|
"name": tool_call.name,
|
||||||
@@ -943,13 +904,10 @@ 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_trackers and progress_callback is not None:
|
if file_edit_tracker is not None and progress_callback is not None:
|
||||||
await invoke_file_edit_progress(
|
await invoke_file_edit_progress(
|
||||||
progress_callback,
|
progress_callback,
|
||||||
[build_file_edit_end_event(
|
[build_file_edit_end_event(file_edit_tracker)],
|
||||||
file_edit_tracker,
|
|
||||||
params if isinstance(params, dict) else None,
|
|
||||||
) for file_edit_tracker in file_edit_trackers],
|
|
||||||
)
|
)
|
||||||
|
|
||||||
detail = "" if result is None else str(result)
|
detail = "" if result is None else str(result)
|
||||||
|
|||||||
@@ -1,352 +0,0 @@
|
|||||||
"""Apply file edits by providing structured edit instructions."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import difflib
|
|
||||||
import re
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from nanobot.agent.tools.base import tool_parameters
|
|
||||||
from nanobot.agent.tools.filesystem import _FsTool
|
|
||||||
from nanobot.agent.tools.schema import (
|
|
||||||
ArraySchema,
|
|
||||||
BooleanSchema,
|
|
||||||
ObjectSchema,
|
|
||||||
StringSchema,
|
|
||||||
tool_parameters_schema,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
|
||||||
class _PatchSummary:
|
|
||||||
action: str
|
|
||||||
path: str
|
|
||||||
added: int = 0
|
|
||||||
deleted: int = 0
|
|
||||||
|
|
||||||
|
|
||||||
class _PatchError(ValueError):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
_ABSOLUTE_WINDOWS_RE = re.compile(r"^[A-Za-z]:[\\/]")
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_relative_path(path: str) -> str:
|
|
||||||
normalized = path.strip()
|
|
||||||
if not normalized:
|
|
||||||
raise _PatchError("patch path cannot be empty")
|
|
||||||
if "\0" in normalized:
|
|
||||||
raise _PatchError(f"patch path contains a null byte: {path!r}")
|
|
||||||
if normalized.startswith(("~", "/", "\\")) or _ABSOLUTE_WINDOWS_RE.match(normalized):
|
|
||||||
raise _PatchError(f"patch path must be relative: {path}")
|
|
||||||
if any(part == ".." for part in re.split(r"[\\/]+", normalized)):
|
|
||||||
raise _PatchError(f"patch path must not contain '..': {path}")
|
|
||||||
return normalized
|
|
||||||
|
|
||||||
|
|
||||||
def _lines_to_text(lines: list[str]) -> str:
|
|
||||||
if not lines:
|
|
||||||
return ""
|
|
||||||
return "\n".join(lines) + "\n"
|
|
||||||
|
|
||||||
|
|
||||||
def _text_line_count(text: str) -> int:
|
|
||||||
if not text:
|
|
||||||
return 0
|
|
||||||
return len(text.splitlines())
|
|
||||||
|
|
||||||
|
|
||||||
def _line_diff_stats(before: str, after: str) -> tuple[int, int]:
|
|
||||||
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 _format_summary(summary: _PatchSummary) -> str:
|
|
||||||
stats = ""
|
|
||||||
if summary.added or summary.deleted:
|
|
||||||
stats = f" (+{summary.added}/-{summary.deleted})"
|
|
||||||
return f"- {summary.action} {summary.path}{stats}"
|
|
||||||
|
|
||||||
|
|
||||||
@tool_parameters(
|
|
||||||
tool_parameters_schema(
|
|
||||||
edits=ArraySchema(
|
|
||||||
items=ObjectSchema(
|
|
||||||
path=StringSchema("Relative path to the file to edit."),
|
|
||||||
action=StringSchema(
|
|
||||||
"Operation type: replace (find and replace text), add (append new content or create file), delete (remove text).",
|
|
||||||
enum=["replace", "add", "delete"],
|
|
||||||
),
|
|
||||||
old_text=StringSchema(
|
|
||||||
"Exact text to search for in the file. Required for replace and delete.",
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
new_text=StringSchema(
|
|
||||||
"Text to replace with or append. Required for replace and add.",
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
required=["path", "action"],
|
|
||||||
),
|
|
||||||
description="List of edits to apply. Each edit specifies a file and the change to make.",
|
|
||||||
min_items=1,
|
|
||||||
max_items=20,
|
|
||||||
),
|
|
||||||
dry_run=BooleanSchema(
|
|
||||||
description="Validate and summarize the patch without writing files.",
|
|
||||||
default=False,
|
|
||||||
),
|
|
||||||
required=["edits"],
|
|
||||||
)
|
|
||||||
)
|
|
||||||
class ApplyPatchTool(_FsTool):
|
|
||||||
"""Apply file edits by providing structured edit instructions."""
|
|
||||||
_scopes = {"core", "subagent"}
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "apply_patch"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Default tool for code edits. Supports multi-file changes in a single call. "
|
|
||||||
"Provide a list of structured edits, each specifying a file path, action (replace/add/delete), and the text to change. "
|
|
||||||
"Paths must be relative. Set dry_run=true to validate and preview without writing files. "
|
|
||||||
"Use edit_file only for small exact replacements on a single file."
|
|
||||||
)
|
|
||||||
|
|
||||||
async def execute(
|
|
||||||
self,
|
|
||||||
edits: list[dict] | None = None,
|
|
||||||
dry_run: bool = False,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> str:
|
|
||||||
try:
|
|
||||||
if not edits:
|
|
||||||
raise _PatchError("must provide edits")
|
|
||||||
|
|
||||||
writes: dict[Path, str] = {}
|
|
||||||
deletes: set[Path] = set()
|
|
||||||
summaries: list[_PatchSummary] = []
|
|
||||||
|
|
||||||
for edit in edits:
|
|
||||||
if not isinstance(edit, dict):
|
|
||||||
raise _PatchError("each edit must be an object")
|
|
||||||
raw_path = edit.get("path")
|
|
||||||
if not isinstance(raw_path, str):
|
|
||||||
raise _PatchError("path required for edit")
|
|
||||||
path = _validate_relative_path(raw_path)
|
|
||||||
action = edit.get("action")
|
|
||||||
if not isinstance(action, str):
|
|
||||||
raise _PatchError(f"action required for edit: {path}")
|
|
||||||
source = self._resolve(path)
|
|
||||||
|
|
||||||
if action == "add":
|
|
||||||
new_text = edit.get("new_text")
|
|
||||||
if new_text is None:
|
|
||||||
raise _PatchError(f"new_text required for add: {path}")
|
|
||||||
|
|
||||||
pending = writes.get(source)
|
|
||||||
if pending is not None:
|
|
||||||
content = pending
|
|
||||||
exists = True
|
|
||||||
elif source.exists():
|
|
||||||
raw = source.read_bytes()
|
|
||||||
try:
|
|
||||||
content = raw.decode("utf-8")
|
|
||||||
except UnicodeDecodeError:
|
|
||||||
raise _PatchError(f"file is not UTF-8 text: {path}")
|
|
||||||
exists = True
|
|
||||||
else:
|
|
||||||
content = ""
|
|
||||||
exists = False
|
|
||||||
|
|
||||||
if exists:
|
|
||||||
uses_crlf = "\r\n" in content
|
|
||||||
new_norm = content.replace("\r\n", "\n") + new_text.replace("\r\n", "\n")
|
|
||||||
if new_norm and not new_norm.endswith("\n"):
|
|
||||||
new_norm += "\n"
|
|
||||||
if uses_crlf:
|
|
||||||
new_norm = new_norm.replace("\n", "\r\n")
|
|
||||||
writes[source] = new_norm
|
|
||||||
deletes.discard(source)
|
|
||||||
added, deleted = _line_diff_stats(content, new_norm)
|
|
||||||
action_name = "update"
|
|
||||||
else:
|
|
||||||
new_norm = new_text.replace("\r\n", "\n")
|
|
||||||
if new_norm and not new_norm.endswith("\n"):
|
|
||||||
new_norm += "\n"
|
|
||||||
writes[source] = new_norm
|
|
||||||
deletes.discard(source)
|
|
||||||
added = _text_line_count(new_norm)
|
|
||||||
deleted = 0
|
|
||||||
action_name = "add"
|
|
||||||
|
|
||||||
summaries.append(
|
|
||||||
_PatchSummary(
|
|
||||||
action=action_name, path=path, added=added, deleted=deleted
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
elif action == "replace":
|
|
||||||
old_text = edit.get("old_text") or ""
|
|
||||||
if not old_text:
|
|
||||||
raise _PatchError(f"old_text required for replace: {path}")
|
|
||||||
new_text = edit.get("new_text")
|
|
||||||
if new_text is None:
|
|
||||||
raise _PatchError(f"new_text required for replace: {path}")
|
|
||||||
|
|
||||||
pending = writes.get(source)
|
|
||||||
if pending is not None:
|
|
||||||
content = pending
|
|
||||||
elif source.exists():
|
|
||||||
raw = source.read_bytes()
|
|
||||||
try:
|
|
||||||
content = raw.decode("utf-8")
|
|
||||||
except UnicodeDecodeError:
|
|
||||||
raise _PatchError(f"file is not UTF-8 text: {path}")
|
|
||||||
else:
|
|
||||||
raise _PatchError(f"file to update does not exist: {path}")
|
|
||||||
|
|
||||||
if pending is None and not source.is_file():
|
|
||||||
raise _PatchError(f"path to update is not a file: {path}")
|
|
||||||
|
|
||||||
uses_crlf = "\r\n" in content
|
|
||||||
norm_content = content.replace("\r\n", "\n")
|
|
||||||
norm_old = old_text.replace("\r\n", "\n")
|
|
||||||
|
|
||||||
pos = norm_content.find(norm_old)
|
|
||||||
if pos < 0:
|
|
||||||
raise _PatchError(f"old_text not found in {path}")
|
|
||||||
if norm_content.find(norm_old, pos + 1) >= 0:
|
|
||||||
raise _PatchError(f"old_text appears multiple times in {path}")
|
|
||||||
|
|
||||||
new_norm = (
|
|
||||||
norm_content[:pos]
|
|
||||||
+ new_text.replace("\r\n", "\n")
|
|
||||||
+ norm_content[pos + len(norm_old) :]
|
|
||||||
)
|
|
||||||
if new_norm and not new_norm.endswith("\n"):
|
|
||||||
new_norm += "\n"
|
|
||||||
if uses_crlf:
|
|
||||||
new_norm = new_norm.replace("\n", "\r\n")
|
|
||||||
|
|
||||||
writes[source] = new_norm
|
|
||||||
deletes.discard(source)
|
|
||||||
added, deleted = _line_diff_stats(content, new_norm)
|
|
||||||
summaries.append(
|
|
||||||
_PatchSummary(
|
|
||||||
action="update", path=path, added=added, deleted=deleted
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
elif action == "delete":
|
|
||||||
old_text = edit.get("old_text") or ""
|
|
||||||
if not old_text:
|
|
||||||
raise _PatchError(f"old_text required for delete: {path}")
|
|
||||||
|
|
||||||
pending = writes.get(source)
|
|
||||||
if pending is not None:
|
|
||||||
content = pending
|
|
||||||
elif source.exists():
|
|
||||||
raw = source.read_bytes()
|
|
||||||
try:
|
|
||||||
content = raw.decode("utf-8")
|
|
||||||
except UnicodeDecodeError:
|
|
||||||
raise _PatchError(f"file is not UTF-8 text: {path}")
|
|
||||||
else:
|
|
||||||
raise _PatchError(f"file to update does not exist: {path}")
|
|
||||||
|
|
||||||
if pending is None and not source.is_file():
|
|
||||||
raise _PatchError(f"path to update is not a file: {path}")
|
|
||||||
|
|
||||||
uses_crlf = "\r\n" in content
|
|
||||||
norm_content = content.replace("\r\n", "\n")
|
|
||||||
norm_old = old_text.replace("\r\n", "\n")
|
|
||||||
|
|
||||||
pos = norm_content.find(norm_old)
|
|
||||||
if pos < 0:
|
|
||||||
raise _PatchError(f"old_text not found in {path}")
|
|
||||||
if norm_content.find(norm_old, pos + 1) >= 0:
|
|
||||||
raise _PatchError(f"old_text appears multiple times in {path}")
|
|
||||||
|
|
||||||
if norm_old == norm_content:
|
|
||||||
deletes.add(source)
|
|
||||||
writes.pop(source, None)
|
|
||||||
added, deleted = 0, _text_line_count(content)
|
|
||||||
summaries.append(
|
|
||||||
_PatchSummary(
|
|
||||||
action="delete", path=path, added=added, deleted=deleted
|
|
||||||
)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
new_norm = (
|
|
||||||
norm_content[:pos] + norm_content[pos + len(norm_old) :]
|
|
||||||
)
|
|
||||||
if new_norm and not new_norm.endswith("\n"):
|
|
||||||
new_norm += "\n"
|
|
||||||
if uses_crlf:
|
|
||||||
new_norm = new_norm.replace("\n", "\r\n")
|
|
||||||
writes[source] = new_norm
|
|
||||||
deletes.discard(source)
|
|
||||||
added, deleted = _line_diff_stats(content, new_norm)
|
|
||||||
summaries.append(
|
|
||||||
_PatchSummary(
|
|
||||||
action="update", path=path, added=added, deleted=deleted
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
else:
|
|
||||||
raise _PatchError(f"unknown action: {action}")
|
|
||||||
|
|
||||||
if dry_run:
|
|
||||||
return "Patch dry-run succeeded:\n" + "\n".join(
|
|
||||||
_format_summary(summary) for summary in summaries
|
|
||||||
)
|
|
||||||
|
|
||||||
backups: dict[Path, bytes | None] = {}
|
|
||||||
for path in set(writes) | deletes:
|
|
||||||
backups[path] = path.read_bytes() if path.exists() else None
|
|
||||||
|
|
||||||
try:
|
|
||||||
for path in deletes:
|
|
||||||
if path.exists():
|
|
||||||
path.unlink()
|
|
||||||
for path, content in writes.items():
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
path.write_text(content, encoding="utf-8", newline="")
|
|
||||||
except Exception:
|
|
||||||
for path, data in backups.items():
|
|
||||||
if data is None:
|
|
||||||
if path.exists():
|
|
||||||
path.unlink()
|
|
||||||
else:
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
path.write_bytes(data)
|
|
||||||
raise
|
|
||||||
|
|
||||||
for path in set(writes) | deletes:
|
|
||||||
self._file_states.record_write(path)
|
|
||||||
return "Patch applied:\n" + "\n".join(
|
|
||||||
_format_summary(summary) for summary in summaries
|
|
||||||
)
|
|
||||||
except PermissionError as exc:
|
|
||||||
return f"Error: {exc}"
|
|
||||||
except _PatchError as exc:
|
|
||||||
return f"Error applying patch: {exc}"
|
|
||||||
except Exception as exc:
|
|
||||||
return f"Error applying patch: {exc}"
|
|
||||||
@@ -1,127 +0,0 @@
|
|||||||
"""Controlled runner for installed CLI Apps."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from pydantic import Field
|
|
||||||
|
|
||||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
|
||||||
from nanobot.agent.tools.schema import ArraySchema, BooleanSchema, IntegerSchema, StringSchema, tool_parameters_schema
|
|
||||||
from nanobot.cli_apps import CliAppError, CliAppManager, CliAppsRuntimeConfig
|
|
||||||
from nanobot.config.schema import Base
|
|
||||||
|
|
||||||
|
|
||||||
class CliAppsToolConfig(Base):
|
|
||||||
"""CLI Apps tool configuration."""
|
|
||||||
|
|
||||||
enable: bool = True
|
|
||||||
install_timeout: int = Field(default=300, ge=1, le=3600)
|
|
||||||
run_timeout: int = Field(default=60, ge=1, le=600)
|
|
||||||
catalog_ttl_seconds: int = Field(default=3600, ge=60, le=86_400)
|
|
||||||
|
|
||||||
|
|
||||||
@tool_parameters(
|
|
||||||
tool_parameters_schema(
|
|
||||||
required=["name"],
|
|
||||||
name=StringSchema("Installed CLI app registry name, for example gimp, safari, or obsidian."),
|
|
||||||
args=ArraySchema(
|
|
||||||
StringSchema("One command-line argument."),
|
|
||||||
description="Arguments to pass to the CLI entry point. Do not include the entry point itself.",
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
json=BooleanSchema(
|
|
||||||
description="Whether to prepend --json when supported by the CLI.",
|
|
||||||
default=False,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
working_dir=StringSchema("Optional working directory for the CLI call.", nullable=True),
|
|
||||||
timeout=IntegerSchema(
|
|
||||||
description="Timeout in seconds for this CLI call.",
|
|
||||||
minimum=1,
|
|
||||||
maximum=600,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
class CliAppsTool(Tool):
|
|
||||||
"""Run an installed CLI-Anything or public CLI app through a controlled argv subprocess."""
|
|
||||||
|
|
||||||
config_key = "cli_apps"
|
|
||||||
_scopes = {"core", "subagent"}
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def config_cls(cls):
|
|
||||||
return CliAppsToolConfig
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def enabled(cls, ctx: Any) -> bool:
|
|
||||||
return ctx.config.cli_apps.enable
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def create(cls, ctx: Any) -> Tool:
|
|
||||||
cfg = ctx.config.cli_apps
|
|
||||||
return cls(
|
|
||||||
workspace=Path(ctx.workspace),
|
|
||||||
restrict_to_workspace=ctx.config.restrict_to_workspace,
|
|
||||||
runtime=CliAppsRuntimeConfig(
|
|
||||||
install_timeout=cfg.install_timeout,
|
|
||||||
run_timeout=cfg.run_timeout,
|
|
||||||
catalog_ttl_seconds=cfg.catalog_ttl_seconds,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
workspace: Path,
|
|
||||||
restrict_to_workspace: bool = False,
|
|
||||||
runtime: CliAppsRuntimeConfig | None = None,
|
|
||||||
) -> None:
|
|
||||||
self.workspace = workspace
|
|
||||||
self.restrict_to_workspace = restrict_to_workspace
|
|
||||||
self.runtime = runtime or CliAppsRuntimeConfig()
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "run_cli_app"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
try:
|
|
||||||
installed = CliAppManager(workspace=self.workspace, runtime=self.runtime).installed_names()
|
|
||||||
except Exception:
|
|
||||||
installed = []
|
|
||||||
installed_note = (
|
|
||||||
f" Installed Settings CLI Apps: {', '.join(installed)}."
|
|
||||||
if installed
|
|
||||||
else " No Settings CLI Apps are currently installed."
|
|
||||||
)
|
|
||||||
return (
|
|
||||||
"Run a CLI App that the user explicitly installed in Settings or attached as @app. "
|
|
||||||
"Do not use this for ordinary system CLIs such as git, gh, python, npm, or brew; "
|
|
||||||
"unknown names are rejected. Execution uses argv, not shell."
|
|
||||||
+ installed_note
|
|
||||||
)
|
|
||||||
|
|
||||||
async def execute(
|
|
||||||
self,
|
|
||||||
name: str,
|
|
||||||
args: list[str] | None = None,
|
|
||||||
json: bool | None = False,
|
|
||||||
working_dir: str | None = None,
|
|
||||||
timeout: int | None = None,
|
|
||||||
) -> str:
|
|
||||||
manager = CliAppManager(workspace=self.workspace, runtime=self.runtime)
|
|
||||||
try:
|
|
||||||
return manager.run(
|
|
||||||
name,
|
|
||||||
args=args or [],
|
|
||||||
json_output=bool(json),
|
|
||||||
working_dir=working_dir,
|
|
||||||
timeout=timeout,
|
|
||||||
restrict_to_workspace=self.restrict_to_workspace,
|
|
||||||
)
|
|
||||||
except CliAppError as exc:
|
|
||||||
return f"Error: {exc.message}"
|
|
||||||
@@ -1,591 +0,0 @@
|
|||||||
"""Session support for long-running exec workflows."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import shutil
|
|
||||||
import time
|
|
||||||
import uuid
|
|
||||||
from contextlib import suppress
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
|
||||||
from nanobot.agent.tools.schema import BooleanSchema, IntegerSchema, StringSchema, tool_parameters_schema
|
|
||||||
|
|
||||||
|
|
||||||
DEFAULT_YIELD_MS = 1000
|
|
||||||
MAX_YIELD_MS = 30_000
|
|
||||||
DEFAULT_WAIT_FOR_MS = 10_000
|
|
||||||
MAX_WAIT_FOR_MS = 120_000
|
|
||||||
DEFAULT_MAX_OUTPUT_CHARS = 10_000
|
|
||||||
MAX_OUTPUT_CHARS = 50_000
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
|
||||||
class _SessionPoll:
|
|
||||||
output: str
|
|
||||||
done: bool
|
|
||||||
exit_code: int | None
|
|
||||||
elapsed_s: float = 0.0
|
|
||||||
timed_out: bool = False
|
|
||||||
terminated: bool = False
|
|
||||||
stdin_closed: bool = False
|
|
||||||
truncated_chars: int = 0
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
|
||||||
class ExecSessionInfo:
|
|
||||||
session_id: str
|
|
||||||
command: str
|
|
||||||
cwd: str
|
|
||||||
elapsed_s: float
|
|
||||||
idle_s: float
|
|
||||||
remaining_s: float
|
|
||||||
returncode: int | None
|
|
||||||
|
|
||||||
|
|
||||||
class _ExecSession:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
session_id: str,
|
|
||||||
process: asyncio.subprocess.Process,
|
|
||||||
command: str,
|
|
||||||
cwd: str,
|
|
||||||
timeout: int,
|
|
||||||
) -> None:
|
|
||||||
self.session_id = session_id
|
|
||||||
self.process = process
|
|
||||||
self.command = command
|
|
||||||
self.cwd = cwd
|
|
||||||
self.started_at = time.monotonic()
|
|
||||||
self.deadline = time.monotonic() + timeout
|
|
||||||
self.last_access = time.monotonic()
|
|
||||||
self._chunks: list[str] = []
|
|
||||||
self._lock = asyncio.Lock()
|
|
||||||
self._timed_out = False
|
|
||||||
self._stdout_task = asyncio.create_task(self._read_stream(process.stdout, ""))
|
|
||||||
self._stderr_task = asyncio.create_task(self._read_stream(process.stderr, "STDERR:\n"))
|
|
||||||
|
|
||||||
async def _read_stream(
|
|
||||||
self,
|
|
||||||
stream: asyncio.StreamReader | None,
|
|
||||||
prefix: str,
|
|
||||||
) -> None:
|
|
||||||
if stream is None:
|
|
||||||
return
|
|
||||||
first = True
|
|
||||||
while True:
|
|
||||||
chunk = await stream.read(4096)
|
|
||||||
if not chunk:
|
|
||||||
break
|
|
||||||
text = chunk.decode("utf-8", errors="replace")
|
|
||||||
if prefix and first:
|
|
||||||
text = prefix + text
|
|
||||||
first = False
|
|
||||||
async with self._lock:
|
|
||||||
self._chunks.append(text)
|
|
||||||
|
|
||||||
async def write(self, chars: str) -> str | None:
|
|
||||||
if self.process.returncode is not None:
|
|
||||||
return "session has already exited"
|
|
||||||
if self.process.stdin is None:
|
|
||||||
return "session stdin is not available"
|
|
||||||
try:
|
|
||||||
self.process.stdin.write(chars.encode("utf-8"))
|
|
||||||
await self.process.stdin.drain()
|
|
||||||
except (BrokenPipeError, ConnectionResetError):
|
|
||||||
return "session stdin is closed"
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def close_stdin(self) -> str | None:
|
|
||||||
if self.process.returncode is not None:
|
|
||||||
return "session has already exited"
|
|
||||||
if self.process.stdin is None:
|
|
||||||
return "session stdin is not available"
|
|
||||||
self.process.stdin.close()
|
|
||||||
with suppress(BrokenPipeError, ConnectionResetError):
|
|
||||||
await self.process.stdin.wait_closed()
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def poll(
|
|
||||||
self,
|
|
||||||
yield_time_ms: int,
|
|
||||||
max_output_chars: int,
|
|
||||||
*,
|
|
||||||
terminated: bool = False,
|
|
||||||
stdin_closed: bool = False,
|
|
||||||
) -> _SessionPoll:
|
|
||||||
self.last_access = time.monotonic()
|
|
||||||
if yield_time_ms > 0 and self.process.returncode is None:
|
|
||||||
await asyncio.sleep(min(yield_time_ms, MAX_YIELD_MS) / 1000)
|
|
||||||
|
|
||||||
if self.process.returncode is None and time.monotonic() >= self.deadline:
|
|
||||||
self._timed_out = True
|
|
||||||
await self.kill()
|
|
||||||
|
|
||||||
if self.process.returncode is not None:
|
|
||||||
with suppress(asyncio.TimeoutError):
|
|
||||||
await asyncio.wait_for(
|
|
||||||
asyncio.gather(self._stdout_task, self._stderr_task),
|
|
||||||
timeout=2.0,
|
|
||||||
)
|
|
||||||
|
|
||||||
async with self._lock:
|
|
||||||
output = "".join(self._chunks)
|
|
||||||
self._chunks.clear()
|
|
||||||
|
|
||||||
output, truncated = _truncate_output(output, max_output_chars)
|
|
||||||
return _SessionPoll(
|
|
||||||
output=output,
|
|
||||||
done=self.process.returncode is not None,
|
|
||||||
exit_code=self.process.returncode,
|
|
||||||
elapsed_s=max(0.0, time.monotonic() - self.started_at),
|
|
||||||
timed_out=self._timed_out,
|
|
||||||
terminated=terminated,
|
|
||||||
stdin_closed=stdin_closed,
|
|
||||||
truncated_chars=truncated,
|
|
||||||
)
|
|
||||||
|
|
||||||
async def kill(self) -> None:
|
|
||||||
if self.process.returncode is not None:
|
|
||||||
return
|
|
||||||
self.process.kill()
|
|
||||||
with suppress(asyncio.TimeoutError):
|
|
||||||
await asyncio.wait_for(self.process.wait(), timeout=5.0)
|
|
||||||
|
|
||||||
|
|
||||||
class ExecSessionManager:
|
|
||||||
def __init__(self, *, max_sessions: int = 8, idle_timeout: int = 1800) -> None:
|
|
||||||
self.max_sessions = max_sessions
|
|
||||||
self.idle_timeout = idle_timeout
|
|
||||||
self._sessions: dict[str, _ExecSession] = {}
|
|
||||||
self._lock = asyncio.Lock()
|
|
||||||
|
|
||||||
async def start(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
command: str,
|
|
||||||
cwd: str,
|
|
||||||
env: dict[str, str],
|
|
||||||
timeout: int,
|
|
||||||
shell_program: str | None,
|
|
||||||
login: bool,
|
|
||||||
yield_time_ms: int,
|
|
||||||
max_output_chars: int,
|
|
||||||
) -> tuple[str, _SessionPoll]:
|
|
||||||
async with self._lock:
|
|
||||||
await self._cleanup_locked()
|
|
||||||
if len(self._sessions) >= self.max_sessions:
|
|
||||||
raise RuntimeError(f"maximum exec sessions reached ({self.max_sessions})")
|
|
||||||
process = await self._spawn(command, cwd, env, shell_program, login)
|
|
||||||
session_id = uuid.uuid4().hex[:12]
|
|
||||||
session = _ExecSession(
|
|
||||||
session_id=session_id,
|
|
||||||
process=process,
|
|
||||||
command=command,
|
|
||||||
cwd=cwd,
|
|
||||||
timeout=timeout,
|
|
||||||
)
|
|
||||||
self._sessions[session_id] = session
|
|
||||||
|
|
||||||
poll = await session.poll(yield_time_ms, max_output_chars)
|
|
||||||
if poll.done:
|
|
||||||
async with self._lock:
|
|
||||||
self._sessions.pop(session_id, None)
|
|
||||||
return session_id, poll
|
|
||||||
|
|
||||||
async def write(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
session_id: str,
|
|
||||||
chars: str | None,
|
|
||||||
close_stdin: bool,
|
|
||||||
terminate: bool,
|
|
||||||
yield_time_ms: int,
|
|
||||||
max_output_chars: int,
|
|
||||||
) -> _SessionPoll:
|
|
||||||
async with self._lock:
|
|
||||||
await self._cleanup_locked()
|
|
||||||
session = self._sessions.get(session_id)
|
|
||||||
if session is None:
|
|
||||||
raise KeyError(session_id)
|
|
||||||
|
|
||||||
if chars:
|
|
||||||
error = await session.write(chars)
|
|
||||||
if error:
|
|
||||||
raise RuntimeError(error)
|
|
||||||
stdin_closed = False
|
|
||||||
if close_stdin:
|
|
||||||
error = await session.close_stdin()
|
|
||||||
if error:
|
|
||||||
raise RuntimeError(error)
|
|
||||||
stdin_closed = True
|
|
||||||
if terminate:
|
|
||||||
await session.kill()
|
|
||||||
poll = await session.poll(
|
|
||||||
yield_time_ms,
|
|
||||||
max_output_chars,
|
|
||||||
terminated=terminate,
|
|
||||||
stdin_closed=stdin_closed,
|
|
||||||
)
|
|
||||||
if poll.done:
|
|
||||||
async with self._lock:
|
|
||||||
self._sessions.pop(session_id, None)
|
|
||||||
return poll
|
|
||||||
|
|
||||||
async def list(self) -> list[ExecSessionInfo]:
|
|
||||||
async with self._lock:
|
|
||||||
await self._cleanup_locked()
|
|
||||||
now = time.monotonic()
|
|
||||||
return [
|
|
||||||
ExecSessionInfo(
|
|
||||||
session_id=session_id,
|
|
||||||
command=session.command,
|
|
||||||
cwd=session.cwd,
|
|
||||||
elapsed_s=max(0.0, now - session.started_at),
|
|
||||||
idle_s=max(0.0, now - session.last_access),
|
|
||||||
remaining_s=max(0.0, session.deadline - now),
|
|
||||||
returncode=session.process.returncode,
|
|
||||||
)
|
|
||||||
for session_id, session in sorted(self._sessions.items())
|
|
||||||
]
|
|
||||||
|
|
||||||
async def _cleanup_locked(self) -> None:
|
|
||||||
now = time.monotonic()
|
|
||||||
stale = [
|
|
||||||
session_id
|
|
||||||
for session_id, session in self._sessions.items()
|
|
||||||
if now - session.last_access > self.idle_timeout
|
|
||||||
]
|
|
||||||
for session_id in stale:
|
|
||||||
session = self._sessions.pop(session_id)
|
|
||||||
await session.kill()
|
|
||||||
|
|
||||||
async def _spawn(
|
|
||||||
self,
|
|
||||||
command: str,
|
|
||||||
cwd: str,
|
|
||||||
env: dict[str, str],
|
|
||||||
shell_program: str | None,
|
|
||||||
login: bool,
|
|
||||||
) -> asyncio.subprocess.Process:
|
|
||||||
from nanobot.agent.tools import shell
|
|
||||||
|
|
||||||
if shell._IS_WINDOWS:
|
|
||||||
return await asyncio.create_subprocess_shell(
|
|
||||||
command,
|
|
||||||
stdin=asyncio.subprocess.PIPE,
|
|
||||||
stdout=asyncio.subprocess.PIPE,
|
|
||||||
stderr=asyncio.subprocess.PIPE,
|
|
||||||
cwd=cwd,
|
|
||||||
env=env,
|
|
||||||
)
|
|
||||||
shell_program = shell_program or shutil.which("bash") or "/bin/bash"
|
|
||||||
args = [shell_program]
|
|
||||||
if login and shell_program.rsplit("/", 1)[-1] in {"bash", "zsh"}:
|
|
||||||
args.append("-l")
|
|
||||||
args.extend(["-c", command])
|
|
||||||
return await asyncio.create_subprocess_exec(
|
|
||||||
*args,
|
|
||||||
stdin=asyncio.subprocess.PIPE,
|
|
||||||
stdout=asyncio.subprocess.PIPE,
|
|
||||||
stderr=asyncio.subprocess.PIPE,
|
|
||||||
cwd=cwd,
|
|
||||||
env=env,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
DEFAULT_EXEC_SESSION_MANAGER = ExecSessionManager()
|
|
||||||
|
|
||||||
|
|
||||||
def clamp_session_int(value: int | None, default: int, minimum: int, maximum: int) -> int:
|
|
||||||
if value is None:
|
|
||||||
return default
|
|
||||||
return min(max(value, minimum), maximum)
|
|
||||||
|
|
||||||
|
|
||||||
def _truncate_output(output: str, max_output_chars: int) -> tuple[str, int]:
|
|
||||||
if len(output) <= max_output_chars:
|
|
||||||
return output, 0
|
|
||||||
half = max_output_chars // 2
|
|
||||||
omitted = len(output) - max_output_chars
|
|
||||||
return (
|
|
||||||
output[:half]
|
|
||||||
+ f"\n\n... ({omitted:,} chars truncated) ...\n\n"
|
|
||||||
+ output[-half:],
|
|
||||||
omitted,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def format_session_poll(session_id: str, poll: _SessionPoll) -> str:
|
|
||||||
parts = [poll.output] if poll.output else []
|
|
||||||
if poll.truncated_chars:
|
|
||||||
parts.append(f"(output truncated by {poll.truncated_chars:,} chars)")
|
|
||||||
if poll.timed_out:
|
|
||||||
parts.append("Error: Command timed out; session was terminated.")
|
|
||||||
if poll.terminated and not poll.timed_out:
|
|
||||||
parts.append("Session terminated.")
|
|
||||||
if poll.stdin_closed:
|
|
||||||
parts.append("Stdin closed.")
|
|
||||||
if poll.done:
|
|
||||||
parts.append(f"Exit code: {poll.exit_code}")
|
|
||||||
else:
|
|
||||||
parts.append(f"Process running. session_id: {session_id}")
|
|
||||||
parts.append(f"Elapsed: {poll.elapsed_s:.1f}s")
|
|
||||||
return "\n".join(parts) if parts else "(no output yet)"
|
|
||||||
|
|
||||||
|
|
||||||
@tool_parameters(
|
|
||||||
tool_parameters_schema(
|
|
||||||
session_id=StringSchema("Session id returned by exec when yield_time_ms is used."),
|
|
||||||
chars=StringSchema(
|
|
||||||
"Bytes/text to write to stdin. Omit or pass an empty string to only poll recent output.",
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
close_stdin=BooleanSchema(
|
|
||||||
description="Close stdin after writing chars. Useful for commands waiting for EOF.",
|
|
||||||
default=False,
|
|
||||||
),
|
|
||||||
terminate=BooleanSchema(
|
|
||||||
description="Terminate the running exec session.",
|
|
||||||
default=False,
|
|
||||||
),
|
|
||||||
yield_time_ms=IntegerSchema(
|
|
||||||
DEFAULT_YIELD_MS,
|
|
||||||
description="Milliseconds to wait before returning recent output (default 1000, max 30000).",
|
|
||||||
minimum=0,
|
|
||||||
maximum=MAX_YIELD_MS,
|
|
||||||
),
|
|
||||||
wait_for=StringSchema(
|
|
||||||
"Optional text to wait for in output before returning. "
|
|
||||||
"Useful for interactive commands and dev servers.",
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
wait_timeout_ms=IntegerSchema(
|
|
||||||
DEFAULT_WAIT_FOR_MS,
|
|
||||||
description="Maximum milliseconds to wait for wait_for text (default 10000, max 120000).",
|
|
||||||
minimum=0,
|
|
||||||
maximum=MAX_WAIT_FOR_MS,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
max_output_chars=IntegerSchema(
|
|
||||||
DEFAULT_MAX_OUTPUT_CHARS,
|
|
||||||
description="Maximum output characters to return from this poll (default 10000, max 50000).",
|
|
||||||
minimum=1000,
|
|
||||||
maximum=MAX_OUTPUT_CHARS,
|
|
||||||
),
|
|
||||||
max_output_tokens=IntegerSchema(
|
|
||||||
DEFAULT_MAX_OUTPUT_CHARS,
|
|
||||||
description="Compatibility alias for max_output_chars. The current runtime uses a character budget.",
|
|
||||||
minimum=1000,
|
|
||||||
maximum=MAX_OUTPUT_CHARS,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
required=["session_id"],
|
|
||||||
)
|
|
||||||
)
|
|
||||||
class WriteStdinTool(Tool):
|
|
||||||
"""Write to or poll a running exec session."""
|
|
||||||
|
|
||||||
_scopes = {"core", "subagent"}
|
|
||||||
config_key = "exec"
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def config_cls(cls):
|
|
||||||
from nanobot.agent.tools.shell import ExecToolConfig
|
|
||||||
|
|
||||||
return ExecToolConfig
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def enabled(cls, ctx: Any) -> bool:
|
|
||||||
return ctx.config.exec.enable
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
manager: ExecSessionManager | None = None,
|
|
||||||
) -> None:
|
|
||||||
self._manager = manager or DEFAULT_EXEC_SESSION_MANAGER
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def create(cls, ctx: Any) -> Tool:
|
|
||||||
return cls()
|
|
||||||
|
|
||||||
@property
|
|
||||||
def exclusive(self) -> bool:
|
|
||||||
return True
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "write_stdin"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Interact with a running exec session created by exec with "
|
|
||||||
"yield_time_ms. Use chars='' to poll without writing, chars to send "
|
|
||||||
"stdin, close_stdin=true to send EOF, or terminate=true to stop the "
|
|
||||||
"process. Use wait_for with wait_timeout_ms for dev servers, test "
|
|
||||||
"watchers, and prompts where you need to wait for expected output. "
|
|
||||||
"Do not use this to start new commands; start them with exec."
|
|
||||||
)
|
|
||||||
|
|
||||||
async def execute(
|
|
||||||
self,
|
|
||||||
session_id: str,
|
|
||||||
chars: str | None = None,
|
|
||||||
close_stdin: bool = False,
|
|
||||||
terminate: bool = False,
|
|
||||||
yield_time_ms: int | None = None,
|
|
||||||
wait_for: str | None = None,
|
|
||||||
wait_timeout_ms: int | None = None,
|
|
||||||
max_output_chars: int | None = None,
|
|
||||||
max_output_tokens: int | None = None,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> str:
|
|
||||||
try:
|
|
||||||
if max_output_chars is None:
|
|
||||||
max_output_chars = max_output_tokens
|
|
||||||
output_limit = clamp_session_int(
|
|
||||||
max_output_chars,
|
|
||||||
DEFAULT_MAX_OUTPUT_CHARS,
|
|
||||||
1000,
|
|
||||||
MAX_OUTPUT_CHARS,
|
|
||||||
)
|
|
||||||
if wait_for:
|
|
||||||
return await self._wait_for_output(
|
|
||||||
session_id=session_id,
|
|
||||||
chars=chars,
|
|
||||||
close_stdin=close_stdin,
|
|
||||||
terminate=terminate,
|
|
||||||
wait_for=wait_for,
|
|
||||||
wait_timeout_ms=clamp_session_int(
|
|
||||||
wait_timeout_ms,
|
|
||||||
DEFAULT_WAIT_FOR_MS,
|
|
||||||
0,
|
|
||||||
MAX_WAIT_FOR_MS,
|
|
||||||
),
|
|
||||||
max_output_chars=output_limit,
|
|
||||||
)
|
|
||||||
poll = await self._manager.write(
|
|
||||||
session_id=session_id,
|
|
||||||
chars=chars,
|
|
||||||
close_stdin=close_stdin,
|
|
||||||
terminate=terminate,
|
|
||||||
yield_time_ms=clamp_session_int(yield_time_ms, DEFAULT_YIELD_MS, 0, MAX_YIELD_MS),
|
|
||||||
max_output_chars=output_limit,
|
|
||||||
)
|
|
||||||
return format_session_poll(session_id, poll)
|
|
||||||
except KeyError:
|
|
||||||
return f"Error: exec session not found: {session_id}"
|
|
||||||
except Exception as exc:
|
|
||||||
return f"Error writing to exec session: {exc}"
|
|
||||||
|
|
||||||
async def _wait_for_output(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
session_id: str,
|
|
||||||
chars: str | None,
|
|
||||||
close_stdin: bool,
|
|
||||||
terminate: bool,
|
|
||||||
wait_for: str,
|
|
||||||
wait_timeout_ms: int,
|
|
||||||
max_output_chars: int,
|
|
||||||
) -> str:
|
|
||||||
deadline = time.monotonic() + (wait_timeout_ms / 1000)
|
|
||||||
aggregate: list[str] = []
|
|
||||||
first = True
|
|
||||||
poll: _SessionPoll | None = None
|
|
||||||
|
|
||||||
while True:
|
|
||||||
remaining_ms = max(0, int((deadline - time.monotonic()) * 1000))
|
|
||||||
step_ms = min(500, remaining_ms)
|
|
||||||
poll = await self._manager.write(
|
|
||||||
session_id=session_id,
|
|
||||||
chars=chars if first else None,
|
|
||||||
close_stdin=close_stdin if first else False,
|
|
||||||
terminate=terminate if first else False,
|
|
||||||
yield_time_ms=step_ms,
|
|
||||||
max_output_chars=max_output_chars,
|
|
||||||
)
|
|
||||||
first = False
|
|
||||||
if poll.output:
|
|
||||||
aggregate.append(poll.output)
|
|
||||||
joined = "".join(aggregate)
|
|
||||||
if wait_for in joined:
|
|
||||||
poll.output = joined
|
|
||||||
return format_session_poll(session_id, poll)
|
|
||||||
if poll.done or remaining_ms <= 0:
|
|
||||||
poll.output = "".join(aggregate)
|
|
||||||
result = format_session_poll(session_id, poll)
|
|
||||||
if wait_for not in poll.output:
|
|
||||||
result += f"\nWait target not observed: {wait_for!r}"
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
@tool_parameters(tool_parameters_schema())
|
|
||||||
class ListExecSessionsTool(Tool):
|
|
||||||
"""List active exec sessions."""
|
|
||||||
|
|
||||||
_scopes = {"core", "subagent"}
|
|
||||||
config_key = "exec"
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def config_cls(cls):
|
|
||||||
from nanobot.agent.tools.shell import ExecToolConfig
|
|
||||||
|
|
||||||
return ExecToolConfig
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def enabled(cls, ctx: Any) -> bool:
|
|
||||||
return ctx.config.exec.enable
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
manager: ExecSessionManager | None = None,
|
|
||||||
) -> None:
|
|
||||||
self._manager = manager or DEFAULT_EXEC_SESSION_MANAGER
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def create(cls, ctx: Any) -> Tool:
|
|
||||||
return cls()
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "list_exec_sessions"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"List active long-running exec sessions, including session_id, cwd, "
|
|
||||||
"elapsed time, idle time, remaining timeout, and command preview. "
|
|
||||||
"Use this to recover a session_id after context shifts before "
|
|
||||||
"polling, writing stdin, or terminating with write_stdin."
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def read_only(self) -> bool:
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def execute(self, **kwargs: Any) -> str:
|
|
||||||
try:
|
|
||||||
sessions = await self._manager.list()
|
|
||||||
if not sessions:
|
|
||||||
return "No active exec sessions."
|
|
||||||
lines = []
|
|
||||||
for info in sessions:
|
|
||||||
command = " ".join(info.command.split())
|
|
||||||
if len(command) > 120:
|
|
||||||
command = command[:119] + "..."
|
|
||||||
status = "exited" if info.returncode is not None else "running"
|
|
||||||
lines.append(
|
|
||||||
f"{info.session_id} | {status} | elapsed={info.elapsed_s:.1f}s "
|
|
||||||
f"| idle={info.idle_s:.1f}s | remaining={info.remaining_s:.1f}s "
|
|
||||||
f"| cwd={info.cwd} | {command}"
|
|
||||||
)
|
|
||||||
return "\n".join(lines)
|
|
||||||
except Exception as exc:
|
|
||||||
return f"Error listing exec sessions: {exc}"
|
|
||||||
@@ -132,10 +132,6 @@ def _parse_page_range(pages: str, total: int) -> tuple[int, int]:
|
|||||||
minimum=1,
|
minimum=1,
|
||||||
),
|
),
|
||||||
pages=StringSchema("Page range for PDF files, e.g. '1-5' (default: all, max 20 pages)"),
|
pages=StringSchema("Page range for PDF files, e.g. '1-5' (default: all, max 20 pages)"),
|
||||||
force=BooleanSchema(
|
|
||||||
description="Bypass same-file read deduplication and return content again.",
|
|
||||||
default=False,
|
|
||||||
),
|
|
||||||
required=["path"],
|
required=["path"],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -158,11 +154,7 @@ class ReadFileTool(_FsTool):
|
|||||||
"Text output format: LINE_NUM|CONTENT. "
|
"Text output format: LINE_NUM|CONTENT. "
|
||||||
"Images return visual content for analysis. "
|
"Images return visual content for analysis. "
|
||||||
"Supports PDF, DOCX, XLSX, PPTX documents. "
|
"Supports PDF, DOCX, XLSX, PPTX documents. "
|
||||||
"Use find_files/list_dir first when the path is uncertain. "
|
|
||||||
"Read the relevant range before editing so replacements or patches "
|
|
||||||
"are based on current content. "
|
|
||||||
"Use offset and limit for large text files. "
|
"Use offset and limit for large text files. "
|
||||||
"Use force=true to re-read content even if unchanged. "
|
|
||||||
"Reads exceeding ~128K chars are truncated."
|
"Reads exceeding ~128K chars are truncated."
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -170,15 +162,7 @@ class ReadFileTool(_FsTool):
|
|||||||
def read_only(self) -> bool:
|
def read_only(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
async def execute(
|
async def execute(self, path: str | None = None, offset: int = 1, limit: int | None = None, pages: str | None = None, **kwargs: Any) -> Any:
|
||||||
self,
|
|
||||||
path: str | None = None,
|
|
||||||
offset: int = 1,
|
|
||||||
limit: int | None = None,
|
|
||||||
pages: str | None = None,
|
|
||||||
force: bool = False,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> Any:
|
|
||||||
try:
|
try:
|
||||||
if not path:
|
if not path:
|
||||||
return "Error reading file: Unknown path"
|
return "Error reading file: Unknown path"
|
||||||
@@ -218,13 +202,7 @@ class ReadFileTool(_FsTool):
|
|||||||
current_mtime = os.path.getmtime(fp)
|
current_mtime = os.path.getmtime(fp)
|
||||||
except OSError:
|
except OSError:
|
||||||
current_mtime = 0.0
|
current_mtime = 0.0
|
||||||
if (
|
if entry and entry.can_dedup and entry.offset == offset and entry.limit == limit:
|
||||||
not force
|
|
||||||
and entry
|
|
||||||
and entry.can_dedup
|
|
||||||
and entry.offset == offset
|
|
||||||
and entry.limit == limit
|
|
||||||
):
|
|
||||||
if current_mtime != entry.mtime:
|
if current_mtime != entry.mtime:
|
||||||
# File was modified externally - force full read and mark as not dedupable
|
# File was modified externally - force full read and mark as not dedupable
|
||||||
entry.can_dedup = False
|
entry.can_dedup = False
|
||||||
@@ -387,10 +365,9 @@ class WriteFileTool(_FsTool):
|
|||||||
@property
|
@property
|
||||||
def description(self) -> str:
|
def description(self) -> str:
|
||||||
return (
|
return (
|
||||||
"Create a new file or intentionally replace an entire file with "
|
"Write content to a file. Overwrites if the file already exists; "
|
||||||
"the provided content. Overwrites existing files and creates parent "
|
"creates parent directories as needed. "
|
||||||
"directories as needed. For code changes or partial edits, prefer "
|
"For partial edits, prefer edit_file instead."
|
||||||
"apply_patch; use edit_file only for small exact replacements."
|
|
||||||
)
|
)
|
||||||
|
|
||||||
async def execute(self, path: str | None = None, content: str | None = None, **kwargs: Any) -> str:
|
async def execute(self, path: str | None = None, content: str | None = None, **kwargs: Any) -> str:
|
||||||
@@ -680,24 +657,6 @@ def _find_match(content: str, old_text: str) -> tuple[str | None, int]:
|
|||||||
old_text=StringSchema("The text to find and replace"),
|
old_text=StringSchema("The text to find and replace"),
|
||||||
new_text=StringSchema("The text to replace with"),
|
new_text=StringSchema("The text to replace with"),
|
||||||
replace_all=BooleanSchema(description="Replace all occurrences (default false)"),
|
replace_all=BooleanSchema(description="Replace all occurrences (default false)"),
|
||||||
occurrence=IntegerSchema(
|
|
||||||
1,
|
|
||||||
description="Optional 1-based occurrence to replace when old_text appears multiple times.",
|
|
||||||
minimum=1,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
line_hint=IntegerSchema(
|
|
||||||
1,
|
|
||||||
description="Optional 1-based line hint used to choose the nearest match.",
|
|
||||||
minimum=1,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
expected_replacements=IntegerSchema(
|
|
||||||
1,
|
|
||||||
description="Optional guard for the number of replacements that must be made.",
|
|
||||||
minimum=1,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
required=["path", "old_text", "new_text"],
|
required=["path", "old_text", "new_text"],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -715,13 +674,10 @@ class EditFileTool(_FsTool):
|
|||||||
@property
|
@property
|
||||||
def description(self) -> str:
|
def description(self) -> str:
|
||||||
return (
|
return (
|
||||||
"Perform a small, exact replacement in one file by replacing "
|
"Edit a file by replacing old_text with new_text. "
|
||||||
"old_text with new_text. Use this for narrow text substitutions "
|
"Tolerates minor whitespace/indentation differences and curly/straight quote mismatches. "
|
||||||
"with old_text copied from read_file. For multi-file, structural, "
|
"If old_text matches multiple times, you must provide more context "
|
||||||
"or generated code edits, prefer apply_patch. If old_text matches "
|
"or set replace_all=true. Shows a diff of the closest match on failure."
|
||||||
"multiple times, provide more context or set occurrence, line_hint, "
|
|
||||||
"replace_all, and expected_replacements. Shows closest-match "
|
|
||||||
"diagnostics on failure."
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -732,8 +688,7 @@ class EditFileTool(_FsTool):
|
|||||||
async def execute(
|
async def execute(
|
||||||
self, path: str | None = None, old_text: str | None = None,
|
self, path: str | None = None, old_text: str | None = None,
|
||||||
new_text: str | None = None,
|
new_text: str | None = None,
|
||||||
replace_all: bool = False, occurrence: int | None = None,
|
replace_all: bool = False, **kwargs: Any,
|
||||||
line_hint: int | None = None, expected_replacements: int | None = None, **kwargs: Any,
|
|
||||||
) -> str:
|
) -> str:
|
||||||
try:
|
try:
|
||||||
if not path:
|
if not path:
|
||||||
@@ -742,12 +697,10 @@ class EditFileTool(_FsTool):
|
|||||||
raise ValueError("Unknown old_text")
|
raise ValueError("Unknown old_text")
|
||||||
if new_text is None:
|
if new_text is None:
|
||||||
raise ValueError("Unknown new_text")
|
raise ValueError("Unknown new_text")
|
||||||
if occurrence is not None and occurrence < 1:
|
|
||||||
return "Error: occurrence must be >= 1."
|
# .ipynb detection
|
||||||
if line_hint is not None and line_hint < 1:
|
if path.endswith(".ipynb"):
|
||||||
return "Error: line_hint must be >= 1."
|
return "Error: This is a Jupyter notebook. Use the notebook_edit tool instead of edit_file."
|
||||||
if expected_replacements is not None and expected_replacements < 1:
|
|
||||||
return "Error: expected_replacements must be >= 1."
|
|
||||||
|
|
||||||
fp = self._resolve(path)
|
fp = self._resolve(path)
|
||||||
|
|
||||||
@@ -790,42 +743,15 @@ class EditFileTool(_FsTool):
|
|||||||
if not matches:
|
if not matches:
|
||||||
return self._not_found_msg(old_text, content, path)
|
return self._not_found_msg(old_text, content, path)
|
||||||
count = len(matches)
|
count = len(matches)
|
||||||
if replace_all and occurrence is not None:
|
|
||||||
return "Error: occurrence cannot be used with replace_all=true."
|
|
||||||
if replace_all and line_hint is not None:
|
|
||||||
return "Error: line_hint cannot be used with replace_all=true."
|
|
||||||
if occurrence is not None and line_hint is not None:
|
|
||||||
return "Error: line_hint cannot be used with occurrence."
|
|
||||||
if count > 1 and not replace_all:
|
if count > 1 and not replace_all:
|
||||||
if occurrence is not None:
|
line_numbers = [match.line for match in matches]
|
||||||
if occurrence > count:
|
preview = ", ".join(f"line {n}" for n in line_numbers[:3])
|
||||||
return (
|
if len(line_numbers) > 3:
|
||||||
f"Error: occurrence {occurrence} is out of range; "
|
preview += ", ..."
|
||||||
f"old_text appears {count} times."
|
location_hint = f" at {preview}" if preview else ""
|
||||||
)
|
|
||||||
elif line_hint is not None:
|
|
||||||
nearest = min(matches, key=lambda match: abs(match.line - line_hint))
|
|
||||||
distance = abs(nearest.line - line_hint)
|
|
||||||
if sum(1 for match in matches if abs(match.line - line_hint) == distance) > 1:
|
|
||||||
return (
|
|
||||||
f"Error: line_hint {line_hint} is ambiguous; "
|
|
||||||
f"old_text appears {count} times."
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
line_numbers = [match.line for match in matches]
|
|
||||||
preview = ", ".join(f"line {n}" for n in line_numbers[:3])
|
|
||||||
if len(line_numbers) > 3:
|
|
||||||
preview += ", ..."
|
|
||||||
location_hint = f" at {preview}" if preview else ""
|
|
||||||
return (
|
|
||||||
f"Warning: old_text appears {count} times{location_hint}. "
|
|
||||||
"Provide more context, set occurrence to choose one match, "
|
|
||||||
"or set replace_all=true."
|
|
||||||
)
|
|
||||||
elif occurrence is not None and occurrence > count:
|
|
||||||
return (
|
return (
|
||||||
f"Error: occurrence {occurrence} is out of range; "
|
f"Warning: old_text appears {count} times{location_hint}. "
|
||||||
f"old_text appears {count} time."
|
"Provide more context to make it unique, or set replace_all=true."
|
||||||
)
|
)
|
||||||
|
|
||||||
norm_new = new_text.replace("\r\n", "\n")
|
norm_new = new_text.replace("\r\n", "\n")
|
||||||
@@ -834,17 +760,7 @@ class EditFileTool(_FsTool):
|
|||||||
if fp.suffix.lower() not in self._MARKDOWN_EXTS:
|
if fp.suffix.lower() not in self._MARKDOWN_EXTS:
|
||||||
norm_new = self._strip_trailing_ws(norm_new)
|
norm_new = self._strip_trailing_ws(norm_new)
|
||||||
|
|
||||||
if replace_all:
|
selected = matches if replace_all else matches[:1]
|
||||||
selected = matches
|
|
||||||
elif line_hint is not None:
|
|
||||||
selected = [min(matches, key=lambda match: abs(match.line - line_hint))]
|
|
||||||
else:
|
|
||||||
selected = [matches[occurrence - 1 if occurrence else 0]]
|
|
||||||
if expected_replacements is not None and len(selected) != expected_replacements:
|
|
||||||
return (
|
|
||||||
f"Error: expected {expected_replacements} replacements but "
|
|
||||||
f"would make {len(selected)}."
|
|
||||||
)
|
|
||||||
new_content = content
|
new_content = content
|
||||||
for match in reversed(selected):
|
for match in reversed(selected):
|
||||||
replacement = _preserve_quote_style(norm_old, match.text, norm_new)
|
replacement = _preserve_quote_style(norm_old, match.text, norm_new)
|
||||||
|
|||||||
@@ -17,9 +17,11 @@ 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,
|
||||||
|
GeminiImageGenerationClient,
|
||||||
ImageGenerationError,
|
ImageGenerationError,
|
||||||
ImageGenerationProvider,
|
MiniMaxImageGenerationClient,
|
||||||
get_image_gen_provider,
|
OpenRouterImageGenerationClient,
|
||||||
)
|
)
|
||||||
from nanobot.utils.artifacts import (
|
from nanobot.utils.artifacts import (
|
||||||
ArtifactError,
|
ArtifactError,
|
||||||
@@ -117,18 +119,37 @@ 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) -> ImageGenerationProvider | None:
|
def _provider_client(
|
||||||
|
self,
|
||||||
|
) -> OpenRouterImageGenerationClient | AIHubMixImageGenerationClient | MiniMaxImageGenerationClient | GeminiImageGenerationClient | 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,
|
||||||
}
|
}
|
||||||
return cls(**kwargs)
|
if self.config.provider == "openrouter":
|
||||||
|
return OpenRouterImageGenerationClient(**kwargs)
|
||||||
|
if self.config.provider == "aihubmix":
|
||||||
|
return AIHubMixImageGenerationClient(**kwargs)
|
||||||
|
if self.config.provider == "minimax":
|
||||||
|
return MiniMaxImageGenerationClient(**kwargs)
|
||||||
|
if self.config.provider == "gemini":
|
||||||
|
return GeminiImageGenerationClient(**kwargs)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _missing_api_key_error(self) -> str:
|
||||||
|
provider = self.config.provider
|
||||||
|
if provider == "openrouter":
|
||||||
|
return "Error: OpenRouter API key is not configured. Set providers.openrouter.apiKey."
|
||||||
|
if provider == "aihubmix":
|
||||||
|
return "Error: AIHubMix API key is not configured. Set providers.aihubmix.apiKey."
|
||||||
|
if provider == "minimax":
|
||||||
|
return "Error: MiniMax API key is not configured. Set providers.minimax.apiKey."
|
||||||
|
if provider == "gemini":
|
||||||
|
return "Error: Gemini API key is not configured. Set providers.gemini.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()
|
||||||
@@ -167,6 +188,9 @@ class ImageGenerationTool(Tool):
|
|||||||
client = self._provider_client()
|
client = self._provider_client()
|
||||||
if client is None:
|
if client is None:
|
||||||
return f"Error: unsupported image generation provider '{self.config.provider}'"
|
return f"Error: unsupported image generation provider '{self.config.provider}'"
|
||||||
|
provider = self._provider_config()
|
||||||
|
if not provider or not provider.api_key:
|
||||||
|
return self._missing_api_key_error()
|
||||||
|
|
||||||
requested = count or 1
|
requested = count or 1
|
||||||
if requested > self.config.max_images_per_turn:
|
if requested > self.config.max_images_per_turn:
|
||||||
|
|||||||
@@ -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. "
|
"Optional list of existing file paths to attach for proactive or cross-channel delivery. "
|
||||||
"Use artifact paths returned by generate_image here when delivering generated images."
|
"Do not use this to resend generate_image outputs in the current chat."
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
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, use the message tool "
|
"When generate_image creates images in the current chat, the final assistant reply "
|
||||||
"with the artifact paths in the media parameter to deliver the images to the user. "
|
"automatically attaches them; do not call message just to announce or resend them. "
|
||||||
"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."
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,162 @@
|
|||||||
|
"""NotebookEditTool — edit Jupyter .ipynb notebooks."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from nanobot.agent.tools.base import tool_parameters
|
||||||
|
from nanobot.agent.tools.schema import IntegerSchema, StringSchema, tool_parameters_schema
|
||||||
|
from nanobot.agent.tools.filesystem import _FsTool
|
||||||
|
|
||||||
|
|
||||||
|
def _new_cell(source: str, cell_type: str = "code", generate_id: bool = False) -> dict:
|
||||||
|
cell: dict[str, Any] = {
|
||||||
|
"cell_type": cell_type,
|
||||||
|
"source": source,
|
||||||
|
"metadata": {},
|
||||||
|
}
|
||||||
|
if cell_type == "code":
|
||||||
|
cell["outputs"] = []
|
||||||
|
cell["execution_count"] = None
|
||||||
|
if generate_id:
|
||||||
|
cell["id"] = uuid.uuid4().hex[:8]
|
||||||
|
return cell
|
||||||
|
|
||||||
|
|
||||||
|
def _make_empty_notebook() -> dict:
|
||||||
|
return {
|
||||||
|
"nbformat": 4,
|
||||||
|
"nbformat_minor": 5,
|
||||||
|
"metadata": {
|
||||||
|
"kernelspec": {"display_name": "Python 3", "language": "python", "name": "python3"},
|
||||||
|
"language_info": {"name": "python"},
|
||||||
|
},
|
||||||
|
"cells": [],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@tool_parameters(
|
||||||
|
tool_parameters_schema(
|
||||||
|
path=StringSchema("Path to the .ipynb notebook file"),
|
||||||
|
cell_index=IntegerSchema(0, description="0-based index of the cell to edit", minimum=0),
|
||||||
|
new_source=StringSchema("New source content for the cell"),
|
||||||
|
cell_type=StringSchema(
|
||||||
|
"Cell type: 'code' or 'markdown' (default: code)",
|
||||||
|
enum=["code", "markdown"],
|
||||||
|
),
|
||||||
|
edit_mode=StringSchema(
|
||||||
|
"Mode: 'replace' (default), 'insert' (after target), or 'delete'",
|
||||||
|
enum=["replace", "insert", "delete"],
|
||||||
|
),
|
||||||
|
required=["path", "cell_index"],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
class NotebookEditTool(_FsTool):
|
||||||
|
"""Edit Jupyter notebook cells: replace, insert, or delete."""
|
||||||
|
_scopes = {"core"}
|
||||||
|
|
||||||
|
_VALID_CELL_TYPES = frozenset({"code", "markdown"})
|
||||||
|
_VALID_EDIT_MODES = frozenset({"replace", "insert", "delete"})
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self) -> str:
|
||||||
|
return "notebook_edit"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def description(self) -> str:
|
||||||
|
return (
|
||||||
|
"Edit a Jupyter notebook (.ipynb) cell. "
|
||||||
|
"Modes: replace (default) replaces cell content, "
|
||||||
|
"insert adds a new cell after the target index, "
|
||||||
|
"delete removes the cell at the index. "
|
||||||
|
"cell_index is 0-based."
|
||||||
|
)
|
||||||
|
|
||||||
|
async def execute(
|
||||||
|
self,
|
||||||
|
path: str | None = None,
|
||||||
|
cell_index: int = 0,
|
||||||
|
new_source: str = "",
|
||||||
|
cell_type: str = "code",
|
||||||
|
edit_mode: str = "replace",
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> str:
|
||||||
|
try:
|
||||||
|
if not path:
|
||||||
|
return "Error: path is required"
|
||||||
|
|
||||||
|
if not path.endswith(".ipynb"):
|
||||||
|
return "Error: notebook_edit only works on .ipynb files. Use edit_file for other files."
|
||||||
|
|
||||||
|
if edit_mode not in self._VALID_EDIT_MODES:
|
||||||
|
return (
|
||||||
|
f"Error: Invalid edit_mode '{edit_mode}'. "
|
||||||
|
"Use one of: replace, insert, delete."
|
||||||
|
)
|
||||||
|
|
||||||
|
if cell_type not in self._VALID_CELL_TYPES:
|
||||||
|
return (
|
||||||
|
f"Error: Invalid cell_type '{cell_type}'. "
|
||||||
|
"Use one of: code, markdown."
|
||||||
|
)
|
||||||
|
|
||||||
|
fp = self._resolve(path)
|
||||||
|
|
||||||
|
# Create new notebook if file doesn't exist and mode is insert
|
||||||
|
if not fp.exists():
|
||||||
|
if edit_mode != "insert":
|
||||||
|
return f"Error: File not found: {path}"
|
||||||
|
nb = _make_empty_notebook()
|
||||||
|
cell = _new_cell(new_source, cell_type, generate_id=True)
|
||||||
|
nb["cells"].append(cell)
|
||||||
|
fp.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
fp.write_text(json.dumps(nb, indent=1, ensure_ascii=False), encoding="utf-8")
|
||||||
|
return f"Successfully created {fp} with 1 cell"
|
||||||
|
|
||||||
|
try:
|
||||||
|
nb = json.loads(fp.read_text(encoding="utf-8"))
|
||||||
|
except (json.JSONDecodeError, UnicodeDecodeError) as e:
|
||||||
|
return f"Error: Failed to parse notebook: {e}"
|
||||||
|
|
||||||
|
cells = nb.get("cells", [])
|
||||||
|
nbformat_minor = nb.get("nbformat_minor", 0)
|
||||||
|
generate_id = nb.get("nbformat", 0) >= 4 and nbformat_minor >= 5
|
||||||
|
|
||||||
|
if edit_mode == "delete":
|
||||||
|
if cell_index < 0 or cell_index >= len(cells):
|
||||||
|
return f"Error: cell_index {cell_index} out of range (notebook has {len(cells)} cells)"
|
||||||
|
cells.pop(cell_index)
|
||||||
|
nb["cells"] = cells
|
||||||
|
fp.write_text(json.dumps(nb, indent=1, ensure_ascii=False), encoding="utf-8")
|
||||||
|
return f"Successfully deleted cell {cell_index} from {fp}"
|
||||||
|
|
||||||
|
if edit_mode == "insert":
|
||||||
|
insert_at = min(cell_index + 1, len(cells))
|
||||||
|
cell = _new_cell(new_source, cell_type, generate_id=generate_id)
|
||||||
|
cells.insert(insert_at, cell)
|
||||||
|
nb["cells"] = cells
|
||||||
|
fp.write_text(json.dumps(nb, indent=1, ensure_ascii=False), encoding="utf-8")
|
||||||
|
return f"Successfully inserted cell at index {insert_at} in {fp}"
|
||||||
|
|
||||||
|
# Default: replace
|
||||||
|
if cell_index < 0 or cell_index >= len(cells):
|
||||||
|
return f"Error: cell_index {cell_index} out of range (notebook has {len(cells)} cells)"
|
||||||
|
cells[cell_index]["source"] = new_source
|
||||||
|
if cell_type and cells[cell_index].get("cell_type") != cell_type:
|
||||||
|
cells[cell_index]["cell_type"] = cell_type
|
||||||
|
if cell_type == "code":
|
||||||
|
cells[cell_index].setdefault("outputs", [])
|
||||||
|
cells[cell_index].setdefault("execution_count", None)
|
||||||
|
elif "outputs" in cells[cell_index]:
|
||||||
|
del cells[cell_index]["outputs"]
|
||||||
|
cells[cell_index].pop("execution_count", None)
|
||||||
|
nb["cells"] = cells
|
||||||
|
fp.write_text(json.dumps(nb, indent=1, ensure_ascii=False), encoding="utf-8")
|
||||||
|
return f"Successfully edited cell {cell_index} in {fp}"
|
||||||
|
|
||||||
|
except PermissionError as e:
|
||||||
|
return f"Error: {e}"
|
||||||
|
except Exception as e:
|
||||||
|
return f"Error editing notebook: {e}"
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Search tools: file discovery and grep."""
|
"""Search tools: grep."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -12,7 +12,6 @@ from typing import Any, Iterable, TypeVar
|
|||||||
from nanobot.agent.tools.filesystem import ListDirTool, _FsTool
|
from nanobot.agent.tools.filesystem import ListDirTool, _FsTool
|
||||||
|
|
||||||
_DEFAULT_HEAD_LIMIT = 250
|
_DEFAULT_HEAD_LIMIT = 250
|
||||||
_DEFAULT_FILE_HEAD_LIMIT = 200
|
|
||||||
T = TypeVar("T")
|
T = TypeVar("T")
|
||||||
_TYPE_GLOB_MAP = {
|
_TYPE_GLOB_MAP = {
|
||||||
"py": ("*.py", "*.pyi"),
|
"py": ("*.py", "*.pyi"),
|
||||||
@@ -89,14 +88,6 @@ def _matches_type(name: str, file_type: str | None) -> bool:
|
|||||||
return any(fnmatch.fnmatch(name.lower(), pattern.lower()) for pattern in patterns)
|
return any(fnmatch.fnmatch(name.lower(), pattern.lower()) for pattern in patterns)
|
||||||
|
|
||||||
|
|
||||||
def _matches_query(rel_path: str, query: str | None) -> bool:
|
|
||||||
if not query:
|
|
||||||
return True
|
|
||||||
haystack = rel_path.lower()
|
|
||||||
terms = [part for part in query.lower().split() if part]
|
|
||||||
return all(term in haystack for term in terms)
|
|
||||||
|
|
||||||
|
|
||||||
class _SearchTool(_FsTool):
|
class _SearchTool(_FsTool):
|
||||||
_IGNORE_DIRS = set(ListDirTool._IGNORE_DIRS)
|
_IGNORE_DIRS = set(ListDirTool._IGNORE_DIRS)
|
||||||
|
|
||||||
@@ -118,163 +109,6 @@ class _SearchTool(_FsTool):
|
|||||||
yield current / filename
|
yield current / filename
|
||||||
|
|
||||||
|
|
||||||
class FindFilesTool(_SearchTool):
|
|
||||||
"""Find files by path fragment, glob, or type."""
|
|
||||||
_scopes = {"core", "subagent"}
|
|
||||||
|
|
||||||
@property
|
|
||||||
def name(self) -> str:
|
|
||||||
return "find_files"
|
|
||||||
|
|
||||||
@property
|
|
||||||
def description(self) -> str:
|
|
||||||
return (
|
|
||||||
"Find files by path fragment, glob, or file type. "
|
|
||||||
"Use this before read_file when you need to locate files, and "
|
|
||||||
"prefer it over shell find/ls for ordinary workspace discovery. "
|
|
||||||
"Returns workspace-relative paths and skips common dependency/build "
|
|
||||||
"directories."
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def read_only(self) -> bool:
|
|
||||||
return True
|
|
||||||
|
|
||||||
@property
|
|
||||||
def parameters(self) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"path": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "Directory or file to search in (default '.')",
|
|
||||||
},
|
|
||||||
"query": {
|
|
||||||
"type": "string",
|
|
||||||
"description": (
|
|
||||||
"Optional case-insensitive path fragment search. "
|
|
||||||
"Whitespace-separated terms must all be present."
|
|
||||||
),
|
|
||||||
},
|
|
||||||
"glob": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "Optional file filter, e.g. '*.py' or 'tests/**/test_*.py'",
|
|
||||||
},
|
|
||||||
"type": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "Optional file type shorthand, e.g. 'py', 'ts', 'md', 'json'",
|
|
||||||
},
|
|
||||||
"include_dirs": {
|
|
||||||
"type": "boolean",
|
|
||||||
"description": "Include matching directories as well as files (default false)",
|
|
||||||
},
|
|
||||||
"sort": {
|
|
||||||
"type": "string",
|
|
||||||
"enum": ["path", "modified"],
|
|
||||||
"description": "Sort by path or most recently modified first (default path)",
|
|
||||||
},
|
|
||||||
"head_limit": {
|
|
||||||
"type": "integer",
|
|
||||||
"description": "Maximum number of paths to return (default 200, 0 for all, max 1000)",
|
|
||||||
"minimum": 0,
|
|
||||||
"maximum": 1000,
|
|
||||||
},
|
|
||||||
"offset": {
|
|
||||||
"type": "integer",
|
|
||||||
"description": "Skip the first N results before applying head_limit",
|
|
||||||
"minimum": 0,
|
|
||||||
"maximum": 100000,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
def _iter_paths(self, root: Path, *, include_dirs: bool) -> Iterable[Path]:
|
|
||||||
if root.is_file():
|
|
||||||
yield root
|
|
||||||
return
|
|
||||||
if include_dirs:
|
|
||||||
yield root
|
|
||||||
for dirpath, dirnames, filenames in os.walk(root):
|
|
||||||
dirnames[:] = sorted(d for d in dirnames if d not in self._IGNORE_DIRS)
|
|
||||||
current = Path(dirpath)
|
|
||||||
if include_dirs and current != root:
|
|
||||||
yield current
|
|
||||||
for filename in sorted(filenames):
|
|
||||||
yield current / filename
|
|
||||||
|
|
||||||
async def execute(
|
|
||||||
self,
|
|
||||||
path: str = ".",
|
|
||||||
query: str | None = None,
|
|
||||||
glob: str | None = None,
|
|
||||||
type: str | None = None,
|
|
||||||
include_dirs: bool = False,
|
|
||||||
sort: str = "path",
|
|
||||||
head_limit: int | None = None,
|
|
||||||
offset: int = 0,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> str:
|
|
||||||
try:
|
|
||||||
target = self._resolve(path or ".")
|
|
||||||
if not target.exists():
|
|
||||||
return f"Error: Path not found: {path}"
|
|
||||||
if not (target.is_dir() or target.is_file()):
|
|
||||||
return f"Error: Unsupported path: {path}"
|
|
||||||
|
|
||||||
if sort not in {"path", "modified"}:
|
|
||||||
return "Error: sort must be 'path' or 'modified'"
|
|
||||||
|
|
||||||
limit = (
|
|
||||||
_DEFAULT_FILE_HEAD_LIMIT
|
|
||||||
if head_limit is None
|
|
||||||
else None if head_limit == 0 else head_limit
|
|
||||||
)
|
|
||||||
root = target if target.is_dir() else target.parent
|
|
||||||
matches: list[tuple[str, float]] = []
|
|
||||||
|
|
||||||
for candidate in self._iter_paths(target, include_dirs=include_dirs):
|
|
||||||
if candidate.is_dir() and not include_dirs:
|
|
||||||
continue
|
|
||||||
rel_path = candidate.relative_to(root).as_posix()
|
|
||||||
display_path = self._display_path(candidate, root)
|
|
||||||
name = candidate.name
|
|
||||||
|
|
||||||
if glob and not _match_glob(rel_path, name, glob):
|
|
||||||
continue
|
|
||||||
if candidate.is_file() and not _matches_type(name, type):
|
|
||||||
continue
|
|
||||||
if candidate.is_dir() and type:
|
|
||||||
continue
|
|
||||||
if not _matches_query(display_path, query):
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
mtime = candidate.stat().st_mtime
|
|
||||||
except OSError:
|
|
||||||
mtime = 0.0
|
|
||||||
suffix = "/" if candidate.is_dir() else ""
|
|
||||||
matches.append((display_path + suffix, mtime))
|
|
||||||
|
|
||||||
if sort == "modified":
|
|
||||||
matches.sort(key=lambda item: (-item[1], item[0]))
|
|
||||||
else:
|
|
||||||
matches.sort(key=lambda item: item[0])
|
|
||||||
|
|
||||||
paths = [item[0] for item in matches]
|
|
||||||
paged, truncated = _paginate(paths, limit, offset)
|
|
||||||
if not paged:
|
|
||||||
return "No files found"
|
|
||||||
|
|
||||||
result = "\n".join(paged)
|
|
||||||
note = _pagination_note(limit, offset, truncated)
|
|
||||||
if note:
|
|
||||||
result += "\n\n" + note
|
|
||||||
return result
|
|
||||||
except PermissionError as e:
|
|
||||||
return f"Error: {e}"
|
|
||||||
except Exception as e:
|
|
||||||
return f"Error finding files: {e}"
|
|
||||||
|
|
||||||
|
|
||||||
class GrepTool(_SearchTool):
|
class GrepTool(_SearchTool):
|
||||||
"""Search file contents using a regex-like pattern."""
|
"""Search file contents using a regex-like pattern."""
|
||||||
_scopes = {"core", "subagent"}
|
_scopes = {"core", "subagent"}
|
||||||
@@ -291,8 +125,7 @@ class GrepTool(_SearchTool):
|
|||||||
return (
|
return (
|
||||||
"Search file contents with a regex pattern. "
|
"Search file contents with a regex pattern. "
|
||||||
"Default output_mode is files_with_matches (file paths only); "
|
"Default output_mode is files_with_matches (file paths only); "
|
||||||
"use content mode for matching lines with context. Prefer this "
|
"use content mode for matching lines with context. "
|
||||||
"over shell grep for ordinary workspace searches. "
|
|
||||||
"Skips binary and files >2 MB. Supports glob/type filtering."
|
"Skips binary and files >2 MB. Supports glob/type filtering."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
+51
-223
@@ -8,7 +8,6 @@ import re
|
|||||||
import shutil
|
import shutil
|
||||||
import sys
|
import sys
|
||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
from dataclasses import dataclass
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -16,17 +15,8 @@ from loguru import logger
|
|||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
|
|
||||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||||
from nanobot.agent.tools.exec_session import (
|
|
||||||
DEFAULT_MAX_OUTPUT_CHARS,
|
|
||||||
DEFAULT_YIELD_MS,
|
|
||||||
DEFAULT_EXEC_SESSION_MANAGER,
|
|
||||||
MAX_OUTPUT_CHARS,
|
|
||||||
MAX_YIELD_MS,
|
|
||||||
clamp_session_int,
|
|
||||||
format_session_poll,
|
|
||||||
)
|
|
||||||
from nanobot.agent.tools.sandbox import wrap_command
|
from nanobot.agent.tools.sandbox import wrap_command
|
||||||
from nanobot.agent.tools.schema import BooleanSchema, IntegerSchema, StringSchema, tool_parameters_schema
|
from nanobot.agent.tools.schema import IntegerSchema, StringSchema, tool_parameters_schema
|
||||||
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
|
||||||
|
|
||||||
@@ -54,22 +44,10 @@ class ExecToolConfig(Base):
|
|||||||
deny_patterns: list[str] = Field(default_factory=list)
|
deny_patterns: list[str] = Field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
|
||||||
class _PreparedCommand:
|
|
||||||
command: str
|
|
||||||
cwd: str
|
|
||||||
env: dict[str, str]
|
|
||||||
timeout: int
|
|
||||||
shell_program: str | None
|
|
||||||
login: bool
|
|
||||||
|
|
||||||
|
|
||||||
@tool_parameters(
|
@tool_parameters(
|
||||||
tool_parameters_schema(
|
tool_parameters_schema(
|
||||||
command=StringSchema("The shell command to execute"),
|
command=StringSchema("The shell command to execute"),
|
||||||
cmd=StringSchema("Compatibility alias for command"),
|
|
||||||
working_dir=StringSchema("Optional working directory for the command"),
|
working_dir=StringSchema("Optional working directory for the command"),
|
||||||
workdir=StringSchema("Compatibility alias for working_dir"),
|
|
||||||
timeout=IntegerSchema(
|
timeout=IntegerSchema(
|
||||||
60,
|
60,
|
||||||
description=(
|
description=(
|
||||||
@@ -79,44 +57,7 @@ class _PreparedCommand:
|
|||||||
minimum=1,
|
minimum=1,
|
||||||
maximum=600,
|
maximum=600,
|
||||||
),
|
),
|
||||||
shell=StringSchema(
|
required=["command"],
|
||||||
"Optional shell binary to launch. On Unix, supports sh, bash, or zsh.",
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
login=BooleanSchema(
|
|
||||||
description="Whether to run bash/zsh with login shell semantics (default true).",
|
|
||||||
default=True,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
yield_time_ms=IntegerSchema(
|
|
||||||
description=(
|
|
||||||
"Optional milliseconds to wait before returning output. "
|
|
||||||
"When set, a still-running command returns a session_id that "
|
|
||||||
"can be polled or written to with write_stdin. Omit this field "
|
|
||||||
"to keep one-shot exec behavior."
|
|
||||||
),
|
|
||||||
minimum=0,
|
|
||||||
maximum=MAX_YIELD_MS,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
max_output_chars=IntegerSchema(
|
|
||||||
description=(
|
|
||||||
"Maximum output characters to return when yield_time_ms is used "
|
|
||||||
"(default 10000, max 50000)."
|
|
||||||
),
|
|
||||||
minimum=1000,
|
|
||||||
maximum=MAX_OUTPUT_CHARS,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
max_output_tokens=IntegerSchema(
|
|
||||||
description=(
|
|
||||||
"Compatibility alias for max_output_chars. The current runtime "
|
|
||||||
"uses a character budget."
|
|
||||||
),
|
|
||||||
minimum=1000,
|
|
||||||
maximum=MAX_OUTPUT_CHARS,
|
|
||||||
nullable=True,
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
class ExecTool(Tool):
|
class ExecTool(Tool):
|
||||||
@@ -157,7 +98,6 @@ class ExecTool(Tool):
|
|||||||
sandbox: str = "",
|
sandbox: str = "",
|
||||||
path_append: str = "",
|
path_append: str = "",
|
||||||
allowed_env_keys: list[str] | None = None,
|
allowed_env_keys: list[str] | None = None,
|
||||||
session_manager: Any | None = None,
|
|
||||||
):
|
):
|
||||||
self.timeout = timeout
|
self.timeout = timeout
|
||||||
self.working_dir = working_dir
|
self.working_dir = working_dir
|
||||||
@@ -185,7 +125,6 @@ class ExecTool(Tool):
|
|||||||
self.restrict_to_workspace = restrict_to_workspace
|
self.restrict_to_workspace = restrict_to_workspace
|
||||||
self.path_append = path_append
|
self.path_append = path_append
|
||||||
self.allowed_env_keys = allowed_env_keys or []
|
self.allowed_env_keys = allowed_env_keys or []
|
||||||
self._session_manager = session_manager or DEFAULT_EXEC_SESSION_MANAGER
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def name(self) -> str:
|
def name(self) -> str:
|
||||||
@@ -211,15 +150,10 @@ class ExecTool(Tool):
|
|||||||
def description(self) -> str:
|
def description(self) -> str:
|
||||||
return (
|
return (
|
||||||
"Execute a shell command and return its output. "
|
"Execute a shell command and return its output. "
|
||||||
"Use this for tests, builds, package commands, git commands, and "
|
"Prefer read_file/write_file/edit_file over cat/echo/sed, "
|
||||||
"other process execution. Prefer read_file/find_files/grep for "
|
"and grep/glob over shell find/grep. "
|
||||||
"inspection and apply_patch/write_file/edit_file for file changes "
|
|
||||||
"instead of cat, shell find/grep, echo, or sed. "
|
|
||||||
"Use -y or --yes flags to avoid interactive prompts. "
|
"Use -y or --yes flags to avoid interactive prompts. "
|
||||||
"For long-running or interactive commands, pass yield_time_ms; "
|
"Output is truncated at 10 000 chars; timeout defaults to 60s."
|
||||||
"if the command keeps running, exec returns a session_id that can "
|
|
||||||
"be polled or written to with write_stdin. Output is truncated at "
|
|
||||||
"10 000 chars; timeout defaults to 60s."
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -227,111 +161,9 @@ class ExecTool(Tool):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
async def execute(
|
async def execute(
|
||||||
self, command: str | None = None, cmd: str | None = None,
|
self, command: str, working_dir: str | None = None,
|
||||||
working_dir: str | None = None, workdir: str | None = None,
|
timeout: int | None = None, **kwargs: Any,
|
||||||
timeout: int | None = None, shell: str | None = None,
|
|
||||||
login: bool | None = None, yield_time_ms: int | None = None,
|
|
||||||
max_output_chars: int | None = None,
|
|
||||||
max_output_tokens: int | None = None,
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> str:
|
) -> str:
|
||||||
command = command or cmd
|
|
||||||
working_dir = working_dir or workdir
|
|
||||||
if not command:
|
|
||||||
return "Error: Missing command. Provide command or cmd."
|
|
||||||
if max_output_chars is None:
|
|
||||||
max_output_chars = max_output_tokens
|
|
||||||
|
|
||||||
prepared = self._prepare_command(command, working_dir, timeout, shell, login)
|
|
||||||
if isinstance(prepared, str):
|
|
||||||
return prepared
|
|
||||||
|
|
||||||
if yield_time_ms is not None:
|
|
||||||
return await self._execute_session(prepared, yield_time_ms, max_output_chars)
|
|
||||||
|
|
||||||
try:
|
|
||||||
process = await self._spawn(
|
|
||||||
prepared.command,
|
|
||||||
prepared.cwd,
|
|
||||||
prepared.env,
|
|
||||||
prepared.shell_program,
|
|
||||||
prepared.login,
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
stdout, stderr = await asyncio.wait_for(
|
|
||||||
process.communicate(),
|
|
||||||
timeout=prepared.timeout,
|
|
||||||
)
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
await self._kill_process(process)
|
|
||||||
return f"Error: Command timed out after {prepared.timeout} seconds"
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
await self._kill_process(process)
|
|
||||||
raise
|
|
||||||
|
|
||||||
output_parts = []
|
|
||||||
|
|
||||||
if stdout:
|
|
||||||
output_parts.append(stdout.decode("utf-8", errors="replace"))
|
|
||||||
|
|
||||||
if stderr:
|
|
||||||
stderr_text = stderr.decode("utf-8", errors="replace")
|
|
||||||
if stderr_text.strip():
|
|
||||||
output_parts.append(f"STDERR:\n{stderr_text}")
|
|
||||||
|
|
||||||
output_parts.append(f"\nExit code: {process.returncode}")
|
|
||||||
|
|
||||||
result = "\n".join(output_parts) if output_parts else "(no output)"
|
|
||||||
|
|
||||||
max_len = clamp_session_int(max_output_chars, self._MAX_OUTPUT, 1000, MAX_OUTPUT_CHARS)
|
|
||||||
if len(result) > max_len:
|
|
||||||
half = max_len // 2
|
|
||||||
result = (
|
|
||||||
result[:half]
|
|
||||||
+ f"\n\n... ({len(result) - max_len:,} chars truncated) ...\n\n"
|
|
||||||
+ result[-half:]
|
|
||||||
)
|
|
||||||
|
|
||||||
return result
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
return f"Error executing command: {str(e)}"
|
|
||||||
|
|
||||||
async def _execute_session(
|
|
||||||
self,
|
|
||||||
prepared: _PreparedCommand,
|
|
||||||
yield_time_ms: int | None,
|
|
||||||
max_output_chars: int | None,
|
|
||||||
) -> str:
|
|
||||||
try:
|
|
||||||
session_id, poll = await self._session_manager.start(
|
|
||||||
command=prepared.command,
|
|
||||||
cwd=prepared.cwd,
|
|
||||||
env=prepared.env,
|
|
||||||
timeout=prepared.timeout,
|
|
||||||
shell_program=prepared.shell_program,
|
|
||||||
login=prepared.login,
|
|
||||||
yield_time_ms=clamp_session_int(yield_time_ms, DEFAULT_YIELD_MS, 0, MAX_YIELD_MS),
|
|
||||||
max_output_chars=clamp_session_int(
|
|
||||||
max_output_chars,
|
|
||||||
DEFAULT_MAX_OUTPUT_CHARS,
|
|
||||||
1000,
|
|
||||||
MAX_OUTPUT_CHARS,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
return format_session_poll(session_id, poll)
|
|
||||||
except Exception as exc:
|
|
||||||
return f"Error executing command: {exc}"
|
|
||||||
|
|
||||||
def _prepare_command(
|
|
||||||
self,
|
|
||||||
command: str,
|
|
||||||
working_dir: str | None = None,
|
|
||||||
timeout: int | None = None,
|
|
||||||
shell: str | None = None,
|
|
||||||
login: bool | None = None,
|
|
||||||
) -> _PreparedCommand | str:
|
|
||||||
cwd = working_dir or self.working_dir or os.getcwd()
|
cwd = working_dir or self.working_dir or os.getcwd()
|
||||||
|
|
||||||
# Prevent an LLM-supplied working_dir from escaping the configured
|
# Prevent an LLM-supplied working_dir from escaping the configured
|
||||||
@@ -379,24 +211,52 @@ class ExecTool(Tool):
|
|||||||
env["NANOBOT_PATH_APPEND"] = self.path_append
|
env["NANOBOT_PATH_APPEND"] = self.path_append
|
||||||
command = f'export PATH="$PATH{os.pathsep}$NANOBOT_PATH_APPEND"; {command}'
|
command = f'export PATH="$PATH{os.pathsep}$NANOBOT_PATH_APPEND"; {command}'
|
||||||
|
|
||||||
shell_program, shell_error = self._resolve_shell(shell)
|
try:
|
||||||
if shell_error:
|
process = await self._spawn(command, cwd, env)
|
||||||
return shell_error
|
|
||||||
|
|
||||||
return _PreparedCommand(
|
try:
|
||||||
command=command,
|
stdout, stderr = await asyncio.wait_for(
|
||||||
cwd=cwd,
|
process.communicate(),
|
||||||
env=env,
|
timeout=effective_timeout,
|
||||||
timeout=effective_timeout,
|
)
|
||||||
shell_program=shell_program,
|
except asyncio.TimeoutError:
|
||||||
login=True if login is None else login,
|
await self._kill_process(process)
|
||||||
)
|
return f"Error: Command timed out after {effective_timeout} seconds"
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
await self._kill_process(process)
|
||||||
|
raise
|
||||||
|
|
||||||
|
output_parts = []
|
||||||
|
|
||||||
|
if stdout:
|
||||||
|
output_parts.append(stdout.decode("utf-8", errors="replace"))
|
||||||
|
|
||||||
|
if stderr:
|
||||||
|
stderr_text = stderr.decode("utf-8", errors="replace")
|
||||||
|
if stderr_text.strip():
|
||||||
|
output_parts.append(f"STDERR:\n{stderr_text}")
|
||||||
|
|
||||||
|
output_parts.append(f"\nExit code: {process.returncode}")
|
||||||
|
|
||||||
|
result = "\n".join(output_parts) if output_parts else "(no output)"
|
||||||
|
|
||||||
|
max_len = self._MAX_OUTPUT
|
||||||
|
if len(result) > max_len:
|
||||||
|
half = max_len // 2
|
||||||
|
result = (
|
||||||
|
result[:half]
|
||||||
|
+ f"\n\n... ({len(result) - max_len:,} chars truncated) ...\n\n"
|
||||||
|
+ result[-half:]
|
||||||
|
)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
return f"Error executing command: {str(e)}"
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _spawn(
|
async def _spawn(
|
||||||
command: str, cwd: str, env: dict[str, str],
|
command: str, cwd: str, env: dict[str, str],
|
||||||
shell_program: str | None = None,
|
|
||||||
login: bool = True,
|
|
||||||
) -> asyncio.subprocess.Process:
|
) -> asyncio.subprocess.Process:
|
||||||
"""Launch *command* in a platform-appropriate shell."""
|
"""Launch *command* in a platform-appropriate shell."""
|
||||||
if _IS_WINDOWS:
|
if _IS_WINDOWS:
|
||||||
@@ -406,52 +266,20 @@ 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,
|
||||||
env=env,
|
env=env,
|
||||||
)
|
)
|
||||||
shell_program = shell_program or shutil.which("bash") or "/bin/bash"
|
bash = shutil.which("bash") or "/bin/bash"
|
||||||
args = [shell_program]
|
|
||||||
shell_name = Path(shell_program).name.lower()
|
|
||||||
if login and shell_name in {"bash", "bash.exe", "zsh", "zsh.exe"}:
|
|
||||||
args.append("-l")
|
|
||||||
args.extend(["-c", command])
|
|
||||||
return await asyncio.create_subprocess_exec(
|
return await asyncio.create_subprocess_exec(
|
||||||
*args,
|
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,
|
||||||
env=env,
|
env=env,
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _resolve_shell(shell: str | None) -> tuple[str | None, str | None]:
|
|
||||||
if not shell:
|
|
||||||
return None, None
|
|
||||||
if _IS_WINDOWS:
|
|
||||||
return None, "Error: shell parameter is not supported on Windows"
|
|
||||||
if "\0" in shell or "\n" in shell or "\r" in shell:
|
|
||||||
return None, "Error: shell contains invalid characters"
|
|
||||||
allowed = {"sh", "bash", "zsh"}
|
|
||||||
path = Path(shell).expanduser()
|
|
||||||
if path.is_absolute():
|
|
||||||
if path.name not in allowed:
|
|
||||||
return None, f"Error: unsupported shell {shell!r}. Allowed: bash, sh, zsh"
|
|
||||||
if not path.is_file() or not os.access(path, os.X_OK):
|
|
||||||
return None, f"Error: shell is not executable: {shell}"
|
|
||||||
return str(path), None
|
|
||||||
if "/" in shell or "\\" in shell:
|
|
||||||
return None, "Error: shell must be a shell name or absolute path"
|
|
||||||
if shell not in allowed:
|
|
||||||
return None, f"Error: unsupported shell {shell!r}. Allowed: bash, sh, zsh"
|
|
||||||
resolved = shutil.which(shell)
|
|
||||||
if not resolved:
|
|
||||||
return None, f"Error: shell not found: {shell}"
|
|
||||||
return resolved, None
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _kill_process(process: asyncio.subprocess.Process) -> None:
|
async def _kill_process(process: asyncio.subprocess.Process) -> None:
|
||||||
"""Kill a subprocess and reap it to prevent zombies."""
|
"""Kill a subprocess and reap it to prevent zombies."""
|
||||||
@@ -588,7 +416,7 @@ class ExecTool(Tool):
|
|||||||
# Windows: match drive-root paths like `C:\` as well as `C:\path\to\file`, and UNC paths like `\\server\share`
|
# Windows: match drive-root paths like `C:\` as well as `C:\path\to\file`, and UNC paths like `\\server\share`
|
||||||
# NOTE: `*` is required so `C:\` (nothing after the slash) is still extracted.
|
# NOTE: `*` is required so `C:\` (nothing after the slash) is still extracted.
|
||||||
win_paths = re.findall(
|
win_paths = re.findall(
|
||||||
r"(?<![A-Za-z])(?:[A-Za-z]:[^\s\"'|><;]*|\\\\[^\s\"'|><;]+(?:\\[^\s\"'|><;]+)*)",
|
r"(?:[A-Za-z]:[^\s\"'|><;]*|\\\\[^\s\"'|><;]+(?:\\[^\s\"'|><;]+)*)",
|
||||||
command
|
command
|
||||||
)
|
)
|
||||||
posix_paths = re.findall(r"(?:^|[\s|>'\"])(/[^\s\"'>;|<]+)", command) # POSIX: /absolute only
|
posix_paths = re.findall(r"(?:^|[\s|>'\"])(/[^\s\"'>;|<]+)", command) # POSIX: /absolute only
|
||||||
|
|||||||
+18
-99
@@ -8,7 +8,7 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from typing import Any, Callable
|
from typing import Any, Callable
|
||||||
from urllib.parse import quote, urljoin, urlparse
|
from urllib.parse import quote, urlparse
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -78,82 +78,9 @@ def _validate_url(url: str) -> tuple[bool, str]:
|
|||||||
def _validate_url_safe(url: str) -> tuple[bool, str]:
|
def _validate_url_safe(url: str) -> tuple[bool, str]:
|
||||||
"""Validate URL with SSRF protection: scheme, domain, and resolved IP check."""
|
"""Validate URL with SSRF protection: scheme, domain, and resolved IP check."""
|
||||||
from nanobot.security.network import validate_url_target
|
from nanobot.security.network import validate_url_target
|
||||||
|
|
||||||
return validate_url_target(url)
|
return validate_url_target(url)
|
||||||
|
|
||||||
|
|
||||||
async def _get_with_safe_redirects(
|
|
||||||
client: httpx.AsyncClient,
|
|
||||||
url: str,
|
|
||||||
headers: dict[str, str] | None = None,
|
|
||||||
) -> tuple[httpx.Response | None, str | None]:
|
|
||||||
"""GET a URL while validating every redirect target before requesting it."""
|
|
||||||
current_url = url
|
|
||||||
for _ in range(MAX_REDIRECTS + 1):
|
|
||||||
is_valid, error_msg = _validate_url_safe(current_url)
|
|
||||||
if not is_valid:
|
|
||||||
return None, f"Redirect blocked: {error_msg}"
|
|
||||||
|
|
||||||
response = await client.get(current_url, headers=headers, follow_redirects=False)
|
|
||||||
is_redirect = 300 <= response.status_code < 400
|
|
||||||
if not is_redirect:
|
|
||||||
return response, None
|
|
||||||
|
|
||||||
location = response.headers.get("location")
|
|
||||||
if not location:
|
|
||||||
return response, None
|
|
||||||
|
|
||||||
next_url = urljoin(str(response.url), location)
|
|
||||||
is_valid, error_msg = _validate_url_safe(next_url)
|
|
||||||
if not is_valid:
|
|
||||||
await response.aclose()
|
|
||||||
return None, f"Redirect blocked: {error_msg}"
|
|
||||||
|
|
||||||
await response.aclose()
|
|
||||||
current_url = next_url
|
|
||||||
|
|
||||||
return None, f"Too many redirects: exceeded limit of {MAX_REDIRECTS}"
|
|
||||||
|
|
||||||
|
|
||||||
async def _stream_with_safe_redirects(
|
|
||||||
client: httpx.AsyncClient,
|
|
||||||
url: str,
|
|
||||||
headers: dict[str, str] | None = None,
|
|
||||||
) -> tuple[httpx.Response | None, Any | None, str | None]:
|
|
||||||
"""Open a streamed response while validating every redirect target first."""
|
|
||||||
current_url = url
|
|
||||||
for _ in range(MAX_REDIRECTS + 1):
|
|
||||||
is_valid, error_msg = _validate_url_safe(current_url)
|
|
||||||
if not is_valid:
|
|
||||||
return None, None, f"Redirect blocked: {error_msg}"
|
|
||||||
|
|
||||||
stream = client.stream(
|
|
||||||
"GET",
|
|
||||||
current_url,
|
|
||||||
headers=headers,
|
|
||||||
follow_redirects=False,
|
|
||||||
)
|
|
||||||
response = await stream.__aenter__()
|
|
||||||
is_redirect = 300 <= response.status_code < 400
|
|
||||||
if not is_redirect:
|
|
||||||
return response, stream, None
|
|
||||||
|
|
||||||
location = response.headers.get("location")
|
|
||||||
if not location:
|
|
||||||
return response, stream, None
|
|
||||||
|
|
||||||
next_url = urljoin(str(response.url), location)
|
|
||||||
is_valid, error_msg = _validate_url_safe(next_url)
|
|
||||||
if not is_valid:
|
|
||||||
await stream.__aexit__(None, None, None)
|
|
||||||
return None, None, f"Redirect blocked: {error_msg}"
|
|
||||||
|
|
||||||
await stream.__aexit__(None, None, None)
|
|
||||||
current_url = next_url
|
|
||||||
|
|
||||||
return None, None, f"Too many redirects: exceeded limit of {MAX_REDIRECTS}"
|
|
||||||
|
|
||||||
|
|
||||||
def _format_results(query: str, items: list[dict[str, Any]], n: int) -> str:
|
def _format_results(query: str, items: list[dict[str, Any]], n: int) -> str:
|
||||||
"""Format provider results into shared plaintext output."""
|
"""Format provider results into shared plaintext output."""
|
||||||
if not items:
|
if not items:
|
||||||
@@ -561,26 +488,19 @@ class WebFetchTool(Tool):
|
|||||||
|
|
||||||
# Detect and fetch images directly to avoid Jina's textual image captioning
|
# Detect and fetch images directly to avoid Jina's textual image captioning
|
||||||
try:
|
try:
|
||||||
async with httpx.AsyncClient(proxy=self.proxy, timeout=15.0) as client:
|
async with httpx.AsyncClient(proxy=self.proxy, follow_redirects=True, max_redirects=MAX_REDIRECTS, timeout=15.0) as client:
|
||||||
r, stream, redirect_error = await _stream_with_safe_redirects(
|
async with client.stream("GET", url, headers={"User-Agent": self.user_agent}) as r:
|
||||||
client,
|
from nanobot.security.network import validate_resolved_url
|
||||||
url,
|
|
||||||
headers={"User-Agent": self.user_agent},
|
redir_ok, redir_err = validate_resolved_url(str(r.url))
|
||||||
)
|
if not redir_ok:
|
||||||
if redirect_error:
|
return json.dumps({"error": f"Redirect blocked: {redir_err}", "url": url}, ensure_ascii=False)
|
||||||
return json.dumps({"error": redirect_error, "url": url}, ensure_ascii=False)
|
|
||||||
if r is None:
|
|
||||||
return json.dumps({"error": "Fetch failed", "url": url}, ensure_ascii=False)
|
|
||||||
|
|
||||||
try:
|
|
||||||
ctype = r.headers.get("content-type", "")
|
ctype = r.headers.get("content-type", "")
|
||||||
if ctype.startswith("image/"):
|
if ctype.startswith("image/"):
|
||||||
r.raise_for_status()
|
r.raise_for_status()
|
||||||
raw = await r.aread()
|
raw = await r.aread()
|
||||||
return build_image_content_blocks(raw, ctype, url, f"(Image fetched from: {url})")
|
return build_image_content_blocks(raw, ctype, url, f"(Image fetched from: {url})")
|
||||||
finally:
|
|
||||||
if stream is not None:
|
|
||||||
await stream.__aexit__(None, None, None)
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug("Pre-fetch image detection failed for {}: {}", url, e)
|
logger.debug("Pre-fetch image detection failed for {}: {}", url, e)
|
||||||
|
|
||||||
@@ -629,22 +549,23 @@ class WebFetchTool(Tool):
|
|||||||
|
|
||||||
async def _fetch_readability(self, url: str, extract_mode: str, max_chars: int) -> Any:
|
async def _fetch_readability(self, url: str, extract_mode: str, max_chars: int) -> Any:
|
||||||
"""Local fallback using readability-lxml."""
|
"""Local fallback using readability-lxml."""
|
||||||
|
from readability import Document
|
||||||
|
|
||||||
try:
|
try:
|
||||||
async with httpx.AsyncClient(
|
async with httpx.AsyncClient(
|
||||||
|
follow_redirects=True,
|
||||||
|
max_redirects=MAX_REDIRECTS,
|
||||||
timeout=30.0,
|
timeout=30.0,
|
||||||
proxy=self.proxy,
|
proxy=self.proxy,
|
||||||
) as client:
|
) as client:
|
||||||
r, redirect_error = await _get_with_safe_redirects(
|
r = await client.get(url, headers={"User-Agent": self.user_agent})
|
||||||
client,
|
|
||||||
url,
|
|
||||||
headers={"User-Agent": self.user_agent},
|
|
||||||
)
|
|
||||||
if redirect_error:
|
|
||||||
return json.dumps({"error": redirect_error, "url": url}, ensure_ascii=False)
|
|
||||||
if r is None:
|
|
||||||
return json.dumps({"error": "Fetch failed", "url": url}, ensure_ascii=False)
|
|
||||||
r.raise_for_status()
|
r.raise_for_status()
|
||||||
|
|
||||||
|
from nanobot.security.network import validate_resolved_url
|
||||||
|
redir_ok, redir_err = validate_resolved_url(str(r.url))
|
||||||
|
if not redir_ok:
|
||||||
|
return json.dumps({"error": f"Redirect blocked: {redir_err}", "url": url}, ensure_ascii=False)
|
||||||
|
|
||||||
ctype = r.headers.get("content-type", "")
|
ctype = r.headers.get("content-type", "")
|
||||||
if ctype.startswith("image/"):
|
if ctype.startswith("image/"):
|
||||||
return build_image_content_blocks(r.content, ctype, url, f"(Image fetched from: {url})")
|
return build_image_content_blocks(r.content, ctype, url, f"(Image fetched from: {url})")
|
||||||
@@ -652,8 +573,6 @@ class WebFetchTool(Tool):
|
|||||||
if "application/json" in ctype:
|
if "application/json" in ctype:
|
||||||
text, extractor = json.dumps(r.json(), indent=2, ensure_ascii=False), "json"
|
text, extractor = json.dumps(r.json(), indent=2, ensure_ascii=False), "json"
|
||||||
elif "text/html" in ctype or r.text[:256].lower().startswith(("<!doctype", "<html")):
|
elif "text/html" in ctype or r.text[:256].lower().startswith(("<!doctype", "<html")):
|
||||||
from readability import Document
|
|
||||||
|
|
||||||
doc = Document(r.text)
|
doc = Document(r.text)
|
||||||
content = self._to_markdown(doc.summary()) if extract_mode == "markdown" else _strip_tags(doc.summary())
|
content = self._to_markdown(doc.summary()) if extract_mode == "markdown" else _strip_tags(doc.summary())
|
||||||
text = f"# {doc.title()}\n\n{content}" if doc.title() else content
|
text = f"# {doc.title()}\n\n{content}" if doc.title() else content
|
||||||
|
|||||||
@@ -70,47 +70,34 @@ 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_channel_names, discover_enabled
|
from nanobot.channels.registry import discover_all
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
# Collect enabled module names first, then only import those.
|
for name, cls in discover_all().items():
|
||||||
# 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
|
||||||
if (
|
enabled = (
|
||||||
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)
|
||||||
):
|
)
|
||||||
enabled_names.add(name)
|
if not enabled:
|
||||||
|
|
||||||
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
|
||||||
static_path = _default_webui_dist()
|
static_path = _default_webui_dist()
|
||||||
if static_path is not None:
|
if static_path is not None:
|
||||||
kwargs["static_dist_path"] = static_path
|
kwargs["static_dist_path"] = static_path
|
||||||
kwargs["workspace_path"] = self.config.workspace_path
|
|
||||||
if self._webui_runtime_model_name is not None:
|
if self._webui_runtime_model_name is not None:
|
||||||
kwargs["runtime_model_name"] = self._webui_runtime_model_name
|
kwargs["runtime_model_name"] = self._webui_runtime_model_name
|
||||||
channel = cls(section, self.bus, **kwargs)
|
channel = cls(section, self.bus, **kwargs)
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
"""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
|
||||||
@@ -36,14 +37,12 @@ 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(enabled_names: set[str] | None = None) -> dict[str, type[BaseChannel]]:
|
def discover_plugins() -> 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
|
||||||
@@ -52,44 +51,21 @@ def discover_plugins(enabled_names: set[str] | None = None) -> dict[str, type[Ba
|
|||||||
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.
|
||||||
"""
|
"""
|
||||||
names = discover_channel_names()
|
builtin: dict[str, type[BaseChannel]] = {}
|
||||||
return discover_enabled(set(names), _names=names, _include_all_external=True)
|
for modname in discover_channel_names():
|
||||||
|
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}
|
||||||
|
|||||||
+235
-246
@@ -34,35 +34,18 @@ from nanobot.bus.events import OUTBOUND_META_AGENT_UI, OutboundMessage
|
|||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.channels.base import BaseChannel
|
from nanobot.channels.base import BaseChannel
|
||||||
from nanobot.command.builtin import builtin_command_palette
|
from nanobot.command.builtin import builtin_command_palette
|
||||||
from nanobot.config.paths import get_media_dir, get_workspace_path
|
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.webui.settings_api import (
|
from nanobot.utils.webui_thread_disk import delete_webui_thread
|
||||||
WebUISettingsError,
|
from nanobot.utils.webui_transcript import append_transcript_object, build_webui_thread_response
|
||||||
settings_payload,
|
from nanobot.utils.webui_turn_helpers import websocket_turn_wall_started_at
|
||||||
update_agent_settings,
|
|
||||||
update_image_generation_settings,
|
|
||||||
update_provider_settings,
|
|
||||||
update_web_search_settings,
|
|
||||||
)
|
|
||||||
from nanobot.webui.cli_apps_api import (
|
|
||||||
cli_apps_action,
|
|
||||||
cli_apps_payload,
|
|
||||||
normalize_cli_app_mentions,
|
|
||||||
)
|
|
||||||
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
|
||||||
@@ -239,6 +222,47 @@ 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:]}"
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_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()
|
||||||
@@ -425,16 +449,6 @@ _MEDIA_ALLOWED_MIMES: frozenset[str] = frozenset({
|
|||||||
"video/webm",
|
"video/webm",
|
||||||
"video/quicktime",
|
"video/quicktime",
|
||||||
})
|
})
|
||||||
_MARKDOWN_LOCAL_IMAGE_RE = re.compile(
|
|
||||||
r"!\[([^\]]*)\]\((<[^>]+>|[^)\s]+)(\s+(?:\"[^\"]*\"|'[^']*'))?\)"
|
|
||||||
)
|
|
||||||
_INLINE_MARKDOWN_IMAGE_EXTS: frozenset[str] = frozenset({
|
|
||||||
".png",
|
|
||||||
".jpg",
|
|
||||||
".jpeg",
|
|
||||||
".webp",
|
|
||||||
".gif",
|
|
||||||
})
|
|
||||||
|
|
||||||
|
|
||||||
def _issue_route_secret_matches(headers: Any, configured_secret: str) -> bool:
|
def _issue_route_secret_matches(headers: Any, configured_secret: str) -> bool:
|
||||||
@@ -464,7 +478,6 @@ class WebSocketChannel(BaseChannel):
|
|||||||
*,
|
*,
|
||||||
session_manager: "SessionManager | None" = None,
|
session_manager: "SessionManager | None" = None,
|
||||||
static_dist_path: Path | None = None,
|
static_dist_path: Path | None = None,
|
||||||
workspace_path: Path | None = None,
|
|
||||||
runtime_model_name: Callable[[], str | None] | None = None,
|
runtime_model_name: Callable[[], str | None] | None = None,
|
||||||
):
|
):
|
||||||
if isinstance(config, dict):
|
if isinstance(config, dict):
|
||||||
@@ -477,10 +490,6 @@ class WebSocketChannel(BaseChannel):
|
|||||||
self._conn_chats: dict[Any, set[str]] = {}
|
self._conn_chats: dict[Any, set[str]] = {}
|
||||||
# connection -> default chat_id for legacy frames that omit routing.
|
# connection -> default chat_id for legacy frames that omit routing.
|
||||||
self._conn_default: dict[Any, str] = {}
|
self._conn_default: dict[Any, str] = {}
|
||||||
# Chat IDs that opted into WebUI-specific rendering by sending a typed
|
|
||||||
# envelope with ``webui: true``. Raw WebSocket clients keep the legacy
|
|
||||||
# wire shape.
|
|
||||||
self._webui_chats: set[str] = set()
|
|
||||||
# Single-use tokens consumed at WebSocket handshake.
|
# Single-use tokens consumed at WebSocket handshake.
|
||||||
self._issued_tokens: dict[str, float] = {}
|
self._issued_tokens: dict[str, float] = {}
|
||||||
# Multi-use tokens for HTTP routes served beside WS; checked but not consumed.
|
# Multi-use tokens for HTTP routes served beside WS; checked but not consumed.
|
||||||
@@ -491,14 +500,7 @@ class WebSocketChannel(BaseChannel):
|
|||||||
self._static_dist_path: Path | None = (
|
self._static_dist_path: Path | None = (
|
||||||
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._workspace_path = (
|
|
||||||
Path(workspace_path).expanduser()
|
|
||||||
if workspace_path is not None
|
|
||||||
else get_workspace_path()
|
|
||||||
).resolve(strict=False)
|
|
||||||
self._runtime_model_name = runtime_model_name
|
self._runtime_model_name = runtime_model_name
|
||||||
self._settings_restart_sections: set[str] = set()
|
|
||||||
self._stream_text_buffers: dict[tuple[str, str], list[str]] = {}
|
|
||||||
# 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
|
||||||
@@ -661,12 +663,6 @@ 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)
|
||||||
|
|
||||||
@@ -676,24 +672,6 @@ 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)
|
|
||||||
|
|
||||||
if got == "/api/settings/cli-apps":
|
|
||||||
return self._handle_settings_cli_apps(request)
|
|
||||||
|
|
||||||
if got == "/api/settings/cli-apps/install":
|
|
||||||
return await self._handle_settings_cli_apps_action(request, "install")
|
|
||||||
|
|
||||||
if got == "/api/settings/cli-apps/update":
|
|
||||||
return await self._handle_settings_cli_apps_action(request, "update")
|
|
||||||
|
|
||||||
if got == "/api/settings/cli-apps/uninstall":
|
|
||||||
return await self._handle_settings_cli_apps_action(request, "uninstall")
|
|
||||||
|
|
||||||
if got == "/api/settings/cli-apps/test":
|
|
||||||
return await self._handle_settings_cli_apps_action(request, "test")
|
|
||||||
|
|
||||||
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))
|
||||||
@@ -805,141 +783,221 @@ 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 = [
|
||||||
for s in sessions:
|
{k: v for k, v in s.items() if k != "path"}
|
||||||
key = s.get("key")
|
for s in sessions
|
||||||
if not (isinstance(key, str) and key.startswith("websocket:")):
|
if isinstance(s.get("key"), str) and s["key"].startswith("websocket:")
|
||||||
continue
|
]
|
||||||
row = {k: v for k, v in s.items() if k != "path"}
|
|
||||||
chat_id = key.split(":", 1)[1]
|
|
||||||
started_at = websocket_turn_wall_started_at(chat_id)
|
|
||||||
if started_at is not None:
|
|
||||||
row["run_started_at"] = started_at
|
|
||||||
cleaned.append(row)
|
|
||||||
return _http_json_response({"sessions": cleaned})
|
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:
|
||||||
|
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
|
||||||
|
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._with_settings_restart_state(settings_payload()))
|
return _http_json_response(self._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)
|
||||||
try:
|
config = load_config()
|
||||||
payload = update_agent_settings(query)
|
defaults = config.agents.defaults
|
||||||
except WebUISettingsError as e:
|
changed = False
|
||||||
return _http_error(e.status, e.message)
|
|
||||||
return _http_json_response(
|
model = _query_first(query, "model")
|
||||||
self._with_settings_restart_state(payload, section="runtime")
|
if model is not None:
|
||||||
)
|
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)
|
||||||
|
spec = find_by_name(provider)
|
||||||
|
if (
|
||||||
|
provider_config is None
|
||||||
|
or spec is None
|
||||||
|
or not _provider_configured_for_settings(spec, provider_config)
|
||||||
|
):
|
||||||
|
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)
|
||||||
try:
|
provider_name = (_query_first(query, "provider") or "").strip()
|
||||||
payload = update_provider_settings(query)
|
if not provider_name:
|
||||||
except WebUISettingsError as e:
|
return _http_error(400, "provider is required")
|
||||||
return _http_error(e.status, e.message)
|
spec = find_by_name(provider_name)
|
||||||
return _http_json_response(self._with_settings_restart_state(payload, section="image"))
|
if spec is None or spec.is_oauth:
|
||||||
|
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")
|
||||||
query = _parse_query(request.path)
|
from nanobot.config.loader import load_config, save_config
|
||||||
try:
|
|
||||||
payload = update_web_search_settings(query)
|
|
||||||
except WebUISettingsError as e:
|
|
||||||
return _http_error(e.status, e.message)
|
|
||||||
return _http_json_response(self._with_settings_restart_state(payload, section="web"))
|
|
||||||
|
|
||||||
def _handle_settings_image_generation_update(self, request: WsRequest) -> Response:
|
|
||||||
if not self._check_api_token(request):
|
|
||||||
return _http_error(401, "Unauthorized")
|
|
||||||
query = _parse_query(request.path)
|
query = _parse_query(request.path)
|
||||||
try:
|
provider_name = (_query_first(query, "provider") or "").strip().lower()
|
||||||
payload = update_image_generation_settings(query)
|
provider_option = _WEB_SEARCH_PROVIDER_BY_NAME.get(provider_name)
|
||||||
except WebUISettingsError as e:
|
if provider_option is None:
|
||||||
return _http_error(e.status, e.message)
|
return _http_error(400, "unknown web search provider")
|
||||||
return _http_json_response(self._with_settings_restart_state(payload, section="image"))
|
|
||||||
|
|
||||||
def _handle_settings_cli_apps(self, request: WsRequest) -> Response:
|
config = load_config()
|
||||||
if not self._check_api_token(request):
|
search_config = config.tools.web.search
|
||||||
return _http_error(401, "Unauthorized")
|
previous_provider = search_config.provider
|
||||||
try:
|
changed = False
|
||||||
payload = cli_apps_payload()
|
|
||||||
except Exception:
|
|
||||||
self.logger.exception("failed to load CLI Apps payload")
|
|
||||||
return _http_error(500, "failed to load CLI Apps")
|
|
||||||
return _http_json_response(payload)
|
|
||||||
|
|
||||||
async def _handle_settings_cli_apps_action(self, request: WsRequest, action: str) -> Response:
|
def set_value(attr: str, value: str | None) -> None:
|
||||||
if not self._check_api_token(request):
|
nonlocal changed
|
||||||
return _http_error(401, "Unauthorized")
|
if getattr(search_config, attr) != value:
|
||||||
query = _parse_query(request.path)
|
setattr(search_config, attr, value)
|
||||||
try:
|
changed = True
|
||||||
payload = await asyncio.to_thread(cli_apps_action, action, query)
|
|
||||||
except WebUISettingsError as e:
|
if search_config.provider != provider_name:
|
||||||
return _http_error(e.status, e.message)
|
search_config.provider = provider_name
|
||||||
except Exception as e:
|
changed = True
|
||||||
status = getattr(e, "status", 500)
|
|
||||||
message = getattr(e, "message", str(e))
|
credential = provider_option["credential"]
|
||||||
if status >= 500:
|
if credential == "none":
|
||||||
self.logger.exception("CLI Apps action '{}' failed", action)
|
set_value("api_key", "")
|
||||||
return _http_error(status, message)
|
set_value("base_url", "")
|
||||||
return _http_json_response(payload)
|
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:
|
||||||
@@ -982,7 +1040,6 @@ class WebSocketChannel(BaseChannel):
|
|||||||
data = build_webui_thread_response(
|
data = build_webui_thread_response(
|
||||||
decoded_key,
|
decoded_key,
|
||||||
augment_user_media=self._augment_transcript_user_media,
|
augment_user_media=self._augment_transcript_user_media,
|
||||||
augment_assistant_text=self._rewrite_local_markdown_images,
|
|
||||||
)
|
)
|
||||||
if data is None:
|
if data is None:
|
||||||
return _http_error(404, "webui thread not found")
|
return _http_error(404, "webui thread not found")
|
||||||
@@ -1029,9 +1086,6 @@ class WebSocketChannel(BaseChannel):
|
|||||||
}
|
}
|
||||||
if media:
|
if media:
|
||||||
user_obj["media_paths"] = list(media)
|
user_obj["media_paths"] = list(media)
|
||||||
cli_apps = meta.get("cli_apps")
|
|
||||||
if isinstance(cli_apps, list) and cli_apps:
|
|
||||||
user_obj["cli_apps"] = cli_apps
|
|
||||||
self._try_append_webui_transcript(chat_id, user_obj)
|
self._try_append_webui_transcript(chat_id, user_obj)
|
||||||
await super()._handle_message(
|
await super()._handle_message(
|
||||||
sender_id,
|
sender_id,
|
||||||
@@ -1121,46 +1175,6 @@ class WebSocketChannel(BaseChannel):
|
|||||||
return None
|
return None
|
||||||
return {"url": signed, "name": path.name}
|
return {"url": signed, "name": path.name}
|
||||||
|
|
||||||
def _markdown_image_url_for_local_path(self, raw_url: str) -> str | None:
|
|
||||||
url = raw_url.strip()
|
|
||||||
if url.startswith("<") and url.endswith(">"):
|
|
||||||
url = url[1:-1].strip()
|
|
||||||
if not url or url.startswith(("/api/media/", "#")):
|
|
||||||
return None
|
|
||||||
parsed = urlparse(url)
|
|
||||||
if parsed.scheme or parsed.netloc:
|
|
||||||
return None
|
|
||||||
if parsed.query or parsed.fragment:
|
|
||||||
return None
|
|
||||||
path_text = unquote(url)
|
|
||||||
if Path(path_text).suffix.lower() not in _INLINE_MARKDOWN_IMAGE_EXTS:
|
|
||||||
return None
|
|
||||||
candidate = Path(path_text).expanduser()
|
|
||||||
if not candidate.is_absolute():
|
|
||||||
candidate = self._workspace_path / candidate
|
|
||||||
try:
|
|
||||||
resolved = candidate.resolve(strict=False)
|
|
||||||
resolved.relative_to(self._workspace_path)
|
|
||||||
except (OSError, ValueError):
|
|
||||||
return None
|
|
||||||
if not resolved.is_file():
|
|
||||||
return None
|
|
||||||
signed = self._sign_or_stage_media_path(resolved)
|
|
||||||
return signed["url"] if signed else None
|
|
||||||
|
|
||||||
def _rewrite_local_markdown_images(self, text: str) -> str:
|
|
||||||
if "![" not in text:
|
|
||||||
return text
|
|
||||||
|
|
||||||
def replace(match: re.Match[str]) -> str:
|
|
||||||
signed_url = self._markdown_image_url_for_local_path(match.group(2))
|
|
||||||
if not signed_url:
|
|
||||||
return match.group(0)
|
|
||||||
title = match.group(3) or ""
|
|
||||||
return f""
|
|
||||||
|
|
||||||
return _MARKDOWN_LOCAL_IMAGE_RE.sub(replace, text)
|
|
||||||
|
|
||||||
def _handle_media_fetch(self, sig: str, payload: str) -> Response:
|
def _handle_media_fetch(self, sig: str, payload: str) -> Response:
|
||||||
"""Serve a single media file previously signed via
|
"""Serve a single media file previously signed via
|
||||||
:meth:`_sign_media_path`. Validates the signature, decodes the
|
:meth:`_sign_media_path`. Validates the signature, decodes the
|
||||||
@@ -1532,10 +1546,6 @@ class WebSocketChannel(BaseChannel):
|
|||||||
metadata: dict[str, Any] = {"remote": getattr(connection, "remote_address", None)}
|
metadata: dict[str, Any] = {"remote": getattr(connection, "remote_address", None)}
|
||||||
if envelope.get("webui") is True:
|
if envelope.get("webui") is True:
|
||||||
metadata["webui"] = True
|
metadata["webui"] = True
|
||||||
self._webui_chats.add(cid)
|
|
||||||
cli_apps = normalize_cli_app_mentions(envelope.get("cli_apps"))
|
|
||||||
if cli_apps:
|
|
||||||
metadata["cli_apps"] = cli_apps
|
|
||||||
image_generation = envelope.get("image_generation")
|
image_generation = envelope.get("image_generation")
|
||||||
if isinstance(image_generation, dict) and image_generation.get("enabled") is True:
|
if isinstance(image_generation, dict) and image_generation.get("enabled") is True:
|
||||||
aspect_ratio = image_generation.get("aspect_ratio")
|
aspect_ratio = image_generation.get("aspect_ratio")
|
||||||
@@ -1569,7 +1579,6 @@ class WebSocketChannel(BaseChannel):
|
|||||||
self._subs.clear()
|
self._subs.clear()
|
||||||
self._conn_chats.clear()
|
self._conn_chats.clear()
|
||||||
self._conn_default.clear()
|
self._conn_default.clear()
|
||||||
self._webui_chats.clear()
|
|
||||||
self._issued_tokens.clear()
|
self._issued_tokens.clear()
|
||||||
self._api_tokens.clear()
|
self._api_tokens.clear()
|
||||||
|
|
||||||
@@ -1648,12 +1657,10 @@ class WebSocketChannel(BaseChannel):
|
|||||||
await self._safe_send_to(connection, raw, label=" ")
|
await self._safe_send_to(connection, raw, label=" ")
|
||||||
return
|
return
|
||||||
text = msg.content
|
text = msg.content
|
||||||
should_rewrite_images = msg.chat_id in self._webui_chats
|
|
||||||
wire_text = self._rewrite_local_markdown_images(text) if should_rewrite_images else text
|
|
||||||
payload: dict[str, Any] = {
|
payload: dict[str, Any] = {
|
||||||
"event": "message",
|
"event": "message",
|
||||||
"chat_id": msg.chat_id,
|
"chat_id": msg.chat_id,
|
||||||
"text": wire_text,
|
"text": text,
|
||||||
}
|
}
|
||||||
if msg.media:
|
if msg.media:
|
||||||
payload["media"] = msg.media
|
payload["media"] = msg.media
|
||||||
@@ -1681,9 +1688,7 @@ class WebSocketChannel(BaseChannel):
|
|||||||
payload["kind"] = "tool_hint"
|
payload["kind"] = "tool_hint"
|
||||||
elif msg.metadata.get("_progress"):
|
elif msg.metadata.get("_progress"):
|
||||||
payload["kind"] = "progress"
|
payload["kind"] = "progress"
|
||||||
transcript_payload = dict(payload)
|
self._try_append_webui_transcript(msg.chat_id, payload)
|
||||||
transcript_payload["text"] = text
|
|
||||||
self._try_append_webui_transcript(msg.chat_id, transcript_payload)
|
|
||||||
raw = json.dumps(payload, ensure_ascii=False)
|
raw = json.dumps(payload, ensure_ascii=False)
|
||||||
for connection in conns:
|
for connection in conns:
|
||||||
await self._safe_send_to(connection, raw, label=" ")
|
await self._safe_send_to(connection, raw, label=" ")
|
||||||
@@ -1748,33 +1753,17 @@ class WebSocketChannel(BaseChannel):
|
|||||||
if not conns:
|
if not conns:
|
||||||
return
|
return
|
||||||
meta = metadata or {}
|
meta = metadata or {}
|
||||||
stream_key = (chat_id, str(meta.get("_stream_id") or ""))
|
|
||||||
should_rewrite_images = chat_id in self._webui_chats
|
|
||||||
transcript_body: dict[str, Any] | None = None
|
|
||||||
if meta.get("_stream_end"):
|
if meta.get("_stream_end"):
|
||||||
body: dict[str, Any] = {"event": "stream_end", "chat_id": chat_id}
|
body: dict[str, Any] = {"event": "stream_end", "chat_id": chat_id}
|
||||||
if should_rewrite_images:
|
|
||||||
buffered = self._stream_text_buffers.pop(stream_key, [])
|
|
||||||
if delta:
|
|
||||||
buffered.append(delta)
|
|
||||||
full_text = "".join(buffered)
|
|
||||||
rewritten = self._rewrite_local_markdown_images(full_text)
|
|
||||||
if rewritten != full_text or delta:
|
|
||||||
body["text"] = rewritten
|
|
||||||
transcript_body = {**body, "text": full_text}
|
|
||||||
else:
|
else:
|
||||||
body = {
|
body = {
|
||||||
"event": "delta",
|
"event": "delta",
|
||||||
"chat_id": chat_id,
|
"chat_id": chat_id,
|
||||||
"text": delta,
|
"text": delta,
|
||||||
}
|
}
|
||||||
if should_rewrite_images:
|
|
||||||
self._stream_text_buffers.setdefault(stream_key, []).append(delta)
|
|
||||||
if meta.get("_stream_id") is not None:
|
if meta.get("_stream_id") is not None:
|
||||||
body["stream_id"] = meta["_stream_id"]
|
body["stream_id"] = meta["_stream_id"]
|
||||||
if transcript_body is not None:
|
self._try_append_webui_transcript(chat_id, body)
|
||||||
transcript_body["stream_id"] = meta["_stream_id"]
|
|
||||||
self._try_append_webui_transcript(chat_id, transcript_body or body)
|
|
||||||
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=" stream ")
|
await self._safe_send_to(connection, raw, label=" stream ")
|
||||||
|
|||||||
+6
-163
@@ -79,12 +79,6 @@ BASE_INFO: dict[str, str] = {"channel_version": WEIXIN_CHANNEL_VERSION}
|
|||||||
ERRCODE_SESSION_EXPIRED = -14
|
ERRCODE_SESSION_EXPIRED = -14
|
||||||
SESSION_PAUSE_DURATION_S = 60 * 60
|
SESSION_PAUSE_DURATION_S = 60 * 60
|
||||||
|
|
||||||
# iLink context_token is observed to expire server-side after ~90-160s of
|
|
||||||
# agent inactivity (openclaw/openclaw#61174). Proactively refresh before
|
|
||||||
# sending if the cached token is older than this threshold.
|
|
||||||
CONTEXT_TOKEN_MAX_AGE_S = 60
|
|
||||||
|
|
||||||
|
|
||||||
# Retry constants (matching the reference plugin's monitor.ts)
|
# Retry constants (matching the reference plugin's monitor.ts)
|
||||||
MAX_CONSECUTIVE_FAILURES = 3
|
MAX_CONSECUTIVE_FAILURES = 3
|
||||||
BACKOFF_DELAY_S = 30
|
BACKOFF_DELAY_S = 30
|
||||||
@@ -165,8 +159,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
self._session_pause_until: float = 0.0
|
self._session_pause_until: float = 0.0
|
||||||
self._typing_tasks: dict[str, asyncio.Task] = {}
|
self._typing_tasks: dict[str, asyncio.Task] = {}
|
||||||
self._typing_tickets: dict[str, dict[str, Any]] = {}
|
self._typing_tickets: dict[str, dict[str, Any]] = {}
|
||||||
self._context_token_at: dict[str, float] = {}
|
|
||||||
self._pending_tool_hints: dict[str, list[str]] = {}
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# State persistence
|
# State persistence
|
||||||
@@ -494,7 +486,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
except Exception:
|
except Exception:
|
||||||
if not self._running:
|
if not self._running:
|
||||||
break
|
break
|
||||||
self.logger.exception("WeChat poll loop error")
|
|
||||||
consecutive_failures += 1
|
consecutive_failures += 1
|
||||||
if consecutive_failures >= MAX_CONSECUTIVE_FAILURES:
|
if consecutive_failures >= MAX_CONSECUTIVE_FAILURES:
|
||||||
consecutive_failures = 0
|
consecutive_failures = 0
|
||||||
@@ -504,7 +495,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
async def stop(self) -> None:
|
async def stop(self) -> None:
|
||||||
self._running = False
|
self._running = False
|
||||||
self._pending_tool_hints.clear()
|
|
||||||
if self._poll_task and not self._poll_task.done():
|
if self._poll_task and not self._poll_task.done():
|
||||||
self._poll_task.cancel()
|
self._poll_task.cancel()
|
||||||
for chat_id in list(self._typing_tasks):
|
for chat_id in list(self._typing_tasks):
|
||||||
@@ -555,7 +545,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
# Check for API-level errors (monitor.ts checks both ret and errcode)
|
# Check for API-level errors (monitor.ts checks both ret and errcode)
|
||||||
ret = data.get("ret", 0)
|
ret = data.get("ret", 0)
|
||||||
errcode = data.get("errcode", 0)
|
errcode = data.get("errcode", 0)
|
||||||
|
|
||||||
is_error = (ret is not None and ret != 0) or (errcode is not None and errcode != 0)
|
is_error = (ret is not None and ret != 0) or (errcode is not None and errcode != 0)
|
||||||
|
|
||||||
if is_error:
|
if is_error:
|
||||||
@@ -586,10 +575,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
# Process messages (WeixinMessage[] from types.ts)
|
# Process messages (WeixinMessage[] from types.ts)
|
||||||
msgs: list[dict] = data.get("msgs", []) or []
|
msgs: list[dict] = data.get("msgs", []) or []
|
||||||
for msg in msgs:
|
for msg in msgs:
|
||||||
try:
|
with suppress(Exception):
|
||||||
await self._process_message(msg)
|
await self._process_message(msg)
|
||||||
except Exception:
|
|
||||||
self.logger.exception("Failed to process WeChat message")
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Inbound message processing (matches inbound.ts + process-message.ts)
|
# Inbound message processing (matches inbound.ts + process-message.ts)
|
||||||
@@ -623,7 +610,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
ctx_token = msg.get("context_token", "")
|
ctx_token = msg.get("context_token", "")
|
||||||
if ctx_token:
|
if ctx_token:
|
||||||
self._context_tokens[from_user_id] = ctx_token
|
self._context_tokens[from_user_id] = ctx_token
|
||||||
self._context_token_at[from_user_id] = time.time()
|
|
||||||
self._save_state()
|
self._save_state()
|
||||||
|
|
||||||
# Parse item_list (WeixinMessage.item_list — types.ts:161)
|
# Parse item_list (WeixinMessage.item_list — types.ts:161)
|
||||||
@@ -929,99 +915,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
}
|
}
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
async def _refresh_context_token_if_stale(
|
|
||||||
self, chat_id: str, context_token: str
|
|
||||||
) -> str:
|
|
||||||
"""Return a fresh context_token if the cached one is too old.
|
|
||||||
|
|
||||||
iLink context_token expires server-side after a short idle period
|
|
||||||
(empirically ~90s). Proactively refreshing before sending prevents
|
|
||||||
silent message loss on long agent turns or cron pushes.
|
|
||||||
"""
|
|
||||||
if not context_token:
|
|
||||||
return context_token
|
|
||||||
|
|
||||||
now = time.time()
|
|
||||||
cached_at = self._context_token_at.get(chat_id, 0)
|
|
||||||
age = now - cached_at
|
|
||||||
|
|
||||||
if age < CONTEXT_TOKEN_MAX_AGE_S:
|
|
||||||
return context_token
|
|
||||||
|
|
||||||
self.logger.debug(
|
|
||||||
"WeChat context_token for {} is {:.0f}s old; refreshing via getconfig",
|
|
||||||
chat_id,
|
|
||||||
age,
|
|
||||||
)
|
|
||||||
|
|
||||||
body: dict[str, Any] = {
|
|
||||||
"ilink_user_id": chat_id,
|
|
||||||
"context_token": context_token,
|
|
||||||
"base_info": BASE_INFO,
|
|
||||||
}
|
|
||||||
try:
|
|
||||||
data = await self._api_post("ilink/bot/getconfig", body)
|
|
||||||
except Exception as e:
|
|
||||||
self.logger.warning("WeChat getconfig failed for {}: {}", chat_id, e)
|
|
||||||
return context_token
|
|
||||||
|
|
||||||
if data.get("ret", 0) != 0:
|
|
||||||
self.logger.warning(
|
|
||||||
"WeChat getconfig returned ret={} for {}: {}",
|
|
||||||
data.get("ret"),
|
|
||||||
chat_id,
|
|
||||||
data.get("errmsg", ""),
|
|
||||||
)
|
|
||||||
return context_token
|
|
||||||
|
|
||||||
new_token = str(data.get("context_token", "") or "")
|
|
||||||
if new_token and new_token != context_token:
|
|
||||||
self.logger.info(
|
|
||||||
"WeChat context_token refreshed for {} (age {:.0f}s -> fresh)",
|
|
||||||
chat_id,
|
|
||||||
age,
|
|
||||||
)
|
|
||||||
self._context_tokens[chat_id] = new_token
|
|
||||||
self._context_token_at[chat_id] = now
|
|
||||||
self._save_state()
|
|
||||||
return new_token
|
|
||||||
|
|
||||||
return context_token
|
|
||||||
|
|
||||||
async def _flush_tool_hints(self, chat_id: str) -> None:
|
|
||||||
"""Send any buffered tool hints for *chat_id* as a single message.
|
|
||||||
|
|
||||||
Tool hints are coalesced to reduce message count and avoid hitting the
|
|
||||||
WeChat iLink rate limit (~7 msgs / 5 min). Failures are logged but
|
|
||||||
not raised so that the main message send is never blocked.
|
|
||||||
"""
|
|
||||||
hints = self._pending_tool_hints.pop(chat_id, None)
|
|
||||||
if not hints:
|
|
||||||
return
|
|
||||||
|
|
||||||
self.logger.info(
|
|
||||||
"Flushing {} buffered tool hint(s) for {}",
|
|
||||||
len(hints),
|
|
||||||
chat_id,
|
|
||||||
)
|
|
||||||
|
|
||||||
ctx_token = self._context_tokens.get(chat_id, "")
|
|
||||||
ctx_token = await self._refresh_context_token_if_stale(chat_id, ctx_token)
|
|
||||||
if not ctx_token:
|
|
||||||
self.logger.warning(
|
|
||||||
"Dropped {} buffered tool hint(s) for {}: no context_token",
|
|
||||||
len(hints),
|
|
||||||
chat_id,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
try:
|
|
||||||
await self._send_text(chat_id, "\n\n".join(hints), ctx_token)
|
|
||||||
except Exception:
|
|
||||||
self.logger.exception(
|
|
||||||
"Failed to flush buffered tool hints for {}", chat_id
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _send_typing(self, user_id: str, typing_ticket: str, status: int) -> None:
|
async def _send_typing(self, user_id: str, typing_ticket: str, status: int) -> None:
|
||||||
"""Best-effort sendtyping wrapper."""
|
"""Best-effort sendtyping wrapper."""
|
||||||
if not typing_ticket:
|
if not typing_ticket:
|
||||||
@@ -1051,47 +944,11 @@ class WeixinChannel(BaseChannel):
|
|||||||
self._assert_session_active()
|
self._assert_session_active()
|
||||||
|
|
||||||
is_progress = bool((msg.metadata or {}).get("_progress", False))
|
is_progress = bool((msg.metadata or {}).get("_progress", False))
|
||||||
|
|
||||||
# Buffer tool hints to coalesce consecutive ones and avoid burning
|
|
||||||
# WeChat iLink rate-limit quota (~7 msgs / 5 min).
|
|
||||||
if is_progress and (msg.metadata or {}).get("_tool_hint"):
|
|
||||||
if not self.send_tool_hints:
|
|
||||||
return
|
|
||||||
self._pending_tool_hints.setdefault(msg.chat_id, []).append(msg.content)
|
|
||||||
self.logger.debug(
|
|
||||||
"Buffered tool hint for {} (count={})",
|
|
||||||
msg.chat_id,
|
|
||||||
len(self._pending_tool_hints[msg.chat_id]),
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
# Reasoning deltas are invisible in WeChat (there is no reasoning
|
|
||||||
# UI). Skip them entirely — do not send and do not flush buffer.
|
|
||||||
if is_progress and (msg.metadata or {}).get("_reasoning_delta"):
|
|
||||||
self.logger.debug(
|
|
||||||
"Dropped invisible reasoning delta for {}", msg.chat_id
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
content = msg.content.strip()
|
|
||||||
|
|
||||||
# Empty progress messages (e.g. after_iteration tool_events) must
|
|
||||||
# NOT act as separators — they have no visible content.
|
|
||||||
if is_progress and not content and not (msg.media or []):
|
|
||||||
self.logger.debug(
|
|
||||||
"Skipped empty progress message for {} (no visible content)",
|
|
||||||
msg.chat_id,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
# Flush buffered hints before sending any visible message.
|
|
||||||
await self._flush_tool_hints(msg.chat_id)
|
|
||||||
|
|
||||||
if not is_progress:
|
if not is_progress:
|
||||||
await self._stop_typing(msg.chat_id, clear_remote=True)
|
await self._stop_typing(msg.chat_id, clear_remote=True)
|
||||||
|
|
||||||
|
content = msg.content.strip()
|
||||||
ctx_token = self._context_tokens.get(msg.chat_id, "")
|
ctx_token = self._context_tokens.get(msg.chat_id, "")
|
||||||
ctx_token = await self._refresh_context_token_if_stale(msg.chat_id, ctx_token)
|
|
||||||
if not ctx_token:
|
if not ctx_token:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"WeChat context_token missing for chat_id={msg.chat_id}, cannot send"
|
f"WeChat context_token missing for chat_id={msg.chat_id}, cannot send"
|
||||||
@@ -1180,18 +1037,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
with suppress(Exception):
|
with suppress(Exception):
|
||||||
await self._send_typing(msg.chat_id, typing_ticket, TYPING_STATUS_CANCEL)
|
await self._send_typing(msg.chat_id, typing_ticket, TYPING_STATUS_CANCEL)
|
||||||
|
|
||||||
async def send_delta(
|
|
||||||
self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None
|
|
||||||
) -> None:
|
|
||||||
"""Weixin iLink does not support native streaming deltas.
|
|
||||||
|
|
||||||
We only hook ``_stream_end`` so buffered tool hints are flushed even
|
|
||||||
when the final answer carries the ``_streamed`` flag and bypasses
|
|
||||||
:meth:`send`.
|
|
||||||
"""
|
|
||||||
if metadata and metadata.get("_stream_end"):
|
|
||||||
await self._flush_tool_hints(chat_id)
|
|
||||||
|
|
||||||
async def _start_typing(self, chat_id: str, context_token: str = "") -> None:
|
async def _start_typing(self, chat_id: str, context_token: str = "") -> None:
|
||||||
"""Start typing indicator immediately when a message is received."""
|
"""Start typing indicator immediately when a message is received."""
|
||||||
if not self._client or not self._token or not chat_id:
|
if not self._client or not self._token or not chat_id:
|
||||||
@@ -1275,11 +1120,10 @@ class WeixinChannel(BaseChannel):
|
|||||||
}
|
}
|
||||||
|
|
||||||
data = await self._api_post("ilink/bot/sendmessage", body)
|
data = await self._api_post("ilink/bot/sendmessage", body)
|
||||||
ret = data.get("ret", 0)
|
|
||||||
errcode = data.get("errcode", 0)
|
errcode = data.get("errcode", 0)
|
||||||
if (ret is not None and ret != 0) or (errcode is not None and errcode != 0):
|
if errcode and errcode != 0:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"WeChat send text error (ret={ret}, errcode={errcode}): {data.get('errmsg', '')}"
|
f"WeChat send text error (code {errcode}): {data.get('errmsg', '')}"
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _send_media_file(
|
async def _send_media_file(
|
||||||
@@ -1426,11 +1270,10 @@ class WeixinChannel(BaseChannel):
|
|||||||
}
|
}
|
||||||
|
|
||||||
data = await self._api_post("ilink/bot/sendmessage", body)
|
data = await self._api_post("ilink/bot/sendmessage", body)
|
||||||
ret = data.get("ret", 0)
|
|
||||||
errcode = data.get("errcode", 0)
|
errcode = data.get("errcode", 0)
|
||||||
if (ret is not None and ret != 0) or (errcode is not None and errcode != 0):
|
if errcode and errcode != 0:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"WeChat send media error (ret={ret}, errcode={errcode}): {data.get('errmsg', '')}"
|
f"WeChat send media error (code {errcode}): {data.get('errmsg', '')}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+12
-6
@@ -620,7 +620,6 @@ 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:
|
||||||
@@ -640,7 +639,12 @@ def serve(
|
|||||||
agent_loop = AgentLoop.from_config(
|
agent_loop = AgentLoop.from_config(
|
||||||
runtime_config, bus,
|
runtime_config, bus,
|
||||||
session_manager=session_manager,
|
session_manager=session_manager,
|
||||||
image_generation_provider_configs=image_gen_provider_configs(runtime_config),
|
image_generation_provider_configs={
|
||||||
|
"openrouter": runtime_config.providers.openrouter,
|
||||||
|
"aihubmix": runtime_config.providers.aihubmix,
|
||||||
|
"minimax": runtime_config.providers.minimax,
|
||||||
|
"gemini": runtime_config.providers.gemini,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
console.print(f"[red]Error: {exc}[/red]")
|
console.print(f"[red]Error: {exc}[/red]")
|
||||||
@@ -720,7 +724,6 @@ 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
|
||||||
@@ -751,7 +754,12 @@ 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_gen_provider_configs(config),
|
image_generation_provider_configs={
|
||||||
|
"openrouter": config.providers.openrouter,
|
||||||
|
"aihubmix": config.providers.aihubmix,
|
||||||
|
"minimax": config.providers.minimax,
|
||||||
|
"gemini": config.providers.gemini,
|
||||||
|
},
|
||||||
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,
|
||||||
@@ -1118,7 +1126,6 @@ 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)
|
||||||
@@ -1142,7 +1149,6 @@ def agent(
|
|||||||
agent_loop = AgentLoop.from_config(
|
agent_loop = AgentLoop.from_config(
|
||||||
config, bus,
|
config, bus,
|
||||||
cron_service=cron,
|
cron_service=cron,
|
||||||
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]")
|
||||||
|
|||||||
@@ -1,13 +0,0 @@
|
|||||||
"""CLI Apps integration helpers."""
|
|
||||||
|
|
||||||
from nanobot.cli_apps.service import (
|
|
||||||
CliAppError,
|
|
||||||
CliAppManager,
|
|
||||||
CliAppsRuntimeConfig,
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"CliAppError",
|
|
||||||
"CliAppManager",
|
|
||||||
"CliAppsRuntimeConfig",
|
|
||||||
]
|
|
||||||
@@ -1,955 +0,0 @@
|
|||||||
"""CLI-Anything catalog, install state, and safe CLI execution."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
|
||||||
import os
|
|
||||||
import re
|
|
||||||
import shlex
|
|
||||||
import shutil
|
|
||||||
import subprocess
|
|
||||||
import sys
|
|
||||||
import time
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
|
||||||
from urllib.parse import urlparse
|
|
||||||
|
|
||||||
import httpx
|
|
||||||
|
|
||||||
from nanobot.config.paths import get_runtime_subdir
|
|
||||||
|
|
||||||
CLI_ANYTHING_REGISTRY_URL = "https://hkuds.github.io/CLI-Anything/registry.json"
|
|
||||||
CLI_ANYTHING_PUBLIC_REGISTRY_URL = "https://hkuds.github.io/CLI-Anything/public_registry.json"
|
|
||||||
CLI_ANYTHING_RAW_BASE = "https://raw.githubusercontent.com/HKUDS/CLI-Anything/main"
|
|
||||||
CLI_ANYTHING_RAW_SKILLS_BASE = f"{CLI_ANYTHING_RAW_BASE}/skills/"
|
|
||||||
|
|
||||||
_MAX_TOOL_OUTPUT_CHARS = 12_000
|
|
||||||
_MAX_ARTIFACT_SCAN_PATHS = 4_000
|
|
||||||
_MAX_ARTIFACT_REPORT = 12
|
|
||||||
_SAFE_NAME_RE = re.compile(r"[^a-z0-9_-]+")
|
|
||||||
_MENTION_RE = re.compile(r"(^|[\s([{])@([a-z0-9_-]+)\b", re.IGNORECASE)
|
|
||||||
_SHELL_META_CHARS = ("|", "&&", "||", ";", "$(", "`", ">", "<")
|
|
||||||
_ARTIFACT_EXTENSIONS = frozenset({
|
|
||||||
".csv",
|
|
||||||
".drawio",
|
|
||||||
".gif",
|
|
||||||
".html",
|
|
||||||
".jpeg",
|
|
||||||
".jpg",
|
|
||||||
".json",
|
|
||||||
".md",
|
|
||||||
".pdf",
|
|
||||||
".png",
|
|
||||||
".svg",
|
|
||||||
".txt",
|
|
||||||
".vsdx",
|
|
||||||
".webp",
|
|
||||||
".xml",
|
|
||||||
})
|
|
||||||
_INLINE_ARTIFACT_EXTENSIONS = frozenset({".gif", ".jpeg", ".jpg", ".png", ".webp"})
|
|
||||||
_ARTIFACT_IGNORE_DIRS = frozenset({
|
|
||||||
".git",
|
|
||||||
".hg",
|
|
||||||
".mypy_cache",
|
|
||||||
".nanobot",
|
|
||||||
".pytest_cache",
|
|
||||||
".ruff_cache",
|
|
||||||
".venv",
|
|
||||||
"__pycache__",
|
|
||||||
"build",
|
|
||||||
"dist",
|
|
||||||
"node_modules",
|
|
||||||
"venv",
|
|
||||||
})
|
|
||||||
|
|
||||||
|
|
||||||
class CliAppError(ValueError):
|
|
||||||
"""User-facing CLI Apps failure."""
|
|
||||||
|
|
||||||
def __init__(self, message: str, *, status: int = 400) -> None:
|
|
||||||
super().__init__(message)
|
|
||||||
self.message = message
|
|
||||||
self.status = status
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
|
||||||
class CliAppsRuntimeConfig:
|
|
||||||
"""Runtime knobs for CLI Apps."""
|
|
||||||
|
|
||||||
install_timeout: int = 300
|
|
||||||
run_timeout: int = 60
|
|
||||||
catalog_ttl_seconds: int = 3600
|
|
||||||
|
|
||||||
|
|
||||||
_BRANDS: dict[str, tuple[str, str]] = {
|
|
||||||
"1password-cli": ("1password", "#3B66BC"),
|
|
||||||
"audacity": ("audacity", "#0000CC"),
|
|
||||||
"blender": ("blender", "#E87D0D"),
|
|
||||||
"browser": ("googlechrome", "#4285F4"),
|
|
||||||
"calibre": ("calibre", "#45B29D"),
|
|
||||||
"chromadb": ("chroma", "#FFDE2D"),
|
|
||||||
"comfyui": ("comfyui", "#111827"),
|
|
||||||
"contentful": ("contentful", "#2478CC"),
|
|
||||||
"dify": ("dify", "#155EEF"),
|
|
||||||
"drawio": ("diagramsdotnet", "#F08705"),
|
|
||||||
"elevenlabs": ("elevenlabs", "#000000"),
|
|
||||||
"eth2-quickstart": ("ethereum", "#627EEA"),
|
|
||||||
"firefly-iii": ("fireflyiii", "#CD5029"),
|
|
||||||
"freecad": ("freecad", "#418FDE"),
|
|
||||||
"generate-veo-video": ("googlegemini", "#8E75B2"),
|
|
||||||
"gimp": ("gimp", "#5C5543"),
|
|
||||||
"godot": ("godotengine", "#478CBF"),
|
|
||||||
"hacker-feeds-cli": ("rss", "#FFA500"),
|
|
||||||
"inkscape": ("inkscape", "#000000"),
|
|
||||||
"intelwatch": ("intel", "#0071C5"),
|
|
||||||
"iterm2": ("iterm2", "#000000"),
|
|
||||||
"jimeng": ("bytedance", "#3C8CFF"),
|
|
||||||
"kdenlive": ("kdenlive", "#527EB2"),
|
|
||||||
"krita": ("krita", "#3BABFF"),
|
|
||||||
"libreoffice": ("libreoffice", "#18A303"),
|
|
||||||
"mailchimp": ("mailchimp", "#FFE01B"),
|
|
||||||
"mermaid": ("mermaid", "#FF3670"),
|
|
||||||
"minimax": ("minimax", "#111827"),
|
|
||||||
"musescore": ("musescore", "#1A70B8"),
|
|
||||||
"n8n": ("n8n", "#EA4B71"),
|
|
||||||
"notebooklm": ("googlenotebooklm", "#4285F4"),
|
|
||||||
"obs-studio": ("obsstudio", "#302E31"),
|
|
||||||
"obsidian": ("obsidian", "#7C3AED"),
|
|
||||||
"ollama": ("ollama", "#000000"),
|
|
||||||
"pm2": ("pm2", "#2B037A"),
|
|
||||||
"qgis": ("qgis", "#589632"),
|
|
||||||
"safari": ("safari", "#006CFF"),
|
|
||||||
"sanity": ("sanity", "#F03E2F"),
|
|
||||||
"sentry": ("sentry", "#362D59"),
|
|
||||||
"sketch": ("sketch", "#F7B500"),
|
|
||||||
"shopify": ("shopify", "#7AB55C"),
|
|
||||||
"nsight-graphics": ("nvidia", "#76B900"),
|
|
||||||
"unrealinsights": ("unrealengine", "#0E1128"),
|
|
||||||
"ueatelier": ("unrealengine", "#0E1128"),
|
|
||||||
"ve-twini": ("x", "#000000"),
|
|
||||||
"wecom": ("wechat", "#07C160"),
|
|
||||||
"suno": ("suno", "#000000"),
|
|
||||||
"lldb": ("llvm", "#262D3A"),
|
|
||||||
"android-cli": ("android", "#3DDC84"),
|
|
||||||
"adguardhome": ("adguard", "#68BC71"),
|
|
||||||
"zotero": ("zotero", "#CC2936"),
|
|
||||||
"zoom": ("zoom", "#0B5CFF"),
|
|
||||||
}
|
|
||||||
|
|
||||||
_BRAND_DOMAINS: dict[str, tuple[str, str]] = {
|
|
||||||
"3mf": ("3mf.io", "#00A1DE"),
|
|
||||||
"anygen": ("anygen.com", "#111827"),
|
|
||||||
"clibrowser": ("github.com/allthingssecurity/clibrowser", "#24292F"),
|
|
||||||
"cloudanalyzer": ("github.com/rsasaki0109/CloudAnalyzer", "#2563EB"),
|
|
||||||
"cloudcompare": ("cloudcompare.org", "#4D83C3"),
|
|
||||||
"deployhq": ("deployhq.com", "#00A2D9"),
|
|
||||||
"exa": ("exa.ai", "#111827"),
|
|
||||||
"feishu": ("larksuite.com", "#00A5FF"),
|
|
||||||
"inkstitch": ("inkstitch.org", "#222222"),
|
|
||||||
"macrocli": ("github.com/HKUDS/CLI-Anything/tree/main/macrocli", "#24292F"),
|
|
||||||
"mubu": ("mubu.com", "#16A085"),
|
|
||||||
"nslogger": ("github.com/fpillet/NSLogger", "#24292F"),
|
|
||||||
"novita": ("novita.ai", "#7C3AED"),
|
|
||||||
"openscreen": ("openscreen.com", "#2563EB"),
|
|
||||||
"py4csr": ("github.com/yanmingyu92/py4csr", "#24292F"),
|
|
||||||
"quietshrink": ("github.com/achiya-automation/quietshrink", "#111827"),
|
|
||||||
"renderdoc": ("renderdoc.org", "#2C7DB8"),
|
|
||||||
"rms": ("rms.teltonika-networks.com", "#0054A6"),
|
|
||||||
"sbox": ("sbox.game", "#F59E0B"),
|
|
||||||
"seaclip": ("github.com/SeaClip-Lite/SeaClip", "#0284C7"),
|
|
||||||
"shotcut": ("shotcut.org", "#3B82F6"),
|
|
||||||
"slay-the-spire-ii": ("megacrit.com", "#B91C1C"),
|
|
||||||
"stata": ("stata.com", "#1F4E79"),
|
|
||||||
"unimol-tools": ("github.com/deepmodeling/Uni-Mol", "#4F46E5"),
|
|
||||||
"videocaptioner": ("github.com/WEIFENG2333/VideoCaptioner", "#2563EB"),
|
|
||||||
"wiremock": ("wiremock.org", "#FF6A00"),
|
|
||||||
}
|
|
||||||
|
|
||||||
_BRAND_ALIASES: dict[str, str] = {
|
|
||||||
"1password": "1password-cli",
|
|
||||||
"dify-workflow": "dify",
|
|
||||||
"feishu-lark": "feishu",
|
|
||||||
"lark-cli": "feishu",
|
|
||||||
"minimax-cli": "minimax",
|
|
||||||
"obsidian-cli": "obsidian",
|
|
||||||
"slay-the-spire-2": "slay-the-spire-ii",
|
|
||||||
"slay-the-spire-ii": "slay-the-spire-ii",
|
|
||||||
"unimol-tools": "unimol-tools",
|
|
||||||
"unimol": "unimol-tools",
|
|
||||||
"veo": "generate-veo-video",
|
|
||||||
}
|
|
||||||
|
|
||||||
_BRAND_TRAILING_WORDS = ("cli", "workflow", "workflows", "app", "apps", "tool", "tools")
|
|
||||||
|
|
||||||
|
|
||||||
def _now() -> float:
|
|
||||||
return time.time()
|
|
||||||
|
|
||||||
|
|
||||||
def _safe_skill_name(name: str) -> str:
|
|
||||||
clean = _SAFE_NAME_RE.sub("-", name.lower()).strip("-")
|
|
||||||
return f"cli-app-{clean or 'app'}"
|
|
||||||
|
|
||||||
|
|
||||||
def _has_shell_meta(command: str) -> bool:
|
|
||||||
return any(char in command for char in _SHELL_META_CHARS)
|
|
||||||
|
|
||||||
|
|
||||||
def _command_exists(command: str) -> bool:
|
|
||||||
try:
|
|
||||||
parts = shlex.split(command)
|
|
||||||
except ValueError:
|
|
||||||
return False
|
|
||||||
if not parts:
|
|
||||||
return False
|
|
||||||
return shutil.which(parts[0]) is not None
|
|
||||||
|
|
||||||
|
|
||||||
def _is_pip_install_command(command: str) -> bool:
|
|
||||||
try:
|
|
||||||
tokens = shlex.split(command)
|
|
||||||
except ValueError:
|
|
||||||
return False
|
|
||||||
return (
|
|
||||||
len(tokens) >= 3
|
|
||||||
and tokens[:2] == ["pip", "install"]
|
|
||||||
) or (
|
|
||||||
len(tokens) >= 5
|
|
||||||
and tokens[1:4] == ["-m", "pip", "install"]
|
|
||||||
and tokens[0] in {"python", "python3", sys.executable}
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _pip_uninstall_args_from_command(command: str) -> list[str] | None:
|
|
||||||
if not command or _has_shell_meta(command):
|
|
||||||
return None
|
|
||||||
try:
|
|
||||||
tokens = shlex.split(command)
|
|
||||||
except ValueError:
|
|
||||||
return None
|
|
||||||
if tokens[:2] == ["pip", "uninstall"]:
|
|
||||||
args = tokens[2:]
|
|
||||||
elif (
|
|
||||||
len(tokens) >= 5
|
|
||||||
and tokens[1:4] == ["-m", "pip", "uninstall"]
|
|
||||||
and tokens[0] in {"python", "python3", sys.executable}
|
|
||||||
):
|
|
||||||
args = tokens[4:]
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
packages = [arg for arg in args if arg not in {"-y", "--yes"}]
|
|
||||||
if not packages or any(arg.startswith("-") for arg in packages):
|
|
||||||
return None
|
|
||||||
return packages
|
|
||||||
|
|
||||||
|
|
||||||
def _brand_key(value: str) -> str:
|
|
||||||
return _SAFE_NAME_RE.sub("-", value.lower()).replace("_", "-").strip("-")
|
|
||||||
|
|
||||||
|
|
||||||
def _brand_candidates(app: dict[str, Any]) -> list[str]:
|
|
||||||
values = [
|
|
||||||
str(app.get("name") or ""),
|
|
||||||
str(app.get("display_name") or ""),
|
|
||||||
str(app.get("entry_point") or "").removeprefix("cli-anything-"),
|
|
||||||
]
|
|
||||||
seen: set[str] = set()
|
|
||||||
candidates: list[str] = []
|
|
||||||
for value in values:
|
|
||||||
key = _brand_key(value)
|
|
||||||
while key and key not in seen:
|
|
||||||
seen.add(key)
|
|
||||||
candidates.append(key)
|
|
||||||
parts = key.split("-")
|
|
||||||
if len(parts) <= 1 or parts[-1] not in _BRAND_TRAILING_WORDS:
|
|
||||||
break
|
|
||||||
key = "-".join(parts[:-1])
|
|
||||||
return candidates
|
|
||||||
|
|
||||||
|
|
||||||
def _brand_payload(app: dict[str, Any]) -> tuple[str | None, str | None]:
|
|
||||||
brand = None
|
|
||||||
domain_brand = None
|
|
||||||
for candidate in _brand_candidates(app):
|
|
||||||
key = _BRAND_ALIASES.get(candidate, candidate)
|
|
||||||
brand = _BRANDS.get(key)
|
|
||||||
if brand:
|
|
||||||
break
|
|
||||||
domain_brand = _BRAND_DOMAINS.get(key)
|
|
||||||
if domain_brand:
|
|
||||||
break
|
|
||||||
if not brand:
|
|
||||||
if not domain_brand:
|
|
||||||
return None, None
|
|
||||||
domain, color = domain_brand
|
|
||||||
return f"https://www.google.com/s2/favicons?domain={domain}&sz=64", color
|
|
||||||
slug, color = brand
|
|
||||||
return f"https://cdn.simpleicons.org/{slug}/{color.lstrip('#')}", color
|
|
||||||
|
|
||||||
|
|
||||||
def _read_json(path: Path) -> dict[str, Any] | None:
|
|
||||||
try:
|
|
||||||
data = json.loads(path.read_text(encoding="utf-8"))
|
|
||||||
except (OSError, json.JSONDecodeError):
|
|
||||||
return None
|
|
||||||
return data if isinstance(data, dict) else None
|
|
||||||
|
|
||||||
|
|
||||||
def _write_json(path: Path, data: dict[str, Any]) -> None:
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
payload = json.dumps(data, indent=2, ensure_ascii=False)
|
|
||||||
tmp_path = path.with_name(f".{path.name}.{os.getpid()}.{int(_now() * 1_000_000)}.tmp")
|
|
||||||
try:
|
|
||||||
tmp_path.write_text(payload, encoding="utf-8")
|
|
||||||
tmp_path.replace(path)
|
|
||||||
finally:
|
|
||||||
if tmp_path.exists():
|
|
||||||
tmp_path.unlink()
|
|
||||||
|
|
||||||
|
|
||||||
def _safe_skill_path(value: str) -> str | None:
|
|
||||||
if not value.startswith("skills/"):
|
|
||||||
return None
|
|
||||||
parts = value.split("/")
|
|
||||||
if any(part in {"", ".", ".."} for part in parts):
|
|
||||||
return None
|
|
||||||
return value if parts[-1] == "SKILL.md" else None
|
|
||||||
|
|
||||||
|
|
||||||
def _skill_content_url(skill_md: str) -> str | None:
|
|
||||||
safe_path = _safe_skill_path(skill_md)
|
|
||||||
if safe_path:
|
|
||||||
return f"{CLI_ANYTHING_RAW_BASE}/{safe_path}"
|
|
||||||
parsed = urlparse(skill_md)
|
|
||||||
if parsed.scheme != "https" or parsed.netloc != "raw.githubusercontent.com":
|
|
||||||
return None
|
|
||||||
if not skill_md.startswith(CLI_ANYTHING_RAW_SKILLS_BASE):
|
|
||||||
return None
|
|
||||||
suffix = skill_md.removeprefix(f"{CLI_ANYTHING_RAW_BASE}/")
|
|
||||||
return skill_md if _safe_skill_path(suffix) else None
|
|
||||||
|
|
||||||
|
|
||||||
def _truncate(text: str, limit: int = _MAX_TOOL_OUTPUT_CHARS) -> str:
|
|
||||||
if len(text) <= limit:
|
|
||||||
return text
|
|
||||||
omitted = len(text) - limit
|
|
||||||
return text[:limit] + f"\n\n... truncated {omitted} characters ..."
|
|
||||||
|
|
||||||
|
|
||||||
class CliAppManager:
|
|
||||||
"""Manage CLI-Anything registry entries and local install state."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
workspace: Path,
|
|
||||||
data_dir: Path | None = None,
|
|
||||||
runtime: CliAppsRuntimeConfig | None = None,
|
|
||||||
) -> None:
|
|
||||||
self.workspace = Path(workspace).expanduser()
|
|
||||||
self.data_dir = Path(data_dir) if data_dir is not None else get_runtime_subdir("cli-apps")
|
|
||||||
self.runtime = runtime or CliAppsRuntimeConfig()
|
|
||||||
|
|
||||||
@property
|
|
||||||
def installed_path(self) -> Path:
|
|
||||||
return self.data_dir / "installed.json"
|
|
||||||
|
|
||||||
def _cache_path(self, source: str) -> Path:
|
|
||||||
return self.data_dir / f"{source}_registry_cache.json"
|
|
||||||
|
|
||||||
def _load_installed(self) -> dict[str, Any]:
|
|
||||||
data = _read_json(self.installed_path) or {}
|
|
||||||
apps = data.get("apps") if isinstance(data.get("apps"), dict) else data
|
|
||||||
return apps if isinstance(apps, dict) else {}
|
|
||||||
|
|
||||||
def _save_installed(self, installed: dict[str, Any]) -> None:
|
|
||||||
_write_json(self.installed_path, {"schema_version": 1, "apps": installed})
|
|
||||||
|
|
||||||
def installed_names(self) -> list[str]:
|
|
||||||
"""Return registry names explicitly installed through CLI Apps."""
|
|
||||||
return sorted(str(name) for name in self._load_installed())
|
|
||||||
|
|
||||||
def _fetch_registry(
|
|
||||||
self,
|
|
||||||
url: str,
|
|
||||||
cache_path: Path,
|
|
||||||
*,
|
|
||||||
force_refresh: bool = False,
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
cached = _read_json(cache_path)
|
|
||||||
if (
|
|
||||||
not force_refresh
|
|
||||||
and cached
|
|
||||||
and _now() - float(cached.get("_cached_at", 0)) < self.runtime.catalog_ttl_seconds
|
|
||||||
):
|
|
||||||
data = cached.get("data")
|
|
||||||
if isinstance(data, dict):
|
|
||||||
return data
|
|
||||||
|
|
||||||
try:
|
|
||||||
response = httpx.get(url, timeout=15.0, follow_redirects=True)
|
|
||||||
response.raise_for_status()
|
|
||||||
data = response.json()
|
|
||||||
if not isinstance(data, dict):
|
|
||||||
raise ValueError("registry response must be an object")
|
|
||||||
except Exception:
|
|
||||||
if cached and isinstance(cached.get("data"), dict):
|
|
||||||
return cached["data"]
|
|
||||||
raise
|
|
||||||
|
|
||||||
_write_json(cache_path, {"_cached_at": _now(), "data": data})
|
|
||||||
return data
|
|
||||||
|
|
||||||
def catalog(self, *, force_refresh: bool = False) -> tuple[list[dict[str, Any]], str | None]:
|
|
||||||
registries = [
|
|
||||||
(
|
|
||||||
"harness",
|
|
||||||
self._fetch_registry(
|
|
||||||
CLI_ANYTHING_REGISTRY_URL,
|
|
||||||
self._cache_path("harness"),
|
|
||||||
force_refresh=force_refresh,
|
|
||||||
),
|
|
||||||
),
|
|
||||||
(
|
|
||||||
"public",
|
|
||||||
self._fetch_registry(
|
|
||||||
CLI_ANYTHING_PUBLIC_REGISTRY_URL,
|
|
||||||
self._cache_path("public"),
|
|
||||||
force_refresh=force_refresh,
|
|
||||||
),
|
|
||||||
),
|
|
||||||
]
|
|
||||||
apps_by_name: dict[str, dict[str, Any]] = {}
|
|
||||||
updated_values: list[str] = []
|
|
||||||
for source, registry in registries:
|
|
||||||
meta = registry.get("meta")
|
|
||||||
if isinstance(meta, dict) and isinstance(meta.get("updated"), str):
|
|
||||||
updated_values.append(meta["updated"])
|
|
||||||
for row in registry.get("clis", []):
|
|
||||||
if not isinstance(row, dict) or not row.get("name"):
|
|
||||||
continue
|
|
||||||
entry = dict(row)
|
|
||||||
entry["_source"] = source
|
|
||||||
key = str(entry["name"]).lower()
|
|
||||||
previous = apps_by_name.get(key)
|
|
||||||
if previous:
|
|
||||||
previous_source = str(previous.get("_source") or source)
|
|
||||||
merged_source = (
|
|
||||||
previous_source if previous_source == source else f"{previous_source}+{source}"
|
|
||||||
)
|
|
||||||
apps_by_name[key] = {**previous, **entry, "_source": merged_source}
|
|
||||||
else:
|
|
||||||
apps_by_name[key] = entry
|
|
||||||
return list(apps_by_name.values()), max(updated_values) if updated_values else None
|
|
||||||
|
|
||||||
def get_app(self, name: str, *, force_refresh: bool = False) -> dict[str, Any]:
|
|
||||||
wanted = name.lower()
|
|
||||||
for app in self.catalog(force_refresh=force_refresh)[0]:
|
|
||||||
if str(app.get("name", "")).lower() == wanted:
|
|
||||||
return app
|
|
||||||
raise CliAppError(f"CLI app '{name}' not found", status=404)
|
|
||||||
|
|
||||||
def mentioned_installed_apps(self, text: str) -> list[dict[str, str]]:
|
|
||||||
"""Return installed CLI Apps referenced as ``@name`` in user text."""
|
|
||||||
if "@" not in text:
|
|
||||||
return []
|
|
||||||
installed = self._load_installed()
|
|
||||||
if not installed:
|
|
||||||
return []
|
|
||||||
installed_by_name = {
|
|
||||||
str(name).lower(): (str(name), data if isinstance(data, dict) else {})
|
|
||||||
for name, data in installed.items()
|
|
||||||
}
|
|
||||||
seen: set[str] = set()
|
|
||||||
mentions: list[dict[str, str]] = []
|
|
||||||
for match in _MENTION_RE.finditer(text):
|
|
||||||
wanted = str(match.group(2)).lower()
|
|
||||||
if wanted in seen or wanted not in installed_by_name:
|
|
||||||
continue
|
|
||||||
installed_name, data = installed_by_name[wanted]
|
|
||||||
seen.add(wanted)
|
|
||||||
entry_point = str(data.get("entry_point") or "")
|
|
||||||
mentions.append(
|
|
||||||
{
|
|
||||||
"name": installed_name,
|
|
||||||
"entry_point": entry_point,
|
|
||||||
"source": str(data.get("source") or ""),
|
|
||||||
"skill": f"skills/{_safe_skill_name(installed_name)}/SKILL.md",
|
|
||||||
"tool": "run_cli_app",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
return mentions
|
|
||||||
|
|
||||||
def _strategy(self, app: dict[str, Any]) -> str:
|
|
||||||
package_manager = str(app.get("package_manager") or "").lower()
|
|
||||||
install_strategy = str(app.get("install_strategy") or "").lower()
|
|
||||||
if package_manager == "bundled" or install_strategy == "bundled":
|
|
||||||
return "bundled"
|
|
||||||
if package_manager in {"npm", "brew", "uv", "pip"}:
|
|
||||||
return package_manager
|
|
||||||
if app.get("npm_package"):
|
|
||||||
return "npm"
|
|
||||||
install_cmd = str(app.get("install_cmd") or "")
|
|
||||||
if _is_pip_install_command(install_cmd):
|
|
||||||
return "pip"
|
|
||||||
return "unsupported"
|
|
||||||
|
|
||||||
def _install_supported(self, app: dict[str, Any]) -> bool:
|
|
||||||
if self._strategy(app) == "unsupported":
|
|
||||||
return False
|
|
||||||
install_cmd = str(app.get("install_cmd") or "")
|
|
||||||
return not _has_shell_meta(install_cmd)
|
|
||||||
|
|
||||||
def _skill_path(self, name: str) -> Path:
|
|
||||||
return self.workspace / "skills" / _safe_skill_name(name) / "SKILL.md"
|
|
||||||
|
|
||||||
def _app_payload(
|
|
||||||
self,
|
|
||||||
app: dict[str, Any],
|
|
||||||
installed: dict[str, Any],
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
name = str(app["name"])
|
|
||||||
entry_point = str(app.get("entry_point") or "")
|
|
||||||
install_supported = self._install_supported(app)
|
|
||||||
is_installed = name in installed
|
|
||||||
available = bool(entry_point and shutil.which(entry_point))
|
|
||||||
if is_installed and available:
|
|
||||||
status = "installed"
|
|
||||||
elif is_installed:
|
|
||||||
status = "missing"
|
|
||||||
elif not install_supported:
|
|
||||||
status = "unsupported"
|
|
||||||
elif available:
|
|
||||||
status = "available"
|
|
||||||
else:
|
|
||||||
status = "not_installed"
|
|
||||||
logo_url, brand_color = _brand_payload(app)
|
|
||||||
return {
|
|
||||||
"name": name,
|
|
||||||
"display_name": app.get("display_name") or name,
|
|
||||||
"category": app.get("category") or "uncategorized",
|
|
||||||
"description": app.get("description") or "",
|
|
||||||
"requires": app.get("requires") or "",
|
|
||||||
"source": app.get("_source") or "harness",
|
|
||||||
"entry_point": entry_point,
|
|
||||||
"install_supported": install_supported,
|
|
||||||
"installed": is_installed,
|
|
||||||
"available": available,
|
|
||||||
"status": status,
|
|
||||||
"logo_url": logo_url,
|
|
||||||
"brand_color": brand_color,
|
|
||||||
"skill_installed": self._skill_path(name).is_file(),
|
|
||||||
}
|
|
||||||
|
|
||||||
def payload(self, *, force_refresh: bool = False) -> dict[str, Any]:
|
|
||||||
apps, updated = self.catalog(force_refresh=force_refresh)
|
|
||||||
installed = self._load_installed()
|
|
||||||
rows = [self._app_payload(app, installed) for app in apps]
|
|
||||||
rows.sort(key=lambda item: (str(item["category"]), str(item["display_name"]).lower()))
|
|
||||||
return {
|
|
||||||
"apps": rows,
|
|
||||||
"installed_count": sum(1 for item in rows if item["installed"]),
|
|
||||||
"catalog_updated_at": updated,
|
|
||||||
}
|
|
||||||
|
|
||||||
def _pip_package_from_install(self, app: dict[str, Any]) -> str | None:
|
|
||||||
install_cmd = str(app.get("install_cmd") or "")
|
|
||||||
try:
|
|
||||||
tokens = shlex.split(install_cmd)
|
|
||||||
except ValueError:
|
|
||||||
return None
|
|
||||||
if tokens[:2] == ["pip", "install"]:
|
|
||||||
args = tokens[2:]
|
|
||||||
elif len(tokens) >= 5 and tokens[1:4] == ["-m", "pip", "install"]:
|
|
||||||
args = tokens[4:]
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
args = [arg for arg in args if not arg.startswith("-")]
|
|
||||||
if len(args) != 1 or args[0].startswith("git+"):
|
|
||||||
return None
|
|
||||||
return args[0]
|
|
||||||
|
|
||||||
def _pip_install_argv(self, app: dict[str, Any], *, update: bool = False) -> list[str]:
|
|
||||||
install_cmd = str(app.get("install_cmd") or "")
|
|
||||||
if not _is_pip_install_command(install_cmd) or _has_shell_meta(install_cmd):
|
|
||||||
raise CliAppError("unsupported pip install command")
|
|
||||||
tokens = shlex.split(install_cmd)
|
|
||||||
args = tokens[2:] if tokens[:2] == ["pip", "install"] else tokens[4:]
|
|
||||||
prefix = [sys.executable, "-m", "pip", "install"]
|
|
||||||
if update:
|
|
||||||
prefix.extend(["--upgrade", "--force-reinstall"])
|
|
||||||
return prefix + args
|
|
||||||
|
|
||||||
def _pip_uninstall_argv(self, app: dict[str, Any]) -> list[str]:
|
|
||||||
uninstall_cmd = str(app.get("uninstall_cmd") or "")
|
|
||||||
packages = _pip_uninstall_args_from_command(uninstall_cmd)
|
|
||||||
if packages:
|
|
||||||
return [sys.executable, "-m", "pip", "uninstall", "-y", *packages]
|
|
||||||
package = str(app.get("pip_package") or "").strip() or self._pip_package_from_install(app)
|
|
||||||
if not package:
|
|
||||||
entry_point = str(app.get("entry_point") or "").strip()
|
|
||||||
package = entry_point if entry_point.startswith("cli-anything-") else f"cli-anything-{_brand_key(str(app['name']))}"
|
|
||||||
return [sys.executable, "-m", "pip", "uninstall", "-y", package]
|
|
||||||
|
|
||||||
def _npm_argv(self, app: dict[str, Any], action: str) -> list[str]:
|
|
||||||
npm = shutil.which("npm")
|
|
||||||
if not npm:
|
|
||||||
raise CliAppError("npm is not installed")
|
|
||||||
package = str(app.get("npm_package") or "")
|
|
||||||
if not package:
|
|
||||||
raise CliAppError("registry entry has no npm_package")
|
|
||||||
if action == "install":
|
|
||||||
return [npm, "install", "-g", package]
|
|
||||||
if action == "update":
|
|
||||||
return [npm, "install", "-g", package + "@latest"]
|
|
||||||
return [npm, "uninstall", "-g", package]
|
|
||||||
|
|
||||||
def _split_safe_command(self, app: dict[str, Any], key: str, expected: str) -> list[str]:
|
|
||||||
command = str(app.get(key) or "")
|
|
||||||
if not command:
|
|
||||||
raise CliAppError(f"no {key} is defined for {app['name']}")
|
|
||||||
if _has_shell_meta(command):
|
|
||||||
raise CliAppError("script-style install commands are disabled in this MVP")
|
|
||||||
try:
|
|
||||||
argv = shlex.split(command)
|
|
||||||
except ValueError as exc:
|
|
||||||
raise CliAppError(f"invalid command: {exc}") from exc
|
|
||||||
if not argv or argv[0] != expected:
|
|
||||||
raise CliAppError(f"unsupported {expected} command")
|
|
||||||
return argv
|
|
||||||
|
|
||||||
def _argv_for_action(self, app: dict[str, Any], action: str) -> list[str] | None:
|
|
||||||
strategy = self._strategy(app)
|
|
||||||
if strategy == "pip":
|
|
||||||
if action == "install":
|
|
||||||
return self._pip_install_argv(app)
|
|
||||||
if action == "update":
|
|
||||||
return self._pip_install_argv(app, update=True)
|
|
||||||
return self._pip_uninstall_argv(app)
|
|
||||||
if strategy == "npm":
|
|
||||||
return self._npm_argv(app, action)
|
|
||||||
if strategy == "brew":
|
|
||||||
key = {"install": "install_cmd", "update": "update_cmd", "uninstall": "uninstall_cmd"}[action]
|
|
||||||
return self._split_safe_command(app, key, "brew")
|
|
||||||
if strategy == "uv":
|
|
||||||
key = {"install": "install_cmd", "update": "update_cmd", "uninstall": "uninstall_cmd"}[action]
|
|
||||||
return self._split_safe_command(app, key, "uv")
|
|
||||||
if strategy == "bundled":
|
|
||||||
return None
|
|
||||||
raise CliAppError("this CLI app uses an unsupported install strategy")
|
|
||||||
|
|
||||||
def _run_argv(self, argv: list[str], *, timeout: int) -> subprocess.CompletedProcess[str]:
|
|
||||||
return subprocess.run(
|
|
||||||
argv,
|
|
||||||
capture_output=True,
|
|
||||||
text=True,
|
|
||||||
timeout=timeout,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _installed_entry(self, app: dict[str, Any]) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"version": app.get("version") or "unknown",
|
|
||||||
"entry_point": app.get("entry_point") or "",
|
|
||||||
"source": app.get("_source") or "harness",
|
|
||||||
"strategy": self._strategy(app),
|
|
||||||
"installed_at": int(_now()),
|
|
||||||
}
|
|
||||||
|
|
||||||
def _fetch_skill_content(self, app: dict[str, Any]) -> str | None:
|
|
||||||
skill_md = str(app.get("skill_md") or "").strip()
|
|
||||||
if not skill_md:
|
|
||||||
return None
|
|
||||||
url = _skill_content_url(skill_md)
|
|
||||||
if not url:
|
|
||||||
return None
|
|
||||||
try:
|
|
||||||
response = httpx.get(url, timeout=15.0, follow_redirects=True)
|
|
||||||
response.raise_for_status()
|
|
||||||
text = response.text
|
|
||||||
except Exception:
|
|
||||||
return None
|
|
||||||
if "SKILL.md" not in url and not text.lstrip().startswith("---"):
|
|
||||||
return None
|
|
||||||
return text if len(text) < 250_000 else None
|
|
||||||
|
|
||||||
def _fallback_skill(self, app: dict[str, Any]) -> str:
|
|
||||||
name = str(app.get("name") or "unknown")
|
|
||||||
display = str(app.get("display_name") or name)
|
|
||||||
entry = str(app.get("entry_point") or f"cli-anything-{name}")
|
|
||||||
description = str(app.get("description") or f"Use {display} from nanobot.")
|
|
||||||
return f"""---
|
|
||||||
name: {_safe_skill_name(name)}
|
|
||||||
description: >-
|
|
||||||
{description}
|
|
||||||
---
|
|
||||||
|
|
||||||
# {display}
|
|
||||||
|
|
||||||
Use this skill when the user asks nanobot to operate {display} through its installed CLI app.
|
|
||||||
|
|
||||||
If the user attached `@{name}` in chat, treat that as the selected app for the current turn.
|
|
||||||
|
|
||||||
## Commands
|
|
||||||
|
|
||||||
```bash
|
|
||||||
{entry} --help
|
|
||||||
{entry} --json --help
|
|
||||||
```
|
|
||||||
|
|
||||||
Prefer machine-readable output when the CLI supports `--json`.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def _with_nanobot_skill_note(self, content: str, app: dict[str, Any]) -> str:
|
|
||||||
marker = "<!-- nanobot-cli-app-note -->"
|
|
||||||
if marker in content:
|
|
||||||
return content
|
|
||||||
name = str(app.get("name") or "unknown")
|
|
||||||
note = f"""{marker}
|
|
||||||
## Nanobot execution
|
|
||||||
|
|
||||||
Use the `run_cli_app` tool with `name="{name}"` for command execution. Do not invoke this CLI through shell unless the user explicitly asks. Prefer this skill when Runtime Context mentions `@{name}` as a CLI App Attachment.
|
|
||||||
"""
|
|
||||||
lines = content.splitlines(keepends=True)
|
|
||||||
if lines and lines[0].strip() == "---":
|
|
||||||
for index, line in enumerate(lines[1:], start=1):
|
|
||||||
if line.strip() == "---":
|
|
||||||
return "".join(lines[: index + 1]) + "\n" + note + "\n" + "".join(lines[index + 1 :])
|
|
||||||
return note + "\n" + content
|
|
||||||
|
|
||||||
def install_skill(self, app: dict[str, Any]) -> Path:
|
|
||||||
path = self._skill_path(str(app["name"]))
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
content = self._fetch_skill_content(app) or self._fallback_skill(app)
|
|
||||||
content = self._with_nanobot_skill_note(content, app)
|
|
||||||
path.write_text(content, encoding="utf-8")
|
|
||||||
return path
|
|
||||||
|
|
||||||
def remove_skill(self, name: str) -> None:
|
|
||||||
skill_dir = self._skill_path(name).parent
|
|
||||||
if skill_dir.is_dir():
|
|
||||||
shutil.rmtree(skill_dir)
|
|
||||||
|
|
||||||
def _record_installed(self, app: dict[str, Any]) -> None:
|
|
||||||
installed = self._load_installed()
|
|
||||||
installed[str(app["name"])] = self._installed_entry(app)
|
|
||||||
self._save_installed(installed)
|
|
||||||
self.install_skill(app)
|
|
||||||
|
|
||||||
def install(self, name: str) -> dict[str, Any]:
|
|
||||||
app = self.get_app(name)
|
|
||||||
if not self._install_supported(app):
|
|
||||||
raise CliAppError("this CLI app uses an unsupported install strategy")
|
|
||||||
strategy = self._strategy(app)
|
|
||||||
if strategy == "bundled":
|
|
||||||
detect_cmd = str(app.get("detect_cmd") or app.get("entry_point") or "")
|
|
||||||
if detect_cmd and _command_exists(detect_cmd):
|
|
||||||
self._record_installed(app)
|
|
||||||
return self.payload() | {"last_action": {"ok": True, "message": f"CLI for {app['display_name']} is available."}}
|
|
||||||
note = app.get("install_notes") or f"{app['display_name']} is bundled with its parent app."
|
|
||||||
raise CliAppError(str(note))
|
|
||||||
argv = self._argv_for_action(app, "install")
|
|
||||||
assert argv is not None
|
|
||||||
result = self._run_argv(argv, timeout=self.runtime.install_timeout)
|
|
||||||
if result.returncode != 0:
|
|
||||||
raise CliAppError(_truncate(result.stderr or result.stdout or "install failed"), status=500)
|
|
||||||
self._record_installed(app)
|
|
||||||
return self.payload() | {"last_action": {"ok": True, "message": f"Installed CLI for {app['display_name']}."}}
|
|
||||||
|
|
||||||
def update(self, name: str) -> dict[str, Any]:
|
|
||||||
app = self.get_app(name, force_refresh=True)
|
|
||||||
if str(app["name"]) not in self._load_installed():
|
|
||||||
raise CliAppError("CLI app is not installed")
|
|
||||||
if self._strategy(app) == "bundled":
|
|
||||||
self._record_installed(app)
|
|
||||||
return self.payload() | {"last_action": {"ok": True, "message": f"Checked {app['display_name']}."}}
|
|
||||||
argv = self._argv_for_action(app, "update")
|
|
||||||
assert argv is not None
|
|
||||||
result = self._run_argv(argv, timeout=self.runtime.install_timeout)
|
|
||||||
if result.returncode != 0:
|
|
||||||
raise CliAppError(_truncate(result.stderr or result.stdout or "update failed"), status=500)
|
|
||||||
self._record_installed(app)
|
|
||||||
return self.payload() | {"last_action": {"ok": True, "message": f"Updated CLI for {app['display_name']}."}}
|
|
||||||
|
|
||||||
def uninstall(self, name: str) -> dict[str, Any]:
|
|
||||||
app = self.get_app(name)
|
|
||||||
installed = self._load_installed()
|
|
||||||
if str(app["name"]) not in installed:
|
|
||||||
raise CliAppError("CLI app is not installed")
|
|
||||||
if self._strategy(app) != "bundled":
|
|
||||||
argv = self._argv_for_action(app, "uninstall")
|
|
||||||
assert argv is not None
|
|
||||||
result = self._run_argv(argv, timeout=self.runtime.install_timeout)
|
|
||||||
if result.returncode != 0:
|
|
||||||
raise CliAppError(_truncate(result.stderr or result.stdout or "uninstall failed"), status=500)
|
|
||||||
installed.pop(str(app["name"]), None)
|
|
||||||
self._save_installed(installed)
|
|
||||||
self.remove_skill(str(app["name"]))
|
|
||||||
return self.payload() | {"last_action": {"ok": True, "message": f"Uninstalled CLI for {app['display_name']}."}}
|
|
||||||
|
|
||||||
def test(self, name: str) -> dict[str, Any]:
|
|
||||||
app = self.get_app(name)
|
|
||||||
entry = str(app.get("entry_point") or "")
|
|
||||||
resolved = shutil.which(entry)
|
|
||||||
if not entry or not resolved:
|
|
||||||
raise CliAppError(f"{entry or name} is not available on PATH")
|
|
||||||
result = self._run_argv([resolved, "--help"], timeout=min(self.runtime.run_timeout, 30))
|
|
||||||
ok = result.returncode == 0
|
|
||||||
output = _truncate((result.stdout or result.stderr or "").strip(), 3000)
|
|
||||||
return self.payload() | {
|
|
||||||
"last_action": {
|
|
||||||
"ok": ok,
|
|
||||||
"message": f"{entry} --help exited {result.returncode}",
|
|
||||||
"output": output,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
def _resolve_cwd(
|
|
||||||
self,
|
|
||||||
working_dir: str | None,
|
|
||||||
*,
|
|
||||||
restrict_to_workspace: bool,
|
|
||||||
) -> Path:
|
|
||||||
cwd = Path(working_dir).expanduser() if working_dir else self.workspace
|
|
||||||
cwd = cwd.resolve(strict=False)
|
|
||||||
workspace = self.workspace.resolve(strict=False)
|
|
||||||
if restrict_to_workspace and cwd != workspace and not cwd.is_relative_to(workspace):
|
|
||||||
raise CliAppError("working_dir is outside the configured workspace")
|
|
||||||
return cwd
|
|
||||||
|
|
||||||
def _iter_artifact_candidates(self, cwd: Path) -> list[Path]:
|
|
||||||
if not cwd.is_dir():
|
|
||||||
return []
|
|
||||||
out: list[Path] = []
|
|
||||||
stack = [cwd]
|
|
||||||
scanned = 0
|
|
||||||
while stack and scanned < _MAX_ARTIFACT_SCAN_PATHS:
|
|
||||||
directory = stack.pop()
|
|
||||||
try:
|
|
||||||
entries = sorted(directory.iterdir(), key=lambda path: path.name.lower())
|
|
||||||
except OSError:
|
|
||||||
continue
|
|
||||||
for path in entries:
|
|
||||||
if scanned >= _MAX_ARTIFACT_SCAN_PATHS:
|
|
||||||
break
|
|
||||||
scanned += 1
|
|
||||||
try:
|
|
||||||
if path.is_dir() and not path.is_symlink():
|
|
||||||
if path.name not in _ARTIFACT_IGNORE_DIRS:
|
|
||||||
stack.append(path)
|
|
||||||
continue
|
|
||||||
if path.is_file() and path.suffix.lower() in _ARTIFACT_EXTENSIONS:
|
|
||||||
out.append(path.resolve(strict=False))
|
|
||||||
except OSError:
|
|
||||||
continue
|
|
||||||
return out
|
|
||||||
|
|
||||||
def _artifact_snapshot(self, cwd: Path) -> dict[Path, tuple[int, int]]:
|
|
||||||
snapshot: dict[Path, tuple[int, int]] = {}
|
|
||||||
for path in self._iter_artifact_candidates(cwd):
|
|
||||||
try:
|
|
||||||
stat = path.stat()
|
|
||||||
except OSError:
|
|
||||||
continue
|
|
||||||
snapshot[path] = (stat.st_mtime_ns, stat.st_size)
|
|
||||||
return snapshot
|
|
||||||
|
|
||||||
def _changed_artifacts(
|
|
||||||
self,
|
|
||||||
cwd: Path,
|
|
||||||
before: dict[Path, tuple[int, int]],
|
|
||||||
) -> list[Path]:
|
|
||||||
changed: list[tuple[int, Path]] = []
|
|
||||||
for path, stamp in self._artifact_snapshot(cwd).items():
|
|
||||||
if before.get(path) == stamp:
|
|
||||||
continue
|
|
||||||
changed.append((stamp[0], path))
|
|
||||||
changed.sort(key=lambda item: (item[0], item[1].name.lower()))
|
|
||||||
return [path for _, path in changed[-_MAX_ARTIFACT_REPORT:]]
|
|
||||||
|
|
||||||
def _format_artifact_path(self, cwd: Path, path: Path) -> str:
|
|
||||||
try:
|
|
||||||
return path.relative_to(cwd).as_posix()
|
|
||||||
except ValueError:
|
|
||||||
return path.name
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _format_artifact_size(path: Path) -> str:
|
|
||||||
try:
|
|
||||||
size = path.stat().st_size
|
|
||||||
except OSError:
|
|
||||||
return "unknown size"
|
|
||||||
if size < 1024:
|
|
||||||
return f"{size} B"
|
|
||||||
if size < 1024 * 1024:
|
|
||||||
return f"{size / 1024:.1f} KB"
|
|
||||||
return f"{size / (1024 * 1024):.1f} MB"
|
|
||||||
|
|
||||||
def _format_artifact_lines(self, cwd: Path, paths: list[Path]) -> list[str]:
|
|
||||||
lines: list[str] = []
|
|
||||||
for path in paths:
|
|
||||||
rel = self._format_artifact_path(cwd, path)
|
|
||||||
ext = path.suffix.lower()
|
|
||||||
kind = (
|
|
||||||
"previewable image"
|
|
||||||
if ext in _INLINE_ARTIFACT_EXTENSIONS
|
|
||||||
else ext.lstrip(".") or "file"
|
|
||||||
)
|
|
||||||
lines.append(f"- {rel} ({kind}, {self._format_artifact_size(path)})")
|
|
||||||
return lines
|
|
||||||
|
|
||||||
def run(
|
|
||||||
self,
|
|
||||||
name: str,
|
|
||||||
args: list[str] | None = None,
|
|
||||||
*,
|
|
||||||
json_output: bool = False,
|
|
||||||
working_dir: str | None = None,
|
|
||||||
timeout: int | None = None,
|
|
||||||
restrict_to_workspace: bool = False,
|
|
||||||
) -> str:
|
|
||||||
app = self.get_app(name)
|
|
||||||
installed = self._load_installed()
|
|
||||||
if str(app["name"]) not in installed:
|
|
||||||
raise CliAppError(f"CLI app '{name}' is not installed")
|
|
||||||
cwd = self._resolve_cwd(working_dir, restrict_to_workspace=restrict_to_workspace)
|
|
||||||
entry = str(installed[str(app["name"])].get("entry_point") or app.get("entry_point") or "")
|
|
||||||
resolved = shutil.which(entry)
|
|
||||||
if not entry or not resolved:
|
|
||||||
raise CliAppError(f"{entry or name} is not available on PATH")
|
|
||||||
clean_args = [str(arg) for arg in (args or [])]
|
|
||||||
if json_output and "--json" not in clean_args:
|
|
||||||
clean_args = ["--json", *clean_args]
|
|
||||||
effective_timeout = max(1, min(timeout or self.runtime.run_timeout, 600))
|
|
||||||
artifact_snapshot = self._artifact_snapshot(cwd)
|
|
||||||
try:
|
|
||||||
result = subprocess.run(
|
|
||||||
[resolved, *clean_args],
|
|
||||||
cwd=str(cwd),
|
|
||||||
capture_output=True,
|
|
||||||
text=True,
|
|
||||||
timeout=effective_timeout,
|
|
||||||
env=os.environ.copy(),
|
|
||||||
)
|
|
||||||
except subprocess.TimeoutExpired:
|
|
||||||
return f"CLI app '{name}' timed out after {effective_timeout}s"
|
|
||||||
output = [
|
|
||||||
f"CLI app '{name}' exited {result.returncode}.",
|
|
||||||
f"Command: {entry} {' '.join(shlex.quote(arg) for arg in clean_args)}".rstrip(),
|
|
||||||
]
|
|
||||||
if result.stdout:
|
|
||||||
output.append("\nSTDOUT:\n" + result.stdout.rstrip())
|
|
||||||
if result.stderr:
|
|
||||||
output.append("\nSTDERR:\n" + result.stderr.rstrip())
|
|
||||||
artifacts = self._changed_artifacts(cwd, artifact_snapshot)
|
|
||||||
if artifacts:
|
|
||||||
output.append(
|
|
||||||
"\nArtifacts created or updated:\n"
|
|
||||||
+ "\n".join(self._format_artifact_lines(cwd, artifacts))
|
|
||||||
)
|
|
||||||
if any(path.suffix.lower() in _INLINE_ARTIFACT_EXTENSIONS for path in artifacts):
|
|
||||||
output.append(
|
|
||||||
"\nTo show a preview in WebUI, reference a raster artifact with Markdown "
|
|
||||||
"using its workspace-relative path, for example ``."
|
|
||||||
)
|
|
||||||
return _truncate("\n".join(output))
|
|
||||||
@@ -1,62 +0,0 @@
|
|||||||
"""CLI Apps helpers shared by the agent loop and settings surfaces."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Mapping
|
|
||||||
|
|
||||||
|
|
||||||
def session_extra(metadata: Mapping[str, Any] | None) -> dict[str, Any]:
|
|
||||||
"""Return persisted session kwargs for CLI app attachments."""
|
|
||||||
cli_apps = metadata.get("cli_apps") if isinstance(metadata, Mapping) else None
|
|
||||||
return {"cli_apps": cli_apps} if isinstance(cli_apps, list) and cli_apps else {}
|
|
||||||
|
|
||||||
|
|
||||||
def runtime_lines(message: Any, workspace: Path, *, skip: bool = False) -> list[str]:
|
|
||||||
"""Return model-visible CLI app annotations for the current turn."""
|
|
||||||
if skip:
|
|
||||||
return []
|
|
||||||
text = message.content if isinstance(getattr(message, "content", None), str) else ""
|
|
||||||
metadata = message.metadata if isinstance(getattr(message, "metadata", None), Mapping) else None
|
|
||||||
return _cli_app_runtime_lines(text, metadata, workspace)
|
|
||||||
|
|
||||||
|
|
||||||
def _cli_app_runtime_lines(
|
|
||||||
text: str,
|
|
||||||
metadata: Mapping[str, Any] | None,
|
|
||||||
workspace: Path,
|
|
||||||
) -> list[str]:
|
|
||||||
structured = metadata.get("cli_apps") if isinstance(metadata, Mapping) else None
|
|
||||||
if isinstance(structured, list):
|
|
||||||
mentions = [
|
|
||||||
item for item in structured
|
|
||||||
if isinstance(item, Mapping) and isinstance(item.get("name"), str)
|
|
||||||
]
|
|
||||||
if mentions:
|
|
||||||
return [
|
|
||||||
"CLI App Attachment: "
|
|
||||||
f"@{str(item['name']).strip().lower()} "
|
|
||||||
f"(installed; tool=run_cli_app; "
|
|
||||||
f"entry_point={str(item.get('entry_point') or 'unknown')}; "
|
|
||||||
f"skill=skills/cli-app-{str(item['name']).strip().lower()}/SKILL.md). "
|
|
||||||
"Read the skill when useful, then run this app with `run_cli_app`; do not bypass it with shell."
|
|
||||||
for item in mentions
|
|
||||||
if str(item.get("name") or "").strip()
|
|
||||||
]
|
|
||||||
if "@" not in text:
|
|
||||||
return []
|
|
||||||
try:
|
|
||||||
from nanobot.cli_apps import CliAppManager
|
|
||||||
|
|
||||||
mentions = CliAppManager(workspace=workspace).mentioned_installed_apps(text)
|
|
||||||
except Exception:
|
|
||||||
return []
|
|
||||||
return [
|
|
||||||
"CLI App Mention: "
|
|
||||||
f"@{item['name']} "
|
|
||||||
f"(installed; tool={item['tool']}; "
|
|
||||||
f"entry_point={item['entry_point'] or 'unknown'}; "
|
|
||||||
f"skill={item['skill']}). "
|
|
||||||
"Read the skill when useful, then run this app with `run_cli_app`; do not bypass it with shell."
|
|
||||||
for item in mentions
|
|
||||||
]
|
|
||||||
@@ -11,7 +11,6 @@ from pydantic_settings import BaseSettings
|
|||||||
from nanobot.cron.types import CronSchedule
|
from nanobot.cron.types import CronSchedule
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nanobot.agent.tools.cli_apps import CliAppsToolConfig
|
|
||||||
from nanobot.agent.tools.image_generation import ImageGenerationToolConfig
|
from nanobot.agent.tools.image_generation import ImageGenerationToolConfig
|
||||||
from nanobot.agent.tools.self import MyToolConfig
|
from nanobot.agent.tools.self import MyToolConfig
|
||||||
from nanobot.agent.tools.shell import ExecToolConfig
|
from nanobot.agent.tools.shell import ExecToolConfig
|
||||||
@@ -191,7 +190,6 @@ 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)
|
||||||
@@ -209,10 +207,8 @@ 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 (硅基流动)
|
||||||
novita: ProviderConfig = Field(default_factory=ProviderConfig) # Novita AI
|
|
||||||
volcengine: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine (火山引擎)
|
volcengine: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine (火山引擎)
|
||||||
volcengine_coding_plan: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine Coding Plan
|
volcengine_coding_plan: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine Coding Plan
|
||||||
byteplus: ProviderConfig = Field(default_factory=ProviderConfig) # BytePlus (VolcEngine international)
|
byteplus: ProviderConfig = Field(default_factory=ProviderConfig) # BytePlus (VolcEngine international)
|
||||||
@@ -277,7 +273,6 @@ class ToolsConfig(Base):
|
|||||||
|
|
||||||
web: WebToolsConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.web", "WebToolsConfig"))
|
web: WebToolsConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.web", "WebToolsConfig"))
|
||||||
exec: ExecToolConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.shell", "ExecToolConfig"))
|
exec: ExecToolConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.shell", "ExecToolConfig"))
|
||||||
cli_apps: CliAppsToolConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.cli_apps", "CliAppsToolConfig"))
|
|
||||||
my: MyToolConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.self", "MyToolConfig"))
|
my: MyToolConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.self", "MyToolConfig"))
|
||||||
image_generation: ImageGenerationToolConfig = Field(
|
image_generation: ImageGenerationToolConfig = Field(
|
||||||
default_factory=lambda: _lazy_default("nanobot.agent.tools.image_generation", "ImageGenerationToolConfig"),
|
default_factory=lambda: _lazy_default("nanobot.agent.tools.image_generation", "ImageGenerationToolConfig"),
|
||||||
@@ -464,7 +459,6 @@ def _resolve_tool_config_refs() -> None:
|
|||||||
"""
|
"""
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
from nanobot.agent.tools.cli_apps import CliAppsToolConfig
|
|
||||||
from nanobot.agent.tools.image_generation import ImageGenerationToolConfig
|
from nanobot.agent.tools.image_generation import ImageGenerationToolConfig
|
||||||
from nanobot.agent.tools.self import MyToolConfig
|
from nanobot.agent.tools.self import MyToolConfig
|
||||||
from nanobot.agent.tools.shell import ExecToolConfig
|
from nanobot.agent.tools.shell import ExecToolConfig
|
||||||
@@ -473,7 +467,6 @@ def _resolve_tool_config_refs() -> None:
|
|||||||
# Re-export into this module's namespace
|
# Re-export into this module's namespace
|
||||||
mod = sys.modules[__name__]
|
mod = sys.modules[__name__]
|
||||||
mod.ExecToolConfig = ExecToolConfig # type: ignore[attr-defined]
|
mod.ExecToolConfig = ExecToolConfig # type: ignore[attr-defined]
|
||||||
mod.CliAppsToolConfig = CliAppsToolConfig # type: ignore[attr-defined]
|
|
||||||
mod.WebToolsConfig = WebToolsConfig # type: ignore[attr-defined]
|
mod.WebToolsConfig = WebToolsConfig # type: ignore[attr-defined]
|
||||||
mod.WebSearchConfig = WebSearchConfig # type: ignore[attr-defined]
|
mod.WebSearchConfig = WebSearchConfig # type: ignore[attr-defined]
|
||||||
mod.WebFetchConfig = WebFetchConfig # type: ignore[attr-defined]
|
mod.WebFetchConfig = WebFetchConfig # type: ignore[attr-defined]
|
||||||
|
|||||||
@@ -1,18 +1,6 @@
|
|||||||
"""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
|
|
||||||
|
|||||||
+6
-2
@@ -8,7 +8,6 @@ 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)
|
||||||
@@ -64,7 +63,12 @@ class Nanobot:
|
|||||||
|
|
||||||
loop = AgentLoop.from_config(
|
loop = AgentLoop.from_config(
|
||||||
config,
|
config,
|
||||||
image_generation_provider_configs=image_gen_provider_configs(config),
|
image_generation_provider_configs={
|
||||||
|
"openrouter": config.providers.openrouter,
|
||||||
|
"aihubmix": config.providers.aihubmix,
|
||||||
|
"minimax": config.providers.minimax,
|
||||||
|
"gemini": config.providers.gemini,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
return cls(loop)
|
return cls(loop)
|
||||||
|
|
||||||
|
|||||||
@@ -590,7 +590,6 @@ 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,
|
||||||
@@ -599,12 +598,11 @@ 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 or on_tool_call_delta:
|
if on_content_delta or on_thinking_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(
|
||||||
@@ -613,22 +611,7 @@ class AnthropicProvider(LLMProvider):
|
|||||||
)
|
)
|
||||||
except StopAsyncIteration:
|
except StopAsyncIteration:
|
||||||
break
|
break
|
||||||
if chunk.type == "content_block_start":
|
if (
|
||||||
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"
|
||||||
):
|
):
|
||||||
@@ -642,20 +625,6 @@ 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,7 +158,6 @@ 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(
|
||||||
@@ -170,7 +169,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, on_tool_call_delta)
|
await consume_sdk_stream(stream, on_content_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 a tool-capable stop.
|
"""Tools execute only when has_tool_calls AND finish_reason is ``tool_calls`` / ``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", "function_call", "stop")
|
return self.finish_reason in ("tool_calls", "stop")
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -501,7 +501,6 @@ 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.
|
||||||
|
|
||||||
@@ -515,7 +514,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_tool_call_delta
|
_ = on_thinking_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,
|
||||||
@@ -545,7 +544,6 @@ 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:
|
||||||
@@ -563,7 +561,6 @@ 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,9 +704,8 @@ 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_tool_call_delta
|
_ = on_thinking_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] = []
|
||||||
|
|||||||
@@ -207,9 +207,8 @@ 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
|
||||||
client.api_key = token
|
self._client.api_key = token
|
||||||
return token
|
return token
|
||||||
|
|
||||||
async def chat(
|
async def chat(
|
||||||
@@ -244,7 +243,6 @@ 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(
|
||||||
@@ -257,5 +255,4 @@ 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,
|
|
||||||
)
|
)
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -40,7 +40,6 @@ 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
|
||||||
@@ -71,7 +70,6 @@ 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):
|
||||||
@@ -80,7 +78,6 @@ 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:
|
||||||
@@ -103,18 +100,9 @@ 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(
|
return await self._call_codex(messages, tools, model, reasoning_effort, tool_choice, on_content_delta)
|
||||||
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
|
||||||
@@ -150,7 +138,6 @@ 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:
|
||||||
@@ -161,7 +148,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, on_tool_call_delta)
|
return await consume_sse(response, on_content_delta)
|
||||||
|
|
||||||
|
|
||||||
def _prompt_cache_key(messages: list[dict[str, Any]]) -> str:
|
def _prompt_cache_key(messages: list[dict[str, Any]]) -> str:
|
||||||
|
|||||||
@@ -11,15 +11,25 @@ import secrets
|
|||||||
import string
|
import string
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from collections import deque
|
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from ipaddress import ip_address
|
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,
|
||||||
@@ -29,15 +39,8 @@ 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",
|
||||||
@@ -75,43 +78,41 @@ _THINKING_STYLE_MAP: dict[str, Any] = {
|
|||||||
"enable_thinking": lambda on: {"enable_thinking": on},
|
"enable_thinking": lambda on: {"enable_thinking": on},
|
||||||
"reasoning_split": lambda on: {"reasoning_split": on},
|
"reasoning_split": lambda on: {"reasoning_split": on},
|
||||||
}
|
}
|
||||||
_GATEWAY_REASONING_STYLE_MAP: dict[str, Any] = {
|
|
||||||
"reasoning_effort": lambda effort: {"reasoning": {"effort": effort}},
|
|
||||||
}
|
|
||||||
_MODEL_THINKING_STYLES: dict[str, str] = {
|
|
||||||
**dict.fromkeys(_KIMI_THINKING_MODELS, "thinking_type"),
|
|
||||||
**dict.fromkeys(_MIMO_THINKING_MODELS, "thinking_type"),
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _model_slug(model_name: str) -> str:
|
def _is_kimi_thinking_model(model_name: str) -> bool:
|
||||||
return model_name.lower().rsplit("/", 1)[-1]
|
"""Return True if model_name refers to a Kimi thinking-capable model.
|
||||||
|
|
||||||
|
Supports two forms:
|
||||||
|
- Exact match: e.g. kimi-k2.5 / kimi-k2.6 in _KIMI_THINKING_MODELS
|
||||||
|
- Slug match: moonshotai/kimi-k2.5 -> the part after the last "/"
|
||||||
|
is checked against _KIMI_THINKING_MODELS
|
||||||
|
|
||||||
|
This covers both the native Moonshot provider (bare slug) and
|
||||||
|
OpenRouter-style names (``"publisher/slug"``).
|
||||||
|
"""
|
||||||
|
name = model_name.lower()
|
||||||
|
if name in _KIMI_THINKING_MODELS:
|
||||||
|
return True
|
||||||
|
if "/" in name and name.rsplit("/", 1)[1] in _KIMI_THINKING_MODELS:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def _model_thinking_style(model_name: str) -> str:
|
def _is_mimo_thinking_model(model_name: str) -> bool:
|
||||||
return _MODEL_THINKING_STYLES.get(_model_slug(model_name), "")
|
"""Return True if model_name refers to a MiMo thinking-capable model.
|
||||||
|
|
||||||
|
Mirrors _is_kimi_thinking_model: gateway providers (e.g. OpenRouter
|
||||||
def _thinking_styles_for(spec: ProviderSpec | None, model_name: str) -> list[str]:
|
routing ``xiaomi/mimo-v2.5-pro``) have no ``thinking_style`` on their
|
||||||
styles: list[str] = []
|
spec, so the spec-driven branch in _build_kwargs misses them. The
|
||||||
if spec and spec.thinking_style:
|
model-name path catches those cases.
|
||||||
styles.append(spec.thinking_style)
|
"""
|
||||||
model_style = _model_thinking_style(model_name)
|
name = model_name.lower()
|
||||||
if model_style and model_style not in styles:
|
if name in _MIMO_THINKING_MODELS:
|
||||||
styles.append(model_style)
|
return True
|
||||||
return styles
|
if "/" in name and name.rsplit("/", 1)[1] in _MIMO_THINKING_MODELS:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
def _thinking_extra_body(style: str, thinking_enabled: bool) -> dict[str, Any] | None:
|
|
||||||
builder = _THINKING_STYLE_MAP.get(style)
|
|
||||||
return builder(thinking_enabled) if builder else None
|
|
||||||
|
|
||||||
|
|
||||||
def _gateway_reasoning_extra_body(style: str, effort: str | None) -> dict[str, Any] | None:
|
|
||||||
if not effort:
|
|
||||||
return None
|
|
||||||
builder = _GATEWAY_REASONING_STYLE_MAP.get(style)
|
|
||||||
return builder(effort) if builder else None
|
|
||||||
|
|
||||||
|
|
||||||
def _openai_compat_timeout_s() -> float:
|
def _openai_compat_timeout_s() -> float:
|
||||||
@@ -301,75 +302,42 @@ 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
|
||||||
self._default_headers = {"x-session-affinity": uuid.uuid4().hex}
|
default_headers = {"x-session-affinity": uuid.uuid4().hex}
|
||||||
if _uses_openrouter_attribution(spec, effective_base):
|
if _uses_openrouter_attribution(spec, effective_base):
|
||||||
self._default_headers.update(_DEFAULT_OPENROUTER_HEADERS)
|
default_headers.update(_DEFAULT_OPENROUTER_HEADERS)
|
||||||
if extra_headers:
|
if extra_headers:
|
||||||
self._default_headers.update(extra_headers)
|
default_headers.update(extra_headers)
|
||||||
self._api_key_for_client = api_key or "no-key"
|
|
||||||
self._is_local = _is_local_endpoint(spec, effective_base)
|
|
||||||
|
|
||||||
# Lazy-init: the OpenAI client and its httpx transport are expensive
|
|
||||||
# to create (~700 ms on Windows). Defer until first use.
|
|
||||||
self._client: AsyncOpenAIType | None = None
|
|
||||||
self._client_lock = asyncio.Lock()
|
|
||||||
|
|
||||||
# Responses API circuit breaker: skip after repeated failures,
|
|
||||||
# probe again after _RESPONSES_PROBE_INTERVAL_S seconds.
|
|
||||||
self._responses_failures: dict[str, int] = {}
|
|
||||||
self._responses_tripped_at: dict[str, float] = {}
|
|
||||||
|
|
||||||
def _build_client(self) -> None:
|
|
||||||
"""Create the OpenAI client using the current module-level AsyncOpenAI."""
|
|
||||||
import httpx
|
|
||||||
|
|
||||||
|
# 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.
|
||||||
timeout_s = _openai_compat_timeout_s()
|
timeout_s = _openai_compat_timeout_s()
|
||||||
http_client: httpx.AsyncClient | None = None
|
http_client: httpx.AsyncClient | None = None
|
||||||
if self._is_local:
|
if _is_local_endpoint(spec, effective_base):
|
||||||
# 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(
|
http_client = httpx.AsyncClient(
|
||||||
limits=httpx.Limits(keepalive_expiry=0),
|
limits=httpx.Limits(keepalive_expiry=0),
|
||||||
timeout=timeout_s,
|
timeout=timeout_s,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._client = AsyncOpenAI(
|
self._client = AsyncOpenAI(
|
||||||
api_key=self._api_key_for_client,
|
api_key=api_key or "no-key",
|
||||||
base_url=self._effective_base,
|
base_url=effective_base,
|
||||||
default_headers=self._default_headers,
|
default_headers=default_headers,
|
||||||
max_retries=0,
|
max_retries=0,
|
||||||
timeout=timeout_s,
|
timeout=timeout_s,
|
||||||
http_client=http_client,
|
http_client=http_client,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _ensure_client(self):
|
# Responses API circuit breaker: skip after repeated failures,
|
||||||
"""Return the shared OpenAI client, creating it on first call."""
|
# probe again after _RESPONSES_PROBE_INTERVAL_S seconds.
|
||||||
if self._client is not None:
|
self._responses_failures: dict[str, int] = {}
|
||||||
return self._client
|
self._responses_tripped_at: dict[str, float] = {}
|
||||||
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."""
|
||||||
@@ -464,7 +432,6 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
"""Strip non-standard keys, normalize tool_call IDs."""
|
"""Strip non-standard keys, normalize tool_call IDs."""
|
||||||
sanitized = LLMProvider._sanitize_request_messages(messages, _ALLOWED_MSG_KEYS)
|
sanitized = LLMProvider._sanitize_request_messages(messages, _ALLOWED_MSG_KEYS)
|
||||||
id_map: dict[str, str] = {}
|
id_map: dict[str, str] = {}
|
||||||
pending_tool_ids: dict[str, deque[str]] = {}
|
|
||||||
force_string_content = bool(self._spec and self._spec.name == "deepseek")
|
force_string_content = bool(self._spec and self._spec.name == "deepseek")
|
||||||
|
|
||||||
def map_id(value: Any) -> Any:
|
def map_id(value: Any) -> Any:
|
||||||
@@ -472,49 +439,15 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
return value
|
return value
|
||||||
return id_map.setdefault(value, self._normalize_tool_call_id(value))
|
return id_map.setdefault(value, self._normalize_tool_call_id(value))
|
||||||
|
|
||||||
def unique_tool_id(value: Any, used_ids: set[str], idx: int) -> str:
|
|
||||||
if isinstance(value, str) and value:
|
|
||||||
base = map_id(value)
|
|
||||||
else:
|
|
||||||
base = _short_tool_id()
|
|
||||||
if not isinstance(base, str) or not base:
|
|
||||||
base = _short_tool_id()
|
|
||||||
if base not in used_ids:
|
|
||||||
return base
|
|
||||||
seed = value if isinstance(value, str) and value else base
|
|
||||||
salt = 1
|
|
||||||
while True:
|
|
||||||
candidate = self._normalize_tool_call_id(f"{seed}:{idx}:{salt}")
|
|
||||||
if isinstance(candidate, str) and candidate not in used_ids:
|
|
||||||
return candidate
|
|
||||||
salt += 1
|
|
||||||
|
|
||||||
def map_tool_result_id(value: Any) -> Any:
|
|
||||||
if not isinstance(value, str):
|
|
||||||
return value
|
|
||||||
queue = pending_tool_ids.get(value)
|
|
||||||
if queue:
|
|
||||||
mapped = queue.popleft()
|
|
||||||
if not queue:
|
|
||||||
pending_tool_ids.pop(value, None)
|
|
||||||
return mapped
|
|
||||||
return map_id(value)
|
|
||||||
|
|
||||||
for clean in sanitized:
|
for clean in sanitized:
|
||||||
if isinstance(clean.get("tool_calls"), list):
|
if isinstance(clean.get("tool_calls"), list):
|
||||||
normalized = []
|
normalized = []
|
||||||
used_ids: set[str] = set()
|
for tc in clean["tool_calls"]:
|
||||||
for idx, tc in enumerate(clean["tool_calls"]):
|
|
||||||
if not isinstance(tc, dict):
|
if not isinstance(tc, dict):
|
||||||
normalized.append(tc)
|
normalized.append(tc)
|
||||||
continue
|
continue
|
||||||
tc_clean = dict(tc)
|
tc_clean = dict(tc)
|
||||||
raw_id = tc_clean.get("id")
|
tc_clean["id"] = map_id(tc_clean.get("id"))
|
||||||
mapped_id = unique_tool_id(raw_id, used_ids, idx)
|
|
||||||
tc_clean["id"] = mapped_id
|
|
||||||
used_ids.add(mapped_id)
|
|
||||||
if isinstance(raw_id, str) and raw_id:
|
|
||||||
pending_tool_ids.setdefault(raw_id, deque()).append(mapped_id)
|
|
||||||
function = tc_clean.get("function")
|
function = tc_clean.get("function")
|
||||||
if isinstance(function, dict):
|
if isinstance(function, dict):
|
||||||
function_clean = dict(function)
|
function_clean = dict(function)
|
||||||
@@ -532,7 +465,7 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
# that mix non-empty content with tool_calls.
|
# that mix non-empty content with tool_calls.
|
||||||
clean["content"] = None
|
clean["content"] = None
|
||||||
if "tool_call_id" in clean and clean["tool_call_id"]:
|
if "tool_call_id" in clean and clean["tool_call_id"]:
|
||||||
clean["tool_call_id"] = map_tool_result_id(clean["tool_call_id"])
|
clean["tool_call_id"] = map_id(clean["tool_call_id"])
|
||||||
if (
|
if (
|
||||||
force_string_content
|
force_string_content
|
||||||
and not (clean.get("role") == "assistant" and clean.get("tool_calls"))
|
and not (clean.get("role") == "assistant" and clean.get("tool_calls"))
|
||||||
@@ -619,27 +552,39 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
if wire_effort and semantic_effort != "none":
|
if wire_effort and semantic_effort != "none":
|
||||||
kwargs["reasoning_effort"] = wire_effort
|
kwargs["reasoning_effort"] = wire_effort
|
||||||
|
|
||||||
# Only send thinking controls when reasoning_effort is explicit so
|
# Provider-specific thinking parameters.
|
||||||
# omitting the config preserves each provider's default.
|
# Only sent when reasoning_effort is explicitly configured so that
|
||||||
if reasoning_effort is not None:
|
# the provider default is preserved otherwise.
|
||||||
|
# The mapping is driven by ProviderSpec.thinking_style so that adding
|
||||||
|
# a new provider never requires touching this function.
|
||||||
|
if spec and spec.thinking_style and reasoning_effort is not None:
|
||||||
thinking_enabled = semantic_effort not in ("none", "minimal")
|
thinking_enabled = semantic_effort not in ("none", "minimal")
|
||||||
for thinking_style in _thinking_styles_for(spec, model_name):
|
extra = _THINKING_STYLE_MAP.get(spec.thinking_style, lambda _: None)(thinking_enabled)
|
||||||
extra = _thinking_extra_body(thinking_style, thinking_enabled)
|
if extra:
|
||||||
if extra:
|
kwargs.setdefault("extra_body", {}).update(extra)
|
||||||
kwargs.setdefault("extra_body", {}).update(extra)
|
|
||||||
gateway_style = getattr(spec, "gateway_reasoning_style", "") if spec else ""
|
|
||||||
if gateway_style and _model_thinking_style(model_name):
|
|
||||||
extra = _gateway_reasoning_extra_body(gateway_style, semantic_effort)
|
|
||||||
if extra:
|
|
||||||
kwargs.setdefault("extra_body", {}).update(extra)
|
|
||||||
|
|
||||||
# Moonshot rejects requests that carry both 'reasoning_effort'
|
# Model-level thinking injection for Kimi thinking-capable models.
|
||||||
# and the native 'thinking' param. We already expressed the
|
# Strip any provider prefix (e.g. "moonshotai/") before the set lookup
|
||||||
# user's intent via the provider-native shape, so drop the
|
# so that OpenRouter-style names like "moonshotai/kimi-k2.5" are handled
|
||||||
# redundant wire-level kwarg. Only kimi models need this —
|
# identically to bare names like "kimi-k2.5".
|
||||||
# Xiaomi's API accepts both params.
|
if reasoning_effort is not None and _is_kimi_thinking_model(model_name):
|
||||||
if _model_slug(model_name) in _KIMI_THINKING_MODELS:
|
thinking_enabled = semantic_effort not in ("none", "minimal")
|
||||||
kwargs.pop("reasoning_effort", None)
|
kwargs.setdefault("extra_body", {}).update(
|
||||||
|
{"thinking": {"type": "enabled" if thinking_enabled else "disabled"}}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Model-level thinking injection for MiMo thinking-capable models.
|
||||||
|
# Same shape as Kimi: gateway providers (OpenRouter, etc.) lack the
|
||||||
|
# xiaomi_mimo spec's thinking_style, so the spec-driven branch above
|
||||||
|
# misses them — match by model name to catch "xiaomi/mimo-v2.5-pro"
|
||||||
|
# and friends. (Direct xiaomi_mimo requests are also covered here;
|
||||||
|
# both branches write the same payload, so the dict update is a
|
||||||
|
# safe no-op for already-handled cases.)
|
||||||
|
if reasoning_effort is not None and _is_mimo_thinking_model(model_name):
|
||||||
|
thinking_enabled = semantic_effort not in ("none", "minimal")
|
||||||
|
kwargs.setdefault("extra_body", {}).update(
|
||||||
|
{"thinking": {"type": "enabled" if thinking_enabled else "disabled"}}
|
||||||
|
)
|
||||||
|
|
||||||
if tools:
|
if tools:
|
||||||
kwargs["tools"] = tools
|
kwargs["tools"] = tools
|
||||||
@@ -654,7 +599,8 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
and semantic_effort not in ("none", "minimal")
|
and semantic_effort not in ("none", "minimal")
|
||||||
and (
|
and (
|
||||||
(spec and spec.thinking_style)
|
(spec and spec.thinking_style)
|
||||||
or _model_thinking_style(model_name)
|
or _is_kimi_thinking_model(model_name)
|
||||||
|
or _is_mimo_thinking_model(model_name)
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
implicit_deepseek_thinking = (
|
implicit_deepseek_thinking = (
|
||||||
@@ -1053,21 +999,6 @@ 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)
|
||||||
@@ -1098,7 +1029,6 @@ 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
|
||||||
|
|
||||||
@@ -1117,19 +1047,8 @@ 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 (getattr(delta, "tool_calls", None) or []) if delta else []:
|
for tc in (delta.tool_calls 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))
|
|
||||||
|
|
||||||
# Some providers (e.g. Zhipu/GLM) reuse the same tool_call id for
|
|
||||||
# parallel tool calls in streaming mode. Deduplicate before building
|
|
||||||
# the response so downstream tool messages don't collide.
|
|
||||||
_seen_tc_ids: set[str] = set()
|
|
||||||
for b in tc_bufs.values():
|
|
||||||
if not b["id"] or b["id"] in _seen_tc_ids:
|
|
||||||
b["id"] = _short_tool_id()
|
|
||||||
_seen_tc_ids.add(b["id"])
|
|
||||||
|
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
content="".join(content_parts) or None,
|
content="".join(content_parts) or None,
|
||||||
@@ -1245,7 +1164,6 @@ 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:
|
||||||
@@ -1285,9 +1203,7 @@ 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):
|
||||||
@@ -1310,16 +1226,9 @@ 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(
|
||||||
@@ -1343,12 +1252,6 @@ 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)
|
||||||
@@ -1376,28 +1279,6 @@ 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(
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ def convert_messages(messages: list[dict[str, Any]]) -> tuple[str, list[dict[str
|
|||||||
"""
|
"""
|
||||||
system_prompt = ""
|
system_prompt = ""
|
||||||
input_items: list[dict[str, Any]] = []
|
input_items: list[dict[str, Any]] = []
|
||||||
used_item_ids: set[str] = set()
|
|
||||||
|
|
||||||
for idx, msg in enumerate(messages):
|
for idx, msg in enumerate(messages):
|
||||||
role = msg.get("role")
|
role = msg.get("role")
|
||||||
@@ -31,19 +30,17 @@ def convert_messages(messages: list[dict[str, Any]]) -> tuple[str, list[dict[str
|
|||||||
|
|
||||||
if role == "assistant":
|
if role == "assistant":
|
||||||
if isinstance(content, str) and content:
|
if isinstance(content, str) and content:
|
||||||
message_id = _unique_item_id(f"msg_{idx}", used_item_ids)
|
|
||||||
input_items.append({
|
input_items.append({
|
||||||
"type": "message", "role": "assistant",
|
"type": "message", "role": "assistant",
|
||||||
"content": [{"type": "output_text", "text": content}],
|
"content": [{"type": "output_text", "text": content}],
|
||||||
"status": "completed", "id": message_id,
|
"status": "completed", "id": f"msg_{idx}",
|
||||||
})
|
})
|
||||||
for tool_call in msg.get("tool_calls", []) or []:
|
for tool_call in msg.get("tool_calls", []) or []:
|
||||||
fn = tool_call.get("function") or {}
|
fn = tool_call.get("function") or {}
|
||||||
call_id, item_id = split_tool_call_id(tool_call.get("id"))
|
call_id, item_id = split_tool_call_id(tool_call.get("id"))
|
||||||
response_item_id = _unique_item_id(item_id or f"fc_{idx}", used_item_ids)
|
|
||||||
input_items.append({
|
input_items.append({
|
||||||
"type": "function_call",
|
"type": "function_call",
|
||||||
"id": response_item_id,
|
"id": item_id or f"fc_{idx}",
|
||||||
"call_id": call_id or f"call_{idx}",
|
"call_id": call_id or f"call_{idx}",
|
||||||
"name": fn.get("name"),
|
"name": fn.get("name"),
|
||||||
"arguments": fn.get("arguments") or "{}",
|
"arguments": fn.get("arguments") or "{}",
|
||||||
@@ -100,20 +97,6 @@ def convert_tools(tools: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|||||||
return converted
|
return converted
|
||||||
|
|
||||||
|
|
||||||
def _unique_item_id(item_id: str, used: set[str]) -> str:
|
|
||||||
"""Return a Responses input item id that is unique within one request."""
|
|
||||||
if item_id not in used:
|
|
||||||
used.add(item_id)
|
|
||||||
return item_id
|
|
||||||
|
|
||||||
suffix = 2
|
|
||||||
while f"{item_id}_{suffix}" in used:
|
|
||||||
suffix += 1
|
|
||||||
unique = f"{item_id}_{suffix}"
|
|
||||||
used.add(unique)
|
|
||||||
return unique
|
|
||||||
|
|
||||||
|
|
||||||
def split_tool_call_id(tool_call_id: Any) -> tuple[str, str | None]:
|
def split_tool_call_id(tool_call_id: Any) -> tuple[str, str | None]:
|
||||||
"""Split a compound ``call_id|item_id`` string.
|
"""Split a compound ``call_id|item_id`` string.
|
||||||
|
|
||||||
|
|||||||
@@ -62,7 +62,6 @@ 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 = ""
|
||||||
@@ -83,12 +82,6 @@ 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
|
||||||
@@ -97,14 +90,7 @@ 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:
|
||||||
delta = event.get("delta") or ""
|
tool_call_buffers[call_id]["arguments"] += 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:
|
||||||
@@ -224,7 +210,6 @@ 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 = ""
|
||||||
@@ -247,12 +232,6 @@ 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
|
||||||
@@ -261,14 +240,7 @@ 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:
|
||||||
delta = getattr(event, "delta", "") or ""
|
tool_call_buffers[call_id]["arguments"] += 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:
|
||||||
|
|||||||
@@ -71,11 +71,6 @@ class ProviderSpec:
|
|||||||
# "reasoning_split" — {"reasoning_split": true/false} (MiniMax)
|
# "reasoning_split" — {"reasoning_split": true/false} (MiniMax)
|
||||||
thinking_style: str = ""
|
thinking_style: str = ""
|
||||||
|
|
||||||
# Gateway-native reasoning control to pair with model-level thinking styles.
|
|
||||||
# "reasoning_effort" — {"reasoning": {"effort": <none|minimal|...>}}
|
|
||||||
# (OpenRouter)
|
|
||||||
gateway_reasoning_style: str = ""
|
|
||||||
|
|
||||||
# When True, treat the "reasoning" response field as formal content
|
# When True, treat the "reasoning" response field as formal content
|
||||||
# when "content" is empty. Only set this for providers (e.g. StepFun)
|
# when "content" is empty. Only set this for providers (e.g. StepFun)
|
||||||
# whose API returns the actual answer in "reasoning" instead of "content".
|
# whose API returns the actual answer in "reasoning" instead of "content".
|
||||||
@@ -147,7 +142,6 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
|
|||||||
detect_by_base_keyword="openrouter",
|
detect_by_base_keyword="openrouter",
|
||||||
default_api_base="https://openrouter.ai/api/v1",
|
default_api_base="https://openrouter.ai/api/v1",
|
||||||
supports_prompt_caching=True,
|
supports_prompt_caching=True,
|
||||||
gateway_reasoning_style="reasoning_effort",
|
|
||||||
),
|
),
|
||||||
# Hugging Face Inference Providers: OpenAI-compatible router for chat models.
|
# Hugging Face Inference Providers: OpenAI-compatible router for chat models.
|
||||||
ProviderSpec(
|
ProviderSpec(
|
||||||
@@ -161,18 +155,6 @@ 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".
|
||||||
@@ -199,18 +181,6 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
|
|||||||
default_api_base="https://api.siliconflow.cn/v1",
|
default_api_base="https://api.siliconflow.cn/v1",
|
||||||
),
|
),
|
||||||
|
|
||||||
# Novita AI: OpenAI-compatible gateway for hosted model APIs.
|
|
||||||
ProviderSpec(
|
|
||||||
name="novita",
|
|
||||||
keywords=("novita",),
|
|
||||||
env_key="NOVITA_API_KEY",
|
|
||||||
display_name="Novita AI",
|
|
||||||
backend="openai_compat",
|
|
||||||
is_gateway=True,
|
|
||||||
detect_by_base_keyword="novita",
|
|
||||||
default_api_base="https://api.novita.ai/openai",
|
|
||||||
),
|
|
||||||
|
|
||||||
# VolcEngine (火山引擎): OpenAI-compatible gateway, pay-per-use models
|
# VolcEngine (火山引擎): OpenAI-compatible gateway, pay-per-use models
|
||||||
ProviderSpec(
|
ProviderSpec(
|
||||||
name="volcengine",
|
name="volcengine",
|
||||||
@@ -420,16 +390,6 @@ 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(
|
||||||
|
|||||||
@@ -165,23 +165,6 @@ class Session:
|
|||||||
image_placeholder_text(p) for p in media if isinstance(p, str) and p
|
image_placeholder_text(p) for p in media if isinstance(p, str) and p
|
||||||
)
|
)
|
||||||
content = f"{content}\n{breadcrumbs}" if content else breadcrumbs
|
content = f"{content}\n{breadcrumbs}" if content else breadcrumbs
|
||||||
cli_apps = message.get("cli_apps")
|
|
||||||
if role == "user" and isinstance(cli_apps, list) and cli_apps and isinstance(content, str):
|
|
||||||
cli_lines: list[str] = []
|
|
||||||
for item in cli_apps[:8]:
|
|
||||||
if not isinstance(item, dict):
|
|
||||||
continue
|
|
||||||
name = str(item.get("name") or "").strip().lower()
|
|
||||||
if not name:
|
|
||||||
continue
|
|
||||||
entry = str(item.get("entry_point") or "unknown").strip() or "unknown"
|
|
||||||
cli_lines.append(
|
|
||||||
f"[CLI App Attachment: @{name}; tool=run_cli_app; entry_point={entry}; "
|
|
||||||
f"skill=skills/cli-app-{name}/SKILL.md]"
|
|
||||||
)
|
|
||||||
if cli_lines:
|
|
||||||
breadcrumbs = "\n".join(cli_lines)
|
|
||||||
content = f"{content}\n{breadcrumbs}" if content else breadcrumbs
|
|
||||||
if include_timestamps:
|
if include_timestamps:
|
||||||
content = self._annotate_message_time(message, content)
|
content = self._annotate_message_time(message, content)
|
||||||
if role == "assistant" and isinstance(content, str) and not content.strip():
|
if role == "assistant" and isinstance(content, str) and not content.strip():
|
||||||
|
|||||||
@@ -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.
|
||||||
- After generating images, call the `message` tool with the artifact paths in the `media` parameter to deliver them to the user.
|
- 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.
|
||||||
|
|
||||||
## Prompt Rules
|
## Prompt Rules
|
||||||
|
|
||||||
@@ -42,6 +42,73 @@ 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.
|
||||||
|
|
||||||
|
For Gemini, the image tool supports two model families. Imagen 4 (`imagen-4.0-generate-001`) supports text-to-image only. Gemini Flash (`gemini-2.5-flash-image`) also supports reference-image edits. Configuration:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"gemini": {
|
||||||
|
"apiKey": "AIza..."
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "gemini",
|
||||||
|
"model": "imagen-4.0-generate-001"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
For Gemini models, `defaultImageSize` has no effect; use `defaultAspectRatio` instead. Imagen 4 supports `1:1`, `9:16`, `16:9`, `3:4`, and `4:3`.
|
||||||
|
|
||||||
## Examples
|
## Examples
|
||||||
|
|
||||||
Generate a new image:
|
Generate a new image:
|
||||||
|
|||||||
@@ -1,9 +1,5 @@
|
|||||||
# Agent Instructions
|
# Agent Instructions
|
||||||
|
|
||||||
## Workspace Guidance
|
|
||||||
|
|
||||||
Use this file for project-specific preferences, recurring workflow conventions, and instructions you want the agent to remember for this workspace. Keep durable facts about the user in `USER.md`, personality/style guidance in `SOUL.md`, and long-term memory in `memory/MEMORY.md`.
|
|
||||||
|
|
||||||
## Scheduled Reminders
|
## Scheduled Reminders
|
||||||
|
|
||||||
Before scheduling reminders, check available skills and follow skill guidance first.
|
Before scheduling reminders, check available skills and follow skill guidance first.
|
||||||
@@ -14,10 +10,10 @@ Get USER_ID and CHANNEL from the current session (e.g., `8281248569` and `telegr
|
|||||||
|
|
||||||
## Heartbeat Tasks
|
## Heartbeat Tasks
|
||||||
|
|
||||||
`HEARTBEAT.md` is checked on the configured heartbeat interval. Use file tools to manage periodic tasks.
|
`HEARTBEAT.md` is checked on the configured heartbeat interval. Use file tools to manage periodic tasks:
|
||||||
|
|
||||||
- Use `apply_patch` for normal task-list updates, especially when adding, removing, or changing multiple lines.
|
- **Add**: `edit_file` to append new tasks
|
||||||
- Use `edit_file` only for small exact replacements copied from the current `HEARTBEAT.md`.
|
- **Remove**: `edit_file` to delete completed tasks
|
||||||
- Use `write_file` for first creation or intentional full-file rewrites.
|
- **Rewrite**: `write_file` to replace all tasks
|
||||||
|
|
||||||
When the user asks for a recurring/periodic task, update `HEARTBEAT.md` instead of creating a one-time cron reminder.
|
When the user asks for a recurring/periodic task, update `HEARTBEAT.md` instead of creating a one-time cron reminder.
|
||||||
|
|||||||
@@ -0,0 +1,28 @@
|
|||||||
|
# Tool Usage Notes
|
||||||
|
|
||||||
|
Tool signatures are provided automatically via function calling.
|
||||||
|
This file documents non-obvious constraints and usage patterns.
|
||||||
|
|
||||||
|
## exec — Safety Limits
|
||||||
|
|
||||||
|
- Commands have a configurable timeout (default 60s)
|
||||||
|
- Dangerous commands are blocked (rm -rf, format, dd, shutdown, etc.)
|
||||||
|
- Output is truncated at 10,000 characters
|
||||||
|
- `restrictToWorkspace` config can limit file access to the workspace
|
||||||
|
|
||||||
|
## grep — Content Search
|
||||||
|
|
||||||
|
- Use `grep` to search file contents inside the workspace
|
||||||
|
- Default behavior returns only matching file paths (`output_mode="files_with_matches"`)
|
||||||
|
- Supports optional `glob` filtering (e.g. `glob="*.py"`) plus `context_before` / `context_after`
|
||||||
|
- Supports `type="py"`, `type="ts"`, `type="md"` and similar shorthand filters
|
||||||
|
- Use `fixed_strings=true` for literal keywords containing regex characters
|
||||||
|
- Use `output_mode="files_with_matches"` to get only matching file paths
|
||||||
|
- Use `output_mode="count"` to size a search before reading full matches
|
||||||
|
- Use `head_limit` and `offset` to page across results
|
||||||
|
- Prefer this over `exec` for code and history searches
|
||||||
|
- Binary or oversized files may be skipped to keep results readable
|
||||||
|
|
||||||
|
## cron — Scheduled Reminders
|
||||||
|
|
||||||
|
- Please refer to cron skill for usage.
|
||||||
@@ -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 'generate_image' creates images, call 'message' with the artifact paths in the 'media' parameter to deliver them to the user.
|
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.
|
||||||
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,67 +0,0 @@
|
|||||||
# Tool Usage Notes
|
|
||||||
|
|
||||||
Tool signatures are provided automatically via function calling. This section
|
|
||||||
documents the general tool contract and non-obvious usage patterns.
|
|
||||||
|
|
||||||
## General Tool Contract
|
|
||||||
|
|
||||||
- Use the narrowest structured tool that directly matches the task.
|
|
||||||
- Use read-only discovery before writes when state is uncertain.
|
|
||||||
- Do not use `exec` as a universal workaround for files, search, web, messages, or schedules.
|
|
||||||
- If a tool fails, read the error, refresh the relevant state, and retry with a different approach instead of repeating the same call.
|
|
||||||
- After meaningful changes, verify with the smallest reliable check: re-read changed state, run targeted tests, or inspect command output.
|
|
||||||
- Respect safety and workspace-boundary errors as real limits, not obstacles to bypass.
|
|
||||||
|
|
||||||
## Discovery and Reading
|
|
||||||
|
|
||||||
- Use `find_files` or `list_dir` to locate workspace paths before `read_file` when a path is uncertain.
|
|
||||||
- Use `grep` for content search inside the workspace; prefer it over shell grep for ordinary searches.
|
|
||||||
- `grep` defaults to `output_mode="files_with_matches"`; use `output_mode="content"` for matching lines with context.
|
|
||||||
- Use `fixed_strings=true` for literal keywords containing regex characters.
|
|
||||||
- Use `output_mode="count"` to size a broad search before reading full matches.
|
|
||||||
- Use `head_limit` and `offset` to page across large result sets.
|
|
||||||
- Binary or oversized files may be skipped to keep results readable.
|
|
||||||
|
|
||||||
## File and Coding Workflows
|
|
||||||
|
|
||||||
- For code or config changes, the default loop is: locate (`find_files`/`grep`), inspect (`read_file`), edit (`apply_patch`), then verify (`exec` or re-read).
|
|
||||||
- Use `apply_patch` as the default code editing tool, especially for multi-file changes, structural edits, generated code, moves, adds, or deletes.
|
|
||||||
- Use `apply_patch dry_run=true` when the patch is uncertain and you want validation plus a change summary before writing.
|
|
||||||
- Use `edit_file` only for small exact replacements in one file, with `old_text` copied from `read_file`; add `occurrence`, `line_hint`, or `expected_replacements` when ambiguity matters.
|
|
||||||
- Use `write_file` for new files or intentional full-file rewrites, not routine partial edits.
|
|
||||||
- If `apply_patch` or `edit_file` fails, re-read with `force=true`, narrow the context, and try a smaller patch rather than switching to shell `sed` or `echo`.
|
|
||||||
|
|
||||||
## Process Execution
|
|
||||||
|
|
||||||
- Use `exec` for tests, builds, package commands, git commands, and other process execution.
|
|
||||||
- Prefer dedicated file/search tools over `cat`, shell `find`, shell `grep`, `sed`, or `echo` for ordinary workspace inspection and edits.
|
|
||||||
- Use non-interactive flags such as `-y` or `--yes` when available.
|
|
||||||
- Commands have a configurable timeout (default 60s), dangerous commands are blocked, and output is truncated.
|
|
||||||
- For long-running or interactive commands, pass `yield_time_ms`; if the process keeps running, continue with `write_stdin`.
|
|
||||||
- Use `write_stdin` to poll, provide stdin, close stdin, wait for expected output with `wait_for`, or terminate an existing exec session.
|
|
||||||
- Use `list_exec_sessions` to recover active session IDs after context shifts.
|
|
||||||
|
|
||||||
## CLI App Attachments
|
|
||||||
|
|
||||||
- When Runtime Context lists a `CLI App Attachment` or `CLI App Mention`, treat the `@name` as an app capability the user intentionally attached to the current turn.
|
|
||||||
- If the task may need app-specific behavior, read the listed skill first, then call `run_cli_app` with that `name`.
|
|
||||||
- Do not run an attached CLI app through shell or generic process tools unless the user explicitly asks for that lower-level path.
|
|
||||||
- If the app CLI is missing, lacks local desktop/app/API prerequisites, or cannot complete the requested action, explain that concrete blocker and what was attempted.
|
|
||||||
|
|
||||||
## Web and External Information
|
|
||||||
|
|
||||||
- Use web tools when the user asks for current information, a specific URL, or information likely to have changed.
|
|
||||||
- Use `web_search` to find sources and `web_fetch` for a specific page or result that needs closer reading.
|
|
||||||
- Do not invent freshness-sensitive facts when tools can verify them.
|
|
||||||
|
|
||||||
## Messaging and Media
|
|
||||||
|
|
||||||
- Use `message` to send content or local media to the user/channel.
|
|
||||||
- `read_file` only reads content for your analysis; it does not deliver a file to the user.
|
|
||||||
- When sending an existing local file, attach it through the message/media mechanism instead of pasting file contents unless the user asked for text.
|
|
||||||
|
|
||||||
## Scheduling and Background Work
|
|
||||||
|
|
||||||
- Use `cron` for scheduled reminders or recurring jobs; do not run `nanobot cron` through `exec`.
|
|
||||||
- For heartbeat tasks, update `HEARTBEAT.md` according to the agent instructions.
|
|
||||||
- Do not write reminders only to memory files when the user expects an actual notification.
|
|
||||||
@@ -1,42 +1,6 @@
|
|||||||
"""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,6 +21,8 @@ _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."""
|
||||||
@@ -113,10 +115,48 @@ 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. "
|
||||||
"Call the message tool with the artifact paths in the media parameter "
|
"For the current chat, reply naturally; the runtime attaches generated images automatically. "
|
||||||
"to deliver the images to the user. Keep raw paths internal unless the "
|
"Do not call message just to announce or resend them. Keep raw paths internal unless the user asks for debug details."
|
||||||
"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
|
||||||
|
|||||||
@@ -3,16 +3,14 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import difflib
|
import difflib
|
||||||
import re
|
import json
|
||||||
import time
|
from dataclasses import dataclass
|
||||||
from dataclasses import dataclass, field
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Awaitable, Callable
|
from typing import Any
|
||||||
|
|
||||||
TRACKED_FILE_EDIT_TOOLS = frozenset({"write_file", "edit_file", "apply_patch"})
|
|
||||||
|
TRACKED_FILE_EDIT_TOOLS = frozenset({"write_file", "edit_file", "notebook_edit"})
|
||||||
_MAX_SNAPSHOT_BYTES = 2 * 1024 * 1024
|
_MAX_SNAPSHOT_BYTES = 2 * 1024 * 1024
|
||||||
_LIVE_EMIT_INTERVAL_S = 0.18
|
|
||||||
_LIVE_EMIT_LINE_STEP = 24
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
@@ -105,8 +103,6 @@ def line_diff_stats(before: str | None, after: str | None) -> tuple[int, int]:
|
|||||||
"""Return ``(added, deleted)`` for a UTF-8 text line-level diff."""
|
"""Return ``(added, deleted)`` for a UTF-8 text line-level diff."""
|
||||||
if before is None or after is None:
|
if before is None or after is None:
|
||||||
return 0, 0
|
return 0, 0
|
||||||
if before == "":
|
|
||||||
return _text_line_count(after), 0
|
|
||||||
before_lines = before.replace("\r\n", "\n").splitlines()
|
before_lines = before.replace("\r\n", "\n").splitlines()
|
||||||
after_lines = after.replace("\r\n", "\n").splitlines()
|
after_lines = after.replace("\r\n", "\n").splitlines()
|
||||||
added = 0
|
added = 0
|
||||||
@@ -122,28 +118,6 @@ def line_diff_stats(before: str | None, after: str | None) -> tuple[int, int]:
|
|||||||
return added, deleted
|
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(
|
def prepare_file_edit_tracker(
|
||||||
*,
|
*,
|
||||||
call_id: str,
|
call_id: str,
|
||||||
@@ -152,108 +126,19 @@ def prepare_file_edit_tracker(
|
|||||||
workspace: Path | None,
|
workspace: Path | None,
|
||||||
params: dict[str, Any] | None,
|
params: dict[str, Any] | None,
|
||||||
) -> FileEditTracker | None:
|
) -> FileEditTracker | None:
|
||||||
trackers = prepare_file_edit_trackers(
|
|
||||||
call_id=call_id,
|
|
||||||
tool_name=tool_name,
|
|
||||||
tool=tool,
|
|
||||||
workspace=workspace,
|
|
||||||
params=params,
|
|
||||||
)
|
|
||||||
return trackers[0] if trackers else None
|
|
||||||
|
|
||||||
|
|
||||||
def prepare_file_edit_trackers(
|
|
||||||
*,
|
|
||||||
call_id: str,
|
|
||||||
tool_name: str,
|
|
||||||
tool: Any,
|
|
||||||
workspace: Path | None,
|
|
||||||
params: dict[str, Any] | None,
|
|
||||||
) -> list[FileEditTracker]:
|
|
||||||
if not is_file_edit_tool(tool_name):
|
if not is_file_edit_tool(tool_name):
|
||||||
return []
|
return None
|
||||||
paths = resolve_file_edit_paths(tool_name, tool, workspace, params)
|
|
||||||
trackers: list[FileEditTracker] = []
|
|
||||||
seen: set[Path] = set()
|
|
||||||
for path in paths:
|
|
||||||
try:
|
|
||||||
resolved = path.resolve()
|
|
||||||
except Exception:
|
|
||||||
resolved = path
|
|
||||||
if resolved in seen:
|
|
||||||
continue
|
|
||||||
seen.add(resolved)
|
|
||||||
before = read_file_snapshot(path)
|
|
||||||
trackers.append(FileEditTracker(
|
|
||||||
call_id=str(call_id or ""),
|
|
||||||
tool=tool_name,
|
|
||||||
path=path,
|
|
||||||
display_path=display_file_edit_path(path, workspace),
|
|
||||||
before=before,
|
|
||||||
))
|
|
||||||
return trackers
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_file_edit_paths(
|
|
||||||
tool_name: str,
|
|
||||||
tool: Any,
|
|
||||||
workspace: Path | None,
|
|
||||||
params: dict[str, Any] | None,
|
|
||||||
) -> list[Path]:
|
|
||||||
if tool_name == "apply_patch":
|
|
||||||
return _resolve_apply_patch_paths(tool, workspace, params)
|
|
||||||
path = resolve_file_edit_path(tool, workspace, params)
|
path = resolve_file_edit_path(tool, workspace, params)
|
||||||
if path is None:
|
if path is None:
|
||||||
return []
|
return None
|
||||||
return [path]
|
before = read_file_snapshot(path)
|
||||||
|
return FileEditTracker(
|
||||||
|
call_id=str(call_id or ""),
|
||||||
def _resolve_apply_patch_paths(
|
tool=tool_name,
|
||||||
tool: Any,
|
path=path,
|
||||||
workspace: Path | None,
|
display_path=display_file_edit_path(path, workspace),
|
||||||
params: dict[str, Any] | None,
|
before=before,
|
||||||
) -> list[Path]:
|
)
|
||||||
if not isinstance(params, dict):
|
|
||||||
return []
|
|
||||||
edits = params.get("edits")
|
|
||||||
if not isinstance(edits, list) or not edits:
|
|
||||||
return []
|
|
||||||
if params.get("dry_run") is True:
|
|
||||||
return []
|
|
||||||
|
|
||||||
resolved: list[Path] = []
|
|
||||||
seen: set[Path] = set()
|
|
||||||
for edit in edits:
|
|
||||||
if not isinstance(edit, dict):
|
|
||||||
continue
|
|
||||||
raw_path = edit.get("path")
|
|
||||||
if not isinstance(raw_path, str) or not raw_path.strip():
|
|
||||||
continue
|
|
||||||
path = _resolve_raw_file_edit_path(tool, workspace, raw_path)
|
|
||||||
if path is not None and path not in seen:
|
|
||||||
seen.add(path)
|
|
||||||
resolved.append(path)
|
|
||||||
return resolved
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve_raw_file_edit_path(
|
|
||||||
tool: Any,
|
|
||||||
workspace: Path | None,
|
|
||||||
raw_path: str,
|
|
||||||
) -> Path | 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 build_file_edit_start_event(
|
def build_file_edit_start_event(
|
||||||
@@ -275,22 +160,12 @@ def build_file_edit_start_event(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_file_edit_end_event(
|
def build_file_edit_end_event(tracker: FileEditTracker) -> dict[str, Any]:
|
||||||
tracker: FileEditTracker,
|
|
||||||
params: dict[str, Any] | None = None,
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
after = read_file_snapshot(tracker.path)
|
after = read_file_snapshot(tracker.path)
|
||||||
counted = False
|
|
||||||
if tracker.before.countable and after.countable:
|
if tracker.before.countable and after.countable:
|
||||||
added, deleted = line_diff_stats(tracker.before.text, after.text)
|
added, deleted = line_diff_stats(tracker.before.text, after.text)
|
||||||
counted = True
|
|
||||||
else:
|
else:
|
||||||
predicted_after = _predict_after_text(tracker.tool, params or {}, tracker.before)
|
added, deleted = 0, 0
|
||||||
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(
|
return _event_payload(
|
||||||
tracker,
|
tracker,
|
||||||
phase="end",
|
phase="end",
|
||||||
@@ -298,14 +173,11 @@ def build_file_edit_end_event(
|
|||||||
added=added,
|
added=added,
|
||||||
deleted=deleted,
|
deleted=deleted,
|
||||||
approximate=False,
|
approximate=False,
|
||||||
binary=(after.binary or after.oversized or after.unreadable) and not counted,
|
binary=after.binary or after.oversized or after.unreadable,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_file_edit_error_event(
|
def build_file_edit_error_event(tracker: FileEditTracker, error: str | None = None) -> dict[str, Any]:
|
||||||
tracker: FileEditTracker,
|
|
||||||
error: str | None = None,
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
payload = _event_payload(
|
payload = _event_payload(
|
||||||
tracker,
|
tracker,
|
||||||
phase="error",
|
phase="error",
|
||||||
@@ -319,594 +191,6 @@ def build_file_edit_error_event(
|
|||||||
return payload
|
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 == "apply_patch":
|
|
||||||
await self._update_apply_patch(state)
|
|
||||||
return
|
|
||||||
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 _update_apply_patch(self, state: _StreamingFileEditState) -> None:
|
|
||||||
if _json_bool_true(state.arguments, "dry_run"):
|
|
||||||
return
|
|
||||||
tool = self._tools.get("apply_patch") if hasattr(self._tools, "get") else None
|
|
||||||
events: list[dict[str, Any]] = []
|
|
||||||
now = time.monotonic()
|
|
||||||
|
|
||||||
path_matches = list(re.finditer(r'"path"\s*:\s*"([^"]+)"', state.arguments))
|
|
||||||
if not path_matches:
|
|
||||||
return
|
|
||||||
|
|
||||||
for i, m in enumerate(path_matches):
|
|
||||||
raw_path = m.group(1)
|
|
||||||
path = _resolve_raw_file_edit_path(tool, self._workspace, raw_path)
|
|
||||||
if path is None:
|
|
||||||
continue
|
|
||||||
|
|
||||||
segment_start = m.start()
|
|
||||||
segment_end = path_matches[i + 1].start() if i + 1 < len(path_matches) else len(state.arguments)
|
|
||||||
segment = state.arguments[segment_start:segment_end]
|
|
||||||
|
|
||||||
action_match = re.search(r'"action"\s*:\s*"(replace|add|delete)"', segment)
|
|
||||||
action = action_match.group(1) if action_match else "replace"
|
|
||||||
|
|
||||||
old_text = _extract_json_string_prefix(segment, "old_text") or ""
|
|
||||||
new_text = _extract_json_string_prefix(segment, "new_text") or ""
|
|
||||||
|
|
||||||
added = _text_line_count(new_text) if action in ("replace", "add") else 0
|
|
||||||
deleted = _text_line_count(old_text) if action in ("replace", "delete") else 0
|
|
||||||
delete_file = action == "delete"
|
|
||||||
|
|
||||||
file_state = state.patch_files.get(raw_path)
|
|
||||||
if file_state is None:
|
|
||||||
tracker = FileEditTracker(
|
|
||||||
call_id=state.call_id or state.key,
|
|
||||||
tool="apply_patch",
|
|
||||||
path=path,
|
|
||||||
display_path=display_file_edit_path(path, self._workspace),
|
|
||||||
before=read_file_snapshot(path),
|
|
||||||
)
|
|
||||||
file_state = _StreamingPatchFileState(tracker=tracker)
|
|
||||||
state.patch_files[raw_path] = file_state
|
|
||||||
if delete_file and added == 0 and deleted == 0 and file_state.tracker.before.countable:
|
|
||||||
deleted = _text_line_count(file_state.tracker.before.text or "")
|
|
||||||
if not file_state.should_emit(added, deleted, now):
|
|
||||||
continue
|
|
||||||
file_state.mark_emitted(added, deleted, now)
|
|
||||||
events.append(build_file_edit_live_event(
|
|
||||||
file_state.tracker,
|
|
||||||
added=added,
|
|
||||||
deleted=deleted,
|
|
||||||
))
|
|
||||||
if events:
|
|
||||||
await self._emit(events)
|
|
||||||
|
|
||||||
async def flush(self) -> None:
|
|
||||||
events: list[dict[str, Any]] = []
|
|
||||||
now = time.monotonic()
|
|
||||||
for state in self._states.values():
|
|
||||||
for file_state in state.patch_files.values():
|
|
||||||
added, deleted = file_state.last_added, file_state.last_deleted
|
|
||||||
if not file_state.emitted_once:
|
|
||||||
continue
|
|
||||||
if (
|
|
||||||
file_state.last_emitted_added == added
|
|
||||||
and file_state.last_emitted_deleted == deleted
|
|
||||||
):
|
|
||||||
continue
|
|
||||||
file_state.mark_emitted(added, deleted, now)
|
|
||||||
events.append(build_file_edit_live_event(
|
|
||||||
file_state.tracker,
|
|
||||||
added=added,
|
|
||||||
deleted=deleted,
|
|
||||||
))
|
|
||||||
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."""
|
|
||||||
used_canonicals: set[str] = set()
|
|
||||||
for tool_call in final_tool_calls:
|
|
||||||
canonical = self.canonical_call_id_for(tool_call)
|
|
||||||
if canonical and canonical not in used_canonicals:
|
|
||||||
try:
|
|
||||||
tool_call.id = canonical
|
|
||||||
used_canonicals.add(canonical)
|
|
||||||
except (AttributeError, TypeError):
|
|
||||||
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():
|
|
||||||
for file_state in state.patch_files.values():
|
|
||||||
if any(state.matches_final_tool_call(tool_call) for tool_call in final_tool_calls):
|
|
||||||
continue
|
|
||||||
events.append(build_file_edit_error_event(file_state.tracker, error))
|
|
||||||
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 _StreamingPatchFileState:
|
|
||||||
tracker: FileEditTracker
|
|
||||||
emitted_once: bool = False
|
|
||||||
last_emitted_added: int = -1
|
|
||||||
last_emitted_deleted: int = -1
|
|
||||||
last_emit_at: float = 0.0
|
|
||||||
last_added: int = 0
|
|
||||||
last_deleted: int = 0
|
|
||||||
|
|
||||||
def should_emit(self, added: int, deleted: int, now: float) -> bool:
|
|
||||||
self.last_added = added
|
|
||||||
self.last_deleted = deleted
|
|
||||||
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_added = added
|
|
||||||
self.last_deleted = deleted
|
|
||||||
self.last_emitted_added = added
|
|
||||||
self.last_emitted_deleted = deleted
|
|
||||||
self.last_emit_at = now
|
|
||||||
|
|
||||||
|
|
||||||
@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")
|
|
||||||
)
|
|
||||||
patch_files: dict[str, _StreamingPatchFileState] = field(default_factory=dict)
|
|
||||||
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()
|
|
||||||
self.patch_files.clear()
|
|
||||||
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
|
|
||||||
if self.name == "apply_patch":
|
|
||||||
arguments = getattr(tool_call, "arguments", None)
|
|
||||||
if not isinstance(arguments, dict):
|
|
||||||
return False
|
|
||||||
edits = arguments.get("edits")
|
|
||||||
if not isinstance(edits, list):
|
|
||||||
return False
|
|
||||||
return '"edits"' in self.arguments
|
|
||||||
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 _json_bool_true(source: str, key: str) -> bool:
|
|
||||||
return re.search(rf'"{re.escape(key)}"\s*:\s*true\b', source) is not None
|
|
||||||
|
|
||||||
|
|
||||||
def _extract_json_string_prefix(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:
|
|
||||||
break
|
|
||||||
try:
|
|
||||||
out.append(chr(int(digits, 16)))
|
|
||||||
except ValueError:
|
|
||||||
break
|
|
||||||
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 "".join(out)
|
|
||||||
|
|
||||||
|
|
||||||
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(
|
def _event_payload(
|
||||||
tracker: FileEditTracker,
|
tracker: FileEditTracker,
|
||||||
*,
|
*,
|
||||||
@@ -922,7 +206,6 @@ def _event_payload(
|
|||||||
"call_id": tracker.call_id,
|
"call_id": tracker.call_id,
|
||||||
"tool": tracker.tool,
|
"tool": tracker.tool,
|
||||||
"path": tracker.display_path,
|
"path": tracker.display_path,
|
||||||
"absolute_path": tracker.path.as_posix(),
|
|
||||||
"phase": phase,
|
"phase": phase,
|
||||||
"added": max(0, int(added)),
|
"added": max(0, int(added)),
|
||||||
"deleted": max(0, int(deleted)),
|
"deleted": max(0, int(deleted)),
|
||||||
@@ -958,4 +241,71 @@ def _predict_after_text(
|
|||||||
return before_text.replace(old_text, new_text)
|
return before_text.replace(old_text, new_text)
|
||||||
return before_text.replace(old_text, new_text, 1)
|
return before_text.replace(old_text, new_text, 1)
|
||||||
return None
|
return None
|
||||||
|
if tool_name == "notebook_edit":
|
||||||
|
return _predict_notebook_after_text(params, before_text)
|
||||||
return None
|
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
|
||||||
|
|||||||
@@ -576,7 +576,7 @@ def build_status_content(
|
|||||||
|
|
||||||
|
|
||||||
def sync_workspace_templates(workspace: Path, silent: bool = False) -> list[str]:
|
def sync_workspace_templates(workspace: Path, silent: bool = False) -> list[str]:
|
||||||
"""Sync bundled templates to workspace. Creates missing files without overwriting user files."""
|
"""Sync bundled templates to workspace. Only creates missing files."""
|
||||||
from importlib.resources import files as pkg_files
|
from importlib.resources import files as pkg_files
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -589,11 +589,10 @@ def sync_workspace_templates(workspace: Path, silent: bool = False) -> list[str]
|
|||||||
added: list[str] = []
|
added: list[str] = []
|
||||||
|
|
||||||
def _write(src, dest: Path):
|
def _write(src, dest: Path):
|
||||||
content = src.read_text(encoding="utf-8") if src else ""
|
|
||||||
if dest.exists():
|
if dest.exists():
|
||||||
return
|
return
|
||||||
dest.parent.mkdir(parents=True, exist_ok=True)
|
dest.parent.mkdir(parents=True, exist_ok=True)
|
||||||
dest.write_text(content, encoding="utf-8")
|
dest.write_text(src.read_text(encoding="utf-8") if src else "", encoding="utf-8")
|
||||||
added.append(str(dest.relative_to(workspace)))
|
added.append(str(dest.relative_to(workspace)))
|
||||||
|
|
||||||
for item in tpl.iterdir():
|
for item in tpl.iterdir():
|
||||||
|
|||||||
@@ -0,0 +1,74 @@
|
|||||||
|
"""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]))
|
||||||
@@ -11,10 +11,8 @@ _TOOL_FORMATS: dict[str, tuple[list[str], str, bool, bool]] = {
|
|||||||
"read_file": (["path", "file_path"], "read {}", True, False),
|
"read_file": (["path", "file_path"], "read {}", True, False),
|
||||||
"write_file": (["path", "file_path"], "write {}", True, False),
|
"write_file": (["path", "file_path"], "write {}", True, False),
|
||||||
"edit": (["file_path", "path"], "edit {}", True, False),
|
"edit": (["file_path", "path"], "edit {}", True, False),
|
||||||
"find_files": (["query", "glob", "path"], "find {}", False, False),
|
|
||||||
"grep": (["pattern"], 'grep "{}"', False, False),
|
"grep": (["pattern"], 'grep "{}"', False, False),
|
||||||
"exec": (["command"], "$ {}", False, True),
|
"exec": (["command"], "$ {}", False, True),
|
||||||
"list_exec_sessions": ([], "exec sessions", False, False),
|
|
||||||
"web_search": (["query"], 'search "{}"', False, False),
|
"web_search": (["query"], 'search "{}"', False, False),
|
||||||
"web_fetch": (["url"], "fetch {}", True, False),
|
"web_fetch": (["url"], "fetch {}", True, False),
|
||||||
"list_dir": (["path"], "ls {}", True, False),
|
"list_dir": (["path"], "ls {}", True, False),
|
||||||
@@ -83,8 +81,6 @@ def _extract_arg(tc, key_args: list[str]) -> str | None:
|
|||||||
|
|
||||||
def _fmt_known(tc, fmt: tuple, max_length: int = 40) -> str:
|
def _fmt_known(tc, fmt: tuple, max_length: int = 40) -> str:
|
||||||
"""Format a registered tool using its template."""
|
"""Format a registered tool using its template."""
|
||||||
if not fmt[0] and "{}" not in fmt[1]:
|
|
||||||
return fmt[1]
|
|
||||||
val = _extract_arg(tc, fmt[0])
|
val = _extract_arg(tc, fmt[0])
|
||||||
if val is None:
|
if val is None:
|
||||||
return tc.name
|
return tc.name
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Legacy WebUI JSON snapshot path helpers (JSON file); transcripts use transcript."""
|
"""Legacy WebUI JSON snapshot path helpers (JSON file); transcripts use webui_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.webui.transcript import delete_webui_transcript
|
from nanobot.utils.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,93 +99,21 @@ 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") not in {"start", "end", "error"}:
|
if event.get("phase") != "start":
|
||||||
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
|
||||||
|
|
||||||
|
|
||||||
_PHASE_RANK = {"start": 1, "end": 2, "error": 3}
|
|
||||||
|
|
||||||
|
|
||||||
def _normalize_tool_events(events: Any) -> list[dict[str, Any]]:
|
|
||||||
if not isinstance(events, list):
|
|
||||||
return []
|
|
||||||
out: list[dict[str, Any]] = []
|
|
||||||
for event in events:
|
|
||||||
if not event or not isinstance(event, dict):
|
|
||||||
continue
|
|
||||||
if event.get("phase") not in {"start", "end", "error"}:
|
|
||||||
continue
|
|
||||||
if not isinstance(event.get("name"), str):
|
|
||||||
fn = event.get("function")
|
|
||||||
if not (isinstance(fn, dict) and isinstance(fn.get("name"), str)):
|
|
||||||
continue
|
|
||||||
out.append(dict(event))
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def _tool_event_key(event: dict[str, Any]) -> str:
|
|
||||||
call_id = event.get("call_id")
|
|
||||||
if isinstance(call_id, str) and call_id:
|
|
||||||
return f"call:{call_id}"
|
|
||||||
return _format_tool_call_trace(event) or json.dumps(event, sort_keys=True, ensure_ascii=False)
|
|
||||||
|
|
||||||
|
|
||||||
def _merge_tool_events(previous: Any, incoming: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|
||||||
if not isinstance(previous, list) or not previous:
|
|
||||||
return incoming
|
|
||||||
if not incoming:
|
|
||||||
return [dict(event) for event in previous if isinstance(event, dict)]
|
|
||||||
merged = [dict(event) for event in previous if isinstance(event, dict)]
|
|
||||||
index_by_key = {_tool_event_key(event): idx for idx, event in enumerate(merged)}
|
|
||||||
for event in incoming:
|
|
||||||
key = _tool_event_key(event)
|
|
||||||
existing_index = index_by_key.get(key)
|
|
||||||
if existing_index is None:
|
|
||||||
index_by_key[key] = len(merged)
|
|
||||||
merged.append(event)
|
|
||||||
continue
|
|
||||||
existing = merged[existing_index]
|
|
||||||
incoming_rank = _PHASE_RANK.get(str(event.get("phase")), 0)
|
|
||||||
existing_rank = _PHASE_RANK.get(str(existing.get("phase")), 0)
|
|
||||||
if incoming_rank >= existing_rank:
|
|
||||||
merged[existing_index] = {**existing, **event}
|
|
||||||
return merged
|
|
||||||
|
|
||||||
|
|
||||||
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]],
|
||||||
*,
|
*,
|
||||||
augment_user_media: Callable[[list[str]], list[dict[str, Any]]] | None = None,
|
augment_user_media: Callable[[list[str]], list[dict[str, Any]]] | None = None,
|
||||||
augment_assistant_text: Callable[[str], str] | None = None,
|
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Fold JSONL records into ``UIMessage``-shaped dicts for the WebUI.
|
"""Fold JSONL records into ``UIMessage``-shaped dicts for the WebUI.
|
||||||
|
|
||||||
@@ -216,17 +144,6 @@ def replay_transcript_to_ui_messages(
|
|||||||
def _ensure_activity_segment() -> str:
|
def _ensure_activity_segment() -> str:
|
||||||
return active_activity_segment_id or _new_activity_segment()
|
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]
|
||||||
@@ -326,7 +243,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
|
nonlocal active_activity_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] = {
|
||||||
@@ -345,50 +262,35 @@ def replay_transcript_to_ui_messages(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
active_activity_segment_id = None
|
active_activity_segment_id = None
|
||||||
active_file_edit_segment_id = None
|
|
||||||
|
|
||||||
def _file_edit_key(edit: dict[str, Any]) -> str:
|
def _file_edit_key(edit: dict[str, Any]) -> str:
|
||||||
call_id = str(edit.get("call_id") or "")
|
return "|".join(
|
||||||
tool = str(edit.get("tool") or "")
|
str(edit.get(k) or "")
|
||||||
if call_id:
|
for k in ("call_id", "tool", "path")
|
||||||
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:
|
def upsert_file_edits(edits: list[dict[str, Any]], idx: int) -> None:
|
||||||
nonlocal active_file_edit_segment_id
|
nonlocal active_file_edit_segment_id
|
||||||
if not edits:
|
if not edits:
|
||||||
return
|
return
|
||||||
segment = active_file_edit_segment_id
|
last = messages[-1] if messages else None
|
||||||
target_index = find_file_edit_trace_index(segment, edits)
|
if (
|
||||||
if target_index is not None:
|
active_file_edit_segment_id
|
||||||
last = messages[target_index]
|
and last
|
||||||
segment = str(last.get("activitySegmentId") or segment or _new_activity_segment(activate=False))
|
and last.get("kind") == "trace"
|
||||||
active_file_edit_segment_id = segment
|
and last.get("fileEdits")
|
||||||
|
):
|
||||||
|
segment = active_file_edit_segment_id
|
||||||
else:
|
else:
|
||||||
if not segment:
|
segment = _new_activity_segment(activate=False)
|
||||||
segment = _new_activity_segment(activate=False)
|
|
||||||
active_file_edit_segment_id = segment
|
active_file_edit_segment_id = segment
|
||||||
|
if not (
|
||||||
|
last
|
||||||
|
and last.get("kind") == "trace"
|
||||||
|
and not last.get("isStreaming")
|
||||||
|
and last.get("fileEdits")
|
||||||
|
and last.get("activitySegmentId") == segment
|
||||||
|
):
|
||||||
messages.append(
|
messages.append(
|
||||||
{
|
{
|
||||||
"id": _new_id("tr", idx),
|
"id": _new_id("tr", idx),
|
||||||
@@ -401,11 +303,7 @@ def replay_transcript_to_ui_messages(
|
|||||||
"createdAt": _ts_base + idx,
|
"createdAt": _ts_base + idx,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
target_index = len(messages) - 1
|
last = 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 [])
|
existing = list(last.get("fileEdits") or [])
|
||||||
index_by_key = {
|
index_by_key = {
|
||||||
_file_edit_key(edit): pos
|
_file_edit_key(edit): pos
|
||||||
@@ -418,14 +316,11 @@ def replay_transcript_to_ui_messages(
|
|||||||
key = _file_edit_key(edit)
|
key = _file_edit_key(edit)
|
||||||
if key in index_by_key:
|
if key in index_by_key:
|
||||||
pos = index_by_key[key]
|
pos = index_by_key[key]
|
||||||
merged = {**existing[pos], **edit}
|
existing[pos] = {**existing[pos], **edit}
|
||||||
if edit.get("path") and not edit.get("pending"):
|
|
||||||
merged.pop("pending", None)
|
|
||||||
existing[pos] = merged
|
|
||||||
else:
|
else:
|
||||||
index_by_key[key] = len(existing)
|
index_by_key[key] = len(existing)
|
||||||
existing.append(dict(edit))
|
existing.append(dict(edit))
|
||||||
messages[target_index] = {
|
messages[-1] = {
|
||||||
**last,
|
**last,
|
||||||
"fileEdits": existing,
|
"fileEdits": existing,
|
||||||
"activitySegmentId": last.get("activitySegmentId") or segment,
|
"activitySegmentId": last.get("activitySegmentId") or segment,
|
||||||
@@ -455,9 +350,6 @@ def replay_transcript_to_ui_messages(
|
|||||||
row["media"] = media_att
|
row["media"] = media_att
|
||||||
if all(m.get("kind") == "image" for m in media_att):
|
if all(m.get("kind") == "image" for m in media_att):
|
||||||
row["images"] = [{"url": m.get("url"), "name": m.get("name")} for m in media_att]
|
row["images"] = [{"url": m.get("url"), "name": m.get("name")} for m in media_att]
|
||||||
cli_apps = rec.get("cli_apps")
|
|
||||||
if isinstance(cli_apps, list) and cli_apps:
|
|
||||||
row["cliApps"] = [dict(app) for app in cli_apps if isinstance(app, dict)]
|
|
||||||
messages.append(row)
|
messages.append(row)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -473,7 +365,6 @@ def replay_transcript_to_ui_messages(
|
|||||||
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:
|
||||||
@@ -502,24 +393,6 @@ def replay_transcript_to_ui_messages(
|
|||||||
buffer_message_id = None
|
buffer_message_id = None
|
||||||
buffer_parts = []
|
buffer_parts = []
|
||||||
continue
|
continue
|
||||||
final_text = rec.get("text")
|
|
||||||
if isinstance(final_text, str):
|
|
||||||
if buffer_message_id is None:
|
|
||||||
buffer_message_id = _new_id("buf", idx)
|
|
||||||
messages.append(
|
|
||||||
{
|
|
||||||
"id": buffer_message_id,
|
|
||||||
"role": "assistant",
|
|
||||||
"content": final_text,
|
|
||||||
"isStreaming": True,
|
|
||||||
"createdAt": _ts_base + idx,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
for i, m in enumerate(messages):
|
|
||||||
if m.get("id") == buffer_message_id:
|
|
||||||
messages[i] = {**m, "content": final_text, "isStreaming": True}
|
|
||||||
break
|
|
||||||
buffer_message_id = None
|
buffer_message_id = None
|
||||||
buffer_parts = []
|
buffer_parts = []
|
||||||
continue
|
continue
|
||||||
@@ -530,7 +403,6 @@ 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
|
||||||
|
|
||||||
@@ -552,12 +424,10 @@ 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
|
||||||
if kind in ("tool_hint", "progress"):
|
if kind in ("tool_hint", "progress"):
|
||||||
structured_events = _normalize_tool_events(rec.get("tool_events"))
|
|
||||||
structured = tool_trace_lines_from_events(rec.get("tool_events"))
|
structured = tool_trace_lines_from_events(rec.get("tool_events"))
|
||||||
text = rec.get("text")
|
text = rec.get("text")
|
||||||
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 [])
|
||||||
@@ -572,22 +442,13 @@ def replay_transcript_to_ui_messages(
|
|||||||
and (last.get("activitySegmentId") in (None, segment))
|
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")])
|
||||||
if structured:
|
merged_traces = prev_traces + trace_lines
|
||||||
merged_traces, added = _merge_unique_tool_trace_lines(prev_traces, structured)
|
messages[-1] = {
|
||||||
if not added and not structured_events:
|
|
||||||
continue
|
|
||||||
else:
|
|
||||||
merged_traces = prev_traces + trace_lines
|
|
||||||
merged = {
|
|
||||||
**last,
|
**last,
|
||||||
"traces": merged_traces,
|
"traces": merged_traces,
|
||||||
"content": merged_traces[-1],
|
"content": trace_lines[-1],
|
||||||
"toolEvents": _merge_tool_events(last.get("toolEvents"), structured_events)
|
|
||||||
if structured_events
|
|
||||||
else last.get("toolEvents"),
|
|
||||||
"activitySegmentId": last.get("activitySegmentId") or segment,
|
"activitySegmentId": last.get("activitySegmentId") or segment,
|
||||||
}
|
}
|
||||||
messages[-1] = merged
|
|
||||||
else:
|
else:
|
||||||
messages.append(
|
messages.append(
|
||||||
{
|
{
|
||||||
@@ -596,7 +457,6 @@ def replay_transcript_to_ui_messages(
|
|||||||
"kind": "trace",
|
"kind": "trace",
|
||||||
"content": trace_lines[-1],
|
"content": trace_lines[-1],
|
||||||
"traces": trace_lines,
|
"traces": trace_lines,
|
||||||
**({"toolEvents": structured_events} if structured_events else {}),
|
|
||||||
"activitySegmentId": segment,
|
"activitySegmentId": segment,
|
||||||
"createdAt": _ts_base + idx,
|
"createdAt": _ts_base + idx,
|
||||||
},
|
},
|
||||||
@@ -645,14 +505,7 @@ def replay_transcript_to_ui_messages(
|
|||||||
buffer_parts = []
|
buffer_parts = []
|
||||||
continue
|
continue
|
||||||
|
|
||||||
for i, m in enumerate(messages):
|
for m in messages:
|
||||||
if (
|
|
||||||
augment_assistant_text is not None
|
|
||||||
and m.get("role") == "assistant"
|
|
||||||
and m.get("kind") != "trace"
|
|
||||||
and isinstance(m.get("content"), str)
|
|
||||||
):
|
|
||||||
messages[i] = {**m, "content": augment_assistant_text(m["content"])}
|
|
||||||
m.pop("isStreaming", None)
|
m.pop("isStreaming", None)
|
||||||
m.pop("reasoningStreaming", None)
|
m.pop("reasoningStreaming", None)
|
||||||
return messages
|
return messages
|
||||||
@@ -662,17 +515,12 @@ def build_webui_thread_response(
|
|||||||
session_key: str,
|
session_key: str,
|
||||||
*,
|
*,
|
||||||
augment_user_media: Callable[[list[str]], list[dict[str, Any]]] | None = None,
|
augment_user_media: Callable[[list[str]], list[dict[str, Any]]] | None = None,
|
||||||
augment_assistant_text: Callable[[str], str] | None = None,
|
|
||||||
) -> dict[str, Any] | None:
|
) -> dict[str, Any] | None:
|
||||||
"""Return a payload compatible with ``WebuiThreadPersistedPayload``."""
|
"""Return a payload compatible with ``WebuiThreadPersistedPayload``."""
|
||||||
lines = read_transcript_lines(session_key)
|
lines = read_transcript_lines(session_key)
|
||||||
if not lines:
|
if not lines:
|
||||||
return None
|
return None
|
||||||
msgs = replay_transcript_to_ui_messages(
|
msgs = replay_transcript_to_ui_messages(lines, augment_user_media=augment_user_media)
|
||||||
lines,
|
|
||||||
augment_user_media=augment_user_media,
|
|
||||||
augment_assistant_text=augment_assistant_text,
|
|
||||||
)
|
|
||||||
return {
|
return {
|
||||||
"schemaVersion": WEBUI_TRANSCRIPT_SCHEMA_VERSION,
|
"schemaVersion": WEBUI_TRANSCRIPT_SCHEMA_VERSION,
|
||||||
"sessionKey": session_key,
|
"sessionKey": session_key,
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Session turn helpers for WebUI-capable WebSocket sessions.
|
"""Outbound helpers for the WebSocket/WebUI wire contract.
|
||||||
|
|
||||||
AgentLoop uses these without importing a concrete channel plugin; only
|
AgentLoop uses these without importing a concrete channel plugin; only
|
||||||
``channel == "websocket"`` messages are affected.
|
``channel == "websocket"`` messages are affected.
|
||||||
@@ -1,2 +0,0 @@
|
|||||||
"""Backend helpers for the bundled WebUI surface."""
|
|
||||||
|
|
||||||
@@ -1,93 +0,0 @@
|
|||||||
"""CLI Apps helpers for the WebUI HTTP and message surfaces."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import re
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from nanobot.cli_apps import CliAppError, CliAppManager, CliAppsRuntimeConfig
|
|
||||||
from nanobot.config.loader import load_config
|
|
||||||
|
|
||||||
QueryParams = dict[str, list[str]]
|
|
||||||
|
|
||||||
_CLI_APP_NAME_RE = re.compile(r"^[a-z0-9][a-z0-9_-]{0,63}$", re.IGNORECASE)
|
|
||||||
_CLI_APP_ATTACHMENT_KEYS = (
|
|
||||||
"name",
|
|
||||||
"display_name",
|
|
||||||
"category",
|
|
||||||
"entry_point",
|
|
||||||
"logo_url",
|
|
||||||
"brand_color",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _clip_ws_string(value: Any, limit: int = 240) -> str | None:
|
|
||||||
if not isinstance(value, str):
|
|
||||||
return None
|
|
||||||
text = value.strip()
|
|
||||||
if not text:
|
|
||||||
return None
|
|
||||||
return text[:limit]
|
|
||||||
|
|
||||||
|
|
||||||
def normalize_cli_app_mentions(raw: Any) -> list[dict[str, str]]:
|
|
||||||
"""Sanitize structured CLI app mentions sent by the WebUI."""
|
|
||||||
if not isinstance(raw, list):
|
|
||||||
return []
|
|
||||||
out: list[dict[str, str]] = []
|
|
||||||
seen: set[str] = set()
|
|
||||||
for item in raw[:8]:
|
|
||||||
if not isinstance(item, dict):
|
|
||||||
continue
|
|
||||||
name = _clip_ws_string(item.get("name"), 64)
|
|
||||||
if not name or _CLI_APP_NAME_RE.match(name) is None:
|
|
||||||
continue
|
|
||||||
key = name.lower()
|
|
||||||
if key in seen:
|
|
||||||
continue
|
|
||||||
seen.add(key)
|
|
||||||
row: dict[str, str] = {"name": key}
|
|
||||||
for field in _CLI_APP_ATTACHMENT_KEYS[1:]:
|
|
||||||
value = _clip_ws_string(item.get(field), 512 if field == "logo_url" else 160)
|
|
||||||
if value:
|
|
||||||
row[field] = value
|
|
||||||
out.append(row)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def _query_first(query: QueryParams, key: str) -> str | None:
|
|
||||||
values = query.get(key)
|
|
||||||
return values[0] if values else None
|
|
||||||
|
|
||||||
|
|
||||||
def _manager() -> CliAppManager:
|
|
||||||
config = load_config()
|
|
||||||
cli_cfg = config.tools.cli_apps
|
|
||||||
return CliAppManager(
|
|
||||||
workspace=config.workspace_path,
|
|
||||||
runtime=CliAppsRuntimeConfig(
|
|
||||||
install_timeout=cli_cfg.install_timeout,
|
|
||||||
run_timeout=cli_cfg.run_timeout,
|
|
||||||
catalog_ttl_seconds=cli_cfg.catalog_ttl_seconds,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def cli_apps_payload() -> dict[str, Any]:
|
|
||||||
return _manager().payload()
|
|
||||||
|
|
||||||
|
|
||||||
def cli_apps_action(action: str, query: QueryParams) -> dict[str, Any]:
|
|
||||||
name = (_query_first(query, "name") or "").strip()
|
|
||||||
if not name:
|
|
||||||
raise CliAppError("missing CLI app name")
|
|
||||||
manager = _manager()
|
|
||||||
if action == "install":
|
|
||||||
return manager.install(name)
|
|
||||||
if action == "update":
|
|
||||||
return manager.update(name)
|
|
||||||
if action == "uninstall":
|
|
||||||
return manager.uninstall(name)
|
|
||||||
if action == "test":
|
|
||||||
return manager.test(name)
|
|
||||||
raise CliAppError(f"unknown CLI app action '{action}'", status=404)
|
|
||||||
@@ -1,613 +0,0 @@
|
|||||||
"""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_oauth:
|
|
||||||
return False
|
|
||||||
if spec.is_local or spec.is_direct:
|
|
||||||
return False
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
def _provider_configured_for_settings(spec: Any, provider_config: Any) -> bool:
|
|
||||||
if spec.is_oauth:
|
|
||||||
return True
|
|
||||||
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)
|
|
||||||
@@ -1,193 +0,0 @@
|
|||||||
"""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
|
|
||||||
|
|
||||||
@@ -139,13 +139,6 @@ class TestLoadBootstrapFiles:
|
|||||||
for name in ContextBuilder.BOOTSTRAP_FILES:
|
for name in ContextBuilder.BOOTSTRAP_FILES:
|
||||||
assert f"## {name}" in result
|
assert f"## {name}" in result
|
||||||
|
|
||||||
def test_legacy_tools_md_is_not_bootstrapped(self, tmp_path):
|
|
||||||
(tmp_path / "TOOLS.md").write_text("workspace tool notes", encoding="utf-8")
|
|
||||||
builder = _builder(tmp_path)
|
|
||||||
result = builder._load_bootstrap_files()
|
|
||||||
assert "TOOLS.md" not in result
|
|
||||||
assert "workspace tool notes" not in result
|
|
||||||
|
|
||||||
def test_utf8_content(self, tmp_path):
|
def test_utf8_content(self, tmp_path):
|
||||||
(tmp_path / "AGENTS.md").write_text("用中文回复", encoding="utf-8")
|
(tmp_path / "AGENTS.md").write_text("用中文回复", encoding="utf-8")
|
||||||
builder = _builder(tmp_path)
|
builder = _builder(tmp_path)
|
||||||
@@ -178,37 +171,6 @@ class TestIsTemplateContent:
|
|||||||
assert ContextBuilder._is_template_content("totally different", "memory/MEMORY.md") is False
|
assert ContextBuilder._is_template_content("totally different", "memory/MEMORY.md") is False
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Bundled bootstrap templates
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
class TestBundledToolContract:
|
|
||||||
def test_tool_contract_balances_general_and_coding_workflows(self):
|
|
||||||
from importlib.resources import files as pkg_files
|
|
||||||
|
|
||||||
tpl = pkg_files("nanobot") / "templates" / "agent" / "tool_contract.md"
|
|
||||||
content = tpl.read_text(encoding="utf-8")
|
|
||||||
|
|
||||||
assert "## General Tool Contract" in content
|
|
||||||
assert "Use the narrowest structured tool" in content
|
|
||||||
assert "Do not use `exec` as a universal workaround" in content
|
|
||||||
assert "## File and Coding Workflows" in content
|
|
||||||
assert "apply_patch" in content
|
|
||||||
assert "## Web and External Information" in content
|
|
||||||
assert "## Messaging and Media" in content
|
|
||||||
assert "## Scheduling and Background Work" in content
|
|
||||||
assert "pure coding" not in content.lower()
|
|
||||||
|
|
||||||
def test_tool_contract_is_injected_without_workspace_file(self, tmp_path):
|
|
||||||
builder = _builder(tmp_path)
|
|
||||||
prompt = builder.build_system_prompt()
|
|
||||||
|
|
||||||
assert "# Tool Usage Notes" in prompt
|
|
||||||
assert "## General Tool Contract" in prompt
|
|
||||||
assert "Do not use `exec` as a universal workaround" in prompt
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# _build_user_content
|
# _build_user_content
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -362,21 +324,6 @@ class TestBuildMessages:
|
|||||||
assert "Other chat goal." not in str(without_goal[-1]["content"])
|
assert "Other chat goal." not in str(without_goal[-1]["content"])
|
||||||
assert "Goal (active):" not in str(without_goal[-1]["content"])
|
assert "Goal (active):" not in str(without_goal[-1]["content"])
|
||||||
|
|
||||||
def test_current_runtime_lines_are_injected(self, tmp_path):
|
|
||||||
builder = _builder(tmp_path)
|
|
||||||
messages = builder.build_messages(
|
|
||||||
[],
|
|
||||||
"please use @zoom tonight",
|
|
||||||
current_runtime_lines=[
|
|
||||||
"CLI App Attachment: @zoom (installed; tool=run_cli_app; entry_point=cli-anything-zoom).",
|
|
||||||
],
|
|
||||||
)
|
|
||||||
user_msg = str(messages[-1]["content"])
|
|
||||||
|
|
||||||
assert "CLI App Attachment: @zoom" in user_msg
|
|
||||||
assert "tool=run_cli_app" in user_msg
|
|
||||||
assert "entry_point=cli-anything-zoom" in user_msg
|
|
||||||
|
|
||||||
def test_consecutive_same_role_merged(self, tmp_path):
|
def test_consecutive_same_role_merged(self, tmp_path):
|
||||||
builder = _builder(tmp_path)
|
builder = _builder(tmp_path)
|
||||||
history = [{"role": "user", "content": "previous user message"}]
|
history = [{"role": "user", "content": "previous user message"}]
|
||||||
|
|||||||
@@ -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 "When 'generate_image' creates images" in prompt
|
assert "the runtime attaches those artifacts to the final assistant reply automatically" in prompt
|
||||||
assert "call 'message' with the artifact paths in the 'media' parameter" in prompt
|
assert "do not call 'message' just to announce or resend them" in prompt
|
||||||
assert "Wait for the tool results, then answer once" in prompt
|
assert "Wait for the tool results, then answer once" in prompt
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -29,15 +29,14 @@ class FakeImageClient:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_outbound_no_longer_carries_generated_media(
|
async def test_generated_image_media_is_attached_to_final_assistant_message(
|
||||||
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.get_image_gen_provider",
|
"nanobot.agent.tools.image_generation.OpenRouterImageGenerationClient",
|
||||||
lambda name: FakeImageClient if name == "openrouter" else None,
|
FakeImageClient,
|
||||||
)
|
)
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
provider.get_default_model.return_value = "test-model"
|
provider.get_default_model.return_value = "test-model"
|
||||||
@@ -82,6 +81,9 @@ async def test_outbound_no_longer_carries_generated_media(
|
|||||||
|
|
||||||
assert result is not None
|
assert result is not None
|
||||||
assert result.content == "Done"
|
assert result.content == "Done"
|
||||||
# OutboundMessage no longer carries generated media —
|
assert len(result.media) == 1
|
||||||
# the LLM sends images via the message tool instead.
|
assert Path(result.media[0]).is_file()
|
||||||
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
|
||||||
|
|||||||
@@ -133,7 +133,6 @@ class TestToolEventProgress:
|
|||||||
"call_id": "call-write",
|
"call_id": "call-write",
|
||||||
"tool": "write_file",
|
"tool": "write_file",
|
||||||
"path": "foo.txt",
|
"path": "foo.txt",
|
||||||
"absolute_path": (tmp_path / "foo.txt").resolve().as_posix(),
|
|
||||||
"phase": "start",
|
"phase": "start",
|
||||||
"added": 2,
|
"added": 2,
|
||||||
"deleted": 1,
|
"deleted": 1,
|
||||||
@@ -310,100 +309,6 @@ class TestToolEventProgress:
|
|||||||
await invoke_file_edit_progress(telegram_progress, edit_events)
|
await invoke_file_edit_progress(telegram_progress, edit_events)
|
||||||
assert bus.outbound_size == 0
|
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,
|
||||||
@@ -651,7 +556,7 @@ class TestToolEventProgress:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.session.webui_turns.maybe_generate_webui_title_after_turn",
|
"nanobot.utils.webui_turn_helpers.maybe_generate_webui_title_after_turn",
|
||||||
fake_title_after_turn,
|
fake_title_after_turn,
|
||||||
)
|
)
|
||||||
scheduled_title: list[object] = []
|
scheduled_title: list[object] = []
|
||||||
@@ -698,7 +603,7 @@ class TestToolEventProgress:
|
|||||||
raise AssertionError("command-only turns should not generate titles")
|
raise AssertionError("command-only turns should not generate titles")
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.session.webui_turns.maybe_generate_webui_title_after_turn",
|
"nanobot.utils.webui_turn_helpers.maybe_generate_webui_title_after_turn",
|
||||||
fake_title_after_turn,
|
fake_title_after_turn,
|
||||||
)
|
)
|
||||||
scheduled: list[object] = []
|
scheduled: list[object] = []
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ 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, SessionManager
|
from nanobot.session.manager import Session, SessionManager
|
||||||
from nanobot.session.webui_turns import (
|
from nanobot.utils.webui_turn_helpers import (
|
||||||
TITLE_GENERATION_MAX_TOKENS,
|
TITLE_GENERATION_MAX_TOKENS,
|
||||||
TITLE_GENERATION_REASONING_EFFORT,
|
TITLE_GENERATION_REASONING_EFFORT,
|
||||||
WEBUI_SESSION_METADATA_KEY,
|
WEBUI_SESSION_METADATA_KEY,
|
||||||
@@ -143,7 +143,7 @@ def test_webui_title_update_uses_captured_llm_runtime(
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.session.webui_turns.maybe_generate_webui_title_after_turn",
|
"nanobot.utils.webui_turn_helpers.maybe_generate_webui_title_after_turn",
|
||||||
fake_title_after_turn,
|
fake_title_after_turn,
|
||||||
)
|
)
|
||||||
coordinator = WebuiTurnCoordinator(
|
coordinator = WebuiTurnCoordinator(
|
||||||
|
|||||||
@@ -346,26 +346,6 @@ class TestSyncWorkspaceTemplates:
|
|||||||
content = (workspace / "AGENTS.md").read_text()
|
content = (workspace / "AGENTS.md").read_text()
|
||||||
assert content == "existing content"
|
assert content == "existing content"
|
||||||
|
|
||||||
def test_does_not_create_tools_md(self, tmp_path):
|
|
||||||
"""Tool contract is injected internally, not copied into user workspaces."""
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
|
|
||||||
added = sync_workspace_templates(workspace, silent=True)
|
|
||||||
|
|
||||||
assert "TOOLS.md" not in added
|
|
||||||
assert not (workspace / "TOOLS.md").exists()
|
|
||||||
|
|
||||||
def test_preserves_existing_tools_md_without_overwriting(self, tmp_path):
|
|
||||||
"""Legacy user workspaces may have TOOLS.md; sync should leave it untouched."""
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir(parents=True)
|
|
||||||
tools_path = workspace / "TOOLS.md"
|
|
||||||
tools_path.write_text("custom tool notes", encoding="utf-8")
|
|
||||||
|
|
||||||
sync_workspace_templates(workspace, silent=True)
|
|
||||||
|
|
||||||
assert tools_path.read_text(encoding="utf-8") == "custom tool notes"
|
|
||||||
|
|
||||||
def test_creates_memory_directory(self, tmp_path):
|
def test_creates_memory_directory(self, tmp_path):
|
||||||
"""Should create memory directory structure."""
|
"""Should create memory directory structure."""
|
||||||
workspace = tmp_path / "workspace"
|
workspace = tmp_path / "workspace"
|
||||||
|
|||||||
@@ -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, ToolCallRequest
|
from nanobot.providers.base import LLMResponse
|
||||||
|
|
||||||
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
|
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
|
||||||
|
|
||||||
@@ -77,220 +77,3 @@ 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()
|
|
||||||
|
|||||||
@@ -359,31 +359,6 @@ def test_get_history_synthesizes_breadcrumb_for_image_only_turn():
|
|||||||
assert history[0] == {"role": "user", "content": "[image: /m/pic.png]"}
|
assert history[0] == {"role": "user", "content": "[image: /m/pic.png]"}
|
||||||
|
|
||||||
|
|
||||||
def test_get_history_synthesizes_cli_app_attachment_breadcrumb():
|
|
||||||
session = Session(key="test:cli-app")
|
|
||||||
session.messages.append(
|
|
||||||
{
|
|
||||||
"role": "user",
|
|
||||||
"content": "please use @drawio",
|
|
||||||
"cli_apps": [{
|
|
||||||
"name": "drawio",
|
|
||||||
"entry_point": "cli-anything-drawio",
|
|
||||||
}],
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
history = session.get_history(max_messages=500)
|
|
||||||
|
|
||||||
assert history == [{
|
|
||||||
"role": "user",
|
|
||||||
"content": (
|
|
||||||
"please use @drawio\n"
|
|
||||||
"[CLI App Attachment: @drawio; tool=run_cli_app; "
|
|
||||||
"entry_point=cli-anything-drawio; skill=skills/cli-app-drawio/SKILL.md]"
|
|
||||||
),
|
|
||||||
}]
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_history_ignores_media_kwarg_on_non_user_rows():
|
def test_get_history_ignores_media_kwarg_on_non_user_rows():
|
||||||
"""``media`` only ever appears on user entries in practice, but the
|
"""``media`` only ever appears on user entries in practice, but the
|
||||||
synthesizer must be defensive: assistants / tools with list content
|
synthesizer must be defensive: assistants / tools with list content
|
||||||
|
|||||||
@@ -0,0 +1,34 @@
|
|||||||
|
"""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())]
|
||||||
@@ -111,23 +111,6 @@ 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
|
||||||
|
|
||||||
@@ -169,25 +152,6 @@ 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
|
||||||
|
|
||||||
@@ -216,7 +180,7 @@ async def test_manager_loads_plugin_from_dict_config():
|
|||||||
)
|
)
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
"nanobot.channels.registry.discover_enabled",
|
"nanobot.channels.registry.discover_all",
|
||||||
return_value={"fakeplugin": _FakePlugin},
|
return_value={"fakeplugin": _FakePlugin},
|
||||||
):
|
):
|
||||||
mgr = ChannelManager.__new__(ChannelManager)
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
@@ -246,7 +210,7 @@ async def test_manager_propagates_groq_transcription_api_base_to_channels():
|
|||||||
)
|
)
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
"nanobot.channels.registry.discover_enabled",
|
"nanobot.channels.registry.discover_all",
|
||||||
return_value={"fakeplugin": _FakePlugin},
|
return_value={"fakeplugin": _FakePlugin},
|
||||||
):
|
):
|
||||||
mgr = ChannelManager.__new__(ChannelManager)
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
@@ -282,7 +246,7 @@ async def test_manager_propagates_openai_transcription_api_base_to_channels():
|
|||||||
)
|
)
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
"nanobot.channels.registry.discover_enabled",
|
"nanobot.channels.registry.discover_all",
|
||||||
return_value={"fakeplugin": _FakePlugin},
|
return_value={"fakeplugin": _FakePlugin},
|
||||||
):
|
):
|
||||||
mgr = ChannelManager.__new__(ChannelManager)
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
@@ -534,8 +498,10 @@ async def test_manager_skips_disabled_plugin():
|
|||||||
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
)
|
)
|
||||||
|
|
||||||
ep = _make_entry_point("fakeplugin", _FakePlugin)
|
with patch(
|
||||||
with patch(_EP_TARGET, return_value=[ep]):
|
"nanobot.channels.registry.discover_all",
|
||||||
|
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,8 +29,7 @@ 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, ModelPresetConfig
|
from nanobot.config.schema import Config
|
||||||
from nanobot.webui.settings_api import settings_payload
|
|
||||||
|
|
||||||
# -- Shared helpers (aligned with test_websocket_integration.py) ---------------
|
# -- Shared helpers (aligned with test_websocket_integration.py) ---------------
|
||||||
|
|
||||||
@@ -480,99 +479,6 @@ async def test_send_delta_emits_delta_and_stream_end() -> None:
|
|||||||
assert second["stream_id"] == "sid"
|
assert second["stream_id"] == "sid"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_send_delta_stream_end_rewrites_local_markdown_image(monkeypatch, tmp_path) -> None:
|
|
||||||
bus = MagicMock()
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
(workspace / "diagram.png").write_bytes(b"\x89PNG\r\n\x1a\nimage")
|
|
||||||
media = tmp_path / "media"
|
|
||||||
|
|
||||||
def fake_media_dir(channel: str | None = None):
|
|
||||||
path = media / channel if channel else media
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
return path
|
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.channels.websocket.get_media_dir", fake_media_dir)
|
|
||||||
channel = WebSocketChannel(
|
|
||||||
{"enabled": True, "allowFrom": ["*"], "streaming": True},
|
|
||||||
bus,
|
|
||||||
workspace_path=workspace,
|
|
||||||
)
|
|
||||||
mock_ws = AsyncMock()
|
|
||||||
channel._attach(mock_ws, "chat-1")
|
|
||||||
channel._webui_chats.add("chat-1")
|
|
||||||
|
|
||||||
await channel.send_delta("chat-1", "
|
|
||||||
await channel.send_delta("chat-1", "diagram.png)", {"_stream_delta": True, "_stream_id": "sid"})
|
|
||||||
await channel.send_delta("chat-1", "", {"_stream_end": True, "_stream_id": "sid"})
|
|
||||||
|
|
||||||
assert mock_ws.send.await_count == 3
|
|
||||||
final = json.loads(mock_ws.send.call_args_list[2][0][0])
|
|
||||||
assert final["event"] == "stream_end"
|
|
||||||
assert final["text"].startswith("
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_send_delta_stream_end_rewrites_inline_final_text(monkeypatch, tmp_path) -> None:
|
|
||||||
bus = MagicMock()
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
(workspace / "diagram.png").write_bytes(b"\x89PNG\r\n\x1a\nimage")
|
|
||||||
media = tmp_path / "media"
|
|
||||||
|
|
||||||
def fake_media_dir(channel: str | None = None):
|
|
||||||
path = media / channel if channel else media
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
return path
|
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.channels.websocket.get_media_dir", fake_media_dir)
|
|
||||||
channel = WebSocketChannel(
|
|
||||||
{"enabled": True, "allowFrom": ["*"], "streaming": True},
|
|
||||||
bus,
|
|
||||||
workspace_path=workspace,
|
|
||||||
)
|
|
||||||
mock_ws = AsyncMock()
|
|
||||||
channel._attach(mock_ws, "chat-1")
|
|
||||||
channel._webui_chats.add("chat-1")
|
|
||||||
|
|
||||||
await channel.send_delta(
|
|
||||||
"chat-1",
|
|
||||||
"",
|
|
||||||
{"_stream_delta": True, "_stream_end": True, "_stream_id": "sid"},
|
|
||||||
)
|
|
||||||
|
|
||||||
mock_ws.send.assert_awaited_once()
|
|
||||||
final = json.loads(mock_ws.send.await_args.args[0])
|
|
||||||
assert final["event"] == "stream_end"
|
|
||||||
assert final["text"].startswith("
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_send_delta_stream_end_leaves_non_webui_payload_unchanged(tmp_path) -> None:
|
|
||||||
bus = MagicMock()
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
(workspace / "diagram.png").write_bytes(b"\x89PNG\r\n\x1a\nimage")
|
|
||||||
channel = WebSocketChannel(
|
|
||||||
{"enabled": True, "allowFrom": ["*"], "streaming": True},
|
|
||||||
bus,
|
|
||||||
workspace_path=workspace,
|
|
||||||
)
|
|
||||||
mock_ws = AsyncMock()
|
|
||||||
channel._attach(mock_ws, "chat-1")
|
|
||||||
|
|
||||||
await channel.send_delta(
|
|
||||||
"chat-1",
|
|
||||||
"",
|
|
||||||
{"_stream_delta": True, "_stream_end": True, "_stream_id": "sid"},
|
|
||||||
)
|
|
||||||
|
|
||||||
mock_ws.send.assert_awaited_once()
|
|
||||||
final = json.loads(mock_ws.send.await_args.args[0])
|
|
||||||
assert final == {"event": "stream_end", "chat_id": "chat-1", "stream_id": "sid"}
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_send_reasoning_delta_emits_streaming_frame() -> None:
|
async def test_send_reasoning_delta_emits_streaming_frame() -> None:
|
||||||
bus = MagicMock()
|
bus = MagicMock()
|
||||||
@@ -850,7 +756,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.session import webui_turns as wth
|
from nanobot.utils import webui_turn_helpers 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")
|
||||||
@@ -863,7 +769,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.session import webui_turns as wth
|
from nanobot.utils import webui_turn_helpers as wth
|
||||||
|
|
||||||
wth._WEBSOCKET_TURN_WALL_STARTED_AT.clear()
|
wth._WEBSOCKET_TURN_WALL_STARTED_AT.clear()
|
||||||
try:
|
try:
|
||||||
@@ -1085,11 +991,6 @@ 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)
|
||||||
@@ -1110,53 +1011,21 @@ 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["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["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"]["configured"] is False
|
||||||
assert providers["atomic_chat"]["api_key_required"] is False
|
assert providers["atomic_chat"]["api_key_required"] is False
|
||||||
assert providers["atomic_chat"]["default_api_base"] == "http://localhost:1337/v1"
|
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["openai_codex"]["configured"] is True
|
|
||||||
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
|
||||||
|
|
||||||
@@ -1171,7 +1040,6 @@ 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(
|
local_provider_updated = await _http_get(
|
||||||
@@ -1191,117 +1059,34 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
|
|||||||
updated = await _http_get(
|
updated = await _http_get(
|
||||||
"http://127.0.0.1:"
|
"http://127.0.0.1:"
|
||||||
f"{port}/api/settings/update?model=atomic_chat/test"
|
f"{port}/api/settings/update?model=atomic_chat/test"
|
||||||
"&provider=atomic_chat&timezone=Asia%2FShanghai"
|
"&provider=atomic_chat",
|
||||||
"&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
|
||||||
updated_body = updated.json()
|
assert updated.json()["requires_restart"] is False
|
||||||
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 True
|
assert search_body["requires_restart"] is False
|
||||||
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 == "atomic_chat/test"
|
assert saved.agents.defaults.model == "atomic_chat/test"
|
||||||
assert saved.agents.defaults.provider == "atomic_chat"
|
assert saved.agents.defaults.provider == "atomic_chat"
|
||||||
assert saved.agents.defaults.model_preset == "deep"
|
assert saved.providers.openrouter.api_key == "sk-or-test"
|
||||||
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.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
|
||||||
@@ -1346,7 +1131,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 = settings_payload()
|
body = _ch(bus)._settings_payload()
|
||||||
|
|
||||||
assert body["agent"]["provider"] == "minimax_anthropic"
|
assert body["agent"]["provider"] == "minimax_anthropic"
|
||||||
|
|
||||||
@@ -1763,54 +1548,6 @@ 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"),
|
||||||
[
|
[
|
||||||
@@ -1837,7 +1574,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.webui.transcript import append_transcript_object
|
from nanobot.utils.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"
|
||||||
|
|||||||
@@ -105,43 +105,6 @@ async def test_message_without_media_backward_compatible() -> None:
|
|||||||
assert call.kwargs["media"] is None
|
assert call.kwargs["media"] is None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_message_forwards_normalized_cli_app_attachments() -> None:
|
|
||||||
channel = _make_channel()
|
|
||||||
mock_conn = AsyncMock()
|
|
||||||
envelope = {
|
|
||||||
"type": "message",
|
|
||||||
"chat_id": "abc123",
|
|
||||||
"content": "please use @drawio",
|
|
||||||
"webui": True,
|
|
||||||
"cli_apps": [
|
|
||||||
{
|
|
||||||
"name": "DrawIO",
|
|
||||||
"display_name": "Draw.io",
|
|
||||||
"category": "diagram",
|
|
||||||
"entry_point": "cli-anything-drawio",
|
|
||||||
"logo_url": "https://example.invalid/drawio.svg",
|
|
||||||
"brand_color": "#F08705",
|
|
||||||
},
|
|
||||||
{"name": "bad name", "entry_point": "nope"},
|
|
||||||
],
|
|
||||||
}
|
|
||||||
|
|
||||||
await channel._dispatch_envelope(mock_conn, "client-1", envelope)
|
|
||||||
|
|
||||||
channel._handle_message.assert_awaited_once()
|
|
||||||
metadata = channel._handle_message.call_args.kwargs["metadata"]
|
|
||||||
assert metadata["webui"] is True
|
|
||||||
assert metadata["cli_apps"] == [{
|
|
||||||
"name": "drawio",
|
|
||||||
"display_name": "Draw.io",
|
|
||||||
"category": "diagram",
|
|
||||||
"entry_point": "cli-anything-drawio",
|
|
||||||
"logo_url": "https://example.invalid/drawio.svg",
|
|
||||||
"brand_color": "#F08705",
|
|
||||||
}]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_message_with_single_image_forwards_saved_path(tmp_path) -> None:
|
async def test_message_with_single_image_forwards_saved_path(tmp_path) -> None:
|
||||||
channel = _make_channel()
|
channel = _make_channel()
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ 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
|
||||||
@@ -140,75 +139,6 @@ async def test_sessions_routes_require_bearer_token(
|
|||||||
await server_task
|
await server_task
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_cli_apps_routes_require_token_and_return_payload(
|
|
||||||
bus: MagicMock,
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"nanobot.channels.websocket.cli_apps_payload",
|
|
||||||
lambda: {
|
|
||||||
"apps": [
|
|
||||||
{
|
|
||||||
"name": "gimp",
|
|
||||||
"display_name": "GIMP",
|
|
||||||
"category": "image",
|
|
||||||
"description": "Image editing",
|
|
||||||
"requires": "Python",
|
|
||||||
"source": "harness",
|
|
||||||
"entry_point": "cli-anything-gimp",
|
|
||||||
"install_supported": True,
|
|
||||||
"installed": False,
|
|
||||||
"available": False,
|
|
||||||
"status": "not_installed",
|
|
||||||
"logo_url": None,
|
|
||||||
"brand_color": None,
|
|
||||||
"skill_installed": False,
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"installed_count": 0,
|
|
||||||
"catalog_updated_at": "2026-04-18",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"nanobot.channels.websocket.cli_apps_action",
|
|
||||||
lambda action, query: {
|
|
||||||
"apps": [],
|
|
||||||
"installed_count": 1,
|
|
||||||
"catalog_updated_at": "2026-04-18",
|
|
||||||
"last_action": {"ok": True, "message": f"{action}:{query['name'][0]}"},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
channel = _ch(bus, session_manager=_seed_session(tmp_path), port=29912)
|
|
||||||
server_task = asyncio.create_task(channel.start())
|
|
||||||
await asyncio.sleep(0.3)
|
|
||||||
try:
|
|
||||||
deny = await _http_get("http://127.0.0.1:29912/api/settings/cli-apps")
|
|
||||||
assert deny.status_code == 401
|
|
||||||
|
|
||||||
boot = await _http_get("http://127.0.0.1:29912/webui/bootstrap")
|
|
||||||
token = boot.json()["token"]
|
|
||||||
auth = {"Authorization": f"Bearer {token}"}
|
|
||||||
|
|
||||||
catalog = await _http_get(
|
|
||||||
"http://127.0.0.1:29912/api/settings/cli-apps",
|
|
||||||
headers=auth,
|
|
||||||
)
|
|
||||||
assert catalog.status_code == 200
|
|
||||||
assert catalog.json()["apps"][0]["name"] == "gimp"
|
|
||||||
|
|
||||||
installed = await _http_get(
|
|
||||||
"http://127.0.0.1:29912/api/settings/cli-apps/install?name=gimp",
|
|
||||||
headers=auth,
|
|
||||||
)
|
|
||||||
assert installed.status_code == 200
|
|
||||||
assert installed.json()["last_action"]["message"] == "install:gimp"
|
|
||||||
finally:
|
|
||||||
await channel.stop()
|
|
||||||
await server_task
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_sessions_list_only_returns_websocket_sessions_by_default(
|
async def test_sessions_list_only_returns_websocket_sessions_by_default(
|
||||||
bus: MagicMock, tmp_path: Path
|
bus: MagicMock, tmp_path: Path
|
||||||
@@ -246,62 +176,13 @@ 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.webui.transcript import append_transcript_object
|
from nanobot.utils.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)
|
||||||
|
|||||||
@@ -44,7 +44,6 @@ def _ch(
|
|||||||
bus: Any,
|
bus: Any,
|
||||||
*,
|
*,
|
||||||
session_manager: SessionManager | None = None,
|
session_manager: SessionManager | None = None,
|
||||||
workspace_path: Path | None = None,
|
|
||||||
port: int,
|
port: int,
|
||||||
) -> WebSocketChannel:
|
) -> WebSocketChannel:
|
||||||
return WebSocketChannel(
|
return WebSocketChannel(
|
||||||
@@ -58,7 +57,6 @@ def _ch(
|
|||||||
},
|
},
|
||||||
bus,
|
bus,
|
||||||
session_manager=session_manager,
|
session_manager=session_manager,
|
||||||
workspace_path=workspace_path,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -69,15 +67,6 @@ def bus() -> MagicMock:
|
|||||||
return b
|
return b
|
||||||
|
|
||||||
|
|
||||||
def _fake_media_dir(root: Path):
|
|
||||||
def inner(channel: str | None = None) -> Path:
|
|
||||||
path = root / channel if channel else root
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
return path
|
|
||||||
|
|
||||||
return inner
|
|
||||||
|
|
||||||
|
|
||||||
async def _http_get(
|
async def _http_get(
|
||||||
url: str, headers: dict[str, str] | None = None
|
url: str, headers: dict[str, str] | None = None
|
||||||
) -> httpx.Response:
|
) -> httpx.Response:
|
||||||
@@ -134,45 +123,6 @@ def test_sign_media_path_round_trips_via_hmac(
|
|||||||
assert _b64url_decode(payload).decode() == "a.png"
|
assert _b64url_decode(payload).decode() == "a.png"
|
||||||
|
|
||||||
|
|
||||||
def test_local_markdown_image_is_staged_and_rewritten(
|
|
||||||
bus: MagicMock,
|
|
||||||
tmp_path: Path,
|
|
||||||
) -> None:
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
(workspace / "demo_arch.png").write_bytes(_PNG_BYTES)
|
|
||||||
media = tmp_path / "media"
|
|
||||||
channel = _ch(bus, workspace_path=workspace, port=0)
|
|
||||||
|
|
||||||
with patch("nanobot.channels.websocket.get_media_dir", side_effect=_fake_media_dir(media)):
|
|
||||||
rewritten = channel._rewrite_local_markdown_images(
|
|
||||||
"The result:\n"
|
|
||||||
)
|
|
||||||
|
|
||||||
assert ".iterdir())
|
|
||||||
assert len(staged) == 1
|
|
||||||
assert staged[0].read_bytes() == _PNG_BYTES
|
|
||||||
|
|
||||||
|
|
||||||
def test_local_markdown_image_rejects_workspace_escape(
|
|
||||||
bus: MagicMock,
|
|
||||||
tmp_path: Path,
|
|
||||||
) -> None:
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
outside = tmp_path / "outside.png"
|
|
||||||
outside.write_bytes(_PNG_BYTES)
|
|
||||||
media = tmp_path / "media"
|
|
||||||
channel = _ch(bus, workspace_path=workspace, port=0)
|
|
||||||
text = ""
|
|
||||||
|
|
||||||
with patch("nanobot.channels.websocket.get_media_dir", side_effect=_fake_media_dir(media)):
|
|
||||||
assert channel._rewrite_local_markdown_images(text) == text
|
|
||||||
|
|
||||||
assert not (media / "websocket").exists()
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# /api/media/<sig>/<payload>: the serving handler
|
# /api/media/<sig>/<payload>: the serving handler
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import tempfile
|
import tempfile
|
||||||
import time
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock
|
from unittest.mock import AsyncMock
|
||||||
@@ -375,7 +374,6 @@ async def test_send_uses_typing_start_and_cancel_when_ticket_available() -> None
|
|||||||
channel._client = object()
|
channel._client = object()
|
||||||
channel._token = "token"
|
channel._token = "token"
|
||||||
channel._context_tokens["wx-user"] = "ctx-typing"
|
channel._context_tokens["wx-user"] = "ctx-typing"
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock()
|
channel._send_text = AsyncMock()
|
||||||
channel._api_post = AsyncMock(
|
channel._api_post = AsyncMock(
|
||||||
side_effect=[
|
side_effect=[
|
||||||
@@ -404,7 +402,6 @@ async def test_send_still_sends_text_when_typing_ticket_missing() -> None:
|
|||||||
channel._client = object()
|
channel._client = object()
|
||||||
channel._token = "token"
|
channel._token = "token"
|
||||||
channel._context_tokens["wx-user"] = "ctx-no-ticket"
|
channel._context_tokens["wx-user"] = "ctx-no-ticket"
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock()
|
channel._send_text = AsyncMock()
|
||||||
channel._api_post = AsyncMock(return_value={"ret": 1, "errmsg": "no config"})
|
channel._api_post = AsyncMock(return_value={"ret": 1, "errmsg": "no config"})
|
||||||
|
|
||||||
@@ -1257,526 +1254,3 @@ async def test_send_text_succeeds_on_zero_errcode() -> None:
|
|||||||
await channel._send_text("wx-user", "hello", "ctx-ok")
|
await channel._send_text("wx-user", "hello", "ctx-ok")
|
||||||
|
|
||||||
channel._api_post.assert_awaited_once()
|
channel._api_post.assert_awaited_once()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_send_text_raises_on_nonzero_ret_even_when_errcode_zero() -> None:
|
|
||||||
"""_send_text must raise when the API returns ret != 0, even if errcode is 0.
|
|
||||||
|
|
||||||
The iLink API signals failure through either field. Checking only errcode
|
|
||||||
caused silent message drops (responses generated but never delivered).
|
|
||||||
"""
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel._api_post = AsyncMock(
|
|
||||||
return_value={"ret": -100, "errcode": 0, "errmsg": "internal error"}
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(RuntimeError, match="WeChat send text error.*ret=-100.*errcode=0"):
|
|
||||||
await channel._send_text("wx-user", "hello", "ctx-ok")
|
|
||||||
|
|
||||||
channel._api_post.assert_awaited_once()
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Tests for _poll_once not silently dropping messages on processing errors
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_poll_once_logs_exception_on_process_message_failure(monkeypatch) -> None:
|
|
||||||
"""When _process_message raises, _poll_once must log the error and continue
|
|
||||||
processing remaining messages instead of silently swallowing the exception."""
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = SimpleNamespace(timeout=None)
|
|
||||||
channel._token = "token"
|
|
||||||
channel._get_updates_buf = "old-buf"
|
|
||||||
|
|
||||||
calls = []
|
|
||||||
logged_messages: list[str] = []
|
|
||||||
|
|
||||||
async def _failing_process(msg: dict) -> None:
|
|
||||||
calls.append(msg.get("message_id"))
|
|
||||||
if msg.get("message_id") == "msg-1":
|
|
||||||
raise RuntimeError("processing failed")
|
|
||||||
|
|
||||||
channel._process_message = _failing_process # type: ignore[method-assign]
|
|
||||||
|
|
||||||
monkeypatch.setattr(
|
|
||||||
channel.logger,
|
|
||||||
"exception",
|
|
||||||
lambda message, *args, **kwargs: logged_messages.append(str(message)),
|
|
||||||
)
|
|
||||||
|
|
||||||
channel._api_post = AsyncMock( # type: ignore[method-assign]
|
|
||||||
return_value={
|
|
||||||
"ret": 0,
|
|
||||||
"errcode": 0,
|
|
||||||
"get_updates_buf": "new-buf",
|
|
||||||
"msgs": [
|
|
||||||
{"message_id": "msg-1", "message_type": 1},
|
|
||||||
{"message_id": "msg-2", "message_type": 1},
|
|
||||||
],
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
await channel._poll_once()
|
|
||||||
|
|
||||||
# Both messages should have been attempted
|
|
||||||
assert calls == ["msg-1", "msg-2"]
|
|
||||||
# Buffer should still advance (already updated before processing)
|
|
||||||
assert channel._get_updates_buf == "new-buf"
|
|
||||||
# Error should be logged
|
|
||||||
assert any("Failed to process WeChat message" in m for m in logged_messages)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_poll_loop_logs_exception_and_continues_on_poll_failure(monkeypatch) -> None:
|
|
||||||
"""When _poll_once raises a non-timeout exception, the start() loop must log
|
|
||||||
the error and continue polling instead of exiting silently."""
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.config.token = "token" # skip QR login in start()
|
|
||||||
channel._running = True
|
|
||||||
|
|
||||||
call_count = 0
|
|
||||||
logged_messages: list[str] = []
|
|
||||||
|
|
||||||
async def _failing_poll() -> None:
|
|
||||||
nonlocal call_count
|
|
||||||
call_count += 1
|
|
||||||
if call_count == 1:
|
|
||||||
raise RuntimeError("poll exploded")
|
|
||||||
channel._running = False # Stop after second call
|
|
||||||
|
|
||||||
channel._poll_once = _failing_poll # type: ignore[method-assign]
|
|
||||||
|
|
||||||
monkeypatch.setattr(
|
|
||||||
channel.logger,
|
|
||||||
"exception",
|
|
||||||
lambda message, *args, **kwargs: logged_messages.append(str(message)),
|
|
||||||
)
|
|
||||||
|
|
||||||
# Use a tiny retry delay so the test finishes quickly
|
|
||||||
original_retry = weixin_mod.RETRY_DELAY_S
|
|
||||||
weixin_mod.RETRY_DELAY_S = 0.01
|
|
||||||
try:
|
|
||||||
await channel.start()
|
|
||||||
finally:
|
|
||||||
weixin_mod.RETRY_DELAY_S = original_retry
|
|
||||||
|
|
||||||
assert call_count == 2
|
|
||||||
assert any("WeChat poll loop error" in m for m in logged_messages)
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Tool-hint buffering
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_buffer_single_tool_hint_not_sent_immediately() -> None:
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = True
|
|
||||||
channel._context_tokens["wx-user"] = "ctx-1"
|
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock()
|
|
||||||
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Using tool",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
channel._send_text.assert_not_awaited()
|
|
||||||
assert channel._pending_tool_hints["wx-user"] == ["Using tool"]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_buffer_multiple_tool_hints_flushed_on_final_answer() -> None:
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = True
|
|
||||||
channel._context_tokens["wx-user"] = "ctx-1"
|
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock()
|
|
||||||
|
|
||||||
for hint in ["tool1", "tool2"]:
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": hint,
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Done",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
assert channel._send_text.await_count == 2
|
|
||||||
channel._send_text.assert_any_await("wx-user", "tool1\n\ntool2", "ctx-1")
|
|
||||||
channel._send_text.assert_any_await("wx-user", "Done", "ctx-1")
|
|
||||||
assert "wx-user" not in channel._pending_tool_hints
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_thought_progress_flushes_tool_hints() -> None:
|
|
||||||
"""Thoughts are visible progress messages and must act as separators,
|
|
||||||
flushing buffered tool hints before they are sent."""
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = True
|
|
||||||
channel._context_tokens["wx-user"] = "ctx-1"
|
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock()
|
|
||||||
|
|
||||||
# Buffer a tool hint
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "search 'foo'",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
# Send a thought — progress but not a tool_hint.
|
|
||||||
# It must act as a separator and flush the buffered hint.
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Let me think...",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
# The buffered hint was flushed before the thought was sent.
|
|
||||||
channel._send_text.assert_any_await("wx-user", "search 'foo'", "ctx-1")
|
|
||||||
channel._send_text.assert_any_await("wx-user", "Let me think...", "ctx-1")
|
|
||||||
assert "wx-user" not in channel._pending_tool_hints
|
|
||||||
|
|
||||||
# Final answer arrives with nothing left to flush.
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Done",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
assert channel._send_text.await_count == 3
|
|
||||||
channel._send_text.assert_any_await("wx-user", "Done", "ctx-1")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_reasoning_delta_does_not_flush_tool_hints() -> None:
|
|
||||||
"""Reasoning deltas are invisible in WeChat and must NOT flush buffered
|
|
||||||
tool hints — otherwise hints separated only by hidden reasoning would
|
|
||||||
fail to coalesce."""
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = True
|
|
||||||
channel._context_tokens["wx-user"] = "ctx-1"
|
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock()
|
|
||||||
|
|
||||||
# Buffer a tool hint
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "search 'foo'",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
# Send a reasoning delta — invisible in WeChat, must NOT flush
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Thinking step 1...",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_reasoning_delta": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
# Reasoning is invisible; hint stays buffered, _send_text not called
|
|
||||||
channel._send_text.assert_not_awaited()
|
|
||||||
assert channel._pending_tool_hints["wx-user"] == ["search 'foo'"]
|
|
||||||
|
|
||||||
# Final answer flushes the buffered hint
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Done",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
channel._send_text.assert_any_await("wx-user", "search 'foo'", "ctx-1")
|
|
||||||
channel._send_text.assert_any_await("wx-user", "Done", "ctx-1")
|
|
||||||
assert "wx-user" not in channel._pending_tool_hints
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_empty_progress_message_does_not_flush_tool_hints() -> None:
|
|
||||||
"""Empty progress messages (e.g. after_iteration tool_events) have no
|
|
||||||
visible content and must NOT act as separators."""
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = True
|
|
||||||
channel._context_tokens["wx-user"] = "ctx-1"
|
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock()
|
|
||||||
|
|
||||||
# Buffer a tool hint
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "search 'foo'",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
# Send an empty progress message (no content, no media)
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_events": [{"phase": "end"}]},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
# Nothing should have been sent yet
|
|
||||||
channel._send_text.assert_not_awaited()
|
|
||||||
assert channel._pending_tool_hints["wx-user"] == ["search 'foo'"]
|
|
||||||
|
|
||||||
# Final answer flushes the buffered hint
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Done",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
channel._send_text.assert_any_await("wx-user", "search 'foo'", "ctx-1")
|
|
||||||
channel._send_text.assert_any_await("wx-user", "Done", "ctx-1")
|
|
||||||
assert "wx-user" not in channel._pending_tool_hints
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_buffer_flush_refreshes_context_token() -> None:
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = True
|
|
||||||
channel._context_tokens["wx-user"] = "ctx-old"
|
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._refresh_context_token_if_stale = AsyncMock(return_value="ctx-refreshed")
|
|
||||||
channel._send_text = AsyncMock()
|
|
||||||
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "hint",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Done",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
assert channel._refresh_context_token_if_stale.await_count == 2
|
|
||||||
channel._refresh_context_token_if_stale.assert_any_await("wx-user", "ctx-old")
|
|
||||||
channel._send_text.assert_any_await("wx-user", "hint", "ctx-refreshed")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_buffer_flush_failure_does_not_block_final_answer() -> None:
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = True
|
|
||||||
channel._context_tokens["wx-user"] = "ctx-1"
|
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock(side_effect=[RuntimeError("boom"), None])
|
|
||||||
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "hint",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "Done",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
assert channel._send_text.await_count == 2
|
|
||||||
channel._send_text.assert_any_await("wx-user", "hint", "ctx-1")
|
|
||||||
channel._send_text.assert_any_await("wx-user", "Done", "ctx-1")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_buffer_flushed_on_stream_end() -> None:
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = True
|
|
||||||
channel._context_tokens["wx-user"] = "ctx-1"
|
|
||||||
channel._context_token_at["wx-user"] = time.time()
|
|
||||||
channel._send_text = AsyncMock()
|
|
||||||
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "hint",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
await channel.send_delta("wx-user", "", {"_stream_end": True})
|
|
||||||
|
|
||||||
channel._send_text.assert_awaited_once_with("wx-user", "hint", "ctx-1")
|
|
||||||
assert "wx-user" not in channel._pending_tool_hints
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_stop_clears_buffer() -> None:
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._pending_tool_hints["wx-user"] = ["hint1", "hint2"]
|
|
||||||
await channel.stop()
|
|
||||||
assert "wx-user" not in channel._pending_tool_hints
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_send_tool_hints_false_drops_tool_hints() -> None:
|
|
||||||
channel, _bus = _make_channel()
|
|
||||||
channel._client = object()
|
|
||||||
channel._token = "token"
|
|
||||||
channel.send_tool_hints = False
|
|
||||||
channel._send_text = AsyncMock()
|
|
||||||
|
|
||||||
await channel.send(
|
|
||||||
type(
|
|
||||||
"Msg",
|
|
||||||
(),
|
|
||||||
{
|
|
||||||
"chat_id": "wx-user",
|
|
||||||
"content": "hint",
|
|
||||||
"media": [],
|
|
||||||
"metadata": {"_progress": True, "_tool_hint": True},
|
|
||||||
},
|
|
||||||
)()
|
|
||||||
)
|
|
||||||
|
|
||||||
channel._send_text.assert_not_awaited()
|
|
||||||
assert "wx-user" not in channel._pending_tool_hints
|
|
||||||
|
|||||||
@@ -572,7 +572,6 @@ 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")
|
||||||
|
|
||||||
@@ -612,8 +611,7 @@ 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:
|
||||||
provider = make_provider(config)
|
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"
|
||||||
|
|||||||
@@ -1,408 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
|
||||||
import subprocess
|
|
||||||
import sys
|
|
||||||
import time
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
from nanobot.cli_apps.service import CliAppError, CliAppManager, CliAppsRuntimeConfig
|
|
||||||
|
|
||||||
|
|
||||||
def _write_cache(path: Path, registry: dict) -> None:
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
path.write_text(
|
|
||||||
json.dumps({"_cached_at": time.time(), "data": registry}),
|
|
||||||
encoding="utf-8",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _manager(tmp_path: Path) -> CliAppManager:
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
return CliAppManager(
|
|
||||||
workspace=workspace,
|
|
||||||
data_dir=tmp_path / "data",
|
|
||||||
runtime=CliAppsRuntimeConfig(catalog_ttl_seconds=3600, install_timeout=5, run_timeout=5),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _seed_catalog(manager: CliAppManager) -> None:
|
|
||||||
harness = {
|
|
||||||
"meta": {"updated": "2026-04-16"},
|
|
||||||
"clis": [
|
|
||||||
{
|
|
||||||
"name": "gimp",
|
|
||||||
"display_name": "GIMP",
|
|
||||||
"version": "1.0.0",
|
|
||||||
"description": "Image editing",
|
|
||||||
"category": "image",
|
|
||||||
"requires": "Python 3.10+",
|
|
||||||
"install_cmd": "pip install cli-anything-gimp",
|
|
||||||
"entry_point": "cli-anything-gimp",
|
|
||||||
"skill_md": "skills/cli-anything-gimp/SKILL.md",
|
|
||||||
}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
public = {
|
|
||||||
"meta": {"updated": "2026-04-18"},
|
|
||||||
"clis": [
|
|
||||||
{
|
|
||||||
"name": "gimp",
|
|
||||||
"display_name": "GIMP",
|
|
||||||
"description": "Public duplicate entry",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "jimeng",
|
|
||||||
"display_name": "Jimeng",
|
|
||||||
"version": "latest",
|
|
||||||
"description": "Script install",
|
|
||||||
"category": "ai",
|
|
||||||
"install_strategy": "script",
|
|
||||||
"install_cmd": "curl -fsSL https://example.invalid/install.sh | bash",
|
|
||||||
"entry_point": "dreamina",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "feishu",
|
|
||||||
"display_name": "Feishu/Lark CLI",
|
|
||||||
"version": "latest",
|
|
||||||
"description": "Official Lark CLI",
|
|
||||||
"category": "communication",
|
|
||||||
"package_manager": "npm",
|
|
||||||
"npm_package": "@larksuite/cli",
|
|
||||||
"install_cmd": "npm install -g @larksuite/cli",
|
|
||||||
"entry_point": "lark-cli",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "dify-workflow",
|
|
||||||
"display_name": "Dify Workflow",
|
|
||||||
"version": "latest",
|
|
||||||
"description": "Run Dify workflows",
|
|
||||||
"category": "ai",
|
|
||||||
"install_cmd": "pip install cli-anything-dify-workflow",
|
|
||||||
"entry_point": "cli-anything-dify-workflow",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "shopify",
|
|
||||||
"display_name": "Shopify CLI",
|
|
||||||
"version": "latest",
|
|
||||||
"description": "Shopify",
|
|
||||||
"category": "web",
|
|
||||||
"package_manager": "npm",
|
|
||||||
"npm_package": "@shopify/cli",
|
|
||||||
"install_cmd": "npm install -g @shopify/cli",
|
|
||||||
"entry_point": "shopify",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "clibrowser",
|
|
||||||
"display_name": "clibrowser",
|
|
||||||
"version": "latest",
|
|
||||||
"description": "Cargo install",
|
|
||||||
"category": "web",
|
|
||||||
"install_cmd": "cargo install --git https://example.invalid/clibrowser.git",
|
|
||||||
"entry_point": "clibrowser",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "suno",
|
|
||||||
"display_name": "Suno CLI",
|
|
||||||
"version": "latest",
|
|
||||||
"description": "python3 pip install",
|
|
||||||
"category": "music",
|
|
||||||
"package_manager": "pip",
|
|
||||||
"install_strategy": "command",
|
|
||||||
"install_cmd": "python3 -m pip install git+https://example.invalid/suno-cli.git",
|
|
||||||
"uninstall_cmd": "python3 -m pip uninstall -y suno-cli",
|
|
||||||
"entry_point": "suno",
|
|
||||||
},
|
|
||||||
],
|
|
||||||
}
|
|
||||||
_write_cache(manager._cache_path("harness"), harness)
|
|
||||||
_write_cache(manager._cache_path("public"), public)
|
|
||||||
|
|
||||||
|
|
||||||
def test_payload_merges_catalog_and_marks_unsupported_installs(tmp_path: Path) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
_seed_catalog(manager)
|
|
||||||
|
|
||||||
payload = manager.payload()
|
|
||||||
|
|
||||||
assert payload["catalog_updated_at"] == "2026-04-18"
|
|
||||||
apps = {app["name"]: app for app in payload["apps"]}
|
|
||||||
assert set(apps) == {
|
|
||||||
"clibrowser",
|
|
||||||
"dify-workflow",
|
|
||||||
"feishu",
|
|
||||||
"gimp",
|
|
||||||
"jimeng",
|
|
||||||
"shopify",
|
|
||||||
"suno",
|
|
||||||
}
|
|
||||||
assert apps["gimp"]["install_supported"] is True
|
|
||||||
assert apps["gimp"]["source"] == "harness+public"
|
|
||||||
assert apps["gimp"]["description"] == "Public duplicate entry"
|
|
||||||
assert apps["clibrowser"]["install_supported"] is False
|
|
||||||
assert apps["jimeng"]["install_supported"] is False
|
|
||||||
assert apps["suno"]["install_supported"] is True
|
|
||||||
assert apps["gimp"]["logo_url"]
|
|
||||||
assert apps["dify-workflow"]["logo_url"] == "https://cdn.simpleicons.org/dify/155EEF"
|
|
||||||
assert apps["feishu"]["logo_url"] == (
|
|
||||||
"https://www.google.com/s2/favicons?domain=larksuite.com&sz=64"
|
|
||||||
)
|
|
||||||
assert apps["jimeng"]["logo_url"] == "https://cdn.simpleicons.org/bytedance/3C8CFF"
|
|
||||||
assert apps["clibrowser"]["logo_url"] == (
|
|
||||||
"https://www.google.com/s2/favicons?domain=github.com/allthingssecurity/clibrowser&sz=64"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_install_dispatches_safe_pip_and_installs_skill(
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
_seed_catalog(manager)
|
|
||||||
calls: list[list[str]] = []
|
|
||||||
|
|
||||||
def fake_run(argv: list[str], *, timeout: int) -> subprocess.CompletedProcess[str]:
|
|
||||||
calls.append(argv)
|
|
||||||
return subprocess.CompletedProcess(argv, 0, stdout="ok", stderr="")
|
|
||||||
|
|
||||||
monkeypatch.setattr(manager, "_run_argv", fake_run)
|
|
||||||
monkeypatch.setattr(
|
|
||||||
manager,
|
|
||||||
"_fetch_skill_content",
|
|
||||||
lambda app: "---\nname: cli-anything-gimp\ndescription: GIMP\n---\n# GIMP\n",
|
|
||||||
)
|
|
||||||
|
|
||||||
payload = manager.install("gimp")
|
|
||||||
|
|
||||||
assert calls == [[sys.executable, "-m", "pip", "install", "cli-anything-gimp"]]
|
|
||||||
assert payload["last_action"]["ok"] is True
|
|
||||||
installed = json.loads(manager.installed_path.read_text(encoding="utf-8"))["apps"]
|
|
||||||
assert installed["gimp"]["entry_point"] == "cli-anything-gimp"
|
|
||||||
skill = manager.workspace / "skills" / "cli-app-gimp" / "SKILL.md"
|
|
||||||
assert skill.is_file()
|
|
||||||
assert 'run_cli_app` tool with `name="gimp"' in skill.read_text(encoding="utf-8")
|
|
||||||
|
|
||||||
|
|
||||||
def test_installed_state_writes_atomically_without_temp_leftovers(tmp_path: Path) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
|
|
||||||
manager._save_installed({"gimp": {"entry_point": "cli-anything-gimp"}})
|
|
||||||
manager._save_installed({"zoom": {"entry_point": "cli-anything-zoom"}})
|
|
||||||
|
|
||||||
installed = json.loads(manager.installed_path.read_text(encoding="utf-8"))["apps"]
|
|
||||||
assert set(installed) == {"zoom"}
|
|
||||||
assert not list(manager.installed_path.parent.glob(".installed.json.*.tmp"))
|
|
||||||
|
|
||||||
|
|
||||||
def test_fetch_skill_content_rejects_untrusted_urls(
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
|
|
||||||
def fail_get(*args, **kwargs):
|
|
||||||
raise AssertionError("untrusted skill URL should not be fetched")
|
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.httpx.get", fail_get)
|
|
||||||
|
|
||||||
assert manager._fetch_skill_content({
|
|
||||||
"name": "evil",
|
|
||||||
"skill_md": "https://example.com/SKILL.md",
|
|
||||||
}) is None
|
|
||||||
assert manager._fetch_skill_content({
|
|
||||||
"name": "evil",
|
|
||||||
"skill_md": "skills/../evil/SKILL.md",
|
|
||||||
}) is None
|
|
||||||
|
|
||||||
|
|
||||||
def test_fetch_skill_content_allows_cli_anything_raw_skill_url(
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
seen: list[str] = []
|
|
||||||
|
|
||||||
class Response:
|
|
||||||
text = "---\nname: cli-app-test\ndescription: Test\n---\n# Test\n"
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def raise_for_status() -> None:
|
|
||||||
return None
|
|
||||||
|
|
||||||
def fake_get(url: str, **kwargs):
|
|
||||||
seen.append(url)
|
|
||||||
return Response()
|
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.httpx.get", fake_get)
|
|
||||||
|
|
||||||
content = manager._fetch_skill_content({
|
|
||||||
"name": "gimp",
|
|
||||||
"skill_md": "https://raw.githubusercontent.com/HKUDS/CLI-Anything/main/skills/cli-anything-gimp/SKILL.md",
|
|
||||||
})
|
|
||||||
|
|
||||||
assert content and "# Test" in content
|
|
||||||
assert seen == [
|
|
||||||
"https://raw.githubusercontent.com/HKUDS/CLI-Anything/main/skills/cli-anything-gimp/SKILL.md"
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
def test_uninstall_removes_installed_state_and_generated_skill(
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
_seed_catalog(manager)
|
|
||||||
manager._save_installed({"gimp": {"entry_point": "cli-anything-gimp"}})
|
|
||||||
skill_dir = manager.workspace / "skills" / "cli-app-gimp"
|
|
||||||
skill_dir.mkdir(parents=True)
|
|
||||||
(skill_dir / "SKILL.md").write_text("# GIMP\n", encoding="utf-8")
|
|
||||||
monkeypatch.setattr(
|
|
||||||
manager,
|
|
||||||
"_run_argv",
|
|
||||||
lambda argv, *, timeout: subprocess.CompletedProcess(argv, 0, stdout="ok", stderr=""),
|
|
||||||
)
|
|
||||||
|
|
||||||
payload = manager.uninstall("gimp")
|
|
||||||
|
|
||||||
assert payload["last_action"]["ok"] is True
|
|
||||||
assert "gimp" not in json.loads(manager.installed_path.read_text(encoding="utf-8"))["apps"]
|
|
||||||
assert not skill_dir.exists()
|
|
||||||
|
|
||||||
|
|
||||||
def test_uninstall_uses_safe_python_m_pip_uninstall_command(
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
_seed_catalog(manager)
|
|
||||||
manager._save_installed({"suno": {"entry_point": "suno"}})
|
|
||||||
calls: list[list[str]] = []
|
|
||||||
|
|
||||||
def fake_run(argv: list[str], *, timeout: int) -> subprocess.CompletedProcess[str]:
|
|
||||||
calls.append(argv)
|
|
||||||
return subprocess.CompletedProcess(argv, 0, stdout="ok", stderr="")
|
|
||||||
|
|
||||||
monkeypatch.setattr(manager, "_run_argv", fake_run)
|
|
||||||
|
|
||||||
payload = manager.uninstall("suno")
|
|
||||||
|
|
||||||
assert calls == [[sys.executable, "-m", "pip", "uninstall", "-y", "suno-cli"]]
|
|
||||||
assert payload["last_action"]["ok"] is True
|
|
||||||
|
|
||||||
|
|
||||||
def test_mentioned_installed_apps_only_returns_installed_mentions(tmp_path: Path) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
manager._save_installed(
|
|
||||||
{
|
|
||||||
"gimp": {"entry_point": "cli-anything-gimp", "source": "harness"},
|
|
||||||
"zoom": {"entry_point": "cli-anything-zoom", "source": "public"},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
mentions = manager.mentioned_installed_apps("use @zoom and @krita, then @GIMP")
|
|
||||||
|
|
||||||
assert mentions == [
|
|
||||||
{
|
|
||||||
"name": "zoom",
|
|
||||||
"entry_point": "cli-anything-zoom",
|
|
||||||
"source": "public",
|
|
||||||
"skill": "skills/cli-app-zoom/SKILL.md",
|
|
||||||
"tool": "run_cli_app",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "gimp",
|
|
||||||
"entry_point": "cli-anything-gimp",
|
|
||||||
"source": "harness",
|
|
||||||
"skill": "skills/cli-app-gimp/SKILL.md",
|
|
||||||
"tool": "run_cli_app",
|
|
||||||
},
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
def test_install_rejects_unknown_and_script_strategy(tmp_path: Path) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
_seed_catalog(manager)
|
|
||||||
|
|
||||||
with pytest.raises(CliAppError, match="not found"):
|
|
||||||
manager.install("missing")
|
|
||||||
|
|
||||||
with pytest.raises(CliAppError, match="unsupported"):
|
|
||||||
manager.install("jimeng")
|
|
||||||
|
|
||||||
|
|
||||||
def test_run_installed_cli_uses_argv_without_shell(
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
_seed_catalog(manager)
|
|
||||||
resolved = str(tmp_path / "bin" / "cli-anything-gimp")
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"nanobot.cli_apps.service.shutil.which",
|
|
||||||
lambda entry: resolved if entry == "cli-anything-gimp" else None,
|
|
||||||
)
|
|
||||||
|
|
||||||
def fake_run(argv: list[str], **kwargs: object) -> subprocess.CompletedProcess[str]:
|
|
||||||
assert "shell" not in kwargs or kwargs["shell"] is False
|
|
||||||
return subprocess.CompletedProcess(
|
|
||||||
argv,
|
|
||||||
0,
|
|
||||||
stdout="ARGS=" + repr(argv[1:]),
|
|
||||||
stderr="",
|
|
||||||
)
|
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.subprocess.run", fake_run)
|
|
||||||
manager._save_installed(
|
|
||||||
{
|
|
||||||
"gimp": {
|
|
||||||
"version": "1.0.0",
|
|
||||||
"entry_point": "cli-anything-gimp",
|
|
||||||
"source": "harness",
|
|
||||||
"strategy": "pip",
|
|
||||||
}
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
result = manager.run("gimp", ["project", "list"], json_output=True)
|
|
||||||
|
|
||||||
assert "CLI app 'gimp' exited 0" in result
|
|
||||||
assert "['--json', 'project', 'list']" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_run_reports_created_artifacts(
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
_seed_catalog(manager)
|
|
||||||
resolved = str(tmp_path / "bin" / "cli-anything-gimp")
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"nanobot.cli_apps.service.shutil.which",
|
|
||||||
lambda entry: resolved if entry == "cli-anything-gimp" else None,
|
|
||||||
)
|
|
||||||
|
|
||||||
def fake_run(argv: list[str], **kwargs: object) -> subprocess.CompletedProcess[str]:
|
|
||||||
cwd = Path(str(kwargs["cwd"]))
|
|
||||||
(cwd / "diagram.png").write_bytes(b"\x89PNG\r\n\x1a\nimage")
|
|
||||||
return subprocess.CompletedProcess(argv, 0, stdout="done", stderr="")
|
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.subprocess.run", fake_run)
|
|
||||||
manager._save_installed({"gimp": {"entry_point": "cli-anything-gimp"}})
|
|
||||||
|
|
||||||
result = manager.run("gimp", ["render"])
|
|
||||||
|
|
||||||
assert "Artifacts created or updated:" in result
|
|
||||||
assert "diagram.png (previewable image" in result
|
|
||||||
assert "" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_run_blocks_working_dir_outside_workspace(tmp_path: Path) -> None:
|
|
||||||
manager = _manager(tmp_path)
|
|
||||||
_seed_catalog(manager)
|
|
||||||
manager._save_installed({"gimp": {"entry_point": "cli-anything-gimp"}})
|
|
||||||
|
|
||||||
with pytest.raises(CliAppError, match="outside the configured workspace"):
|
|
||||||
manager.run("gimp", working_dir="/etc", restrict_to_workspace=True)
|
|
||||||
@@ -1,125 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import json
|
|
||||||
import subprocess
|
|
||||||
import time
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from nanobot.agent.tools.cli_apps import CliAppsTool
|
|
||||||
from nanobot.cli_apps.service import CliAppManager, CliAppsRuntimeConfig
|
|
||||||
|
|
||||||
|
|
||||||
def _write_cache(path: Path, registry: dict) -> None:
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
path.write_text(
|
|
||||||
json.dumps({"_cached_at": time.time(), "data": registry}),
|
|
||||||
encoding="utf-8",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_run_cli_app_uses_installed_registry_app(
|
|
||||||
tmp_path: Path,
|
|
||||||
monkeypatch,
|
|
||||||
) -> None:
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
data_dir = tmp_path / "data"
|
|
||||||
registry = {
|
|
||||||
"meta": {"updated": "2026-04-16"},
|
|
||||||
"clis": [
|
|
||||||
{
|
|
||||||
"name": "gimp",
|
|
||||||
"display_name": "GIMP",
|
|
||||||
"version": "1.0.0",
|
|
||||||
"description": "Image editing",
|
|
||||||
"category": "image",
|
|
||||||
"install_cmd": "pip install cli-anything-gimp",
|
|
||||||
"entry_point": "cli-anything-gimp",
|
|
||||||
}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
_write_cache(data_dir / "harness_registry_cache.json", registry)
|
|
||||||
_write_cache(data_dir / "public_registry_cache.json", {"meta": {}, "clis": []})
|
|
||||||
CliAppManager(workspace=workspace, data_dir=data_dir)._save_installed(
|
|
||||||
{"gimp": {"entry_point": "cli-anything-gimp"}}
|
|
||||||
)
|
|
||||||
resolved = str(tmp_path / "bin" / "cli-anything-gimp")
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"nanobot.cli_apps.service.shutil.which",
|
|
||||||
lambda entry: resolved if entry == "cli-anything-gimp" else None,
|
|
||||||
)
|
|
||||||
|
|
||||||
def fake_run(argv: list[str], **kwargs: object) -> subprocess.CompletedProcess[str]:
|
|
||||||
assert "shell" not in kwargs or kwargs["shell"] is False
|
|
||||||
return subprocess.CompletedProcess(
|
|
||||||
argv,
|
|
||||||
0,
|
|
||||||
stdout="tool:" + " ".join(argv[1:]),
|
|
||||||
stderr="",
|
|
||||||
)
|
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.subprocess.run", fake_run)
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.get_runtime_subdir", lambda _name: data_dir)
|
|
||||||
|
|
||||||
tool = CliAppsTool(
|
|
||||||
workspace=workspace,
|
|
||||||
restrict_to_workspace=True,
|
|
||||||
runtime=CliAppsRuntimeConfig(run_timeout=5),
|
|
||||||
)
|
|
||||||
assert tool.name == "run_cli_app"
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
name="gimp",
|
|
||||||
args=["project", "list"],
|
|
||||||
json=True,
|
|
||||||
working_dir=str(workspace),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "CLI app 'gimp' exited 0" in result
|
|
||||||
assert "tool:--json project list" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_run_cli_app_rejects_uninstalled_app(tmp_path: Path, monkeypatch) -> None:
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
data_dir = tmp_path / "data"
|
|
||||||
registry = {
|
|
||||||
"meta": {"updated": "2026-04-16"},
|
|
||||||
"clis": [
|
|
||||||
{
|
|
||||||
"name": "gimp",
|
|
||||||
"display_name": "GIMP",
|
|
||||||
"version": "1.0.0",
|
|
||||||
"description": "Image editing",
|
|
||||||
"category": "image",
|
|
||||||
"install_cmd": "pip install cli-anything-gimp",
|
|
||||||
"entry_point": "cli-anything-gimp",
|
|
||||||
}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
_write_cache(data_dir / "harness_registry_cache.json", registry)
|
|
||||||
_write_cache(data_dir / "public_registry_cache.json", {"meta": {}, "clis": []})
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.get_runtime_subdir", lambda _name: data_dir)
|
|
||||||
tool = CliAppsTool(workspace=workspace, restrict_to_workspace=True)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(name="gimp"))
|
|
||||||
|
|
||||||
assert "not installed" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_run_cli_app_description_names_only_settings_installed_apps(tmp_path: Path, monkeypatch) -> None:
|
|
||||||
workspace = tmp_path / "workspace"
|
|
||||||
workspace.mkdir()
|
|
||||||
data_dir = tmp_path / "data"
|
|
||||||
CliAppManager(workspace=workspace, data_dir=data_dir)._save_installed(
|
|
||||||
{"drawio": {"entry_point": "cli-anything-drawio"}}
|
|
||||||
)
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.get_runtime_subdir", lambda _name: data_dir)
|
|
||||||
|
|
||||||
tool = CliAppsTool(workspace=workspace)
|
|
||||||
|
|
||||||
assert "Settings CLI Apps: drawio" in tool.description
|
|
||||||
assert "ordinary system CLIs such as git, gh" in tool.description
|
|
||||||
@@ -1,64 +0,0 @@
|
|||||||
"""Tests for CLI Apps loop helpers."""
|
|
||||||
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
from nanobot.cli_apps.service import CliAppManager
|
|
||||||
from nanobot.cli_apps.utils import runtime_lines, session_extra
|
|
||||||
|
|
||||||
|
|
||||||
def test_session_extra_returns_cli_apps_only_when_present() -> None:
|
|
||||||
cli_apps = [{"name": "zoom"}]
|
|
||||||
assert session_extra({"cli_apps": cli_apps}) == {"cli_apps": cli_apps}
|
|
||||||
assert session_extra({}) == {}
|
|
||||||
assert session_extra(None) == {}
|
|
||||||
|
|
||||||
|
|
||||||
def test_cli_app_mentions_inject_runtime_metadata(tmp_path, monkeypatch):
|
|
||||||
data_dir = tmp_path / "data"
|
|
||||||
monkeypatch.setattr("nanobot.cli_apps.service.get_runtime_subdir", lambda _name: data_dir)
|
|
||||||
manager = CliAppManager(workspace=tmp_path)
|
|
||||||
manager._save_installed(
|
|
||||||
{
|
|
||||||
"zoom": {
|
|
||||||
"entry_point": "cli-anything-zoom",
|
|
||||||
"source": "harness",
|
|
||||||
},
|
|
||||||
"krita": {
|
|
||||||
"entry_point": "cli-anything-krita",
|
|
||||||
"source": "harness",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
lines = runtime_lines(
|
|
||||||
SimpleNamespace(content="please use @zoom tonight; ignore @krita?", metadata={}),
|
|
||||||
tmp_path,
|
|
||||||
)
|
|
||||||
|
|
||||||
joined = "\n".join(lines)
|
|
||||||
assert "CLI App Mention: @zoom" in joined
|
|
||||||
assert "tool=run_cli_app" in joined
|
|
||||||
assert "entry_point=cli-anything-zoom" in joined
|
|
||||||
assert "skill=skills/cli-app-zoom/SKILL.md" in joined
|
|
||||||
|
|
||||||
|
|
||||||
def test_structured_cli_app_attachment_injects_runtime_metadata(tmp_path):
|
|
||||||
lines = runtime_lines(
|
|
||||||
SimpleNamespace(
|
|
||||||
content="please use @zoom tonight",
|
|
||||||
metadata={
|
|
||||||
"cli_apps": [{
|
|
||||||
"name": "zoom",
|
|
||||||
"entry_point": "cli-anything-zoom",
|
|
||||||
"display_name": "Zoom",
|
|
||||||
}],
|
|
||||||
},
|
|
||||||
),
|
|
||||||
tmp_path,
|
|
||||||
)
|
|
||||||
|
|
||||||
joined = "\n".join(lines)
|
|
||||||
assert "CLI App Attachment: @zoom" in joined
|
|
||||||
assert "tool=run_cli_app" in joined
|
|
||||||
assert "entry_point=cli-anything-zoom" in joined
|
|
||||||
assert "skill=skills/cli-app-zoom/SKILL.md" in joined
|
|
||||||
@@ -192,20 +192,3 @@ def test_match_provider_uses_preset_provider_when_forced() -> None:
|
|||||||
})
|
})
|
||||||
name = config.get_provider_name()
|
name = config.get_provider_name()
|
||||||
assert name == "anthropic"
|
assert name == "anthropic"
|
||||||
|
|
||||||
|
|
||||||
def test_match_provider_routes_forced_novita_model_api_models() -> None:
|
|
||||||
config = Config.model_validate({
|
|
||||||
"providers": {
|
|
||||||
"novita": {"apiKey": "sk-test"},
|
|
||||||
},
|
|
||||||
"agents": {
|
|
||||||
"defaults": {
|
|
||||||
"model": "deepseek-v4-pro",
|
|
||||||
"provider": "novita",
|
|
||||||
}
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
assert config.get_provider_name() == "novita"
|
|
||||||
assert config.get_api_base() == "https://api.novita.ai/openai"
|
|
||||||
|
|||||||
@@ -1,73 +0,0 @@
|
|||||||
"""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,74 +129,6 @@ 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")
|
||||||
|
|||||||
@@ -56,35 +56,6 @@ def test_custom_provider_parse_chunks_accepts_plain_text_chunks() -> None:
|
|||||||
assert result.content == "hello world"
|
assert result.content == "hello world"
|
||||||
|
|
||||||
|
|
||||||
def test_custom_provider_parse_chunks_deduplicates_parallel_tool_call_ids() -> None:
|
|
||||||
chunks = [{
|
|
||||||
"choices": [{
|
|
||||||
"finish_reason": "tool_calls",
|
|
||||||
"delta": {
|
|
||||||
"tool_calls": [
|
|
||||||
{
|
|
||||||
"index": 0,
|
|
||||||
"id": "call_dup",
|
|
||||||
"function": {"name": "read_file", "arguments": '{"path":"a.txt"}'},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"index": 1,
|
|
||||||
"id": "call_dup",
|
|
||||||
"function": {"name": "read_file", "arguments": '{"path":"b.txt"}'},
|
|
||||||
},
|
|
||||||
],
|
|
||||||
},
|
|
||||||
}],
|
|
||||||
}]
|
|
||||||
|
|
||||||
result = OpenAICompatProvider._parse_chunks(chunks)
|
|
||||||
ids = [tool_call.id for tool_call in result.tool_calls or []]
|
|
||||||
|
|
||||||
assert ids[0] == "call_dup"
|
|
||||||
assert len(ids) == 2
|
|
||||||
assert len(set(ids)) == 2
|
|
||||||
|
|
||||||
|
|
||||||
def test_local_provider_502_error_includes_reachability_hint() -> None:
|
def test_local_provider_502_error_includes_reachability_hint() -> None:
|
||||||
spec = find_by_name("ollama")
|
spec = find_by_name("ollama")
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
||||||
|
|||||||
@@ -65,7 +65,6 @@ 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,6 +1,5 @@
|
|||||||
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
|
||||||
|
|
||||||
@@ -9,15 +8,10 @@ import pytest
|
|||||||
|
|
||||||
from nanobot.providers.image_generation import (
|
from nanobot.providers.image_generation import (
|
||||||
AIHubMixImageGenerationClient,
|
AIHubMixImageGenerationClient,
|
||||||
CodexImageGenerationClient,
|
|
||||||
GeminiImageGenerationClient,
|
GeminiImageGenerationClient,
|
||||||
GeneratedImageResponse,
|
GeneratedImageResponse,
|
||||||
ImageGenerationError,
|
ImageGenerationError,
|
||||||
MiniMaxImageGenerationClient,
|
|
||||||
OllamaImageGenerationClient,
|
|
||||||
OpenAIImageGenerationClient,
|
|
||||||
OpenRouterImageGenerationClient,
|
OpenRouterImageGenerationClient,
|
||||||
StepFunImageGenerationClient,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
PNG_BYTES = (
|
PNG_BYTES = (
|
||||||
@@ -30,7 +24,6 @@ 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:
|
||||||
@@ -39,14 +32,12 @@ class FakeResponse:
|
|||||||
payload: dict[str, Any],
|
payload: dict[str, Any],
|
||||||
status_code: int = 200,
|
status_code: int = 200,
|
||||||
content: bytes = b"",
|
content: bytes = b"",
|
||||||
sse_lines: list[str] | None = None,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
self._payload = payload
|
self._payload = payload
|
||||||
self.status_code = status_code
|
self.status_code = status_code
|
||||||
self.text = str(payload)
|
self.text = str(payload)
|
||||||
self.content = content
|
self.content = content
|
||||||
self.request = httpx.Request("POST", "https://openrouter.ai/api/v1/chat/completions")
|
self.request = httpx.Request("POST", "https://openrouter.ai/api/v1/chat/completions")
|
||||||
self._sse_lines = sse_lines
|
|
||||||
|
|
||||||
def json(self) -> dict[str, Any]:
|
def json(self) -> dict[str, Any]:
|
||||||
return self._payload
|
return self._payload
|
||||||
@@ -56,15 +47,6 @@ class FakeResponse:
|
|||||||
response = httpx.Response(self.status_code, request=self.request, text=self.text)
|
response = httpx.Response(self.status_code, request=self.request, text=self.text)
|
||||||
raise httpx.HTTPStatusError("failed", request=self.request, response=response)
|
raise httpx.HTTPStatusError("failed", request=self.request, response=response)
|
||||||
|
|
||||||
async def aiter_lines(self):
|
|
||||||
if self._sse_lines is not None:
|
|
||||||
for line in self._sse_lines:
|
|
||||||
yield line
|
|
||||||
return
|
|
||||||
# Fallback: treat response text as SSE lines
|
|
||||||
for line in self.text.split("\n"):
|
|
||||||
yield line
|
|
||||||
|
|
||||||
|
|
||||||
class FakeClient:
|
class FakeClient:
|
||||||
def __init__(self, response: FakeResponse) -> None:
|
def __init__(self, response: FakeResponse) -> None:
|
||||||
@@ -147,54 +129,6 @@ async def test_openrouter_image_generation_requires_api_key() -> None:
|
|||||||
await client.generate(prompt="draw", model="model")
|
await client.generate(prompt="draw", model="model")
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_ollama_image_generation_payload_and_response() -> None:
|
|
||||||
raw_b64 = PNG_DATA_URL.removeprefix("data:image/png;base64,")
|
|
||||||
fake = FakeClient(FakeResponse({"image": raw_b64}))
|
|
||||||
client = OllamaImageGenerationClient(
|
|
||||||
api_key="ollama-test",
|
|
||||||
api_base="http://localhost:11434/v1/",
|
|
||||||
extra_headers={"X-Test": "1"},
|
|
||||||
extra_body={"seed": 123},
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await client.generate(
|
|
||||||
prompt="a sunset",
|
|
||||||
model="x/z-image-turbo",
|
|
||||||
aspect_ratio="16:9",
|
|
||||||
image_size="1K",
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.images == [PNG_DATA_URL]
|
|
||||||
assert response.content == ""
|
|
||||||
|
|
||||||
call = fake.calls[0]
|
|
||||||
assert call["url"] == "http://localhost:11434/api/generate"
|
|
||||||
assert call["headers"]["Authorization"] == "Bearer ollama-test"
|
|
||||||
assert call["headers"]["X-Test"] == "1"
|
|
||||||
body = call["json"]
|
|
||||||
assert body["model"] == "x/z-image-turbo"
|
|
||||||
assert body["prompt"] == "a sunset"
|
|
||||||
assert body["width"] == 1024
|
|
||||||
assert body["height"] == 576
|
|
||||||
assert body["steps"] == 0
|
|
||||||
assert body["stream"] is False
|
|
||||||
assert body["seed"] == 123
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_ollama_image_generation_rejects_reference_images() -> None:
|
|
||||||
client = OllamaImageGenerationClient(api_key=None)
|
|
||||||
|
|
||||||
with pytest.raises(ImageGenerationError, match="reference images"):
|
|
||||||
await client.generate(
|
|
||||||
prompt="edit this",
|
|
||||||
model="x/z-image-turbo",
|
|
||||||
reference_images=["ref.png"],
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_aihubmix_image_generation_payload_and_response() -> None:
|
async def test_aihubmix_image_generation_payload_and_response() -> None:
|
||||||
raw_b64 = PNG_DATA_URL.removeprefix("data:image/png;base64,")
|
raw_b64 = PNG_DATA_URL.removeprefix("data:image/png;base64,")
|
||||||
@@ -271,20 +205,6 @@ async def test_aihubmix_image_generation_downloads_url_response() -> None:
|
|||||||
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,")
|
RAW_B64 = PNG_DATA_URL.removeprefix("data:image/png;base64,")
|
||||||
|
|
||||||
|
|
||||||
@@ -410,11 +330,6 @@ async def test_gemini_requires_api_key() -> None:
|
|||||||
await client.generate(prompt="draw", model="imagen-4.0-generate-001")
|
await client.generate(prompt="draw", model="imagen-4.0-generate-001")
|
||||||
|
|
||||||
|
|
||||||
def test_gemini_image_client_uses_native_api_base_by_default() -> None:
|
|
||||||
client = GeminiImageGenerationClient(api_key="AIza-test")
|
|
||||||
assert client.api_base == "https://generativelanguage.googleapis.com/v1beta"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_gemini_no_images_raises() -> None:
|
async def test_gemini_no_images_raises() -> None:
|
||||||
fake = FakeClient(FakeResponse({"candidates": [{"content": {"parts": [{"text": "sorry"}]}}]}))
|
fake = FakeClient(FakeResponse({"candidates": [{"content": {"parts": [{"text": "sorry"}]}}]}))
|
||||||
@@ -422,608 +337,3 @@ async def test_gemini_no_images_raises() -> None:
|
|||||||
|
|
||||||
with pytest.raises(ImageGenerationError, match="returned no images"):
|
with pytest.raises(ImageGenerationError, match="returned no images"):
|
||||||
await client.generate(prompt="draw", model="gemini-2.0-flash-preview-image-generation")
|
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,")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_minimax_base64_response_uses_detected_mime() -> None:
|
|
||||||
raw_b64 = base64.b64encode(JPEG_BYTES).decode("ascii")
|
|
||||||
fake = FakeClient(FakeResponse({"data": {"image_base64": [raw_b64]}}))
|
|
||||||
client = MiniMaxImageGenerationClient(api_key="sk-mm-test", client=fake) # type: ignore[arg-type]
|
|
||||||
|
|
||||||
response = await client.generate(prompt="draw", model="image-01")
|
|
||||||
|
|
||||||
assert response.images == [f"data:image/jpeg;base64,{raw_b64}"]
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# 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")
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# OpenAI
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_payload_and_response() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
api_base="https://api.openai.com/v1",
|
|
||||||
extra_headers={"X-Test": "1"},
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await client.generate(
|
|
||||||
prompt="a cat on the moon",
|
|
||||||
model="dall-e-3",
|
|
||||||
aspect_ratio="16:9",
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.images == [PNG_DATA_URL]
|
|
||||||
call = fake.calls[0]
|
|
||||||
assert call["url"] == "https://api.openai.com/v1/images/generations"
|
|
||||||
assert call["headers"]["Authorization"] == "Bearer sk-openai-test"
|
|
||||||
assert call["headers"]["X-Test"] == "1"
|
|
||||||
body = call["json"]
|
|
||||||
assert body["model"] == "dall-e-3"
|
|
||||||
assert body["prompt"] == "a cat on the moon"
|
|
||||||
assert body["response_format"] == "b64_json"
|
|
||||||
assert body["n"] == 1
|
|
||||||
assert body["size"] == "1792x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_b64_json_response_uses_detected_mime() -> None:
|
|
||||||
raw_b64 = base64.b64encode(JPEG_BYTES).decode("ascii")
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": raw_b64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await client.generate(prompt="draw", model="dall-e-3")
|
|
||||||
|
|
||||||
assert response.images == [f"data:image/jpeg;base64,{raw_b64}"]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_url_download_fallback() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"url": "https://cdn.example/image.png"}]}))
|
|
||||||
fake.get_response = FakeResponse({}, content=PNG_BYTES)
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await client.generate(prompt="draw", model="dall-e-3")
|
|
||||||
|
|
||||||
assert response.images[0].startswith("data:image/png;base64,")
|
|
||||||
assert fake.get_calls[0]["url"] == "https://cdn.example/image.png"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_multiple_images() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({
|
|
||||||
"data": [
|
|
||||||
{"b64_json": RAW_B64},
|
|
||||||
{"b64_json": RAW_B64},
|
|
||||||
]
|
|
||||||
}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await client.generate(prompt="draw", model="dall-e-3")
|
|
||||||
|
|
||||||
assert len(response.images) == 2
|
|
||||||
assert response.images == [PNG_DATA_URL, PNG_DATA_URL]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_aspect_ratio_to_size() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(prompt="draw", model="dall-e-3", aspect_ratio="1:1")
|
|
||||||
assert fake.calls[0]["json"]["size"] == "1024x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_dalle3_uses_supported_orientation_sizes() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(prompt="draw", model="dall-e-3", aspect_ratio="3:4")
|
|
||||||
await client.generate(prompt="draw", model="dall-e-3", aspect_ratio="4:3")
|
|
||||||
|
|
||||||
assert fake.calls[0]["json"]["size"] == "1024x1792"
|
|
||||||
assert fake.calls[1]["json"]["size"] == "1792x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_dalle2_uses_square_size_for_non_square_ratios() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(prompt="draw", model="dall-e-2", aspect_ratio="16:9")
|
|
||||||
|
|
||||||
assert fake.calls[0]["json"]["size"] == "1024x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_gpt_image_uses_supported_landscape_size() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(prompt="draw", model="gpt-image-1", aspect_ratio="16:9")
|
|
||||||
|
|
||||||
assert fake.calls[0]["json"]["size"] == "1536x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_gpt_image_uses_supported_orientation_sizes() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(prompt="draw", model="gpt-image-1", aspect_ratio="3:4")
|
|
||||||
await client.generate(prompt="draw", model="gpt-image-1", aspect_ratio="4:3")
|
|
||||||
|
|
||||||
assert fake.calls[0]["json"]["size"] == "1024x1536"
|
|
||||||
assert fake.calls[1]["json"]["size"] == "1536x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_default_size_when_no_aspect_ratio() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(prompt="draw", model="dall-e-3")
|
|
||||||
|
|
||||||
body = fake.calls[0]["json"]
|
|
||||||
assert body["size"] == "1024x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_ignores_explicit_size_unsupported_by_model_family() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(
|
|
||||||
prompt="draw",
|
|
||||||
model="dall-e-3",
|
|
||||||
aspect_ratio="16:9",
|
|
||||||
image_size="1536x1024",
|
|
||||||
)
|
|
||||||
|
|
||||||
body = fake.calls[0]["json"]
|
|
||||||
assert body["size"] == "1792x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_uses_explicit_image_size() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": [{"b64_json": RAW_B64}]}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(
|
|
||||||
prompt="draw",
|
|
||||||
model="dall-e-3",
|
|
||||||
aspect_ratio="16:9",
|
|
||||||
image_size="1024x1024",
|
|
||||||
)
|
|
||||||
|
|
||||||
body = fake.calls[0]["json"]
|
|
||||||
assert body["size"] == "1024x1024"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_requires_api_key() -> None:
|
|
||||||
client = OpenAIImageGenerationClient(api_key=None)
|
|
||||||
|
|
||||||
with pytest.raises(ImageGenerationError, match="API key"):
|
|
||||||
await client.generate(prompt="draw", model="dall-e-3")
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# OpenAI Codex (Responses API)
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_codex_payload_and_response(monkeypatch) -> None:
|
|
||||||
import sys
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class FakeToken:
|
|
||||||
account_id: str = "acct-123"
|
|
||||||
access: str = "oauth-token"
|
|
||||||
|
|
||||||
async def fake_to_thread(fn, *args, **kwargs):
|
|
||||||
return fn(*args, **kwargs)
|
|
||||||
|
|
||||||
monkeypatch.setattr("asyncio.to_thread", fake_to_thread)
|
|
||||||
fake_oauth = SimpleNamespace(get_token=lambda: FakeToken())
|
|
||||||
monkeypatch.setitem(sys.modules, "oauth_cli_kit", fake_oauth)
|
|
||||||
|
|
||||||
sse_lines = [
|
|
||||||
'data: {"type":"response.output_item.added","item":{"id":"ig_1","type":"image_generation_call","status":"in_progress"}}',
|
|
||||||
"",
|
|
||||||
f'data: {{"type":"response.output_item.done","item":{{"id":"ig_1","type":"image_generation_call","result":"{PNG_DATA_URL}","status":"completed"}}}}',
|
|
||||||
"",
|
|
||||||
'data: [DONE]',
|
|
||||||
"",
|
|
||||||
]
|
|
||||||
fake = FakeClient(FakeResponse({}, sse_lines=sse_lines))
|
|
||||||
client = CodexImageGenerationClient(
|
|
||||||
api_key=None,
|
|
||||||
api_base="https://chatgpt.com/backend-api",
|
|
||||||
extra_headers={"X-Test": "1"},
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await client.generate(
|
|
||||||
prompt="draw a cat",
|
|
||||||
model="gpt-5.4",
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.images == [PNG_DATA_URL]
|
|
||||||
assert response.content == ""
|
|
||||||
call = fake.calls[0]
|
|
||||||
assert call["url"] == "https://chatgpt.com/backend-api/codex/responses"
|
|
||||||
assert call["headers"]["Authorization"] == "Bearer oauth-token"
|
|
||||||
assert call["headers"]["chatgpt-account-id"] == "acct-123"
|
|
||||||
assert call["headers"]["OpenAI-Beta"] == "responses=experimental"
|
|
||||||
assert call["headers"]["X-Test"] == "1"
|
|
||||||
body = call["json"]
|
|
||||||
assert body["model"] == "gpt-5.4"
|
|
||||||
assert body["instructions"] == "Generate an image based on the user's request."
|
|
||||||
assert body["input"] == [{"role": "user", "content": "draw a cat"}]
|
|
||||||
assert body["tools"] == [{"type": "image_generation"}]
|
|
||||||
assert body["tool_choice"] == "auto"
|
|
||||||
assert body["store"] is False
|
|
||||||
assert body["stream"] is True
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_codex_strips_model_prefix(monkeypatch) -> None:
|
|
||||||
import sys
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class FakeToken:
|
|
||||||
account_id: str = "acct-123"
|
|
||||||
access: str = "oauth-token"
|
|
||||||
|
|
||||||
async def fake_to_thread(fn, *args, **kwargs):
|
|
||||||
return fn(*args, **kwargs)
|
|
||||||
|
|
||||||
monkeypatch.setattr("asyncio.to_thread", fake_to_thread)
|
|
||||||
fake_oauth = SimpleNamespace(get_token=lambda: FakeToken())
|
|
||||||
monkeypatch.setitem(sys.modules, "oauth_cli_kit", fake_oauth)
|
|
||||||
|
|
||||||
fake = FakeClient(FakeResponse({}, sse_lines=[
|
|
||||||
f'data: {{"type":"response.output_item.done","item":{{"type":"image_generation_call","result":"{PNG_DATA_URL}"}}}}',
|
|
||||||
"",
|
|
||||||
'data: [DONE]',
|
|
||||||
"",
|
|
||||||
]))
|
|
||||||
client = CodexImageGenerationClient(
|
|
||||||
api_key=None, client=fake # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
await client.generate(prompt="draw", model="openai-codex/gpt-5.4")
|
|
||||||
|
|
||||||
assert fake.calls[0]["json"]["model"] == "gpt-5.4"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_codex_requires_oauth(monkeypatch) -> None:
|
|
||||||
async def fake_to_thread(fn, *args, **kwargs):
|
|
||||||
raise RuntimeError("no token")
|
|
||||||
|
|
||||||
monkeypatch.setattr("asyncio.to_thread", fake_to_thread)
|
|
||||||
|
|
||||||
client = CodexImageGenerationClient(api_key=None)
|
|
||||||
|
|
||||||
with pytest.raises(ImageGenerationError, match="OAuth token"):
|
|
||||||
await client.generate(prompt="draw", model="gpt-5.4")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_codex_no_images_raises(monkeypatch) -> None:
|
|
||||||
import sys
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class FakeToken:
|
|
||||||
account_id: str = "acct-123"
|
|
||||||
access: str = "oauth-token"
|
|
||||||
|
|
||||||
async def fake_to_thread(fn, *args, **kwargs):
|
|
||||||
return fn(*args, **kwargs)
|
|
||||||
|
|
||||||
monkeypatch.setattr("asyncio.to_thread", fake_to_thread)
|
|
||||||
fake_oauth = SimpleNamespace(get_token=lambda: FakeToken())
|
|
||||||
monkeypatch.setitem(sys.modules, "oauth_cli_kit", fake_oauth)
|
|
||||||
|
|
||||||
fake = FakeClient(FakeResponse({}, sse_lines=[
|
|
||||||
'data: {"type":"response.completed","response":{"status":"completed"}}',
|
|
||||||
"",
|
|
||||||
'data: [DONE]',
|
|
||||||
"",
|
|
||||||
]))
|
|
||||||
client = CodexImageGenerationClient(
|
|
||||||
api_key=None, client=fake # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(ImageGenerationError, match="returned no images"):
|
|
||||||
await client.generate(prompt="draw", model="gpt-5.4")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_codex_extracts_text_content(monkeypatch) -> None:
|
|
||||||
import sys
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class FakeToken:
|
|
||||||
account_id: str = "acct-123"
|
|
||||||
access: str = "oauth-token"
|
|
||||||
|
|
||||||
async def fake_to_thread(fn, *args, **kwargs):
|
|
||||||
return fn(*args, **kwargs)
|
|
||||||
|
|
||||||
monkeypatch.setattr("asyncio.to_thread", fake_to_thread)
|
|
||||||
fake_oauth = SimpleNamespace(get_token=lambda: FakeToken())
|
|
||||||
monkeypatch.setitem(sys.modules, "oauth_cli_kit", fake_oauth)
|
|
||||||
|
|
||||||
fake = FakeClient(FakeResponse({}, sse_lines=[
|
|
||||||
'data: {"type":"response.output_text.delta","delta":"Here "}',
|
|
||||||
"",
|
|
||||||
'data: {"type":"response.output_text.delta","delta":"is your cat image."}',
|
|
||||||
"",
|
|
||||||
f'data: {{"type":"response.output_item.done","item":{{"type":"image_generation_call","result":"{PNG_DATA_URL}"}}}}',
|
|
||||||
"",
|
|
||||||
'data: [DONE]',
|
|
||||||
"",
|
|
||||||
]))
|
|
||||||
client = CodexImageGenerationClient(
|
|
||||||
api_key=None, client=fake # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await client.generate(prompt="draw a cat", model="gpt-5.4")
|
|
||||||
|
|
||||||
assert response.images == [PNG_DATA_URL]
|
|
||||||
assert response.content == "Here is your cat image."
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_codex_json_result_format(monkeypatch) -> None:
|
|
||||||
"""image_generation_call result can be a dict with image_url key."""
|
|
||||||
import sys
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class FakeToken:
|
|
||||||
account_id: str = "acct-123"
|
|
||||||
access: str = "oauth-token"
|
|
||||||
|
|
||||||
async def fake_to_thread(fn, *args, **kwargs):
|
|
||||||
return fn(*args, **kwargs)
|
|
||||||
|
|
||||||
monkeypatch.setattr("asyncio.to_thread", fake_to_thread)
|
|
||||||
fake_oauth = SimpleNamespace(get_token=lambda: FakeToken())
|
|
||||||
monkeypatch.setitem(sys.modules, "oauth_cli_kit", fake_oauth)
|
|
||||||
|
|
||||||
fake = FakeClient(FakeResponse({}, sse_lines=[
|
|
||||||
f'data: {{"type":"response.output_item.done","item":{{"type":"image_generation_call","result":{{"image_url":"{PNG_DATA_URL}"}}}}}}',
|
|
||||||
"",
|
|
||||||
'data: [DONE]',
|
|
||||||
"",
|
|
||||||
]))
|
|
||||||
client = CodexImageGenerationClient(
|
|
||||||
api_key=None, client=fake # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await client.generate(prompt="draw", model="gpt-5.4")
|
|
||||||
|
|
||||||
assert response.images == [PNG_DATA_URL]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_openai_no_images_raises() -> None:
|
|
||||||
fake = FakeClient(FakeResponse({"data": []}))
|
|
||||||
client = OpenAIImageGenerationClient(
|
|
||||||
api_key="sk-openai-test",
|
|
||||||
client=fake, # type: ignore[arg-type]
|
|
||||||
)
|
|
||||||
|
|
||||||
with pytest.raises(ImageGenerationError, match="returned no images"):
|
|
||||||
await client.generate(prompt="draw", model="dall-e-3")
|
|
||||||
|
|||||||
@@ -164,130 +164,6 @@ 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."""
|
||||||
@@ -326,98 +202,6 @@ 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)
|
||||||
@@ -441,15 +225,6 @@ def test_openrouter_spec_is_gateway() -> None:
|
|||||||
assert spec.default_api_base == "https://openrouter.ai/api/v1"
|
assert spec.default_api_base == "https://openrouter.ai/api/v1"
|
||||||
|
|
||||||
|
|
||||||
def test_novita_spec_uses_openai_compatible_gateway() -> None:
|
|
||||||
spec = find_by_name("novita")
|
|
||||||
assert spec is not None
|
|
||||||
assert spec.is_gateway is True
|
|
||||||
assert spec.backend == "openai_compat"
|
|
||||||
assert spec.env_key == "NOVITA_API_KEY"
|
|
||||||
assert spec.default_api_base == "https://api.novita.ai/openai"
|
|
||||||
|
|
||||||
|
|
||||||
def test_gemma_routes_to_gemini_provider() -> None:
|
def test_gemma_routes_to_gemini_provider() -> None:
|
||||||
"""gemma models (e.g. gemma-3-27b-it) must auto-route to Gemini when GEMINI_API_KEY is set.
|
"""gemma models (e.g. gemma-3-27b-it) must auto-route to Gemini when GEMINI_API_KEY is set.
|
||||||
Users running gemma via the Gemini API endpoint expect automatic provider detection."""
|
Users running gemma via the Gemini API endpoint expect automatic provider detection."""
|
||||||
@@ -458,34 +233,27 @@ def test_gemma_routes_to_gemini_provider() -> None:
|
|||||||
assert "gemma" in spec.keywords
|
assert "gemma" in spec.keywords
|
||||||
|
|
||||||
|
|
||||||
def test_gemini_spec_keeps_openai_compat_base() -> None:
|
def test_openrouter_sets_default_attribution_headers() -> None:
|
||||||
spec = find_by_name("gemini")
|
|
||||||
assert spec is not None
|
|
||||||
assert spec.default_api_base == "https://generativelanguage.googleapis.com/v1beta/openai/"
|
|
||||||
|
|
||||||
|
|
||||||
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 mock_client_cls:
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient:
|
||||||
provider = OpenAICompatProvider(
|
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 = mock_client_cls.call_args.kwargs["default_headers"]
|
headers = MockClient.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
|
||||||
|
|
||||||
|
|
||||||
async def test_openrouter_user_headers_override_default_attribution() -> None:
|
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 mock_client_cls:
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient:
|
||||||
provider = OpenAICompatProvider(
|
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",
|
||||||
@@ -496,9 +264,8 @@ async def test_openrouter_user_headers_override_default_attribution() -> None:
|
|||||||
},
|
},
|
||||||
spec=spec,
|
spec=spec,
|
||||||
)
|
)
|
||||||
await provider._ensure_client()
|
|
||||||
|
|
||||||
headers = mock_client_cls.call_args.kwargs["default_headers"]
|
headers = MockClient.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"
|
||||||
@@ -1022,41 +789,6 @@ def test_openai_compat_keeps_tool_calls_after_consecutive_assistant_messages() -
|
|||||||
assert sanitized[2]["tool_call_id"] == "3ec83c30d"
|
assert sanitized[2]["tool_call_id"] == "3ec83c30d"
|
||||||
|
|
||||||
|
|
||||||
def test_openai_compat_deduplicates_duplicate_tool_call_ids_in_history() -> None:
|
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
|
||||||
provider = OpenAICompatProvider()
|
|
||||||
|
|
||||||
sanitized = provider._sanitize_messages([
|
|
||||||
{"role": "user", "content": "check both files"},
|
|
||||||
{
|
|
||||||
"role": "assistant",
|
|
||||||
"content": "",
|
|
||||||
"tool_calls": [
|
|
||||||
{
|
|
||||||
"id": "ab1b45c2a",
|
|
||||||
"type": "function",
|
|
||||||
"function": {"name": "read_file", "arguments": '{"path":"a.txt"}'},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": "ab1b45c2a",
|
|
||||||
"type": "function",
|
|
||||||
"function": {"name": "read_file", "arguments": '{"path":"b.txt"}'},
|
|
||||||
},
|
|
||||||
],
|
|
||||||
},
|
|
||||||
{"role": "tool", "tool_call_id": "ab1b45c2a", "name": "read_file", "content": "a"},
|
|
||||||
{"role": "tool", "tool_call_id": "ab1b45c2a", "name": "read_file", "content": "b"},
|
|
||||||
{"role": "user", "content": "continue"},
|
|
||||||
])
|
|
||||||
|
|
||||||
tool_call_ids = [tc["id"] for tc in sanitized[1]["tool_calls"]]
|
|
||||||
tool_result_ids = [sanitized[2]["tool_call_id"], sanitized[3]["tool_call_id"]]
|
|
||||||
|
|
||||||
assert tool_call_ids[0] == "ab1b45c2a"
|
|
||||||
assert len(tool_call_ids) == len(set(tool_call_ids)) == 2
|
|
||||||
assert tool_result_ids == tool_call_ids
|
|
||||||
|
|
||||||
|
|
||||||
def test_openai_compat_stringifies_dict_tool_arguments() -> None:
|
def test_openai_compat_stringifies_dict_tool_arguments() -> None:
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
||||||
provider = OpenAICompatProvider()
|
provider = OpenAICompatProvider()
|
||||||
@@ -1426,15 +1158,12 @@ def test_kimi_k25_thinking_enabled() -> None:
|
|||||||
"""kimi-k2.5 with reasoning_effort set should opt in to thinking."""
|
"""kimi-k2.5 with reasoning_effort set should opt in to thinking."""
|
||||||
kw = _build_kwargs_for("moonshot", "kimi-k2.5", reasoning_effort="medium")
|
kw = _build_kwargs_for("moonshot", "kimi-k2.5", reasoning_effort="medium")
|
||||||
assert kw.get("extra_body") == {"thinking": {"type": "enabled"}}
|
assert kw.get("extra_body") == {"thinking": {"type": "enabled"}}
|
||||||
# Moonshot rejects both 'reasoning_effort' and 'thinking' (#3939)
|
|
||||||
assert "reasoning_effort" not in kw
|
|
||||||
|
|
||||||
|
|
||||||
def test_kimi_k25_thinking_disabled_for_minimal() -> None:
|
def test_kimi_k25_thinking_disabled_for_minimal() -> None:
|
||||||
"""reasoning_effort='minimal' maps to thinking disabled for kimi-k2.5."""
|
"""reasoning_effort='minimal' maps to thinking disabled for kimi-k2.5."""
|
||||||
kw = _build_kwargs_for("moonshot", "kimi-k2.5", reasoning_effort="minimal")
|
kw = _build_kwargs_for("moonshot", "kimi-k2.5", reasoning_effort="minimal")
|
||||||
assert kw.get("extra_body") == {"thinking": {"type": "disabled"}}
|
assert kw.get("extra_body") == {"thinking": {"type": "disabled"}}
|
||||||
assert "reasoning_effort" not in kw
|
|
||||||
|
|
||||||
|
|
||||||
def test_kimi_k25_no_extra_body_when_reasoning_effort_none() -> None:
|
def test_kimi_k25_no_extra_body_when_reasoning_effort_none() -> None:
|
||||||
@@ -1444,36 +1173,21 @@ def test_kimi_k25_no_extra_body_when_reasoning_effort_none() -> None:
|
|||||||
|
|
||||||
|
|
||||||
def test_kimi_k25_thinking_enabled_with_openrouter_prefix() -> None:
|
def test_kimi_k25_thinking_enabled_with_openrouter_prefix() -> None:
|
||||||
"""OpenRouter-style model names like moonshotai/kimi-k2.5 must trigger thinking.
|
"""OpenRouter-style model names like moonshotai/kimi-k2.5 must trigger thinking."""
|
||||||
|
|
||||||
OR drops upstream-provider `thinking` fields, so the same intent also has
|
|
||||||
to go through OR's `reasoning.effort` shape (#3851 follow-up).
|
|
||||||
"""
|
|
||||||
kw = _build_kwargs_for("openrouter", "moonshotai/kimi-k2.5", reasoning_effort="medium")
|
kw = _build_kwargs_for("openrouter", "moonshotai/kimi-k2.5", reasoning_effort="medium")
|
||||||
assert kw.get("extra_body") == {
|
assert kw.get("extra_body") == {"thinking": {"type": "enabled"}}
|
||||||
"thinking": {"type": "enabled"},
|
|
||||||
"reasoning": {"effort": "medium"},
|
|
||||||
}
|
|
||||||
# Even via OR, reasoning_effort wire kwarg is dropped for kimi models
|
|
||||||
assert "reasoning_effort" not in kw
|
|
||||||
|
|
||||||
|
|
||||||
def test_kimi_k26_thinking_enabled() -> None:
|
def test_kimi_k26_thinking_enabled() -> None:
|
||||||
"""kimi-k2.6 with reasoning_effort set should opt in to thinking."""
|
"""kimi-k2.6 with reasoning_effort set should opt in to thinking."""
|
||||||
kw = _build_kwargs_for("moonshot", "kimi-k2.6", reasoning_effort="medium")
|
kw = _build_kwargs_for("moonshot", "kimi-k2.6", reasoning_effort="medium")
|
||||||
assert kw.get("extra_body") == {"thinking": {"type": "enabled"}}
|
assert kw.get("extra_body") == {"thinking": {"type": "enabled"}}
|
||||||
assert "reasoning_effort" not in kw
|
|
||||||
|
|
||||||
|
|
||||||
def test_kimi_k26_thinking_enabled_with_openrouter_prefix() -> None:
|
def test_kimi_k26_thinking_enabled_with_openrouter_prefix() -> None:
|
||||||
"""OpenRouter-style names like moonshotai/kimi-k2.6 must trigger thinking
|
"""OpenRouter-style names like moonshotai/kimi-k2.6 must trigger thinking."""
|
||||||
via both upstream `thinking` and OR's `reasoning.effort`."""
|
|
||||||
kw = _build_kwargs_for("openrouter", "moonshotai/kimi-k2.6", reasoning_effort="medium")
|
kw = _build_kwargs_for("openrouter", "moonshotai/kimi-k2.6", reasoning_effort="medium")
|
||||||
assert kw.get("extra_body") == {
|
assert kw.get("extra_body") == {"thinking": {"type": "enabled"}}
|
||||||
"thinking": {"type": "enabled"},
|
|
||||||
"reasoning": {"effort": "medium"},
|
|
||||||
}
|
|
||||||
assert "reasoning_effort" not in kw
|
|
||||||
|
|
||||||
|
|
||||||
def test_moonshot_kimi_k26_temperature_override() -> None:
|
def test_moonshot_kimi_k26_temperature_override() -> None:
|
||||||
@@ -1492,7 +1206,6 @@ def test_kimi_k26_code_preview_thinking_enabled() -> None:
|
|||||||
"""k2.6-code-preview also supports thinking; should behave like k2.5."""
|
"""k2.6-code-preview also supports thinking; should behave like k2.5."""
|
||||||
kw = _build_kwargs_for("moonshot", "k2.6-code-preview", reasoning_effort="high")
|
kw = _build_kwargs_for("moonshot", "k2.6-code-preview", reasoning_effort="high")
|
||||||
assert kw.get("extra_body") == {"thinking": {"type": "enabled"}}
|
assert kw.get("extra_body") == {"thinking": {"type": "enabled"}}
|
||||||
assert "reasoning_effort" not in kw
|
|
||||||
|
|
||||||
|
|
||||||
def test_kimi_k2_series_no_thinking_injection() -> None:
|
def test_kimi_k2_series_no_thinking_injection() -> None:
|
||||||
@@ -1522,7 +1235,6 @@ def test_kimi_k25_thinking_disabled_for_none_string() -> None:
|
|||||||
"""reasoning_effort='none' maps to thinking disabled for kimi-k2.5."""
|
"""reasoning_effort='none' maps to thinking disabled for kimi-k2.5."""
|
||||||
kw = _build_kwargs_for("moonshot", "kimi-k2.5", reasoning_effort="none")
|
kw = _build_kwargs_for("moonshot", "kimi-k2.5", reasoning_effort="none")
|
||||||
assert kw.get("extra_body") == {"thinking": {"type": "disabled"}}
|
assert kw.get("extra_body") == {"thinking": {"type": "disabled"}}
|
||||||
assert "reasoning_effort" not in kw
|
|
||||||
|
|
||||||
|
|
||||||
def test_dashscope_thinking_disabled_for_none_string() -> None:
|
def test_dashscope_thinking_disabled_for_none_string() -> None:
|
||||||
|
|||||||
@@ -44,15 +44,9 @@ 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", ""],
|
["refusal", "content_filter", "error", "length", "function_call", ""],
|
||||||
)
|
)
|
||||||
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,18 +85,17 @@ class TestIsLocalEndpoint:
|
|||||||
class TestLocalKeepaliveConfig:
|
class TestLocalKeepaliveConfig:
|
||||||
"""Verify that local endpoints get keepalive_expiry=0."""
|
"""Verify that local endpoints get keepalive_expiry=0."""
|
||||||
|
|
||||||
async def test_local_spec_disables_keepalive(self):
|
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
|
||||||
|
|
||||||
async def test_lan_ip_disables_keepalive(self):
|
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 = ""
|
||||||
@@ -104,18 +103,16 @@ 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
|
||||||
|
|
||||||
async def test_cloud_keeps_default_keepalive(self):
|
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
|
||||||
|
|||||||
@@ -1,97 +0,0 @@
|
|||||||
"""Tests for the Novita AI 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_novita_config_field_exists() -> None:
|
|
||||||
config = ProvidersConfig()
|
|
||||||
|
|
||||||
assert hasattr(config, "novita")
|
|
||||||
|
|
||||||
|
|
||||||
def test_novita_provider_in_registry() -> None:
|
|
||||||
specs = {spec.name: spec for spec in PROVIDERS}
|
|
||||||
|
|
||||||
assert "novita" in specs
|
|
||||||
novita = specs["novita"]
|
|
||||||
assert novita.backend == "openai_compat"
|
|
||||||
assert novita.env_key == "NOVITA_API_KEY"
|
|
||||||
assert novita.display_name == "Novita AI"
|
|
||||||
assert novita.is_gateway is True
|
|
||||||
assert novita.detect_by_base_keyword == "novita"
|
|
||||||
assert novita.default_api_base == "https://api.novita.ai/openai"
|
|
||||||
assert novita.strip_model_prefix is False
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_by_name_novita() -> None:
|
|
||||||
spec = find_by_name("novita")
|
|
||||||
|
|
||||||
assert spec is not None
|
|
||||||
assert spec.name == "novita"
|
|
||||||
|
|
||||||
|
|
||||||
def test_novita_forced_provider_uses_default_api_base() -> None:
|
|
||||||
config = Config.model_validate({
|
|
||||||
"providers": {
|
|
||||||
"novita": {
|
|
||||||
"apiKey": "novita-key",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"agents": {
|
|
||||||
"defaults": {
|
|
||||||
"model": "deepseek-v4-pro",
|
|
||||||
"provider": "novita",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
assert config.get_provider_name("deepseek-v4-pro") == "novita"
|
|
||||||
assert config.get_api_key("deepseek-v4-pro") == "novita-key"
|
|
||||||
assert config.get_api_base("deepseek-v4-pro") == "https://api.novita.ai/openai"
|
|
||||||
|
|
||||||
|
|
||||||
def test_novita_gateway_routes_unprefixed_models_when_configured() -> None:
|
|
||||||
config = Config.model_validate({
|
|
||||||
"providers": {
|
|
||||||
"novita": {
|
|
||||||
"apiKey": "novita-key",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"agents": {
|
|
||||||
"defaults": {
|
|
||||||
"model": "deepseek-v4-pro",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
assert config.get_provider_name("deepseek-v4-pro") == "novita"
|
|
||||||
assert config.get_api_key("deepseek-v4-pro") == "novita-key"
|
|
||||||
assert config.get_api_base("deepseek-v4-pro") == "https://api.novita.ai/openai"
|
|
||||||
|
|
||||||
|
|
||||||
def test_novita_preserves_model_api_id() -> None:
|
|
||||||
spec = find_by_name("novita")
|
|
||||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
|
||||||
provider = OpenAICompatProvider(
|
|
||||||
api_key="novita-key",
|
|
||||||
default_model="deepseek-v4-pro",
|
|
||||||
spec=spec,
|
|
||||||
)
|
|
||||||
|
|
||||||
kwargs = provider._build_kwargs(
|
|
||||||
messages=[{"role": "user", "content": "hi"}],
|
|
||||||
tools=None,
|
|
||||||
model="deepseek-v4-pro",
|
|
||||||
max_tokens=1024,
|
|
||||||
temperature=0.7,
|
|
||||||
reasoning_effort=None,
|
|
||||||
tool_choice=None,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert kwargs["model"] == "deepseek-v4-pro"
|
|
||||||
assert kwargs["max_tokens"] == 1024
|
|
||||||
assert "max_completion_tokens" not in kwargs
|
|
||||||
@@ -16,15 +16,7 @@ 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(
|
async def fake_request(url, headers, body, verify, on_content_delta=None):
|
||||||
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,18 +8,16 @@ def _assert_openai_compat_timeout(timeout) -> None:
|
|||||||
assert timeout == 120.0
|
assert timeout == 120.0
|
||||||
|
|
||||||
|
|
||||||
async def test_openai_compat_provider_defers_sdk_client_until_first_use() -> None:
|
def test_openai_compat_provider_sets_sdk_timeout() -> 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:
|
||||||
provider = OpenAICompatProvider(api_key="test-key", api_base="https://example.com/v1")
|
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
|
||||||
|
|
||||||
|
|
||||||
async def test_openai_compat_provider_sets_timeout_on_local_http_client() -> None:
|
def test_openai_compat_provider_sets_timeout_on_local_http_client() -> None:
|
||||||
spec = ProviderSpec(
|
spec = ProviderSpec(
|
||||||
name="local",
|
name="local",
|
||||||
keywords=(),
|
keywords=(),
|
||||||
@@ -31,13 +29,11 @@ async def test_openai_compat_provider_sets_timeout_on_local_http_client() -> Non
|
|||||||
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(
|
||||||
"httpx.AsyncClient",
|
"nanobot.providers.openai_compat_provider.httpx.AsyncClient",
|
||||||
return_value=sentinel.http_client,
|
return_value=sentinel.http_client,
|
||||||
) as mock_http_client,
|
) as mock_http_client,
|
||||||
):
|
):
|
||||||
provider = OpenAICompatProvider(spec=spec)
|
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"])
|
||||||
@@ -48,11 +44,10 @@ async def test_openai_compat_provider_sets_timeout_on_local_http_client() -> Non
|
|||||||
assert openai_kwargs["http_client"] is sentinel.http_client
|
assert openai_kwargs["http_client"] is sentinel.http_client
|
||||||
|
|
||||||
|
|
||||||
async def test_openai_compat_provider_timeout_can_be_overridden_by_env(monkeypatch) -> None:
|
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:
|
||||||
provider = OpenAICompatProvider(api_key="test-key", api_base="https://example.com/v1")
|
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
|
||||||
|
|||||||
@@ -155,49 +155,6 @@ class TestConvertMessages:
|
|||||||
assert items[0]["id"] == "fc_1"
|
assert items[0]["id"] == "fc_1"
|
||||||
assert items[0]["name"] == "get_weather"
|
assert items[0]["name"] == "get_weather"
|
||||||
|
|
||||||
def test_duplicate_response_item_ids_are_made_unique(self):
|
|
||||||
"""Codex rejects replayed Responses input items with duplicate ids."""
|
|
||||||
_, items = convert_messages([
|
|
||||||
{
|
|
||||||
"role": "assistant",
|
|
||||||
"content": None,
|
|
||||||
"tool_calls": [{
|
|
||||||
"id": "call_a|rs_same",
|
|
||||||
"function": {"name": "first", "arguments": "{}"},
|
|
||||||
}],
|
|
||||||
},
|
|
||||||
{"role": "tool", "tool_call_id": "call_a|rs_same", "content": "ok"},
|
|
||||||
{
|
|
||||||
"role": "assistant",
|
|
||||||
"content": None,
|
|
||||||
"tool_calls": [{
|
|
||||||
"id": "call_b|rs_same",
|
|
||||||
"function": {"name": "second", "arguments": "{}"},
|
|
||||||
}],
|
|
||||||
},
|
|
||||||
{"role": "tool", "tool_call_id": "call_b|rs_same", "content": "ok"},
|
|
||||||
])
|
|
||||||
function_call_ids = [
|
|
||||||
item["id"] for item in items if item.get("type") == "function_call"
|
|
||||||
]
|
|
||||||
assert function_call_ids == ["rs_same", "rs_same_2"]
|
|
||||||
assert len(function_call_ids) == len(set(function_call_ids))
|
|
||||||
|
|
||||||
def test_fallback_response_item_ids_are_unique_with_multiple_tool_calls(self):
|
|
||||||
_, items = convert_messages([{
|
|
||||||
"role": "assistant",
|
|
||||||
"content": None,
|
|
||||||
"tool_calls": [
|
|
||||||
{"id": "call_a", "function": {"name": "first", "arguments": "{}"}},
|
|
||||||
{"id": "call_b", "function": {"name": "second", "arguments": "{}"}},
|
|
||||||
],
|
|
||||||
}])
|
|
||||||
function_call_ids = [
|
|
||||||
item["id"] for item in items if item.get("type") == "function_call"
|
|
||||||
]
|
|
||||||
assert function_call_ids == ["fc_0", "fc_0_2"]
|
|
||||||
assert len(function_call_ids) == len(set(function_call_ids))
|
|
||||||
|
|
||||||
def test_assistant_with_tool_calls_no_id(self):
|
def test_assistant_with_tool_calls_no_id(self):
|
||||||
"""Fallback IDs when tool_call.id is missing."""
|
"""Fallback IDs when tool_call.id is missing."""
|
||||||
_, items = convert_messages([{
|
_, items = convert_messages([{
|
||||||
@@ -496,56 +453,6 @@ 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,10 +5,9 @@ from nanobot.providers.azure_openai_provider import AzureOpenAIProvider
|
|||||||
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
|
|
||||||
|
|
||||||
async def test_openai_compat_disables_sdk_retries_by_default() -> None:
|
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:
|
||||||
provider = OpenAICompatProvider(api_key="sk-test", default_model="gpt-4o")
|
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
|
||||||
|
|||||||
@@ -1,80 +0,0 @@
|
|||||||
"""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
|
|
||||||
@@ -32,7 +32,7 @@ def _mimo_spec():
|
|||||||
|
|
||||||
|
|
||||||
def _openrouter_spec():
|
def _openrouter_spec():
|
||||||
"""Return the registered OpenRouter ProviderSpec."""
|
"""Return the registered OpenRouter ProviderSpec (no thinking_style)."""
|
||||||
specs = {s.name: s for s in PROVIDERS}
|
specs = {s.name: s for s in PROVIDERS}
|
||||||
return specs["openrouter"]
|
return specs["openrouter"]
|
||||||
|
|
||||||
@@ -77,13 +77,6 @@ def test_xiaomi_mimo_uses_thinking_type_style():
|
|||||||
assert spec.default_api_base == "https://api.xiaomimimo.com/v1"
|
assert spec.default_api_base == "https://api.xiaomimimo.com/v1"
|
||||||
|
|
||||||
|
|
||||||
def test_openrouter_declares_gateway_reasoning_style():
|
|
||||||
"""OpenRouter uses its own reasoning.effort field for routed thinking models."""
|
|
||||||
spec = _openrouter_spec()
|
|
||||||
assert spec.thinking_style == ""
|
|
||||||
assert spec.gateway_reasoning_style == "reasoning_effort"
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# _build_kwargs wire-format
|
# _build_kwargs wire-format
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -149,11 +142,9 @@ def test_mimo_reasoning_effort_unset_preserves_provider_default():
|
|||||||
|
|
||||||
|
|
||||||
def test_mimo_via_openrouter_reasoning_effort_none_disables_thinking():
|
def test_mimo_via_openrouter_reasoning_effort_none_disables_thinking():
|
||||||
"""OpenRouter routes MiMo as "xiaomi/mimo-v2.5-pro" and does NOT forward
|
"""OpenRouter routes MiMo as "xiaomi/mimo-v2.5-pro"; the openrouter spec
|
||||||
extra_body.thinking to upstream, so a disable signal must also reach OR
|
has no thinking_style, so the disable signal must come from the
|
||||||
in its own `reasoning.effort` shape. Verifies both the upstream-MiMo
|
model-name path (#3845)."""
|
||||||
payload (#3845) and the OR-native payload (#3851 follow-up) are sent.
|
|
||||||
"""
|
|
||||||
provider = _openrouter_provider("xiaomi/mimo-v2.5-pro")
|
provider = _openrouter_provider("xiaomi/mimo-v2.5-pro")
|
||||||
kwargs = provider._build_kwargs(
|
kwargs = provider._build_kwargs(
|
||||||
messages=_simple_messages(),
|
messages=_simple_messages(),
|
||||||
@@ -161,15 +152,11 @@ def test_mimo_via_openrouter_reasoning_effort_none_disables_thinking():
|
|||||||
temperature=0.7, reasoning_effort="none", tool_choice=None,
|
temperature=0.7, reasoning_effort="none", tool_choice=None,
|
||||||
)
|
)
|
||||||
assert "reasoning_effort" not in kwargs
|
assert "reasoning_effort" not in kwargs
|
||||||
assert kwargs["extra_body"] == {
|
assert kwargs["extra_body"] == {"thinking": {"type": "disabled"}}
|
||||||
"thinking": {"type": "disabled"},
|
|
||||||
"reasoning": {"effort": "none"},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def test_mimo_via_openrouter_reasoning_effort_medium_enables_thinking():
|
def test_mimo_via_openrouter_reasoning_effort_medium_enables_thinking():
|
||||||
"""Non-none/minimal effort enables thinking and the OR `reasoning.effort`
|
"""Same as the direct path: any non-none/minimal effort enables thinking."""
|
||||||
field mirrors the requested effort level."""
|
|
||||||
provider = _openrouter_provider("xiaomi/mimo-v2.5-pro")
|
provider = _openrouter_provider("xiaomi/mimo-v2.5-pro")
|
||||||
kwargs = provider._build_kwargs(
|
kwargs = provider._build_kwargs(
|
||||||
messages=_simple_messages(),
|
messages=_simple_messages(),
|
||||||
@@ -177,10 +164,7 @@ def test_mimo_via_openrouter_reasoning_effort_medium_enables_thinking():
|
|||||||
temperature=0.7, reasoning_effort="medium", tool_choice=None,
|
temperature=0.7, reasoning_effort="medium", tool_choice=None,
|
||||||
)
|
)
|
||||||
assert kwargs.get("reasoning_effort") == "medium"
|
assert kwargs.get("reasoning_effort") == "medium"
|
||||||
assert kwargs["extra_body"] == {
|
assert kwargs["extra_body"] == {"thinking": {"type": "enabled"}}
|
||||||
"thinking": {"type": "enabled"},
|
|
||||||
"reasoning": {"effort": "medium"},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def test_mimo_via_openrouter_bare_slug_also_matches():
|
def test_mimo_via_openrouter_bare_slug_also_matches():
|
||||||
@@ -192,16 +176,12 @@ def test_mimo_via_openrouter_bare_slug_also_matches():
|
|||||||
tools=None, model=None, max_tokens=100,
|
tools=None, model=None, max_tokens=100,
|
||||||
temperature=0.7, reasoning_effort="none", tool_choice=None,
|
temperature=0.7, reasoning_effort="none", tool_choice=None,
|
||||||
)
|
)
|
||||||
assert kwargs["extra_body"] == {
|
assert kwargs["extra_body"] == {"thinking": {"type": "disabled"}}
|
||||||
"thinking": {"type": "disabled"},
|
|
||||||
"reasoning": {"effort": "none"},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def test_mimo_flash_via_openrouter_does_not_inject_thinking():
|
def test_mimo_flash_via_openrouter_does_not_inject_thinking():
|
||||||
"""mimo-v2-flash has no thinking mode per Xiaomi docs; the allowlist
|
"""mimo-v2-flash has no thinking mode per Xiaomi docs; the allowlist
|
||||||
excludes it, so neither the upstream `thinking` field nor OR's
|
excludes it, so no thinking field should be injected on the gateway path."""
|
||||||
`reasoning.effort` should be injected on the gateway path."""
|
|
||||||
provider = _openrouter_provider("xiaomi/mimo-v2-flash")
|
provider = _openrouter_provider("xiaomi/mimo-v2-flash")
|
||||||
kwargs = provider._build_kwargs(
|
kwargs = provider._build_kwargs(
|
||||||
messages=_simple_messages(),
|
messages=_simple_messages(),
|
||||||
@@ -220,18 +200,3 @@ def test_non_mimo_model_via_openrouter_unaffected():
|
|||||||
temperature=0.7, reasoning_effort="none", tool_choice=None,
|
temperature=0.7, reasoning_effort="none", tool_choice=None,
|
||||||
)
|
)
|
||||||
assert "extra_body" not in kwargs
|
assert "extra_body" not in kwargs
|
||||||
|
|
||||||
|
|
||||||
def test_kimi_via_openrouter_also_injects_reasoning_effort():
|
|
||||||
"""Kimi has the same gateway problem as MiMo: OR drops the upstream
|
|
||||||
`thinking` field. The same OR-reasoning injection should fire."""
|
|
||||||
provider = _openrouter_provider("moonshotai/kimi-k2.5")
|
|
||||||
kwargs = provider._build_kwargs(
|
|
||||||
messages=_simple_messages(),
|
|
||||||
tools=None, model=None, max_tokens=100,
|
|
||||||
temperature=0.7, reasoning_effort="none", tool_choice=None,
|
|
||||||
)
|
|
||||||
assert kwargs["extra_body"] == {
|
|
||||||
"thinking": {"type": "disabled"},
|
|
||||||
"reasoning": {"effort": "none"},
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,330 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
|
|
||||||
from nanobot.agent.tools.apply_patch import ApplyPatchTool
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_replace(tmp_path):
|
|
||||||
target = tmp_path / "calc.py"
|
|
||||||
target.write_text("def add(a, b):\n return a + b\n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "calc.py",
|
|
||||||
"action": "replace",
|
|
||||||
"old_text": " return a + b",
|
|
||||||
"new_text": " return a - b",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "update calc.py" in result
|
|
||||||
assert target.read_text() == "def add(a, b):\n return a - b\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_add_new_file(tmp_path):
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "config.py",
|
|
||||||
"action": "add",
|
|
||||||
"new_text": "DEBUG = True",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "add config.py" in result
|
|
||||||
assert (tmp_path / "config.py").read_text() == "DEBUG = True\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_preserves_new_file_trailing_blank_lines(tmp_path):
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "notes.txt",
|
|
||||||
"action": "add",
|
|
||||||
"new_text": "one\n\n",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "add notes.txt" in result
|
|
||||||
assert (tmp_path / "notes.txt").read_text() == "one\n\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_add_to_existing_file(tmp_path):
|
|
||||||
target = tmp_path / "log.py"
|
|
||||||
target.write_text("import logging\n\nlogger = logging.getLogger(__name__)\n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "log.py",
|
|
||||||
"action": "add",
|
|
||||||
"new_text": "def debug(msg):\n logger.debug(msg)",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "update log.py" in result
|
|
||||||
assert (
|
|
||||||
target.read_text()
|
|
||||||
== "import logging\n\nlogger = logging.getLogger(__name__)\ndef debug(msg):\n logger.debug(msg)\n"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_delete(tmp_path):
|
|
||||||
target = tmp_path / "utils.py"
|
|
||||||
target.write_text("def unused():\n pass\ndef used():\n return 1\n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "utils.py",
|
|
||||||
"action": "delete",
|
|
||||||
"old_text": "def unused():\n pass\n",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "update utils.py" in result
|
|
||||||
assert target.read_text() == "def used():\n return 1\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_delete_entire_file(tmp_path):
|
|
||||||
target = tmp_path / "obsolete.txt"
|
|
||||||
target.write_text("remove me\n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "obsolete.txt",
|
|
||||||
"action": "delete",
|
|
||||||
"old_text": "remove me\n",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "delete obsolete.txt" in result
|
|
||||||
assert not target.exists()
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_delete_substring_with_surrounding_whitespace(tmp_path):
|
|
||||||
target = tmp_path / "keep_whitespace.txt"
|
|
||||||
target.write_text(" token \n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "keep_whitespace.txt",
|
|
||||||
"action": "delete",
|
|
||||||
"old_text": "token",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "update keep_whitespace.txt" in result
|
|
||||||
assert target.exists()
|
|
||||||
assert target.read_text() == " \n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_batch_multiple_files(tmp_path):
|
|
||||||
a = tmp_path / "a.py"
|
|
||||||
a.write_text("X = 1\n")
|
|
||||||
b = tmp_path / "b.py"
|
|
||||||
b.write_text("from a import X\nprint(X)\n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "a.py",
|
|
||||||
"action": "replace",
|
|
||||||
"old_text": "X = 1",
|
|
||||||
"new_text": "Y = 1",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"path": "b.py",
|
|
||||||
"action": "replace",
|
|
||||||
"old_text": "from a import X",
|
|
||||||
"new_text": "from a import Y",
|
|
||||||
},
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "update a.py" in result
|
|
||||||
assert "update b.py" in result
|
|
||||||
assert a.read_text() == "Y = 1\n"
|
|
||||||
assert b.read_text() == "from a import Y\nprint(X)\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_rejects_ambiguous_old_text(tmp_path):
|
|
||||||
target = tmp_path / "repeated.txt"
|
|
||||||
target.write_text("target\nmiddle\ntarget\n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "repeated.txt",
|
|
||||||
"action": "replace",
|
|
||||||
"old_text": "target",
|
|
||||||
"new_text": "changed",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "old_text appears multiple times" in result
|
|
||||||
assert target.read_text() == "target\nmiddle\ntarget\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_dry_run_validates_without_writing(tmp_path):
|
|
||||||
target = tmp_path / "dry.txt"
|
|
||||||
target.write_text("before\n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "dry.txt",
|
|
||||||
"action": "replace",
|
|
||||||
"old_text": "before",
|
|
||||||
"new_text": "after",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"path": "added.txt",
|
|
||||||
"action": "add",
|
|
||||||
"new_text": "new",
|
|
||||||
},
|
|
||||||
],
|
|
||||||
dry_run=True,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "Patch dry-run succeeded" in result
|
|
||||||
assert target.read_text() == "before\n"
|
|
||||||
assert not (tmp_path / "added.txt").exists()
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_rejects_absolute_and_parent_paths(tmp_path):
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
absolute = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "/tmp/owned.txt",
|
|
||||||
"action": "add",
|
|
||||||
"new_text": "nope",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
parent = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "../owned.txt",
|
|
||||||
"action": "add",
|
|
||||||
"new_text": "nope",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
windows_absolute = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": r"C:\owned.txt",
|
|
||||||
"action": "add",
|
|
||||||
"new_text": "nope",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
windows_parent = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": r"..\owned.txt",
|
|
||||||
"action": "add",
|
|
||||||
"new_text": "nope",
|
|
||||||
}
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "must be relative" in absolute
|
|
||||||
assert "must not contain '..'" in parent
|
|
||||||
assert "must be relative" in windows_absolute
|
|
||||||
assert "must not contain '..'" in windows_parent
|
|
||||||
assert not (tmp_path.parent / "owned.txt").exists()
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_reports_invalid_edit_shapes(tmp_path):
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
missing_path = asyncio.run(tool.execute(edits=[{"action": "add", "new_text": "x"}]))
|
|
||||||
missing_action = asyncio.run(tool.execute(edits=[{"path": "x.txt", "new_text": "x"}]))
|
|
||||||
non_object = asyncio.run(tool.execute(edits=["not an object"])) # type: ignore[list-item]
|
|
||||||
|
|
||||||
assert "path required for edit" in missing_path
|
|
||||||
assert "action required for edit: x.txt" in missing_action
|
|
||||||
assert "each edit must be an object" in non_object
|
|
||||||
|
|
||||||
|
|
||||||
def test_apply_patch_edits_rolls_back_when_late_operation_fails(tmp_path):
|
|
||||||
first = tmp_path / "first.txt"
|
|
||||||
first.write_text("before\n")
|
|
||||||
tool = ApplyPatchTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
tool.execute(
|
|
||||||
edits=[
|
|
||||||
{
|
|
||||||
"path": "first.txt",
|
|
||||||
"action": "replace",
|
|
||||||
"old_text": "before",
|
|
||||||
"new_text": "after",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"path": "missing.txt",
|
|
||||||
"action": "delete",
|
|
||||||
"old_text": "remove me",
|
|
||||||
},
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "file to update does not exist: missing.txt" in result
|
|
||||||
assert first.read_text() == "before\n"
|
|
||||||
@@ -1,5 +1,5 @@
|
|||||||
"""Tests for EditFileTool enhancements: read-before-edit tracking, path suggestions,
|
"""Tests for EditFileTool enhancements: read-before-edit tracking, path suggestions,
|
||||||
notebook JSON editing, and create-file semantics."""
|
.ipynb detection, and create-file semantics."""
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
@@ -108,27 +108,22 @@ class TestEditCreateFile:
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# .ipynb editing
|
# .ipynb detection
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
class TestEditIpynbFiles:
|
class TestEditIpynbDetection:
|
||||||
"""edit_file edits notebooks as normal JSON files."""
|
"""edit_file should refuse .ipynb and suggest notebook_edit."""
|
||||||
|
|
||||||
@pytest.fixture()
|
@pytest.fixture()
|
||||||
def tool(self, tmp_path):
|
def tool(self, tmp_path):
|
||||||
return EditFileTool(workspace=tmp_path)
|
return EditFileTool(workspace=tmp_path)
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_ipynb_can_be_edited_as_json(self, tool, tmp_path):
|
async def test_ipynb_rejected_with_suggestion(self, tool, tmp_path):
|
||||||
f = tmp_path / "analysis.ipynb"
|
f = tmp_path / "analysis.ipynb"
|
||||||
f.write_text('{"cells": []}', encoding="utf-8")
|
f.write_text('{"cells": []}', encoding="utf-8")
|
||||||
result = await tool.execute(
|
result = await tool.execute(path=str(f), old_text="x", new_text="y")
|
||||||
path=str(f),
|
assert "notebook" in result.lower()
|
||||||
old_text='"cells": []',
|
|
||||||
new_text='"cells": [{"cell_type": "markdown", "source": "hi"}]',
|
|
||||||
)
|
|
||||||
assert "Successfully edited" in result
|
|
||||||
assert '"source": "hi"' in f.read_text(encoding="utf-8")
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ strategy, and sandbox behaviour per platform — without actually running
|
|||||||
platform-specific binaries (all subprocess calls are mocked).
|
platform-specific binaries (all subprocess calls are mocked).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import sys
|
import sys
|
||||||
from unittest.mock import AsyncMock, patch
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
@@ -109,9 +108,6 @@ class TestSpawnUnix:
|
|||||||
assert "-c" in args
|
assert "-c" in args
|
||||||
assert "echo hi" in args
|
assert "echo hi" in args
|
||||||
|
|
||||||
kwargs = mock_exec.call_args[1]
|
|
||||||
assert kwargs["stdin"] == asyncio.subprocess.DEVNULL
|
|
||||||
|
|
||||||
|
|
||||||
class TestSpawnWindows:
|
class TestSpawnWindows:
|
||||||
|
|
||||||
@@ -128,9 +124,6 @@ class TestSpawnWindows:
|
|||||||
args = mock_shell.call_args[0]
|
args = mock_shell.call_args[0]
|
||||||
assert "dir" in args
|
assert "dir" in args
|
||||||
|
|
||||||
kwargs = mock_shell.call_args[1]
|
|
||||||
assert kwargs["stdin"] == asyncio.subprocess.DEVNULL
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_passes_cwd_and_env(self):
|
async def test_passes_cwd_and_env(self):
|
||||||
env = {"PATH": "/usr/bin"}
|
env = {"PATH": "/usr/bin"}
|
||||||
@@ -162,7 +155,7 @@ class TestPathAppendPlatform:
|
|||||||
captured_cmd = None
|
captured_cmd = None
|
||||||
captured_env = {}
|
captured_env = {}
|
||||||
|
|
||||||
async def capture_spawn(cmd, cwd, env, shell_program=None, login=True):
|
async def capture_spawn(cmd, cwd, env):
|
||||||
nonlocal captured_cmd
|
nonlocal captured_cmd
|
||||||
captured_cmd = cmd
|
captured_cmd = cmd
|
||||||
captured_env.update(env)
|
captured_env.update(env)
|
||||||
@@ -190,7 +183,7 @@ class TestPathAppendPlatform:
|
|||||||
|
|
||||||
captured_env = {}
|
captured_env = {}
|
||||||
|
|
||||||
async def capture_spawn(cmd, cwd, env, shell_program=None, login=True):
|
async def capture_spawn(cmd, cwd, env):
|
||||||
captured_env.update(env)
|
captured_env.update(env)
|
||||||
return mock_proc
|
return mock_proc
|
||||||
|
|
||||||
|
|||||||
@@ -1,361 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import re
|
|
||||||
import shlex
|
|
||||||
import subprocess
|
|
||||||
import sys
|
|
||||||
|
|
||||||
from nanobot.agent.tools.shell import ExecTool
|
|
||||||
from nanobot.agent.tools.exec_session import ExecSessionManager, ListExecSessionsTool, WriteStdinTool
|
|
||||||
|
|
||||||
|
|
||||||
def _python_command(code: str) -> str:
|
|
||||||
if sys.platform == "win32":
|
|
||||||
return f"{subprocess.list2cmdline([sys.executable])} -u -c {subprocess.list2cmdline([code])}"
|
|
||||||
return f"{shlex.quote(sys.executable)} -u -c {shlex.quote(code)}"
|
|
||||||
|
|
||||||
|
|
||||||
def _session_id(output: str) -> str:
|
|
||||||
match = re.search(r"session_id:\s*([0-9a-f]+)", output)
|
|
||||||
assert match, output
|
|
||||||
return match.group(1)
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_keeps_one_shot_behavior_without_yield_time_ms(tmp_path):
|
|
||||||
async def run() -> str:
|
|
||||||
tool = ExecTool(working_dir=str(tmp_path), timeout=5)
|
|
||||||
return await tool.execute(command="echo hello")
|
|
||||||
|
|
||||||
result = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "hello" in result
|
|
||||||
assert "Exit code: 0" in result
|
|
||||||
assert "session_id:" not in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_accepts_command_aliases(tmp_path):
|
|
||||||
async def run() -> str:
|
|
||||||
tool = ExecTool(working_dir="/")
|
|
||||||
return await tool.execute(
|
|
||||||
cmd=_python_command("import os; print(os.getcwd())"),
|
|
||||||
workdir=str(tmp_path),
|
|
||||||
)
|
|
||||||
|
|
||||||
result = asyncio.run(run())
|
|
||||||
|
|
||||||
assert str(tmp_path) in result
|
|
||||||
assert "Exit code: 0" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_returns_completed_session_output_when_yield_time_ms_is_used(tmp_path):
|
|
||||||
async def run() -> str:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
|
|
||||||
result = await tool.execute(command="echo hello", yield_time_ms=1000)
|
|
||||||
if "session_id:" in result:
|
|
||||||
sid = _session_id(result)
|
|
||||||
result += "\n" + await stdin_tool.execute(
|
|
||||||
session_id=sid,
|
|
||||||
chars="",
|
|
||||||
yield_time_ms=1000,
|
|
||||||
)
|
|
||||||
return result
|
|
||||||
|
|
||||||
result = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "hello" in result
|
|
||||||
assert "Exit code: 0" in result
|
|
||||||
assert "session_id:" not in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_session_accepts_max_output_tokens_alias(tmp_path):
|
|
||||||
async def run() -> str:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
command = _python_command("print('A' * 2000)")
|
|
||||||
return await tool.execute(
|
|
||||||
command=command,
|
|
||||||
yield_time_ms=1000,
|
|
||||||
max_output_tokens=1000,
|
|
||||||
)
|
|
||||||
|
|
||||||
result = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "chars truncated" in result
|
|
||||||
assert "Exit code: 0" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_one_shot_accepts_max_output_tokens_alias(tmp_path):
|
|
||||||
async def run() -> str:
|
|
||||||
tool = ExecTool(working_dir=str(tmp_path), timeout=5)
|
|
||||||
command = _python_command("print('A' * 2000)")
|
|
||||||
return await tool.execute(command=command, max_output_tokens=1000)
|
|
||||||
|
|
||||||
result = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "chars truncated" in result
|
|
||||||
assert "Exit code: 0" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_accepts_supported_shell_parameter(tmp_path):
|
|
||||||
async def run() -> str:
|
|
||||||
tool = ExecTool(working_dir=str(tmp_path), timeout=5)
|
|
||||||
return await tool.execute(command="echo shell-ok", shell="sh", login=False)
|
|
||||||
|
|
||||||
if sys.platform == "win32":
|
|
||||||
return
|
|
||||||
result = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "shell-ok" in result
|
|
||||||
assert "Exit code: 0" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_rejects_unsupported_shell(tmp_path):
|
|
||||||
async def run() -> str:
|
|
||||||
tool = ExecTool(working_dir=str(tmp_path), timeout=5)
|
|
||||||
return await tool.execute(command="echo no", shell="python")
|
|
||||||
|
|
||||||
if sys.platform == "win32":
|
|
||||||
return
|
|
||||||
result = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "unsupported shell" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_can_continue_with_stdin(tmp_path):
|
|
||||||
async def run() -> tuple[str, str]:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
command = _python_command(
|
|
||||||
"import sys; print('ready', flush=True); "
|
|
||||||
"line=sys.stdin.readline(); print('got:' + line.strip(), flush=True)"
|
|
||||||
)
|
|
||||||
|
|
||||||
initial = await exec_tool.execute(command=command, yield_time_ms=500)
|
|
||||||
sid = _session_id(initial)
|
|
||||||
result = await stdin_tool.execute(session_id=sid, chars="ping\n", yield_time_ms=1000)
|
|
||||||
return initial, result
|
|
||||||
|
|
||||||
initial, result = asyncio.run(run())
|
|
||||||
assert "ready" in initial
|
|
||||||
assert "Process running" in initial
|
|
||||||
assert "Elapsed:" in initial
|
|
||||||
assert "got:ping" in result
|
|
||||||
assert "Exit code: 0" in result
|
|
||||||
assert "Elapsed:" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_write_stdin_can_close_stdin(tmp_path):
|
|
||||||
async def run() -> tuple[str, str]:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
command = _python_command(
|
|
||||||
"import sys; print('ready', flush=True); "
|
|
||||||
"data=sys.stdin.read(); print('got:' + data, flush=True)"
|
|
||||||
)
|
|
||||||
|
|
||||||
initial = await exec_tool.execute(command=command, yield_time_ms=500)
|
|
||||||
sid = _session_id(initial)
|
|
||||||
result = await stdin_tool.execute(
|
|
||||||
session_id=sid,
|
|
||||||
chars="payload",
|
|
||||||
close_stdin=True,
|
|
||||||
yield_time_ms=1000,
|
|
||||||
)
|
|
||||||
return initial, result
|
|
||||||
|
|
||||||
initial, result = asyncio.run(run())
|
|
||||||
assert "ready" in initial
|
|
||||||
assert "got:payload" in result
|
|
||||||
assert "Stdin closed." in result
|
|
||||||
assert "Exit code: 0" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_write_stdin_can_terminate_session(tmp_path):
|
|
||||||
async def run() -> tuple[str, str]:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
exec_tool = ExecTool(working_dir=str(tmp_path), timeout=30, session_manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
command = _python_command(
|
|
||||||
"import time; print('ready', flush=True); time.sleep(30)"
|
|
||||||
)
|
|
||||||
|
|
||||||
initial = await exec_tool.execute(command=command, yield_time_ms=500)
|
|
||||||
sid = _session_id(initial)
|
|
||||||
result = await stdin_tool.execute(
|
|
||||||
session_id=sid,
|
|
||||||
terminate=True,
|
|
||||||
yield_time_ms=0,
|
|
||||||
)
|
|
||||||
return initial, result
|
|
||||||
|
|
||||||
initial, result = asyncio.run(run())
|
|
||||||
assert "ready" in initial
|
|
||||||
assert "Session terminated." in result
|
|
||||||
assert "Exit code:" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_write_stdin_accepts_max_output_tokens_alias(tmp_path):
|
|
||||||
async def run() -> tuple[str, str, str]:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
command = _python_command(
|
|
||||||
"import time; print('A' * 2000, flush=True); time.sleep(5)"
|
|
||||||
)
|
|
||||||
|
|
||||||
initial = await exec_tool.execute(command=command, yield_time_ms=0)
|
|
||||||
sid = _session_id(initial)
|
|
||||||
poll = await stdin_tool.execute(
|
|
||||||
session_id=sid,
|
|
||||||
yield_time_ms=500,
|
|
||||||
max_output_tokens=1000,
|
|
||||||
)
|
|
||||||
cleanup = await stdin_tool.execute(session_id=sid, terminate=True, yield_time_ms=0)
|
|
||||||
return initial, poll, cleanup
|
|
||||||
|
|
||||||
initial, poll, cleanup = asyncio.run(run())
|
|
||||||
assert "Process running" in initial
|
|
||||||
assert "chars truncated" in poll
|
|
||||||
assert "Session terminated." in cleanup
|
|
||||||
|
|
||||||
|
|
||||||
def test_write_stdin_preserves_completed_session_output_until_polled(tmp_path):
|
|
||||||
async def run() -> tuple[str, str]:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
command = _python_command(
|
|
||||||
"import time; print('ready', flush=True); "
|
|
||||||
"time.sleep(1.0); print('done', flush=True)"
|
|
||||||
)
|
|
||||||
|
|
||||||
initial = await exec_tool.execute(command=command, yield_time_ms=300)
|
|
||||||
sid = _session_id(initial)
|
|
||||||
await asyncio.sleep(1.2)
|
|
||||||
final = await stdin_tool.execute(session_id=sid, chars="", yield_time_ms=0)
|
|
||||||
return initial, final
|
|
||||||
|
|
||||||
initial, final = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "ready" in initial
|
|
||||||
assert "done" in final
|
|
||||||
assert "Exit code: 0" in final
|
|
||||||
|
|
||||||
|
|
||||||
def test_write_stdin_can_wait_for_expected_output(tmp_path):
|
|
||||||
async def run() -> tuple[str, str, str]:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
command = _python_command(
|
|
||||||
"import time; print('booting', flush=True); "
|
|
||||||
"time.sleep(0.4); print('ready', flush=True); time.sleep(5)"
|
|
||||||
)
|
|
||||||
|
|
||||||
initial = await exec_tool.execute(command=command, yield_time_ms=100)
|
|
||||||
sid = _session_id(initial)
|
|
||||||
waited = await stdin_tool.execute(
|
|
||||||
session_id=sid,
|
|
||||||
wait_for="ready",
|
|
||||||
wait_timeout_ms=3000,
|
|
||||||
yield_time_ms=0,
|
|
||||||
)
|
|
||||||
cleanup = await stdin_tool.execute(session_id=sid, terminate=True, yield_time_ms=0)
|
|
||||||
return initial, waited, cleanup
|
|
||||||
|
|
||||||
initial, waited, cleanup = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "Process running" in initial
|
|
||||||
assert "booting" in initial + waited
|
|
||||||
assert "ready" in waited
|
|
||||||
assert "Wait target not observed" not in waited
|
|
||||||
assert "Session terminated." in cleanup
|
|
||||||
|
|
||||||
|
|
||||||
def test_write_stdin_wait_for_reports_timeout_without_killing_session(tmp_path):
|
|
||||||
async def run() -> tuple[str, str, str]:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
command = _python_command(
|
|
||||||
"import time; print('booting', flush=True); time.sleep(5)"
|
|
||||||
)
|
|
||||||
|
|
||||||
initial = await exec_tool.execute(command=command, yield_time_ms=100)
|
|
||||||
sid = _session_id(initial)
|
|
||||||
waited = await stdin_tool.execute(
|
|
||||||
session_id=sid,
|
|
||||||
wait_for="never-ready",
|
|
||||||
wait_timeout_ms=200,
|
|
||||||
yield_time_ms=0,
|
|
||||||
)
|
|
||||||
cleanup = await stdin_tool.execute(session_id=sid, terminate=True, yield_time_ms=0)
|
|
||||||
return initial, waited, cleanup
|
|
||||||
|
|
||||||
initial, waited, cleanup = asyncio.run(run())
|
|
||||||
|
|
||||||
assert "Process running" in initial
|
|
||||||
assert "booting" in initial + waited
|
|
||||||
assert "Process running" in waited
|
|
||||||
assert "Wait target not observed: 'never-ready'" in waited
|
|
||||||
assert "Session terminated." in cleanup
|
|
||||||
|
|
||||||
|
|
||||||
def test_exec_session_mode_reuses_exec_safety_guard(tmp_path):
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
tool = ExecTool(
|
|
||||||
working_dir=str(tmp_path),
|
|
||||||
deny_patterns=[r"echo\s+blocked"],
|
|
||||||
session_manager=manager,
|
|
||||||
)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(command="echo blocked", yield_time_ms=0))
|
|
||||||
|
|
||||||
assert "blocked by deny pattern" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_write_stdin_reports_missing_session(tmp_path):
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
tool = WriteStdinTool(manager=manager)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(session_id="missing", chars=""))
|
|
||||||
|
|
||||||
assert "exec session not found" in result
|
|
||||||
|
|
||||||
|
|
||||||
def test_list_exec_sessions_reports_running_commands(tmp_path):
|
|
||||||
async def run() -> tuple[str, str, str]:
|
|
||||||
manager = ExecSessionManager()
|
|
||||||
exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
|
|
||||||
list_tool = ListExecSessionsTool(manager=manager)
|
|
||||||
stdin_tool = WriteStdinTool(manager=manager)
|
|
||||||
command = _python_command(
|
|
||||||
"import time; print('ready', flush=True); time.sleep(5)"
|
|
||||||
)
|
|
||||||
|
|
||||||
initial = await exec_tool.execute(command=command, yield_time_ms=500)
|
|
||||||
sid = _session_id(initial)
|
|
||||||
listing = await list_tool.execute()
|
|
||||||
cleanup = await stdin_tool.execute(session_id=sid, terminate=True, yield_time_ms=0)
|
|
||||||
return sid, listing, cleanup
|
|
||||||
|
|
||||||
sid, listing, cleanup = asyncio.run(run())
|
|
||||||
|
|
||||||
assert sid in listing
|
|
||||||
assert "running" in listing
|
|
||||||
assert "elapsed=" in listing
|
|
||||||
assert "remaining=" in listing
|
|
||||||
assert str(tmp_path) in listing
|
|
||||||
assert "Session terminated." in cleanup
|
|
||||||
|
|
||||||
|
|
||||||
def test_list_exec_sessions_reports_empty_state():
|
|
||||||
result = asyncio.run(ListExecSessionsTool(manager=ExecSessionManager()).execute())
|
|
||||||
|
|
||||||
assert result == "No active exec sessions."
|
|
||||||
@@ -1,216 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
|
|
||||||
from nanobot.agent.tools.filesystem import EditFileTool, ReadFileTool
|
|
||||||
|
|
||||||
|
|
||||||
def test_read_file_force_bypasses_dedup(tmp_path):
|
|
||||||
target = tmp_path / "data.txt"
|
|
||||||
target.write_text("alpha\n")
|
|
||||||
tool = ReadFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
first = asyncio.run(tool.execute(path=str(target)))
|
|
||||||
second = asyncio.run(tool.execute(path=str(target)))
|
|
||||||
forced = asyncio.run(tool.execute(path=str(target), force=True))
|
|
||||||
|
|
||||||
assert "alpha" in first
|
|
||||||
assert "unchanged" in second.lower()
|
|
||||||
assert "alpha" in forced
|
|
||||||
assert "unchanged" not in forced.lower()
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_can_select_occurrence(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("one\nsame\ntwo\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
occurrence=2,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "Successfully edited" in result
|
|
||||||
assert target.read_text() == "one\nsame\ntwo\nchanged\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_expected_replacements_guards_replace_all(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
replace_all=True,
|
|
||||||
expected_replacements=1,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "expected 1 replacements but would make 2" in result
|
|
||||||
assert target.read_text() == "same\nsame\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_expected_replacements_allows_replace_all_when_count_matches(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
replace_all=True,
|
|
||||||
expected_replacements=2,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "Successfully edited" in result
|
|
||||||
assert target.read_text() == "changed\nchanged\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_can_select_nearest_line_hint(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("one\nsame\ntwo\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
line_hint=4,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "Successfully edited" in result
|
|
||||||
assert target.read_text() == "one\nsame\ntwo\nchanged\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_can_edit_ipynb_as_json(tmp_path):
|
|
||||||
target = tmp_path / "analysis.ipynb"
|
|
||||||
target.write_text('{"cells": []}')
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text='"cells": []',
|
|
||||||
new_text='"cells": [{"cell_type": "markdown", "source": "hi"}]',
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "Successfully edited" in result
|
|
||||||
assert '"source": "hi"' in target.read_text()
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_multiple_match_hint_mentions_occurrence(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "old_text appears 2 times" in result
|
|
||||||
assert "occurrence" in result
|
|
||||||
assert target.read_text() == "same\nsame\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_rejects_ambiguous_line_hint(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\nmiddle\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
line_hint=2,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "line_hint 2 is ambiguous" in result
|
|
||||||
assert target.read_text() == "same\nmiddle\nsame\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_rejects_occurrence_with_replace_all(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
occurrence=1,
|
|
||||||
replace_all=True,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "occurrence cannot be used with replace_all" in result
|
|
||||||
assert target.read_text() == "same\nsame\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_rejects_line_hint_with_replace_all(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
line_hint=1,
|
|
||||||
replace_all=True,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "line_hint cannot be used with replace_all" in result
|
|
||||||
assert target.read_text() == "same\nsame\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_rejects_line_hint_with_occurrence(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\nsame\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
occurrence=1,
|
|
||||||
line_hint=1,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "line_hint cannot be used with occurrence" in result
|
|
||||||
assert target.read_text() == "same\nsame\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_rejects_zero_occurrence(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
occurrence=0,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "occurrence must be >= 1" in result
|
|
||||||
assert target.read_text() == "same\n"
|
|
||||||
|
|
||||||
|
|
||||||
def test_edit_file_rejects_zero_line_hint(tmp_path):
|
|
||||||
target = tmp_path / "duplicate.txt"
|
|
||||||
target.write_text("same\n")
|
|
||||||
tool = EditFileTool(workspace=tmp_path)
|
|
||||||
|
|
||||||
result = asyncio.run(tool.execute(
|
|
||||||
path=str(target),
|
|
||||||
old_text="same",
|
|
||||||
new_text="changed",
|
|
||||||
line_hint=0,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert "line_hint must be >= 1" in result
|
|
||||||
assert target.read_text() == "same\n"
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user