mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-09 13:58:36 +03:00
Compare commits
77
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5257453c4c | ||
|
|
a4dfbdf996 | ||
|
|
949a10f536 | ||
|
|
2a6c616080 | ||
|
|
1bcd5f9742 | ||
|
|
26947db479 | ||
|
|
0514233217 | ||
|
|
345c393e53 | ||
|
|
faf2b07923 | ||
|
|
efd42cc236 | ||
|
|
3823042290 | ||
|
|
5bdb7a90b1 | ||
|
|
bc8fbd1ce4 | ||
|
|
6aad945719 | ||
|
|
f450c6ef6c | ||
|
|
8956df3668 | ||
|
|
0506e6c1c1 | ||
|
|
b94d4c0509 | ||
|
|
d0c68157b1 | ||
|
|
351e3720b6 | ||
|
|
c3c1424db3 | ||
|
|
929ee09499 | ||
|
|
3f21e83af8 | ||
|
|
8682b017e2 | ||
|
|
7fad14802e | ||
|
|
842b8b255d | ||
|
|
758c4e74c9 | ||
|
|
f08de72f18 | ||
|
|
1814272583 | ||
|
|
5e99b81c6e | ||
|
|
d9a5080d66 | ||
|
|
55501057ac | ||
|
|
2dce5e07c1 | ||
|
|
5635907e33 | ||
|
|
a0684978fb | ||
|
|
1a4ad67628 | ||
|
|
ed2ca759e7 | ||
|
|
79a915307c | ||
|
|
2abd990b89 | ||
|
|
0207b541df | ||
|
|
b1d5475681 | ||
|
|
e04e1c24ff | ||
|
|
c8c520cc9a | ||
|
|
bee89df422 | ||
|
|
17d21c8e64 | ||
|
|
aebe928cf0 | ||
|
|
a42a4e9d83 | ||
|
|
c15f63a320 | ||
|
|
9652e67204 | ||
|
|
f8c580d015 | ||
|
|
5968b408dc | ||
|
|
e464a81545 | ||
|
|
0ba71298e6 | ||
|
|
cf25a582ba | ||
|
|
5ff9146a24 | ||
|
|
1331084873 | ||
|
|
ace3fd6049 | ||
|
|
5bf0f6fe7d | ||
|
|
e7d371ec1e | ||
|
|
33abe915e7 | ||
|
|
813de554c9 | ||
|
|
f0f0bf02d7 | ||
|
|
5e9fa28ff2 | ||
|
|
3f71014b7c | ||
|
|
fab14696a9 | ||
|
|
4a7d7b8823 | ||
|
|
13d6c0ae52 | ||
|
|
ef10df9acb | ||
|
|
9f19297056 | ||
|
|
9d69ba9f56 | ||
|
|
f5cf0bfdee | ||
|
|
6e428b7939 | ||
|
|
37060dea0b | ||
|
|
6b3997c463 | ||
|
|
e868fb32d2 | ||
|
|
f958eb4cc9 | ||
|
|
80219baf25 |
@@ -21,13 +21,23 @@
|
|||||||
## 📢 News
|
## 📢 News
|
||||||
|
|
||||||
> [!IMPORTANT]
|
> [!IMPORTANT]
|
||||||
> **Security note:** Due to `litellm` supply chain poisoning, **please check your Python environment ASAP** and refer to this [advisory](https://github.com/HKUDS/nanobot/discussions/2445) for details. We have fully removed the `litellm` dependency in [this commit](https://github.com/HKUDS/nanobot/commit/3dfdab7).
|
> **Security note:** Due to `litellm` supply chain poisoning, **please check your Python environment ASAP** and refer to this [advisory](https://github.com/HKUDS/nanobot/discussions/2445) for details. We have fully removed the `litellm` since **v0.1.4.post6**.
|
||||||
|
|
||||||
|
- **2026-03-27** 🚀 Released **v0.1.4.post6** — architecture decoupling, litellm removal, end-to-end streaming, WeChat channel, and a security fix. Please see [release notes](https://github.com/HKUDS/nanobot/releases/tag/v0.1.4.post6) for details.
|
||||||
|
- **2026-03-26** 🏗️ Agent runner extracted and lifecycle hooks unified; stream delta coalescing at boundaries.
|
||||||
|
- **2026-03-25** 🌏 StepFun provider, configurable timezone, Gemini thought signatures.
|
||||||
|
- **2026-03-24** 🔧 WeChat compatibility, Feishu CardKit streaming, test suite restructured.
|
||||||
|
- **2026-03-23** 🔧 Command routing refactored for plugins, WhatsApp/WeChat media, unified channel login CLI.
|
||||||
|
- **2026-03-22** ⚡ End-to-end streaming, WeChat channel, Anthropic cache optimization, `/status` command.
|
||||||
- **2026-03-21** 🔒 Replace `litellm` with native `openai` + `anthropic` SDKs. Please see [commit](https://github.com/HKUDS/nanobot/commit/3dfdab7).
|
- **2026-03-21** 🔒 Replace `litellm` with native `openai` + `anthropic` SDKs. Please see [commit](https://github.com/HKUDS/nanobot/commit/3dfdab7).
|
||||||
- **2026-03-20** 🧙 Interactive setup wizard — pick your provider, model autocomplete, and you're good to go.
|
- **2026-03-20** 🧙 Interactive setup wizard — pick your provider, model autocomplete, and you're good to go.
|
||||||
- **2026-03-19** 💬 Telegram gets more resilient under load; Feishu now renders code blocks properly.
|
- **2026-03-19** 💬 Telegram gets more resilient under load; Feishu now renders code blocks properly.
|
||||||
- **2026-03-18** 📷 Telegram can now send media via URL. Cron schedules show human-readable details.
|
- **2026-03-18** 📷 Telegram can now send media via URL. Cron schedules show human-readable details.
|
||||||
- **2026-03-17** ✨ Feishu formatting glow-up, Slack reacts when done, custom endpoints support extra headers, and image handling is more reliable.
|
- **2026-03-17** ✨ Feishu formatting glow-up, Slack reacts when done, custom endpoints support extra headers, and image handling is more reliable.
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary>Earlier news</summary>
|
||||||
|
|
||||||
- **2026-03-16** 🚀 Released **v0.1.4.post5** — a refinement-focused release with stronger reliability and channel support, and a more dependable day-to-day experience. Please see [release notes](https://github.com/HKUDS/nanobot/releases/tag/v0.1.4.post5) for details.
|
- **2026-03-16** 🚀 Released **v0.1.4.post5** — a refinement-focused release with stronger reliability and channel support, and a more dependable day-to-day experience. Please see [release notes](https://github.com/HKUDS/nanobot/releases/tag/v0.1.4.post5) for details.
|
||||||
- **2026-03-15** 🧩 DingTalk rich media, smarter built-in skills, and cleaner model compatibility.
|
- **2026-03-15** 🧩 DingTalk rich media, smarter built-in skills, and cleaner model compatibility.
|
||||||
- **2026-03-14** 💬 Channel plugins, Feishu replies, and steadier MCP, QQ, and media handling.
|
- **2026-03-14** 💬 Channel plugins, Feishu replies, and steadier MCP, QQ, and media handling.
|
||||||
@@ -39,10 +49,6 @@
|
|||||||
- **2026-03-08** 🚀 Released **v0.1.4.post4** — a reliability-packed release with safer defaults, better multi-instance support, sturdier MCP, and major channel and provider improvements. Please see [release notes](https://github.com/HKUDS/nanobot/releases/tag/v0.1.4.post4) for details.
|
- **2026-03-08** 🚀 Released **v0.1.4.post4** — a reliability-packed release with safer defaults, better multi-instance support, sturdier MCP, and major channel and provider improvements. Please see [release notes](https://github.com/HKUDS/nanobot/releases/tag/v0.1.4.post4) for details.
|
||||||
- **2026-03-07** 🚀 Azure OpenAI provider, WhatsApp media, QQ group chats, and more Telegram/Feishu polish.
|
- **2026-03-07** 🚀 Azure OpenAI provider, WhatsApp media, QQ group chats, and more Telegram/Feishu polish.
|
||||||
- **2026-03-06** 🪄 Lighter providers, smarter media handling, and sturdier memory and CLI compatibility.
|
- **2026-03-06** 🪄 Lighter providers, smarter media handling, and sturdier memory and CLI compatibility.
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary>Earlier news</summary>
|
|
||||||
|
|
||||||
- **2026-03-05** ⚡️ Telegram draft streaming, MCP SSE support, and broader channel reliability fixes.
|
- **2026-03-05** ⚡️ Telegram draft streaming, MCP SSE support, and broader channel reliability fixes.
|
||||||
- **2026-03-04** 🛠️ Dependency cleanup, safer file reads, and another round of test and Cron fixes.
|
- **2026-03-04** 🛠️ Dependency cleanup, safer file reads, and another round of test and Cron fixes.
|
||||||
- **2026-03-03** 🧠 Cleaner user-message merging, safer multimodal saves, and stronger Cron guards.
|
- **2026-03-03** 🧠 Cleaner user-message merging, safer multimodal saves, and stronger Cron guards.
|
||||||
@@ -109,6 +115,8 @@
|
|||||||
- [Configuration](#️-configuration)
|
- [Configuration](#️-configuration)
|
||||||
- [Multiple Instances](#-multiple-instances)
|
- [Multiple Instances](#-multiple-instances)
|
||||||
- [CLI Reference](#-cli-reference)
|
- [CLI Reference](#-cli-reference)
|
||||||
|
- [Python SDK](#-python-sdk)
|
||||||
|
- [OpenAI-Compatible API](#-openai-compatible-api)
|
||||||
- [Docker](#-docker)
|
- [Docker](#-docker)
|
||||||
- [Linux Service](#-linux-service)
|
- [Linux Service](#-linux-service)
|
||||||
- [Project Structure](#-project-structure)
|
- [Project Structure](#-project-structure)
|
||||||
@@ -505,14 +513,17 @@ nanobot gateway
|
|||||||
</details>
|
</details>
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>Feishu (飞书)</b></summary>
|
<summary><b>Feishu</b></summary>
|
||||||
|
|
||||||
Uses **WebSocket** long connection — no public IP required.
|
Uses **WebSocket** long connection — no public IP required.
|
||||||
|
|
||||||
**1. Create a Feishu bot**
|
**1. Create a Feishu bot**
|
||||||
- Visit [Feishu Open Platform](https://open.feishu.cn/app)
|
- Visit [Feishu Open Platform](https://open.feishu.cn/app)
|
||||||
- Create a new app → Enable **Bot** capability
|
- Create a new app → Enable **Bot** capability
|
||||||
- **Permissions**: Add `im:message` (send messages) and `im:message.p2p_msg:readonly` (receive messages)
|
- **Permissions**:
|
||||||
|
- `im:message` (send messages) and `im:message.p2p_msg:readonly` (receive messages)
|
||||||
|
- **Streaming replies** (default in nanobot): add **`cardkit:card:write`** (often labeled **Create and update cards** in the Feishu developer console). Required for CardKit entities and streamed assistant text. Older apps may not have it yet — open **Permission management**, enable the scope, then **publish** a new app version if the console requires it.
|
||||||
|
- If you **cannot** add `cardkit:card:write`, set `"streaming": false` under `channels.feishu` (see below). The bot still works; replies use normal interactive cards without token-by-token streaming.
|
||||||
- **Events**: Add `im.message.receive_v1` (receive messages)
|
- **Events**: Add `im.message.receive_v1` (receive messages)
|
||||||
- Select **Long Connection** mode (requires running nanobot first to establish connection)
|
- Select **Long Connection** mode (requires running nanobot first to establish connection)
|
||||||
- Get **App ID** and **App Secret** from "Credentials & Basic Info"
|
- Get **App ID** and **App Secret** from "Credentials & Basic Info"
|
||||||
@@ -530,12 +541,14 @@ Uses **WebSocket** long connection — no public IP required.
|
|||||||
"encryptKey": "",
|
"encryptKey": "",
|
||||||
"verificationToken": "",
|
"verificationToken": "",
|
||||||
"allowFrom": ["ou_YOUR_OPEN_ID"],
|
"allowFrom": ["ou_YOUR_OPEN_ID"],
|
||||||
"groupPolicy": "mention"
|
"groupPolicy": "mention",
|
||||||
|
"streaming": true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> `streaming` defaults to `true`. Use `false` if your app does not have **`cardkit:card:write`** (see permissions above).
|
||||||
> `encryptKey` and `verificationToken` are optional for Long Connection mode.
|
> `encryptKey` and `verificationToken` are optional for Long Connection mode.
|
||||||
> `allowFrom`: Add your open_id (find it in nanobot logs when you message the bot). Use `["*"]` to allow all users.
|
> `allowFrom`: Add your open_id (find it in nanobot logs when you message the bot). Use `["*"]` to allow all users.
|
||||||
> `groupPolicy`: `"mention"` (default — respond only when @mentioned), `"open"` (respond to all group messages). Private chats always respond.
|
> `groupPolicy`: `"mention"` (default — respond only when @mentioned), `"open"` (respond to all group messages). Private chats always respond.
|
||||||
@@ -733,14 +746,10 @@ nanobot gateway
|
|||||||
|
|
||||||
Uses **HTTP long-poll** with QR-code login via the ilinkai personal WeChat API. No local WeChat desktop client is required.
|
Uses **HTTP long-poll** with QR-code login via the ilinkai personal WeChat API. No local WeChat desktop client is required.
|
||||||
|
|
||||||
> Weixin support is available from source checkout, but is not included in the current PyPI release yet.
|
**1. Install with WeChat support**
|
||||||
|
|
||||||
**1. Install from source**
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
git clone https://github.com/HKUDS/nanobot.git
|
pip install "nanobot-ai[weixin]"
|
||||||
cd nanobot
|
|
||||||
pip install -e ".[weixin]"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**2. Configure**
|
**2. Configure**
|
||||||
@@ -846,6 +855,7 @@ Config file: `~/.nanobot/config.json`
|
|||||||
> - **VolcEngine / BytePlus Coding Plan**: Use dedicated providers `volcengineCodingPlan` or `byteplusCodingPlan` instead of the pay-per-use `volcengine` / `byteplus` providers.
|
> - **VolcEngine / BytePlus Coding Plan**: Use dedicated providers `volcengineCodingPlan` or `byteplusCodingPlan` instead of the pay-per-use `volcengine` / `byteplus` providers.
|
||||||
> - **Zhipu Coding Plan**: If you're on Zhipu's coding plan, set `"apiBase": "https://open.bigmodel.cn/api/coding/paas/v4"` in your zhipu provider config.
|
> - **Zhipu Coding Plan**: If you're on Zhipu's coding plan, set `"apiBase": "https://open.bigmodel.cn/api/coding/paas/v4"` in your zhipu provider config.
|
||||||
> - **Alibaba Cloud BaiLian**: If you're using Alibaba Cloud BaiLian's OpenAI-compatible endpoint, set `"apiBase": "https://dashscope.aliyuncs.com/compatible-mode/v1"` in your dashscope provider config.
|
> - **Alibaba Cloud BaiLian**: If you're using Alibaba Cloud BaiLian's OpenAI-compatible endpoint, set `"apiBase": "https://dashscope.aliyuncs.com/compatible-mode/v1"` in your dashscope provider config.
|
||||||
|
> - **Step Fun (Mainland China)**: If your API key is from Step Fun's mainland China platform (stepfun.com), set `"apiBase": "https://api.stepfun.com/v1"` in your stepfun provider config.
|
||||||
|
|
||||||
| Provider | Purpose | Get API Key |
|
| Provider | Purpose | Get API Key |
|
||||||
|----------|---------|-------------|
|
|----------|---------|-------------|
|
||||||
@@ -867,6 +877,7 @@ Config file: `~/.nanobot/config.json`
|
|||||||
| `zhipu` | LLM (Zhipu GLM) | [open.bigmodel.cn](https://open.bigmodel.cn) |
|
| `zhipu` | LLM (Zhipu GLM) | [open.bigmodel.cn](https://open.bigmodel.cn) |
|
||||||
| `ollama` | LLM (local, Ollama) | — |
|
| `ollama` | LLM (local, Ollama) | — |
|
||||||
| `mistral` | LLM | [docs.mistral.ai](https://docs.mistral.ai/) |
|
| `mistral` | LLM | [docs.mistral.ai](https://docs.mistral.ai/) |
|
||||||
|
| `stepfun` | LLM (Step Fun/阶跃星辰) | [platform.stepfun.com](https://platform.stepfun.com) |
|
||||||
| `ovms` | LLM (local, OpenVINO Model Server) | [docs.openvino.ai](https://docs.openvino.ai/2026/model-server/ovms_docs_llm_quickstart.html) |
|
| `ovms` | LLM (local, OpenVINO Model Server) | [docs.openvino.ai](https://docs.openvino.ai/2026/model-server/ovms_docs_llm_quickstart.html) |
|
||||||
| `vllm` | LLM (local, any OpenAI-compatible server) | — |
|
| `vllm` | LLM (local, any OpenAI-compatible server) | — |
|
||||||
| `openai_codex` | LLM (Codex, OAuth) | `nanobot provider login openai-codex` |
|
| `openai_codex` | LLM (Codex, OAuth) | `nanobot provider login openai-codex` |
|
||||||
@@ -1154,9 +1165,43 @@ That's it! Environment variables, model routing, config matching, and `nanobot s
|
|||||||
| `detect_by_key_prefix` | Detect gateway by API key prefix | `"sk-or-"` |
|
| `detect_by_key_prefix` | Detect gateway by API key prefix | `"sk-or-"` |
|
||||||
| `detect_by_base_keyword` | Detect gateway by API base URL | `"openrouter"` |
|
| `detect_by_base_keyword` | Detect gateway by API base URL | `"openrouter"` |
|
||||||
| `strip_model_prefix` | Strip provider prefix before sending to gateway | `True` (for AiHubMix) |
|
| `strip_model_prefix` | Strip provider prefix before sending to gateway | `True` (for AiHubMix) |
|
||||||
|
| `supports_max_completion_tokens` | Use `max_completion_tokens` instead of `max_tokens`; required for providers that reject both being set simultaneously (e.g. VolcEngine) | `True` |
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
### Channel Settings
|
||||||
|
|
||||||
|
Global settings that apply to all channels. Configure under the `channels` section in `~/.nanobot/config.json`:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"channels": {
|
||||||
|
"sendProgress": true,
|
||||||
|
"sendToolHints": false,
|
||||||
|
"sendMaxRetries": 3,
|
||||||
|
"telegram": { ... }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
| Setting | Default | Description |
|
||||||
|
|---------|---------|-------------|
|
||||||
|
| `sendProgress` | `true` | Stream agent's text progress to the channel |
|
||||||
|
| `sendToolHints` | `false` | Stream tool-call hints (e.g. `read_file("…")`) |
|
||||||
|
| `sendMaxRetries` | `3` | Max delivery attempts per outbound message, including the initial send (0-10 configured, minimum 1 actual attempt) |
|
||||||
|
|
||||||
|
#### Retry Behavior
|
||||||
|
|
||||||
|
When a channel send operation raises an error, nanobot retries with exponential backoff:
|
||||||
|
|
||||||
|
- **Attempt 1**: Initial send
|
||||||
|
- **Attempts 2-4**: Retry delays are 1s, 2s, 4s
|
||||||
|
- **Attempts 5+**: Retry delay caps at 4s
|
||||||
|
- **Transient failures** (network hiccups, temporary API limits): Retry usually succeeds
|
||||||
|
- **Permanent failures** (invalid token, channel banned): All retries fail
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> When a channel is completely unavailable, there's no way to notify the user since we cannot reach them through that channel. Monitor logs for "Failed to send to {channel} after N attempts" to detect persistent delivery failures.
|
||||||
|
|
||||||
### Web Search
|
### Web Search
|
||||||
|
|
||||||
@@ -1342,9 +1387,33 @@ MCP tools are automatically discovered and registered on startup. The LLM can us
|
|||||||
| `tools.restrictToWorkspace` | `false` | When `true`, restricts **all** agent tools (shell, file read/write/edit, list) to the workspace directory. Prevents path traversal and out-of-scope access. |
|
| `tools.restrictToWorkspace` | `false` | When `true`, restricts **all** agent tools (shell, file read/write/edit, list) to the workspace directory. Prevents path traversal and out-of-scope access. |
|
||||||
| `tools.exec.enable` | `true` | When `false`, the shell `exec` tool is not registered at all. Use this to completely disable shell command execution. |
|
| `tools.exec.enable` | `true` | When `false`, the shell `exec` tool is not registered at all. Use this to completely disable shell command execution. |
|
||||||
| `tools.exec.pathAppend` | `""` | Extra directories to append to `PATH` when running shell commands (e.g. `/usr/sbin` for `ufw`). |
|
| `tools.exec.pathAppend` | `""` | Extra directories to append to `PATH` when running shell commands (e.g. `/usr/sbin` for `ufw`). |
|
||||||
|
| `tools.exec.commandWrapper` | `""` | Sandbox wrapper command template. See [Exec Tool Sandbox](docs/COMMAND_WRAPPER.md) for details and examples. |
|
||||||
|
|
||||||
| `channels.*.allowFrom` | `[]` (deny all) | Whitelist of user IDs. Empty denies all; use `["*"]` to allow everyone. |
|
| `channels.*.allowFrom` | `[]` (deny all) | Whitelist of user IDs. Empty denies all; use `["*"]` to allow everyone. |
|
||||||
|
|
||||||
|
|
||||||
|
### Timezone
|
||||||
|
|
||||||
|
Time is context. Context should be precise.
|
||||||
|
|
||||||
|
By default, nanobot uses `UTC` for runtime time context. If you want the agent to think in your local time, set `agents.defaults.timezone` to a valid [IANA timezone name](https://en.wikipedia.org/wiki/List_of_tz_database_time_zones):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"agents": {
|
||||||
|
"defaults": {
|
||||||
|
"timezone": "Asia/Shanghai"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
This affects runtime time strings shown to the model, such as runtime context and heartbeat prompts. It also becomes the default timezone for cron schedules when a cron expression omits `tz`, and for one-shot `at` times when the ISO datetime has no explicit offset.
|
||||||
|
|
||||||
|
Common examples: `UTC`, `America/New_York`, `America/Los_Angeles`, `Europe/London`, `Europe/Berlin`, `Asia/Tokyo`, `Asia/Shanghai`, `Asia/Singapore`, `Australia/Sydney`.
|
||||||
|
|
||||||
|
> Need another timezone? Browse the full [IANA Time Zone Database](https://en.wikipedia.org/wiki/List_of_tz_database_time_zones).
|
||||||
|
|
||||||
## 🧩 Multiple Instances
|
## 🧩 Multiple Instances
|
||||||
|
|
||||||
Run multiple nanobot instances simultaneously with separate configs and runtime data. Use `--config` as the main entrypoint. Optionally pass `--workspace` during `onboard` when you want to initialize or update the saved workspace for a specific instance.
|
Run multiple nanobot instances simultaneously with separate configs and runtime data. Use `--config` as the main entrypoint. Optionally pass `--workspace` during `onboard` when you want to initialize or update the saved workspace for a specific instance.
|
||||||
@@ -1476,6 +1545,7 @@ nanobot gateway --config ~/.nanobot-telegram/config.json --workspace /tmp/nanobo
|
|||||||
| `nanobot agent` | Interactive chat mode |
|
| `nanobot agent` | Interactive chat mode |
|
||||||
| `nanobot agent --no-markdown` | Show plain-text replies |
|
| `nanobot agent --no-markdown` | Show plain-text replies |
|
||||||
| `nanobot agent --logs` | Show runtime logs during chat |
|
| `nanobot agent --logs` | Show runtime logs during chat |
|
||||||
|
| `nanobot serve` | Start the OpenAI-compatible API |
|
||||||
| `nanobot gateway` | Start the gateway |
|
| `nanobot gateway` | Start the gateway |
|
||||||
| `nanobot status` | Show status |
|
| `nanobot status` | Show status |
|
||||||
| `nanobot provider login openai-codex` | OAuth login for providers |
|
| `nanobot provider login openai-codex` | OAuth login for providers |
|
||||||
@@ -1504,6 +1574,110 @@ The agent can also manage this file itself — ask it to "add a periodic task" a
|
|||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
## 🐍 Python SDK
|
||||||
|
|
||||||
|
Use nanobot as a library — no CLI, no gateway, just Python:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from nanobot import Nanobot
|
||||||
|
|
||||||
|
bot = Nanobot.from_config()
|
||||||
|
result = await bot.run("Summarize the README")
|
||||||
|
print(result.content)
|
||||||
|
```
|
||||||
|
|
||||||
|
Each call carries a `session_key` for conversation isolation — different keys get independent history:
|
||||||
|
|
||||||
|
```python
|
||||||
|
await bot.run("hi", session_key="user-alice")
|
||||||
|
await bot.run("hi", session_key="task-42")
|
||||||
|
```
|
||||||
|
|
||||||
|
Add lifecycle hooks to observe or customize the agent:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from nanobot.agent import AgentHook, AgentHookContext
|
||||||
|
|
||||||
|
class AuditHook(AgentHook):
|
||||||
|
async def before_execute_tools(self, ctx: AgentHookContext) -> None:
|
||||||
|
for tc in ctx.tool_calls:
|
||||||
|
print(f"[tool] {tc.name}")
|
||||||
|
|
||||||
|
result = await bot.run("Hello", hooks=[AuditHook()])
|
||||||
|
```
|
||||||
|
|
||||||
|
See [docs/PYTHON_SDK.md](docs/PYTHON_SDK.md) for the full SDK reference.
|
||||||
|
|
||||||
|
## 🔌 OpenAI-Compatible API
|
||||||
|
|
||||||
|
nanobot can expose a minimal OpenAI-compatible endpoint for local integrations:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
pip install "nanobot-ai[api]"
|
||||||
|
nanobot serve
|
||||||
|
```
|
||||||
|
|
||||||
|
By default, the API binds to `127.0.0.1:8900`. You can change this in `config.json`.
|
||||||
|
|
||||||
|
### Behavior
|
||||||
|
|
||||||
|
- Session isolation: pass `"session_id"` in the request body to isolate conversations; omit for a shared default session (`api:default`)
|
||||||
|
- Single-message input: each request must contain exactly one `user` message
|
||||||
|
- Fixed model: omit `model`, or pass the same model shown by `/v1/models`
|
||||||
|
- No streaming: `stream=true` is not supported
|
||||||
|
|
||||||
|
### Endpoints
|
||||||
|
|
||||||
|
- `GET /health`
|
||||||
|
- `GET /v1/models`
|
||||||
|
- `POST /v1/chat/completions`
|
||||||
|
|
||||||
|
### curl
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://127.0.0.1:8900/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "hi"}],
|
||||||
|
"session_id": "my-session"
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
### Python (`requests`)
|
||||||
|
|
||||||
|
```python
|
||||||
|
import requests
|
||||||
|
|
||||||
|
resp = requests.post(
|
||||||
|
"http://127.0.0.1:8900/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"messages": [{"role": "user", "content": "hi"}],
|
||||||
|
"session_id": "my-session", # optional: isolate conversation
|
||||||
|
},
|
||||||
|
timeout=120,
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
print(resp.json()["choices"][0]["message"]["content"])
|
||||||
|
```
|
||||||
|
|
||||||
|
### Python (`openai`)
|
||||||
|
|
||||||
|
```python
|
||||||
|
from openai import OpenAI
|
||||||
|
|
||||||
|
client = OpenAI(
|
||||||
|
base_url="http://127.0.0.1:8900/v1",
|
||||||
|
api_key="dummy",
|
||||||
|
)
|
||||||
|
|
||||||
|
resp = client.chat.completions.create(
|
||||||
|
model="MiniMax-M2.7",
|
||||||
|
messages=[{"role": "user", "content": "hi"}],
|
||||||
|
extra_body={"session_id": "my-session"}, # optional: isolate conversation
|
||||||
|
)
|
||||||
|
print(resp.choices[0].message.content)
|
||||||
|
```
|
||||||
|
|
||||||
## 🐳 Docker
|
## 🐳 Docker
|
||||||
|
|
||||||
> [!TIP]
|
> [!TIP]
|
||||||
|
|||||||
+4
-3
@@ -1,5 +1,6 @@
|
|||||||
#!/bin/bash
|
#!/bin/bash
|
||||||
# Count core agent lines (excluding channels/, cli/, providers/ adapters)
|
# Count core agent lines (excluding channels/, cli/, api/, providers/ adapters,
|
||||||
|
# and the high-level Python SDK facade)
|
||||||
cd "$(dirname "$0")" || exit 1
|
cd "$(dirname "$0")" || exit 1
|
||||||
|
|
||||||
echo "nanobot core agent line count"
|
echo "nanobot core agent line count"
|
||||||
@@ -15,7 +16,7 @@ root=$(cat nanobot/__init__.py nanobot/__main__.py | wc -l)
|
|||||||
printf " %-16s %5s lines\n" "(root)" "$root"
|
printf " %-16s %5s lines\n" "(root)" "$root"
|
||||||
|
|
||||||
echo ""
|
echo ""
|
||||||
total=$(find nanobot -name "*.py" ! -path "*/channels/*" ! -path "*/cli/*" ! -path "*/command/*" ! -path "*/providers/*" ! -path "*/skills/*" | xargs cat | wc -l)
|
total=$(find nanobot -name "*.py" ! -path "*/channels/*" ! -path "*/cli/*" ! -path "*/api/*" ! -path "*/command/*" ! -path "*/providers/*" ! -path "*/skills/*" ! -path "nanobot/nanobot.py" | xargs cat | wc -l)
|
||||||
echo " Core total: $total lines"
|
echo " Core total: $total lines"
|
||||||
echo ""
|
echo ""
|
||||||
echo " (excludes: channels/, cli/, command/, providers/, skills/)"
|
echo " (excludes: channels/, cli/, api/, command/, providers/, skills/, nanobot.py)"
|
||||||
|
|||||||
@@ -0,0 +1,86 @@
|
|||||||
|
# Exec Tool Sandbox (`commandWrapper`)
|
||||||
|
|
||||||
|
The `tools.exec.commandWrapper` config option wraps every shell command in a user-defined template before execution. This allows you to add a sandbox layer (e.g. bubblewrap, firejail, nsjail) without any code changes to nanobot.
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"commandWrapper": "<template>"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Leave empty (the default) to run commands directly with no wrapper.
|
||||||
|
|
||||||
|
## Placeholders
|
||||||
|
|
||||||
|
Two placeholders are available in the template:
|
||||||
|
|
||||||
|
| Placeholder | Value |
|
||||||
|
|---|---|
|
||||||
|
| `{command}` | The original shell command generated by the LLM |
|
||||||
|
| `{cwd}` | Absolute path of the working directory |
|
||||||
|
|
||||||
|
nanobot performs plain string replacement — it does not parse, validate, or shell-escape the values. The wrapper template is trusted configuration.
|
||||||
|
|
||||||
|
## Examples
|
||||||
|
|
||||||
|
### bubblewrap
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"commandWrapper": "bwrap --ro-bind /usr /usr --ro-bind-try /bin /bin --ro-bind-try /lib /lib --ro-bind-try /lib64 /lib64 --proc /proc --dev /dev --tmpfs /tmp --bind {cwd} {cwd} --chdir {cwd} -- sh -c \"{command}\""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Requires: `apt install bubblewrap` (or equivalent for your distro).
|
||||||
|
|
||||||
|
### firejail
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"commandWrapper": "firejail --noprofile --private={cwd} -- {command}"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### nsjail
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"commandWrapper": "nsjail -Mo --chroot /sandbox --cwd {cwd} -- {command}"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Caveats
|
||||||
|
|
||||||
|
> [!WARNING]
|
||||||
|
> **Do not wrap `{command}` in shell quotes.** If the original command contains the same quote character, the shell will break the quoting context. For example, `sh -c '{command}'` will fail on any command that contains single quotes.
|
||||||
|
|
||||||
|
This is an inherent limitation of the template approach — nanobot substitutes `{command}` as a raw string and cannot safely shell-quote it (the command may contain compound syntax like `&&`, `|`, `;` that must be preserved for the inner shell).
|
||||||
|
|
||||||
|
### Interaction with `create_subprocess_shell`
|
||||||
|
|
||||||
|
nanobot executes the wrapped command via `create_subprocess_shell`, which adds an outer shell layer. Keep this in mind when designing your template:
|
||||||
|
|
||||||
|
- **Without `sh -c`** (e.g. `firejail ... -- {command}`): The outer shell parses `{command}` directly. Compound commands with `&&` and `|` work as expected because they are parsed by the outer shell before the sandbox tool receives them.
|
||||||
|
- **With `sh -c`** (e.g. `bwrap ... -- sh -c "{command}"`): The command is passed through two shell layers. This is only needed if the sandbox tool requires a single command argument but you want to support compound syntax.
|
||||||
|
|
||||||
|
### `restrict_to_workspace` is independent
|
||||||
|
|
||||||
|
The `tools.restrictToWorkspace` setting and `commandWrapper` are orthogonal features. The workspace restriction guards against path traversal in the original command (before wrapping). The sandbox wrapper provides OS-level isolation. You can use either or both — they address different threat models.
|
||||||
@@ -0,0 +1,136 @@
|
|||||||
|
# Python SDK
|
||||||
|
|
||||||
|
Use nanobot programmatically — load config, run the agent, get results.
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
```python
|
||||||
|
import asyncio
|
||||||
|
from nanobot import Nanobot
|
||||||
|
|
||||||
|
async def main():
|
||||||
|
bot = Nanobot.from_config()
|
||||||
|
result = await bot.run("What time is it in Tokyo?")
|
||||||
|
print(result.content)
|
||||||
|
|
||||||
|
asyncio.run(main())
|
||||||
|
```
|
||||||
|
|
||||||
|
## API
|
||||||
|
|
||||||
|
### `Nanobot.from_config(config_path?, *, workspace?)`
|
||||||
|
|
||||||
|
Create a `Nanobot` from a config file.
|
||||||
|
|
||||||
|
| Param | Type | Default | Description |
|
||||||
|
|-------|------|---------|-------------|
|
||||||
|
| `config_path` | `str \| Path \| None` | `None` | Path to `config.json`. Defaults to `~/.nanobot/config.json`. |
|
||||||
|
| `workspace` | `str \| Path \| None` | `None` | Override workspace directory from config. |
|
||||||
|
|
||||||
|
Raises `FileNotFoundError` if an explicit path doesn't exist.
|
||||||
|
|
||||||
|
### `await bot.run(message, *, session_key?, hooks?)`
|
||||||
|
|
||||||
|
Run the agent once. Returns a `RunResult`.
|
||||||
|
|
||||||
|
| Param | Type | Default | Description |
|
||||||
|
|-------|------|---------|-------------|
|
||||||
|
| `message` | `str` | *(required)* | The user message to process. |
|
||||||
|
| `session_key` | `str` | `"sdk:default"` | Session identifier for conversation isolation. Different keys get independent history. |
|
||||||
|
| `hooks` | `list[AgentHook] \| None` | `None` | Lifecycle hooks for this run only. |
|
||||||
|
|
||||||
|
```python
|
||||||
|
# Isolated sessions — each user gets independent conversation history
|
||||||
|
await bot.run("hi", session_key="user-alice")
|
||||||
|
await bot.run("hi", session_key="user-bob")
|
||||||
|
```
|
||||||
|
|
||||||
|
### `RunResult`
|
||||||
|
|
||||||
|
| Field | Type | Description |
|
||||||
|
|-------|------|-------------|
|
||||||
|
| `content` | `str` | The agent's final text response. |
|
||||||
|
| `tools_used` | `list[str]` | Tool names invoked during the run. |
|
||||||
|
| `messages` | `list[dict]` | Raw message history (for debugging). |
|
||||||
|
|
||||||
|
## Hooks
|
||||||
|
|
||||||
|
Hooks let you observe or modify the agent loop without touching internals.
|
||||||
|
|
||||||
|
Subclass `AgentHook` and override any method:
|
||||||
|
|
||||||
|
| Method | When |
|
||||||
|
|--------|------|
|
||||||
|
| `before_iteration(ctx)` | Before each LLM call |
|
||||||
|
| `on_stream(ctx, delta)` | On each streamed token |
|
||||||
|
| `on_stream_end(ctx)` | When streaming finishes |
|
||||||
|
| `before_execute_tools(ctx)` | Before tool execution (inspect `ctx.tool_calls`) |
|
||||||
|
| `after_iteration(ctx, response)` | After each LLM response |
|
||||||
|
| `finalize_content(ctx, content)` | Transform final output text |
|
||||||
|
|
||||||
|
### Example: Audit Hook
|
||||||
|
|
||||||
|
```python
|
||||||
|
from nanobot.agent import AgentHook, AgentHookContext
|
||||||
|
|
||||||
|
class AuditHook(AgentHook):
|
||||||
|
def __init__(self):
|
||||||
|
self.calls = []
|
||||||
|
|
||||||
|
async def before_execute_tools(self, ctx: AgentHookContext) -> None:
|
||||||
|
for tc in ctx.tool_calls:
|
||||||
|
self.calls.append(tc.name)
|
||||||
|
print(f"[audit] {tc.name}({tc.arguments})")
|
||||||
|
|
||||||
|
hook = AuditHook()
|
||||||
|
result = await bot.run("List files in /tmp", hooks=[hook])
|
||||||
|
print(f"Tools used: {hook.calls}")
|
||||||
|
```
|
||||||
|
|
||||||
|
### Composing Hooks
|
||||||
|
|
||||||
|
Pass multiple hooks — they run in order, errors in one don't block others:
|
||||||
|
|
||||||
|
```python
|
||||||
|
result = await bot.run("hi", hooks=[AuditHook(), MetricsHook()])
|
||||||
|
```
|
||||||
|
|
||||||
|
Under the hood this uses `CompositeHook` for fan-out with error isolation.
|
||||||
|
|
||||||
|
### `finalize_content` Pipeline
|
||||||
|
|
||||||
|
Unlike the async methods (fan-out), `finalize_content` is a pipeline — each hook's output feeds the next:
|
||||||
|
|
||||||
|
```python
|
||||||
|
class Censor(AgentHook):
|
||||||
|
def finalize_content(self, ctx, content):
|
||||||
|
return content.replace("secret", "***") if content else content
|
||||||
|
```
|
||||||
|
|
||||||
|
## Full Example
|
||||||
|
|
||||||
|
```python
|
||||||
|
import asyncio
|
||||||
|
from nanobot import Nanobot
|
||||||
|
from nanobot.agent import AgentHook, AgentHookContext
|
||||||
|
|
||||||
|
class TimingHook(AgentHook):
|
||||||
|
async def before_iteration(self, ctx: AgentHookContext) -> None:
|
||||||
|
import time
|
||||||
|
ctx.metadata["_t0"] = time.time()
|
||||||
|
|
||||||
|
async def after_iteration(self, ctx, response) -> None:
|
||||||
|
import time
|
||||||
|
elapsed = time.time() - ctx.metadata.get("_t0", 0)
|
||||||
|
print(f"[timing] iteration took {elapsed:.2f}s")
|
||||||
|
|
||||||
|
async def main():
|
||||||
|
bot = Nanobot.from_config(workspace="/my/project")
|
||||||
|
result = await bot.run(
|
||||||
|
"Explain the main function",
|
||||||
|
hooks=[TimingHook()],
|
||||||
|
)
|
||||||
|
print(result.content)
|
||||||
|
|
||||||
|
asyncio.run(main())
|
||||||
|
```
|
||||||
+5
-1
@@ -2,5 +2,9 @@
|
|||||||
nanobot - A lightweight AI agent framework
|
nanobot - A lightweight AI agent framework
|
||||||
"""
|
"""
|
||||||
|
|
||||||
__version__ = "0.1.4.post5"
|
__version__ = "0.1.4.post6"
|
||||||
__logo__ = "🐈"
|
__logo__ = "🐈"
|
||||||
|
|
||||||
|
from nanobot.nanobot import Nanobot, RunResult
|
||||||
|
|
||||||
|
__all__ = ["Nanobot", "RunResult"]
|
||||||
|
|||||||
@@ -1,8 +1,19 @@
|
|||||||
"""Agent core module."""
|
"""Agent core module."""
|
||||||
|
|
||||||
from nanobot.agent.context import ContextBuilder
|
from nanobot.agent.context import ContextBuilder
|
||||||
|
from nanobot.agent.hook import AgentHook, AgentHookContext, CompositeHook
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.agent.memory import MemoryStore
|
from nanobot.agent.memory import MemoryStore
|
||||||
from nanobot.agent.skills import SkillsLoader
|
from nanobot.agent.skills import SkillsLoader
|
||||||
|
from nanobot.agent.subagent import SubagentManager
|
||||||
|
|
||||||
__all__ = ["AgentLoop", "ContextBuilder", "MemoryStore", "SkillsLoader"]
|
__all__ = [
|
||||||
|
"AgentHook",
|
||||||
|
"AgentHookContext",
|
||||||
|
"AgentLoop",
|
||||||
|
"CompositeHook",
|
||||||
|
"ContextBuilder",
|
||||||
|
"MemoryStore",
|
||||||
|
"SkillsLoader",
|
||||||
|
"SubagentManager",
|
||||||
|
]
|
||||||
|
|||||||
@@ -19,8 +19,9 @@ class ContextBuilder:
|
|||||||
BOOTSTRAP_FILES = ["AGENTS.md", "SOUL.md", "USER.md", "TOOLS.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]"
|
||||||
|
|
||||||
def __init__(self, workspace: Path):
|
def __init__(self, workspace: Path, timezone: str | None = None):
|
||||||
self.workspace = workspace
|
self.workspace = workspace
|
||||||
|
self.timezone = timezone
|
||||||
self.memory = MemoryStore(workspace)
|
self.memory = MemoryStore(workspace)
|
||||||
self.skills = SkillsLoader(workspace)
|
self.skills = SkillsLoader(workspace)
|
||||||
|
|
||||||
@@ -100,9 +101,11 @@ Reply directly with text for conversations. Only use the 'message' tool to send
|
|||||||
IMPORTANT: To send files (images, documents, audio, video) to the user, you MUST call the 'message' tool 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 file", media=["/path/to/file.png"])"""
|
IMPORTANT: To send files (images, documents, audio, video) to the user, you MUST call the 'message' tool 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 file", media=["/path/to/file.png"])"""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _build_runtime_context(channel: str | None, chat_id: str | None) -> str:
|
def _build_runtime_context(
|
||||||
|
channel: str | None, chat_id: str | None, timezone: str | None = None,
|
||||||
|
) -> str:
|
||||||
"""Build untrusted runtime metadata block for injection before the user message."""
|
"""Build untrusted runtime metadata block for injection before the user message."""
|
||||||
lines = [f"Current Time: {current_time_str()}"]
|
lines = [f"Current Time: {current_time_str(timezone)}"]
|
||||||
if channel and chat_id:
|
if channel and chat_id:
|
||||||
lines += [f"Channel: {channel}", f"Chat ID: {chat_id}"]
|
lines += [f"Channel: {channel}", f"Chat ID: {chat_id}"]
|
||||||
return ContextBuilder._RUNTIME_CONTEXT_TAG + "\n" + "\n".join(lines)
|
return ContextBuilder._RUNTIME_CONTEXT_TAG + "\n" + "\n".join(lines)
|
||||||
@@ -130,7 +133,7 @@ IMPORTANT: To send files (images, documents, audio, video) to the user, you MUST
|
|||||||
current_role: str = "user",
|
current_role: str = "user",
|
||||||
) -> 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."""
|
||||||
runtime_ctx = self._build_runtime_context(channel, chat_id)
|
runtime_ctx = self._build_runtime_context(channel, chat_id, self.timezone)
|
||||||
user_content = self._build_user_content(current_message, media)
|
user_content = self._build_user_content(current_message, media)
|
||||||
|
|
||||||
# Merge runtime context and user content into a single user message
|
# Merge runtime context and user content into a single user message
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
"""Shared lifecycle hook primitives for agent runs."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class AgentHookContext:
|
||||||
|
"""Mutable per-iteration state exposed to runner hooks."""
|
||||||
|
|
||||||
|
iteration: int
|
||||||
|
messages: list[dict[str, Any]]
|
||||||
|
response: LLMResponse | None = None
|
||||||
|
usage: dict[str, int] = field(default_factory=dict)
|
||||||
|
tool_calls: list[ToolCallRequest] = field(default_factory=list)
|
||||||
|
tool_results: list[Any] = field(default_factory=list)
|
||||||
|
tool_events: list[dict[str, str]] = field(default_factory=list)
|
||||||
|
final_content: str | None = None
|
||||||
|
stop_reason: str | None = None
|
||||||
|
error: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class AgentHook:
|
||||||
|
"""Minimal lifecycle surface for shared runner customization."""
|
||||||
|
|
||||||
|
def wants_streaming(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def on_stream(self, context: AgentHookContext, delta: str) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def on_stream_end(self, context: AgentHookContext, *, resuming: bool) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def before_execute_tools(self, context: AgentHookContext) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def after_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def finalize_content(self, context: AgentHookContext, content: str | None) -> str | None:
|
||||||
|
return content
|
||||||
|
|
||||||
|
|
||||||
|
class CompositeHook(AgentHook):
|
||||||
|
"""Fan-out hook that delegates to an ordered list of hooks.
|
||||||
|
|
||||||
|
Error isolation: async methods catch and log per-hook exceptions
|
||||||
|
so a faulty custom hook cannot crash the agent loop.
|
||||||
|
``finalize_content`` is a pipeline (no isolation — bugs should surface).
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_hooks",)
|
||||||
|
|
||||||
|
def __init__(self, hooks: list[AgentHook]) -> None:
|
||||||
|
self._hooks = list(hooks)
|
||||||
|
|
||||||
|
def wants_streaming(self) -> bool:
|
||||||
|
return any(h.wants_streaming() for h in self._hooks)
|
||||||
|
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
for h in self._hooks:
|
||||||
|
try:
|
||||||
|
await h.before_iteration(context)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("AgentHook.before_iteration error in {}", type(h).__name__)
|
||||||
|
|
||||||
|
async def on_stream(self, context: AgentHookContext, delta: str) -> None:
|
||||||
|
for h in self._hooks:
|
||||||
|
try:
|
||||||
|
await h.on_stream(context, delta)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("AgentHook.on_stream error in {}", type(h).__name__)
|
||||||
|
|
||||||
|
async def on_stream_end(self, context: AgentHookContext, *, resuming: bool) -> None:
|
||||||
|
for h in self._hooks:
|
||||||
|
try:
|
||||||
|
await h.on_stream_end(context, resuming=resuming)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("AgentHook.on_stream_end error in {}", type(h).__name__)
|
||||||
|
|
||||||
|
async def before_execute_tools(self, context: AgentHookContext) -> None:
|
||||||
|
for h in self._hooks:
|
||||||
|
try:
|
||||||
|
await h.before_execute_tools(context)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("AgentHook.before_execute_tools error in {}", type(h).__name__)
|
||||||
|
|
||||||
|
async def after_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
for h in self._hooks:
|
||||||
|
try:
|
||||||
|
await h.after_iteration(context)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("AgentHook.after_iteration error in {}", type(h).__name__)
|
||||||
|
|
||||||
|
def finalize_content(self, context: AgentHookContext, content: str | None) -> str | None:
|
||||||
|
for h in self._hooks:
|
||||||
|
content = h.finalize_content(context, content)
|
||||||
|
return content
|
||||||
+164
-120
@@ -14,7 +14,9 @@ from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.agent.context import ContextBuilder
|
from nanobot.agent.context import ContextBuilder
|
||||||
|
from nanobot.agent.hook import AgentHook, AgentHookContext, CompositeHook
|
||||||
from nanobot.agent.memory import MemoryConsolidator
|
from nanobot.agent.memory import MemoryConsolidator
|
||||||
|
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||||
from nanobot.agent.subagent import SubagentManager
|
from nanobot.agent.subagent import SubagentManager
|
||||||
from nanobot.agent.tools.cron import CronTool
|
from nanobot.agent.tools.cron import CronTool
|
||||||
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
||||||
@@ -35,6 +37,111 @@ if TYPE_CHECKING:
|
|||||||
from nanobot.cron.service import CronService
|
from nanobot.cron.service import CronService
|
||||||
|
|
||||||
|
|
||||||
|
class _LoopHook(AgentHook):
|
||||||
|
"""Core lifecycle hook for the main agent loop.
|
||||||
|
|
||||||
|
Handles streaming delta relay, progress reporting, tool-call logging,
|
||||||
|
and think-tag stripping for the built-in agent path.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
agent_loop: AgentLoop,
|
||||||
|
on_progress: Callable[..., Awaitable[None]] | None = None,
|
||||||
|
on_stream: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
on_stream_end: Callable[..., Awaitable[None]] | None = None,
|
||||||
|
*,
|
||||||
|
channel: str = "cli",
|
||||||
|
chat_id: str = "direct",
|
||||||
|
message_id: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
self._loop = agent_loop
|
||||||
|
self._on_progress = on_progress
|
||||||
|
self._on_stream = on_stream
|
||||||
|
self._on_stream_end = on_stream_end
|
||||||
|
self._channel = channel
|
||||||
|
self._chat_id = chat_id
|
||||||
|
self._message_id = message_id
|
||||||
|
self._stream_buf = ""
|
||||||
|
|
||||||
|
def wants_streaming(self) -> bool:
|
||||||
|
return self._on_stream is not None
|
||||||
|
|
||||||
|
async def on_stream(self, context: AgentHookContext, delta: str) -> None:
|
||||||
|
from nanobot.utils.helpers import strip_think
|
||||||
|
|
||||||
|
prev_clean = strip_think(self._stream_buf)
|
||||||
|
self._stream_buf += delta
|
||||||
|
new_clean = strip_think(self._stream_buf)
|
||||||
|
incremental = new_clean[len(prev_clean):]
|
||||||
|
if incremental and self._on_stream:
|
||||||
|
await self._on_stream(incremental)
|
||||||
|
|
||||||
|
async def on_stream_end(self, context: AgentHookContext, *, resuming: bool) -> None:
|
||||||
|
if self._on_stream_end:
|
||||||
|
await self._on_stream_end(resuming=resuming)
|
||||||
|
self._stream_buf = ""
|
||||||
|
|
||||||
|
async def before_execute_tools(self, context: AgentHookContext) -> None:
|
||||||
|
if self._on_progress:
|
||||||
|
if not self._on_stream:
|
||||||
|
thought = self._loop._strip_think(
|
||||||
|
context.response.content if context.response else None
|
||||||
|
)
|
||||||
|
if thought:
|
||||||
|
await self._on_progress(thought)
|
||||||
|
tool_hint = self._loop._strip_think(self._loop._tool_hint(context.tool_calls))
|
||||||
|
await self._on_progress(tool_hint, tool_hint=True)
|
||||||
|
for tc in context.tool_calls:
|
||||||
|
args_str = json.dumps(tc.arguments, ensure_ascii=False)
|
||||||
|
logger.info("Tool call: {}({})", tc.name, args_str[:200])
|
||||||
|
self._loop._set_tool_context(self._channel, self._chat_id, self._message_id)
|
||||||
|
|
||||||
|
def finalize_content(self, context: AgentHookContext, content: str | None) -> str | None:
|
||||||
|
return self._loop._strip_think(content)
|
||||||
|
|
||||||
|
|
||||||
|
class _LoopHookChain(AgentHook):
|
||||||
|
"""Run the core loop hook first, then best-effort extra hooks.
|
||||||
|
|
||||||
|
This preserves the historical failure behavior of ``_LoopHook`` while still
|
||||||
|
letting user-supplied hooks opt into ``CompositeHook`` isolation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_primary", "_extras")
|
||||||
|
|
||||||
|
def __init__(self, primary: AgentHook, extra_hooks: list[AgentHook]) -> None:
|
||||||
|
self._primary = primary
|
||||||
|
self._extras = CompositeHook(extra_hooks)
|
||||||
|
|
||||||
|
def wants_streaming(self) -> bool:
|
||||||
|
return self._primary.wants_streaming() or self._extras.wants_streaming()
|
||||||
|
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
await self._primary.before_iteration(context)
|
||||||
|
await self._extras.before_iteration(context)
|
||||||
|
|
||||||
|
async def on_stream(self, context: AgentHookContext, delta: str) -> None:
|
||||||
|
await self._primary.on_stream(context, delta)
|
||||||
|
await self._extras.on_stream(context, delta)
|
||||||
|
|
||||||
|
async def on_stream_end(self, context: AgentHookContext, *, resuming: bool) -> None:
|
||||||
|
await self._primary.on_stream_end(context, resuming=resuming)
|
||||||
|
await self._extras.on_stream_end(context, resuming=resuming)
|
||||||
|
|
||||||
|
async def before_execute_tools(self, context: AgentHookContext) -> None:
|
||||||
|
await self._primary.before_execute_tools(context)
|
||||||
|
await self._extras.before_execute_tools(context)
|
||||||
|
|
||||||
|
async def after_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
await self._primary.after_iteration(context)
|
||||||
|
await self._extras.after_iteration(context)
|
||||||
|
|
||||||
|
def finalize_content(self, context: AgentHookContext, content: str | None) -> str | None:
|
||||||
|
content = self._primary.finalize_content(context, content)
|
||||||
|
return self._extras.finalize_content(context, content)
|
||||||
|
|
||||||
|
|
||||||
class AgentLoop:
|
class AgentLoop:
|
||||||
"""
|
"""
|
||||||
The agent loop is the core processing engine.
|
The agent loop is the core processing engine.
|
||||||
@@ -65,6 +172,8 @@ class AgentLoop:
|
|||||||
session_manager: SessionManager | None = None,
|
session_manager: SessionManager | None = None,
|
||||||
mcp_servers: dict | None = None,
|
mcp_servers: dict | None = None,
|
||||||
channels_config: ChannelsConfig | None = None,
|
channels_config: ChannelsConfig | None = None,
|
||||||
|
timezone: str | None = None,
|
||||||
|
hooks: list[AgentHook] | None = None,
|
||||||
):
|
):
|
||||||
from nanobot.config.schema import ExecToolConfig, WebSearchConfig
|
from nanobot.config.schema import ExecToolConfig, WebSearchConfig
|
||||||
|
|
||||||
@@ -82,10 +191,12 @@ class AgentLoop:
|
|||||||
self.restrict_to_workspace = restrict_to_workspace
|
self.restrict_to_workspace = restrict_to_workspace
|
||||||
self._start_time = time.time()
|
self._start_time = time.time()
|
||||||
self._last_usage: dict[str, int] = {}
|
self._last_usage: dict[str, int] = {}
|
||||||
|
self._extra_hooks: list[AgentHook] = hooks or []
|
||||||
|
|
||||||
self.context = ContextBuilder(workspace)
|
self.context = ContextBuilder(workspace, timezone=timezone)
|
||||||
self.sessions = session_manager or SessionManager(workspace)
|
self.sessions = session_manager or SessionManager(workspace)
|
||||||
self.tools = ToolRegistry()
|
self.tools = ToolRegistry()
|
||||||
|
self.runner = AgentRunner(provider)
|
||||||
self.subagents = SubagentManager(
|
self.subagents = SubagentManager(
|
||||||
provider=provider,
|
provider=provider,
|
||||||
workspace=workspace,
|
workspace=workspace,
|
||||||
@@ -137,13 +248,16 @@ class AgentLoop:
|
|||||||
timeout=self.exec_config.timeout,
|
timeout=self.exec_config.timeout,
|
||||||
restrict_to_workspace=self.restrict_to_workspace,
|
restrict_to_workspace=self.restrict_to_workspace,
|
||||||
path_append=self.exec_config.path_append,
|
path_append=self.exec_config.path_append,
|
||||||
|
command_wrapper=self.exec_config.command_wrapper,
|
||||||
))
|
))
|
||||||
self.tools.register(WebSearchTool(config=self.web_search_config, proxy=self.web_proxy))
|
self.tools.register(WebSearchTool(config=self.web_search_config, proxy=self.web_proxy))
|
||||||
self.tools.register(WebFetchTool(proxy=self.web_proxy))
|
self.tools.register(WebFetchTool(proxy=self.web_proxy))
|
||||||
self.tools.register(MessageTool(send_callback=self.bus.publish_outbound))
|
self.tools.register(MessageTool(send_callback=self.bus.publish_outbound))
|
||||||
self.tools.register(SpawnTool(manager=self.subagents))
|
self.tools.register(SpawnTool(manager=self.subagents))
|
||||||
if self.cron_service:
|
if self.cron_service:
|
||||||
self.tools.register(CronTool(self.cron_service))
|
self.tools.register(
|
||||||
|
CronTool(self.cron_service, default_timezone=self.context.timezone or "UTC")
|
||||||
|
)
|
||||||
|
|
||||||
async def _connect_mcp(self) -> None:
|
async def _connect_mcp(self) -> None:
|
||||||
"""Connect to configured MCP servers (one-time, lazy)."""
|
"""Connect to configured MCP servers (one-time, lazy)."""
|
||||||
@@ -211,124 +325,36 @@ class AgentLoop:
|
|||||||
``resuming=True`` means tool calls follow (spinner should restart);
|
``resuming=True`` means tool calls follow (spinner should restart);
|
||||||
``resuming=False`` means this is the final response.
|
``resuming=False`` means this is the final response.
|
||||||
"""
|
"""
|
||||||
messages = initial_messages
|
loop_hook = _LoopHook(
|
||||||
iteration = 0
|
self,
|
||||||
final_content = None
|
on_progress=on_progress,
|
||||||
tools_used: list[str] = []
|
on_stream=on_stream,
|
||||||
|
on_stream_end=on_stream_end,
|
||||||
|
channel=channel,
|
||||||
|
chat_id=chat_id,
|
||||||
|
message_id=message_id,
|
||||||
|
)
|
||||||
|
hook: AgentHook = (
|
||||||
|
_LoopHookChain(loop_hook, self._extra_hooks)
|
||||||
|
if self._extra_hooks
|
||||||
|
else loop_hook
|
||||||
|
)
|
||||||
|
|
||||||
# Wrap on_stream with stateful think-tag filter so downstream
|
result = await self.runner.run(AgentRunSpec(
|
||||||
# consumers (CLI, channels) never see <think> blocks.
|
initial_messages=initial_messages,
|
||||||
_raw_stream = on_stream
|
tools=self.tools,
|
||||||
_stream_buf = ""
|
model=self.model,
|
||||||
|
max_iterations=self.max_iterations,
|
||||||
async def _filtered_stream(delta: str) -> None:
|
hook=hook,
|
||||||
nonlocal _stream_buf
|
error_message="Sorry, I encountered an error calling the AI model.",
|
||||||
from nanobot.utils.helpers import strip_think
|
concurrent_tools=True,
|
||||||
prev_clean = strip_think(_stream_buf)
|
))
|
||||||
_stream_buf += delta
|
self._last_usage = result.usage
|
||||||
new_clean = strip_think(_stream_buf)
|
if result.stop_reason == "max_iterations":
|
||||||
incremental = new_clean[len(prev_clean):]
|
|
||||||
if incremental and _raw_stream:
|
|
||||||
await _raw_stream(incremental)
|
|
||||||
|
|
||||||
while iteration < self.max_iterations:
|
|
||||||
iteration += 1
|
|
||||||
|
|
||||||
tool_defs = self.tools.get_definitions()
|
|
||||||
|
|
||||||
if on_stream:
|
|
||||||
response = await self.provider.chat_stream_with_retry(
|
|
||||||
messages=messages,
|
|
||||||
tools=tool_defs,
|
|
||||||
model=self.model,
|
|
||||||
on_content_delta=_filtered_stream,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
response = await self.provider.chat_with_retry(
|
|
||||||
messages=messages,
|
|
||||||
tools=tool_defs,
|
|
||||||
model=self.model,
|
|
||||||
)
|
|
||||||
|
|
||||||
usage = response.usage or {}
|
|
||||||
self._last_usage = {
|
|
||||||
"prompt_tokens": int(usage.get("prompt_tokens", 0) or 0),
|
|
||||||
"completion_tokens": int(usage.get("completion_tokens", 0) or 0),
|
|
||||||
}
|
|
||||||
|
|
||||||
if response.has_tool_calls:
|
|
||||||
if on_stream and on_stream_end:
|
|
||||||
await on_stream_end(resuming=True)
|
|
||||||
_stream_buf = ""
|
|
||||||
|
|
||||||
if on_progress:
|
|
||||||
if not on_stream:
|
|
||||||
thought = self._strip_think(response.content)
|
|
||||||
if thought:
|
|
||||||
await on_progress(thought)
|
|
||||||
tool_hint = self._tool_hint(response.tool_calls)
|
|
||||||
tool_hint = self._strip_think(tool_hint)
|
|
||||||
await on_progress(tool_hint, tool_hint=True)
|
|
||||||
|
|
||||||
tool_call_dicts = [
|
|
||||||
tc.to_openai_tool_call()
|
|
||||||
for tc in response.tool_calls
|
|
||||||
]
|
|
||||||
messages = self.context.add_assistant_message(
|
|
||||||
messages, response.content, tool_call_dicts,
|
|
||||||
reasoning_content=response.reasoning_content,
|
|
||||||
thinking_blocks=response.thinking_blocks,
|
|
||||||
)
|
|
||||||
|
|
||||||
for tc in response.tool_calls:
|
|
||||||
tools_used.append(tc.name)
|
|
||||||
args_str = json.dumps(tc.arguments, ensure_ascii=False)
|
|
||||||
logger.info("Tool call: {}({})", tc.name, args_str[:200])
|
|
||||||
|
|
||||||
# Re-bind tool context right before execution so that
|
|
||||||
# concurrent sessions don't clobber each other's routing.
|
|
||||||
self._set_tool_context(channel, chat_id, message_id)
|
|
||||||
|
|
||||||
# Execute all tool calls concurrently — the LLM batches
|
|
||||||
# independent calls in a single response on purpose.
|
|
||||||
# return_exceptions=True ensures all results are collected
|
|
||||||
# even if one tool is cancelled or raises BaseException.
|
|
||||||
results = await asyncio.gather(*(
|
|
||||||
self.tools.execute(tc.name, tc.arguments)
|
|
||||||
for tc in response.tool_calls
|
|
||||||
), return_exceptions=True)
|
|
||||||
|
|
||||||
for tool_call, result in zip(response.tool_calls, results):
|
|
||||||
if isinstance(result, BaseException):
|
|
||||||
result = f"Error: {type(result).__name__}: {result}"
|
|
||||||
messages = self.context.add_tool_result(
|
|
||||||
messages, tool_call.id, tool_call.name, result
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
if on_stream and on_stream_end:
|
|
||||||
await on_stream_end(resuming=False)
|
|
||||||
_stream_buf = ""
|
|
||||||
|
|
||||||
clean = self._strip_think(response.content)
|
|
||||||
if response.finish_reason == "error":
|
|
||||||
logger.error("LLM returned error: {}", (clean or "")[:200])
|
|
||||||
final_content = clean or "Sorry, I encountered an error calling the AI model."
|
|
||||||
break
|
|
||||||
messages = self.context.add_assistant_message(
|
|
||||||
messages, clean, reasoning_content=response.reasoning_content,
|
|
||||||
thinking_blocks=response.thinking_blocks,
|
|
||||||
)
|
|
||||||
final_content = clean
|
|
||||||
break
|
|
||||||
|
|
||||||
if final_content is None and iteration >= self.max_iterations:
|
|
||||||
logger.warning("Max iterations ({}) reached", self.max_iterations)
|
logger.warning("Max iterations ({}) reached", self.max_iterations)
|
||||||
final_content = (
|
elif result.stop_reason == "error":
|
||||||
f"I reached the maximum number of tool call iterations ({self.max_iterations}) "
|
logger.error("LLM returned error: {}", (result.final_content or "")[:200])
|
||||||
"without completing the task. You can try breaking the task into smaller steps."
|
return result.final_content, result.tools_used, result.messages
|
||||||
)
|
|
||||||
|
|
||||||
return final_content, tools_used, messages
|
|
||||||
|
|
||||||
async def run(self) -> None:
|
async def run(self) -> None:
|
||||||
"""Run the agent loop, dispatching messages as tasks to stay responsive to /stop."""
|
"""Run the agent loop, dispatching messages as tasks to stay responsive to /stop."""
|
||||||
@@ -370,17 +396,35 @@ class AgentLoop:
|
|||||||
try:
|
try:
|
||||||
on_stream = on_stream_end = None
|
on_stream = on_stream_end = None
|
||||||
if msg.metadata.get("_wants_stream"):
|
if msg.metadata.get("_wants_stream"):
|
||||||
|
# Split one answer into distinct stream segments.
|
||||||
|
stream_base_id = f"{msg.session_key}:{time.time_ns()}"
|
||||||
|
stream_segment = 0
|
||||||
|
|
||||||
|
def _current_stream_id() -> str:
|
||||||
|
return f"{stream_base_id}:{stream_segment}"
|
||||||
|
|
||||||
async def on_stream(delta: str) -> None:
|
async def on_stream(delta: str) -> None:
|
||||||
|
meta = dict(msg.metadata or {})
|
||||||
|
meta["_stream_delta"] = True
|
||||||
|
meta["_stream_id"] = _current_stream_id()
|
||||||
await self.bus.publish_outbound(OutboundMessage(
|
await self.bus.publish_outbound(OutboundMessage(
|
||||||
channel=msg.channel, chat_id=msg.chat_id,
|
channel=msg.channel, chat_id=msg.chat_id,
|
||||||
content=delta, metadata={"_stream_delta": True},
|
content=delta,
|
||||||
|
metadata=meta,
|
||||||
))
|
))
|
||||||
|
|
||||||
async def on_stream_end(*, resuming: bool = False) -> None:
|
async def on_stream_end(*, resuming: bool = False) -> None:
|
||||||
|
nonlocal stream_segment
|
||||||
|
meta = dict(msg.metadata or {})
|
||||||
|
meta["_stream_end"] = True
|
||||||
|
meta["_resuming"] = resuming
|
||||||
|
meta["_stream_id"] = _current_stream_id()
|
||||||
await self.bus.publish_outbound(OutboundMessage(
|
await self.bus.publish_outbound(OutboundMessage(
|
||||||
channel=msg.channel, chat_id=msg.chat_id,
|
channel=msg.channel, chat_id=msg.chat_id,
|
||||||
content="", metadata={"_stream_end": True, "_resuming": resuming},
|
content="",
|
||||||
|
metadata=meta,
|
||||||
))
|
))
|
||||||
|
stream_segment += 1
|
||||||
|
|
||||||
response = await self._process_message(
|
response = await self._process_message(
|
||||||
msg, on_stream=on_stream, on_stream_end=on_stream_end,
|
msg, on_stream=on_stream, on_stream_end=on_stream_end,
|
||||||
|
|||||||
@@ -0,0 +1,232 @@
|
|||||||
|
"""Shared execution loop for tool-using agents."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||||
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
|
from nanobot.providers.base import LLMProvider, ToolCallRequest
|
||||||
|
from nanobot.utils.helpers import build_assistant_message
|
||||||
|
|
||||||
|
_DEFAULT_MAX_ITERATIONS_MESSAGE = (
|
||||||
|
"I reached the maximum number of tool call iterations ({max_iterations}) "
|
||||||
|
"without completing the task. You can try breaking the task into smaller steps."
|
||||||
|
)
|
||||||
|
_DEFAULT_ERROR_MESSAGE = "Sorry, I encountered an error calling the AI model."
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class AgentRunSpec:
|
||||||
|
"""Configuration for a single agent execution."""
|
||||||
|
|
||||||
|
initial_messages: list[dict[str, Any]]
|
||||||
|
tools: ToolRegistry
|
||||||
|
model: str
|
||||||
|
max_iterations: int
|
||||||
|
temperature: float | None = None
|
||||||
|
max_tokens: int | None = None
|
||||||
|
reasoning_effort: str | None = None
|
||||||
|
hook: AgentHook | None = None
|
||||||
|
error_message: str | None = _DEFAULT_ERROR_MESSAGE
|
||||||
|
max_iterations_message: str | None = None
|
||||||
|
concurrent_tools: bool = False
|
||||||
|
fail_on_tool_error: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class AgentRunResult:
|
||||||
|
"""Outcome of a shared agent execution."""
|
||||||
|
|
||||||
|
final_content: str | None
|
||||||
|
messages: list[dict[str, Any]]
|
||||||
|
tools_used: list[str] = field(default_factory=list)
|
||||||
|
usage: dict[str, int] = field(default_factory=dict)
|
||||||
|
stop_reason: str = "completed"
|
||||||
|
error: str | None = None
|
||||||
|
tool_events: list[dict[str, str]] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentRunner:
|
||||||
|
"""Run a tool-capable LLM loop without product-layer concerns."""
|
||||||
|
|
||||||
|
def __init__(self, provider: LLMProvider):
|
||||||
|
self.provider = provider
|
||||||
|
|
||||||
|
async def run(self, spec: AgentRunSpec) -> AgentRunResult:
|
||||||
|
hook = spec.hook or AgentHook()
|
||||||
|
messages = list(spec.initial_messages)
|
||||||
|
final_content: str | None = None
|
||||||
|
tools_used: list[str] = []
|
||||||
|
usage = {"prompt_tokens": 0, "completion_tokens": 0}
|
||||||
|
error: str | None = None
|
||||||
|
stop_reason = "completed"
|
||||||
|
tool_events: list[dict[str, str]] = []
|
||||||
|
|
||||||
|
for iteration in range(spec.max_iterations):
|
||||||
|
context = AgentHookContext(iteration=iteration, messages=messages)
|
||||||
|
await hook.before_iteration(context)
|
||||||
|
kwargs: dict[str, Any] = {
|
||||||
|
"messages": messages,
|
||||||
|
"tools": spec.tools.get_definitions(),
|
||||||
|
"model": spec.model,
|
||||||
|
}
|
||||||
|
if spec.temperature is not None:
|
||||||
|
kwargs["temperature"] = spec.temperature
|
||||||
|
if spec.max_tokens is not None:
|
||||||
|
kwargs["max_tokens"] = spec.max_tokens
|
||||||
|
if spec.reasoning_effort is not None:
|
||||||
|
kwargs["reasoning_effort"] = spec.reasoning_effort
|
||||||
|
|
||||||
|
if hook.wants_streaming():
|
||||||
|
async def _stream(delta: str) -> None:
|
||||||
|
await hook.on_stream(context, delta)
|
||||||
|
|
||||||
|
response = await self.provider.chat_stream_with_retry(
|
||||||
|
**kwargs,
|
||||||
|
on_content_delta=_stream,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
response = await self.provider.chat_with_retry(**kwargs)
|
||||||
|
|
||||||
|
raw_usage = response.usage or {}
|
||||||
|
usage = {
|
||||||
|
"prompt_tokens": int(raw_usage.get("prompt_tokens", 0) or 0),
|
||||||
|
"completion_tokens": int(raw_usage.get("completion_tokens", 0) or 0),
|
||||||
|
}
|
||||||
|
context.response = response
|
||||||
|
context.usage = usage
|
||||||
|
context.tool_calls = list(response.tool_calls)
|
||||||
|
|
||||||
|
if response.has_tool_calls:
|
||||||
|
if hook.wants_streaming():
|
||||||
|
await hook.on_stream_end(context, resuming=True)
|
||||||
|
|
||||||
|
messages.append(build_assistant_message(
|
||||||
|
response.content or "",
|
||||||
|
tool_calls=[tc.to_openai_tool_call() for tc in response.tool_calls],
|
||||||
|
reasoning_content=response.reasoning_content,
|
||||||
|
thinking_blocks=response.thinking_blocks,
|
||||||
|
))
|
||||||
|
tools_used.extend(tc.name for tc in response.tool_calls)
|
||||||
|
|
||||||
|
await hook.before_execute_tools(context)
|
||||||
|
|
||||||
|
results, new_events, fatal_error = await self._execute_tools(spec, response.tool_calls)
|
||||||
|
tool_events.extend(new_events)
|
||||||
|
context.tool_results = list(results)
|
||||||
|
context.tool_events = list(new_events)
|
||||||
|
if fatal_error is not None:
|
||||||
|
error = f"Error: {type(fatal_error).__name__}: {fatal_error}"
|
||||||
|
stop_reason = "tool_error"
|
||||||
|
context.error = error
|
||||||
|
context.stop_reason = stop_reason
|
||||||
|
await hook.after_iteration(context)
|
||||||
|
break
|
||||||
|
for tool_call, result in zip(response.tool_calls, results):
|
||||||
|
messages.append({
|
||||||
|
"role": "tool",
|
||||||
|
"tool_call_id": tool_call.id,
|
||||||
|
"name": tool_call.name,
|
||||||
|
"content": result,
|
||||||
|
})
|
||||||
|
await hook.after_iteration(context)
|
||||||
|
continue
|
||||||
|
|
||||||
|
if hook.wants_streaming():
|
||||||
|
await hook.on_stream_end(context, resuming=False)
|
||||||
|
|
||||||
|
clean = hook.finalize_content(context, response.content)
|
||||||
|
if response.finish_reason == "error":
|
||||||
|
final_content = clean or spec.error_message or _DEFAULT_ERROR_MESSAGE
|
||||||
|
stop_reason = "error"
|
||||||
|
error = final_content
|
||||||
|
context.final_content = final_content
|
||||||
|
context.error = error
|
||||||
|
context.stop_reason = stop_reason
|
||||||
|
await hook.after_iteration(context)
|
||||||
|
break
|
||||||
|
|
||||||
|
messages.append(build_assistant_message(
|
||||||
|
clean,
|
||||||
|
reasoning_content=response.reasoning_content,
|
||||||
|
thinking_blocks=response.thinking_blocks,
|
||||||
|
))
|
||||||
|
final_content = clean
|
||||||
|
context.final_content = final_content
|
||||||
|
context.stop_reason = stop_reason
|
||||||
|
await hook.after_iteration(context)
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
stop_reason = "max_iterations"
|
||||||
|
template = spec.max_iterations_message or _DEFAULT_MAX_ITERATIONS_MESSAGE
|
||||||
|
final_content = template.format(max_iterations=spec.max_iterations)
|
||||||
|
|
||||||
|
return AgentRunResult(
|
||||||
|
final_content=final_content,
|
||||||
|
messages=messages,
|
||||||
|
tools_used=tools_used,
|
||||||
|
usage=usage,
|
||||||
|
stop_reason=stop_reason,
|
||||||
|
error=error,
|
||||||
|
tool_events=tool_events,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _execute_tools(
|
||||||
|
self,
|
||||||
|
spec: AgentRunSpec,
|
||||||
|
tool_calls: list[ToolCallRequest],
|
||||||
|
) -> tuple[list[Any], list[dict[str, str]], BaseException | None]:
|
||||||
|
if spec.concurrent_tools:
|
||||||
|
tool_results = await asyncio.gather(*(
|
||||||
|
self._run_tool(spec, tool_call)
|
||||||
|
for tool_call in tool_calls
|
||||||
|
))
|
||||||
|
else:
|
||||||
|
tool_results = [
|
||||||
|
await self._run_tool(spec, tool_call)
|
||||||
|
for tool_call in tool_calls
|
||||||
|
]
|
||||||
|
|
||||||
|
results: list[Any] = []
|
||||||
|
events: list[dict[str, str]] = []
|
||||||
|
fatal_error: BaseException | None = None
|
||||||
|
for result, event, error in tool_results:
|
||||||
|
results.append(result)
|
||||||
|
events.append(event)
|
||||||
|
if error is not None and fatal_error is None:
|
||||||
|
fatal_error = error
|
||||||
|
return results, events, fatal_error
|
||||||
|
|
||||||
|
async def _run_tool(
|
||||||
|
self,
|
||||||
|
spec: AgentRunSpec,
|
||||||
|
tool_call: ToolCallRequest,
|
||||||
|
) -> tuple[Any, dict[str, str], BaseException | None]:
|
||||||
|
try:
|
||||||
|
result = await spec.tools.execute(tool_call.name, tool_call.arguments)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except BaseException as exc:
|
||||||
|
event = {
|
||||||
|
"name": tool_call.name,
|
||||||
|
"status": "error",
|
||||||
|
"detail": str(exc),
|
||||||
|
}
|
||||||
|
if spec.fail_on_tool_error:
|
||||||
|
return f"Error: {type(exc).__name__}: {exc}", event, exc
|
||||||
|
return f"Error: {type(exc).__name__}: {exc}", event, None
|
||||||
|
|
||||||
|
detail = "" if result is None else str(result)
|
||||||
|
detail = detail.replace("\n", " ").strip()
|
||||||
|
if not detail:
|
||||||
|
detail = "(empty)"
|
||||||
|
elif len(detail) > 120:
|
||||||
|
detail = detail[:120] + "..."
|
||||||
|
return result, {
|
||||||
|
"name": tool_call.name,
|
||||||
|
"status": "error" if isinstance(result, str) and result.startswith("Error") else "ok",
|
||||||
|
"detail": detail,
|
||||||
|
}, None
|
||||||
+79
-51
@@ -8,6 +8,8 @@ from typing import Any
|
|||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
|
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||||
|
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||||
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
||||||
from nanobot.agent.tools.filesystem import EditFileTool, ListDirTool, ReadFileTool, WriteFileTool
|
from nanobot.agent.tools.filesystem import EditFileTool, ListDirTool, ReadFileTool, WriteFileTool
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
@@ -17,7 +19,21 @@ from nanobot.bus.events import InboundMessage
|
|||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.config.schema import ExecToolConfig
|
from nanobot.config.schema import ExecToolConfig
|
||||||
from nanobot.providers.base import LLMProvider
|
from nanobot.providers.base import LLMProvider
|
||||||
from nanobot.utils.helpers import build_assistant_message
|
|
||||||
|
|
||||||
|
class _SubagentHook(AgentHook):
|
||||||
|
"""Logging-only hook for subagent execution."""
|
||||||
|
|
||||||
|
def __init__(self, task_id: str) -> None:
|
||||||
|
self._task_id = task_id
|
||||||
|
|
||||||
|
async def before_execute_tools(self, context: AgentHookContext) -> None:
|
||||||
|
for tool_call in context.tool_calls:
|
||||||
|
args_str = json.dumps(tool_call.arguments, ensure_ascii=False)
|
||||||
|
logger.debug(
|
||||||
|
"Subagent [{}] executing: {} with arguments: {}",
|
||||||
|
self._task_id, tool_call.name, args_str,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class SubagentManager:
|
class SubagentManager:
|
||||||
@@ -44,6 +60,7 @@ class SubagentManager:
|
|||||||
self.web_proxy = web_proxy
|
self.web_proxy = web_proxy
|
||||||
self.exec_config = exec_config or ExecToolConfig()
|
self.exec_config = exec_config or ExecToolConfig()
|
||||||
self.restrict_to_workspace = restrict_to_workspace
|
self.restrict_to_workspace = restrict_to_workspace
|
||||||
|
self.runner = AgentRunner(provider)
|
||||||
self._running_tasks: dict[str, asyncio.Task[None]] = {}
|
self._running_tasks: dict[str, asyncio.Task[None]] = {}
|
||||||
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
|
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
|
||||||
|
|
||||||
@@ -98,64 +115,54 @@ class SubagentManager:
|
|||||||
tools.register(WriteFileTool(workspace=self.workspace, allowed_dir=allowed_dir))
|
tools.register(WriteFileTool(workspace=self.workspace, allowed_dir=allowed_dir))
|
||||||
tools.register(EditFileTool(workspace=self.workspace, allowed_dir=allowed_dir))
|
tools.register(EditFileTool(workspace=self.workspace, allowed_dir=allowed_dir))
|
||||||
tools.register(ListDirTool(workspace=self.workspace, allowed_dir=allowed_dir))
|
tools.register(ListDirTool(workspace=self.workspace, allowed_dir=allowed_dir))
|
||||||
tools.register(ExecTool(
|
if self.exec_config.enable:
|
||||||
working_dir=str(self.workspace),
|
tools.register(ExecTool(
|
||||||
timeout=self.exec_config.timeout,
|
working_dir=str(self.workspace),
|
||||||
restrict_to_workspace=self.restrict_to_workspace,
|
timeout=self.exec_config.timeout,
|
||||||
path_append=self.exec_config.path_append,
|
restrict_to_workspace=self.restrict_to_workspace,
|
||||||
))
|
path_append=self.exec_config.path_append,
|
||||||
|
command_wrapper=self.exec_config.command_wrapper,
|
||||||
|
))
|
||||||
tools.register(WebSearchTool(config=self.web_search_config, proxy=self.web_proxy))
|
tools.register(WebSearchTool(config=self.web_search_config, proxy=self.web_proxy))
|
||||||
tools.register(WebFetchTool(proxy=self.web_proxy))
|
tools.register(WebFetchTool(proxy=self.web_proxy))
|
||||||
|
|
||||||
system_prompt = self._build_subagent_prompt()
|
system_prompt = self._build_subagent_prompt()
|
||||||
messages: list[dict[str, Any]] = [
|
messages: list[dict[str, Any]] = [
|
||||||
{"role": "system", "content": system_prompt},
|
{"role": "system", "content": system_prompt},
|
||||||
{"role": "user", "content": task},
|
{"role": "user", "content": task},
|
||||||
]
|
]
|
||||||
|
|
||||||
# Run agent loop (limited iterations)
|
result = await self.runner.run(AgentRunSpec(
|
||||||
max_iterations = 15
|
initial_messages=messages,
|
||||||
iteration = 0
|
tools=tools,
|
||||||
final_result: str | None = None
|
model=self.model,
|
||||||
|
max_iterations=15,
|
||||||
while iteration < max_iterations:
|
hook=_SubagentHook(task_id),
|
||||||
iteration += 1
|
max_iterations_message="Task completed but no final response was generated.",
|
||||||
|
error_message=None,
|
||||||
response = await self.provider.chat_with_retry(
|
fail_on_tool_error=True,
|
||||||
messages=messages,
|
))
|
||||||
tools=tools.get_definitions(),
|
if result.stop_reason == "tool_error":
|
||||||
model=self.model,
|
await self._announce_result(
|
||||||
|
task_id,
|
||||||
|
label,
|
||||||
|
task,
|
||||||
|
self._format_partial_progress(result),
|
||||||
|
origin,
|
||||||
|
"error",
|
||||||
)
|
)
|
||||||
|
return
|
||||||
if response.has_tool_calls:
|
if result.stop_reason == "error":
|
||||||
tool_call_dicts = [
|
await self._announce_result(
|
||||||
tc.to_openai_tool_call()
|
task_id,
|
||||||
for tc in response.tool_calls
|
label,
|
||||||
]
|
task,
|
||||||
messages.append(build_assistant_message(
|
result.error or "Error: subagent execution failed.",
|
||||||
response.content or "",
|
origin,
|
||||||
tool_calls=tool_call_dicts,
|
"error",
|
||||||
reasoning_content=response.reasoning_content,
|
)
|
||||||
thinking_blocks=response.thinking_blocks,
|
return
|
||||||
))
|
final_result = result.final_content or "Task completed but no final response was generated."
|
||||||
|
|
||||||
# Execute tools
|
|
||||||
for tool_call in response.tool_calls:
|
|
||||||
args_str = json.dumps(tool_call.arguments, ensure_ascii=False)
|
|
||||||
logger.debug("Subagent [{}] executing: {} with arguments: {}", task_id, tool_call.name, args_str)
|
|
||||||
result = await tools.execute(tool_call.name, tool_call.arguments)
|
|
||||||
messages.append({
|
|
||||||
"role": "tool",
|
|
||||||
"tool_call_id": tool_call.id,
|
|
||||||
"name": tool_call.name,
|
|
||||||
"content": result,
|
|
||||||
})
|
|
||||||
else:
|
|
||||||
final_result = response.content
|
|
||||||
break
|
|
||||||
|
|
||||||
if final_result is None:
|
|
||||||
final_result = "Task completed but no final response was generated."
|
|
||||||
|
|
||||||
logger.info("Subagent [{}] completed successfully", task_id)
|
logger.info("Subagent [{}] completed successfully", task_id)
|
||||||
await self._announce_result(task_id, label, task, final_result, origin, "ok")
|
await self._announce_result(task_id, label, task, final_result, origin, "ok")
|
||||||
@@ -196,7 +203,28 @@ Summarize this naturally for the user. Keep it brief (1-2 sentences). Do not men
|
|||||||
|
|
||||||
await self.bus.publish_inbound(msg)
|
await self.bus.publish_inbound(msg)
|
||||||
logger.debug("Subagent [{}] announced result to {}:{}", task_id, origin['channel'], origin['chat_id'])
|
logger.debug("Subagent [{}] announced result to {}:{}", task_id, origin['channel'], origin['chat_id'])
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _format_partial_progress(result) -> str:
|
||||||
|
completed = [e for e in result.tool_events if e["status"] == "ok"]
|
||||||
|
failure = next((e for e in reversed(result.tool_events) if e["status"] == "error"), None)
|
||||||
|
lines: list[str] = []
|
||||||
|
if completed:
|
||||||
|
lines.append("Completed steps:")
|
||||||
|
for event in completed[-3:]:
|
||||||
|
lines.append(f"- {event['name']}: {event['detail']}")
|
||||||
|
if failure:
|
||||||
|
if lines:
|
||||||
|
lines.append("")
|
||||||
|
lines.append("Failure:")
|
||||||
|
lines.append(f"- {failure['name']}: {failure['detail']}")
|
||||||
|
if result.error and not failure:
|
||||||
|
if lines:
|
||||||
|
lines.append("")
|
||||||
|
lines.append("Failure:")
|
||||||
|
lines.append(f"- {result.error}")
|
||||||
|
return "\n".join(lines) or (result.error or "Error: subagent execution failed.")
|
||||||
|
|
||||||
def _build_subagent_prompt(self) -> str:
|
def _build_subagent_prompt(self) -> str:
|
||||||
"""Build a focused system prompt for the subagent."""
|
"""Build a focused system prompt for the subagent."""
|
||||||
from nanobot.agent.context import ContextBuilder
|
from nanobot.agent.context import ContextBuilder
|
||||||
|
|||||||
+58
-25
@@ -1,7 +1,7 @@
|
|||||||
"""Cron tool for scheduling reminders and tasks."""
|
"""Cron tool for scheduling reminders and tasks."""
|
||||||
|
|
||||||
from contextvars import ContextVar
|
from contextvars import ContextVar
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from nanobot.agent.tools.base import Tool
|
from nanobot.agent.tools.base import Tool
|
||||||
@@ -12,8 +12,9 @@ from nanobot.cron.types import CronJobState, CronSchedule
|
|||||||
class CronTool(Tool):
|
class CronTool(Tool):
|
||||||
"""Tool to schedule reminders and recurring tasks."""
|
"""Tool to schedule reminders and recurring tasks."""
|
||||||
|
|
||||||
def __init__(self, cron_service: CronService):
|
def __init__(self, cron_service: CronService, default_timezone: str = "UTC"):
|
||||||
self._cron = cron_service
|
self._cron = cron_service
|
||||||
|
self._default_timezone = default_timezone
|
||||||
self._channel = ""
|
self._channel = ""
|
||||||
self._chat_id = ""
|
self._chat_id = ""
|
||||||
self._in_cron_context: ContextVar[bool] = ContextVar("cron_in_context", default=False)
|
self._in_cron_context: ContextVar[bool] = ContextVar("cron_in_context", default=False)
|
||||||
@@ -31,13 +32,37 @@ class CronTool(Tool):
|
|||||||
"""Restore previous cron context."""
|
"""Restore previous cron context."""
|
||||||
self._in_cron_context.reset(token)
|
self._in_cron_context.reset(token)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _validate_timezone(tz: str) -> str | None:
|
||||||
|
from zoneinfo import ZoneInfo
|
||||||
|
|
||||||
|
try:
|
||||||
|
ZoneInfo(tz)
|
||||||
|
except (KeyError, Exception):
|
||||||
|
return f"Error: unknown timezone '{tz}'"
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _display_timezone(self, schedule: CronSchedule) -> str:
|
||||||
|
"""Pick the most human-meaningful timezone for display."""
|
||||||
|
return schedule.tz or self._default_timezone
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _format_timestamp(ms: int, tz_name: str) -> str:
|
||||||
|
from zoneinfo import ZoneInfo
|
||||||
|
|
||||||
|
dt = datetime.fromtimestamp(ms / 1000, tz=ZoneInfo(tz_name))
|
||||||
|
return f"{dt.isoformat()} ({tz_name})"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def name(self) -> str:
|
def name(self) -> str:
|
||||||
return "cron"
|
return "cron"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def description(self) -> str:
|
def description(self) -> str:
|
||||||
return "Schedule reminders and recurring tasks. Actions: add, list, remove."
|
return (
|
||||||
|
"Schedule reminders and recurring tasks. Actions: add, list, remove. "
|
||||||
|
f"If tz is omitted, cron expressions and naive ISO times default to {self._default_timezone}."
|
||||||
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def parameters(self) -> dict[str, Any]:
|
def parameters(self) -> dict[str, Any]:
|
||||||
@@ -49,7 +74,7 @@ class CronTool(Tool):
|
|||||||
"enum": ["add", "list", "remove"],
|
"enum": ["add", "list", "remove"],
|
||||||
"description": "Action to perform",
|
"description": "Action to perform",
|
||||||
},
|
},
|
||||||
"message": {"type": "string", "description": "Reminder message (for add)"},
|
"message": {"type": "string", "description": "Instruction for the agent to execute when the job triggers (e.g., 'Send a reminder to WeChat: xxx' or 'Check system status and report')"},
|
||||||
"every_seconds": {
|
"every_seconds": {
|
||||||
"type": "integer",
|
"type": "integer",
|
||||||
"description": "Interval in seconds (for recurring tasks)",
|
"description": "Interval in seconds (for recurring tasks)",
|
||||||
@@ -60,11 +85,17 @@ class CronTool(Tool):
|
|||||||
},
|
},
|
||||||
"tz": {
|
"tz": {
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "IANA timezone for cron expressions (e.g. 'America/Vancouver')",
|
"description": (
|
||||||
|
"Optional IANA timezone for cron expressions "
|
||||||
|
f"(e.g. 'America/Vancouver'). Defaults to {self._default_timezone}."
|
||||||
|
),
|
||||||
},
|
},
|
||||||
"at": {
|
"at": {
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "ISO datetime for one-time execution (e.g. '2026-02-12T10:30:00')",
|
"description": (
|
||||||
|
"ISO datetime for one-time execution "
|
||||||
|
f"(e.g. '2026-02-12T10:30:00'). Naive values default to {self._default_timezone}."
|
||||||
|
),
|
||||||
},
|
},
|
||||||
"job_id": {"type": "string", "description": "Job ID (for remove)"},
|
"job_id": {"type": "string", "description": "Job ID (for remove)"},
|
||||||
},
|
},
|
||||||
@@ -107,26 +138,29 @@ class CronTool(Tool):
|
|||||||
if tz and not cron_expr:
|
if tz and not cron_expr:
|
||||||
return "Error: tz can only be used with cron_expr"
|
return "Error: tz can only be used with cron_expr"
|
||||||
if tz:
|
if tz:
|
||||||
from zoneinfo import ZoneInfo
|
if err := self._validate_timezone(tz):
|
||||||
|
return err
|
||||||
try:
|
|
||||||
ZoneInfo(tz)
|
|
||||||
except (KeyError, Exception):
|
|
||||||
return f"Error: unknown timezone '{tz}'"
|
|
||||||
|
|
||||||
# Build schedule
|
# Build schedule
|
||||||
delete_after = False
|
delete_after = False
|
||||||
if every_seconds:
|
if every_seconds:
|
||||||
schedule = CronSchedule(kind="every", every_ms=every_seconds * 1000)
|
schedule = CronSchedule(kind="every", every_ms=every_seconds * 1000)
|
||||||
elif cron_expr:
|
elif cron_expr:
|
||||||
schedule = CronSchedule(kind="cron", expr=cron_expr, tz=tz)
|
effective_tz = tz or self._default_timezone
|
||||||
|
if err := self._validate_timezone(effective_tz):
|
||||||
|
return err
|
||||||
|
schedule = CronSchedule(kind="cron", expr=cron_expr, tz=effective_tz)
|
||||||
elif at:
|
elif at:
|
||||||
from datetime import datetime
|
from zoneinfo import ZoneInfo
|
||||||
|
|
||||||
try:
|
try:
|
||||||
dt = datetime.fromisoformat(at)
|
dt = datetime.fromisoformat(at)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
return f"Error: invalid ISO datetime format '{at}'. Expected format: YYYY-MM-DDTHH:MM:SS"
|
return f"Error: invalid ISO datetime format '{at}'. Expected format: YYYY-MM-DDTHH:MM:SS"
|
||||||
|
if dt.tzinfo is None:
|
||||||
|
if err := self._validate_timezone(self._default_timezone):
|
||||||
|
return err
|
||||||
|
dt = dt.replace(tzinfo=ZoneInfo(self._default_timezone))
|
||||||
at_ms = int(dt.timestamp() * 1000)
|
at_ms = int(dt.timestamp() * 1000)
|
||||||
schedule = CronSchedule(kind="at", at_ms=at_ms)
|
schedule = CronSchedule(kind="at", at_ms=at_ms)
|
||||||
delete_after = True
|
delete_after = True
|
||||||
@@ -144,8 +178,7 @@ class CronTool(Tool):
|
|||||||
)
|
)
|
||||||
return f"Created job '{job.name}' (id: {job.id})"
|
return f"Created job '{job.name}' (id: {job.id})"
|
||||||
|
|
||||||
@staticmethod
|
def _format_timing(self, schedule: CronSchedule) -> str:
|
||||||
def _format_timing(schedule: CronSchedule) -> str:
|
|
||||||
"""Format schedule as a human-readable timing string."""
|
"""Format schedule as a human-readable timing string."""
|
||||||
if schedule.kind == "cron":
|
if schedule.kind == "cron":
|
||||||
tz = f" ({schedule.tz})" if schedule.tz else ""
|
tz = f" ({schedule.tz})" if schedule.tz else ""
|
||||||
@@ -160,23 +193,23 @@ class CronTool(Tool):
|
|||||||
return f"every {ms // 1000}s"
|
return f"every {ms // 1000}s"
|
||||||
return f"every {ms}ms"
|
return f"every {ms}ms"
|
||||||
if schedule.kind == "at" and schedule.at_ms:
|
if schedule.kind == "at" and schedule.at_ms:
|
||||||
dt = datetime.fromtimestamp(schedule.at_ms / 1000, tz=timezone.utc)
|
return f"at {self._format_timestamp(schedule.at_ms, self._display_timezone(schedule))}"
|
||||||
return f"at {dt.isoformat()}"
|
|
||||||
return schedule.kind
|
return schedule.kind
|
||||||
|
|
||||||
@staticmethod
|
def _format_state(self, state: CronJobState, schedule: CronSchedule) -> list[str]:
|
||||||
def _format_state(state: CronJobState) -> list[str]:
|
|
||||||
"""Format job run state as display lines."""
|
"""Format job run state as display lines."""
|
||||||
lines: list[str] = []
|
lines: list[str] = []
|
||||||
|
display_tz = self._display_timezone(schedule)
|
||||||
if state.last_run_at_ms:
|
if state.last_run_at_ms:
|
||||||
last_dt = datetime.fromtimestamp(state.last_run_at_ms / 1000, tz=timezone.utc)
|
info = (
|
||||||
info = f" Last run: {last_dt.isoformat()} — {state.last_status or 'unknown'}"
|
f" Last run: {self._format_timestamp(state.last_run_at_ms, display_tz)}"
|
||||||
|
f" — {state.last_status or 'unknown'}"
|
||||||
|
)
|
||||||
if state.last_error:
|
if state.last_error:
|
||||||
info += f" ({state.last_error})"
|
info += f" ({state.last_error})"
|
||||||
lines.append(info)
|
lines.append(info)
|
||||||
if state.next_run_at_ms:
|
if state.next_run_at_ms:
|
||||||
next_dt = datetime.fromtimestamp(state.next_run_at_ms / 1000, tz=timezone.utc)
|
lines.append(f" Next run: {self._format_timestamp(state.next_run_at_ms, display_tz)}")
|
||||||
lines.append(f" Next run: {next_dt.isoformat()}")
|
|
||||||
return lines
|
return lines
|
||||||
|
|
||||||
def _list_jobs(self) -> str:
|
def _list_jobs(self) -> str:
|
||||||
@@ -187,7 +220,7 @@ class CronTool(Tool):
|
|||||||
for j in jobs:
|
for j in jobs:
|
||||||
timing = self._format_timing(j.schedule)
|
timing = self._format_timing(j.schedule)
|
||||||
parts = [f"- {j.name} (id: {j.id}, {timing})"]
|
parts = [f"- {j.name} (id: {j.id}, {timing})"]
|
||||||
parts.extend(self._format_state(j.state))
|
parts.extend(self._format_state(j.state, j.schedule))
|
||||||
lines.append("\n".join(parts))
|
lines.append("\n".join(parts))
|
||||||
return "Scheduled jobs:\n" + "\n".join(lines)
|
return "Scheduled jobs:\n" + "\n".join(lines)
|
||||||
|
|
||||||
|
|||||||
@@ -170,7 +170,11 @@ async def connect_mcp_servers(
|
|||||||
timeout: httpx.Timeout | None = None,
|
timeout: httpx.Timeout | None = None,
|
||||||
auth: httpx.Auth | None = None,
|
auth: httpx.Auth | None = None,
|
||||||
) -> httpx.AsyncClient:
|
) -> httpx.AsyncClient:
|
||||||
merged_headers = {**(cfg.headers or {}), **(headers or {})}
|
merged_headers = {
|
||||||
|
"Accept": "application/json, text/event-stream",
|
||||||
|
**(cfg.headers or {}),
|
||||||
|
**(headers or {}),
|
||||||
|
}
|
||||||
return httpx.AsyncClient(
|
return httpx.AsyncClient(
|
||||||
headers=merged_headers or None,
|
headers=merged_headers or None,
|
||||||
follow_redirects=True,
|
follow_redirects=True,
|
||||||
|
|||||||
@@ -23,9 +23,11 @@ class ExecTool(Tool):
|
|||||||
allow_patterns: list[str] | None = None,
|
allow_patterns: list[str] | None = None,
|
||||||
restrict_to_workspace: bool = False,
|
restrict_to_workspace: bool = False,
|
||||||
path_append: str = "",
|
path_append: str = "",
|
||||||
|
command_wrapper: str = "",
|
||||||
):
|
):
|
||||||
self.timeout = timeout
|
self.timeout = timeout
|
||||||
self.working_dir = working_dir
|
self.working_dir = working_dir
|
||||||
|
self.command_wrapper = command_wrapper
|
||||||
self.deny_patterns = deny_patterns or [
|
self.deny_patterns = deny_patterns or [
|
||||||
r"\brm\s+-[rf]{1,2}\b", # rm -r, rm -rf, rm -fr
|
r"\brm\s+-[rf]{1,2}\b", # rm -r, rm -rf, rm -fr
|
||||||
r"\bdel\s+/[fq]\b", # del /f, del /q
|
r"\bdel\s+/[fq]\b", # del /f, del /q
|
||||||
@@ -82,11 +84,16 @@ class ExecTool(Tool):
|
|||||||
self, command: str, working_dir: str | None = None,
|
self, command: str, working_dir: str | None = None,
|
||||||
timeout: int | None = None, **kwargs: Any,
|
timeout: int | None = None, **kwargs: Any,
|
||||||
) -> str:
|
) -> str:
|
||||||
cwd = working_dir or self.working_dir or os.getcwd()
|
cwd = os.path.abspath(working_dir or self.working_dir or os.getcwd())
|
||||||
guard_error = self._guard_command(command, cwd)
|
guard_error = self._guard_command(command, cwd)
|
||||||
if guard_error:
|
if guard_error:
|
||||||
return guard_error
|
return guard_error
|
||||||
|
|
||||||
|
if self.command_wrapper:
|
||||||
|
original_command = command
|
||||||
|
command = self.command_wrapper.replace("{cwd}", cwd).replace("{command}", command)
|
||||||
|
logger.debug("command_wrapper applied: {} -> {}", original_command, command)
|
||||||
|
|
||||||
effective_timeout = min(timeout or self.timeout, self._MAX_TIMEOUT)
|
effective_timeout = min(timeout or self.timeout, self._MAX_TIMEOUT)
|
||||||
|
|
||||||
env = os.environ.copy()
|
env = os.environ.copy()
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
"""OpenAI-compatible HTTP API for nanobot."""
|
||||||
@@ -0,0 +1,193 @@
|
|||||||
|
"""OpenAI-compatible HTTP API server for a fixed nanobot session.
|
||||||
|
|
||||||
|
Provides /v1/chat/completions and /v1/models endpoints.
|
||||||
|
All requests route to a single persistent API session.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
API_SESSION_KEY = "api:default"
|
||||||
|
API_CHAT_ID = "default"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Response helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _error_json(status: int, message: str, err_type: str = "invalid_request_error") -> web.Response:
|
||||||
|
return web.json_response(
|
||||||
|
{"error": {"message": message, "type": err_type, "code": status}},
|
||||||
|
status=status,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _chat_completion_response(content: str, model: str) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"id": f"chatcmpl-{uuid.uuid4().hex[:12]}",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": int(time.time()),
|
||||||
|
"model": model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"message": {"role": "assistant", "content": content},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _response_text(value: Any) -> str:
|
||||||
|
"""Normalize process_direct output to plain assistant text."""
|
||||||
|
if value is None:
|
||||||
|
return ""
|
||||||
|
if hasattr(value, "content"):
|
||||||
|
return str(getattr(value, "content") or "")
|
||||||
|
return str(value)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Route handlers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def handle_chat_completions(request: web.Request) -> web.Response:
|
||||||
|
"""POST /v1/chat/completions"""
|
||||||
|
|
||||||
|
# --- Parse body ---
|
||||||
|
try:
|
||||||
|
body = await request.json()
|
||||||
|
except Exception:
|
||||||
|
return _error_json(400, "Invalid JSON body")
|
||||||
|
|
||||||
|
messages = body.get("messages")
|
||||||
|
if not isinstance(messages, list) or len(messages) != 1:
|
||||||
|
return _error_json(400, "Only a single user message is supported")
|
||||||
|
|
||||||
|
# Stream not yet supported
|
||||||
|
if body.get("stream", False):
|
||||||
|
return _error_json(400, "stream=true is not supported yet. Set stream=false or omit it.")
|
||||||
|
|
||||||
|
message = messages[0]
|
||||||
|
if not isinstance(message, dict) or message.get("role") != "user":
|
||||||
|
return _error_json(400, "Only a single user message is supported")
|
||||||
|
user_content = message.get("content", "")
|
||||||
|
if isinstance(user_content, list):
|
||||||
|
# Multi-modal content array — extract text parts
|
||||||
|
user_content = " ".join(
|
||||||
|
part.get("text", "") for part in user_content if part.get("type") == "text"
|
||||||
|
)
|
||||||
|
|
||||||
|
agent_loop = request.app["agent_loop"]
|
||||||
|
timeout_s: float = request.app.get("request_timeout", 120.0)
|
||||||
|
model_name: str = request.app.get("model_name", "nanobot")
|
||||||
|
if (requested_model := body.get("model")) and requested_model != model_name:
|
||||||
|
return _error_json(400, f"Only configured model '{model_name}' is available")
|
||||||
|
|
||||||
|
session_key = f"api:{body['session_id']}" if body.get("session_id") else API_SESSION_KEY
|
||||||
|
session_locks: dict[str, asyncio.Lock] = request.app["session_locks"]
|
||||||
|
session_lock = session_locks.setdefault(session_key, asyncio.Lock())
|
||||||
|
|
||||||
|
logger.info("API request session_key={} content={}", session_key, user_content[:80])
|
||||||
|
|
||||||
|
_FALLBACK = "I've completed processing but have no response to give."
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with session_lock:
|
||||||
|
try:
|
||||||
|
response = await asyncio.wait_for(
|
||||||
|
agent_loop.process_direct(
|
||||||
|
content=user_content,
|
||||||
|
session_key=session_key,
|
||||||
|
channel="api",
|
||||||
|
chat_id=API_CHAT_ID,
|
||||||
|
),
|
||||||
|
timeout=timeout_s,
|
||||||
|
)
|
||||||
|
response_text = _response_text(response)
|
||||||
|
|
||||||
|
if not response_text or not response_text.strip():
|
||||||
|
logger.warning(
|
||||||
|
"Empty response for session {}, retrying",
|
||||||
|
session_key,
|
||||||
|
)
|
||||||
|
retry_response = await asyncio.wait_for(
|
||||||
|
agent_loop.process_direct(
|
||||||
|
content=user_content,
|
||||||
|
session_key=session_key,
|
||||||
|
channel="api",
|
||||||
|
chat_id=API_CHAT_ID,
|
||||||
|
),
|
||||||
|
timeout=timeout_s,
|
||||||
|
)
|
||||||
|
response_text = _response_text(retry_response)
|
||||||
|
if not response_text or not response_text.strip():
|
||||||
|
logger.warning(
|
||||||
|
"Empty response after retry for session {}, using fallback",
|
||||||
|
session_key,
|
||||||
|
)
|
||||||
|
response_text = _FALLBACK
|
||||||
|
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
return _error_json(504, f"Request timed out after {timeout_s}s")
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Error processing request for session {}", session_key)
|
||||||
|
return _error_json(500, "Internal server error", err_type="server_error")
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Unexpected API lock error for session {}", session_key)
|
||||||
|
return _error_json(500, "Internal server error", err_type="server_error")
|
||||||
|
|
||||||
|
return web.json_response(_chat_completion_response(response_text, model_name))
|
||||||
|
|
||||||
|
|
||||||
|
async def handle_models(request: web.Request) -> web.Response:
|
||||||
|
"""GET /v1/models"""
|
||||||
|
model_name = request.app.get("model_name", "nanobot")
|
||||||
|
return web.json_response({
|
||||||
|
"object": "list",
|
||||||
|
"data": [
|
||||||
|
{
|
||||||
|
"id": model_name,
|
||||||
|
"object": "model",
|
||||||
|
"created": 0,
|
||||||
|
"owned_by": "nanobot",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
async def handle_health(request: web.Request) -> web.Response:
|
||||||
|
"""GET /health"""
|
||||||
|
return web.json_response({"status": "ok"})
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# App factory
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def create_app(agent_loop, model_name: str = "nanobot", request_timeout: float = 120.0) -> web.Application:
|
||||||
|
"""Create the aiohttp application.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
agent_loop: An initialized AgentLoop instance.
|
||||||
|
model_name: Model name reported in responses.
|
||||||
|
request_timeout: Per-request timeout in seconds.
|
||||||
|
"""
|
||||||
|
app = web.Application()
|
||||||
|
app["agent_loop"] = agent_loop
|
||||||
|
app["model_name"] = model_name
|
||||||
|
app["request_timeout"] = request_timeout
|
||||||
|
app["session_locks"] = {} # per-user locks, keyed by session_key
|
||||||
|
|
||||||
|
app.router.add_post("/v1/chat/completions", handle_chat_completions)
|
||||||
|
app.router.add_get("/v1/models", handle_models)
|
||||||
|
app.router.add_get("/health", handle_health)
|
||||||
|
return app
|
||||||
@@ -85,11 +85,22 @@ class BaseChannel(ABC):
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
msg: The message to send.
|
msg: The message to send.
|
||||||
|
|
||||||
|
Implementations should raise on delivery failure so the channel manager
|
||||||
|
can apply any retry policy in one place.
|
||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||||
"""Deliver a streaming text chunk. Override in subclass to enable streaming."""
|
"""Deliver a streaming text chunk.
|
||||||
|
|
||||||
|
Override in subclasses to enable streaming. Implementations should
|
||||||
|
raise on delivery failure so the channel manager can retry.
|
||||||
|
|
||||||
|
Streaming contract: ``_stream_delta`` is a chunk, ``_stream_end`` ends
|
||||||
|
the current segment, and stateful implementations must key buffers by
|
||||||
|
``_stream_id`` rather than only by ``chat_id``.
|
||||||
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|||||||
+412
-291
@@ -1,25 +1,37 @@
|
|||||||
"""Discord channel implementation using Discord Gateway websocket."""
|
"""Discord channel implementation using discord.py."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import importlib.util
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Literal
|
from typing import TYPE_CHECKING, Any, Literal
|
||||||
|
|
||||||
import httpx
|
|
||||||
from pydantic import Field
|
|
||||||
import websockets
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
from pydantic import Field
|
||||||
|
|
||||||
from nanobot.bus.events import OutboundMessage
|
from nanobot.bus.events import 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 build_help_text
|
||||||
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.utils.helpers import split_message
|
from nanobot.utils.helpers import safe_filename, split_message
|
||||||
|
|
||||||
|
DISCORD_AVAILABLE = importlib.util.find_spec("discord") is not None
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import discord
|
||||||
|
from discord import app_commands
|
||||||
|
from discord.abc import Messageable
|
||||||
|
|
||||||
|
if DISCORD_AVAILABLE:
|
||||||
|
import discord
|
||||||
|
from discord import app_commands
|
||||||
|
from discord.abc import Messageable
|
||||||
|
|
||||||
DISCORD_API_BASE = "https://discord.com/api/v10"
|
|
||||||
MAX_ATTACHMENT_BYTES = 20 * 1024 * 1024 # 20MB
|
MAX_ATTACHMENT_BYTES = 20 * 1024 * 1024 # 20MB
|
||||||
MAX_MESSAGE_LEN = 2000 # Discord message character limit
|
MAX_MESSAGE_LEN = 2000 # Discord message character limit
|
||||||
|
TYPING_INTERVAL_S = 8
|
||||||
|
|
||||||
|
|
||||||
class DiscordConfig(Base):
|
class DiscordConfig(Base):
|
||||||
@@ -28,13 +40,205 @@ class DiscordConfig(Base):
|
|||||||
enabled: bool = False
|
enabled: bool = False
|
||||||
token: str = ""
|
token: str = ""
|
||||||
allow_from: list[str] = Field(default_factory=list)
|
allow_from: list[str] = Field(default_factory=list)
|
||||||
gateway_url: str = "wss://gateway.discord.gg/?v=10&encoding=json"
|
|
||||||
intents: int = 37377
|
intents: int = 37377
|
||||||
group_policy: Literal["mention", "open"] = "mention"
|
group_policy: Literal["mention", "open"] = "mention"
|
||||||
|
read_receipt_emoji: str = "👀"
|
||||||
|
working_emoji: str = "🔧"
|
||||||
|
working_emoji_delay: float = 2.0
|
||||||
|
|
||||||
|
|
||||||
|
if DISCORD_AVAILABLE:
|
||||||
|
|
||||||
|
class DiscordBotClient(discord.Client):
|
||||||
|
"""discord.py client that forwards events to the channel."""
|
||||||
|
|
||||||
|
def __init__(self, channel: DiscordChannel, *, intents: discord.Intents) -> None:
|
||||||
|
super().__init__(intents=intents)
|
||||||
|
self._channel = channel
|
||||||
|
self.tree = app_commands.CommandTree(self)
|
||||||
|
self._register_app_commands()
|
||||||
|
|
||||||
|
async def on_ready(self) -> None:
|
||||||
|
self._channel._bot_user_id = str(self.user.id) if self.user else None
|
||||||
|
logger.info("Discord bot connected as user {}", self._channel._bot_user_id)
|
||||||
|
try:
|
||||||
|
synced = await self.tree.sync()
|
||||||
|
logger.info("Discord app commands synced: {}", len(synced))
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Discord app command sync failed: {}", e)
|
||||||
|
|
||||||
|
async def on_message(self, message: discord.Message) -> None:
|
||||||
|
await self._channel._handle_discord_message(message)
|
||||||
|
|
||||||
|
async def _reply_ephemeral(self, interaction: discord.Interaction, text: str) -> bool:
|
||||||
|
"""Send an ephemeral interaction response and report success."""
|
||||||
|
try:
|
||||||
|
await interaction.response.send_message(text, ephemeral=True)
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Discord interaction response failed: {}", e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _forward_slash_command(
|
||||||
|
self,
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
command_text: str,
|
||||||
|
) -> None:
|
||||||
|
sender_id = str(interaction.user.id)
|
||||||
|
channel_id = interaction.channel_id
|
||||||
|
|
||||||
|
if channel_id is None:
|
||||||
|
logger.warning("Discord slash command missing channel_id: {}", command_text)
|
||||||
|
return
|
||||||
|
|
||||||
|
if not self._channel.is_allowed(sender_id):
|
||||||
|
await self._reply_ephemeral(interaction, "You are not allowed to use this bot.")
|
||||||
|
return
|
||||||
|
|
||||||
|
await self._reply_ephemeral(interaction, f"Processing {command_text}...")
|
||||||
|
|
||||||
|
await self._channel._handle_message(
|
||||||
|
sender_id=sender_id,
|
||||||
|
chat_id=str(channel_id),
|
||||||
|
content=command_text,
|
||||||
|
metadata={
|
||||||
|
"interaction_id": str(interaction.id),
|
||||||
|
"guild_id": str(interaction.guild_id) if interaction.guild_id else None,
|
||||||
|
"is_slash_command": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
def _register_app_commands(self) -> None:
|
||||||
|
commands = (
|
||||||
|
("new", "Start a new conversation", "/new"),
|
||||||
|
("stop", "Stop the current task", "/stop"),
|
||||||
|
("restart", "Restart the bot", "/restart"),
|
||||||
|
("status", "Show bot status", "/status"),
|
||||||
|
)
|
||||||
|
|
||||||
|
for name, description, command_text in commands:
|
||||||
|
@self.tree.command(name=name, description=description)
|
||||||
|
async def command_handler(
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
_command_text: str = command_text,
|
||||||
|
) -> None:
|
||||||
|
await self._forward_slash_command(interaction, _command_text)
|
||||||
|
|
||||||
|
@self.tree.command(name="help", description="Show available commands")
|
||||||
|
async def help_command(interaction: discord.Interaction) -> None:
|
||||||
|
sender_id = str(interaction.user.id)
|
||||||
|
if not self._channel.is_allowed(sender_id):
|
||||||
|
await self._reply_ephemeral(interaction, "You are not allowed to use this bot.")
|
||||||
|
return
|
||||||
|
await self._reply_ephemeral(interaction, build_help_text())
|
||||||
|
|
||||||
|
@self.tree.error
|
||||||
|
async def on_app_command_error(
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
error: app_commands.AppCommandError,
|
||||||
|
) -> None:
|
||||||
|
command_name = interaction.command.qualified_name if interaction.command else "?"
|
||||||
|
logger.warning(
|
||||||
|
"Discord app command failed user={} channel={} cmd={} error={}",
|
||||||
|
interaction.user.id,
|
||||||
|
interaction.channel_id,
|
||||||
|
command_name,
|
||||||
|
error,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_outbound(self, msg: OutboundMessage) -> None:
|
||||||
|
"""Send a nanobot outbound message using Discord transport rules."""
|
||||||
|
channel_id = int(msg.chat_id)
|
||||||
|
|
||||||
|
channel = self.get_channel(channel_id)
|
||||||
|
if channel is None:
|
||||||
|
try:
|
||||||
|
channel = await self.fetch_channel(channel_id)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Discord channel {} unavailable: {}", msg.chat_id, e)
|
||||||
|
return
|
||||||
|
|
||||||
|
reference, mention_settings = self._build_reply_context(channel, msg.reply_to)
|
||||||
|
sent_media = False
|
||||||
|
failed_media: list[str] = []
|
||||||
|
|
||||||
|
for index, media_path in enumerate(msg.media or []):
|
||||||
|
if await self._send_file(
|
||||||
|
channel,
|
||||||
|
media_path,
|
||||||
|
reference=reference if index == 0 else None,
|
||||||
|
mention_settings=mention_settings,
|
||||||
|
):
|
||||||
|
sent_media = True
|
||||||
|
else:
|
||||||
|
failed_media.append(Path(media_path).name)
|
||||||
|
|
||||||
|
for index, chunk in enumerate(self._build_chunks(msg.content or "", failed_media, sent_media)):
|
||||||
|
kwargs: dict[str, Any] = {"content": chunk}
|
||||||
|
if index == 0 and reference is not None and not sent_media:
|
||||||
|
kwargs["reference"] = reference
|
||||||
|
kwargs["allowed_mentions"] = mention_settings
|
||||||
|
await channel.send(**kwargs)
|
||||||
|
|
||||||
|
async def _send_file(
|
||||||
|
self,
|
||||||
|
channel: Messageable,
|
||||||
|
file_path: str,
|
||||||
|
*,
|
||||||
|
reference: discord.PartialMessage | None,
|
||||||
|
mention_settings: discord.AllowedMentions,
|
||||||
|
) -> bool:
|
||||||
|
"""Send a file attachment via discord.py."""
|
||||||
|
path = Path(file_path)
|
||||||
|
if not path.is_file():
|
||||||
|
logger.warning("Discord file not found, skipping: {}", file_path)
|
||||||
|
return False
|
||||||
|
|
||||||
|
if path.stat().st_size > MAX_ATTACHMENT_BYTES:
|
||||||
|
logger.warning("Discord file too large (>20MB), skipping: {}", path.name)
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
kwargs: dict[str, Any] = {"file": discord.File(path)}
|
||||||
|
if reference is not None:
|
||||||
|
kwargs["reference"] = reference
|
||||||
|
kwargs["allowed_mentions"] = mention_settings
|
||||||
|
await channel.send(**kwargs)
|
||||||
|
logger.info("Discord file sent: {}", path.name)
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Error sending Discord file {}: {}", path.name, e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_chunks(content: str, failed_media: list[str], sent_media: bool) -> list[str]:
|
||||||
|
"""Build outbound text chunks, including attachment-failure fallback text."""
|
||||||
|
chunks = split_message(content, MAX_MESSAGE_LEN)
|
||||||
|
if chunks or not failed_media or sent_media:
|
||||||
|
return chunks
|
||||||
|
fallback = "\n".join(f"[attachment: {name} - send failed]" for name in failed_media)
|
||||||
|
return split_message(fallback, MAX_MESSAGE_LEN)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_reply_context(
|
||||||
|
channel: Messageable,
|
||||||
|
reply_to: str | None,
|
||||||
|
) -> tuple[discord.PartialMessage | None, discord.AllowedMentions]:
|
||||||
|
"""Build reply context for outbound messages."""
|
||||||
|
mention_settings = discord.AllowedMentions(replied_user=False)
|
||||||
|
if not reply_to:
|
||||||
|
return None, mention_settings
|
||||||
|
try:
|
||||||
|
message_id = int(reply_to)
|
||||||
|
except ValueError:
|
||||||
|
logger.warning("Invalid Discord reply target: {}", reply_to)
|
||||||
|
return None, mention_settings
|
||||||
|
|
||||||
|
return channel.get_partial_message(message_id), mention_settings
|
||||||
|
|
||||||
|
|
||||||
class DiscordChannel(BaseChannel):
|
class DiscordChannel(BaseChannel):
|
||||||
"""Discord channel using Gateway websocket."""
|
"""Discord channel using discord.py."""
|
||||||
|
|
||||||
name = "discord"
|
name = "discord"
|
||||||
display_name = "Discord"
|
display_name = "Discord"
|
||||||
@@ -43,353 +247,270 @@ class DiscordChannel(BaseChannel):
|
|||||||
def default_config(cls) -> dict[str, Any]:
|
def default_config(cls) -> dict[str, Any]:
|
||||||
return DiscordConfig().model_dump(by_alias=True)
|
return DiscordConfig().model_dump(by_alias=True)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _channel_key(channel_or_id: Any) -> str:
|
||||||
|
"""Normalize channel-like objects and ids to a stable string key."""
|
||||||
|
channel_id = getattr(channel_or_id, "id", channel_or_id)
|
||||||
|
return str(channel_id)
|
||||||
|
|
||||||
def __init__(self, config: Any, bus: MessageBus):
|
def __init__(self, config: Any, bus: MessageBus):
|
||||||
if isinstance(config, dict):
|
if isinstance(config, dict):
|
||||||
config = DiscordConfig.model_validate(config)
|
config = DiscordConfig.model_validate(config)
|
||||||
super().__init__(config, bus)
|
super().__init__(config, bus)
|
||||||
self.config: DiscordConfig = config
|
self.config: DiscordConfig = config
|
||||||
self._ws: websockets.WebSocketClientProtocol | None = None
|
self._client: DiscordBotClient | None = None
|
||||||
self._seq: int | None = None
|
self._typing_tasks: dict[str, asyncio.Task[None]] = {}
|
||||||
self._heartbeat_task: asyncio.Task | None = None
|
|
||||||
self._typing_tasks: dict[str, asyncio.Task] = {}
|
|
||||||
self._http: httpx.AsyncClient | None = None
|
|
||||||
self._bot_user_id: str | None = None
|
self._bot_user_id: str | None = None
|
||||||
|
self._pending_reactions: dict[str, Any] = {} # chat_id -> message object
|
||||||
|
self._working_emoji_tasks: dict[str, asyncio.Task[None]] = {}
|
||||||
|
|
||||||
async def start(self) -> None:
|
async def start(self) -> None:
|
||||||
"""Start the Discord gateway connection."""
|
"""Start the Discord client."""
|
||||||
|
if not DISCORD_AVAILABLE:
|
||||||
|
logger.error("discord.py not installed. Run: pip install nanobot-ai[discord]")
|
||||||
|
return
|
||||||
|
|
||||||
if not self.config.token:
|
if not self.config.token:
|
||||||
logger.error("Discord bot token not configured")
|
logger.error("Discord bot token not configured")
|
||||||
return
|
return
|
||||||
|
|
||||||
self._running = True
|
try:
|
||||||
self._http = httpx.AsyncClient(timeout=30.0)
|
intents = discord.Intents.none()
|
||||||
|
intents.value = self.config.intents
|
||||||
|
self._client = DiscordBotClient(self, intents=intents)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to initialize Discord client: {}", e)
|
||||||
|
self._client = None
|
||||||
|
self._running = False
|
||||||
|
return
|
||||||
|
|
||||||
while self._running:
|
self._running = True
|
||||||
try:
|
logger.info("Starting Discord client via discord.py...")
|
||||||
logger.info("Connecting to Discord gateway...")
|
|
||||||
async with websockets.connect(self.config.gateway_url) as ws:
|
try:
|
||||||
self._ws = ws
|
await self._client.start(self.config.token)
|
||||||
await self._gateway_loop()
|
except asyncio.CancelledError:
|
||||||
except asyncio.CancelledError:
|
raise
|
||||||
break
|
except Exception as e:
|
||||||
except Exception as e:
|
logger.error("Discord client startup failed: {}", e)
|
||||||
logger.warning("Discord gateway error: {}", e)
|
finally:
|
||||||
if self._running:
|
self._running = False
|
||||||
logger.info("Reconnecting to Discord gateway in 5 seconds...")
|
await self._reset_runtime_state(close_client=True)
|
||||||
await asyncio.sleep(5)
|
|
||||||
|
|
||||||
async def stop(self) -> None:
|
async def stop(self) -> None:
|
||||||
"""Stop the Discord channel."""
|
"""Stop the Discord channel."""
|
||||||
self._running = False
|
self._running = False
|
||||||
if self._heartbeat_task:
|
await self._reset_runtime_state(close_client=True)
|
||||||
self._heartbeat_task.cancel()
|
|
||||||
self._heartbeat_task = None
|
|
||||||
for task in self._typing_tasks.values():
|
|
||||||
task.cancel()
|
|
||||||
self._typing_tasks.clear()
|
|
||||||
if self._ws:
|
|
||||||
await self._ws.close()
|
|
||||||
self._ws = None
|
|
||||||
if self._http:
|
|
||||||
await self._http.aclose()
|
|
||||||
self._http = None
|
|
||||||
|
|
||||||
async def send(self, msg: OutboundMessage) -> None:
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
"""Send a message through Discord REST API, including file attachments."""
|
"""Send a message through Discord using discord.py."""
|
||||||
if not self._http:
|
client = self._client
|
||||||
logger.warning("Discord HTTP client not initialized")
|
if client is None or not client.is_ready():
|
||||||
|
logger.warning("Discord client not ready; dropping outbound message")
|
||||||
return
|
return
|
||||||
|
|
||||||
url = f"{DISCORD_API_BASE}/channels/{msg.chat_id}/messages"
|
is_progress = bool((msg.metadata or {}).get("_progress"))
|
||||||
headers = {"Authorization": f"Bot {self.config.token}"}
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
sent_media = False
|
await client.send_outbound(msg)
|
||||||
failed_media: list[str] = []
|
except Exception as e:
|
||||||
|
logger.error("Error sending Discord message: {}", e)
|
||||||
# Send file attachments first
|
|
||||||
for media_path in msg.media or []:
|
|
||||||
if await self._send_file(url, headers, media_path, reply_to=msg.reply_to):
|
|
||||||
sent_media = True
|
|
||||||
else:
|
|
||||||
failed_media.append(Path(media_path).name)
|
|
||||||
|
|
||||||
# Send text content
|
|
||||||
chunks = split_message(msg.content or "", MAX_MESSAGE_LEN)
|
|
||||||
if not chunks and failed_media and not sent_media:
|
|
||||||
chunks = split_message(
|
|
||||||
"\n".join(f"[attachment: {name} - send failed]" for name in failed_media),
|
|
||||||
MAX_MESSAGE_LEN,
|
|
||||||
)
|
|
||||||
if not chunks:
|
|
||||||
return
|
|
||||||
|
|
||||||
for i, chunk in enumerate(chunks):
|
|
||||||
payload: dict[str, Any] = {"content": chunk}
|
|
||||||
|
|
||||||
# Let the first successful attachment carry the reply if present.
|
|
||||||
if i == 0 and msg.reply_to and not sent_media:
|
|
||||||
payload["message_reference"] = {"message_id": msg.reply_to}
|
|
||||||
payload["allowed_mentions"] = {"replied_user": False}
|
|
||||||
|
|
||||||
if not await self._send_payload(url, headers, payload):
|
|
||||||
break # Abort remaining chunks on failure
|
|
||||||
finally:
|
finally:
|
||||||
await self._stop_typing(msg.chat_id)
|
if not is_progress:
|
||||||
|
await self._stop_typing(msg.chat_id)
|
||||||
|
await self._clear_reactions(msg.chat_id)
|
||||||
|
|
||||||
async def _send_payload(
|
async def _handle_discord_message(self, message: discord.Message) -> None:
|
||||||
self, url: str, headers: dict[str, str], payload: dict[str, Any]
|
"""Handle incoming Discord messages from discord.py."""
|
||||||
) -> bool:
|
if message.author.bot:
|
||||||
"""Send a single Discord API payload with retry on rate-limit. Returns True on success."""
|
return
|
||||||
for attempt in range(3):
|
|
||||||
|
sender_id = str(message.author.id)
|
||||||
|
channel_id = self._channel_key(message.channel)
|
||||||
|
content = message.content or ""
|
||||||
|
|
||||||
|
if not self._should_accept_inbound(message, sender_id, content):
|
||||||
|
return
|
||||||
|
|
||||||
|
media_paths, attachment_markers = await self._download_attachments(message.attachments)
|
||||||
|
full_content = self._compose_inbound_content(content, attachment_markers)
|
||||||
|
metadata = self._build_inbound_metadata(message)
|
||||||
|
|
||||||
|
await self._start_typing(message.channel)
|
||||||
|
|
||||||
|
# Add read receipt reaction immediately, working emoji after delay
|
||||||
|
channel_id = self._channel_key(message.channel)
|
||||||
|
try:
|
||||||
|
await message.add_reaction(self.config.read_receipt_emoji)
|
||||||
|
self._pending_reactions[channel_id] = message
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Failed to add read receipt reaction: {}", e)
|
||||||
|
|
||||||
|
# Delayed working indicator (cosmetic — not tied to subagent lifecycle)
|
||||||
|
async def _delayed_working_emoji() -> None:
|
||||||
|
await asyncio.sleep(self.config.working_emoji_delay)
|
||||||
try:
|
try:
|
||||||
response = await self._http.post(url, headers=headers, json=payload)
|
await message.add_reaction(self.config.working_emoji)
|
||||||
if response.status_code == 429:
|
except Exception:
|
||||||
data = response.json()
|
pass
|
||||||
retry_after = float(data.get("retry_after", 1.0))
|
|
||||||
logger.warning("Discord rate limited, retrying in {}s", retry_after)
|
|
||||||
await asyncio.sleep(retry_after)
|
|
||||||
continue
|
|
||||||
response.raise_for_status()
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
if attempt == 2:
|
|
||||||
logger.error("Error sending Discord message: {}", e)
|
|
||||||
else:
|
|
||||||
await asyncio.sleep(1)
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def _send_file(
|
self._working_emoji_tasks[channel_id] = asyncio.create_task(_delayed_working_emoji())
|
||||||
|
|
||||||
|
try:
|
||||||
|
await self._handle_message(
|
||||||
|
sender_id=sender_id,
|
||||||
|
chat_id=channel_id,
|
||||||
|
content=full_content,
|
||||||
|
media=media_paths,
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
await self._clear_reactions(channel_id)
|
||||||
|
await self._stop_typing(channel_id)
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def _on_message(self, message: discord.Message) -> None:
|
||||||
|
"""Backward-compatible alias for legacy tests/callers."""
|
||||||
|
await self._handle_discord_message(message)
|
||||||
|
|
||||||
|
def _should_accept_inbound(
|
||||||
self,
|
self,
|
||||||
url: str,
|
message: discord.Message,
|
||||||
headers: dict[str, str],
|
sender_id: str,
|
||||||
file_path: str,
|
content: str,
|
||||||
reply_to: str | None = None,
|
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Send a file attachment via Discord REST API using multipart/form-data."""
|
"""Check if inbound Discord message should be processed."""
|
||||||
path = Path(file_path)
|
|
||||||
if not path.is_file():
|
|
||||||
logger.warning("Discord file not found, skipping: {}", file_path)
|
|
||||||
return False
|
|
||||||
|
|
||||||
if path.stat().st_size > MAX_ATTACHMENT_BYTES:
|
|
||||||
logger.warning("Discord file too large (>20MB), skipping: {}", path.name)
|
|
||||||
return False
|
|
||||||
|
|
||||||
payload_json: dict[str, Any] = {}
|
|
||||||
if reply_to:
|
|
||||||
payload_json["message_reference"] = {"message_id": reply_to}
|
|
||||||
payload_json["allowed_mentions"] = {"replied_user": False}
|
|
||||||
|
|
||||||
for attempt in range(3):
|
|
||||||
try:
|
|
||||||
with open(path, "rb") as f:
|
|
||||||
files = {"files[0]": (path.name, f, "application/octet-stream")}
|
|
||||||
data: dict[str, Any] = {}
|
|
||||||
if payload_json:
|
|
||||||
data["payload_json"] = json.dumps(payload_json)
|
|
||||||
response = await self._http.post(
|
|
||||||
url, headers=headers, files=files, data=data
|
|
||||||
)
|
|
||||||
if response.status_code == 429:
|
|
||||||
resp_data = response.json()
|
|
||||||
retry_after = float(resp_data.get("retry_after", 1.0))
|
|
||||||
logger.warning("Discord rate limited, retrying in {}s", retry_after)
|
|
||||||
await asyncio.sleep(retry_after)
|
|
||||||
continue
|
|
||||||
response.raise_for_status()
|
|
||||||
logger.info("Discord file sent: {}", path.name)
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
if attempt == 2:
|
|
||||||
logger.error("Error sending Discord file {}: {}", path.name, e)
|
|
||||||
else:
|
|
||||||
await asyncio.sleep(1)
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def _gateway_loop(self) -> None:
|
|
||||||
"""Main gateway loop: identify, heartbeat, dispatch events."""
|
|
||||||
if not self._ws:
|
|
||||||
return
|
|
||||||
|
|
||||||
async for raw in self._ws:
|
|
||||||
try:
|
|
||||||
data = json.loads(raw)
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
logger.warning("Invalid JSON from Discord gateway: {}", raw[:100])
|
|
||||||
continue
|
|
||||||
|
|
||||||
op = data.get("op")
|
|
||||||
event_type = data.get("t")
|
|
||||||
seq = data.get("s")
|
|
||||||
payload = data.get("d")
|
|
||||||
|
|
||||||
if seq is not None:
|
|
||||||
self._seq = seq
|
|
||||||
|
|
||||||
if op == 10:
|
|
||||||
# HELLO: start heartbeat and identify
|
|
||||||
interval_ms = payload.get("heartbeat_interval", 45000)
|
|
||||||
await self._start_heartbeat(interval_ms / 1000)
|
|
||||||
await self._identify()
|
|
||||||
elif op == 0 and event_type == "READY":
|
|
||||||
logger.info("Discord gateway READY")
|
|
||||||
# Capture bot user ID for mention detection
|
|
||||||
user_data = payload.get("user") or {}
|
|
||||||
self._bot_user_id = user_data.get("id")
|
|
||||||
logger.info("Discord bot connected as user {}", self._bot_user_id)
|
|
||||||
elif op == 0 and event_type == "MESSAGE_CREATE":
|
|
||||||
await self._handle_message_create(payload)
|
|
||||||
elif op == 7:
|
|
||||||
# RECONNECT: exit loop to reconnect
|
|
||||||
logger.info("Discord gateway requested reconnect")
|
|
||||||
break
|
|
||||||
elif op == 9:
|
|
||||||
# INVALID_SESSION: reconnect
|
|
||||||
logger.warning("Discord gateway invalid session")
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _identify(self) -> None:
|
|
||||||
"""Send IDENTIFY payload."""
|
|
||||||
if not self._ws:
|
|
||||||
return
|
|
||||||
|
|
||||||
identify = {
|
|
||||||
"op": 2,
|
|
||||||
"d": {
|
|
||||||
"token": self.config.token,
|
|
||||||
"intents": self.config.intents,
|
|
||||||
"properties": {
|
|
||||||
"os": "nanobot",
|
|
||||||
"browser": "nanobot",
|
|
||||||
"device": "nanobot",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
await self._ws.send(json.dumps(identify))
|
|
||||||
|
|
||||||
async def _start_heartbeat(self, interval_s: float) -> None:
|
|
||||||
"""Start or restart the heartbeat loop."""
|
|
||||||
if self._heartbeat_task:
|
|
||||||
self._heartbeat_task.cancel()
|
|
||||||
|
|
||||||
async def heartbeat_loop() -> None:
|
|
||||||
while self._running and self._ws:
|
|
||||||
payload = {"op": 1, "d": self._seq}
|
|
||||||
try:
|
|
||||||
await self._ws.send(json.dumps(payload))
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning("Discord heartbeat failed: {}", e)
|
|
||||||
break
|
|
||||||
await asyncio.sleep(interval_s)
|
|
||||||
|
|
||||||
self._heartbeat_task = asyncio.create_task(heartbeat_loop())
|
|
||||||
|
|
||||||
async def _handle_message_create(self, payload: dict[str, Any]) -> None:
|
|
||||||
"""Handle incoming Discord messages."""
|
|
||||||
author = payload.get("author") or {}
|
|
||||||
if author.get("bot"):
|
|
||||||
return
|
|
||||||
|
|
||||||
sender_id = str(author.get("id", ""))
|
|
||||||
channel_id = str(payload.get("channel_id", ""))
|
|
||||||
content = payload.get("content") or ""
|
|
||||||
guild_id = payload.get("guild_id")
|
|
||||||
|
|
||||||
if not sender_id or not channel_id:
|
|
||||||
return
|
|
||||||
|
|
||||||
if not self.is_allowed(sender_id):
|
if not self.is_allowed(sender_id):
|
||||||
return
|
return False
|
||||||
|
if message.guild is not None and not self._should_respond_in_group(message, content):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
# Check group channel policy (DMs always respond if is_allowed passes)
|
async def _download_attachments(
|
||||||
if guild_id is not None:
|
self,
|
||||||
if not self._should_respond_in_group(payload, content):
|
attachments: list[discord.Attachment],
|
||||||
return
|
) -> tuple[list[str], list[str]]:
|
||||||
|
"""Download supported attachments and return paths + display markers."""
|
||||||
content_parts = [content] if content else []
|
|
||||||
media_paths: list[str] = []
|
media_paths: list[str] = []
|
||||||
|
markers: list[str] = []
|
||||||
media_dir = get_media_dir("discord")
|
media_dir = get_media_dir("discord")
|
||||||
|
|
||||||
for attachment in payload.get("attachments") or []:
|
for attachment in attachments:
|
||||||
url = attachment.get("url")
|
filename = attachment.filename or "attachment"
|
||||||
filename = attachment.get("filename") or "attachment"
|
if attachment.size and attachment.size > MAX_ATTACHMENT_BYTES:
|
||||||
size = attachment.get("size") or 0
|
markers.append(f"[attachment: {filename} - too large]")
|
||||||
if not url or not self._http:
|
|
||||||
continue
|
|
||||||
if size and size > MAX_ATTACHMENT_BYTES:
|
|
||||||
content_parts.append(f"[attachment: {filename} - too large]")
|
|
||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
media_dir.mkdir(parents=True, exist_ok=True)
|
media_dir.mkdir(parents=True, exist_ok=True)
|
||||||
file_path = media_dir / f"{attachment.get('id', 'file')}_{filename.replace('/', '_')}"
|
safe_name = safe_filename(filename)
|
||||||
resp = await self._http.get(url)
|
file_path = media_dir / f"{attachment.id}_{safe_name}"
|
||||||
resp.raise_for_status()
|
await attachment.save(file_path)
|
||||||
file_path.write_bytes(resp.content)
|
|
||||||
media_paths.append(str(file_path))
|
media_paths.append(str(file_path))
|
||||||
content_parts.append(f"[attachment: {file_path}]")
|
markers.append(f"[attachment: {file_path.name}]")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("Failed to download Discord attachment: {}", e)
|
logger.warning("Failed to download Discord attachment: {}", e)
|
||||||
content_parts.append(f"[attachment: {filename} - download failed]")
|
markers.append(f"[attachment: {filename} - download failed]")
|
||||||
|
|
||||||
reply_to = (payload.get("referenced_message") or {}).get("id")
|
return media_paths, markers
|
||||||
|
|
||||||
await self._start_typing(channel_id)
|
@staticmethod
|
||||||
|
def _compose_inbound_content(content: str, attachment_markers: list[str]) -> str:
|
||||||
|
"""Combine message text with attachment markers."""
|
||||||
|
content_parts = [content] if content else []
|
||||||
|
content_parts.extend(attachment_markers)
|
||||||
|
return "\n".join(part for part in content_parts if part) or "[empty message]"
|
||||||
|
|
||||||
await self._handle_message(
|
@staticmethod
|
||||||
sender_id=sender_id,
|
def _build_inbound_metadata(message: discord.Message) -> dict[str, str | None]:
|
||||||
chat_id=channel_id,
|
"""Build metadata for inbound Discord messages."""
|
||||||
content="\n".join(p for p in content_parts if p) or "[empty message]",
|
reply_to = str(message.reference.message_id) if message.reference and message.reference.message_id else None
|
||||||
media=media_paths,
|
return {
|
||||||
metadata={
|
"message_id": str(message.id),
|
||||||
"message_id": str(payload.get("id", "")),
|
"guild_id": str(message.guild.id) if message.guild else None,
|
||||||
"guild_id": guild_id,
|
"reply_to": reply_to,
|
||||||
"reply_to": reply_to,
|
}
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
def _should_respond_in_group(self, payload: dict[str, Any], content: str) -> bool:
|
def _should_respond_in_group(self, message: discord.Message, content: str) -> bool:
|
||||||
"""Check if bot should respond in a group channel based on policy."""
|
"""Check if the bot should respond in a guild channel based on policy."""
|
||||||
if self.config.group_policy == "open":
|
if self.config.group_policy == "open":
|
||||||
return True
|
return True
|
||||||
|
|
||||||
if self.config.group_policy == "mention":
|
if self.config.group_policy == "mention":
|
||||||
# Check if bot was mentioned in the message
|
bot_user_id = self._bot_user_id
|
||||||
if self._bot_user_id:
|
if bot_user_id is None:
|
||||||
# Check mentions array
|
logger.debug("Discord message in {} ignored (bot identity unavailable)", message.channel.id)
|
||||||
mentions = payload.get("mentions") or []
|
return False
|
||||||
for mention in mentions:
|
|
||||||
if str(mention.get("id")) == self._bot_user_id:
|
if any(str(user.id) == bot_user_id for user in message.mentions):
|
||||||
return True
|
return True
|
||||||
# Also check content for mention format <@USER_ID>
|
if f"<@{bot_user_id}>" in content or f"<@!{bot_user_id}>" in content:
|
||||||
if f"<@{self._bot_user_id}>" in content or f"<@!{self._bot_user_id}>" in content:
|
return True
|
||||||
return True
|
|
||||||
logger.debug("Discord message in {} ignored (bot not mentioned)", payload.get("channel_id"))
|
logger.debug("Discord message in {} ignored (bot not mentioned)", message.channel.id)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
async def _start_typing(self, channel_id: str) -> None:
|
async def _start_typing(self, channel: Messageable) -> None:
|
||||||
"""Start periodic typing indicator for a channel."""
|
"""Start periodic typing indicator for a channel."""
|
||||||
|
channel_id = self._channel_key(channel)
|
||||||
await self._stop_typing(channel_id)
|
await self._stop_typing(channel_id)
|
||||||
|
|
||||||
async def typing_loop() -> None:
|
async def typing_loop() -> None:
|
||||||
url = f"{DISCORD_API_BASE}/channels/{channel_id}/typing"
|
|
||||||
headers = {"Authorization": f"Bot {self.config.token}"}
|
|
||||||
while self._running:
|
while self._running:
|
||||||
try:
|
try:
|
||||||
await self._http.post(url, headers=headers)
|
async with channel.typing():
|
||||||
|
await asyncio.sleep(TYPING_INTERVAL_S)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
return
|
return
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug("Discord typing indicator failed for {}: {}", channel_id, e)
|
logger.debug("Discord typing indicator failed for {}: {}", channel_id, e)
|
||||||
return
|
return
|
||||||
await asyncio.sleep(8)
|
|
||||||
|
|
||||||
self._typing_tasks[channel_id] = asyncio.create_task(typing_loop())
|
self._typing_tasks[channel_id] = asyncio.create_task(typing_loop())
|
||||||
|
|
||||||
async def _stop_typing(self, channel_id: str) -> None:
|
async def _stop_typing(self, channel_id: str) -> None:
|
||||||
"""Stop typing indicator for a channel."""
|
"""Stop typing indicator for a channel."""
|
||||||
task = self._typing_tasks.pop(channel_id, None)
|
task = self._typing_tasks.pop(self._channel_key(channel_id), None)
|
||||||
if task:
|
if task is None:
|
||||||
|
return
|
||||||
|
task.cancel()
|
||||||
|
try:
|
||||||
|
await task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
async def _clear_reactions(self, chat_id: str) -> None:
|
||||||
|
"""Remove all pending reactions after bot replies."""
|
||||||
|
# Cancel delayed working emoji if it hasn't fired yet
|
||||||
|
task = self._working_emoji_tasks.pop(chat_id, None)
|
||||||
|
if task and not task.done():
|
||||||
task.cancel()
|
task.cancel()
|
||||||
|
|
||||||
|
msg_obj = self._pending_reactions.pop(chat_id, None)
|
||||||
|
if msg_obj is None:
|
||||||
|
return
|
||||||
|
bot_user = self._client.user if self._client else None
|
||||||
|
for emoji in (self.config.read_receipt_emoji, self.config.working_emoji):
|
||||||
|
try:
|
||||||
|
await msg_obj.remove_reaction(emoji, bot_user)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def _cancel_all_typing(self) -> None:
|
||||||
|
"""Stop all typing tasks."""
|
||||||
|
channel_ids = list(self._typing_tasks)
|
||||||
|
for channel_id in channel_ids:
|
||||||
|
await self._stop_typing(channel_id)
|
||||||
|
|
||||||
|
async def _reset_runtime_state(self, close_client: bool) -> None:
|
||||||
|
"""Reset client and typing state."""
|
||||||
|
await self._cancel_all_typing()
|
||||||
|
if close_client and self._client is not None and not self._client.is_closed():
|
||||||
|
try:
|
||||||
|
await self._client.close()
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Discord client close failed: {}", e)
|
||||||
|
self._client = None
|
||||||
|
self._bot_user_id = None
|
||||||
|
|||||||
@@ -51,6 +51,10 @@ class EmailConfig(Base):
|
|||||||
subject_prefix: str = "Re: "
|
subject_prefix: str = "Re: "
|
||||||
allow_from: list[str] = Field(default_factory=list)
|
allow_from: list[str] = Field(default_factory=list)
|
||||||
|
|
||||||
|
# Email authentication verification (anti-spoofing)
|
||||||
|
verify_dkim: bool = True # Require Authentication-Results with dkim=pass
|
||||||
|
verify_spf: bool = True # Require Authentication-Results with spf=pass
|
||||||
|
|
||||||
|
|
||||||
class EmailChannel(BaseChannel):
|
class EmailChannel(BaseChannel):
|
||||||
"""
|
"""
|
||||||
@@ -123,6 +127,12 @@ class EmailChannel(BaseChannel):
|
|||||||
return
|
return
|
||||||
|
|
||||||
self._running = True
|
self._running = True
|
||||||
|
if not self.config.verify_dkim and not self.config.verify_spf:
|
||||||
|
logger.warning(
|
||||||
|
"Email channel: DKIM and SPF verification are both DISABLED. "
|
||||||
|
"Emails with spoofed From headers will be accepted. "
|
||||||
|
"Set verify_dkim=true and verify_spf=true for anti-spoofing protection."
|
||||||
|
)
|
||||||
logger.info("Starting Email channel (IMAP polling mode)...")
|
logger.info("Starting Email channel (IMAP polling mode)...")
|
||||||
|
|
||||||
poll_seconds = max(5, int(self.config.poll_interval_seconds))
|
poll_seconds = max(5, int(self.config.poll_interval_seconds))
|
||||||
@@ -360,6 +370,23 @@ class EmailChannel(BaseChannel):
|
|||||||
if not sender:
|
if not sender:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
# --- Anti-spoofing: verify Authentication-Results ---
|
||||||
|
spf_pass, dkim_pass = self._check_authentication_results(parsed)
|
||||||
|
if self.config.verify_spf and not spf_pass:
|
||||||
|
logger.warning(
|
||||||
|
"Email from {} rejected: SPF verification failed "
|
||||||
|
"(no 'spf=pass' in Authentication-Results header)",
|
||||||
|
sender,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
if self.config.verify_dkim and not dkim_pass:
|
||||||
|
logger.warning(
|
||||||
|
"Email from {} rejected: DKIM verification failed "
|
||||||
|
"(no 'dkim=pass' in Authentication-Results header)",
|
||||||
|
sender,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
subject = self._decode_header_value(parsed.get("Subject", ""))
|
subject = self._decode_header_value(parsed.get("Subject", ""))
|
||||||
date_value = parsed.get("Date", "")
|
date_value = parsed.get("Date", "")
|
||||||
message_id = parsed.get("Message-ID", "").strip()
|
message_id = parsed.get("Message-ID", "").strip()
|
||||||
@@ -370,7 +397,7 @@ class EmailChannel(BaseChannel):
|
|||||||
|
|
||||||
body = body[: self.config.max_body_chars]
|
body = body[: self.config.max_body_chars]
|
||||||
content = (
|
content = (
|
||||||
f"Email received.\n"
|
f"[EMAIL-CONTEXT] Email received.\n"
|
||||||
f"From: {sender}\n"
|
f"From: {sender}\n"
|
||||||
f"Subject: {subject}\n"
|
f"Subject: {subject}\n"
|
||||||
f"Date: {date_value}\n\n"
|
f"Date: {date_value}\n\n"
|
||||||
@@ -493,6 +520,23 @@ class EmailChannel(BaseChannel):
|
|||||||
return cls._html_to_text(payload).strip()
|
return cls._html_to_text(payload).strip()
|
||||||
return payload.strip()
|
return payload.strip()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _check_authentication_results(parsed_msg: Any) -> tuple[bool, bool]:
|
||||||
|
"""Parse Authentication-Results headers for SPF and DKIM verdicts.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A tuple of (spf_pass, dkim_pass) booleans.
|
||||||
|
"""
|
||||||
|
spf_pass = False
|
||||||
|
dkim_pass = False
|
||||||
|
for ar_header in parsed_msg.get_all("Authentication-Results") or []:
|
||||||
|
ar_lower = ar_header.lower()
|
||||||
|
if re.search(r"\bspf\s*=\s*pass\b", ar_lower):
|
||||||
|
spf_pass = True
|
||||||
|
if re.search(r"\bdkim\s*=\s*pass\b", ar_lower):
|
||||||
|
dkim_pass = True
|
||||||
|
return spf_pass, dkim_pass
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _html_to_text(raw_html: str) -> str:
|
def _html_to_text(raw_html: str) -> str:
|
||||||
text = re.sub(r"<\s*br\s*/?>", "\n", raw_html, flags=re.IGNORECASE)
|
text = re.sub(r"<\s*br\s*/?>", "\n", raw_html, flags=re.IGNORECASE)
|
||||||
|
|||||||
+161
-5
@@ -5,7 +5,10 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
import threading
|
import threading
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Literal
|
from typing import Any, Literal
|
||||||
|
|
||||||
@@ -248,6 +251,19 @@ class FeishuConfig(Base):
|
|||||||
react_emoji: str = "THUMBSUP"
|
react_emoji: str = "THUMBSUP"
|
||||||
group_policy: Literal["open", "mention"] = "mention"
|
group_policy: Literal["open", "mention"] = "mention"
|
||||||
reply_to_message: bool = False # If True, bot replies quote the user's original message
|
reply_to_message: bool = False # If True, bot replies quote the user's original message
|
||||||
|
streaming: bool = True
|
||||||
|
|
||||||
|
|
||||||
|
_STREAM_ELEMENT_ID = "streaming_md"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _FeishuStreamBuf:
|
||||||
|
"""Per-chat streaming accumulator using CardKit streaming API."""
|
||||||
|
text: str = ""
|
||||||
|
card_id: str | None = None
|
||||||
|
sequence: int = 0
|
||||||
|
last_edit: float = 0.0
|
||||||
|
|
||||||
|
|
||||||
class FeishuChannel(BaseChannel):
|
class FeishuChannel(BaseChannel):
|
||||||
@@ -265,6 +281,8 @@ class FeishuChannel(BaseChannel):
|
|||||||
name = "feishu"
|
name = "feishu"
|
||||||
display_name = "Feishu"
|
display_name = "Feishu"
|
||||||
|
|
||||||
|
_STREAM_EDIT_INTERVAL = 0.5 # throttle between CardKit streaming updates
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def default_config(cls) -> dict[str, Any]:
|
def default_config(cls) -> dict[str, Any]:
|
||||||
return FeishuConfig().model_dump(by_alias=True)
|
return FeishuConfig().model_dump(by_alias=True)
|
||||||
@@ -279,6 +297,7 @@ class FeishuChannel(BaseChannel):
|
|||||||
self._ws_thread: threading.Thread | None = None
|
self._ws_thread: threading.Thread | None = None
|
||||||
self._processed_message_ids: OrderedDict[str, None] = OrderedDict() # Ordered dedup cache
|
self._processed_message_ids: OrderedDict[str, None] = OrderedDict() # Ordered dedup cache
|
||||||
self._loop: asyncio.AbstractEventLoop | None = None
|
self._loop: asyncio.AbstractEventLoop | None = None
|
||||||
|
self._stream_bufs: dict[str, _FeishuStreamBuf] = {}
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _register_optional_event(builder: Any, method_name: str, handler: Any) -> Any:
|
def _register_optional_event(builder: Any, method_name: str, handler: Any) -> Any:
|
||||||
@@ -906,8 +925,8 @@ class FeishuChannel(BaseChannel):
|
|||||||
logger.error("Error replying to Feishu message {}: {}", parent_message_id, e)
|
logger.error("Error replying to Feishu message {}: {}", parent_message_id, e)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def _send_message_sync(self, receive_id_type: str, receive_id: str, msg_type: str, content: str) -> bool:
|
def _send_message_sync(self, receive_id_type: str, receive_id: str, msg_type: str, content: str) -> str | None:
|
||||||
"""Send a single message (text/image/file/interactive) synchronously."""
|
"""Send a single message and return the message_id on success."""
|
||||||
from lark_oapi.api.im.v1 import CreateMessageRequest, CreateMessageRequestBody
|
from lark_oapi.api.im.v1 import CreateMessageRequest, CreateMessageRequestBody
|
||||||
try:
|
try:
|
||||||
request = CreateMessageRequest.builder() \
|
request = CreateMessageRequest.builder() \
|
||||||
@@ -925,13 +944,149 @@ class FeishuChannel(BaseChannel):
|
|||||||
"Failed to send Feishu {} message: code={}, msg={}, log_id={}",
|
"Failed to send Feishu {} message: code={}, msg={}, log_id={}",
|
||||||
msg_type, response.code, response.msg, response.get_log_id()
|
msg_type, response.code, response.msg, response.get_log_id()
|
||||||
)
|
)
|
||||||
return False
|
return None
|
||||||
logger.debug("Feishu {} message sent to {}", msg_type, receive_id)
|
msg_id = getattr(response.data, "message_id", None)
|
||||||
return True
|
logger.debug("Feishu {} message sent to {}: {}", msg_type, receive_id, msg_id)
|
||||||
|
return msg_id
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error sending Feishu {} message: {}", msg_type, e)
|
logger.error("Error sending Feishu {} message: {}", msg_type, e)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _create_streaming_card_sync(self, receive_id_type: str, chat_id: str) -> str | None:
|
||||||
|
"""Create a CardKit streaming card, send it to chat, return card_id."""
|
||||||
|
from lark_oapi.api.cardkit.v1 import CreateCardRequest, CreateCardRequestBody
|
||||||
|
card_json = {
|
||||||
|
"schema": "2.0",
|
||||||
|
"config": {"wide_screen_mode": True, "update_multi": True, "streaming_mode": True},
|
||||||
|
"body": {"elements": [{"tag": "markdown", "content": "", "element_id": _STREAM_ELEMENT_ID}]},
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
request = CreateCardRequest.builder().request_body(
|
||||||
|
CreateCardRequestBody.builder()
|
||||||
|
.type("card_json")
|
||||||
|
.data(json.dumps(card_json, ensure_ascii=False))
|
||||||
|
.build()
|
||||||
|
).build()
|
||||||
|
response = self._client.cardkit.v1.card.create(request)
|
||||||
|
if not response.success():
|
||||||
|
logger.warning("Failed to create streaming card: code={}, msg={}", response.code, response.msg)
|
||||||
|
return None
|
||||||
|
card_id = getattr(response.data, "card_id", None)
|
||||||
|
if card_id:
|
||||||
|
message_id = self._send_message_sync(
|
||||||
|
receive_id_type, chat_id, "interactive",
|
||||||
|
json.dumps({"type": "card", "data": {"card_id": card_id}}),
|
||||||
|
)
|
||||||
|
if message_id:
|
||||||
|
return card_id
|
||||||
|
logger.warning("Created streaming card {} but failed to send it to {}", card_id, chat_id)
|
||||||
|
return None
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Error creating streaming card: {}", e)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _stream_update_text_sync(self, card_id: str, content: str, sequence: int) -> bool:
|
||||||
|
"""Stream-update the markdown element on a CardKit card (typewriter effect)."""
|
||||||
|
from lark_oapi.api.cardkit.v1 import ContentCardElementRequest, ContentCardElementRequestBody
|
||||||
|
try:
|
||||||
|
request = ContentCardElementRequest.builder() \
|
||||||
|
.card_id(card_id) \
|
||||||
|
.element_id(_STREAM_ELEMENT_ID) \
|
||||||
|
.request_body(
|
||||||
|
ContentCardElementRequestBody.builder()
|
||||||
|
.content(content).sequence(sequence).build()
|
||||||
|
).build()
|
||||||
|
response = self._client.cardkit.v1.card_element.content(request)
|
||||||
|
if not response.success():
|
||||||
|
logger.warning("Failed to stream-update card {}: code={}, msg={}", card_id, response.code, response.msg)
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Error stream-updating card {}: {}", card_id, e)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def _close_streaming_mode_sync(self, card_id: str, sequence: int) -> bool:
|
||||||
|
"""Turn off CardKit streaming_mode so the chat list preview exits the streaming placeholder.
|
||||||
|
|
||||||
|
Per Feishu docs, streaming cards keep a generating-style summary in the session list until
|
||||||
|
streaming_mode is set to false via card settings (after final content update).
|
||||||
|
Sequence must strictly exceed the previous card OpenAPI operation on this entity.
|
||||||
|
"""
|
||||||
|
from lark_oapi.api.cardkit.v1 import SettingsCardRequest, SettingsCardRequestBody
|
||||||
|
settings_payload = json.dumps({"config": {"streaming_mode": False}}, ensure_ascii=False)
|
||||||
|
try:
|
||||||
|
request = SettingsCardRequest.builder() \
|
||||||
|
.card_id(card_id) \
|
||||||
|
.request_body(
|
||||||
|
SettingsCardRequestBody.builder()
|
||||||
|
.settings(settings_payload)
|
||||||
|
.sequence(sequence)
|
||||||
|
.uuid(str(uuid.uuid4()))
|
||||||
|
.build()
|
||||||
|
).build()
|
||||||
|
response = self._client.cardkit.v1.card.settings(request)
|
||||||
|
if not response.success():
|
||||||
|
logger.warning(
|
||||||
|
"Failed to close streaming on card {}: code={}, msg={}",
|
||||||
|
card_id, response.code, response.msg,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Error closing streaming on card {}: {}", card_id, e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||||
|
"""Progressive streaming via CardKit: create card on first delta, stream-update on subsequent."""
|
||||||
|
if not self._client:
|
||||||
|
return
|
||||||
|
meta = metadata or {}
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
rid_type = "chat_id" if chat_id.startswith("oc_") else "open_id"
|
||||||
|
|
||||||
|
# --- stream end: final update or fallback ---
|
||||||
|
if meta.get("_stream_end"):
|
||||||
|
buf = self._stream_bufs.pop(chat_id, None)
|
||||||
|
if not buf or not buf.text:
|
||||||
|
return
|
||||||
|
if buf.card_id:
|
||||||
|
buf.sequence += 1
|
||||||
|
await loop.run_in_executor(
|
||||||
|
None, self._stream_update_text_sync, buf.card_id, buf.text, buf.sequence,
|
||||||
|
)
|
||||||
|
# Required so the chat list preview exits the streaming placeholder (Feishu streaming card docs).
|
||||||
|
buf.sequence += 1
|
||||||
|
await loop.run_in_executor(
|
||||||
|
None, self._close_streaming_mode_sync, buf.card_id, buf.sequence,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
for chunk in self._split_elements_by_table_limit(self._build_card_elements(buf.text)):
|
||||||
|
card = json.dumps({"config": {"wide_screen_mode": True}, "elements": chunk}, ensure_ascii=False)
|
||||||
|
await loop.run_in_executor(None, self._send_message_sync, rid_type, chat_id, "interactive", card)
|
||||||
|
return
|
||||||
|
|
||||||
|
# --- accumulate delta ---
|
||||||
|
buf = self._stream_bufs.get(chat_id)
|
||||||
|
if buf is None:
|
||||||
|
buf = _FeishuStreamBuf()
|
||||||
|
self._stream_bufs[chat_id] = buf
|
||||||
|
buf.text += delta
|
||||||
|
if not buf.text.strip():
|
||||||
|
return
|
||||||
|
|
||||||
|
now = time.monotonic()
|
||||||
|
if buf.card_id is None:
|
||||||
|
card_id = await loop.run_in_executor(None, self._create_streaming_card_sync, rid_type, chat_id)
|
||||||
|
if card_id:
|
||||||
|
buf.card_id = card_id
|
||||||
|
buf.sequence = 1
|
||||||
|
await loop.run_in_executor(None, self._stream_update_text_sync, card_id, buf.text, 1)
|
||||||
|
buf.last_edit = now
|
||||||
|
elif (now - buf.last_edit) >= self._STREAM_EDIT_INTERVAL:
|
||||||
|
buf.sequence += 1
|
||||||
|
await loop.run_in_executor(None, self._stream_update_text_sync, buf.card_id, buf.text, buf.sequence)
|
||||||
|
buf.last_edit = now
|
||||||
|
|
||||||
async def send(self, msg: OutboundMessage) -> None:
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
"""Send a message through Feishu, including media (images/files) if present."""
|
"""Send a message through Feishu, including media (images/files) if present."""
|
||||||
if not self._client:
|
if not self._client:
|
||||||
@@ -1031,6 +1186,7 @@ class FeishuChannel(BaseChannel):
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error sending Feishu message: {}", e)
|
logger.error("Error sending Feishu message: {}", e)
|
||||||
|
raise
|
||||||
|
|
||||||
def _on_message_sync(self, data: Any) -> None:
|
def _on_message_sync(self, data: Any) -> None:
|
||||||
"""
|
"""
|
||||||
|
|||||||
+111
-13
@@ -7,10 +7,14 @@ from typing import Any
|
|||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
|
from nanobot.bus.events import 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.config.schema import Config
|
from nanobot.config.schema import Config
|
||||||
|
|
||||||
|
# Retry delays for message sending (exponential backoff: 1s, 2s, 4s)
|
||||||
|
_SEND_RETRY_DELAYS = (1, 2, 4)
|
||||||
|
|
||||||
|
|
||||||
class ChannelManager:
|
class ChannelManager:
|
||||||
"""
|
"""
|
||||||
@@ -114,12 +118,20 @@ class ChannelManager:
|
|||||||
"""Dispatch outbound messages to the appropriate channel."""
|
"""Dispatch outbound messages to the appropriate channel."""
|
||||||
logger.info("Outbound dispatcher started")
|
logger.info("Outbound dispatcher started")
|
||||||
|
|
||||||
|
# Buffer for messages that couldn't be processed during delta coalescing
|
||||||
|
# (since asyncio.Queue doesn't support push_front)
|
||||||
|
pending: list[OutboundMessage] = []
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
msg = await asyncio.wait_for(
|
# First check pending buffer before waiting on queue
|
||||||
self.bus.consume_outbound(),
|
if pending:
|
||||||
timeout=1.0
|
msg = pending.pop(0)
|
||||||
)
|
else:
|
||||||
|
msg = await asyncio.wait_for(
|
||||||
|
self.bus.consume_outbound(),
|
||||||
|
timeout=1.0
|
||||||
|
)
|
||||||
|
|
||||||
if msg.metadata.get("_progress"):
|
if msg.metadata.get("_progress"):
|
||||||
if msg.metadata.get("_tool_hint") and not self.config.channels.send_tool_hints:
|
if msg.metadata.get("_tool_hint") and not self.config.channels.send_tool_hints:
|
||||||
@@ -127,17 +139,15 @@ class ChannelManager:
|
|||||||
if not msg.metadata.get("_tool_hint") and not self.config.channels.send_progress:
|
if not msg.metadata.get("_tool_hint") and not self.config.channels.send_progress:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
# Coalesce consecutive _stream_delta messages for the same (channel, chat_id)
|
||||||
|
# to reduce API calls and improve streaming latency
|
||||||
|
if msg.metadata.get("_stream_delta") and not msg.metadata.get("_stream_end"):
|
||||||
|
msg, extra_pending = self._coalesce_stream_deltas(msg)
|
||||||
|
pending.extend(extra_pending)
|
||||||
|
|
||||||
channel = self.channels.get(msg.channel)
|
channel = self.channels.get(msg.channel)
|
||||||
if channel:
|
if channel:
|
||||||
try:
|
await self._send_with_retry(channel, msg)
|
||||||
if msg.metadata.get("_stream_delta") or msg.metadata.get("_stream_end"):
|
|
||||||
await channel.send_delta(msg.chat_id, msg.content, msg.metadata)
|
|
||||||
elif msg.metadata.get("_streamed"):
|
|
||||||
pass
|
|
||||||
else:
|
|
||||||
await channel.send(msg)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error("Error sending to {}: {}", msg.channel, e)
|
|
||||||
else:
|
else:
|
||||||
logger.warning("Unknown channel: {}", msg.channel)
|
logger.warning("Unknown channel: {}", msg.channel)
|
||||||
|
|
||||||
@@ -146,6 +156,94 @@ class ChannelManager:
|
|||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
break
|
break
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _send_once(channel: BaseChannel, msg: OutboundMessage) -> None:
|
||||||
|
"""Send one outbound message without retry policy."""
|
||||||
|
if msg.metadata.get("_stream_delta") or msg.metadata.get("_stream_end"):
|
||||||
|
await channel.send_delta(msg.chat_id, msg.content, msg.metadata)
|
||||||
|
elif not msg.metadata.get("_streamed"):
|
||||||
|
await channel.send(msg)
|
||||||
|
|
||||||
|
def _coalesce_stream_deltas(
|
||||||
|
self, first_msg: OutboundMessage
|
||||||
|
) -> tuple[OutboundMessage, list[OutboundMessage]]:
|
||||||
|
"""Merge consecutive _stream_delta messages for the same (channel, chat_id).
|
||||||
|
|
||||||
|
This reduces the number of API calls when the queue has accumulated multiple
|
||||||
|
deltas, which happens when LLM generates faster than the channel can process.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple of (merged_message, list_of_non_matching_messages)
|
||||||
|
"""
|
||||||
|
target_key = (first_msg.channel, first_msg.chat_id)
|
||||||
|
combined_content = first_msg.content
|
||||||
|
final_metadata = dict(first_msg.metadata or {})
|
||||||
|
non_matching: list[OutboundMessage] = []
|
||||||
|
|
||||||
|
# Only merge consecutive deltas. As soon as we hit any other message,
|
||||||
|
# stop and hand that boundary back to the dispatcher via `pending`.
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
next_msg = self.bus.outbound.get_nowait()
|
||||||
|
except asyncio.QueueEmpty:
|
||||||
|
break
|
||||||
|
|
||||||
|
# Check if this message belongs to the same stream
|
||||||
|
same_target = (next_msg.channel, next_msg.chat_id) == target_key
|
||||||
|
is_delta = next_msg.metadata and next_msg.metadata.get("_stream_delta")
|
||||||
|
is_end = next_msg.metadata and next_msg.metadata.get("_stream_end")
|
||||||
|
|
||||||
|
if same_target and is_delta and not final_metadata.get("_stream_end"):
|
||||||
|
# Accumulate content
|
||||||
|
combined_content += next_msg.content
|
||||||
|
# If we see _stream_end, remember it and stop coalescing this stream
|
||||||
|
if is_end:
|
||||||
|
final_metadata["_stream_end"] = True
|
||||||
|
# Stream ended - stop coalescing this stream
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
# First non-matching message defines the coalescing boundary.
|
||||||
|
non_matching.append(next_msg)
|
||||||
|
break
|
||||||
|
|
||||||
|
merged = OutboundMessage(
|
||||||
|
channel=first_msg.channel,
|
||||||
|
chat_id=first_msg.chat_id,
|
||||||
|
content=combined_content,
|
||||||
|
metadata=final_metadata,
|
||||||
|
)
|
||||||
|
return merged, non_matching
|
||||||
|
|
||||||
|
async def _send_with_retry(self, channel: BaseChannel, msg: OutboundMessage) -> None:
|
||||||
|
"""Send a message with retry on failure using exponential backoff.
|
||||||
|
|
||||||
|
Note: CancelledError is re-raised to allow graceful shutdown.
|
||||||
|
"""
|
||||||
|
max_attempts = max(self.config.channels.send_max_retries, 1)
|
||||||
|
|
||||||
|
for attempt in range(max_attempts):
|
||||||
|
try:
|
||||||
|
await self._send_once(channel, msg)
|
||||||
|
return # Send succeeded
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise # Propagate cancellation for graceful shutdown
|
||||||
|
except Exception as e:
|
||||||
|
if attempt == max_attempts - 1:
|
||||||
|
logger.error(
|
||||||
|
"Failed to send to {} after {} attempts: {} - {}",
|
||||||
|
msg.channel, max_attempts, type(e).__name__, e
|
||||||
|
)
|
||||||
|
return
|
||||||
|
delay = _SEND_RETRY_DELAYS[min(attempt, len(_SEND_RETRY_DELAYS) - 1)]
|
||||||
|
logger.warning(
|
||||||
|
"Send to {} failed (attempt {}/{}): {}, retrying in {}s",
|
||||||
|
msg.channel, attempt + 1, max_attempts, type(e).__name__, delay
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
await asyncio.sleep(delay)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise # Propagate cancellation during sleep
|
||||||
|
|
||||||
def get_channel(self, name: str) -> BaseChannel | None:
|
def get_channel(self, name: str) -> BaseChannel | None:
|
||||||
"""Get a channel by name."""
|
"""Get a channel by name."""
|
||||||
return self.channels.get(name)
|
return self.channels.get(name)
|
||||||
|
|||||||
+116
-8
@@ -3,6 +3,8 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import mimetypes
|
import mimetypes
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Literal, TypeAlias
|
from typing import Any, Literal, TypeAlias
|
||||||
|
|
||||||
@@ -28,8 +30,8 @@ try:
|
|||||||
RoomSendError,
|
RoomSendError,
|
||||||
RoomTypingError,
|
RoomTypingError,
|
||||||
SyncError,
|
SyncError,
|
||||||
UploadError,
|
UploadError, RoomSendResponse,
|
||||||
)
|
)
|
||||||
from nio.crypto.attachments import decrypt_attachment
|
from nio.crypto.attachments import decrypt_attachment
|
||||||
from nio.exceptions import EncryptionError
|
from nio.exceptions import EncryptionError
|
||||||
except ImportError as e:
|
except ImportError as e:
|
||||||
@@ -97,6 +99,22 @@ MATRIX_HTML_CLEANER = nh3.Cleaner(
|
|||||||
link_rel="noopener noreferrer",
|
link_rel="noopener noreferrer",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _StreamBuf:
|
||||||
|
"""
|
||||||
|
Represents a buffer for managing LLM response stream data.
|
||||||
|
|
||||||
|
:ivar text: Stores the text content of the buffer.
|
||||||
|
:type text: str
|
||||||
|
:ivar event_id: Identifier for the associated event. None indicates no
|
||||||
|
specific event association.
|
||||||
|
:type event_id: str | None
|
||||||
|
:ivar last_edit: Timestamp of the most recent edit to the buffer.
|
||||||
|
:type last_edit: float
|
||||||
|
"""
|
||||||
|
text: str = ""
|
||||||
|
event_id: str | None = None
|
||||||
|
last_edit: float = 0.0
|
||||||
|
|
||||||
def _render_markdown_html(text: str) -> str | None:
|
def _render_markdown_html(text: str) -> str | None:
|
||||||
"""Render markdown to sanitized HTML; returns None for plain text."""
|
"""Render markdown to sanitized HTML; returns None for plain text."""
|
||||||
@@ -114,12 +132,47 @@ def _render_markdown_html(text: str) -> str | None:
|
|||||||
return formatted
|
return formatted
|
||||||
|
|
||||||
|
|
||||||
def _build_matrix_text_content(text: str) -> dict[str, object]:
|
def _build_matrix_text_content(
|
||||||
"""Build Matrix m.text payload with optional HTML formatted_body."""
|
text: str,
|
||||||
|
event_id: str | None = None,
|
||||||
|
thread_relates_to: dict[str, object] | None = None,
|
||||||
|
) -> dict[str, object]:
|
||||||
|
"""
|
||||||
|
Constructs and returns a dictionary representing the matrix text content with optional
|
||||||
|
HTML formatting and reference to an existing event for replacement. This function is
|
||||||
|
primarily used to create content payloads compatible with the Matrix messaging protocol.
|
||||||
|
|
||||||
|
:param text: The plain text content to include in the message.
|
||||||
|
:type text: str
|
||||||
|
:param event_id: Optional ID of the event to replace. If provided, the function will
|
||||||
|
include information indicating that the message is a replacement of the specified
|
||||||
|
event.
|
||||||
|
:type event_id: str | None
|
||||||
|
:param thread_relates_to: Optional Matrix thread relation metadata. For edits this is
|
||||||
|
stored in ``m.new_content`` so the replacement remains in the same thread.
|
||||||
|
:type thread_relates_to: dict[str, object] | None
|
||||||
|
:return: A dictionary containing the matrix text content, potentially enriched with
|
||||||
|
HTML formatting and replacement metadata if applicable.
|
||||||
|
:rtype: dict[str, object]
|
||||||
|
"""
|
||||||
content: dict[str, object] = {"msgtype": "m.text", "body": text, "m.mentions": {}}
|
content: dict[str, object] = {"msgtype": "m.text", "body": text, "m.mentions": {}}
|
||||||
if html := _render_markdown_html(text):
|
if html := _render_markdown_html(text):
|
||||||
content["format"] = MATRIX_HTML_FORMAT
|
content["format"] = MATRIX_HTML_FORMAT
|
||||||
content["formatted_body"] = html
|
content["formatted_body"] = html
|
||||||
|
if event_id:
|
||||||
|
content["m.new_content"] = {
|
||||||
|
"body": text,
|
||||||
|
"msgtype": "m.text",
|
||||||
|
}
|
||||||
|
content["m.relates_to"] = {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": event_id,
|
||||||
|
}
|
||||||
|
if thread_relates_to:
|
||||||
|
content["m.new_content"]["m.relates_to"] = thread_relates_to
|
||||||
|
elif thread_relates_to:
|
||||||
|
content["m.relates_to"] = thread_relates_to
|
||||||
|
|
||||||
return content
|
return content
|
||||||
|
|
||||||
|
|
||||||
@@ -159,7 +212,8 @@ class MatrixConfig(Base):
|
|||||||
allow_from: list[str] = Field(default_factory=list)
|
allow_from: list[str] = Field(default_factory=list)
|
||||||
group_policy: Literal["open", "mention", "allowlist"] = "open"
|
group_policy: Literal["open", "mention", "allowlist"] = "open"
|
||||||
group_allow_from: list[str] = Field(default_factory=list)
|
group_allow_from: list[str] = Field(default_factory=list)
|
||||||
allow_room_mentions: bool = False
|
allow_room_mentions: bool = False,
|
||||||
|
streaming: bool = False
|
||||||
|
|
||||||
|
|
||||||
class MatrixChannel(BaseChannel):
|
class MatrixChannel(BaseChannel):
|
||||||
@@ -167,6 +221,8 @@ class MatrixChannel(BaseChannel):
|
|||||||
|
|
||||||
name = "matrix"
|
name = "matrix"
|
||||||
display_name = "Matrix"
|
display_name = "Matrix"
|
||||||
|
_STREAM_EDIT_INTERVAL = 2 # min seconds between edit_message_text calls
|
||||||
|
monotonic_time = time.monotonic
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def default_config(cls) -> dict[str, Any]:
|
def default_config(cls) -> dict[str, Any]:
|
||||||
@@ -192,6 +248,8 @@ class MatrixChannel(BaseChannel):
|
|||||||
)
|
)
|
||||||
self._server_upload_limit_bytes: int | None = None
|
self._server_upload_limit_bytes: int | None = None
|
||||||
self._server_upload_limit_checked = False
|
self._server_upload_limit_checked = False
|
||||||
|
self._stream_bufs: dict[str, _StreamBuf] = {}
|
||||||
|
|
||||||
|
|
||||||
async def start(self) -> None:
|
async def start(self) -> None:
|
||||||
"""Start Matrix client and begin sync loop."""
|
"""Start Matrix client and begin sync loop."""
|
||||||
@@ -297,14 +355,17 @@ class MatrixChannel(BaseChannel):
|
|||||||
room = getattr(self.client, "rooms", {}).get(room_id)
|
room = getattr(self.client, "rooms", {}).get(room_id)
|
||||||
return bool(getattr(room, "encrypted", False))
|
return bool(getattr(room, "encrypted", False))
|
||||||
|
|
||||||
async def _send_room_content(self, room_id: str, content: dict[str, Any]) -> None:
|
async def _send_room_content(self, room_id: str,
|
||||||
|
content: dict[str, Any]) -> None | RoomSendResponse | RoomSendError:
|
||||||
"""Send m.room.message with E2EE options."""
|
"""Send m.room.message with E2EE options."""
|
||||||
if not self.client:
|
if not self.client:
|
||||||
return
|
return None
|
||||||
kwargs: dict[str, Any] = {"room_id": room_id, "message_type": "m.room.message", "content": content}
|
kwargs: dict[str, Any] = {"room_id": room_id, "message_type": "m.room.message", "content": content}
|
||||||
|
|
||||||
if self.config.e2ee_enabled:
|
if self.config.e2ee_enabled:
|
||||||
kwargs["ignore_unverified_devices"] = True
|
kwargs["ignore_unverified_devices"] = True
|
||||||
await self.client.room_send(**kwargs)
|
response = await self.client.room_send(**kwargs)
|
||||||
|
return response
|
||||||
|
|
||||||
async def _resolve_server_upload_limit_bytes(self) -> int | None:
|
async def _resolve_server_upload_limit_bytes(self) -> int | None:
|
||||||
"""Query homeserver upload limit once per channel lifecycle."""
|
"""Query homeserver upload limit once per channel lifecycle."""
|
||||||
@@ -414,6 +475,53 @@ class MatrixChannel(BaseChannel):
|
|||||||
if not is_progress:
|
if not is_progress:
|
||||||
await self._stop_typing_keepalive(msg.chat_id, clear_typing=True)
|
await self._stop_typing_keepalive(msg.chat_id, clear_typing=True)
|
||||||
|
|
||||||
|
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||||
|
meta = metadata or {}
|
||||||
|
relates_to = self._build_thread_relates_to(metadata)
|
||||||
|
|
||||||
|
if meta.get("_stream_end"):
|
||||||
|
buf = self._stream_bufs.pop(chat_id, None)
|
||||||
|
if not buf or not buf.event_id or not buf.text:
|
||||||
|
return
|
||||||
|
|
||||||
|
await self._stop_typing_keepalive(chat_id, clear_typing=True)
|
||||||
|
|
||||||
|
content = _build_matrix_text_content(
|
||||||
|
buf.text,
|
||||||
|
buf.event_id,
|
||||||
|
thread_relates_to=relates_to,
|
||||||
|
)
|
||||||
|
await self._send_room_content(chat_id, content)
|
||||||
|
return
|
||||||
|
|
||||||
|
buf = self._stream_bufs.get(chat_id)
|
||||||
|
if buf is None:
|
||||||
|
buf = _StreamBuf()
|
||||||
|
self._stream_bufs[chat_id] = buf
|
||||||
|
buf.text += delta
|
||||||
|
|
||||||
|
if not buf.text.strip():
|
||||||
|
return
|
||||||
|
|
||||||
|
now = self.monotonic_time()
|
||||||
|
|
||||||
|
if not buf.last_edit or (now - buf.last_edit) >= self._STREAM_EDIT_INTERVAL:
|
||||||
|
try:
|
||||||
|
content = _build_matrix_text_content(
|
||||||
|
buf.text,
|
||||||
|
buf.event_id,
|
||||||
|
thread_relates_to=relates_to,
|
||||||
|
)
|
||||||
|
response = await self._send_room_content(chat_id, content)
|
||||||
|
buf.last_edit = now
|
||||||
|
if not buf.event_id:
|
||||||
|
# we are editing the same message all the time, so only the first time the event id needs to be set
|
||||||
|
buf.event_id = response.event_id
|
||||||
|
except Exception:
|
||||||
|
await self._stop_typing_keepalive(chat_id, clear_typing=True)
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
def _register_event_callbacks(self) -> None:
|
def _register_event_callbacks(self) -> None:
|
||||||
self.client.add_event_callback(self._on_message, RoomMessageText)
|
self.client.add_event_callback(self._on_message, RoomMessageText)
|
||||||
self.client.add_event_callback(self._on_media_message, MATRIX_MEDIA_EVENT_FILTER)
|
self.client.add_event_callback(self._on_media_message, MATRIX_MEDIA_EVENT_FILTER)
|
||||||
|
|||||||
@@ -374,6 +374,7 @@ class MochatChannel(BaseChannel):
|
|||||||
content, msg.reply_to)
|
content, msg.reply_to)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Failed to send Mochat message: {}", e)
|
logger.error("Failed to send Mochat message: {}", e)
|
||||||
|
raise
|
||||||
|
|
||||||
# ---- config / init helpers ---------------------------------------------
|
# ---- config / init helpers ---------------------------------------------
|
||||||
|
|
||||||
|
|||||||
@@ -145,6 +145,7 @@ class SlackChannel(BaseChannel):
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error sending Slack message: {}", e)
|
logger.error("Error sending Slack message: {}", e)
|
||||||
|
raise
|
||||||
|
|
||||||
async def _on_socket_request(
|
async def _on_socket_request(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ from typing import Any, Literal
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
from telegram import BotCommand, ReactionTypeEmoji, ReplyParameters, Update
|
from telegram import BotCommand, ReactionTypeEmoji, ReplyParameters, Update
|
||||||
from telegram.error import TimedOut
|
from telegram.error import BadRequest, TimedOut
|
||||||
from telegram.ext import Application, CommandHandler, ContextTypes, MessageHandler, filters
|
from telegram.ext import Application, CommandHandler, ContextTypes, MessageHandler, filters
|
||||||
from telegram.request import HTTPXRequest
|
from telegram.request import HTTPXRequest
|
||||||
|
|
||||||
@@ -163,6 +163,7 @@ class _StreamBuf:
|
|||||||
text: str = ""
|
text: str = ""
|
||||||
message_id: int | None = None
|
message_id: int | None = None
|
||||||
last_edit: float = 0.0
|
last_edit: float = 0.0
|
||||||
|
stream_id: str | None = None
|
||||||
|
|
||||||
|
|
||||||
class TelegramConfig(Base):
|
class TelegramConfig(Base):
|
||||||
@@ -476,6 +477,11 @@ class TelegramChannel(BaseChannel):
|
|||||||
)
|
)
|
||||||
except Exception as e2:
|
except Exception as e2:
|
||||||
logger.error("Error sending Telegram message: {}", e2)
|
logger.error("Error sending Telegram message: {}", e2)
|
||||||
|
raise
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_not_modified_error(exc: Exception) -> bool:
|
||||||
|
return isinstance(exc, BadRequest) and "message is not modified" in str(exc).lower()
|
||||||
|
|
||||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||||
"""Progressive message editing: send on first delta, edit on subsequent ones."""
|
"""Progressive message editing: send on first delta, edit on subsequent ones."""
|
||||||
@@ -483,11 +489,14 @@ class TelegramChannel(BaseChannel):
|
|||||||
return
|
return
|
||||||
meta = metadata or {}
|
meta = metadata or {}
|
||||||
int_chat_id = int(chat_id)
|
int_chat_id = int(chat_id)
|
||||||
|
stream_id = meta.get("_stream_id")
|
||||||
|
|
||||||
if meta.get("_stream_end"):
|
if meta.get("_stream_end"):
|
||||||
buf = self._stream_bufs.pop(chat_id, None)
|
buf = self._stream_bufs.get(chat_id)
|
||||||
if not buf or not buf.message_id or not buf.text:
|
if not buf or not buf.message_id or not buf.text:
|
||||||
return
|
return
|
||||||
|
if stream_id is not None and buf.stream_id is not None and buf.stream_id != stream_id:
|
||||||
|
return
|
||||||
self._stop_typing(chat_id)
|
self._stop_typing(chat_id)
|
||||||
try:
|
try:
|
||||||
html = _markdown_to_telegram_html(buf.text)
|
html = _markdown_to_telegram_html(buf.text)
|
||||||
@@ -497,6 +506,10 @@ class TelegramChannel(BaseChannel):
|
|||||||
text=html, parse_mode="HTML",
|
text=html, parse_mode="HTML",
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
if self._is_not_modified_error(e):
|
||||||
|
logger.debug("Final stream edit already applied for {}", chat_id)
|
||||||
|
self._stream_bufs.pop(chat_id, None)
|
||||||
|
return
|
||||||
logger.debug("Final stream edit failed (HTML), trying plain: {}", e)
|
logger.debug("Final stream edit failed (HTML), trying plain: {}", e)
|
||||||
try:
|
try:
|
||||||
await self._call_with_retry(
|
await self._call_with_retry(
|
||||||
@@ -504,14 +517,22 @@ class TelegramChannel(BaseChannel):
|
|||||||
chat_id=int_chat_id, message_id=buf.message_id,
|
chat_id=int_chat_id, message_id=buf.message_id,
|
||||||
text=buf.text,
|
text=buf.text,
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception as e2:
|
||||||
pass
|
if self._is_not_modified_error(e2):
|
||||||
|
logger.debug("Final stream plain edit already applied for {}", chat_id)
|
||||||
|
self._stream_bufs.pop(chat_id, None)
|
||||||
|
return
|
||||||
|
logger.warning("Final stream edit failed: {}", e2)
|
||||||
|
raise # Let ChannelManager handle retry
|
||||||
|
self._stream_bufs.pop(chat_id, None)
|
||||||
return
|
return
|
||||||
|
|
||||||
buf = self._stream_bufs.get(chat_id)
|
buf = self._stream_bufs.get(chat_id)
|
||||||
if buf is None:
|
if buf is None or (stream_id is not None and buf.stream_id is not None and buf.stream_id != stream_id):
|
||||||
buf = _StreamBuf()
|
buf = _StreamBuf(stream_id=stream_id)
|
||||||
self._stream_bufs[chat_id] = buf
|
self._stream_bufs[chat_id] = buf
|
||||||
|
elif buf.stream_id is None:
|
||||||
|
buf.stream_id = stream_id
|
||||||
buf.text += delta
|
buf.text += delta
|
||||||
|
|
||||||
if not buf.text.strip():
|
if not buf.text.strip():
|
||||||
@@ -528,6 +549,7 @@ class TelegramChannel(BaseChannel):
|
|||||||
buf.last_edit = now
|
buf.last_edit = now
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("Stream initial send failed: {}", e)
|
logger.warning("Stream initial send failed: {}", e)
|
||||||
|
raise # Let ChannelManager handle retry
|
||||||
elif (now - buf.last_edit) >= self._STREAM_EDIT_INTERVAL:
|
elif (now - buf.last_edit) >= self._STREAM_EDIT_INTERVAL:
|
||||||
try:
|
try:
|
||||||
await self._call_with_retry(
|
await self._call_with_retry(
|
||||||
@@ -536,8 +558,12 @@ class TelegramChannel(BaseChannel):
|
|||||||
text=buf.text,
|
text=buf.text,
|
||||||
)
|
)
|
||||||
buf.last_edit = now
|
buf.last_edit = now
|
||||||
except Exception:
|
except Exception as e:
|
||||||
pass
|
if self._is_not_modified_error(e):
|
||||||
|
buf.last_edit = now
|
||||||
|
return
|
||||||
|
logger.warning("Stream edit failed: {}", e)
|
||||||
|
raise # Let ChannelManager handle retry
|
||||||
|
|
||||||
async def _on_start(self, update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
async def _on_start(self, update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||||
"""Handle /start command."""
|
"""Handle /start command."""
|
||||||
@@ -890,7 +916,12 @@ class TelegramChannel(BaseChannel):
|
|||||||
|
|
||||||
async def _on_error(self, update: object, context: ContextTypes.DEFAULT_TYPE) -> None:
|
async def _on_error(self, update: object, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||||
"""Log polling / handler errors instead of silently swallowing them."""
|
"""Log polling / handler errors instead of silently swallowing them."""
|
||||||
logger.error("Telegram error: {}", context.error)
|
from telegram.error import NetworkError, TimedOut
|
||||||
|
|
||||||
|
if isinstance(context.error, (NetworkError, TimedOut)):
|
||||||
|
logger.warning("Telegram network issue: {}", str(context.error))
|
||||||
|
else:
|
||||||
|
logger.error("Telegram error: {}", context.error)
|
||||||
|
|
||||||
def _get_extension(
|
def _get_extension(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -368,3 +368,4 @@ class WecomChannel(BaseChannel):
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error sending WeCom message: {}", e)
|
logger.error("Error sending WeCom message: {}", e)
|
||||||
|
raise
|
||||||
|
|||||||
+362
-95
@@ -15,6 +15,7 @@ import hashlib
|
|||||||
import json
|
import json
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import os
|
import os
|
||||||
|
import random
|
||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
@@ -53,7 +54,26 @@ MESSAGE_TYPE_BOT = 2
|
|||||||
MESSAGE_STATE_FINISH = 2
|
MESSAGE_STATE_FINISH = 2
|
||||||
|
|
||||||
WEIXIN_MAX_MESSAGE_LEN = 4000
|
WEIXIN_MAX_MESSAGE_LEN = 4000
|
||||||
WEIXIN_CHANNEL_VERSION = "1.0.3"
|
WEIXIN_CHANNEL_VERSION = "2.1.1"
|
||||||
|
ILINK_APP_ID = "bot"
|
||||||
|
|
||||||
|
|
||||||
|
def _build_client_version(version: str) -> int:
|
||||||
|
"""Encode semantic version as 0x00MMNNPP (major/minor/patch in one uint32)."""
|
||||||
|
parts = version.split(".")
|
||||||
|
|
||||||
|
def _as_int(idx: int) -> int:
|
||||||
|
try:
|
||||||
|
return int(parts[idx])
|
||||||
|
except Exception:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
major = _as_int(0)
|
||||||
|
minor = _as_int(1)
|
||||||
|
patch = _as_int(2)
|
||||||
|
return ((major & 0xFF) << 16) | ((minor & 0xFF) << 8) | (patch & 0xFF)
|
||||||
|
|
||||||
|
ILINK_APP_CLIENT_VERSION = _build_client_version(WEIXIN_CHANNEL_VERSION)
|
||||||
BASE_INFO: dict[str, str] = {"channel_version": WEIXIN_CHANNEL_VERSION}
|
BASE_INFO: dict[str, str] = {"channel_version": WEIXIN_CHANNEL_VERSION}
|
||||||
|
|
||||||
# Session-expired error code
|
# Session-expired error code
|
||||||
@@ -65,18 +85,32 @@ MAX_CONSECUTIVE_FAILURES = 3
|
|||||||
BACKOFF_DELAY_S = 30
|
BACKOFF_DELAY_S = 30
|
||||||
RETRY_DELAY_S = 2
|
RETRY_DELAY_S = 2
|
||||||
MAX_QR_REFRESH_COUNT = 3
|
MAX_QR_REFRESH_COUNT = 3
|
||||||
|
TYPING_STATUS_TYPING = 1
|
||||||
|
TYPING_STATUS_CANCEL = 2
|
||||||
|
TYPING_TICKET_TTL_S = 24 * 60 * 60
|
||||||
|
TYPING_KEEPALIVE_INTERVAL_S = 5
|
||||||
|
CONFIG_CACHE_INITIAL_RETRY_S = 2
|
||||||
|
CONFIG_CACHE_MAX_RETRY_S = 60 * 60
|
||||||
|
|
||||||
# Default long-poll timeout; overridden by server via longpolling_timeout_ms.
|
# Default long-poll timeout; overridden by server via longpolling_timeout_ms.
|
||||||
DEFAULT_LONG_POLL_TIMEOUT_S = 35
|
DEFAULT_LONG_POLL_TIMEOUT_S = 35
|
||||||
|
|
||||||
# Media-type codes for getuploadurl (1=image, 2=video, 3=file)
|
# Media-type codes for getuploadurl (1=image, 2=video, 3=file, 4=voice)
|
||||||
UPLOAD_MEDIA_IMAGE = 1
|
UPLOAD_MEDIA_IMAGE = 1
|
||||||
UPLOAD_MEDIA_VIDEO = 2
|
UPLOAD_MEDIA_VIDEO = 2
|
||||||
UPLOAD_MEDIA_FILE = 3
|
UPLOAD_MEDIA_FILE = 3
|
||||||
|
UPLOAD_MEDIA_VOICE = 4
|
||||||
|
|
||||||
# File extensions considered as images / videos for outbound media
|
# File extensions considered as images / videos for outbound media
|
||||||
_IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".webp", ".tiff", ".ico", ".svg"}
|
_IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".webp", ".tiff", ".ico", ".svg"}
|
||||||
_VIDEO_EXTS = {".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv"}
|
_VIDEO_EXTS = {".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv"}
|
||||||
|
_VOICE_EXTS = {".mp3", ".wav", ".amr", ".silk", ".ogg", ".m4a", ".aac", ".flac"}
|
||||||
|
|
||||||
|
|
||||||
|
def _has_downloadable_media_locator(media: dict[str, Any] | None) -> bool:
|
||||||
|
if not isinstance(media, dict):
|
||||||
|
return False
|
||||||
|
return bool(str(media.get("encrypt_query_param", "") or "") or str(media.get("full_url", "") or "").strip())
|
||||||
|
|
||||||
|
|
||||||
class WeixinConfig(Base):
|
class WeixinConfig(Base):
|
||||||
@@ -124,6 +158,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
self._poll_task: asyncio.Task | None = None
|
self._poll_task: asyncio.Task | None = None
|
||||||
self._next_poll_timeout_s: int = DEFAULT_LONG_POLL_TIMEOUT_S
|
self._next_poll_timeout_s: int = DEFAULT_LONG_POLL_TIMEOUT_S
|
||||||
self._session_pause_until: float = 0.0
|
self._session_pause_until: float = 0.0
|
||||||
|
self._typing_tickets: dict[str, dict[str, Any]] = {}
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# State persistence
|
# State persistence
|
||||||
@@ -162,8 +197,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
if base_url:
|
if base_url:
|
||||||
self.config.base_url = base_url
|
self.config.base_url = base_url
|
||||||
return bool(self._token)
|
return bool(self._token)
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.warning("Failed to load WeChat state: {}", e)
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def _save_state(self) -> None:
|
def _save_state(self) -> None:
|
||||||
@@ -176,8 +210,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
"base_url": self.config.base_url,
|
"base_url": self.config.base_url,
|
||||||
}
|
}
|
||||||
state_file.write_text(json.dumps(data, ensure_ascii=False))
|
state_file.write_text(json.dumps(data, ensure_ascii=False))
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.warning("Failed to save WeChat state: {}", e)
|
pass
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# HTTP helpers (matches api.ts buildHeaders / apiFetch)
|
# HTTP helpers (matches api.ts buildHeaders / apiFetch)
|
||||||
@@ -199,6 +233,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
"X-WECHAT-UIN": self._random_wechat_uin(),
|
"X-WECHAT-UIN": self._random_wechat_uin(),
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
"AuthorizationType": "ilink_bot_token",
|
"AuthorizationType": "ilink_bot_token",
|
||||||
|
"iLink-App-Id": ILINK_APP_ID,
|
||||||
|
"iLink-App-ClientVersion": str(ILINK_APP_CLIENT_VERSION),
|
||||||
}
|
}
|
||||||
if auth and self._token:
|
if auth and self._token:
|
||||||
headers["Authorization"] = f"Bearer {self._token}"
|
headers["Authorization"] = f"Bearer {self._token}"
|
||||||
@@ -206,6 +242,15 @@ class WeixinChannel(BaseChannel):
|
|||||||
headers["SKRouteTag"] = str(self.config.route_tag).strip()
|
headers["SKRouteTag"] = str(self.config.route_tag).strip()
|
||||||
return headers
|
return headers
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_retryable_media_download_error(err: Exception) -> bool:
|
||||||
|
if isinstance(err, httpx.TimeoutException | httpx.TransportError):
|
||||||
|
return True
|
||||||
|
if isinstance(err, httpx.HTTPStatusError):
|
||||||
|
status_code = err.response.status_code if err.response is not None else 0
|
||||||
|
return status_code >= 500
|
||||||
|
return False
|
||||||
|
|
||||||
async def _api_get(
|
async def _api_get(
|
||||||
self,
|
self,
|
||||||
endpoint: str,
|
endpoint: str,
|
||||||
@@ -223,6 +268,25 @@ class WeixinChannel(BaseChannel):
|
|||||||
resp.raise_for_status()
|
resp.raise_for_status()
|
||||||
return resp.json()
|
return resp.json()
|
||||||
|
|
||||||
|
async def _api_get_with_base(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
base_url: str,
|
||||||
|
endpoint: str,
|
||||||
|
params: dict | None = None,
|
||||||
|
auth: bool = True,
|
||||||
|
extra_headers: dict[str, str] | None = None,
|
||||||
|
) -> dict:
|
||||||
|
"""GET helper that allows overriding base_url for QR redirect polling."""
|
||||||
|
assert self._client is not None
|
||||||
|
url = f"{base_url.rstrip('/')}/{endpoint}"
|
||||||
|
hdrs = self._make_headers(auth=auth)
|
||||||
|
if extra_headers:
|
||||||
|
hdrs.update(extra_headers)
|
||||||
|
resp = await self._client.get(url, params=params, headers=hdrs)
|
||||||
|
resp.raise_for_status()
|
||||||
|
return resp.json()
|
||||||
|
|
||||||
async def _api_post(
|
async def _api_post(
|
||||||
self,
|
self,
|
||||||
endpoint: str,
|
endpoint: str,
|
||||||
@@ -259,23 +323,27 @@ class WeixinChannel(BaseChannel):
|
|||||||
async def _qr_login(self) -> bool:
|
async def _qr_login(self) -> bool:
|
||||||
"""Perform QR code login flow. Returns True on success."""
|
"""Perform QR code login flow. Returns True on success."""
|
||||||
try:
|
try:
|
||||||
logger.info("Starting WeChat QR code login...")
|
|
||||||
refresh_count = 0
|
refresh_count = 0
|
||||||
qrcode_id, scan_url = await self._fetch_qr_code()
|
qrcode_id, scan_url = await self._fetch_qr_code()
|
||||||
self._print_qr_code(scan_url)
|
self._print_qr_code(scan_url)
|
||||||
|
current_poll_base_url = self.config.base_url
|
||||||
|
|
||||||
logger.info("Waiting for QR code scan...")
|
|
||||||
while self._running:
|
while self._running:
|
||||||
try:
|
try:
|
||||||
# Reference plugin sends iLink-App-ClientVersion header for
|
status_data = await self._api_get_with_base(
|
||||||
# QR status polling (login-qr.ts:81).
|
base_url=current_poll_base_url,
|
||||||
status_data = await self._api_get(
|
endpoint="ilink/bot/get_qrcode_status",
|
||||||
"ilink/bot/get_qrcode_status",
|
|
||||||
params={"qrcode": qrcode_id},
|
params={"qrcode": qrcode_id},
|
||||||
auth=False,
|
auth=False,
|
||||||
extra_headers={"iLink-App-ClientVersion": "1"},
|
|
||||||
)
|
)
|
||||||
except httpx.TimeoutException:
|
except Exception as e:
|
||||||
|
if self._is_retryable_qr_poll_error(e):
|
||||||
|
await asyncio.sleep(1)
|
||||||
|
continue
|
||||||
|
raise
|
||||||
|
|
||||||
|
if not isinstance(status_data, dict):
|
||||||
|
await asyncio.sleep(1)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
status = status_data.get("status", "")
|
status = status_data.get("status", "")
|
||||||
@@ -298,8 +366,15 @@ class WeixinChannel(BaseChannel):
|
|||||||
else:
|
else:
|
||||||
logger.error("Login confirmed but no bot_token in response")
|
logger.error("Login confirmed but no bot_token in response")
|
||||||
return False
|
return False
|
||||||
elif status == "scaned":
|
elif status == "scaned_but_redirect":
|
||||||
logger.info("QR code scanned, waiting for confirmation...")
|
redirect_host = str(status_data.get("redirect_host", "") or "").strip()
|
||||||
|
if redirect_host:
|
||||||
|
if redirect_host.startswith("http://") or redirect_host.startswith("https://"):
|
||||||
|
redirected_base = redirect_host
|
||||||
|
else:
|
||||||
|
redirected_base = f"https://{redirect_host}"
|
||||||
|
if redirected_base != current_poll_base_url:
|
||||||
|
current_poll_base_url = redirected_base
|
||||||
elif status == "expired":
|
elif status == "expired":
|
||||||
refresh_count += 1
|
refresh_count += 1
|
||||||
if refresh_count > MAX_QR_REFRESH_COUNT:
|
if refresh_count > MAX_QR_REFRESH_COUNT:
|
||||||
@@ -309,14 +384,9 @@ class WeixinChannel(BaseChannel):
|
|||||||
MAX_QR_REFRESH_COUNT,
|
MAX_QR_REFRESH_COUNT,
|
||||||
)
|
)
|
||||||
return False
|
return False
|
||||||
logger.warning(
|
|
||||||
"QR code expired, refreshing... ({}/{})",
|
|
||||||
refresh_count,
|
|
||||||
MAX_QR_REFRESH_COUNT,
|
|
||||||
)
|
|
||||||
qrcode_id, scan_url = await self._fetch_qr_code()
|
qrcode_id, scan_url = await self._fetch_qr_code()
|
||||||
|
current_poll_base_url = self.config.base_url
|
||||||
self._print_qr_code(scan_url)
|
self._print_qr_code(scan_url)
|
||||||
logger.info("New QR code generated, waiting for scan...")
|
|
||||||
continue
|
continue
|
||||||
# status == "wait" — keep polling
|
# status == "wait" — keep polling
|
||||||
|
|
||||||
@@ -327,6 +397,16 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_retryable_qr_poll_error(err: Exception) -> bool:
|
||||||
|
if isinstance(err, httpx.TimeoutException | httpx.TransportError):
|
||||||
|
return True
|
||||||
|
if isinstance(err, httpx.HTTPStatusError):
|
||||||
|
status_code = err.response.status_code if err.response is not None else 0
|
||||||
|
if status_code >= 500:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _print_qr_code(url: str) -> None:
|
def _print_qr_code(url: str) -> None:
|
||||||
try:
|
try:
|
||||||
@@ -337,7 +417,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
qr.make(fit=True)
|
qr.make(fit=True)
|
||||||
qr.print_ascii(invert=True)
|
qr.print_ascii(invert=True)
|
||||||
except ImportError:
|
except ImportError:
|
||||||
logger.info("QR code URL (install 'qrcode' for terminal display): {}", url)
|
|
||||||
print(f"\nLogin URL: {url}\n")
|
print(f"\nLogin URL: {url}\n")
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@@ -399,12 +478,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
if not self._running:
|
if not self._running:
|
||||||
break
|
break
|
||||||
consecutive_failures += 1
|
consecutive_failures += 1
|
||||||
logger.error(
|
|
||||||
"WeChat poll error ({}/{}): {}",
|
|
||||||
consecutive_failures,
|
|
||||||
MAX_CONSECUTIVE_FAILURES,
|
|
||||||
e,
|
|
||||||
)
|
|
||||||
if consecutive_failures >= MAX_CONSECUTIVE_FAILURES:
|
if consecutive_failures >= MAX_CONSECUTIVE_FAILURES:
|
||||||
consecutive_failures = 0
|
consecutive_failures = 0
|
||||||
await asyncio.sleep(BACKOFF_DELAY_S)
|
await asyncio.sleep(BACKOFF_DELAY_S)
|
||||||
@@ -419,8 +492,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
await self._client.aclose()
|
await self._client.aclose()
|
||||||
self._client = None
|
self._client = None
|
||||||
self._save_state()
|
self._save_state()
|
||||||
logger.info("WeChat channel stopped")
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Polling (matches monitor.ts monitorWeixinProvider)
|
# Polling (matches monitor.ts monitorWeixinProvider)
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@@ -446,10 +517,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
async def _poll_once(self) -> None:
|
async def _poll_once(self) -> None:
|
||||||
remaining = self._session_pause_remaining_s()
|
remaining = self._session_pause_remaining_s()
|
||||||
if remaining > 0:
|
if remaining > 0:
|
||||||
logger.warning(
|
|
||||||
"WeChat session paused, waiting {} min before next poll.",
|
|
||||||
max((remaining + 59) // 60, 1),
|
|
||||||
)
|
|
||||||
await asyncio.sleep(remaining)
|
await asyncio.sleep(remaining)
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -499,8 +566,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
for msg in msgs:
|
for msg in msgs:
|
||||||
try:
|
try:
|
||||||
await self._process_message(msg)
|
await self._process_message(msg)
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.error("Error processing WeChat message: {}", e)
|
pass
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Inbound message processing (matches inbound.ts + process-message.ts)
|
# Inbound message processing (matches inbound.ts + process-message.ts)
|
||||||
@@ -536,6 +603,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
item_list: list[dict] = msg.get("item_list") or []
|
item_list: list[dict] = msg.get("item_list") or []
|
||||||
content_parts: list[str] = []
|
content_parts: list[str] = []
|
||||||
media_paths: list[str] = []
|
media_paths: list[str] = []
|
||||||
|
has_top_level_downloadable_media = False
|
||||||
|
|
||||||
for item in item_list:
|
for item in item_list:
|
||||||
item_type = item.get("type", 0)
|
item_type = item.get("type", 0)
|
||||||
@@ -572,6 +640,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
elif item_type == ITEM_IMAGE:
|
elif item_type == ITEM_IMAGE:
|
||||||
image_item = item.get("image_item") or {}
|
image_item = item.get("image_item") or {}
|
||||||
|
if _has_downloadable_media_locator(image_item.get("media")):
|
||||||
|
has_top_level_downloadable_media = True
|
||||||
file_path = await self._download_media_item(image_item, "image")
|
file_path = await self._download_media_item(image_item, "image")
|
||||||
if file_path:
|
if file_path:
|
||||||
content_parts.append(f"[image]\n[Image: source: {file_path}]")
|
content_parts.append(f"[image]\n[Image: source: {file_path}]")
|
||||||
@@ -586,6 +656,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
if voice_text:
|
if voice_text:
|
||||||
content_parts.append(f"[voice] {voice_text}")
|
content_parts.append(f"[voice] {voice_text}")
|
||||||
else:
|
else:
|
||||||
|
if _has_downloadable_media_locator(voice_item.get("media")):
|
||||||
|
has_top_level_downloadable_media = True
|
||||||
file_path = await self._download_media_item(voice_item, "voice")
|
file_path = await self._download_media_item(voice_item, "voice")
|
||||||
if file_path:
|
if file_path:
|
||||||
transcription = await self.transcribe_audio(file_path)
|
transcription = await self.transcribe_audio(file_path)
|
||||||
@@ -599,6 +671,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
elif item_type == ITEM_FILE:
|
elif item_type == ITEM_FILE:
|
||||||
file_item = item.get("file_item") or {}
|
file_item = item.get("file_item") or {}
|
||||||
|
if _has_downloadable_media_locator(file_item.get("media")):
|
||||||
|
has_top_level_downloadable_media = True
|
||||||
file_name = file_item.get("file_name", "unknown")
|
file_name = file_item.get("file_name", "unknown")
|
||||||
file_path = await self._download_media_item(
|
file_path = await self._download_media_item(
|
||||||
file_item,
|
file_item,
|
||||||
@@ -613,6 +687,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
elif item_type == ITEM_VIDEO:
|
elif item_type == ITEM_VIDEO:
|
||||||
video_item = item.get("video_item") or {}
|
video_item = item.get("video_item") or {}
|
||||||
|
if _has_downloadable_media_locator(video_item.get("media")):
|
||||||
|
has_top_level_downloadable_media = True
|
||||||
file_path = await self._download_media_item(video_item, "video")
|
file_path = await self._download_media_item(video_item, "video")
|
||||||
if file_path:
|
if file_path:
|
||||||
content_parts.append(f"[video]\n[Video: source: {file_path}]")
|
content_parts.append(f"[video]\n[Video: source: {file_path}]")
|
||||||
@@ -620,17 +696,56 @@ class WeixinChannel(BaseChannel):
|
|||||||
else:
|
else:
|
||||||
content_parts.append("[video]")
|
content_parts.append("[video]")
|
||||||
|
|
||||||
|
# Fallback: when no top-level media was downloaded, try quoted/referenced media.
|
||||||
|
# This aligns with the reference plugin behavior that checks ref_msg.message_item
|
||||||
|
# when main item_list has no downloadable media.
|
||||||
|
if not media_paths and not has_top_level_downloadable_media:
|
||||||
|
ref_media_item: dict[str, Any] | None = None
|
||||||
|
for item in item_list:
|
||||||
|
if item.get("type", 0) != ITEM_TEXT:
|
||||||
|
continue
|
||||||
|
ref = item.get("ref_msg") or {}
|
||||||
|
candidate = ref.get("message_item") or {}
|
||||||
|
if candidate.get("type", 0) in (ITEM_IMAGE, ITEM_VOICE, ITEM_FILE, ITEM_VIDEO):
|
||||||
|
ref_media_item = candidate
|
||||||
|
break
|
||||||
|
|
||||||
|
if ref_media_item:
|
||||||
|
ref_type = ref_media_item.get("type", 0)
|
||||||
|
if ref_type == ITEM_IMAGE:
|
||||||
|
image_item = ref_media_item.get("image_item") or {}
|
||||||
|
file_path = await self._download_media_item(image_item, "image")
|
||||||
|
if file_path:
|
||||||
|
content_parts.append(f"[image]\n[Image: source: {file_path}]")
|
||||||
|
media_paths.append(file_path)
|
||||||
|
elif ref_type == ITEM_VOICE:
|
||||||
|
voice_item = ref_media_item.get("voice_item") or {}
|
||||||
|
file_path = await self._download_media_item(voice_item, "voice")
|
||||||
|
if file_path:
|
||||||
|
transcription = await self.transcribe_audio(file_path)
|
||||||
|
if transcription:
|
||||||
|
content_parts.append(f"[voice] {transcription}")
|
||||||
|
else:
|
||||||
|
content_parts.append(f"[voice]\n[Audio: source: {file_path}]")
|
||||||
|
media_paths.append(file_path)
|
||||||
|
elif ref_type == ITEM_FILE:
|
||||||
|
file_item = ref_media_item.get("file_item") or {}
|
||||||
|
file_name = file_item.get("file_name", "unknown")
|
||||||
|
file_path = await self._download_media_item(file_item, "file", file_name)
|
||||||
|
if file_path:
|
||||||
|
content_parts.append(f"[file: {file_name}]\n[File: source: {file_path}]")
|
||||||
|
media_paths.append(file_path)
|
||||||
|
elif ref_type == ITEM_VIDEO:
|
||||||
|
video_item = ref_media_item.get("video_item") or {}
|
||||||
|
file_path = await self._download_media_item(video_item, "video")
|
||||||
|
if file_path:
|
||||||
|
content_parts.append(f"[video]\n[Video: source: {file_path}]")
|
||||||
|
media_paths.append(file_path)
|
||||||
|
|
||||||
content = "\n".join(content_parts)
|
content = "\n".join(content_parts)
|
||||||
if not content:
|
if not content:
|
||||||
return
|
return
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"WeChat inbound: from={} items={} bodyLen={}",
|
|
||||||
from_user_id,
|
|
||||||
",".join(str(i.get("type", 0)) for i in item_list),
|
|
||||||
len(content),
|
|
||||||
)
|
|
||||||
|
|
||||||
await self._handle_message(
|
await self._handle_message(
|
||||||
sender_id=from_user_id,
|
sender_id=from_user_id,
|
||||||
chat_id=from_user_id,
|
chat_id=from_user_id,
|
||||||
@@ -652,9 +767,10 @@ class WeixinChannel(BaseChannel):
|
|||||||
"""Download + AES-decrypt a media item. Returns local path or None."""
|
"""Download + AES-decrypt a media item. Returns local path or None."""
|
||||||
try:
|
try:
|
||||||
media = typed_item.get("media") or {}
|
media = typed_item.get("media") or {}
|
||||||
encrypt_query_param = media.get("encrypt_query_param", "")
|
encrypt_query_param = str(media.get("encrypt_query_param", "") or "")
|
||||||
|
full_url = str(media.get("full_url", "") or "").strip()
|
||||||
|
|
||||||
if not encrypt_query_param:
|
if not encrypt_query_param and not full_url:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Resolve AES key (media-download.ts:43-45, pic-decrypt.ts:40-52)
|
# Resolve AES key (media-download.ts:43-45, pic-decrypt.ts:40-52)
|
||||||
@@ -671,21 +787,50 @@ class WeixinChannel(BaseChannel):
|
|||||||
elif media_aes_key_b64:
|
elif media_aes_key_b64:
|
||||||
aes_key_b64 = media_aes_key_b64
|
aes_key_b64 = media_aes_key_b64
|
||||||
|
|
||||||
# Build CDN download URL with proper URL-encoding (cdn-url.ts:7)
|
# Reference protocol behavior: VOICE/FILE/VIDEO require aes_key;
|
||||||
cdn_url = (
|
# only IMAGE may be downloaded as plain bytes when key is missing.
|
||||||
f"{self.config.cdn_base_url}/download"
|
if media_type != "image" and not aes_key_b64:
|
||||||
f"?encrypted_query_param={quote(encrypt_query_param)}"
|
return None
|
||||||
)
|
|
||||||
|
|
||||||
assert self._client is not None
|
assert self._client is not None
|
||||||
resp = await self._client.get(cdn_url)
|
fallback_url = ""
|
||||||
resp.raise_for_status()
|
if encrypt_query_param:
|
||||||
data = resp.content
|
fallback_url = (
|
||||||
|
f"{self.config.cdn_base_url}/download"
|
||||||
|
f"?encrypted_query_param={quote(encrypt_query_param)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
download_candidates: list[tuple[str, str]] = []
|
||||||
|
if full_url:
|
||||||
|
download_candidates.append(("full_url", full_url))
|
||||||
|
if fallback_url and (not full_url or fallback_url != full_url):
|
||||||
|
download_candidates.append(("encrypt_query_param", fallback_url))
|
||||||
|
|
||||||
|
data = b""
|
||||||
|
for idx, (download_source, cdn_url) in enumerate(download_candidates):
|
||||||
|
try:
|
||||||
|
resp = await self._client.get(cdn_url)
|
||||||
|
resp.raise_for_status()
|
||||||
|
data = resp.content
|
||||||
|
break
|
||||||
|
except Exception as e:
|
||||||
|
has_more_candidates = idx + 1 < len(download_candidates)
|
||||||
|
should_fallback = (
|
||||||
|
download_source == "full_url"
|
||||||
|
and has_more_candidates
|
||||||
|
and self._is_retryable_media_download_error(e)
|
||||||
|
)
|
||||||
|
if should_fallback:
|
||||||
|
logger.warning(
|
||||||
|
"WeChat media download failed via full_url, falling back to encrypt_query_param: type={} err={}",
|
||||||
|
media_type,
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
raise
|
||||||
|
|
||||||
if aes_key_b64 and data:
|
if aes_key_b64 and data:
|
||||||
data = _decrypt_aes_ecb(data, aes_key_b64)
|
data = _decrypt_aes_ecb(data, aes_key_b64)
|
||||||
elif not aes_key_b64:
|
|
||||||
logger.debug("No AES key for {} item, using raw bytes", media_type)
|
|
||||||
|
|
||||||
if not data:
|
if not data:
|
||||||
return None
|
return None
|
||||||
@@ -694,12 +839,12 @@ class WeixinChannel(BaseChannel):
|
|||||||
ext = _ext_for_type(media_type)
|
ext = _ext_for_type(media_type)
|
||||||
if not filename:
|
if not filename:
|
||||||
ts = int(time.time())
|
ts = int(time.time())
|
||||||
h = abs(hash(encrypt_query_param)) % 100000
|
hash_seed = encrypt_query_param or full_url
|
||||||
|
h = abs(hash(hash_seed)) % 100000
|
||||||
filename = f"{media_type}_{ts}_{h}{ext}"
|
filename = f"{media_type}_{ts}_{h}{ext}"
|
||||||
safe_name = os.path.basename(filename)
|
safe_name = os.path.basename(filename)
|
||||||
file_path = media_dir / safe_name
|
file_path = media_dir / safe_name
|
||||||
file_path.write_bytes(data)
|
file_path.write_bytes(data)
|
||||||
logger.debug("Downloaded WeChat {} to {}", media_type, file_path)
|
|
||||||
return str(file_path)
|
return str(file_path)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -710,14 +855,76 @@ class WeixinChannel(BaseChannel):
|
|||||||
# Outbound (matches send.ts buildTextMessageReq + sendMessageWeixin)
|
# Outbound (matches send.ts buildTextMessageReq + sendMessageWeixin)
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _get_typing_ticket(self, user_id: str, context_token: str = "") -> str:
|
||||||
|
"""Get typing ticket with per-user refresh + failure backoff cache."""
|
||||||
|
now = time.time()
|
||||||
|
entry = self._typing_tickets.get(user_id)
|
||||||
|
if entry and now < float(entry.get("next_fetch_at", 0)):
|
||||||
|
return str(entry.get("ticket", "") or "")
|
||||||
|
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"ilink_user_id": user_id,
|
||||||
|
"context_token": context_token or None,
|
||||||
|
"base_info": BASE_INFO,
|
||||||
|
}
|
||||||
|
data = await self._api_post("ilink/bot/getconfig", body)
|
||||||
|
if data.get("ret", 0) == 0:
|
||||||
|
ticket = str(data.get("typing_ticket", "") or "")
|
||||||
|
self._typing_tickets[user_id] = {
|
||||||
|
"ticket": ticket,
|
||||||
|
"ever_succeeded": True,
|
||||||
|
"next_fetch_at": now + (random.random() * TYPING_TICKET_TTL_S),
|
||||||
|
"retry_delay_s": CONFIG_CACHE_INITIAL_RETRY_S,
|
||||||
|
}
|
||||||
|
return ticket
|
||||||
|
|
||||||
|
prev_delay = float(entry.get("retry_delay_s", CONFIG_CACHE_INITIAL_RETRY_S)) if entry else CONFIG_CACHE_INITIAL_RETRY_S
|
||||||
|
next_delay = min(prev_delay * 2, CONFIG_CACHE_MAX_RETRY_S)
|
||||||
|
if entry:
|
||||||
|
entry["next_fetch_at"] = now + next_delay
|
||||||
|
entry["retry_delay_s"] = next_delay
|
||||||
|
return str(entry.get("ticket", "") or "")
|
||||||
|
|
||||||
|
self._typing_tickets[user_id] = {
|
||||||
|
"ticket": "",
|
||||||
|
"ever_succeeded": False,
|
||||||
|
"next_fetch_at": now + CONFIG_CACHE_INITIAL_RETRY_S,
|
||||||
|
"retry_delay_s": CONFIG_CACHE_INITIAL_RETRY_S,
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
|
||||||
|
async def _send_typing(self, user_id: str, typing_ticket: str, status: int) -> None:
|
||||||
|
"""Best-effort sendtyping wrapper."""
|
||||||
|
if not typing_ticket:
|
||||||
|
return
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"ilink_user_id": user_id,
|
||||||
|
"typing_ticket": typing_ticket,
|
||||||
|
"status": status,
|
||||||
|
"base_info": BASE_INFO,
|
||||||
|
}
|
||||||
|
await self._api_post("ilink/bot/sendtyping", body)
|
||||||
|
|
||||||
|
async def _typing_keepalive_loop(self, user_id: str, typing_ticket: str, stop_event: asyncio.Event) -> None:
|
||||||
|
try:
|
||||||
|
while not stop_event.is_set():
|
||||||
|
await asyncio.sleep(TYPING_KEEPALIVE_INTERVAL_S)
|
||||||
|
if stop_event.is_set():
|
||||||
|
break
|
||||||
|
try:
|
||||||
|
await self._send_typing(user_id, typing_ticket, TYPING_STATUS_TYPING)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
pass
|
||||||
|
|
||||||
async def send(self, msg: OutboundMessage) -> None:
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
if not self._client or not self._token:
|
if not self._client or not self._token:
|
||||||
logger.warning("WeChat client not initialized or not authenticated")
|
logger.warning("WeChat client not initialized or not authenticated")
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
self._assert_session_active()
|
self._assert_session_active()
|
||||||
except RuntimeError as e:
|
except RuntimeError:
|
||||||
logger.warning("WeChat send blocked: {}", e)
|
|
||||||
return
|
return
|
||||||
|
|
||||||
content = msg.content.strip()
|
content = msg.content.strip()
|
||||||
@@ -729,28 +936,62 @@ class WeixinChannel(BaseChannel):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
# --- Send media files first (following Telegram channel pattern) ---
|
typing_ticket = ""
|
||||||
for media_path in (msg.media or []):
|
try:
|
||||||
try:
|
typing_ticket = await self._get_typing_ticket(msg.chat_id, ctx_token)
|
||||||
await self._send_media_file(msg.chat_id, media_path, ctx_token)
|
except Exception:
|
||||||
except Exception as e:
|
typing_ticket = ""
|
||||||
filename = Path(media_path).name
|
|
||||||
logger.error("Failed to send WeChat media {}: {}", media_path, e)
|
|
||||||
# Notify user about failure via text
|
|
||||||
await self._send_text(
|
|
||||||
msg.chat_id, f"[Failed to send: {filename}]", ctx_token,
|
|
||||||
)
|
|
||||||
|
|
||||||
# --- Send text content ---
|
if typing_ticket:
|
||||||
if not content:
|
try:
|
||||||
return
|
await self._send_typing(msg.chat_id, typing_ticket, TYPING_STATUS_TYPING)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
typing_keepalive_stop = asyncio.Event()
|
||||||
|
typing_keepalive_task: asyncio.Task | None = None
|
||||||
|
if typing_ticket:
|
||||||
|
typing_keepalive_task = asyncio.create_task(
|
||||||
|
self._typing_keepalive_loop(msg.chat_id, typing_ticket, typing_keepalive_stop)
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
# --- Send media files first (following Telegram channel pattern) ---
|
||||||
|
for media_path in (msg.media or []):
|
||||||
|
try:
|
||||||
|
await self._send_media_file(msg.chat_id, media_path, ctx_token)
|
||||||
|
except Exception as e:
|
||||||
|
filename = Path(media_path).name
|
||||||
|
logger.error("Failed to send WeChat media {}: {}", media_path, e)
|
||||||
|
# Notify user about failure via text
|
||||||
|
await self._send_text(
|
||||||
|
msg.chat_id, f"[Failed to send: {filename}]", ctx_token,
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- Send text content ---
|
||||||
|
if not content:
|
||||||
|
return
|
||||||
|
|
||||||
chunks = split_message(content, WEIXIN_MAX_MESSAGE_LEN)
|
chunks = split_message(content, WEIXIN_MAX_MESSAGE_LEN)
|
||||||
for chunk in chunks:
|
for chunk in chunks:
|
||||||
await self._send_text(msg.chat_id, chunk, ctx_token)
|
await self._send_text(msg.chat_id, chunk, ctx_token)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error sending WeChat message: {}", e)
|
logger.error("Error sending WeChat message: {}", e)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
if typing_keepalive_task:
|
||||||
|
typing_keepalive_stop.set()
|
||||||
|
typing_keepalive_task.cancel()
|
||||||
|
try:
|
||||||
|
await typing_keepalive_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if typing_ticket:
|
||||||
|
try:
|
||||||
|
await self._send_typing(msg.chat_id, typing_ticket, TYPING_STATUS_CANCEL)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
async def _send_text(
|
async def _send_text(
|
||||||
self,
|
self,
|
||||||
@@ -824,6 +1065,10 @@ class WeixinChannel(BaseChannel):
|
|||||||
upload_type = UPLOAD_MEDIA_VIDEO
|
upload_type = UPLOAD_MEDIA_VIDEO
|
||||||
item_type = ITEM_VIDEO
|
item_type = ITEM_VIDEO
|
||||||
item_key = "video_item"
|
item_key = "video_item"
|
||||||
|
elif ext in _VOICE_EXTS:
|
||||||
|
upload_type = UPLOAD_MEDIA_VOICE
|
||||||
|
item_type = ITEM_VOICE
|
||||||
|
item_key = "voice_item"
|
||||||
else:
|
else:
|
||||||
upload_type = UPLOAD_MEDIA_FILE
|
upload_type = UPLOAD_MEDIA_FILE
|
||||||
item_type = ITEM_FILE
|
item_type = ITEM_FILE
|
||||||
@@ -837,7 +1082,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
# Matches aesEcbPaddedSize: Math.ceil((size + 1) / 16) * 16
|
# Matches aesEcbPaddedSize: Math.ceil((size + 1) / 16) * 16
|
||||||
padded_size = ((raw_size + 1 + 15) // 16) * 16
|
padded_size = ((raw_size + 1 + 15) // 16) * 16
|
||||||
|
|
||||||
# Step 1: Get upload URL (upload_param) from server
|
# Step 1: Get upload URL from server (prefer upload_full_url, fallback to upload_param)
|
||||||
file_key = os.urandom(16).hex()
|
file_key = os.urandom(16).hex()
|
||||||
upload_body: dict[str, Any] = {
|
upload_body: dict[str, Any] = {
|
||||||
"filekey": file_key,
|
"filekey": file_key,
|
||||||
@@ -852,22 +1097,27 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
assert self._client is not None
|
assert self._client is not None
|
||||||
upload_resp = await self._api_post("ilink/bot/getuploadurl", upload_body)
|
upload_resp = await self._api_post("ilink/bot/getuploadurl", upload_body)
|
||||||
logger.debug("WeChat getuploadurl response: {}", upload_resp)
|
|
||||||
|
|
||||||
upload_param = upload_resp.get("upload_param", "")
|
upload_full_url = str(upload_resp.get("upload_full_url", "") or "").strip()
|
||||||
if not upload_param:
|
upload_param = str(upload_resp.get("upload_param", "") or "")
|
||||||
raise RuntimeError(f"getuploadurl returned no upload_param: {upload_resp}")
|
if not upload_full_url and not upload_param:
|
||||||
|
raise RuntimeError(
|
||||||
|
"getuploadurl returned no upload URL "
|
||||||
|
f"(need upload_full_url or upload_param): {upload_resp}"
|
||||||
|
)
|
||||||
|
|
||||||
# Step 2: AES-128-ECB encrypt and POST to CDN
|
# Step 2: AES-128-ECB encrypt and POST to CDN
|
||||||
aes_key_b64 = base64.b64encode(aes_key_raw).decode()
|
aes_key_b64 = base64.b64encode(aes_key_raw).decode()
|
||||||
encrypted_data = _encrypt_aes_ecb(raw_data, aes_key_b64)
|
encrypted_data = _encrypt_aes_ecb(raw_data, aes_key_b64)
|
||||||
|
|
||||||
cdn_upload_url = (
|
if upload_full_url:
|
||||||
f"{self.config.cdn_base_url}/upload"
|
cdn_upload_url = upload_full_url
|
||||||
f"?encrypted_query_param={quote(upload_param)}"
|
else:
|
||||||
f"&filekey={quote(file_key)}"
|
cdn_upload_url = (
|
||||||
)
|
f"{self.config.cdn_base_url}/upload"
|
||||||
logger.debug("WeChat CDN POST url={} ciphertextSize={}", cdn_upload_url[:80], len(encrypted_data))
|
f"?encrypted_query_param={quote(upload_param)}"
|
||||||
|
f"&filekey={quote(file_key)}"
|
||||||
|
)
|
||||||
|
|
||||||
cdn_resp = await self._client.post(
|
cdn_resp = await self._client.post(
|
||||||
cdn_upload_url,
|
cdn_upload_url,
|
||||||
@@ -883,7 +1133,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
"CDN upload response missing x-encrypted-param header; "
|
"CDN upload response missing x-encrypted-param header; "
|
||||||
f"status={cdn_resp.status_code} headers={dict(cdn_resp.headers)}"
|
f"status={cdn_resp.status_code} headers={dict(cdn_resp.headers)}"
|
||||||
)
|
)
|
||||||
logger.debug("WeChat CDN upload success for {}, got download_param", p.name)
|
|
||||||
|
|
||||||
# Step 3: Send message with the media item
|
# Step 3: Send message with the media item
|
||||||
# aes_key for CDNMedia is the hex key encoded as base64
|
# aes_key for CDNMedia is the hex key encoded as base64
|
||||||
@@ -932,7 +1181,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"WeChat send media error (code {errcode}): {data.get('errmsg', '')}"
|
f"WeChat send media error (code {errcode}): {data.get('errmsg', '')}"
|
||||||
)
|
)
|
||||||
logger.info("WeChat media sent: {} (type={})", p.name, item_key)
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -1004,23 +1252,42 @@ def _decrypt_aes_ecb(data: bytes, aes_key_b64: str) -> bytes:
|
|||||||
logger.warning("Failed to parse AES key, returning raw data: {}", e)
|
logger.warning("Failed to parse AES key, returning raw data: {}", e)
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
decrypted: bytes | None = None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from Crypto.Cipher import AES
|
from Crypto.Cipher import AES
|
||||||
|
|
||||||
cipher = AES.new(key, AES.MODE_ECB)
|
cipher = AES.new(key, AES.MODE_ECB)
|
||||||
return cipher.decrypt(data) # pycryptodome auto-strips PKCS7 with unpad
|
decrypted = cipher.decrypt(data)
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
try:
|
if decrypted is None:
|
||||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
try:
|
||||||
|
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||||
|
|
||||||
cipher_obj = Cipher(algorithms.AES(key), modes.ECB())
|
cipher_obj = Cipher(algorithms.AES(key), modes.ECB())
|
||||||
decryptor = cipher_obj.decryptor()
|
decryptor = cipher_obj.decryptor()
|
||||||
return decryptor.update(data) + decryptor.finalize()
|
decrypted = decryptor.update(data) + decryptor.finalize()
|
||||||
except ImportError:
|
except ImportError:
|
||||||
logger.warning("Cannot decrypt media: install 'pycryptodome' or 'cryptography'")
|
logger.warning("Cannot decrypt media: install 'pycryptodome' or 'cryptography'")
|
||||||
|
return data
|
||||||
|
|
||||||
|
return _pkcs7_unpad_safe(decrypted)
|
||||||
|
|
||||||
|
|
||||||
|
def _pkcs7_unpad_safe(data: bytes, block_size: int = 16) -> bytes:
|
||||||
|
"""Safely remove PKCS7 padding when valid; otherwise return original bytes."""
|
||||||
|
if not data:
|
||||||
return data
|
return data
|
||||||
|
if len(data) % block_size != 0:
|
||||||
|
return data
|
||||||
|
pad_len = data[-1]
|
||||||
|
if pad_len < 1 or pad_len > block_size:
|
||||||
|
return data
|
||||||
|
if data[-pad_len:] != bytes([pad_len]) * pad_len:
|
||||||
|
return data
|
||||||
|
return data[:-pad_len]
|
||||||
|
|
||||||
|
|
||||||
def _ext_for_type(media_type: str) -> str:
|
def _ext_for_type(media_type: str) -> str:
|
||||||
|
|||||||
@@ -146,6 +146,7 @@ class WhatsAppChannel(BaseChannel):
|
|||||||
await self._ws.send(json.dumps(payload, ensure_ascii=False))
|
await self._ws.send(json.dumps(payload, ensure_ascii=False))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error sending WhatsApp message: {}", e)
|
logger.error("Error sending WhatsApp message: {}", e)
|
||||||
|
raise
|
||||||
|
|
||||||
for media_path in msg.media or []:
|
for media_path in msg.media or []:
|
||||||
try:
|
try:
|
||||||
@@ -160,6 +161,7 @@ class WhatsAppChannel(BaseChannel):
|
|||||||
await self._ws.send(json.dumps(payload, ensure_ascii=False))
|
await self._ws.send(json.dumps(payload, ensure_ascii=False))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error sending WhatsApp media {}: {}", media_path, e)
|
logger.error("Error sending WhatsApp media {}: {}", media_path, e)
|
||||||
|
raise
|
||||||
|
|
||||||
async def _handle_bridge_message(self, raw: str) -> None:
|
async def _handle_bridge_message(self, raw: str) -> None:
|
||||||
"""Handle a message from the bridge."""
|
"""Handle a message from the bridge."""
|
||||||
|
|||||||
@@ -491,6 +491,91 @@ def _migrate_cron_store(config: "Config") -> None:
|
|||||||
shutil.move(str(legacy_path), str(new_path))
|
shutil.move(str(legacy_path), str(new_path))
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# OpenAI-Compatible API Server
|
||||||
|
# ============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@app.command()
|
||||||
|
def serve(
|
||||||
|
port: int | None = typer.Option(None, "--port", "-p", help="API server port"),
|
||||||
|
host: str | None = typer.Option(None, "--host", "-H", help="Bind address"),
|
||||||
|
timeout: float | None = typer.Option(None, "--timeout", "-t", help="Per-request timeout (seconds)"),
|
||||||
|
verbose: bool = typer.Option(False, "--verbose", "-v", help="Show nanobot runtime logs"),
|
||||||
|
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
|
||||||
|
config: str | None = typer.Option(None, "--config", "-c", help="Path to config file"),
|
||||||
|
):
|
||||||
|
"""Start the OpenAI-compatible API server (/v1/chat/completions)."""
|
||||||
|
try:
|
||||||
|
from aiohttp import web # noqa: F401
|
||||||
|
except ImportError:
|
||||||
|
console.print("[red]aiohttp is required. Install with: pip install 'nanobot-ai[api]'[/red]")
|
||||||
|
raise typer.Exit(1)
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.api.server import create_app
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.session.manager import SessionManager
|
||||||
|
|
||||||
|
if verbose:
|
||||||
|
logger.enable("nanobot")
|
||||||
|
else:
|
||||||
|
logger.disable("nanobot")
|
||||||
|
|
||||||
|
runtime_config = _load_runtime_config(config, workspace)
|
||||||
|
api_cfg = runtime_config.api
|
||||||
|
host = host if host is not None else api_cfg.host
|
||||||
|
port = port if port is not None else api_cfg.port
|
||||||
|
timeout = timeout if timeout is not None else api_cfg.timeout
|
||||||
|
sync_workspace_templates(runtime_config.workspace_path)
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = _make_provider(runtime_config)
|
||||||
|
session_manager = SessionManager(runtime_config.workspace_path)
|
||||||
|
agent_loop = AgentLoop(
|
||||||
|
bus=bus,
|
||||||
|
provider=provider,
|
||||||
|
workspace=runtime_config.workspace_path,
|
||||||
|
model=runtime_config.agents.defaults.model,
|
||||||
|
max_iterations=runtime_config.agents.defaults.max_tool_iterations,
|
||||||
|
context_window_tokens=runtime_config.agents.defaults.context_window_tokens,
|
||||||
|
web_search_config=runtime_config.tools.web.search,
|
||||||
|
web_proxy=runtime_config.tools.web.proxy or None,
|
||||||
|
exec_config=runtime_config.tools.exec,
|
||||||
|
restrict_to_workspace=runtime_config.tools.restrict_to_workspace,
|
||||||
|
session_manager=session_manager,
|
||||||
|
mcp_servers=runtime_config.tools.mcp_servers,
|
||||||
|
channels_config=runtime_config.channels,
|
||||||
|
timezone=runtime_config.agents.defaults.timezone,
|
||||||
|
)
|
||||||
|
|
||||||
|
model_name = runtime_config.agents.defaults.model
|
||||||
|
console.print(f"{__logo__} Starting OpenAI-compatible API server")
|
||||||
|
console.print(f" [cyan]Endpoint[/cyan] : http://{host}:{port}/v1/chat/completions")
|
||||||
|
console.print(f" [cyan]Model[/cyan] : {model_name}")
|
||||||
|
console.print(" [cyan]Session[/cyan] : api:default")
|
||||||
|
console.print(f" [cyan]Timeout[/cyan] : {timeout}s")
|
||||||
|
if host in {"0.0.0.0", "::"}:
|
||||||
|
console.print(
|
||||||
|
"[yellow]Warning:[/yellow] API is bound to all interfaces. "
|
||||||
|
"Only do this behind a trusted network boundary, firewall, or reverse proxy."
|
||||||
|
)
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
api_app = create_app(agent_loop, model_name=model_name, request_timeout=timeout)
|
||||||
|
|
||||||
|
async def on_startup(_app):
|
||||||
|
await agent_loop._connect_mcp()
|
||||||
|
|
||||||
|
async def on_cleanup(_app):
|
||||||
|
await agent_loop.close_mcp()
|
||||||
|
|
||||||
|
api_app.on_startup.append(on_startup)
|
||||||
|
api_app.on_cleanup.append(on_cleanup)
|
||||||
|
|
||||||
|
web.run_app(api_app, host=host, port=port, print=lambda msg: logger.info(msg))
|
||||||
|
|
||||||
|
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
# Gateway / Server
|
# Gateway / Server
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@@ -549,6 +634,7 @@ def gateway(
|
|||||||
session_manager=session_manager,
|
session_manager=session_manager,
|
||||||
mcp_servers=config.tools.mcp_servers,
|
mcp_servers=config.tools.mcp_servers,
|
||||||
channels_config=config.channels,
|
channels_config=config.channels,
|
||||||
|
timezone=config.agents.defaults.timezone,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Set cron callback (needs agent)
|
# Set cron callback (needs agent)
|
||||||
@@ -659,6 +745,7 @@ def gateway(
|
|||||||
on_notify=on_heartbeat_notify,
|
on_notify=on_heartbeat_notify,
|
||||||
interval_s=hb_cfg.interval_s,
|
interval_s=hb_cfg.interval_s,
|
||||||
enabled=hb_cfg.enabled,
|
enabled=hb_cfg.enabled,
|
||||||
|
timezone=config.agents.defaults.timezone,
|
||||||
)
|
)
|
||||||
|
|
||||||
if channels.enabled_channels:
|
if channels.enabled_channels:
|
||||||
@@ -752,6 +839,7 @@ def agent(
|
|||||||
restrict_to_workspace=config.tools.restrict_to_workspace,
|
restrict_to_workspace=config.tools.restrict_to_workspace,
|
||||||
mcp_servers=config.tools.mcp_servers,
|
mcp_servers=config.tools.mcp_servers,
|
||||||
channels_config=config.channels,
|
channels_config=config.channels,
|
||||||
|
timezone=config.agents.defaults.timezone,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Shared reference for progress callbacks
|
# Shared reference for progress callbacks
|
||||||
|
|||||||
@@ -84,6 +84,16 @@ async def cmd_new(ctx: CommandContext) -> OutboundMessage:
|
|||||||
|
|
||||||
async def cmd_help(ctx: CommandContext) -> OutboundMessage:
|
async def cmd_help(ctx: CommandContext) -> OutboundMessage:
|
||||||
"""Return available slash commands."""
|
"""Return available slash commands."""
|
||||||
|
return OutboundMessage(
|
||||||
|
channel=ctx.msg.channel,
|
||||||
|
chat_id=ctx.msg.chat_id,
|
||||||
|
content=build_help_text(),
|
||||||
|
metadata={"render_as": "text"},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_help_text() -> str:
|
||||||
|
"""Build canonical help text shared across channels."""
|
||||||
lines = [
|
lines = [
|
||||||
"🐈 nanobot commands:",
|
"🐈 nanobot commands:",
|
||||||
"/new — Start a new conversation",
|
"/new — Start a new conversation",
|
||||||
@@ -92,12 +102,7 @@ async def cmd_help(ctx: CommandContext) -> OutboundMessage:
|
|||||||
"/status — Show bot status",
|
"/status — Show bot status",
|
||||||
"/help — Show available commands",
|
"/help — Show available commands",
|
||||||
]
|
]
|
||||||
return OutboundMessage(
|
return "\n".join(lines)
|
||||||
channel=ctx.msg.channel,
|
|
||||||
chat_id=ctx.msg.chat_id,
|
|
||||||
content="\n".join(lines),
|
|
||||||
metadata={"render_as": "text"},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def register_builtin_commands(router: CommandRouter) -> None:
|
def register_builtin_commands(router: CommandRouter) -> None:
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ class ChannelsConfig(Base):
|
|||||||
|
|
||||||
send_progress: bool = True # stream agent's text progress to the channel
|
send_progress: bool = True # stream agent's text progress to the channel
|
||||||
send_tool_hints: bool = False # stream tool-call hints (e.g. read_file("…"))
|
send_tool_hints: bool = False # stream tool-call hints (e.g. read_file("…"))
|
||||||
|
send_max_retries: int = Field(default=3, ge=0, le=10) # Max delivery attempts (initial send included)
|
||||||
|
|
||||||
|
|
||||||
class AgentDefaults(Base):
|
class AgentDefaults(Base):
|
||||||
@@ -40,6 +41,7 @@ class AgentDefaults(Base):
|
|||||||
temperature: float = 0.1
|
temperature: float = 0.1
|
||||||
max_tool_iterations: int = 40
|
max_tool_iterations: int = 40
|
||||||
reasoning_effort: str | None = None # low / medium / high - enables LLM thinking mode
|
reasoning_effort: str | None = None # low / medium / high - enables LLM thinking mode
|
||||||
|
timezone: str = "UTC" # IANA timezone, e.g. "Asia/Shanghai", "America/New_York"
|
||||||
|
|
||||||
|
|
||||||
class AgentsConfig(Base):
|
class AgentsConfig(Base):
|
||||||
@@ -75,6 +77,7 @@ class ProvidersConfig(Base):
|
|||||||
moonshot: ProviderConfig = Field(default_factory=ProviderConfig)
|
moonshot: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
minimax: ProviderConfig = Field(default_factory=ProviderConfig)
|
minimax: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
mistral: ProviderConfig = Field(default_factory=ProviderConfig)
|
mistral: ProviderConfig = Field(default_factory=ProviderConfig)
|
||||||
|
stepfun: ProviderConfig = Field(default_factory=ProviderConfig) # Step Fun (阶跃星辰)
|
||||||
aihubmix: ProviderConfig = Field(default_factory=ProviderConfig) # AiHubMix API gateway
|
aihubmix: ProviderConfig = Field(default_factory=ProviderConfig) # AiHubMix API gateway
|
||||||
siliconflow: ProviderConfig = Field(default_factory=ProviderConfig) # SiliconFlow (硅基流动)
|
siliconflow: ProviderConfig = Field(default_factory=ProviderConfig) # SiliconFlow (硅基流动)
|
||||||
volcengine: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine (火山引擎)
|
volcengine: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine (火山引擎)
|
||||||
@@ -93,6 +96,14 @@ class HeartbeatConfig(Base):
|
|||||||
keep_recent_messages: int = 8
|
keep_recent_messages: int = 8
|
||||||
|
|
||||||
|
|
||||||
|
class ApiConfig(Base):
|
||||||
|
"""OpenAI-compatible API server configuration."""
|
||||||
|
|
||||||
|
host: str = "127.0.0.1" # Safer default: local-only bind.
|
||||||
|
port: int = 8900
|
||||||
|
timeout: float = 120.0 # Per-request timeout in seconds.
|
||||||
|
|
||||||
|
|
||||||
class GatewayConfig(Base):
|
class GatewayConfig(Base):
|
||||||
"""Gateway/server configuration."""
|
"""Gateway/server configuration."""
|
||||||
|
|
||||||
@@ -125,6 +136,7 @@ class ExecToolConfig(Base):
|
|||||||
enable: bool = True
|
enable: bool = True
|
||||||
timeout: int = 60
|
timeout: int = 60
|
||||||
path_append: str = ""
|
path_append: str = ""
|
||||||
|
command_wrapper: str = "" # sandbox wrapper command template; supports {command} and {cwd}
|
||||||
|
|
||||||
class MCPServerConfig(Base):
|
class MCPServerConfig(Base):
|
||||||
"""MCP server connection configuration (stdio or HTTP)."""
|
"""MCP server connection configuration (stdio or HTTP)."""
|
||||||
@@ -153,6 +165,7 @@ class Config(BaseSettings):
|
|||||||
agents: AgentsConfig = Field(default_factory=AgentsConfig)
|
agents: AgentsConfig = Field(default_factory=AgentsConfig)
|
||||||
channels: ChannelsConfig = Field(default_factory=ChannelsConfig)
|
channels: ChannelsConfig = Field(default_factory=ChannelsConfig)
|
||||||
providers: ProvidersConfig = Field(default_factory=ProvidersConfig)
|
providers: ProvidersConfig = Field(default_factory=ProvidersConfig)
|
||||||
|
api: ApiConfig = Field(default_factory=ApiConfig)
|
||||||
gateway: GatewayConfig = Field(default_factory=GatewayConfig)
|
gateway: GatewayConfig = Field(default_factory=GatewayConfig)
|
||||||
tools: ToolsConfig = Field(default_factory=ToolsConfig)
|
tools: ToolsConfig = Field(default_factory=ToolsConfig)
|
||||||
|
|
||||||
|
|||||||
@@ -59,6 +59,7 @@ class HeartbeatService:
|
|||||||
on_notify: Callable[[str], Coroutine[Any, Any, None]] | None = None,
|
on_notify: Callable[[str], Coroutine[Any, Any, None]] | None = None,
|
||||||
interval_s: int = 30 * 60,
|
interval_s: int = 30 * 60,
|
||||||
enabled: bool = True,
|
enabled: bool = True,
|
||||||
|
timezone: str | None = None,
|
||||||
):
|
):
|
||||||
self.workspace = workspace
|
self.workspace = workspace
|
||||||
self.provider = provider
|
self.provider = provider
|
||||||
@@ -67,6 +68,7 @@ class HeartbeatService:
|
|||||||
self.on_notify = on_notify
|
self.on_notify = on_notify
|
||||||
self.interval_s = interval_s
|
self.interval_s = interval_s
|
||||||
self.enabled = enabled
|
self.enabled = enabled
|
||||||
|
self.timezone = timezone
|
||||||
self._running = False
|
self._running = False
|
||||||
self._task: asyncio.Task | None = None
|
self._task: asyncio.Task | None = None
|
||||||
|
|
||||||
@@ -93,7 +95,7 @@ class HeartbeatService:
|
|||||||
messages=[
|
messages=[
|
||||||
{"role": "system", "content": "You are a heartbeat agent. Call the heartbeat tool to report your decision."},
|
{"role": "system", "content": "You are a heartbeat agent. Call the heartbeat tool to report your decision."},
|
||||||
{"role": "user", "content": (
|
{"role": "user", "content": (
|
||||||
f"Current Time: {current_time_str()}\n\n"
|
f"Current Time: {current_time_str(self.timezone)}\n\n"
|
||||||
"Review the following HEARTBEAT.md and decide whether there are active tasks.\n\n"
|
"Review the following HEARTBEAT.md and decide whether there are active tasks.\n\n"
|
||||||
f"{content}"
|
f"{content}"
|
||||||
)},
|
)},
|
||||||
|
|||||||
@@ -0,0 +1,170 @@
|
|||||||
|
"""High-level programmatic interface to nanobot."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from nanobot.agent.hook import AgentHook
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class RunResult:
|
||||||
|
"""Result of a single agent run."""
|
||||||
|
|
||||||
|
content: str
|
||||||
|
tools_used: list[str]
|
||||||
|
messages: list[dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
class Nanobot:
|
||||||
|
"""Programmatic facade for running the nanobot agent.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
bot = Nanobot.from_config()
|
||||||
|
result = await bot.run("Summarize this repo", hooks=[MyHook()])
|
||||||
|
print(result.content)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, loop: AgentLoop) -> None:
|
||||||
|
self._loop = loop
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_config(
|
||||||
|
cls,
|
||||||
|
config_path: str | Path | None = None,
|
||||||
|
*,
|
||||||
|
workspace: str | Path | None = None,
|
||||||
|
) -> Nanobot:
|
||||||
|
"""Create a Nanobot instance from a config file.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config_path: Path to ``config.json``. Defaults to
|
||||||
|
``~/.nanobot/config.json``.
|
||||||
|
workspace: Override the workspace directory from config.
|
||||||
|
"""
|
||||||
|
from nanobot.config.loader import load_config
|
||||||
|
from nanobot.config.schema import Config
|
||||||
|
|
||||||
|
resolved: Path | None = None
|
||||||
|
if config_path is not None:
|
||||||
|
resolved = Path(config_path).expanduser().resolve()
|
||||||
|
if not resolved.exists():
|
||||||
|
raise FileNotFoundError(f"Config not found: {resolved}")
|
||||||
|
|
||||||
|
config: Config = load_config(resolved)
|
||||||
|
if workspace is not None:
|
||||||
|
config.agents.defaults.workspace = str(
|
||||||
|
Path(workspace).expanduser().resolve()
|
||||||
|
)
|
||||||
|
|
||||||
|
provider = _make_provider(config)
|
||||||
|
bus = MessageBus()
|
||||||
|
defaults = config.agents.defaults
|
||||||
|
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=bus,
|
||||||
|
provider=provider,
|
||||||
|
workspace=config.workspace_path,
|
||||||
|
model=defaults.model,
|
||||||
|
max_iterations=defaults.max_tool_iterations,
|
||||||
|
context_window_tokens=defaults.context_window_tokens,
|
||||||
|
web_search_config=config.tools.web.search,
|
||||||
|
web_proxy=config.tools.web.proxy or None,
|
||||||
|
exec_config=config.tools.exec,
|
||||||
|
restrict_to_workspace=config.tools.restrict_to_workspace,
|
||||||
|
mcp_servers=config.tools.mcp_servers,
|
||||||
|
timezone=defaults.timezone,
|
||||||
|
)
|
||||||
|
return cls(loop)
|
||||||
|
|
||||||
|
async def run(
|
||||||
|
self,
|
||||||
|
message: str,
|
||||||
|
*,
|
||||||
|
session_key: str = "sdk:default",
|
||||||
|
hooks: list[AgentHook] | None = None,
|
||||||
|
) -> RunResult:
|
||||||
|
"""Run the agent once and return the result.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
message: The user message to process.
|
||||||
|
session_key: Session identifier for conversation isolation.
|
||||||
|
Different keys get independent history.
|
||||||
|
hooks: Optional lifecycle hooks for this run.
|
||||||
|
"""
|
||||||
|
prev = self._loop._extra_hooks
|
||||||
|
if hooks is not None:
|
||||||
|
self._loop._extra_hooks = list(hooks)
|
||||||
|
try:
|
||||||
|
response = await self._loop.process_direct(
|
||||||
|
message, session_key=session_key,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
self._loop._extra_hooks = prev
|
||||||
|
|
||||||
|
content = (response.content if response else None) or ""
|
||||||
|
return RunResult(content=content, tools_used=[], messages=[])
|
||||||
|
|
||||||
|
|
||||||
|
def _make_provider(config: Any) -> Any:
|
||||||
|
"""Create the LLM provider from config (extracted from CLI)."""
|
||||||
|
from nanobot.providers.base import GenerationSettings
|
||||||
|
from nanobot.providers.registry import find_by_name
|
||||||
|
|
||||||
|
model = config.agents.defaults.model
|
||||||
|
provider_name = config.get_provider_name(model)
|
||||||
|
p = config.get_provider(model)
|
||||||
|
spec = find_by_name(provider_name) if provider_name else None
|
||||||
|
backend = spec.backend if spec else "openai_compat"
|
||||||
|
|
||||||
|
if backend == "azure_openai":
|
||||||
|
if not p or not p.api_key or not p.api_base:
|
||||||
|
raise ValueError("Azure OpenAI requires api_key and api_base in config.")
|
||||||
|
elif backend == "openai_compat" and not model.startswith("bedrock/"):
|
||||||
|
needs_key = not (p and p.api_key)
|
||||||
|
exempt = spec and (spec.is_oauth or spec.is_local or spec.is_direct)
|
||||||
|
if needs_key and not exempt:
|
||||||
|
raise ValueError(f"No API key configured for provider '{provider_name}'.")
|
||||||
|
|
||||||
|
if backend == "openai_codex":
|
||||||
|
from nanobot.providers.openai_codex_provider import OpenAICodexProvider
|
||||||
|
|
||||||
|
provider = OpenAICodexProvider(default_model=model)
|
||||||
|
elif backend == "azure_openai":
|
||||||
|
from nanobot.providers.azure_openai_provider import AzureOpenAIProvider
|
||||||
|
|
||||||
|
provider = AzureOpenAIProvider(
|
||||||
|
api_key=p.api_key, api_base=p.api_base, default_model=model
|
||||||
|
)
|
||||||
|
elif backend == "anthropic":
|
||||||
|
from nanobot.providers.anthropic_provider import AnthropicProvider
|
||||||
|
|
||||||
|
provider = AnthropicProvider(
|
||||||
|
api_key=p.api_key if p else None,
|
||||||
|
api_base=config.get_api_base(model),
|
||||||
|
default_model=model,
|
||||||
|
extra_headers=p.extra_headers if p else None,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
|
|
||||||
|
provider = OpenAICompatProvider(
|
||||||
|
api_key=p.api_key if p else None,
|
||||||
|
api_base=config.get_api_base(model),
|
||||||
|
default_model=model,
|
||||||
|
extra_headers=p.extra_headers if p else None,
|
||||||
|
spec=spec,
|
||||||
|
)
|
||||||
|
|
||||||
|
defaults = config.agents.defaults
|
||||||
|
provider.generation = GenerationSettings(
|
||||||
|
temperature=defaults.temperature,
|
||||||
|
max_tokens=defaults.max_tokens,
|
||||||
|
reasoning_effort=defaults.reasoning_effort,
|
||||||
|
)
|
||||||
|
return provider
|
||||||
@@ -26,6 +26,11 @@ _ALNUM = string.ascii_letters + string.digits
|
|||||||
|
|
||||||
_STANDARD_TC_KEYS = frozenset({"id", "type", "index", "function"})
|
_STANDARD_TC_KEYS = frozenset({"id", "type", "index", "function"})
|
||||||
_STANDARD_FN_KEYS = frozenset({"name", "arguments"})
|
_STANDARD_FN_KEYS = frozenset({"name", "arguments"})
|
||||||
|
_DEFAULT_OPENROUTER_HEADERS = {
|
||||||
|
"HTTP-Referer": "https://github.com/HKUDS/nanobot",
|
||||||
|
"X-OpenRouter-Title": "nanobot",
|
||||||
|
"X-OpenRouter-Categories": "cli-agent,personal-agent",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def _short_tool_id() -> str:
|
def _short_tool_id() -> str:
|
||||||
@@ -89,6 +94,13 @@ def _extract_tc_extras(tc: Any) -> tuple[
|
|||||||
return extra_content, prov, fn_prov
|
return extra_content, prov, fn_prov
|
||||||
|
|
||||||
|
|
||||||
|
def _uses_openrouter_attribution(spec: "ProviderSpec | None", api_base: str | None) -> bool:
|
||||||
|
"""Apply Nanobot attribution headers to OpenRouter requests by default."""
|
||||||
|
if spec and spec.name == "openrouter":
|
||||||
|
return True
|
||||||
|
return bool(api_base and "openrouter" in api_base.lower())
|
||||||
|
|
||||||
|
|
||||||
class OpenAICompatProvider(LLMProvider):
|
class OpenAICompatProvider(LLMProvider):
|
||||||
"""Unified provider for all OpenAI-compatible APIs.
|
"""Unified provider for all OpenAI-compatible APIs.
|
||||||
|
|
||||||
@@ -113,14 +125,16 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
self._setup_env(api_key, api_base)
|
self._setup_env(api_key, api_base)
|
||||||
|
|
||||||
effective_base = api_base or (spec.default_api_base if spec else None) or None
|
effective_base = api_base or (spec.default_api_base if spec else None) or None
|
||||||
|
default_headers = {"x-session-affinity": uuid.uuid4().hex}
|
||||||
|
if _uses_openrouter_attribution(spec, effective_base):
|
||||||
|
default_headers.update(_DEFAULT_OPENROUTER_HEADERS)
|
||||||
|
if extra_headers:
|
||||||
|
default_headers.update(extra_headers)
|
||||||
|
|
||||||
self._client = AsyncOpenAI(
|
self._client = AsyncOpenAI(
|
||||||
api_key=api_key or "no-key",
|
api_key=api_key or "no-key",
|
||||||
base_url=effective_base,
|
base_url=effective_base,
|
||||||
default_headers={
|
default_headers=default_headers,
|
||||||
"x-session-affinity": uuid.uuid4().hex,
|
|
||||||
**(extra_headers or {}),
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _setup_env(self, api_key: str, api_base: str | None) -> None:
|
def _setup_env(self, api_key: str, api_base: str | None) -> None:
|
||||||
@@ -229,10 +243,14 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
kwargs: dict[str, Any] = {
|
kwargs: dict[str, Any] = {
|
||||||
"model": model_name,
|
"model": model_name,
|
||||||
"messages": self._sanitize_messages(self._sanitize_empty_content(messages)),
|
"messages": self._sanitize_messages(self._sanitize_empty_content(messages)),
|
||||||
"max_tokens": max(1, max_tokens),
|
|
||||||
"temperature": temperature,
|
"temperature": temperature,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if spec and getattr(spec, "supports_max_completion_tokens", False):
|
||||||
|
kwargs["max_completion_tokens"] = max(1, max_tokens)
|
||||||
|
else:
|
||||||
|
kwargs["max_tokens"] = max(1, max_tokens)
|
||||||
|
|
||||||
if spec:
|
if spec:
|
||||||
model_lower = model_name.lower()
|
model_lower = model_name.lower()
|
||||||
for pattern, overrides in spec.model_overrides:
|
for pattern, overrides in spec.model_overrides:
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ class ProviderSpec:
|
|||||||
|
|
||||||
# gateway behavior
|
# gateway behavior
|
||||||
strip_model_prefix: bool = False # strip "provider/" before sending to gateway
|
strip_model_prefix: bool = False # strip "provider/" before sending to gateway
|
||||||
|
supports_max_completion_tokens: bool = False
|
||||||
|
|
||||||
# per-model param overrides, e.g. (("kimi-k2.5", {"temperature": 1.0}),)
|
# per-model param overrides, e.g. (("kimi-k2.5", {"temperature": 1.0}),)
|
||||||
model_overrides: tuple[tuple[str, dict[str, Any]], ...] = ()
|
model_overrides: tuple[tuple[str, dict[str, Any]], ...] = ()
|
||||||
@@ -286,6 +287,15 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
|
|||||||
backend="openai_compat",
|
backend="openai_compat",
|
||||||
default_api_base="https://api.mistral.ai/v1",
|
default_api_base="https://api.mistral.ai/v1",
|
||||||
),
|
),
|
||||||
|
# Step Fun (阶跃星辰): OpenAI-compatible API
|
||||||
|
ProviderSpec(
|
||||||
|
name="stepfun",
|
||||||
|
keywords=("stepfun", "step"),
|
||||||
|
env_key="STEPFUN_API_KEY",
|
||||||
|
display_name="Step Fun",
|
||||||
|
backend="openai_compat",
|
||||||
|
default_api_base="https://api.stepfun.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(
|
||||||
|
|||||||
@@ -295,7 +295,7 @@ After initialization, customize the SKILL.md and add resources as needed. If you
|
|||||||
|
|
||||||
### Step 4: Edit the Skill
|
### Step 4: Edit the Skill
|
||||||
|
|
||||||
When editing the (newly-generated or existing) skill, remember that the skill is being created for another instance of the agent to use. Include information that would be beneficial and non-obvious to the agent. Consider what procedural knowledge, domain-specific details, or reusable assets would help another the agent instance execute these tasks more effectively.
|
When editing the (newly-generated or existing) skill, remember that the skill is being created for another instance of the agent to use. Include information that would be beneficial and non-obvious to the agent. Consider what procedural knowledge, domain-specific details, or reusable assets would help another agent instance execute these tasks more effectively.
|
||||||
|
|
||||||
#### Learn Proven Design Patterns
|
#### Learn Proven Design Patterns
|
||||||
|
|
||||||
|
|||||||
@@ -55,11 +55,24 @@ def timestamp() -> str:
|
|||||||
return datetime.now().isoformat()
|
return datetime.now().isoformat()
|
||||||
|
|
||||||
|
|
||||||
def current_time_str() -> str:
|
def current_time_str(timezone: str | None = None) -> str:
|
||||||
"""Human-readable current time with weekday and timezone, e.g. '2026-03-15 22:30 (Saturday) (CST)'."""
|
"""Human-readable current time with weekday and UTC offset.
|
||||||
now = datetime.now().strftime("%Y-%m-%d %H:%M (%A)")
|
|
||||||
tz = time.strftime("%Z") or "UTC"
|
When *timezone* is a valid IANA name (e.g. ``"Asia/Shanghai"``), the time
|
||||||
return f"{now} ({tz})"
|
is converted to that zone. Otherwise falls back to the host local time.
|
||||||
|
"""
|
||||||
|
from zoneinfo import ZoneInfo
|
||||||
|
|
||||||
|
try:
|
||||||
|
tz = ZoneInfo(timezone) if timezone else None
|
||||||
|
except (KeyError, Exception):
|
||||||
|
tz = None
|
||||||
|
|
||||||
|
now = datetime.now(tz=tz) if tz else datetime.now().astimezone()
|
||||||
|
offset = now.strftime("%z")
|
||||||
|
offset_fmt = f"{offset[:3]}:{offset[3:]}" if len(offset) == 5 else offset
|
||||||
|
tz_name = timezone or (time.strftime("%Z") or "UTC")
|
||||||
|
return f"{now.strftime('%Y-%m-%d %H:%M (%A)')} ({tz_name}, UTC{offset_fmt})"
|
||||||
|
|
||||||
|
|
||||||
_UNSAFE_CHARS = re.compile(r'[<>:"/\\|?*]')
|
_UNSAFE_CHARS = re.compile(r'[<>:"/\\|?*]')
|
||||||
@@ -111,8 +124,8 @@ def build_assistant_message(
|
|||||||
msg: dict[str, Any] = {"role": "assistant", "content": content}
|
msg: dict[str, Any] = {"role": "assistant", "content": content}
|
||||||
if tool_calls:
|
if tool_calls:
|
||||||
msg["tool_calls"] = tool_calls
|
msg["tool_calls"] = tool_calls
|
||||||
if reasoning_content is not None:
|
if reasoning_content is not None or thinking_blocks:
|
||||||
msg["reasoning_content"] = reasoning_content
|
msg["reasoning_content"] = reasoning_content if reasoning_content is not None else ""
|
||||||
if thinking_blocks:
|
if thinking_blocks:
|
||||||
msg["thinking_blocks"] = thinking_blocks
|
msg["thinking_blocks"] = thinking_blocks
|
||||||
return msg
|
return msg
|
||||||
|
|||||||
+21
-1
@@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "nanobot-ai"
|
name = "nanobot-ai"
|
||||||
version = "0.1.4.post5"
|
version = "0.1.4.post6"
|
||||||
description = "A lightweight personal AI assistant framework"
|
description = "A lightweight personal AI assistant framework"
|
||||||
readme = { file = "README.md", content-type = "text/markdown" }
|
readme = { file = "README.md", content-type = "text/markdown" }
|
||||||
requires-python = ">=3.11"
|
requires-python = ">=3.11"
|
||||||
@@ -51,6 +51,9 @@ dependencies = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
|
api = [
|
||||||
|
"aiohttp>=3.9.0,<4.0.0",
|
||||||
|
]
|
||||||
wecom = [
|
wecom = [
|
||||||
"wecom-aibot-sdk-python>=0.1.5",
|
"wecom-aibot-sdk-python>=0.1.5",
|
||||||
]
|
]
|
||||||
@@ -64,12 +67,16 @@ matrix = [
|
|||||||
"mistune>=3.0.0,<4.0.0",
|
"mistune>=3.0.0,<4.0.0",
|
||||||
"nh3>=0.2.17,<1.0.0",
|
"nh3>=0.2.17,<1.0.0",
|
||||||
]
|
]
|
||||||
|
discord = [
|
||||||
|
"discord.py>=2.5.2,<3.0.0",
|
||||||
|
]
|
||||||
langsmith = [
|
langsmith = [
|
||||||
"langsmith>=0.1.0",
|
"langsmith>=0.1.0",
|
||||||
]
|
]
|
||||||
dev = [
|
dev = [
|
||||||
"pytest>=9.0.0,<10.0.0",
|
"pytest>=9.0.0,<10.0.0",
|
||||||
"pytest-asyncio>=1.3.0,<2.0.0",
|
"pytest-asyncio>=1.3.0,<2.0.0",
|
||||||
|
"aiohttp>=3.9.0,<4.0.0",
|
||||||
"pytest-cov>=6.0.0,<7.0.0",
|
"pytest-cov>=6.0.0,<7.0.0",
|
||||||
"ruff>=0.1.0",
|
"ruff>=0.1.0",
|
||||||
]
|
]
|
||||||
@@ -120,3 +127,16 @@ ignore = ["E501"]
|
|||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
asyncio_mode = "auto"
|
asyncio_mode = "auto"
|
||||||
testpaths = ["tests"]
|
testpaths = ["tests"]
|
||||||
|
|
||||||
|
[tool.coverage.run]
|
||||||
|
source = ["nanobot"]
|
||||||
|
omit = ["tests/*", "**/tests/*"]
|
||||||
|
|
||||||
|
[tool.coverage.report]
|
||||||
|
exclude_lines = [
|
||||||
|
"pragma: no cover",
|
||||||
|
"def __repr__",
|
||||||
|
"raise NotImplementedError",
|
||||||
|
"if __name__ == .__main__.:",
|
||||||
|
"if TYPE_CHECKING:",
|
||||||
|
]
|
||||||
|
|||||||
@@ -0,0 +1,351 @@
|
|||||||
|
"""Tests for CompositeHook fan-out, error isolation, and integration."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.agent.hook import AgentHook, AgentHookContext, CompositeHook
|
||||||
|
|
||||||
|
|
||||||
|
def _ctx() -> AgentHookContext:
|
||||||
|
return AgentHookContext(iteration=0, messages=[])
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Fan-out: every hook is called in order
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_composite_fans_out_before_iteration():
|
||||||
|
calls: list[str] = []
|
||||||
|
|
||||||
|
class H(AgentHook):
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
calls.append(f"A:{context.iteration}")
|
||||||
|
|
||||||
|
class H2(AgentHook):
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
calls.append(f"B:{context.iteration}")
|
||||||
|
|
||||||
|
hook = CompositeHook([H(), H2()])
|
||||||
|
ctx = _ctx()
|
||||||
|
await hook.before_iteration(ctx)
|
||||||
|
assert calls == ["A:0", "B:0"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_composite_fans_out_all_async_methods():
|
||||||
|
"""Verify all async methods fan out to every hook."""
|
||||||
|
events: list[str] = []
|
||||||
|
|
||||||
|
class RecordingHook(AgentHook):
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
events.append("before_iteration")
|
||||||
|
|
||||||
|
async def on_stream(self, context: AgentHookContext, delta: str) -> None:
|
||||||
|
events.append(f"on_stream:{delta}")
|
||||||
|
|
||||||
|
async def on_stream_end(self, context: AgentHookContext, *, resuming: bool) -> None:
|
||||||
|
events.append(f"on_stream_end:{resuming}")
|
||||||
|
|
||||||
|
async def before_execute_tools(self, context: AgentHookContext) -> None:
|
||||||
|
events.append("before_execute_tools")
|
||||||
|
|
||||||
|
async def after_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
events.append("after_iteration")
|
||||||
|
|
||||||
|
hook = CompositeHook([RecordingHook(), RecordingHook()])
|
||||||
|
ctx = _ctx()
|
||||||
|
|
||||||
|
await hook.before_iteration(ctx)
|
||||||
|
await hook.on_stream(ctx, "hi")
|
||||||
|
await hook.on_stream_end(ctx, resuming=True)
|
||||||
|
await hook.before_execute_tools(ctx)
|
||||||
|
await hook.after_iteration(ctx)
|
||||||
|
|
||||||
|
assert events == [
|
||||||
|
"before_iteration", "before_iteration",
|
||||||
|
"on_stream:hi", "on_stream:hi",
|
||||||
|
"on_stream_end:True", "on_stream_end:True",
|
||||||
|
"before_execute_tools", "before_execute_tools",
|
||||||
|
"after_iteration", "after_iteration",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Error isolation: one hook raises, others still run
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_composite_error_isolation_before_iteration():
|
||||||
|
calls: list[str] = []
|
||||||
|
|
||||||
|
class Bad(AgentHook):
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
raise RuntimeError("boom")
|
||||||
|
|
||||||
|
class Good(AgentHook):
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
calls.append("good")
|
||||||
|
|
||||||
|
hook = CompositeHook([Bad(), Good()])
|
||||||
|
await hook.before_iteration(_ctx())
|
||||||
|
assert calls == ["good"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_composite_error_isolation_on_stream():
|
||||||
|
calls: list[str] = []
|
||||||
|
|
||||||
|
class Bad(AgentHook):
|
||||||
|
async def on_stream(self, context: AgentHookContext, delta: str) -> None:
|
||||||
|
raise RuntimeError("stream-boom")
|
||||||
|
|
||||||
|
class Good(AgentHook):
|
||||||
|
async def on_stream(self, context: AgentHookContext, delta: str) -> None:
|
||||||
|
calls.append(delta)
|
||||||
|
|
||||||
|
hook = CompositeHook([Bad(), Good()])
|
||||||
|
await hook.on_stream(_ctx(), "delta")
|
||||||
|
assert calls == ["delta"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_composite_error_isolation_all_async():
|
||||||
|
"""Error isolation for on_stream_end, before_execute_tools, after_iteration."""
|
||||||
|
calls: list[str] = []
|
||||||
|
|
||||||
|
class Bad(AgentHook):
|
||||||
|
async def on_stream_end(self, context, *, resuming):
|
||||||
|
raise RuntimeError("err")
|
||||||
|
async def before_execute_tools(self, context):
|
||||||
|
raise RuntimeError("err")
|
||||||
|
async def after_iteration(self, context):
|
||||||
|
raise RuntimeError("err")
|
||||||
|
|
||||||
|
class Good(AgentHook):
|
||||||
|
async def on_stream_end(self, context, *, resuming):
|
||||||
|
calls.append("on_stream_end")
|
||||||
|
async def before_execute_tools(self, context):
|
||||||
|
calls.append("before_execute_tools")
|
||||||
|
async def after_iteration(self, context):
|
||||||
|
calls.append("after_iteration")
|
||||||
|
|
||||||
|
hook = CompositeHook([Bad(), Good()])
|
||||||
|
ctx = _ctx()
|
||||||
|
await hook.on_stream_end(ctx, resuming=False)
|
||||||
|
await hook.before_execute_tools(ctx)
|
||||||
|
await hook.after_iteration(ctx)
|
||||||
|
assert calls == ["on_stream_end", "before_execute_tools", "after_iteration"]
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# finalize_content: pipeline semantics (no error isolation)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_composite_finalize_content_pipeline():
|
||||||
|
class Upper(AgentHook):
|
||||||
|
def finalize_content(self, context, content):
|
||||||
|
return content.upper() if content else content
|
||||||
|
|
||||||
|
class Suffix(AgentHook):
|
||||||
|
def finalize_content(self, context, content):
|
||||||
|
return (content + "!") if content else content
|
||||||
|
|
||||||
|
hook = CompositeHook([Upper(), Suffix()])
|
||||||
|
result = hook.finalize_content(_ctx(), "hello")
|
||||||
|
assert result == "HELLO!"
|
||||||
|
|
||||||
|
|
||||||
|
def test_composite_finalize_content_none_passthrough():
|
||||||
|
hook = CompositeHook([AgentHook()])
|
||||||
|
assert hook.finalize_content(_ctx(), None) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_composite_finalize_content_ordering():
|
||||||
|
"""First hook transforms first, result feeds second hook."""
|
||||||
|
steps: list[str] = []
|
||||||
|
|
||||||
|
class H1(AgentHook):
|
||||||
|
def finalize_content(self, context, content):
|
||||||
|
steps.append(f"H1:{content}")
|
||||||
|
return content.upper()
|
||||||
|
|
||||||
|
class H2(AgentHook):
|
||||||
|
def finalize_content(self, context, content):
|
||||||
|
steps.append(f"H2:{content}")
|
||||||
|
return content + "!"
|
||||||
|
|
||||||
|
hook = CompositeHook([H1(), H2()])
|
||||||
|
result = hook.finalize_content(_ctx(), "hi")
|
||||||
|
assert result == "HI!"
|
||||||
|
assert steps == ["H1:hi", "H2:HI"]
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# wants_streaming: any-semantics
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_composite_wants_streaming_any_true():
|
||||||
|
class No(AgentHook):
|
||||||
|
def wants_streaming(self):
|
||||||
|
return False
|
||||||
|
|
||||||
|
class Yes(AgentHook):
|
||||||
|
def wants_streaming(self):
|
||||||
|
return True
|
||||||
|
|
||||||
|
hook = CompositeHook([No(), Yes(), No()])
|
||||||
|
assert hook.wants_streaming() is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_composite_wants_streaming_all_false():
|
||||||
|
hook = CompositeHook([AgentHook(), AgentHook()])
|
||||||
|
assert hook.wants_streaming() is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_composite_wants_streaming_empty():
|
||||||
|
hook = CompositeHook([])
|
||||||
|
assert hook.wants_streaming() is False
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Empty hooks list: behaves like no-op AgentHook
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_composite_empty_hooks_no_ops():
|
||||||
|
hook = CompositeHook([])
|
||||||
|
ctx = _ctx()
|
||||||
|
await hook.before_iteration(ctx)
|
||||||
|
await hook.on_stream(ctx, "delta")
|
||||||
|
await hook.on_stream_end(ctx, resuming=False)
|
||||||
|
await hook.before_execute_tools(ctx)
|
||||||
|
await hook.after_iteration(ctx)
|
||||||
|
assert hook.finalize_content(ctx, "test") == "test"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Integration: AgentLoop with extra hooks
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _make_loop(tmp_path, hooks=None):
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
provider.generation.max_tokens = 4096
|
||||||
|
|
||||||
|
with patch("nanobot.agent.loop.ContextBuilder"), \
|
||||||
|
patch("nanobot.agent.loop.SessionManager"), \
|
||||||
|
patch("nanobot.agent.loop.SubagentManager") as mock_sub_mgr, \
|
||||||
|
patch("nanobot.agent.loop.MemoryConsolidator"):
|
||||||
|
mock_sub_mgr.return_value.cancel_by_session = AsyncMock(return_value=0)
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=bus, provider=provider, workspace=tmp_path, hooks=hooks,
|
||||||
|
)
|
||||||
|
return loop
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_agent_loop_extra_hook_receives_calls(tmp_path):
|
||||||
|
"""Extra hook passed to AgentLoop is called alongside core LoopHook."""
|
||||||
|
from nanobot.providers.base import LLMResponse
|
||||||
|
|
||||||
|
events: list[str] = []
|
||||||
|
|
||||||
|
class TrackingHook(AgentHook):
|
||||||
|
async def before_iteration(self, context):
|
||||||
|
events.append(f"before_iter:{context.iteration}")
|
||||||
|
|
||||||
|
async def after_iteration(self, context):
|
||||||
|
events.append(f"after_iter:{context.iteration}")
|
||||||
|
|
||||||
|
loop = _make_loop(tmp_path, hooks=[TrackingHook()])
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(
|
||||||
|
return_value=LLMResponse(content="done", tool_calls=[], usage={})
|
||||||
|
)
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
|
||||||
|
content, tools_used, messages = await loop._run_agent_loop(
|
||||||
|
[{"role": "user", "content": "hi"}]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert content == "done"
|
||||||
|
assert "before_iter:0" in events
|
||||||
|
assert "after_iter:0" in events
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_agent_loop_extra_hook_error_isolation(tmp_path):
|
||||||
|
"""A faulty extra hook does not crash the agent loop."""
|
||||||
|
from nanobot.providers.base import LLMResponse
|
||||||
|
|
||||||
|
class BadHook(AgentHook):
|
||||||
|
async def before_iteration(self, context):
|
||||||
|
raise RuntimeError("I am broken")
|
||||||
|
|
||||||
|
loop = _make_loop(tmp_path, hooks=[BadHook()])
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(
|
||||||
|
return_value=LLMResponse(content="still works", tool_calls=[], usage={})
|
||||||
|
)
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
|
||||||
|
content, _, _ = await loop._run_agent_loop(
|
||||||
|
[{"role": "user", "content": "hi"}]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert content == "still works"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_agent_loop_extra_hooks_do_not_swallow_loop_hook_errors(tmp_path):
|
||||||
|
"""Extra hooks must not change the core LoopHook failure behavior."""
|
||||||
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
|
||||||
|
loop = _make_loop(tmp_path, hooks=[AgentHook()])
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||||
|
content="working",
|
||||||
|
tool_calls=[ToolCallRequest(id="c1", name="list_dir", arguments={"path": "."})],
|
||||||
|
usage={},
|
||||||
|
))
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
loop.tools.execute = AsyncMock(return_value="ok")
|
||||||
|
|
||||||
|
async def bad_progress(*args, **kwargs):
|
||||||
|
raise RuntimeError("progress failed")
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="progress failed"):
|
||||||
|
await loop._run_agent_loop([], on_progress=bad_progress)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_agent_loop_no_hooks_backward_compat(tmp_path):
|
||||||
|
"""Without hooks param, behavior is identical to before."""
|
||||||
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||||
|
content="working",
|
||||||
|
tool_calls=[ToolCallRequest(id="c1", name="list_dir", arguments={"path": "."})],
|
||||||
|
))
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
loop.tools.execute = AsyncMock(return_value="ok")
|
||||||
|
loop.max_iterations = 2
|
||||||
|
|
||||||
|
content, tools_used, _ = await loop._run_agent_loop([])
|
||||||
|
assert content == (
|
||||||
|
"I reached the maximum number of tool call iterations (2) "
|
||||||
|
"without completing the task. You can try breaking the task into smaller steps."
|
||||||
|
)
|
||||||
|
assert tools_used == ["list_dir", "list_dir"]
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.cron import CronTool
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.cron.service import CronService
|
||||||
|
|
||||||
|
|
||||||
|
def test_agent_loop_registers_cron_tool_with_configured_timezone(tmp_path: Path) -> None:
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=bus,
|
||||||
|
provider=provider,
|
||||||
|
workspace=tmp_path,
|
||||||
|
model="test-model",
|
||||||
|
cron_service=CronService(tmp_path / "cron" / "jobs.json"),
|
||||||
|
timezone="Asia/Shanghai",
|
||||||
|
)
|
||||||
|
|
||||||
|
cron_tool = loop.tools.get("cron")
|
||||||
|
|
||||||
|
assert isinstance(cron_tool, CronTool)
|
||||||
|
assert cron_tool._default_timezone == "Asia/Shanghai"
|
||||||
@@ -0,0 +1,335 @@
|
|||||||
|
"""Tests for the shared agent runner and its integration contracts."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
|
||||||
|
|
||||||
|
def _make_loop(tmp_path):
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
|
||||||
|
with patch("nanobot.agent.loop.ContextBuilder"), \
|
||||||
|
patch("nanobot.agent.loop.SessionManager"), \
|
||||||
|
patch("nanobot.agent.loop.SubagentManager") as MockSubMgr:
|
||||||
|
MockSubMgr.return_value.cancel_by_session = AsyncMock(return_value=0)
|
||||||
|
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path)
|
||||||
|
return loop
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_preserves_reasoning_fields_and_tool_results():
|
||||||
|
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||||
|
|
||||||
|
provider = MagicMock()
|
||||||
|
captured_second_call: list[dict] = []
|
||||||
|
call_count = {"n": 0}
|
||||||
|
|
||||||
|
async def chat_with_retry(*, messages, **kwargs):
|
||||||
|
call_count["n"] += 1
|
||||||
|
if call_count["n"] == 1:
|
||||||
|
return LLMResponse(
|
||||||
|
content="thinking",
|
||||||
|
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
|
||||||
|
reasoning_content="hidden reasoning",
|
||||||
|
thinking_blocks=[{"type": "thinking", "thinking": "step"}],
|
||||||
|
usage={"prompt_tokens": 5, "completion_tokens": 3},
|
||||||
|
)
|
||||||
|
captured_second_call[:] = messages
|
||||||
|
return LLMResponse(content="done", tool_calls=[], usage={})
|
||||||
|
|
||||||
|
provider.chat_with_retry = chat_with_retry
|
||||||
|
tools = MagicMock()
|
||||||
|
tools.get_definitions.return_value = []
|
||||||
|
tools.execute = AsyncMock(return_value="tool result")
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
result = await runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[
|
||||||
|
{"role": "system", "content": "system"},
|
||||||
|
{"role": "user", "content": "do task"},
|
||||||
|
],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=3,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.final_content == "done"
|
||||||
|
assert result.tools_used == ["list_dir"]
|
||||||
|
assert result.tool_events == [
|
||||||
|
{"name": "list_dir", "status": "ok", "detail": "tool result"}
|
||||||
|
]
|
||||||
|
|
||||||
|
assistant_messages = [
|
||||||
|
msg for msg in captured_second_call
|
||||||
|
if msg.get("role") == "assistant" and msg.get("tool_calls")
|
||||||
|
]
|
||||||
|
assert len(assistant_messages) == 1
|
||||||
|
assert assistant_messages[0]["reasoning_content"] == "hidden reasoning"
|
||||||
|
assert assistant_messages[0]["thinking_blocks"] == [{"type": "thinking", "thinking": "step"}]
|
||||||
|
assert any(
|
||||||
|
msg.get("role") == "tool" and msg.get("content") == "tool result"
|
||||||
|
for msg in captured_second_call
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_calls_hooks_in_order():
|
||||||
|
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||||
|
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||||
|
|
||||||
|
provider = MagicMock()
|
||||||
|
call_count = {"n": 0}
|
||||||
|
events: list[tuple] = []
|
||||||
|
|
||||||
|
async def chat_with_retry(**kwargs):
|
||||||
|
call_count["n"] += 1
|
||||||
|
if call_count["n"] == 1:
|
||||||
|
return LLMResponse(
|
||||||
|
content="thinking",
|
||||||
|
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
|
||||||
|
)
|
||||||
|
return LLMResponse(content="done", tool_calls=[], usage={})
|
||||||
|
|
||||||
|
provider.chat_with_retry = chat_with_retry
|
||||||
|
tools = MagicMock()
|
||||||
|
tools.get_definitions.return_value = []
|
||||||
|
tools.execute = AsyncMock(return_value="tool result")
|
||||||
|
|
||||||
|
class RecordingHook(AgentHook):
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
events.append(("before_iteration", context.iteration))
|
||||||
|
|
||||||
|
async def before_execute_tools(self, context: AgentHookContext) -> None:
|
||||||
|
events.append((
|
||||||
|
"before_execute_tools",
|
||||||
|
context.iteration,
|
||||||
|
[tc.name for tc in context.tool_calls],
|
||||||
|
))
|
||||||
|
|
||||||
|
async def after_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
events.append((
|
||||||
|
"after_iteration",
|
||||||
|
context.iteration,
|
||||||
|
context.final_content,
|
||||||
|
list(context.tool_results),
|
||||||
|
list(context.tool_events),
|
||||||
|
context.stop_reason,
|
||||||
|
))
|
||||||
|
|
||||||
|
def finalize_content(self, context: AgentHookContext, content: str | None) -> str | None:
|
||||||
|
events.append(("finalize_content", context.iteration, content))
|
||||||
|
return content.upper() if content else content
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
result = await runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=3,
|
||||||
|
hook=RecordingHook(),
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.final_content == "DONE"
|
||||||
|
assert events == [
|
||||||
|
("before_iteration", 0),
|
||||||
|
("before_execute_tools", 0, ["list_dir"]),
|
||||||
|
(
|
||||||
|
"after_iteration",
|
||||||
|
0,
|
||||||
|
None,
|
||||||
|
["tool result"],
|
||||||
|
[{"name": "list_dir", "status": "ok", "detail": "tool result"}],
|
||||||
|
None,
|
||||||
|
),
|
||||||
|
("before_iteration", 1),
|
||||||
|
("finalize_content", 1, "done"),
|
||||||
|
("after_iteration", 1, "DONE", [], [], "completed"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_streaming_hook_receives_deltas_and_end_signal():
|
||||||
|
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||||
|
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||||
|
|
||||||
|
provider = MagicMock()
|
||||||
|
streamed: list[str] = []
|
||||||
|
endings: list[bool] = []
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(*, on_content_delta, **kwargs):
|
||||||
|
await on_content_delta("he")
|
||||||
|
await on_content_delta("llo")
|
||||||
|
return LLMResponse(content="hello", tool_calls=[], usage={})
|
||||||
|
|
||||||
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
|
tools = MagicMock()
|
||||||
|
tools.get_definitions.return_value = []
|
||||||
|
|
||||||
|
class StreamingHook(AgentHook):
|
||||||
|
def wants_streaming(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def on_stream(self, context: AgentHookContext, delta: str) -> None:
|
||||||
|
streamed.append(delta)
|
||||||
|
|
||||||
|
async def on_stream_end(self, context: AgentHookContext, *, resuming: bool) -> None:
|
||||||
|
endings.append(resuming)
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
result = await runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=1,
|
||||||
|
hook=StreamingHook(),
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.final_content == "hello"
|
||||||
|
assert streamed == ["he", "llo"]
|
||||||
|
assert endings == [False]
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_returns_max_iterations_fallback():
|
||||||
|
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||||
|
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||||
|
content="still working",
|
||||||
|
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
|
||||||
|
))
|
||||||
|
tools = MagicMock()
|
||||||
|
tools.get_definitions.return_value = []
|
||||||
|
tools.execute = AsyncMock(return_value="tool result")
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
result = await runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=2,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.stop_reason == "max_iterations"
|
||||||
|
assert result.final_content == (
|
||||||
|
"I reached the maximum number of tool call iterations (2) "
|
||||||
|
"without completing the task. You can try breaking the task into smaller steps."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_returns_structured_tool_error():
|
||||||
|
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||||
|
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||||
|
content="working",
|
||||||
|
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={})],
|
||||||
|
))
|
||||||
|
tools = MagicMock()
|
||||||
|
tools.get_definitions.return_value = []
|
||||||
|
tools.execute = AsyncMock(side_effect=RuntimeError("boom"))
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
|
||||||
|
result = await runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=2,
|
||||||
|
fail_on_tool_error=True,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.stop_reason == "tool_error"
|
||||||
|
assert result.error == "Error: RuntimeError: boom"
|
||||||
|
assert result.tool_events == [
|
||||||
|
{"name": "list_dir", "status": "error", "detail": "boom"}
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_loop_max_iterations_message_stays_stable(tmp_path):
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||||
|
content="working",
|
||||||
|
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={})],
|
||||||
|
))
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
loop.tools.execute = AsyncMock(return_value="ok")
|
||||||
|
loop.max_iterations = 2
|
||||||
|
|
||||||
|
final_content, _, _ = await loop._run_agent_loop([])
|
||||||
|
|
||||||
|
assert final_content == (
|
||||||
|
"I reached the maximum number of tool call iterations (2) "
|
||||||
|
"without completing the task. You can try breaking the task into smaller steps."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_loop_stream_filter_handles_think_only_prefix_without_crashing(tmp_path):
|
||||||
|
loop = _make_loop(tmp_path)
|
||||||
|
deltas: list[str] = []
|
||||||
|
endings: list[bool] = []
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(*, on_content_delta, **kwargs):
|
||||||
|
await on_content_delta("<think>hidden")
|
||||||
|
await on_content_delta("</think>Hello")
|
||||||
|
return LLMResponse(content="<think>hidden</think>Hello", tool_calls=[], usage={})
|
||||||
|
|
||||||
|
loop.provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
|
||||||
|
async def on_stream(delta: str) -> None:
|
||||||
|
deltas.append(delta)
|
||||||
|
|
||||||
|
async def on_stream_end(*, resuming: bool = False) -> None:
|
||||||
|
endings.append(resuming)
|
||||||
|
|
||||||
|
final_content, _, _ = await loop._run_agent_loop(
|
||||||
|
[],
|
||||||
|
on_stream=on_stream,
|
||||||
|
on_stream_end=on_stream_end,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert final_content == "Hello"
|
||||||
|
assert deltas == ["Hello"]
|
||||||
|
assert endings == [False]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_subagent_max_iterations_announces_existing_fallback(tmp_path, monkeypatch):
|
||||||
|
from nanobot.agent.subagent import SubagentManager
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||||
|
content="working",
|
||||||
|
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={})],
|
||||||
|
))
|
||||||
|
mgr = SubagentManager(provider=provider, workspace=tmp_path, bus=bus)
|
||||||
|
mgr._announce_result = AsyncMock()
|
||||||
|
|
||||||
|
async def fake_execute(self, name, arguments):
|
||||||
|
return "tool result"
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.agent.tools.registry.ToolRegistry.execute", fake_execute)
|
||||||
|
|
||||||
|
await mgr._run_subagent("sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"})
|
||||||
|
|
||||||
|
mgr._announce_result.assert_awaited_once()
|
||||||
|
args = mgr._announce_result.await_args.args
|
||||||
|
assert args[3] == "Task completed but no final response was generated."
|
||||||
|
assert args[5] == "ok"
|
||||||
@@ -3,6 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -116,6 +117,43 @@ class TestDispatch:
|
|||||||
out = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
out = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||||
assert out.content == "hi"
|
assert out.content == "hi"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_dispatch_streaming_preserves_message_metadata(self):
|
||||||
|
from nanobot.bus.events import InboundMessage
|
||||||
|
|
||||||
|
loop, bus = _make_loop()
|
||||||
|
msg = InboundMessage(
|
||||||
|
channel="matrix",
|
||||||
|
sender_id="u1",
|
||||||
|
chat_id="!room:matrix.org",
|
||||||
|
content="hello",
|
||||||
|
metadata={
|
||||||
|
"_wants_stream": True,
|
||||||
|
"thread_root_event_id": "$root1",
|
||||||
|
"thread_reply_to_event_id": "$reply1",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def fake_process(_msg, *, on_stream=None, on_stream_end=None, **kwargs):
|
||||||
|
assert on_stream is not None
|
||||||
|
assert on_stream_end is not None
|
||||||
|
await on_stream("hi")
|
||||||
|
await on_stream_end(resuming=False)
|
||||||
|
return None
|
||||||
|
|
||||||
|
loop._process_message = fake_process
|
||||||
|
|
||||||
|
await loop._dispatch(msg)
|
||||||
|
first = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||||
|
second = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||||
|
|
||||||
|
assert first.metadata["thread_root_event_id"] == "$root1"
|
||||||
|
assert first.metadata["thread_reply_to_event_id"] == "$reply1"
|
||||||
|
assert first.metadata["_stream_delta"] is True
|
||||||
|
assert second.metadata["thread_root_event_id"] == "$root1"
|
||||||
|
assert second.metadata["thread_reply_to_event_id"] == "$reply1"
|
||||||
|
assert second.metadata["_stream_end"] is True
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_processing_lock_serializes(self):
|
async def test_processing_lock_serializes(self):
|
||||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||||
@@ -221,3 +259,116 @@ class TestSubagentCancellation:
|
|||||||
assert len(assistant_messages) == 1
|
assert len(assistant_messages) == 1
|
||||||
assert assistant_messages[0]["reasoning_content"] == "hidden reasoning"
|
assert assistant_messages[0]["reasoning_content"] == "hidden reasoning"
|
||||||
assert assistant_messages[0]["thinking_blocks"] == [{"type": "thinking", "thinking": "step"}]
|
assert assistant_messages[0]["thinking_blocks"] == [{"type": "thinking", "thinking": "step"}]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_subagent_exec_tool_not_registered_when_disabled(self, tmp_path):
|
||||||
|
from nanobot.agent.subagent import SubagentManager
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.config.schema import ExecToolConfig
|
||||||
|
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
mgr = SubagentManager(
|
||||||
|
provider=provider,
|
||||||
|
workspace=tmp_path,
|
||||||
|
bus=bus,
|
||||||
|
exec_config=ExecToolConfig(enable=False),
|
||||||
|
)
|
||||||
|
mgr._announce_result = AsyncMock()
|
||||||
|
|
||||||
|
async def fake_run(spec):
|
||||||
|
assert spec.tools.get("exec") is None
|
||||||
|
return SimpleNamespace(
|
||||||
|
stop_reason="done",
|
||||||
|
final_content="done",
|
||||||
|
error=None,
|
||||||
|
tool_events=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr.runner.run = AsyncMock(side_effect=fake_run)
|
||||||
|
|
||||||
|
await mgr._run_subagent("sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"})
|
||||||
|
|
||||||
|
mgr.runner.run.assert_awaited_once()
|
||||||
|
mgr._announce_result.assert_awaited_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_subagent_announces_error_when_tool_execution_fails(self, monkeypatch, tmp_path):
|
||||||
|
from nanobot.agent.subagent import SubagentManager
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||||
|
content="thinking",
|
||||||
|
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={})],
|
||||||
|
))
|
||||||
|
mgr = SubagentManager(provider=provider, workspace=tmp_path, bus=bus)
|
||||||
|
mgr._announce_result = AsyncMock()
|
||||||
|
|
||||||
|
calls = {"n": 0}
|
||||||
|
|
||||||
|
async def fake_execute(self, name, arguments):
|
||||||
|
calls["n"] += 1
|
||||||
|
if calls["n"] == 1:
|
||||||
|
return "first result"
|
||||||
|
raise RuntimeError("boom")
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.agent.tools.registry.ToolRegistry.execute", fake_execute)
|
||||||
|
|
||||||
|
await mgr._run_subagent("sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"})
|
||||||
|
|
||||||
|
mgr._announce_result.assert_awaited_once()
|
||||||
|
args = mgr._announce_result.await_args.args
|
||||||
|
assert "Completed steps:" in args[3]
|
||||||
|
assert "- list_dir: first result" in args[3]
|
||||||
|
assert "Failure:" in args[3]
|
||||||
|
assert "- list_dir: boom" in args[3]
|
||||||
|
assert args[5] == "error"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_cancel_by_session_cancels_running_subagent_tool(self, monkeypatch, tmp_path):
|
||||||
|
from nanobot.agent.subagent import SubagentManager
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||||
|
content="thinking",
|
||||||
|
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={})],
|
||||||
|
))
|
||||||
|
mgr = SubagentManager(provider=provider, workspace=tmp_path, bus=bus)
|
||||||
|
mgr._announce_result = AsyncMock()
|
||||||
|
|
||||||
|
started = asyncio.Event()
|
||||||
|
cancelled = asyncio.Event()
|
||||||
|
|
||||||
|
async def fake_execute(self, name, arguments):
|
||||||
|
started.set()
|
||||||
|
try:
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
cancelled.set()
|
||||||
|
raise
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.agent.tools.registry.ToolRegistry.execute", fake_execute)
|
||||||
|
|
||||||
|
task = asyncio.create_task(
|
||||||
|
mgr._run_subagent("sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"})
|
||||||
|
)
|
||||||
|
mgr._running_tasks["sub-1"] = task
|
||||||
|
mgr._session_tasks["test:c1"] = {"sub-1"}
|
||||||
|
|
||||||
|
await started.wait()
|
||||||
|
|
||||||
|
count = await mgr.cancel_by_session("test:c1")
|
||||||
|
|
||||||
|
assert count == 1
|
||||||
|
assert cancelled.is_set()
|
||||||
|
assert task.cancelled()
|
||||||
|
mgr._announce_result.assert_not_awaited()
|
||||||
|
|||||||
@@ -0,0 +1,298 @@
|
|||||||
|
"""Tests for ChannelManager delta coalescing to reduce streaming latency."""
|
||||||
|
import asyncio
|
||||||
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.bus.events import OutboundMessage
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.channels.base import BaseChannel
|
||||||
|
from nanobot.channels.manager import ChannelManager
|
||||||
|
from nanobot.config.schema import Config
|
||||||
|
|
||||||
|
|
||||||
|
class MockChannel(BaseChannel):
|
||||||
|
"""Mock channel for testing."""
|
||||||
|
|
||||||
|
name = "mock"
|
||||||
|
display_name = "Mock"
|
||||||
|
|
||||||
|
def __init__(self, config, bus):
|
||||||
|
super().__init__(config, bus)
|
||||||
|
self._send_delta_mock = AsyncMock()
|
||||||
|
self._send_mock = AsyncMock()
|
||||||
|
|
||||||
|
async def start(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg):
|
||||||
|
"""Implement abstract method."""
|
||||||
|
return await self._send_mock(msg)
|
||||||
|
|
||||||
|
async def send_delta(self, chat_id, delta, metadata=None):
|
||||||
|
"""Override send_delta for testing."""
|
||||||
|
return await self._send_delta_mock(chat_id, delta, metadata)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def config():
|
||||||
|
"""Create a minimal config for testing."""
|
||||||
|
return Config()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def bus():
|
||||||
|
"""Create a message bus for testing."""
|
||||||
|
return MessageBus()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def manager(config, bus):
|
||||||
|
"""Create a channel manager with a mock channel."""
|
||||||
|
manager = ChannelManager(config, bus)
|
||||||
|
manager.channels["mock"] = MockChannel({}, bus)
|
||||||
|
return manager
|
||||||
|
|
||||||
|
|
||||||
|
class TestDeltaCoalescing:
|
||||||
|
"""Tests for _stream_delta message coalescing."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_single_delta_not_coalesced(self, manager, bus):
|
||||||
|
"""A single delta should be sent as-is."""
|
||||||
|
msg = OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="Hello",
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
)
|
||||||
|
await bus.publish_outbound(msg)
|
||||||
|
|
||||||
|
# Process one message
|
||||||
|
async def process_one():
|
||||||
|
try:
|
||||||
|
m = await asyncio.wait_for(bus.consume_outbound(), timeout=0.1)
|
||||||
|
if m.metadata.get("_stream_delta"):
|
||||||
|
m, pending = manager._coalesce_stream_deltas(m)
|
||||||
|
# Put pending back (none expected)
|
||||||
|
for p in pending:
|
||||||
|
await bus.publish_outbound(p)
|
||||||
|
channel = manager.channels.get(m.channel)
|
||||||
|
if channel:
|
||||||
|
await channel.send_delta(m.chat_id, m.content, m.metadata)
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
await process_one()
|
||||||
|
|
||||||
|
manager.channels["mock"]._send_delta_mock.assert_called_once_with(
|
||||||
|
"chat1", "Hello", {"_stream_delta": True}
|
||||||
|
)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_multiple_deltas_coalesced(self, manager, bus):
|
||||||
|
"""Multiple consecutive deltas for same chat should be merged."""
|
||||||
|
# Put multiple deltas in queue
|
||||||
|
for text in ["Hello", " ", "world", "!"]:
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content=text,
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
))
|
||||||
|
|
||||||
|
# Process using coalescing logic
|
||||||
|
first_msg = await bus.consume_outbound()
|
||||||
|
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||||
|
|
||||||
|
# Should have merged all deltas
|
||||||
|
assert merged.content == "Hello world!"
|
||||||
|
assert merged.metadata.get("_stream_delta") is True
|
||||||
|
# No pending messages (all were coalesced)
|
||||||
|
assert len(pending) == 0
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_deltas_different_chats_not_coalesced(self, manager, bus):
|
||||||
|
"""Deltas for different chats should not be merged."""
|
||||||
|
# Put deltas for different chats
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="Hello",
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
))
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat2",
|
||||||
|
content="World",
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
))
|
||||||
|
|
||||||
|
first_msg = await bus.consume_outbound()
|
||||||
|
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||||
|
|
||||||
|
# First chat should not include second chat's content
|
||||||
|
assert merged.content == "Hello"
|
||||||
|
assert merged.chat_id == "chat1"
|
||||||
|
# Second chat should be in pending
|
||||||
|
assert len(pending) == 1
|
||||||
|
assert pending[0].chat_id == "chat2"
|
||||||
|
assert pending[0].content == "World"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stream_end_terminates_coalescing(self, manager, bus):
|
||||||
|
"""_stream_end should stop coalescing and be included in final message."""
|
||||||
|
# Put deltas with stream_end at the end
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="Hello",
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
))
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content=" world",
|
||||||
|
metadata={"_stream_delta": True, "_stream_end": True},
|
||||||
|
))
|
||||||
|
|
||||||
|
first_msg = await bus.consume_outbound()
|
||||||
|
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||||
|
|
||||||
|
# Should have merged content
|
||||||
|
assert merged.content == "Hello world"
|
||||||
|
# Should have stream_end flag
|
||||||
|
assert merged.metadata.get("_stream_end") is True
|
||||||
|
# No pending
|
||||||
|
assert len(pending) == 0
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_coalescing_stops_at_first_non_matching_boundary(self, manager, bus):
|
||||||
|
"""Only consecutive deltas should be merged; later deltas stay queued."""
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="Hello",
|
||||||
|
metadata={"_stream_delta": True, "_stream_id": "seg-1"},
|
||||||
|
))
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="",
|
||||||
|
metadata={"_stream_end": True, "_stream_id": "seg-1"},
|
||||||
|
))
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="world",
|
||||||
|
metadata={"_stream_delta": True, "_stream_id": "seg-2"},
|
||||||
|
))
|
||||||
|
|
||||||
|
first_msg = await bus.consume_outbound()
|
||||||
|
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||||
|
|
||||||
|
assert merged.content == "Hello"
|
||||||
|
assert merged.metadata.get("_stream_end") is None
|
||||||
|
assert len(pending) == 1
|
||||||
|
assert pending[0].metadata.get("_stream_end") is True
|
||||||
|
assert pending[0].metadata.get("_stream_id") == "seg-1"
|
||||||
|
|
||||||
|
# The next stream segment must remain in queue order for later dispatch.
|
||||||
|
remaining = await bus.consume_outbound()
|
||||||
|
assert remaining.content == "world"
|
||||||
|
assert remaining.metadata.get("_stream_id") == "seg-2"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_non_delta_message_preserved(self, manager, bus):
|
||||||
|
"""Non-delta messages should be preserved in pending list."""
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="Delta",
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
))
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="Final message",
|
||||||
|
metadata={}, # Not a delta
|
||||||
|
))
|
||||||
|
|
||||||
|
first_msg = await bus.consume_outbound()
|
||||||
|
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||||
|
|
||||||
|
assert merged.content == "Delta"
|
||||||
|
assert len(pending) == 1
|
||||||
|
assert pending[0].content == "Final message"
|
||||||
|
assert pending[0].metadata.get("_stream_delta") is None
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_empty_queue_stops_coalescing(self, manager, bus):
|
||||||
|
"""Coalescing should stop when queue is empty."""
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="Only message",
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
))
|
||||||
|
|
||||||
|
first_msg = await bus.consume_outbound()
|
||||||
|
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||||
|
|
||||||
|
assert merged.content == "Only message"
|
||||||
|
assert len(pending) == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestDispatchOutboundWithCoalescing:
|
||||||
|
"""Tests for the full _dispatch_outbound flow with coalescing."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_dispatch_coalesces_and_processes_pending(self, manager, bus):
|
||||||
|
"""_dispatch_outbound should coalesce deltas and process pending messages."""
|
||||||
|
# Put multiple deltas followed by a regular message
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="A",
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
))
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="B",
|
||||||
|
metadata={"_stream_delta": True},
|
||||||
|
))
|
||||||
|
await bus.publish_outbound(OutboundMessage(
|
||||||
|
channel="mock",
|
||||||
|
chat_id="chat1",
|
||||||
|
content="Final",
|
||||||
|
metadata={}, # Regular message
|
||||||
|
))
|
||||||
|
|
||||||
|
# Run one iteration of dispatch logic manually
|
||||||
|
pending = []
|
||||||
|
processed = []
|
||||||
|
|
||||||
|
# First iteration: should coalesce A+B
|
||||||
|
if pending:
|
||||||
|
msg = pending.pop(0)
|
||||||
|
else:
|
||||||
|
msg = await bus.consume_outbound()
|
||||||
|
|
||||||
|
if msg.metadata.get("_stream_delta") and not msg.metadata.get("_stream_end"):
|
||||||
|
msg, extra_pending = manager._coalesce_stream_deltas(msg)
|
||||||
|
pending.extend(extra_pending)
|
||||||
|
|
||||||
|
channel = manager.channels.get(msg.channel)
|
||||||
|
if channel:
|
||||||
|
await channel.send_delta(msg.chat_id, msg.content, msg.metadata)
|
||||||
|
processed.append(("delta", msg.content))
|
||||||
|
|
||||||
|
# Should have sent coalesced delta
|
||||||
|
assert processed == [("delta", "AB")]
|
||||||
|
# Should have pending regular message
|
||||||
|
assert len(pending) == 1
|
||||||
|
assert pending[0].content == "Final"
|
||||||
@@ -2,8 +2,9 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import patch
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
@@ -262,3 +263,618 @@ def test_builtin_channel_init_from_dict():
|
|||||||
ch = TelegramChannel({"enabled": False, "token": "test-tok", "allowFrom": ["*"]}, bus)
|
ch = TelegramChannel({"enabled": False, "token": "test-tok", "allowFrom": ["*"]}, bus)
|
||||||
assert ch.config.token == "test-tok"
|
assert ch.config.token == "test-tok"
|
||||||
assert ch.config.allow_from == ["*"]
|
assert ch.config.allow_from == ["*"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_channels_config_send_max_retries_default():
|
||||||
|
"""ChannelsConfig should have send_max_retries with default value of 3."""
|
||||||
|
cfg = ChannelsConfig()
|
||||||
|
assert hasattr(cfg, 'send_max_retries')
|
||||||
|
assert cfg.send_max_retries == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_channels_config_send_max_retries_upper_bound():
|
||||||
|
"""send_max_retries should be bounded to prevent resource exhaustion."""
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
# Value too high should be rejected
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
ChannelsConfig(send_max_retries=100)
|
||||||
|
|
||||||
|
# Negative should be rejected
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
ChannelsConfig(send_max_retries=-1)
|
||||||
|
|
||||||
|
# Boundary values should be allowed
|
||||||
|
cfg_min = ChannelsConfig(send_max_retries=0)
|
||||||
|
assert cfg_min.send_max_retries == 0
|
||||||
|
|
||||||
|
cfg_max = ChannelsConfig(send_max_retries=10)
|
||||||
|
assert cfg_max.send_max_retries == 10
|
||||||
|
|
||||||
|
# Value above upper bound should be rejected
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
ChannelsConfig(send_max_retries=11)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# _send_with_retry
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_with_retry_succeeds_first_try():
|
||||||
|
"""_send_with_retry should succeed on first try and not retry."""
|
||||||
|
call_count = 0
|
||||||
|
|
||||||
|
class _FailingChannel(BaseChannel):
|
||||||
|
name = "failing"
|
||||||
|
display_name = "Failing"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
# Succeeds on first try
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(send_max_retries=3),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"failing": _FailingChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
msg = OutboundMessage(channel="failing", chat_id="123", content="test")
|
||||||
|
await mgr._send_with_retry(mgr.channels["failing"], msg)
|
||||||
|
|
||||||
|
assert call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_with_retry_retries_on_failure():
|
||||||
|
"""_send_with_retry should retry on failure up to max_retries times."""
|
||||||
|
call_count = 0
|
||||||
|
|
||||||
|
class _FailingChannel(BaseChannel):
|
||||||
|
name = "failing"
|
||||||
|
display_name = "Failing"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
raise RuntimeError("simulated failure")
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(send_max_retries=3),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"failing": _FailingChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
msg = OutboundMessage(channel="failing", chat_id="123", content="test")
|
||||||
|
|
||||||
|
# Patch asyncio.sleep to avoid actual delays
|
||||||
|
with patch("nanobot.channels.manager.asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
|
||||||
|
await mgr._send_with_retry(mgr.channels["failing"], msg)
|
||||||
|
|
||||||
|
assert call_count == 3 # 3 total attempts (initial + 2 retries)
|
||||||
|
assert mock_sleep.call_count == 2 # 2 sleeps between retries
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_with_retry_no_retry_when_max_is_zero():
|
||||||
|
"""_send_with_retry should not retry when send_max_retries is 0."""
|
||||||
|
call_count = 0
|
||||||
|
|
||||||
|
class _FailingChannel(BaseChannel):
|
||||||
|
name = "failing"
|
||||||
|
display_name = "Failing"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
raise RuntimeError("simulated failure")
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(send_max_retries=0),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"failing": _FailingChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
msg = OutboundMessage(channel="failing", chat_id="123", content="test")
|
||||||
|
|
||||||
|
with patch("nanobot.channels.manager.asyncio.sleep", new_callable=AsyncMock):
|
||||||
|
await mgr._send_with_retry(mgr.channels["failing"], msg)
|
||||||
|
|
||||||
|
assert call_count == 1 # Called once but no retry (max(0, 1) = 1)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_with_retry_calls_send_delta():
|
||||||
|
"""_send_with_retry should call send_delta when metadata has _stream_delta."""
|
||||||
|
send_delta_called = False
|
||||||
|
|
||||||
|
class _StreamingChannel(BaseChannel):
|
||||||
|
name = "streaming"
|
||||||
|
display_name = "Streaming"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
pass # Should not be called
|
||||||
|
|
||||||
|
async def send_delta(self, chat_id: str, delta: str, metadata: dict | None = None) -> None:
|
||||||
|
nonlocal send_delta_called
|
||||||
|
send_delta_called = True
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(send_max_retries=3),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"streaming": _StreamingChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
msg = OutboundMessage(
|
||||||
|
channel="streaming", chat_id="123", content="test delta",
|
||||||
|
metadata={"_stream_delta": True}
|
||||||
|
)
|
||||||
|
await mgr._send_with_retry(mgr.channels["streaming"], msg)
|
||||||
|
|
||||||
|
assert send_delta_called is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_with_retry_skips_send_when_streamed():
|
||||||
|
"""_send_with_retry should not call send when metadata has _streamed flag."""
|
||||||
|
send_called = False
|
||||||
|
send_delta_called = False
|
||||||
|
|
||||||
|
class _StreamedChannel(BaseChannel):
|
||||||
|
name = "streamed"
|
||||||
|
display_name = "Streamed"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
nonlocal send_called
|
||||||
|
send_called = True
|
||||||
|
|
||||||
|
async def send_delta(self, chat_id: str, delta: str, metadata: dict | None = None) -> None:
|
||||||
|
nonlocal send_delta_called
|
||||||
|
send_delta_called = True
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(send_max_retries=3),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"streamed": _StreamedChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
# _streamed means message was already sent via send_delta, so skip send
|
||||||
|
msg = OutboundMessage(
|
||||||
|
channel="streamed", chat_id="123", content="test",
|
||||||
|
metadata={"_streamed": True}
|
||||||
|
)
|
||||||
|
await mgr._send_with_retry(mgr.channels["streamed"], msg)
|
||||||
|
|
||||||
|
assert send_called is False
|
||||||
|
assert send_delta_called is False
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_with_retry_propagates_cancelled_error():
|
||||||
|
"""_send_with_retry should re-raise CancelledError for graceful shutdown."""
|
||||||
|
class _CancellingChannel(BaseChannel):
|
||||||
|
name = "cancelling"
|
||||||
|
display_name = "Cancelling"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
raise asyncio.CancelledError("simulated cancellation")
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(send_max_retries=3),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"cancelling": _CancellingChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
msg = OutboundMessage(channel="cancelling", chat_id="123", content="test")
|
||||||
|
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await mgr._send_with_retry(mgr.channels["cancelling"], msg)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_with_retry_propagates_cancelled_error_during_sleep():
|
||||||
|
"""_send_with_retry should re-raise CancelledError during sleep."""
|
||||||
|
call_count = 0
|
||||||
|
|
||||||
|
class _FailingChannel(BaseChannel):
|
||||||
|
name = "failing"
|
||||||
|
display_name = "Failing"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
raise RuntimeError("simulated failure")
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(send_max_retries=3),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"failing": _FailingChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
msg = OutboundMessage(channel="failing", chat_id="123", content="test")
|
||||||
|
|
||||||
|
# Mock sleep to raise CancelledError
|
||||||
|
async def cancel_during_sleep(_):
|
||||||
|
raise asyncio.CancelledError("cancelled during sleep")
|
||||||
|
|
||||||
|
with patch("nanobot.channels.manager.asyncio.sleep", side_effect=cancel_during_sleep):
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await mgr._send_with_retry(mgr.channels["failing"], msg)
|
||||||
|
|
||||||
|
# Should have attempted once before sleep was cancelled
|
||||||
|
assert call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# ChannelManager - lifecycle and getters
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class _ChannelWithAllowFrom(BaseChannel):
|
||||||
|
"""Channel with configurable allow_from."""
|
||||||
|
name = "withallow"
|
||||||
|
display_name = "With Allow"
|
||||||
|
|
||||||
|
def __init__(self, config, bus, allow_from):
|
||||||
|
super().__init__(config, bus)
|
||||||
|
self.config.allow_from = allow_from
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class _StartableChannel(BaseChannel):
|
||||||
|
"""Channel that tracks start/stop calls."""
|
||||||
|
name = "startable"
|
||||||
|
display_name = "Startable"
|
||||||
|
|
||||||
|
def __init__(self, config, bus):
|
||||||
|
super().__init__(config, bus)
|
||||||
|
self.started = False
|
||||||
|
self.stopped = False
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
self.started = True
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
self.stopped = True
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_validate_allow_from_raises_on_empty_list():
|
||||||
|
"""_validate_allow_from should raise SystemExit when allow_from is empty list."""
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.channels = {"test": _ChannelWithAllowFrom(fake_config, None, [])}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
with pytest.raises(SystemExit) as exc_info:
|
||||||
|
mgr._validate_allow_from()
|
||||||
|
|
||||||
|
assert "empty allowFrom" in str(exc_info.value)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_validate_allow_from_passes_with_asterisk():
|
||||||
|
"""_validate_allow_from should not raise when allow_from contains '*'."""
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.channels = {"test": _ChannelWithAllowFrom(fake_config, None, ["*"])}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
# Should not raise
|
||||||
|
mgr._validate_allow_from()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_channel_returns_channel_if_exists():
|
||||||
|
"""get_channel should return the channel if it exists."""
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"telegram": _StartableChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
assert mgr.get_channel("telegram") is not None
|
||||||
|
assert mgr.get_channel("nonexistent") is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_status_returns_running_state():
|
||||||
|
"""get_status should return enabled and running state for each channel."""
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
ch = _StartableChannel(fake_config, mgr.bus)
|
||||||
|
mgr.channels = {"startable": ch}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
status = mgr.get_status()
|
||||||
|
|
||||||
|
assert status["startable"]["enabled"] is True
|
||||||
|
assert status["startable"]["running"] is False # Not started yet
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_enabled_channels_returns_channel_names():
|
||||||
|
"""enabled_channels should return list of enabled channel names."""
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {
|
||||||
|
"telegram": _StartableChannel(fake_config, mgr.bus),
|
||||||
|
"slack": _StartableChannel(fake_config, mgr.bus),
|
||||||
|
}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
enabled = mgr.enabled_channels
|
||||||
|
|
||||||
|
assert "telegram" in enabled
|
||||||
|
assert "slack" in enabled
|
||||||
|
assert len(enabled) == 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stop_all_cancels_dispatcher_and_stops_channels():
|
||||||
|
"""stop_all should cancel the dispatch task and stop all channels."""
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
|
||||||
|
ch = _StartableChannel(fake_config, mgr.bus)
|
||||||
|
mgr.channels = {"startable": ch}
|
||||||
|
|
||||||
|
# Create a real cancelled task
|
||||||
|
async def dummy_task():
|
||||||
|
while True:
|
||||||
|
await asyncio.sleep(1)
|
||||||
|
|
||||||
|
dispatch_task = asyncio.create_task(dummy_task())
|
||||||
|
mgr._dispatch_task = dispatch_task
|
||||||
|
|
||||||
|
await mgr.stop_all()
|
||||||
|
|
||||||
|
# Task should be cancelled
|
||||||
|
assert dispatch_task.cancelled()
|
||||||
|
# Channel should be stopped
|
||||||
|
assert ch.stopped is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_channel_logs_error_on_failure():
|
||||||
|
"""_start_channel should log error when channel start fails."""
|
||||||
|
class _FailingChannel(BaseChannel):
|
||||||
|
name = "failing"
|
||||||
|
display_name = "Failing"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
raise RuntimeError("connection failed")
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
ch = _FailingChannel(fake_config, mgr.bus)
|
||||||
|
|
||||||
|
# Should not raise, just log error
|
||||||
|
await mgr._start_channel("failing", ch)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stop_all_handles_channel_exception():
|
||||||
|
"""stop_all should handle exceptions when stopping channels gracefully."""
|
||||||
|
class _StopFailingChannel(BaseChannel):
|
||||||
|
name = "stopfailing"
|
||||||
|
display_name = "Stop Failing"
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
raise RuntimeError("stop failed")
|
||||||
|
|
||||||
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {"stopfailing": _StopFailingChannel(fake_config, mgr.bus)}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
# Should not raise even if channel.stop() raises
|
||||||
|
await mgr.stop_all()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_all_no_channels_logs_warning():
|
||||||
|
"""start_all should log warning when no channels are enabled."""
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
mgr.channels = {} # No channels
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
# Should return early without creating dispatch task
|
||||||
|
await mgr.start_all()
|
||||||
|
|
||||||
|
assert mgr._dispatch_task is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_all_creates_dispatch_task():
|
||||||
|
"""start_all should create the dispatch task when channels exist."""
|
||||||
|
fake_config = SimpleNamespace(
|
||||||
|
channels=ChannelsConfig(),
|
||||||
|
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
mgr = ChannelManager.__new__(ChannelManager)
|
||||||
|
mgr.config = fake_config
|
||||||
|
mgr.bus = MessageBus()
|
||||||
|
|
||||||
|
ch = _StartableChannel(fake_config, mgr.bus)
|
||||||
|
mgr.channels = {"startable": ch}
|
||||||
|
mgr._dispatch_task = None
|
||||||
|
|
||||||
|
# Cancel immediately after start to avoid running forever
|
||||||
|
async def cancel_after_start():
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
if mgr._dispatch_task:
|
||||||
|
mgr._dispatch_task.cancel()
|
||||||
|
|
||||||
|
cancel_task = asyncio.create_task(cancel_after_start())
|
||||||
|
|
||||||
|
try:
|
||||||
|
await mgr.start_all()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
cancel_task.cancel()
|
||||||
|
try:
|
||||||
|
await cancel_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Dispatch task should have been created
|
||||||
|
assert mgr._dispatch_task is not None
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,676 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from pathlib import Path
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
discord = pytest.importorskip("discord")
|
||||||
|
|
||||||
|
from nanobot.bus.events import OutboundMessage
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.channels.discord import DiscordBotClient, DiscordChannel, DiscordConfig
|
||||||
|
from nanobot.command.builtin import build_help_text
|
||||||
|
|
||||||
|
|
||||||
|
# Minimal Discord client test double used to control startup/readiness behavior.
|
||||||
|
class _FakeDiscordClient:
|
||||||
|
instances: list["_FakeDiscordClient"] = []
|
||||||
|
start_error: Exception | None = None
|
||||||
|
|
||||||
|
def __init__(self, owner, *, intents) -> None:
|
||||||
|
self.owner = owner
|
||||||
|
self.intents = intents
|
||||||
|
self.closed = False
|
||||||
|
self.ready = True
|
||||||
|
self.channels: dict[int, object] = {}
|
||||||
|
self.user = SimpleNamespace(id=999)
|
||||||
|
self.__class__.instances.append(self)
|
||||||
|
|
||||||
|
async def start(self, token: str) -> None:
|
||||||
|
self.token = token
|
||||||
|
if self.__class__.start_error is not None:
|
||||||
|
raise self.__class__.start_error
|
||||||
|
|
||||||
|
async def close(self) -> None:
|
||||||
|
self.closed = True
|
||||||
|
|
||||||
|
def is_closed(self) -> bool:
|
||||||
|
return self.closed
|
||||||
|
|
||||||
|
def is_ready(self) -> bool:
|
||||||
|
return self.ready
|
||||||
|
|
||||||
|
def get_channel(self, channel_id: int):
|
||||||
|
return self.channels.get(channel_id)
|
||||||
|
|
||||||
|
async def send_outbound(self, msg: OutboundMessage) -> None:
|
||||||
|
channel = self.get_channel(int(msg.chat_id))
|
||||||
|
if channel is None:
|
||||||
|
return
|
||||||
|
await channel.send(content=msg.content)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeAttachment:
|
||||||
|
# Attachment double that can simulate successful or failing save() calls.
|
||||||
|
def __init__(self, attachment_id: int, filename: str, *, size: int = 1, fail: bool = False) -> None:
|
||||||
|
self.id = attachment_id
|
||||||
|
self.filename = filename
|
||||||
|
self.size = size
|
||||||
|
self._fail = fail
|
||||||
|
|
||||||
|
async def save(self, path: str | Path) -> None:
|
||||||
|
if self._fail:
|
||||||
|
raise RuntimeError("save failed")
|
||||||
|
Path(path).write_bytes(b"attachment")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakePartialMessage:
|
||||||
|
# Lightweight stand-in for Discord partial message references used in replies.
|
||||||
|
def __init__(self, message_id: int) -> None:
|
||||||
|
self.id = message_id
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeChannel:
|
||||||
|
# Channel double that records outbound payloads and typing activity.
|
||||||
|
def __init__(self, channel_id: int = 123) -> None:
|
||||||
|
self.id = channel_id
|
||||||
|
self.sent_payloads: list[dict] = []
|
||||||
|
self.trigger_typing_calls = 0
|
||||||
|
self.typing_enter_hook = None
|
||||||
|
|
||||||
|
async def send(self, **kwargs) -> None:
|
||||||
|
payload = dict(kwargs)
|
||||||
|
if "file" in payload:
|
||||||
|
payload["file_name"] = payload["file"].filename
|
||||||
|
del payload["file"]
|
||||||
|
self.sent_payloads.append(payload)
|
||||||
|
|
||||||
|
def get_partial_message(self, message_id: int) -> _FakePartialMessage:
|
||||||
|
return _FakePartialMessage(message_id)
|
||||||
|
|
||||||
|
def typing(self):
|
||||||
|
channel = self
|
||||||
|
|
||||||
|
class _TypingContext:
|
||||||
|
async def __aenter__(self):
|
||||||
|
channel.trigger_typing_calls += 1
|
||||||
|
if channel.typing_enter_hook is not None:
|
||||||
|
await channel.typing_enter_hook()
|
||||||
|
|
||||||
|
async def __aexit__(self, exc_type, exc, tb):
|
||||||
|
return False
|
||||||
|
|
||||||
|
return _TypingContext()
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeInteractionResponse:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.messages: list[dict] = []
|
||||||
|
self._done = False
|
||||||
|
|
||||||
|
async def send_message(self, content: str, *, ephemeral: bool = False) -> None:
|
||||||
|
self.messages.append({"content": content, "ephemeral": ephemeral})
|
||||||
|
self._done = True
|
||||||
|
|
||||||
|
def is_done(self) -> bool:
|
||||||
|
return self._done
|
||||||
|
|
||||||
|
|
||||||
|
def _make_interaction(
|
||||||
|
*,
|
||||||
|
user_id: int = 123,
|
||||||
|
channel_id: int | None = 456,
|
||||||
|
guild_id: int | None = None,
|
||||||
|
interaction_id: int = 999,
|
||||||
|
):
|
||||||
|
return SimpleNamespace(
|
||||||
|
user=SimpleNamespace(id=user_id),
|
||||||
|
channel_id=channel_id,
|
||||||
|
guild_id=guild_id,
|
||||||
|
id=interaction_id,
|
||||||
|
command=SimpleNamespace(qualified_name="new"),
|
||||||
|
response=_FakeInteractionResponse(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_message(
|
||||||
|
*,
|
||||||
|
author_id: int = 123,
|
||||||
|
author_bot: bool = False,
|
||||||
|
channel_id: int = 456,
|
||||||
|
message_id: int = 789,
|
||||||
|
content: str = "hello",
|
||||||
|
guild_id: int | None = None,
|
||||||
|
mentions: list[object] | None = None,
|
||||||
|
attachments: list[object] | None = None,
|
||||||
|
reply_to: int | None = None,
|
||||||
|
):
|
||||||
|
# Factory for incoming Discord message objects with optional guild/reply/attachments.
|
||||||
|
guild = SimpleNamespace(id=guild_id) if guild_id is not None else None
|
||||||
|
reference = SimpleNamespace(message_id=reply_to) if reply_to is not None else None
|
||||||
|
return SimpleNamespace(
|
||||||
|
author=SimpleNamespace(id=author_id, bot=author_bot),
|
||||||
|
channel=_FakeChannel(channel_id),
|
||||||
|
content=content,
|
||||||
|
guild=guild,
|
||||||
|
mentions=mentions or [],
|
||||||
|
attachments=attachments or [],
|
||||||
|
reference=reference,
|
||||||
|
id=message_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_returns_when_token_missing() -> None:
|
||||||
|
# If no token is configured, startup should no-op and leave channel stopped.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
|
||||||
|
await channel.start()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert channel._client is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_returns_when_discord_dependency_missing(monkeypatch) -> None:
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, token="token", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.DISCORD_AVAILABLE", False)
|
||||||
|
|
||||||
|
await channel.start()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert channel._client is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_handles_client_construction_failure(monkeypatch) -> None:
|
||||||
|
# Construction errors from the Discord client should be swallowed and keep state clean.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, token="token", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _boom(owner, *, intents):
|
||||||
|
raise RuntimeError("bad client")
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.DiscordBotClient", _boom)
|
||||||
|
|
||||||
|
await channel.start()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert channel._client is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_handles_client_start_failure(monkeypatch) -> None:
|
||||||
|
# If client.start fails, the partially created client should be closed and detached.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, token="token", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
|
||||||
|
_FakeDiscordClient.instances.clear()
|
||||||
|
_FakeDiscordClient.start_error = RuntimeError("connect failed")
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.DiscordBotClient", _FakeDiscordClient)
|
||||||
|
|
||||||
|
await channel.start()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert channel._client is None
|
||||||
|
assert _FakeDiscordClient.instances[0].intents.value == channel.config.intents
|
||||||
|
assert _FakeDiscordClient.instances[0].closed is True
|
||||||
|
|
||||||
|
_FakeDiscordClient.start_error = None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stop_is_safe_after_partial_start(monkeypatch) -> None:
|
||||||
|
# stop() should close/discard the client even when startup was only partially completed.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, token="token", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
client = _FakeDiscordClient(channel, intents=None)
|
||||||
|
channel._client = client
|
||||||
|
channel._running = True
|
||||||
|
|
||||||
|
await channel.stop()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert client.closed is True
|
||||||
|
assert channel._client is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_ignores_bot_messages() -> None:
|
||||||
|
# Incoming bot-authored messages must be ignored to prevent feedback loops.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
channel._handle_message = lambda **kwargs: handled.append(kwargs) # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await channel._on_message(_make_message(author_bot=True))
|
||||||
|
|
||||||
|
assert handled == []
|
||||||
|
|
||||||
|
# If inbound handling raises, typing should be stopped for that channel.
|
||||||
|
async def fail_handle(**kwargs) -> None:
|
||||||
|
raise RuntimeError("boom")
|
||||||
|
|
||||||
|
channel._handle_message = fail_handle # type: ignore[method-assign]
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="boom"):
|
||||||
|
await channel._on_message(_make_message(author_id=123, channel_id=456))
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_accepts_allowlisted_dm() -> None:
|
||||||
|
# Allowed direct messages should be forwarded with normalized metadata.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["123"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await channel._on_message(_make_message(author_id=123, channel_id=456, message_id=789))
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["chat_id"] == "456"
|
||||||
|
assert handled[0]["metadata"] == {"message_id": "789", "guild_id": None, "reply_to": None}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_ignores_unmentioned_guild_message() -> None:
|
||||||
|
# With mention-only group policy, guild messages without a bot mention are dropped.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, allow_from=["*"], group_policy="mention"),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
channel._bot_user_id = "999"
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await channel._on_message(_make_message(guild_id=1, content="hello everyone"))
|
||||||
|
|
||||||
|
assert handled == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_accepts_mentioned_guild_message() -> None:
|
||||||
|
# Mentioned guild messages should be accepted and preserve reply threading metadata.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, allow_from=["*"], group_policy="mention"),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
channel._bot_user_id = "999"
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await channel._on_message(
|
||||||
|
_make_message(
|
||||||
|
guild_id=1,
|
||||||
|
content="<@999> hello",
|
||||||
|
mentions=[SimpleNamespace(id=999)],
|
||||||
|
reply_to=321,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["metadata"]["reply_to"] == "321"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_downloads_attachments(tmp_path, monkeypatch) -> None:
|
||||||
|
# Attachment downloads should be saved and referenced in forwarded content/media.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.get_media_dir", lambda _name: tmp_path)
|
||||||
|
|
||||||
|
await channel._on_message(
|
||||||
|
_make_message(
|
||||||
|
attachments=[_FakeAttachment(12, "photo.png")],
|
||||||
|
content="see file",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["media"] == [str(tmp_path / "12_photo.png")]
|
||||||
|
assert "[attachment:" in handled[0]["content"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_marks_failed_attachment_download(tmp_path, monkeypatch) -> None:
|
||||||
|
# Failed attachment downloads should emit a readable placeholder and no media path.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.get_media_dir", lambda _name: tmp_path)
|
||||||
|
|
||||||
|
await channel._on_message(
|
||||||
|
_make_message(
|
||||||
|
attachments=[_FakeAttachment(12, "photo.png", fail=True)],
|
||||||
|
content="",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["media"] == []
|
||||||
|
assert handled[0]["content"] == "[attachment: photo.png - download failed]"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_warns_when_client_not_ready() -> None:
|
||||||
|
# Sending without a running/ready client should be a safe no-op.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
|
||||||
|
await channel.send(OutboundMessage(channel="discord", chat_id="123", content="hello"))
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_skips_when_channel_not_cached() -> None:
|
||||||
|
# Outbound sends should be skipped when the destination channel is not resolvable.
|
||||||
|
owner = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = DiscordBotClient(owner, intents=discord.Intents.none())
|
||||||
|
fetch_calls: list[int] = []
|
||||||
|
|
||||||
|
async def fetch_channel(channel_id: int):
|
||||||
|
fetch_calls.append(channel_id)
|
||||||
|
raise RuntimeError("not found")
|
||||||
|
|
||||||
|
client.fetch_channel = fetch_channel # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await client.send_outbound(OutboundMessage(channel="discord", chat_id="123", content="hello"))
|
||||||
|
|
||||||
|
assert client.get_channel(123) is None
|
||||||
|
assert fetch_calls == [123]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_fetches_channel_when_not_cached() -> None:
|
||||||
|
owner = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = DiscordBotClient(owner, intents=discord.Intents.none())
|
||||||
|
target = _FakeChannel(channel_id=123)
|
||||||
|
|
||||||
|
async def fetch_channel(channel_id: int):
|
||||||
|
return target if channel_id == 123 else None
|
||||||
|
|
||||||
|
client.fetch_channel = fetch_channel # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await client.send_outbound(OutboundMessage(channel="discord", chat_id="123", content="hello"))
|
||||||
|
|
||||||
|
assert target.sent_payloads == [{"content": "hello"}]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_slash_new_forwards_when_user_is_allowlisted() -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["123"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||||
|
interaction = _make_interaction(user_id=123, channel_id=456, interaction_id=321)
|
||||||
|
|
||||||
|
new_cmd = client.tree.get_command("new")
|
||||||
|
assert new_cmd is not None
|
||||||
|
await new_cmd.callback(interaction)
|
||||||
|
|
||||||
|
assert interaction.response.messages == [
|
||||||
|
{"content": "Processing /new...", "ephemeral": True}
|
||||||
|
]
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["content"] == "/new"
|
||||||
|
assert handled[0]["sender_id"] == "123"
|
||||||
|
assert handled[0]["chat_id"] == "456"
|
||||||
|
assert handled[0]["metadata"]["interaction_id"] == "321"
|
||||||
|
assert handled[0]["metadata"]["is_slash_command"] is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_slash_new_is_blocked_for_disallowed_user() -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["999"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||||
|
interaction = _make_interaction(user_id=123, channel_id=456)
|
||||||
|
|
||||||
|
new_cmd = client.tree.get_command("new")
|
||||||
|
assert new_cmd is not None
|
||||||
|
await new_cmd.callback(interaction)
|
||||||
|
|
||||||
|
assert interaction.response.messages == [
|
||||||
|
{"content": "You are not allowed to use this bot.", "ephemeral": True}
|
||||||
|
]
|
||||||
|
assert handled == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("slash_name", ["stop", "restart", "status"])
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_slash_commands_forward_via_handle_message(slash_name: str) -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||||
|
interaction = _make_interaction()
|
||||||
|
interaction.command.qualified_name = slash_name
|
||||||
|
|
||||||
|
cmd = client.tree.get_command(slash_name)
|
||||||
|
assert cmd is not None
|
||||||
|
await cmd.callback(interaction)
|
||||||
|
|
||||||
|
assert interaction.response.messages == [
|
||||||
|
{"content": f"Processing /{slash_name}...", "ephemeral": True}
|
||||||
|
]
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["content"] == f"/{slash_name}"
|
||||||
|
assert handled[0]["metadata"]["is_slash_command"] is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_slash_help_returns_ephemeral_help_text() -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||||
|
interaction = _make_interaction()
|
||||||
|
interaction.command.qualified_name = "help"
|
||||||
|
|
||||||
|
help_cmd = client.tree.get_command("help")
|
||||||
|
assert help_cmd is not None
|
||||||
|
await help_cmd.callback(interaction)
|
||||||
|
|
||||||
|
assert interaction.response.messages == [
|
||||||
|
{"content": build_help_text(), "ephemeral": True}
|
||||||
|
]
|
||||||
|
assert handled == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_client_send_outbound_chunks_text_replies_and_uploads_files(tmp_path) -> None:
|
||||||
|
# Outbound payloads should upload files, attach reply references, and chunk long text.
|
||||||
|
owner = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = DiscordBotClient(owner, intents=discord.Intents.none())
|
||||||
|
target = _FakeChannel(channel_id=123)
|
||||||
|
client.get_channel = lambda channel_id: target if channel_id == 123 else None # type: ignore[method-assign]
|
||||||
|
|
||||||
|
file_path = tmp_path / "demo.txt"
|
||||||
|
file_path.write_text("hi")
|
||||||
|
|
||||||
|
await client.send_outbound(
|
||||||
|
OutboundMessage(
|
||||||
|
channel="discord",
|
||||||
|
chat_id="123",
|
||||||
|
content="a" * 2100,
|
||||||
|
reply_to="55",
|
||||||
|
media=[str(file_path)],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(target.sent_payloads) == 3
|
||||||
|
assert target.sent_payloads[0]["file_name"] == "demo.txt"
|
||||||
|
assert target.sent_payloads[0]["reference"].id == 55
|
||||||
|
assert target.sent_payloads[1]["content"] == "a" * 2000
|
||||||
|
assert target.sent_payloads[2]["content"] == "a" * 100
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_client_send_outbound_reports_failed_attachments_when_no_text(tmp_path) -> None:
|
||||||
|
# If all attachment sends fail and no text exists, emit a failure placeholder message.
|
||||||
|
owner = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = DiscordBotClient(owner, intents=discord.Intents.none())
|
||||||
|
target = _FakeChannel(channel_id=123)
|
||||||
|
client.get_channel = lambda channel_id: target if channel_id == 123 else None # type: ignore[method-assign]
|
||||||
|
|
||||||
|
missing_file = tmp_path / "missing.txt"
|
||||||
|
|
||||||
|
await client.send_outbound(
|
||||||
|
OutboundMessage(
|
||||||
|
channel="discord",
|
||||||
|
chat_id="123",
|
||||||
|
content="",
|
||||||
|
media=[str(missing_file)],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert target.sent_payloads == [{"content": "[attachment: missing.txt - send failed]"}]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_stops_typing_after_send() -> None:
|
||||||
|
# Active typing indicators should be cancelled/cleared after a successful send.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = _FakeDiscordClient(channel, intents=None)
|
||||||
|
channel._client = client
|
||||||
|
channel._running = True
|
||||||
|
|
||||||
|
start = asyncio.Event()
|
||||||
|
release = asyncio.Event()
|
||||||
|
|
||||||
|
async def slow_typing() -> None:
|
||||||
|
start.set()
|
||||||
|
await release.wait()
|
||||||
|
|
||||||
|
typing_channel = _FakeChannel(channel_id=123)
|
||||||
|
typing_channel.typing_enter_hook = slow_typing
|
||||||
|
|
||||||
|
await channel._start_typing(typing_channel)
|
||||||
|
await start.wait()
|
||||||
|
|
||||||
|
await channel.send(OutboundMessage(channel="discord", chat_id="123", content="hello"))
|
||||||
|
release.set()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
|
|
||||||
|
# Progress messages should keep typing active until a final (non-progress) send.
|
||||||
|
start = asyncio.Event()
|
||||||
|
release = asyncio.Event()
|
||||||
|
|
||||||
|
async def slow_typing_progress() -> None:
|
||||||
|
start.set()
|
||||||
|
await release.wait()
|
||||||
|
|
||||||
|
typing_channel = _FakeChannel(channel_id=123)
|
||||||
|
typing_channel.typing_enter_hook = slow_typing_progress
|
||||||
|
|
||||||
|
await channel._start_typing(typing_channel)
|
||||||
|
await start.wait()
|
||||||
|
|
||||||
|
await channel.send(
|
||||||
|
OutboundMessage(
|
||||||
|
channel="discord",
|
||||||
|
chat_id="123",
|
||||||
|
content="progress",
|
||||||
|
metadata={"_progress": True},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "123" in channel._typing_tasks
|
||||||
|
|
||||||
|
await channel.send(OutboundMessage(channel="discord", chat_id="123", content="final"))
|
||||||
|
release.set()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_typing_uses_typing_context_when_trigger_typing_missing() -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
channel._running = True
|
||||||
|
|
||||||
|
entered = asyncio.Event()
|
||||||
|
release = asyncio.Event()
|
||||||
|
|
||||||
|
class _TypingCtx:
|
||||||
|
async def __aenter__(self):
|
||||||
|
entered.set()
|
||||||
|
|
||||||
|
async def __aexit__(self, exc_type, exc, tb):
|
||||||
|
return False
|
||||||
|
|
||||||
|
class _NoTriggerChannel:
|
||||||
|
def __init__(self, channel_id: int = 123) -> None:
|
||||||
|
self.id = channel_id
|
||||||
|
|
||||||
|
def typing(self):
|
||||||
|
async def _waiter():
|
||||||
|
await release.wait()
|
||||||
|
# Hold the loop so task remains active until explicitly stopped.
|
||||||
|
class _Ctx(_TypingCtx):
|
||||||
|
async def __aenter__(self):
|
||||||
|
await super().__aenter__()
|
||||||
|
await _waiter()
|
||||||
|
return _Ctx()
|
||||||
|
|
||||||
|
typing_channel = _NoTriggerChannel(channel_id=123)
|
||||||
|
await channel._start_typing(typing_channel) # type: ignore[arg-type]
|
||||||
|
await entered.wait()
|
||||||
|
|
||||||
|
assert "123" in channel._typing_tasks
|
||||||
|
|
||||||
|
await channel._stop_typing("123")
|
||||||
|
release.set()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
@@ -10,8 +10,8 @@ from nanobot.channels.email import EmailChannel
|
|||||||
from nanobot.channels.email import EmailConfig
|
from nanobot.channels.email import EmailConfig
|
||||||
|
|
||||||
|
|
||||||
def _make_config() -> EmailConfig:
|
def _make_config(**overrides) -> EmailConfig:
|
||||||
return EmailConfig(
|
defaults = dict(
|
||||||
enabled=True,
|
enabled=True,
|
||||||
consent_granted=True,
|
consent_granted=True,
|
||||||
imap_host="imap.example.com",
|
imap_host="imap.example.com",
|
||||||
@@ -23,19 +23,27 @@ def _make_config() -> EmailConfig:
|
|||||||
smtp_username="bot@example.com",
|
smtp_username="bot@example.com",
|
||||||
smtp_password="secret",
|
smtp_password="secret",
|
||||||
mark_seen=True,
|
mark_seen=True,
|
||||||
|
# Disable auth verification by default so existing tests are unaffected
|
||||||
|
verify_dkim=False,
|
||||||
|
verify_spf=False,
|
||||||
)
|
)
|
||||||
|
defaults.update(overrides)
|
||||||
|
return EmailConfig(**defaults)
|
||||||
|
|
||||||
|
|
||||||
def _make_raw_email(
|
def _make_raw_email(
|
||||||
from_addr: str = "alice@example.com",
|
from_addr: str = "alice@example.com",
|
||||||
subject: str = "Hello",
|
subject: str = "Hello",
|
||||||
body: str = "This is the body.",
|
body: str = "This is the body.",
|
||||||
|
auth_results: str | None = None,
|
||||||
) -> bytes:
|
) -> bytes:
|
||||||
msg = EmailMessage()
|
msg = EmailMessage()
|
||||||
msg["From"] = from_addr
|
msg["From"] = from_addr
|
||||||
msg["To"] = "bot@example.com"
|
msg["To"] = "bot@example.com"
|
||||||
msg["Subject"] = subject
|
msg["Subject"] = subject
|
||||||
msg["Message-ID"] = "<m1@example.com>"
|
msg["Message-ID"] = "<m1@example.com>"
|
||||||
|
if auth_results:
|
||||||
|
msg["Authentication-Results"] = auth_results
|
||||||
msg.set_content(body)
|
msg.set_content(body)
|
||||||
return msg.as_bytes()
|
return msg.as_bytes()
|
||||||
|
|
||||||
@@ -481,3 +489,164 @@ def test_fetch_messages_between_dates_uses_imap_since_before_without_mark_seen(m
|
|||||||
assert fake.search_args is not None
|
assert fake.search_args is not None
|
||||||
assert fake.search_args[1:] == ("SINCE", "06-Feb-2026", "BEFORE", "07-Feb-2026")
|
assert fake.search_args[1:] == ("SINCE", "06-Feb-2026", "BEFORE", "07-Feb-2026")
|
||||||
assert fake.store_calls == []
|
assert fake.store_calls == []
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Security: Anti-spoofing tests for Authentication-Results verification
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _make_fake_imap(raw: bytes):
|
||||||
|
"""Return a FakeIMAP class pre-loaded with the given raw email."""
|
||||||
|
class FakeIMAP:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.store_calls: list[tuple[bytes, str, str]] = []
|
||||||
|
|
||||||
|
def login(self, _user: str, _pw: str):
|
||||||
|
return "OK", [b"logged in"]
|
||||||
|
|
||||||
|
def select(self, _mailbox: str):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def search(self, *_args):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
|
return "OK", [(b"1 (UID 500 BODY[] {200})", raw), b")"]
|
||||||
|
|
||||||
|
def store(self, imap_id: bytes, op: str, flags: str):
|
||||||
|
self.store_calls.append((imap_id, op, flags))
|
||||||
|
return "OK", [b""]
|
||||||
|
|
||||||
|
def logout(self):
|
||||||
|
return "BYE", [b""]
|
||||||
|
|
||||||
|
return FakeIMAP()
|
||||||
|
|
||||||
|
|
||||||
|
def test_spoofed_email_rejected_when_verify_enabled(monkeypatch) -> None:
|
||||||
|
"""An email without Authentication-Results should be rejected when verify_dkim=True."""
|
||||||
|
raw = _make_raw_email(subject="Spoofed", body="Malicious payload")
|
||||||
|
fake = _make_fake_imap(raw)
|
||||||
|
monkeypatch.setattr("nanobot.channels.email.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
||||||
|
|
||||||
|
cfg = _make_config(verify_dkim=True, verify_spf=True)
|
||||||
|
channel = EmailChannel(cfg, MessageBus())
|
||||||
|
items = channel._fetch_new_messages()
|
||||||
|
|
||||||
|
assert len(items) == 0, "Spoofed email without auth headers should be rejected"
|
||||||
|
|
||||||
|
|
||||||
|
def test_email_with_valid_auth_results_accepted(monkeypatch) -> None:
|
||||||
|
"""An email with spf=pass and dkim=pass should be accepted."""
|
||||||
|
raw = _make_raw_email(
|
||||||
|
subject="Legit",
|
||||||
|
body="Hello from verified sender",
|
||||||
|
auth_results="mx.example.com; spf=pass smtp.mailfrom=alice@example.com; dkim=pass header.d=example.com",
|
||||||
|
)
|
||||||
|
fake = _make_fake_imap(raw)
|
||||||
|
monkeypatch.setattr("nanobot.channels.email.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
||||||
|
|
||||||
|
cfg = _make_config(verify_dkim=True, verify_spf=True)
|
||||||
|
channel = EmailChannel(cfg, MessageBus())
|
||||||
|
items = channel._fetch_new_messages()
|
||||||
|
|
||||||
|
assert len(items) == 1
|
||||||
|
assert items[0]["sender"] == "alice@example.com"
|
||||||
|
assert items[0]["subject"] == "Legit"
|
||||||
|
|
||||||
|
|
||||||
|
def test_email_with_partial_auth_rejected(monkeypatch) -> None:
|
||||||
|
"""An email with only spf=pass but no dkim=pass should be rejected when verify_dkim=True."""
|
||||||
|
raw = _make_raw_email(
|
||||||
|
subject="Partial",
|
||||||
|
body="Only SPF passes",
|
||||||
|
auth_results="mx.example.com; spf=pass smtp.mailfrom=alice@example.com; dkim=fail",
|
||||||
|
)
|
||||||
|
fake = _make_fake_imap(raw)
|
||||||
|
monkeypatch.setattr("nanobot.channels.email.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
||||||
|
|
||||||
|
cfg = _make_config(verify_dkim=True, verify_spf=True)
|
||||||
|
channel = EmailChannel(cfg, MessageBus())
|
||||||
|
items = channel._fetch_new_messages()
|
||||||
|
|
||||||
|
assert len(items) == 0, "Email with dkim=fail should be rejected"
|
||||||
|
|
||||||
|
|
||||||
|
def test_backward_compat_verify_disabled(monkeypatch) -> None:
|
||||||
|
"""When verify_dkim=False and verify_spf=False, emails without auth headers are accepted."""
|
||||||
|
raw = _make_raw_email(subject="NoAuth", body="No auth headers present")
|
||||||
|
fake = _make_fake_imap(raw)
|
||||||
|
monkeypatch.setattr("nanobot.channels.email.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
||||||
|
|
||||||
|
cfg = _make_config(verify_dkim=False, verify_spf=False)
|
||||||
|
channel = EmailChannel(cfg, MessageBus())
|
||||||
|
items = channel._fetch_new_messages()
|
||||||
|
|
||||||
|
assert len(items) == 1, "With verification disabled, emails should be accepted as before"
|
||||||
|
|
||||||
|
|
||||||
|
def test_email_content_tagged_with_email_context(monkeypatch) -> None:
|
||||||
|
"""Email content should be prefixed with [EMAIL-CONTEXT] for LLM isolation."""
|
||||||
|
raw = _make_raw_email(subject="Tagged", body="Check the tag")
|
||||||
|
fake = _make_fake_imap(raw)
|
||||||
|
monkeypatch.setattr("nanobot.channels.email.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
||||||
|
|
||||||
|
cfg = _make_config(verify_dkim=False, verify_spf=False)
|
||||||
|
channel = EmailChannel(cfg, MessageBus())
|
||||||
|
items = channel._fetch_new_messages()
|
||||||
|
|
||||||
|
assert len(items) == 1
|
||||||
|
assert items[0]["content"].startswith("[EMAIL-CONTEXT]"), (
|
||||||
|
"Email content must be tagged with [EMAIL-CONTEXT]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_check_authentication_results_method() -> None:
|
||||||
|
"""Unit test for the _check_authentication_results static method."""
|
||||||
|
from email.parser import BytesParser
|
||||||
|
from email import policy
|
||||||
|
|
||||||
|
# No Authentication-Results header
|
||||||
|
msg_no_auth = EmailMessage()
|
||||||
|
msg_no_auth["From"] = "alice@example.com"
|
||||||
|
msg_no_auth.set_content("test")
|
||||||
|
parsed = BytesParser(policy=policy.default).parsebytes(msg_no_auth.as_bytes())
|
||||||
|
spf, dkim = EmailChannel._check_authentication_results(parsed)
|
||||||
|
assert spf is False
|
||||||
|
assert dkim is False
|
||||||
|
|
||||||
|
# Both pass
|
||||||
|
msg_both = EmailMessage()
|
||||||
|
msg_both["From"] = "alice@example.com"
|
||||||
|
msg_both["Authentication-Results"] = (
|
||||||
|
"mx.google.com; spf=pass smtp.mailfrom=example.com; dkim=pass header.d=example.com"
|
||||||
|
)
|
||||||
|
msg_both.set_content("test")
|
||||||
|
parsed = BytesParser(policy=policy.default).parsebytes(msg_both.as_bytes())
|
||||||
|
spf, dkim = EmailChannel._check_authentication_results(parsed)
|
||||||
|
assert spf is True
|
||||||
|
assert dkim is True
|
||||||
|
|
||||||
|
# SPF pass, DKIM fail
|
||||||
|
msg_spf_only = EmailMessage()
|
||||||
|
msg_spf_only["From"] = "alice@example.com"
|
||||||
|
msg_spf_only["Authentication-Results"] = (
|
||||||
|
"mx.google.com; spf=pass smtp.mailfrom=example.com; dkim=fail"
|
||||||
|
)
|
||||||
|
msg_spf_only.set_content("test")
|
||||||
|
parsed = BytesParser(policy=policy.default).parsebytes(msg_spf_only.as_bytes())
|
||||||
|
spf, dkim = EmailChannel._check_authentication_results(parsed)
|
||||||
|
assert spf is True
|
||||||
|
assert dkim is False
|
||||||
|
|
||||||
|
# DKIM pass, SPF fail
|
||||||
|
msg_dkim_only = EmailMessage()
|
||||||
|
msg_dkim_only["From"] = "alice@example.com"
|
||||||
|
msg_dkim_only["Authentication-Results"] = (
|
||||||
|
"mx.google.com; spf=fail smtp.mailfrom=example.com; dkim=pass header.d=example.com"
|
||||||
|
)
|
||||||
|
msg_dkim_only.set_content("test")
|
||||||
|
parsed = BytesParser(policy=policy.default).parsebytes(msg_dkim_only.as_bytes())
|
||||||
|
spf, dkim = EmailChannel._check_authentication_results(parsed)
|
||||||
|
assert spf is False
|
||||||
|
assert dkim is True
|
||||||
|
|||||||
@@ -0,0 +1,258 @@
|
|||||||
|
"""Tests for Feishu streaming (send_delta) via CardKit streaming API."""
|
||||||
|
import time
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.channels.feishu import FeishuChannel, FeishuConfig, _FeishuStreamBuf
|
||||||
|
|
||||||
|
|
||||||
|
def _make_channel(streaming: bool = True) -> FeishuChannel:
|
||||||
|
config = FeishuConfig(
|
||||||
|
enabled=True,
|
||||||
|
app_id="cli_test",
|
||||||
|
app_secret="secret",
|
||||||
|
allow_from=["*"],
|
||||||
|
streaming=streaming,
|
||||||
|
)
|
||||||
|
ch = FeishuChannel(config, MessageBus())
|
||||||
|
ch._client = MagicMock()
|
||||||
|
ch._loop = None
|
||||||
|
return ch
|
||||||
|
|
||||||
|
|
||||||
|
def _mock_create_card_response(card_id: str = "card_stream_001"):
|
||||||
|
resp = MagicMock()
|
||||||
|
resp.success.return_value = True
|
||||||
|
resp.data = SimpleNamespace(card_id=card_id)
|
||||||
|
return resp
|
||||||
|
|
||||||
|
|
||||||
|
def _mock_send_response(message_id: str = "om_stream_001"):
|
||||||
|
resp = MagicMock()
|
||||||
|
resp.success.return_value = True
|
||||||
|
resp.data = SimpleNamespace(message_id=message_id)
|
||||||
|
return resp
|
||||||
|
|
||||||
|
|
||||||
|
def _mock_content_response(success: bool = True):
|
||||||
|
resp = MagicMock()
|
||||||
|
resp.success.return_value = success
|
||||||
|
resp.code = 0 if success else 99999
|
||||||
|
resp.msg = "ok" if success else "error"
|
||||||
|
return resp
|
||||||
|
|
||||||
|
|
||||||
|
class TestFeishuStreamingConfig:
|
||||||
|
def test_streaming_default_true(self):
|
||||||
|
assert FeishuConfig().streaming is True
|
||||||
|
|
||||||
|
def test_supports_streaming_when_enabled(self):
|
||||||
|
ch = _make_channel(streaming=True)
|
||||||
|
assert ch.supports_streaming is True
|
||||||
|
|
||||||
|
def test_supports_streaming_disabled(self):
|
||||||
|
ch = _make_channel(streaming=False)
|
||||||
|
assert ch.supports_streaming is False
|
||||||
|
|
||||||
|
|
||||||
|
class TestCreateStreamingCard:
|
||||||
|
def test_returns_card_id_on_success(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card.create.return_value = _mock_create_card_response("card_123")
|
||||||
|
ch._client.im.v1.message.create.return_value = _mock_send_response()
|
||||||
|
result = ch._create_streaming_card_sync("chat_id", "oc_chat1")
|
||||||
|
assert result == "card_123"
|
||||||
|
ch._client.cardkit.v1.card.create.assert_called_once()
|
||||||
|
ch._client.im.v1.message.create.assert_called_once()
|
||||||
|
|
||||||
|
def test_returns_none_on_failure(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
resp = MagicMock()
|
||||||
|
resp.success.return_value = False
|
||||||
|
resp.code = 99999
|
||||||
|
resp.msg = "error"
|
||||||
|
ch._client.cardkit.v1.card.create.return_value = resp
|
||||||
|
assert ch._create_streaming_card_sync("chat_id", "oc_chat1") is None
|
||||||
|
|
||||||
|
def test_returns_none_on_exception(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card.create.side_effect = RuntimeError("network")
|
||||||
|
assert ch._create_streaming_card_sync("chat_id", "oc_chat1") is None
|
||||||
|
|
||||||
|
def test_returns_none_when_card_send_fails(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card.create.return_value = _mock_create_card_response("card_123")
|
||||||
|
resp = MagicMock()
|
||||||
|
resp.success.return_value = False
|
||||||
|
resp.code = 99999
|
||||||
|
resp.msg = "error"
|
||||||
|
resp.get_log_id.return_value = "log1"
|
||||||
|
ch._client.im.v1.message.create.return_value = resp
|
||||||
|
assert ch._create_streaming_card_sync("chat_id", "oc_chat1") is None
|
||||||
|
|
||||||
|
|
||||||
|
class TestCloseStreamingMode:
|
||||||
|
def test_returns_true_on_success(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card.settings.return_value = _mock_content_response(True)
|
||||||
|
assert ch._close_streaming_mode_sync("card_1", 10) is True
|
||||||
|
|
||||||
|
def test_returns_false_on_failure(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card.settings.return_value = _mock_content_response(False)
|
||||||
|
assert ch._close_streaming_mode_sync("card_1", 10) is False
|
||||||
|
|
||||||
|
def test_returns_false_on_exception(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card.settings.side_effect = RuntimeError("err")
|
||||||
|
assert ch._close_streaming_mode_sync("card_1", 10) is False
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamUpdateText:
|
||||||
|
def test_returns_true_on_success(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response(True)
|
||||||
|
assert ch._stream_update_text_sync("card_1", "hello", 1) is True
|
||||||
|
|
||||||
|
def test_returns_false_on_failure(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response(False)
|
||||||
|
assert ch._stream_update_text_sync("card_1", "hello", 1) is False
|
||||||
|
|
||||||
|
def test_returns_false_on_exception(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card_element.content.side_effect = RuntimeError("err")
|
||||||
|
assert ch._stream_update_text_sync("card_1", "hello", 1) is False
|
||||||
|
|
||||||
|
|
||||||
|
class TestSendDelta:
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_first_delta_creates_card_and_sends(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.cardkit.v1.card.create.return_value = _mock_create_card_response("card_new")
|
||||||
|
ch._client.im.v1.message.create.return_value = _mock_send_response("om_new")
|
||||||
|
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response()
|
||||||
|
|
||||||
|
await ch.send_delta("oc_chat1", "Hello ")
|
||||||
|
|
||||||
|
assert "oc_chat1" in ch._stream_bufs
|
||||||
|
buf = ch._stream_bufs["oc_chat1"]
|
||||||
|
assert buf.text == "Hello "
|
||||||
|
assert buf.card_id == "card_new"
|
||||||
|
assert buf.sequence == 1
|
||||||
|
ch._client.cardkit.v1.card.create.assert_called_once()
|
||||||
|
ch._client.im.v1.message.create.assert_called_once()
|
||||||
|
ch._client.cardkit.v1.card_element.content.assert_called_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_second_delta_within_interval_skips_update(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
buf = _FeishuStreamBuf(text="Hello ", card_id="card_1", sequence=1, last_edit=time.monotonic())
|
||||||
|
ch._stream_bufs["oc_chat1"] = buf
|
||||||
|
|
||||||
|
await ch.send_delta("oc_chat1", "world")
|
||||||
|
|
||||||
|
assert buf.text == "Hello world"
|
||||||
|
ch._client.cardkit.v1.card_element.content.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_delta_after_interval_updates_text(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
buf = _FeishuStreamBuf(text="Hello ", card_id="card_1", sequence=1, last_edit=time.monotonic() - 1.0)
|
||||||
|
ch._stream_bufs["oc_chat1"] = buf
|
||||||
|
|
||||||
|
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response()
|
||||||
|
await ch.send_delta("oc_chat1", "world")
|
||||||
|
|
||||||
|
assert buf.text == "Hello world"
|
||||||
|
assert buf.sequence == 2
|
||||||
|
ch._client.cardkit.v1.card_element.content.assert_called_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stream_end_sends_final_update(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._stream_bufs["oc_chat1"] = _FeishuStreamBuf(
|
||||||
|
text="Final content", card_id="card_1", sequence=3, last_edit=0.0,
|
||||||
|
)
|
||||||
|
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response()
|
||||||
|
ch._client.cardkit.v1.card.settings.return_value = _mock_content_response()
|
||||||
|
|
||||||
|
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||||
|
|
||||||
|
assert "oc_chat1" not in ch._stream_bufs
|
||||||
|
ch._client.cardkit.v1.card_element.content.assert_called_once()
|
||||||
|
ch._client.cardkit.v1.card.settings.assert_called_once()
|
||||||
|
settings_call = ch._client.cardkit.v1.card.settings.call_args[0][0]
|
||||||
|
assert settings_call.body.sequence == 5 # after final content seq 4
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stream_end_fallback_when_no_card_id(self):
|
||||||
|
"""If card creation failed, stream_end falls back to a plain card message."""
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._stream_bufs["oc_chat1"] = _FeishuStreamBuf(
|
||||||
|
text="Fallback content", card_id=None, sequence=0, last_edit=0.0,
|
||||||
|
)
|
||||||
|
ch._client.im.v1.message.create.return_value = _mock_send_response("om_fb")
|
||||||
|
|
||||||
|
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||||
|
|
||||||
|
assert "oc_chat1" not in ch._stream_bufs
|
||||||
|
ch._client.cardkit.v1.card_element.content.assert_not_called()
|
||||||
|
ch._client.im.v1.message.create.assert_called_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stream_end_without_buf_is_noop(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||||
|
ch._client.cardkit.v1.card_element.content.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_empty_delta_skips_send(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
await ch.send_delta("oc_chat1", " ")
|
||||||
|
|
||||||
|
assert "oc_chat1" in ch._stream_bufs
|
||||||
|
ch._client.cardkit.v1.card.create.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_no_client_returns_early(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client = None
|
||||||
|
await ch.send_delta("oc_chat1", "text")
|
||||||
|
assert "oc_chat1" not in ch._stream_bufs
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_sequence_increments_correctly(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
buf = _FeishuStreamBuf(text="a", card_id="card_1", sequence=5, last_edit=0.0)
|
||||||
|
ch._stream_bufs["oc_chat1"] = buf
|
||||||
|
|
||||||
|
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response()
|
||||||
|
await ch.send_delta("oc_chat1", "b")
|
||||||
|
assert buf.sequence == 6
|
||||||
|
|
||||||
|
buf.last_edit = 0.0 # reset to bypass throttle
|
||||||
|
await ch.send_delta("oc_chat1", "c")
|
||||||
|
assert buf.sequence == 7
|
||||||
|
|
||||||
|
|
||||||
|
class TestSendMessageReturnsId:
|
||||||
|
def test_returns_message_id_on_success(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
ch._client.im.v1.message.create.return_value = _mock_send_response("om_abc")
|
||||||
|
result = ch._send_message_sync("chat_id", "oc_chat1", "text", '{"text":"hi"}')
|
||||||
|
assert result == "om_abc"
|
||||||
|
|
||||||
|
def test_returns_none_on_failure(self):
|
||||||
|
ch = _make_channel()
|
||||||
|
resp = MagicMock()
|
||||||
|
resp.success.return_value = False
|
||||||
|
resp.code = 99999
|
||||||
|
resp.msg = "error"
|
||||||
|
resp.get_log_id.return_value = "log1"
|
||||||
|
ch._client.im.v1.message.create.return_value = resp
|
||||||
|
result = ch._send_message_sync("chat_id", "oc_chat1", "text", '{"text":"hi"}')
|
||||||
|
assert result is None
|
||||||
@@ -3,6 +3,9 @@ from pathlib import Path
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from nio import RoomSendResponse
|
||||||
|
|
||||||
|
from nanobot.channels.matrix import _build_matrix_text_content
|
||||||
|
|
||||||
# Check optional matrix dependencies before importing
|
# Check optional matrix dependencies before importing
|
||||||
try:
|
try:
|
||||||
@@ -65,6 +68,7 @@ class _FakeAsyncClient:
|
|||||||
self.raise_on_send = False
|
self.raise_on_send = False
|
||||||
self.raise_on_typing = False
|
self.raise_on_typing = False
|
||||||
self.raise_on_upload = False
|
self.raise_on_upload = False
|
||||||
|
self.room_send_response: RoomSendResponse | None = RoomSendResponse(event_id="", room_id="")
|
||||||
|
|
||||||
def add_event_callback(self, callback, event_type) -> None:
|
def add_event_callback(self, callback, event_type) -> None:
|
||||||
self.callbacks.append((callback, event_type))
|
self.callbacks.append((callback, event_type))
|
||||||
@@ -87,7 +91,7 @@ class _FakeAsyncClient:
|
|||||||
message_type: str,
|
message_type: str,
|
||||||
content: dict[str, object],
|
content: dict[str, object],
|
||||||
ignore_unverified_devices: object = _ROOM_SEND_UNSET,
|
ignore_unverified_devices: object = _ROOM_SEND_UNSET,
|
||||||
) -> None:
|
) -> RoomSendResponse:
|
||||||
call: dict[str, object] = {
|
call: dict[str, object] = {
|
||||||
"room_id": room_id,
|
"room_id": room_id,
|
||||||
"message_type": message_type,
|
"message_type": message_type,
|
||||||
@@ -98,6 +102,7 @@ class _FakeAsyncClient:
|
|||||||
self.room_send_calls.append(call)
|
self.room_send_calls.append(call)
|
||||||
if self.raise_on_send:
|
if self.raise_on_send:
|
||||||
raise RuntimeError("send failed")
|
raise RuntimeError("send failed")
|
||||||
|
return self.room_send_response
|
||||||
|
|
||||||
async def room_typing(
|
async def room_typing(
|
||||||
self,
|
self,
|
||||||
@@ -520,6 +525,7 @@ async def test_on_message_room_mention_requires_opt_in() -> None:
|
|||||||
source={"content": {"m.mentions": {"room": True}}},
|
source={"content": {"m.mentions": {"room": True}}},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
channel.config.allow_room_mentions = False
|
||||||
await channel._on_message(room, room_mention_event)
|
await channel._on_message(room, room_mention_event)
|
||||||
assert handled == []
|
assert handled == []
|
||||||
assert client.typing_calls == []
|
assert client.typing_calls == []
|
||||||
@@ -1322,3 +1328,302 @@ async def test_send_keeps_plaintext_only_for_plain_text() -> None:
|
|||||||
"body": text,
|
"body": text,
|
||||||
"m.mentions": {},
|
"m.mentions": {},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_basic_text() -> None:
|
||||||
|
"""Test basic text content without HTML formatting."""
|
||||||
|
result = _build_matrix_text_content("Hello, World!")
|
||||||
|
expected = {
|
||||||
|
"msgtype": "m.text",
|
||||||
|
"body": "Hello, World!",
|
||||||
|
"m.mentions": {}
|
||||||
|
}
|
||||||
|
assert expected == result
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_with_markdown() -> None:
|
||||||
|
"""Test text content with markdown that renders to HTML."""
|
||||||
|
text = "*Hello* **World**"
|
||||||
|
result = _build_matrix_text_content(text)
|
||||||
|
assert "msgtype" in result
|
||||||
|
assert "body" in result
|
||||||
|
assert result["body"] == text
|
||||||
|
assert "format" in result
|
||||||
|
assert result["format"] == "org.matrix.custom.html"
|
||||||
|
assert "formatted_body" in result
|
||||||
|
assert isinstance(result["formatted_body"], str)
|
||||||
|
assert len(result["formatted_body"]) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_with_event_id() -> None:
|
||||||
|
"""Test text content with event_id for message replacement."""
|
||||||
|
event_id = "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
result = _build_matrix_text_content("Updated message", event_id)
|
||||||
|
assert "msgtype" in result
|
||||||
|
assert "body" in result
|
||||||
|
assert result["m.new_content"]
|
||||||
|
assert result["m.new_content"]["body"] == "Updated message"
|
||||||
|
assert result["m.relates_to"]["rel_type"] == "m.replace"
|
||||||
|
assert result["m.relates_to"]["event_id"] == event_id
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_with_event_id_preserves_thread_relation() -> None:
|
||||||
|
"""Thread relations for edits should stay inside m.new_content."""
|
||||||
|
relates_to = {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": "$root1",
|
||||||
|
"m.in_reply_to": {"event_id": "$reply1"},
|
||||||
|
"is_falling_back": True,
|
||||||
|
}
|
||||||
|
result = _build_matrix_text_content("Updated message", "event-1", relates_to)
|
||||||
|
|
||||||
|
assert result["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "event-1",
|
||||||
|
}
|
||||||
|
assert result["m.new_content"]["m.relates_to"] == relates_to
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_no_event_id() -> None:
|
||||||
|
"""Test that when event_id is not provided, no extra properties are added."""
|
||||||
|
result = _build_matrix_text_content("Regular message")
|
||||||
|
|
||||||
|
# Basic required properties should be present
|
||||||
|
assert "msgtype" in result
|
||||||
|
assert "body" in result
|
||||||
|
assert result["body"] == "Regular message"
|
||||||
|
|
||||||
|
# Extra properties for replacement should NOT be present
|
||||||
|
assert "m.relates_to" not in result
|
||||||
|
assert "m.new_content" not in result
|
||||||
|
assert "format" not in result
|
||||||
|
assert "formatted_body" not in result
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_plain_text_no_html() -> None:
|
||||||
|
"""Test plain text that should not include HTML formatting."""
|
||||||
|
result = _build_matrix_text_content("Simple plain text")
|
||||||
|
assert "msgtype" in result
|
||||||
|
assert "body" in result
|
||||||
|
assert "format" not in result
|
||||||
|
assert "formatted_body" not in result
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_room_content_returns_room_send_response():
|
||||||
|
"""Test that _send_room_content returns the response from client.room_send."""
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
room_id = "!test_room:matrix.org"
|
||||||
|
content = {"msgtype": "m.text", "body": "Hello World"}
|
||||||
|
|
||||||
|
result = await channel._send_room_content(room_id, content)
|
||||||
|
|
||||||
|
assert result is client.room_send_response
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_creates_stream_buffer_and_sends_initial_message() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello")
|
||||||
|
|
||||||
|
assert "!room:matrix.org" in channel._stream_bufs
|
||||||
|
buf = channel._stream_bufs["!room:matrix.org"]
|
||||||
|
assert buf.text == "Hello"
|
||||||
|
assert buf.event_id == "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
assert client.room_send_calls[0]["content"]["body"] == "Hello"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_appends_without_sending_before_edit_interval(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
|
||||||
|
now = 100.0
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: now)
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello")
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", " world")
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
|
||||||
|
buf = channel._stream_bufs["!room:matrix.org"]
|
||||||
|
assert buf.text == "Hello world"
|
||||||
|
assert buf.event_id == "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_edits_again_after_interval(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
|
||||||
|
times = [100.0, 102.0, 104.0, 106.0, 108.0]
|
||||||
|
times.reverse()
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: times and times.pop())
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello")
|
||||||
|
await channel.send_delta("!room:matrix.org", " world")
|
||||||
|
|
||||||
|
assert len(client.room_send_calls) == 2
|
||||||
|
first_content = client.room_send_calls[0]["content"]
|
||||||
|
second_content = client.room_send_calls[1]["content"]
|
||||||
|
|
||||||
|
assert "body" in first_content
|
||||||
|
assert first_content["body"] == "Hello"
|
||||||
|
assert "m.relates_to" not in first_content
|
||||||
|
|
||||||
|
assert "body" in second_content
|
||||||
|
assert "m.relates_to" in second_content
|
||||||
|
assert second_content["body"] == "Hello world"
|
||||||
|
assert second_content["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_stream_end_replaces_existing_message() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
channel._stream_bufs["!room:matrix.org"] = matrix_module._StreamBuf(
|
||||||
|
text="Final text",
|
||||||
|
event_id="event-1",
|
||||||
|
last_edit=100.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True})
|
||||||
|
|
||||||
|
assert "!room:matrix.org" not in channel._stream_bufs
|
||||||
|
assert client.typing_calls[-1] == ("!room:matrix.org", False, TYPING_NOTICE_TIMEOUT_MS)
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
assert client.room_send_calls[0]["content"]["body"] == "Final text"
|
||||||
|
assert client.room_send_calls[0]["content"]["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "event-1",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_starts_threaded_stream_inside_thread() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "event-1"
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
"thread_root_event_id": "$root1",
|
||||||
|
"thread_reply_to_event_id": "$reply1",
|
||||||
|
}
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello", metadata)
|
||||||
|
|
||||||
|
assert client.room_send_calls[0]["content"]["m.relates_to"] == {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": "$root1",
|
||||||
|
"m.in_reply_to": {"event_id": "$reply1"},
|
||||||
|
"is_falling_back": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_threaded_edit_keeps_replace_and_thread_relation(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "event-1"
|
||||||
|
|
||||||
|
times = [100.0, 102.0, 104.0]
|
||||||
|
times.reverse()
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: times and times.pop())
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
"thread_root_event_id": "$root1",
|
||||||
|
"thread_reply_to_event_id": "$reply1",
|
||||||
|
}
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello", metadata)
|
||||||
|
await channel.send_delta("!room:matrix.org", " world", metadata)
|
||||||
|
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True, **metadata})
|
||||||
|
|
||||||
|
edit_content = client.room_send_calls[1]["content"]
|
||||||
|
final_content = client.room_send_calls[2]["content"]
|
||||||
|
|
||||||
|
assert edit_content["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "event-1",
|
||||||
|
}
|
||||||
|
assert edit_content["m.new_content"]["m.relates_to"] == {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": "$root1",
|
||||||
|
"m.in_reply_to": {"event_id": "$reply1"},
|
||||||
|
"is_falling_back": True,
|
||||||
|
}
|
||||||
|
assert final_content["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "event-1",
|
||||||
|
}
|
||||||
|
assert final_content["m.new_content"]["m.relates_to"] == {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": "$root1",
|
||||||
|
"m.in_reply_to": {"event_id": "$reply1"},
|
||||||
|
"is_falling_back": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_stream_end_noop_when_buffer_missing() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True})
|
||||||
|
|
||||||
|
assert client.room_send_calls == []
|
||||||
|
assert client.typing_calls == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_on_error_stops_typing(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
client.raise_on_send = True
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
now = 100.0
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: now)
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello", {"room_id": "!room:matrix.org"})
|
||||||
|
|
||||||
|
assert "!room:matrix.org" in channel._stream_bufs
|
||||||
|
assert channel._stream_bufs["!room:matrix.org"].text == "Hello"
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
|
||||||
|
assert len(client.typing_calls) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_ignores_whitespace_only_delta(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
now = 100.0
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: now)
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", " ")
|
||||||
|
|
||||||
|
assert "!room:matrix.org" in channel._stream_bufs
|
||||||
|
assert channel._stream_bufs["!room:matrix.org"].text == " "
|
||||||
|
assert client.room_send_calls == []
|
||||||
@@ -13,7 +13,7 @@ except ImportError:
|
|||||||
|
|
||||||
from nanobot.bus.events import OutboundMessage
|
from nanobot.bus.events import OutboundMessage
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.channels.telegram import TELEGRAM_REPLY_CONTEXT_MAX_LEN, TelegramChannel
|
from nanobot.channels.telegram import TELEGRAM_REPLY_CONTEXT_MAX_LEN, TelegramChannel, _StreamBuf
|
||||||
from nanobot.channels.telegram import TelegramConfig
|
from nanobot.channels.telegram import TelegramConfig
|
||||||
|
|
||||||
|
|
||||||
@@ -50,8 +50,9 @@ class _FakeBot:
|
|||||||
async def set_my_commands(self, commands) -> None:
|
async def set_my_commands(self, commands) -> None:
|
||||||
self.commands = commands
|
self.commands = commands
|
||||||
|
|
||||||
async def send_message(self, **kwargs) -> None:
|
async def send_message(self, **kwargs):
|
||||||
self.sent_messages.append(kwargs)
|
self.sent_messages.append(kwargs)
|
||||||
|
return SimpleNamespace(message_id=len(self.sent_messages))
|
||||||
|
|
||||||
async def send_photo(self, **kwargs) -> None:
|
async def send_photo(self, **kwargs) -> None:
|
||||||
self.sent_media.append({"kind": "photo", **kwargs})
|
self.sent_media.append({"kind": "photo", **kwargs})
|
||||||
@@ -271,13 +272,132 @@ async def test_send_text_gives_up_after_max_retries() -> None:
|
|||||||
orig_delay = tg_mod._SEND_RETRY_BASE_DELAY
|
orig_delay = tg_mod._SEND_RETRY_BASE_DELAY
|
||||||
tg_mod._SEND_RETRY_BASE_DELAY = 0.01
|
tg_mod._SEND_RETRY_BASE_DELAY = 0.01
|
||||||
try:
|
try:
|
||||||
await channel._send_text(123, "hello", None, {})
|
with pytest.raises(TimedOut):
|
||||||
|
await channel._send_text(123, "hello", None, {})
|
||||||
finally:
|
finally:
|
||||||
tg_mod._SEND_RETRY_BASE_DELAY = orig_delay
|
tg_mod._SEND_RETRY_BASE_DELAY = orig_delay
|
||||||
|
|
||||||
assert channel._app.bot.sent_messages == []
|
assert channel._app.bot.sent_messages == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_error_logs_network_issues_as_warning(monkeypatch) -> None:
|
||||||
|
from telegram.error import NetworkError
|
||||||
|
|
||||||
|
channel = TelegramChannel(
|
||||||
|
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
recorded: list[tuple[str, str]] = []
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.channels.telegram.logger.warning",
|
||||||
|
lambda message, error: recorded.append(("warning", message.format(error))),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.channels.telegram.logger.error",
|
||||||
|
lambda message, error: recorded.append(("error", message.format(error))),
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._on_error(object(), SimpleNamespace(error=NetworkError("proxy disconnected")))
|
||||||
|
|
||||||
|
assert recorded == [("warning", "Telegram network issue: proxy disconnected")]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_error_keeps_non_network_exceptions_as_error(monkeypatch) -> None:
|
||||||
|
channel = TelegramChannel(
|
||||||
|
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
recorded: list[tuple[str, str]] = []
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.channels.telegram.logger.warning",
|
||||||
|
lambda message, error: recorded.append(("warning", message.format(error))),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"nanobot.channels.telegram.logger.error",
|
||||||
|
lambda message, error: recorded.append(("error", message.format(error))),
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._on_error(object(), SimpleNamespace(error=RuntimeError("boom")))
|
||||||
|
|
||||||
|
assert recorded == [("error", "Telegram error: boom")]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_stream_end_raises_and_keeps_buffer_on_failure() -> None:
|
||||||
|
channel = TelegramChannel(
|
||||||
|
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
channel._app = _FakeApp(lambda: None)
|
||||||
|
channel._app.bot.edit_message_text = AsyncMock(side_effect=RuntimeError("boom"))
|
||||||
|
channel._stream_bufs["123"] = _StreamBuf(text="hello", message_id=7, last_edit=0.0)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="boom"):
|
||||||
|
await channel.send_delta("123", "", {"_stream_end": True})
|
||||||
|
|
||||||
|
assert "123" in channel._stream_bufs
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_stream_end_treats_not_modified_as_success() -> None:
|
||||||
|
from telegram.error import BadRequest
|
||||||
|
|
||||||
|
channel = TelegramChannel(
|
||||||
|
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
channel._app = _FakeApp(lambda: None)
|
||||||
|
channel._app.bot.edit_message_text = AsyncMock(side_effect=BadRequest("Message is not modified"))
|
||||||
|
channel._stream_bufs["123"] = _StreamBuf(text="hello", message_id=7, last_edit=0.0, stream_id="s:0")
|
||||||
|
|
||||||
|
await channel.send_delta("123", "", {"_stream_end": True, "_stream_id": "s:0"})
|
||||||
|
|
||||||
|
assert "123" not in channel._stream_bufs
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_new_stream_id_replaces_stale_buffer() -> None:
|
||||||
|
channel = TelegramChannel(
|
||||||
|
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
channel._app = _FakeApp(lambda: None)
|
||||||
|
channel._stream_bufs["123"] = _StreamBuf(
|
||||||
|
text="hello",
|
||||||
|
message_id=7,
|
||||||
|
last_edit=0.0,
|
||||||
|
stream_id="old:0",
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel.send_delta("123", "world", {"_stream_delta": True, "_stream_id": "new:0"})
|
||||||
|
|
||||||
|
buf = channel._stream_bufs["123"]
|
||||||
|
assert buf.text == "world"
|
||||||
|
assert buf.stream_id == "new:0"
|
||||||
|
assert buf.message_id == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_incremental_edit_treats_not_modified_as_success() -> None:
|
||||||
|
from telegram.error import BadRequest
|
||||||
|
|
||||||
|
channel = TelegramChannel(
|
||||||
|
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
channel._app = _FakeApp(lambda: None)
|
||||||
|
channel._stream_bufs["123"] = _StreamBuf(text="hello", message_id=7, last_edit=0.0, stream_id="s:0")
|
||||||
|
channel._app.bot.edit_message_text = AsyncMock(side_effect=BadRequest("Message is not modified"))
|
||||||
|
|
||||||
|
await channel.send_delta("123", "", {"_stream_delta": True, "_stream_id": "s:0"})
|
||||||
|
|
||||||
|
assert channel._stream_bufs["123"].last_edit > 0.0
|
||||||
|
|
||||||
|
|
||||||
def test_derive_topic_session_key_uses_thread_id() -> None:
|
def test_derive_topic_session_key_uses_thread_id() -> None:
|
||||||
message = SimpleNamespace(
|
message = SimpleNamespace(
|
||||||
chat=SimpleNamespace(type="supergroup"),
|
chat=SimpleNamespace(type="supergroup"),
|
||||||
|
|||||||
@@ -1,17 +1,22 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import tempfile
|
import tempfile
|
||||||
|
from pathlib import Path
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
import nanobot.channels.weixin as weixin_mod
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.channels.weixin import (
|
from nanobot.channels.weixin import (
|
||||||
ITEM_IMAGE,
|
ITEM_IMAGE,
|
||||||
ITEM_TEXT,
|
ITEM_TEXT,
|
||||||
MESSAGE_TYPE_BOT,
|
MESSAGE_TYPE_BOT,
|
||||||
WEIXIN_CHANNEL_VERSION,
|
WEIXIN_CHANNEL_VERSION,
|
||||||
|
_decrypt_aes_ecb,
|
||||||
|
_encrypt_aes_ecb,
|
||||||
WeixinChannel,
|
WeixinChannel,
|
||||||
WeixinConfig,
|
WeixinConfig,
|
||||||
)
|
)
|
||||||
@@ -42,10 +47,12 @@ def test_make_headers_includes_route_tag_when_configured() -> None:
|
|||||||
|
|
||||||
assert headers["Authorization"] == "Bearer token"
|
assert headers["Authorization"] == "Bearer token"
|
||||||
assert headers["SKRouteTag"] == "123"
|
assert headers["SKRouteTag"] == "123"
|
||||||
|
assert headers["iLink-App-Id"] == "bot"
|
||||||
|
assert headers["iLink-App-ClientVersion"] == str((2 << 16) | (1 << 8) | 1)
|
||||||
|
|
||||||
|
|
||||||
def test_channel_version_matches_reference_plugin_version() -> None:
|
def test_channel_version_matches_reference_plugin_version() -> None:
|
||||||
assert WEIXIN_CHANNEL_VERSION == "1.0.3"
|
assert WEIXIN_CHANNEL_VERSION == "2.1.1"
|
||||||
|
|
||||||
|
|
||||||
def test_save_and_load_state_persists_context_tokens(tmp_path) -> None:
|
def test_save_and_load_state_persists_context_tokens(tmp_path) -> None:
|
||||||
@@ -169,6 +176,120 @@ async def test_process_message_extracts_media_and_preserves_paths() -> None:
|
|||||||
assert inbound.media == ["/tmp/test.jpg"]
|
assert inbound.media == ["/tmp/test.jpg"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_process_message_falls_back_to_referenced_media_when_no_top_level_media() -> None:
|
||||||
|
channel, bus = _make_channel()
|
||||||
|
channel._download_media_item = AsyncMock(return_value="/tmp/ref.jpg")
|
||||||
|
|
||||||
|
await channel._process_message(
|
||||||
|
{
|
||||||
|
"message_type": 1,
|
||||||
|
"message_id": "m3-ref-fallback",
|
||||||
|
"from_user_id": "wx-user",
|
||||||
|
"context_token": "ctx-3-ref-fallback",
|
||||||
|
"item_list": [
|
||||||
|
{
|
||||||
|
"type": ITEM_TEXT,
|
||||||
|
"text_item": {"text": "reply to image"},
|
||||||
|
"ref_msg": {
|
||||||
|
"message_item": {
|
||||||
|
"type": ITEM_IMAGE,
|
||||||
|
"image_item": {"media": {"encrypt_query_param": "ref-enc"}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
inbound = await asyncio.wait_for(bus.consume_inbound(), timeout=1.0)
|
||||||
|
|
||||||
|
channel._download_media_item.assert_awaited_once_with(
|
||||||
|
{"media": {"encrypt_query_param": "ref-enc"}},
|
||||||
|
"image",
|
||||||
|
)
|
||||||
|
assert inbound.media == ["/tmp/ref.jpg"]
|
||||||
|
assert "reply to image" in inbound.content
|
||||||
|
assert "[image]" in inbound.content
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_process_message_does_not_use_referenced_fallback_when_top_level_media_exists() -> None:
|
||||||
|
channel, bus = _make_channel()
|
||||||
|
channel._download_media_item = AsyncMock(side_effect=["/tmp/top.jpg", "/tmp/ref.jpg"])
|
||||||
|
|
||||||
|
await channel._process_message(
|
||||||
|
{
|
||||||
|
"message_type": 1,
|
||||||
|
"message_id": "m3-ref-no-fallback",
|
||||||
|
"from_user_id": "wx-user",
|
||||||
|
"context_token": "ctx-3-ref-no-fallback",
|
||||||
|
"item_list": [
|
||||||
|
{"type": ITEM_IMAGE, "image_item": {"media": {"encrypt_query_param": "top-enc"}}},
|
||||||
|
{
|
||||||
|
"type": ITEM_TEXT,
|
||||||
|
"text_item": {"text": "has top-level media"},
|
||||||
|
"ref_msg": {
|
||||||
|
"message_item": {
|
||||||
|
"type": ITEM_IMAGE,
|
||||||
|
"image_item": {"media": {"encrypt_query_param": "ref-enc"}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
inbound = await asyncio.wait_for(bus.consume_inbound(), timeout=1.0)
|
||||||
|
|
||||||
|
channel._download_media_item.assert_awaited_once_with(
|
||||||
|
{"media": {"encrypt_query_param": "top-enc"}},
|
||||||
|
"image",
|
||||||
|
)
|
||||||
|
assert inbound.media == ["/tmp/top.jpg"]
|
||||||
|
assert "/tmp/ref.jpg" not in inbound.content
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_process_message_does_not_fallback_when_top_level_media_exists_but_download_fails() -> None:
|
||||||
|
channel, bus = _make_channel()
|
||||||
|
# Top-level image download fails (None), referenced image would succeed if fallback were triggered.
|
||||||
|
channel._download_media_item = AsyncMock(side_effect=[None, "/tmp/ref.jpg"])
|
||||||
|
|
||||||
|
await channel._process_message(
|
||||||
|
{
|
||||||
|
"message_type": 1,
|
||||||
|
"message_id": "m3-ref-no-fallback-on-failure",
|
||||||
|
"from_user_id": "wx-user",
|
||||||
|
"context_token": "ctx-3-ref-no-fallback-on-failure",
|
||||||
|
"item_list": [
|
||||||
|
{"type": ITEM_IMAGE, "image_item": {"media": {"encrypt_query_param": "top-enc"}}},
|
||||||
|
{
|
||||||
|
"type": ITEM_TEXT,
|
||||||
|
"text_item": {"text": "quoted has media"},
|
||||||
|
"ref_msg": {
|
||||||
|
"message_item": {
|
||||||
|
"type": ITEM_IMAGE,
|
||||||
|
"image_item": {"media": {"encrypt_query_param": "ref-enc"}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
inbound = await asyncio.wait_for(bus.consume_inbound(), timeout=1.0)
|
||||||
|
|
||||||
|
# Should only attempt top-level media item; reference fallback must not activate.
|
||||||
|
channel._download_media_item.assert_awaited_once_with(
|
||||||
|
{"media": {"encrypt_query_param": "top-enc"}},
|
||||||
|
"image",
|
||||||
|
)
|
||||||
|
assert inbound.media == []
|
||||||
|
assert "[image]" in inbound.content
|
||||||
|
assert "/tmp/ref.jpg" not in inbound.content
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_send_without_context_token_does_not_send_text() -> None:
|
async def test_send_without_context_token_does_not_send_text() -> None:
|
||||||
channel, _bus = _make_channel()
|
channel, _bus = _make_channel()
|
||||||
@@ -199,6 +320,70 @@ async def test_send_does_not_send_when_session_is_paused() -> None:
|
|||||||
channel._send_text.assert_not_awaited()
|
channel._send_text.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_typing_ticket_fetches_and_caches_per_user() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
channel._api_post = AsyncMock(return_value={"ret": 0, "typing_ticket": "ticket-1"})
|
||||||
|
|
||||||
|
first = await channel._get_typing_ticket("wx-user", "ctx-1")
|
||||||
|
second = await channel._get_typing_ticket("wx-user", "ctx-2")
|
||||||
|
|
||||||
|
assert first == "ticket-1"
|
||||||
|
assert second == "ticket-1"
|
||||||
|
channel._api_post.assert_awaited_once_with(
|
||||||
|
"ilink/bot/getconfig",
|
||||||
|
{"ilink_user_id": "wx-user", "context_token": "ctx-1", "base_info": weixin_mod.BASE_INFO},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_uses_typing_start_and_cancel_when_ticket_available() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
channel._context_tokens["wx-user"] = "ctx-typing"
|
||||||
|
channel._send_text = AsyncMock()
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"ret": 0, "typing_ticket": "ticket-typing"},
|
||||||
|
{"ret": 0},
|
||||||
|
{"ret": 0},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel.send(
|
||||||
|
type("Msg", (), {"chat_id": "wx-user", "content": "pong", "media": [], "metadata": {}})()
|
||||||
|
)
|
||||||
|
|
||||||
|
channel._send_text.assert_awaited_once_with("wx-user", "pong", "ctx-typing")
|
||||||
|
assert channel._api_post.await_count == 3
|
||||||
|
assert channel._api_post.await_args_list[0].args[0] == "ilink/bot/getconfig"
|
||||||
|
assert channel._api_post.await_args_list[1].args[0] == "ilink/bot/sendtyping"
|
||||||
|
assert channel._api_post.await_args_list[1].args[1]["status"] == 1
|
||||||
|
assert channel._api_post.await_args_list[2].args[0] == "ilink/bot/sendtyping"
|
||||||
|
assert channel._api_post.await_args_list[2].args[1]["status"] == 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_still_sends_text_when_typing_ticket_missing() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
channel._context_tokens["wx-user"] = "ctx-no-ticket"
|
||||||
|
channel._send_text = AsyncMock()
|
||||||
|
channel._api_post = AsyncMock(return_value={"ret": 1, "errmsg": "no config"})
|
||||||
|
|
||||||
|
await channel.send(
|
||||||
|
type("Msg", (), {"chat_id": "wx-user", "content": "pong", "media": [], "metadata": {}})()
|
||||||
|
)
|
||||||
|
|
||||||
|
channel._send_text.assert_awaited_once_with("wx-user", "pong", "ctx-no-ticket")
|
||||||
|
channel._api_post.assert_awaited_once()
|
||||||
|
assert channel._api_post.await_args_list[0].args[0] == "ilink/bot/getconfig"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_poll_once_pauses_session_on_expired_errcode() -> None:
|
async def test_poll_once_pauses_session_on_expired_errcode() -> None:
|
||||||
channel, _bus = _make_channel()
|
channel, _bus = _make_channel()
|
||||||
@@ -220,8 +405,12 @@ async def test_qr_login_refreshes_expired_qr_and_then_succeeds() -> None:
|
|||||||
channel._api_get = AsyncMock(
|
channel._api_get = AsyncMock(
|
||||||
side_effect=[
|
side_effect=[
|
||||||
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
||||||
{"status": "expired"},
|
|
||||||
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"status": "expired"},
|
||||||
{
|
{
|
||||||
"status": "confirmed",
|
"status": "confirmed",
|
||||||
"bot_token": "token-2",
|
"bot_token": "token-2",
|
||||||
@@ -247,12 +436,16 @@ async def test_qr_login_returns_false_after_too_many_expired_qr_codes() -> None:
|
|||||||
channel._api_get = AsyncMock(
|
channel._api_get = AsyncMock(
|
||||||
side_effect=[
|
side_effect=[
|
||||||
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
||||||
{"status": "expired"},
|
|
||||||
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
||||||
{"status": "expired"},
|
|
||||||
{"qrcode": "qr-3", "qrcode_img_content": "url-3"},
|
{"qrcode": "qr-3", "qrcode_img_content": "url-3"},
|
||||||
{"status": "expired"},
|
|
||||||
{"qrcode": "qr-4", "qrcode_img_content": "url-4"},
|
{"qrcode": "qr-4", "qrcode_img_content": "url-4"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"status": "expired"},
|
||||||
|
{"status": "expired"},
|
||||||
|
{"status": "expired"},
|
||||||
{"status": "expired"},
|
{"status": "expired"},
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
@@ -262,6 +455,105 @@ async def test_qr_login_returns_false_after_too_many_expired_qr_codes() -> None:
|
|||||||
assert ok is False
|
assert ok is False
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_switches_polling_base_url_on_redirect_status() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
|
||||||
|
|
||||||
|
status_side_effect = [
|
||||||
|
{"status": "scaned_but_redirect", "redirect_host": "idc.redirect.test"},
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-3",
|
||||||
|
"ilink_bot_id": "bot-3",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
channel._api_get = AsyncMock(side_effect=list(status_side_effect))
|
||||||
|
channel._api_get_with_base = AsyncMock(side_effect=list(status_side_effect))
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-3"
|
||||||
|
assert channel._api_get_with_base.await_count == 2
|
||||||
|
first_call = channel._api_get_with_base.await_args_list[0]
|
||||||
|
second_call = channel._api_get_with_base.await_args_list[1]
|
||||||
|
assert first_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
assert second_call.kwargs["base_url"] == "https://idc.redirect.test"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_redirect_without_host_keeps_current_polling_base_url() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
|
||||||
|
|
||||||
|
status_side_effect = [
|
||||||
|
{"status": "scaned_but_redirect"},
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-4",
|
||||||
|
"ilink_bot_id": "bot-4",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
channel._api_get = AsyncMock(side_effect=list(status_side_effect))
|
||||||
|
channel._api_get_with_base = AsyncMock(side_effect=list(status_side_effect))
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-4"
|
||||||
|
assert channel._api_get_with_base.await_count == 2
|
||||||
|
first_call = channel._api_get_with_base.await_args_list[0]
|
||||||
|
second_call = channel._api_get_with_base.await_args_list[1]
|
||||||
|
assert first_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
assert second_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_resets_redirect_base_url_after_qr_refresh() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(side_effect=[("qr-1", "url-1"), ("qr-2", "url-2")])
|
||||||
|
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"status": "scaned_but_redirect", "redirect_host": "idc.redirect.test"},
|
||||||
|
{"status": "expired"},
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-5",
|
||||||
|
"ilink_bot_id": "bot-5",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-5"
|
||||||
|
assert channel._api_get_with_base.await_count == 3
|
||||||
|
first_call = channel._api_get_with_base.await_args_list[0]
|
||||||
|
second_call = channel._api_get_with_base.await_args_list[1]
|
||||||
|
third_call = channel._api_get_with_base.await_args_list[2]
|
||||||
|
assert first_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
assert second_call.kwargs["base_url"] == "https://idc.redirect.test"
|
||||||
|
assert third_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_process_message_skips_bot_messages() -> None:
|
async def test_process_message_skips_bot_messages() -> None:
|
||||||
channel, bus = _make_channel()
|
channel, bus = _make_channel()
|
||||||
@@ -278,3 +570,357 @@ async def test_process_message_skips_bot_messages() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert bus.inbound_size == 0
|
assert bus.inbound_size == 0
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyHttpResponse:
|
||||||
|
def __init__(self, *, headers: dict[str, str] | None = None, status_code: int = 200) -> None:
|
||||||
|
self.headers = headers or {}
|
||||||
|
self.status_code = status_code
|
||||||
|
|
||||||
|
def raise_for_status(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_media_uses_upload_full_url_when_present(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
|
||||||
|
media_file = tmp_path / "photo.jpg"
|
||||||
|
media_file.write_bytes(b"hello-weixin")
|
||||||
|
|
||||||
|
cdn_post = AsyncMock(return_value=_DummyHttpResponse(headers={"x-encrypted-param": "dl-param"}))
|
||||||
|
channel._client = SimpleNamespace(post=cdn_post)
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{
|
||||||
|
"upload_full_url": "https://upload-full.example.test/path?foo=bar",
|
||||||
|
"upload_param": "should-not-be-used",
|
||||||
|
},
|
||||||
|
{"ret": 0},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._send_media_file("wx-user", str(media_file), "ctx-1")
|
||||||
|
|
||||||
|
# first POST call is CDN upload
|
||||||
|
cdn_url = cdn_post.await_args_list[0].args[0]
|
||||||
|
assert cdn_url == "https://upload-full.example.test/path?foo=bar"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_media_falls_back_to_upload_param_url(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
|
||||||
|
media_file = tmp_path / "photo.jpg"
|
||||||
|
media_file.write_bytes(b"hello-weixin")
|
||||||
|
|
||||||
|
cdn_post = AsyncMock(return_value=_DummyHttpResponse(headers={"x-encrypted-param": "dl-param"}))
|
||||||
|
channel._client = SimpleNamespace(post=cdn_post)
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"upload_param": "enc-need-fallback"},
|
||||||
|
{"ret": 0},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._send_media_file("wx-user", str(media_file), "ctx-1")
|
||||||
|
|
||||||
|
cdn_url = cdn_post.await_args_list[0].args[0]
|
||||||
|
assert cdn_url.startswith(f"{channel.config.cdn_base_url}/upload?encrypted_query_param=enc-need-fallback")
|
||||||
|
assert "&filekey=" in cdn_url
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_media_voice_file_uses_voice_item_and_voice_upload_type(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
|
||||||
|
media_file = tmp_path / "voice.mp3"
|
||||||
|
media_file.write_bytes(b"voice-bytes")
|
||||||
|
|
||||||
|
cdn_post = AsyncMock(return_value=_DummyHttpResponse(headers={"x-encrypted-param": "voice-dl-param"}))
|
||||||
|
channel._client = SimpleNamespace(post=cdn_post)
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"upload_full_url": "https://upload-full.example.test/voice?foo=bar"},
|
||||||
|
{"ret": 0},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._send_media_file("wx-user", str(media_file), "ctx-voice")
|
||||||
|
|
||||||
|
getupload_body = channel._api_post.await_args_list[0].args[1]
|
||||||
|
assert getupload_body["media_type"] == 4
|
||||||
|
|
||||||
|
sendmessage_body = channel._api_post.await_args_list[1].args[1]
|
||||||
|
item = sendmessage_body["msg"]["item_list"][0]
|
||||||
|
assert item["type"] == 3
|
||||||
|
assert "voice_item" in item
|
||||||
|
assert "file_item" not in item
|
||||||
|
assert item["voice_item"]["media"]["encrypt_query_param"] == "voice-dl-param"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_typing_uses_keepalive_until_send_finishes() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
channel._context_tokens["wx-user"] = "ctx-typing-loop"
|
||||||
|
async def _api_post_side_effect(endpoint: str, _body: dict | None = None, *, auth: bool = True):
|
||||||
|
if endpoint == "ilink/bot/getconfig":
|
||||||
|
return {"ret": 0, "typing_ticket": "ticket-keepalive"}
|
||||||
|
return {"ret": 0}
|
||||||
|
|
||||||
|
channel._api_post = AsyncMock(side_effect=_api_post_side_effect)
|
||||||
|
|
||||||
|
async def _slow_send_text(*_args, **_kwargs) -> None:
|
||||||
|
await asyncio.sleep(0.03)
|
||||||
|
|
||||||
|
channel._send_text = AsyncMock(side_effect=_slow_send_text)
|
||||||
|
|
||||||
|
old_interval = weixin_mod.TYPING_KEEPALIVE_INTERVAL_S
|
||||||
|
weixin_mod.TYPING_KEEPALIVE_INTERVAL_S = 0.01
|
||||||
|
try:
|
||||||
|
await channel.send(
|
||||||
|
type("Msg", (), {"chat_id": "wx-user", "content": "pong", "media": [], "metadata": {}})()
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
weixin_mod.TYPING_KEEPALIVE_INTERVAL_S = old_interval
|
||||||
|
|
||||||
|
status_calls = [
|
||||||
|
c.args[1]["status"]
|
||||||
|
for c in channel._api_post.await_args_list
|
||||||
|
if c.args and c.args[0] == "ilink/bot/sendtyping"
|
||||||
|
]
|
||||||
|
assert status_calls.count(1) >= 2
|
||||||
|
assert status_calls[-1] == 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_typing_ticket_failure_uses_backoff_and_cached_ticket(monkeypatch) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
|
||||||
|
now = {"value": 1000.0}
|
||||||
|
monkeypatch.setattr(weixin_mod.time, "time", lambda: now["value"])
|
||||||
|
monkeypatch.setattr(weixin_mod.random, "random", lambda: 0.5)
|
||||||
|
|
||||||
|
channel._api_post = AsyncMock(return_value={"ret": 0, "typing_ticket": "ticket-ok"})
|
||||||
|
first = await channel._get_typing_ticket("wx-user", "ctx-1")
|
||||||
|
assert first == "ticket-ok"
|
||||||
|
|
||||||
|
# force refresh window reached
|
||||||
|
now["value"] = now["value"] + (12 * 60 * 60) + 1
|
||||||
|
channel._api_post = AsyncMock(return_value={"ret": 1, "errmsg": "temporary failure"})
|
||||||
|
|
||||||
|
# On refresh failure, should still return cached ticket and apply backoff.
|
||||||
|
second = await channel._get_typing_ticket("wx-user", "ctx-2")
|
||||||
|
assert second == "ticket-ok"
|
||||||
|
assert channel._api_post.await_count == 1
|
||||||
|
|
||||||
|
# Before backoff expiry, no extra fetch should happen.
|
||||||
|
now["value"] += 1
|
||||||
|
third = await channel._get_typing_ticket("wx-user", "ctx-3")
|
||||||
|
assert third == "ticket-ok"
|
||||||
|
assert channel._api_post.await_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_treats_temporary_connect_error_as_wait_and_recovers() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
|
||||||
|
|
||||||
|
request = httpx.Request("GET", "https://ilinkai.weixin.qq.com/ilink/bot/get_qrcode_status")
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
httpx.ConnectError("temporary network", request=request),
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-net-ok",
|
||||||
|
"ilink_bot_id": "bot-id",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-net-ok"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_treats_5xx_gateway_response_error_as_wait_and_recovers() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
|
||||||
|
|
||||||
|
request = httpx.Request("GET", "https://ilinkai.weixin.qq.com/ilink/bot/get_qrcode_status")
|
||||||
|
response = httpx.Response(status_code=524, request=request)
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
httpx.HTTPStatusError("gateway timeout", request=request, response=response),
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-5xx-ok",
|
||||||
|
"ilink_bot_id": "bot-id",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-5xx-ok"
|
||||||
|
|
||||||
|
|
||||||
|
def test_decrypt_aes_ecb_strips_valid_pkcs7_padding() -> None:
|
||||||
|
key_b64 = "MDEyMzQ1Njc4OWFiY2RlZg==" # base64("0123456789abcdef")
|
||||||
|
plaintext = b"hello-weixin-padding"
|
||||||
|
|
||||||
|
ciphertext = _encrypt_aes_ecb(plaintext, key_b64)
|
||||||
|
decrypted = _decrypt_aes_ecb(ciphertext, key_b64)
|
||||||
|
|
||||||
|
assert decrypted == plaintext
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyDownloadResponse:
|
||||||
|
def __init__(self, content: bytes, status_code: int = 200) -> None:
|
||||||
|
self.content = content
|
||||||
|
self.status_code = status_code
|
||||||
|
|
||||||
|
def raise_for_status(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyErrorDownloadResponse(_DummyDownloadResponse):
|
||||||
|
def __init__(self, url: str, status_code: int) -> None:
|
||||||
|
super().__init__(content=b"", status_code=status_code)
|
||||||
|
self._url = url
|
||||||
|
|
||||||
|
def raise_for_status(self) -> None:
|
||||||
|
request = httpx.Request("GET", self._url)
|
||||||
|
response = httpx.Response(self.status_code, request=request)
|
||||||
|
raise httpx.HTTPStatusError(
|
||||||
|
f"download failed with status {self.status_code}",
|
||||||
|
request=request,
|
||||||
|
response=response,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_uses_full_url_when_present(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
full_url = "https://cdn.example.test/download/full"
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(return_value=_DummyDownloadResponse(content=b"raw-image-bytes"))
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {
|
||||||
|
"media": {
|
||||||
|
"full_url": full_url,
|
||||||
|
"encrypt_query_param": "enc-fallback-should-not-be-used",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
saved_path = await channel._download_media_item(item, "image")
|
||||||
|
|
||||||
|
assert saved_path is not None
|
||||||
|
assert Path(saved_path).read_bytes() == b"raw-image-bytes"
|
||||||
|
channel._client.get.assert_awaited_once_with(full_url)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_falls_back_when_full_url_returns_retryable_error(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
full_url = "https://cdn.example.test/download/full?taskid=123"
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
_DummyErrorDownloadResponse(full_url, 500),
|
||||||
|
_DummyDownloadResponse(content=b"fallback-bytes"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {
|
||||||
|
"media": {
|
||||||
|
"full_url": full_url,
|
||||||
|
"encrypt_query_param": "enc-fallback",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
saved_path = await channel._download_media_item(item, "image")
|
||||||
|
|
||||||
|
assert saved_path is not None
|
||||||
|
assert Path(saved_path).read_bytes() == b"fallback-bytes"
|
||||||
|
assert channel._client.get.await_count == 2
|
||||||
|
assert channel._client.get.await_args_list[0].args[0] == full_url
|
||||||
|
fallback_url = channel._client.get.await_args_list[1].args[0]
|
||||||
|
assert fallback_url.startswith(f"{channel.config.cdn_base_url}/download?encrypted_query_param=enc-fallback")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_falls_back_to_encrypt_query_param(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(return_value=_DummyDownloadResponse(content=b"fallback-bytes"))
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {"media": {"encrypt_query_param": "enc-fallback"}}
|
||||||
|
saved_path = await channel._download_media_item(item, "image")
|
||||||
|
|
||||||
|
assert saved_path is not None
|
||||||
|
assert Path(saved_path).read_bytes() == b"fallback-bytes"
|
||||||
|
called_url = channel._client.get.await_args_list[0].args[0]
|
||||||
|
assert called_url.startswith(f"{channel.config.cdn_base_url}/download?encrypted_query_param=enc-fallback")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_does_not_retry_when_full_url_fails_without_fallback(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
full_url = "https://cdn.example.test/download/full"
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(return_value=_DummyErrorDownloadResponse(full_url, 500))
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {"media": {"full_url": full_url}}
|
||||||
|
saved_path = await channel._download_media_item(item, "image")
|
||||||
|
|
||||||
|
assert saved_path is None
|
||||||
|
channel._client.get.assert_awaited_once_with(full_url)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_non_image_requires_aes_key_even_with_full_url(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
full_url = "https://cdn.example.test/download/voice"
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(return_value=_DummyDownloadResponse(content=b"ciphertext-or-unknown"))
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {
|
||||||
|
"media": {
|
||||||
|
"full_url": full_url,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
saved_path = await channel._download_media_item(item, "voice")
|
||||||
|
|
||||||
|
assert saved_path is None
|
||||||
|
channel._client.get.assert_not_awaited()
|
||||||
|
|||||||
+184
-78
@@ -642,27 +642,105 @@ def test_heartbeat_retains_recent_messages_by_default():
|
|||||||
assert config.gateway.heartbeat.keep_recent_messages == 8
|
assert config.gateway.heartbeat.keep_recent_messages == 8
|
||||||
|
|
||||||
|
|
||||||
def test_gateway_uses_workspace_from_config_by_default(monkeypatch, tmp_path: Path) -> None:
|
def _write_instance_config(tmp_path: Path) -> Path:
|
||||||
config_file = tmp_path / "instance" / "config.json"
|
config_file = tmp_path / "instance" / "config.json"
|
||||||
config_file.parent.mkdir(parents=True)
|
config_file.parent.mkdir(parents=True)
|
||||||
config_file.write_text("{}")
|
config_file.write_text("{}")
|
||||||
|
return config_file
|
||||||
|
|
||||||
config = Config()
|
|
||||||
config.agents.defaults.workspace = str(tmp_path / "config-workspace")
|
|
||||||
seen: dict[str, Path] = {}
|
|
||||||
|
|
||||||
|
def _stop_gateway_provider(_config) -> object:
|
||||||
|
raise _StopGatewayError("stop")
|
||||||
|
|
||||||
|
|
||||||
|
def _patch_cli_command_runtime(
|
||||||
|
monkeypatch,
|
||||||
|
config: Config,
|
||||||
|
*,
|
||||||
|
set_config_path=None,
|
||||||
|
sync_templates=None,
|
||||||
|
make_provider=None,
|
||||||
|
message_bus=None,
|
||||||
|
session_manager=None,
|
||||||
|
cron_service=None,
|
||||||
|
get_cron_dir=None,
|
||||||
|
) -> None:
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.config.loader.set_config_path",
|
"nanobot.config.loader.set_config_path",
|
||||||
lambda path: seen.__setitem__("config_path", path),
|
set_config_path or (lambda _path: None),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.cli.commands.sync_workspace_templates",
|
"nanobot.cli.commands.sync_workspace_templates",
|
||||||
lambda path: seen.__setitem__("workspace", path),
|
sync_templates or (lambda _path: None),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.cli.commands._make_provider",
|
"nanobot.cli.commands._make_provider",
|
||||||
lambda _config: (_ for _ in ()).throw(_StopGatewayError("stop")),
|
make_provider or (lambda _config: object()),
|
||||||
|
)
|
||||||
|
|
||||||
|
if message_bus is not None:
|
||||||
|
monkeypatch.setattr("nanobot.bus.queue.MessageBus", message_bus)
|
||||||
|
if session_manager is not None:
|
||||||
|
monkeypatch.setattr("nanobot.session.manager.SessionManager", session_manager)
|
||||||
|
if cron_service is not None:
|
||||||
|
monkeypatch.setattr("nanobot.cron.service.CronService", cron_service)
|
||||||
|
if get_cron_dir is not None:
|
||||||
|
monkeypatch.setattr("nanobot.config.paths.get_cron_dir", get_cron_dir)
|
||||||
|
|
||||||
|
|
||||||
|
def _patch_serve_runtime(monkeypatch, config: Config, seen: dict[str, object]) -> None:
|
||||||
|
pytest.importorskip("aiohttp")
|
||||||
|
|
||||||
|
class _FakeApiApp:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.on_startup: list[object] = []
|
||||||
|
self.on_cleanup: list[object] = []
|
||||||
|
|
||||||
|
class _FakeAgentLoop:
|
||||||
|
def __init__(self, **kwargs) -> None:
|
||||||
|
seen["workspace"] = kwargs["workspace"]
|
||||||
|
|
||||||
|
async def _connect_mcp(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def close_mcp(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _fake_create_app(agent_loop, model_name: str, request_timeout: float):
|
||||||
|
seen["agent_loop"] = agent_loop
|
||||||
|
seen["model_name"] = model_name
|
||||||
|
seen["request_timeout"] = request_timeout
|
||||||
|
return _FakeApiApp()
|
||||||
|
|
||||||
|
def _fake_run_app(api_app, host: str, port: int, print):
|
||||||
|
seen["api_app"] = api_app
|
||||||
|
seen["host"] = host
|
||||||
|
seen["port"] = port
|
||||||
|
|
||||||
|
_patch_cli_command_runtime(
|
||||||
|
monkeypatch,
|
||||||
|
config,
|
||||||
|
message_bus=lambda: object(),
|
||||||
|
session_manager=lambda _workspace: object(),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("nanobot.agent.loop.AgentLoop", _FakeAgentLoop)
|
||||||
|
monkeypatch.setattr("nanobot.api.server.create_app", _fake_create_app)
|
||||||
|
monkeypatch.setattr("aiohttp.web.run_app", _fake_run_app)
|
||||||
|
|
||||||
|
|
||||||
|
def test_gateway_uses_workspace_from_config_by_default(monkeypatch, tmp_path: Path) -> None:
|
||||||
|
config_file = _write_instance_config(tmp_path)
|
||||||
|
config = Config()
|
||||||
|
config.agents.defaults.workspace = str(tmp_path / "config-workspace")
|
||||||
|
seen: dict[str, Path] = {}
|
||||||
|
|
||||||
|
_patch_cli_command_runtime(
|
||||||
|
monkeypatch,
|
||||||
|
config,
|
||||||
|
set_config_path=lambda path: seen.__setitem__("config_path", path),
|
||||||
|
sync_templates=lambda path: seen.__setitem__("workspace", path),
|
||||||
|
make_provider=_stop_gateway_provider,
|
||||||
)
|
)
|
||||||
|
|
||||||
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
||||||
@@ -673,24 +751,17 @@ def test_gateway_uses_workspace_from_config_by_default(monkeypatch, tmp_path: Pa
|
|||||||
|
|
||||||
|
|
||||||
def test_gateway_workspace_option_overrides_config(monkeypatch, tmp_path: Path) -> None:
|
def test_gateway_workspace_option_overrides_config(monkeypatch, tmp_path: Path) -> None:
|
||||||
config_file = tmp_path / "instance" / "config.json"
|
config_file = _write_instance_config(tmp_path)
|
||||||
config_file.parent.mkdir(parents=True)
|
|
||||||
config_file.write_text("{}")
|
|
||||||
|
|
||||||
config = Config()
|
config = Config()
|
||||||
config.agents.defaults.workspace = str(tmp_path / "config-workspace")
|
config.agents.defaults.workspace = str(tmp_path / "config-workspace")
|
||||||
override = tmp_path / "override-workspace"
|
override = tmp_path / "override-workspace"
|
||||||
seen: dict[str, Path] = {}
|
seen: dict[str, Path] = {}
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.set_config_path", lambda _path: None)
|
_patch_cli_command_runtime(
|
||||||
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
monkeypatch,
|
||||||
monkeypatch.setattr(
|
config,
|
||||||
"nanobot.cli.commands.sync_workspace_templates",
|
sync_templates=lambda path: seen.__setitem__("workspace", path),
|
||||||
lambda path: seen.__setitem__("workspace", path),
|
make_provider=_stop_gateway_provider,
|
||||||
)
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"nanobot.cli.commands._make_provider",
|
|
||||||
lambda _config: (_ for _ in ()).throw(_StopGatewayError("stop")),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
result = runner.invoke(
|
result = runner.invoke(
|
||||||
@@ -704,27 +775,23 @@ def test_gateway_workspace_option_overrides_config(monkeypatch, tmp_path: Path)
|
|||||||
|
|
||||||
|
|
||||||
def test_gateway_uses_workspace_directory_for_cron_store(monkeypatch, tmp_path: Path) -> None:
|
def test_gateway_uses_workspace_directory_for_cron_store(monkeypatch, tmp_path: Path) -> None:
|
||||||
config_file = tmp_path / "instance" / "config.json"
|
config_file = _write_instance_config(tmp_path)
|
||||||
config_file.parent.mkdir(parents=True)
|
|
||||||
config_file.write_text("{}")
|
|
||||||
|
|
||||||
config = Config()
|
config = Config()
|
||||||
config.agents.defaults.workspace = str(tmp_path / "config-workspace")
|
config.agents.defaults.workspace = str(tmp_path / "config-workspace")
|
||||||
seen: dict[str, Path] = {}
|
seen: dict[str, Path] = {}
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.set_config_path", lambda _path: None)
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
|
||||||
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
|
||||||
monkeypatch.setattr("nanobot.cli.commands._make_provider", lambda _config: object())
|
|
||||||
monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: object())
|
|
||||||
monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object())
|
|
||||||
|
|
||||||
class _StopCron:
|
class _StopCron:
|
||||||
def __init__(self, store_path: Path) -> None:
|
def __init__(self, store_path: Path) -> None:
|
||||||
seen["cron_store"] = store_path
|
seen["cron_store"] = store_path
|
||||||
raise _StopGatewayError("stop")
|
raise _StopGatewayError("stop")
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.cron.service.CronService", _StopCron)
|
_patch_cli_command_runtime(
|
||||||
|
monkeypatch,
|
||||||
|
config,
|
||||||
|
message_bus=lambda: object(),
|
||||||
|
session_manager=lambda _workspace: object(),
|
||||||
|
cron_service=_StopCron,
|
||||||
|
)
|
||||||
|
|
||||||
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
||||||
|
|
||||||
@@ -735,10 +802,7 @@ def test_gateway_uses_workspace_directory_for_cron_store(monkeypatch, tmp_path:
|
|||||||
def test_gateway_workspace_override_does_not_migrate_legacy_cron(
|
def test_gateway_workspace_override_does_not_migrate_legacy_cron(
|
||||||
monkeypatch, tmp_path: Path
|
monkeypatch, tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
config_file = tmp_path / "instance" / "config.json"
|
config_file = _write_instance_config(tmp_path)
|
||||||
config_file.parent.mkdir(parents=True)
|
|
||||||
config_file.write_text("{}")
|
|
||||||
|
|
||||||
legacy_dir = tmp_path / "global" / "cron"
|
legacy_dir = tmp_path / "global" / "cron"
|
||||||
legacy_dir.mkdir(parents=True)
|
legacy_dir.mkdir(parents=True)
|
||||||
legacy_file = legacy_dir / "jobs.json"
|
legacy_file = legacy_dir / "jobs.json"
|
||||||
@@ -748,20 +812,19 @@ def test_gateway_workspace_override_does_not_migrate_legacy_cron(
|
|||||||
config = Config()
|
config = Config()
|
||||||
seen: dict[str, Path] = {}
|
seen: dict[str, Path] = {}
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.set_config_path", lambda _path: None)
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
|
||||||
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
|
||||||
monkeypatch.setattr("nanobot.cli.commands._make_provider", lambda _config: object())
|
|
||||||
monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: object())
|
|
||||||
monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object())
|
|
||||||
monkeypatch.setattr("nanobot.config.paths.get_cron_dir", lambda: legacy_dir)
|
|
||||||
|
|
||||||
class _StopCron:
|
class _StopCron:
|
||||||
def __init__(self, store_path: Path) -> None:
|
def __init__(self, store_path: Path) -> None:
|
||||||
seen["cron_store"] = store_path
|
seen["cron_store"] = store_path
|
||||||
raise _StopGatewayError("stop")
|
raise _StopGatewayError("stop")
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.cron.service.CronService", _StopCron)
|
_patch_cli_command_runtime(
|
||||||
|
monkeypatch,
|
||||||
|
config,
|
||||||
|
message_bus=lambda: object(),
|
||||||
|
session_manager=lambda _workspace: object(),
|
||||||
|
cron_service=_StopCron,
|
||||||
|
get_cron_dir=lambda: legacy_dir,
|
||||||
|
)
|
||||||
|
|
||||||
result = runner.invoke(
|
result = runner.invoke(
|
||||||
app,
|
app,
|
||||||
@@ -777,10 +840,7 @@ def test_gateway_workspace_override_does_not_migrate_legacy_cron(
|
|||||||
def test_gateway_custom_config_workspace_does_not_migrate_legacy_cron(
|
def test_gateway_custom_config_workspace_does_not_migrate_legacy_cron(
|
||||||
monkeypatch, tmp_path: Path
|
monkeypatch, tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
config_file = tmp_path / "instance" / "config.json"
|
config_file = _write_instance_config(tmp_path)
|
||||||
config_file.parent.mkdir(parents=True)
|
|
||||||
config_file.write_text("{}")
|
|
||||||
|
|
||||||
legacy_dir = tmp_path / "global" / "cron"
|
legacy_dir = tmp_path / "global" / "cron"
|
||||||
legacy_dir.mkdir(parents=True)
|
legacy_dir.mkdir(parents=True)
|
||||||
legacy_file = legacy_dir / "jobs.json"
|
legacy_file = legacy_dir / "jobs.json"
|
||||||
@@ -791,20 +851,19 @@ def test_gateway_custom_config_workspace_does_not_migrate_legacy_cron(
|
|||||||
config.agents.defaults.workspace = str(custom_workspace)
|
config.agents.defaults.workspace = str(custom_workspace)
|
||||||
seen: dict[str, Path] = {}
|
seen: dict[str, Path] = {}
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.set_config_path", lambda _path: None)
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
|
||||||
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
|
||||||
monkeypatch.setattr("nanobot.cli.commands._make_provider", lambda _config: object())
|
|
||||||
monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: object())
|
|
||||||
monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object())
|
|
||||||
monkeypatch.setattr("nanobot.config.paths.get_cron_dir", lambda: legacy_dir)
|
|
||||||
|
|
||||||
class _StopCron:
|
class _StopCron:
|
||||||
def __init__(self, store_path: Path) -> None:
|
def __init__(self, store_path: Path) -> None:
|
||||||
seen["cron_store"] = store_path
|
seen["cron_store"] = store_path
|
||||||
raise _StopGatewayError("stop")
|
raise _StopGatewayError("stop")
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.cron.service.CronService", _StopCron)
|
_patch_cli_command_runtime(
|
||||||
|
monkeypatch,
|
||||||
|
config,
|
||||||
|
message_bus=lambda: object(),
|
||||||
|
session_manager=lambda _workspace: object(),
|
||||||
|
cron_service=_StopCron,
|
||||||
|
get_cron_dir=lambda: legacy_dir,
|
||||||
|
)
|
||||||
|
|
||||||
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
||||||
|
|
||||||
@@ -856,19 +915,14 @@ def test_migrate_cron_store_skips_when_workspace_file_exists(tmp_path: Path) ->
|
|||||||
|
|
||||||
|
|
||||||
def test_gateway_uses_configured_port_when_cli_flag_is_missing(monkeypatch, tmp_path: Path) -> None:
|
def test_gateway_uses_configured_port_when_cli_flag_is_missing(monkeypatch, tmp_path: Path) -> None:
|
||||||
config_file = tmp_path / "instance" / "config.json"
|
config_file = _write_instance_config(tmp_path)
|
||||||
config_file.parent.mkdir(parents=True)
|
|
||||||
config_file.write_text("{}")
|
|
||||||
|
|
||||||
config = Config()
|
config = Config()
|
||||||
config.gateway.port = 18791
|
config.gateway.port = 18791
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.set_config_path", lambda _path: None)
|
_patch_cli_command_runtime(
|
||||||
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
monkeypatch,
|
||||||
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
config,
|
||||||
monkeypatch.setattr(
|
make_provider=_stop_gateway_provider,
|
||||||
"nanobot.cli.commands._make_provider",
|
|
||||||
lambda _config: (_ for _ in ()).throw(_StopGatewayError("stop")),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
||||||
@@ -878,19 +932,14 @@ def test_gateway_uses_configured_port_when_cli_flag_is_missing(monkeypatch, tmp_
|
|||||||
|
|
||||||
|
|
||||||
def test_gateway_cli_port_overrides_configured_port(monkeypatch, tmp_path: Path) -> None:
|
def test_gateway_cli_port_overrides_configured_port(monkeypatch, tmp_path: Path) -> None:
|
||||||
config_file = tmp_path / "instance" / "config.json"
|
config_file = _write_instance_config(tmp_path)
|
||||||
config_file.parent.mkdir(parents=True)
|
|
||||||
config_file.write_text("{}")
|
|
||||||
|
|
||||||
config = Config()
|
config = Config()
|
||||||
config.gateway.port = 18791
|
config.gateway.port = 18791
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.config.loader.set_config_path", lambda _path: None)
|
_patch_cli_command_runtime(
|
||||||
monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config)
|
monkeypatch,
|
||||||
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
config,
|
||||||
monkeypatch.setattr(
|
make_provider=_stop_gateway_provider,
|
||||||
"nanobot.cli.commands._make_provider",
|
|
||||||
lambda _config: (_ for _ in ()).throw(_StopGatewayError("stop")),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
result = runner.invoke(app, ["gateway", "--config", str(config_file), "--port", "18792"])
|
result = runner.invoke(app, ["gateway", "--config", str(config_file), "--port", "18792"])
|
||||||
@@ -899,6 +948,63 @@ def test_gateway_cli_port_overrides_configured_port(monkeypatch, tmp_path: Path)
|
|||||||
assert "port 18792" in result.stdout
|
assert "port 18792" in result.stdout
|
||||||
|
|
||||||
|
|
||||||
|
def test_serve_uses_api_config_defaults_and_workspace_override(
|
||||||
|
monkeypatch, tmp_path: Path
|
||||||
|
) -> None:
|
||||||
|
config_file = _write_instance_config(tmp_path)
|
||||||
|
config = Config()
|
||||||
|
config.agents.defaults.workspace = str(tmp_path / "config-workspace")
|
||||||
|
config.api.host = "127.0.0.2"
|
||||||
|
config.api.port = 18900
|
||||||
|
config.api.timeout = 45.0
|
||||||
|
override_workspace = tmp_path / "override-workspace"
|
||||||
|
seen: dict[str, object] = {}
|
||||||
|
|
||||||
|
_patch_serve_runtime(monkeypatch, config, seen)
|
||||||
|
|
||||||
|
result = runner.invoke(
|
||||||
|
app,
|
||||||
|
["serve", "--config", str(config_file), "--workspace", str(override_workspace)],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.exit_code == 0
|
||||||
|
assert seen["workspace"] == override_workspace
|
||||||
|
assert seen["host"] == "127.0.0.2"
|
||||||
|
assert seen["port"] == 18900
|
||||||
|
assert seen["request_timeout"] == 45.0
|
||||||
|
|
||||||
|
|
||||||
|
def test_serve_cli_options_override_api_config(monkeypatch, tmp_path: Path) -> None:
|
||||||
|
config_file = _write_instance_config(tmp_path)
|
||||||
|
config = Config()
|
||||||
|
config.api.host = "127.0.0.2"
|
||||||
|
config.api.port = 18900
|
||||||
|
config.api.timeout = 45.0
|
||||||
|
seen: dict[str, object] = {}
|
||||||
|
|
||||||
|
_patch_serve_runtime(monkeypatch, config, seen)
|
||||||
|
|
||||||
|
result = runner.invoke(
|
||||||
|
app,
|
||||||
|
[
|
||||||
|
"serve",
|
||||||
|
"--config",
|
||||||
|
str(config_file),
|
||||||
|
"--host",
|
||||||
|
"127.0.0.1",
|
||||||
|
"--port",
|
||||||
|
"18901",
|
||||||
|
"--timeout",
|
||||||
|
"46",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.exit_code == 0
|
||||||
|
assert seen["host"] == "127.0.0.1"
|
||||||
|
assert seen["port"] == 18901
|
||||||
|
assert seen["request_timeout"] == 46.0
|
||||||
|
|
||||||
|
|
||||||
def test_channels_login_requires_channel_name() -> None:
|
def test_channels_login_requires_channel_name() -> None:
|
||||||
result = runner.invoke(app, ["channels", "login"])
|
result = runner.invoke(app, ["channels", "login"])
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
"""Tests for CronTool._list_jobs() output formatting."""
|
"""Tests for CronTool._list_jobs() output formatting."""
|
||||||
|
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from nanobot.agent.tools.cron import CronTool
|
from nanobot.agent.tools.cron import CronTool
|
||||||
from nanobot.cron.service import CronService
|
from nanobot.cron.service import CronService
|
||||||
from nanobot.cron.types import CronJobState, CronSchedule
|
from nanobot.cron.types import CronJobState, CronSchedule
|
||||||
@@ -10,99 +12,120 @@ def _make_tool(tmp_path) -> CronTool:
|
|||||||
return CronTool(service)
|
return CronTool(service)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_tool_with_tz(tmp_path, tz: str) -> CronTool:
|
||||||
|
service = CronService(tmp_path / "cron" / "jobs.json")
|
||||||
|
return CronTool(service, default_timezone=tz)
|
||||||
|
|
||||||
|
|
||||||
# -- _format_timing tests --
|
# -- _format_timing tests --
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_cron_with_tz() -> None:
|
def test_format_timing_cron_with_tz(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
s = CronSchedule(kind="cron", expr="0 9 * * 1-5", tz="America/Denver")
|
s = CronSchedule(kind="cron", expr="0 9 * * 1-5", tz="America/Denver")
|
||||||
assert CronTool._format_timing(s) == "cron: 0 9 * * 1-5 (America/Denver)"
|
assert tool._format_timing(s) == "cron: 0 9 * * 1-5 (America/Denver)"
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_cron_without_tz() -> None:
|
def test_format_timing_cron_without_tz(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
s = CronSchedule(kind="cron", expr="*/5 * * * *")
|
s = CronSchedule(kind="cron", expr="*/5 * * * *")
|
||||||
assert CronTool._format_timing(s) == "cron: */5 * * * *"
|
assert tool._format_timing(s) == "cron: */5 * * * *"
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_every_hours() -> None:
|
def test_format_timing_every_hours(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
s = CronSchedule(kind="every", every_ms=7_200_000)
|
s = CronSchedule(kind="every", every_ms=7_200_000)
|
||||||
assert CronTool._format_timing(s) == "every 2h"
|
assert tool._format_timing(s) == "every 2h"
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_every_minutes() -> None:
|
def test_format_timing_every_minutes(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
s = CronSchedule(kind="every", every_ms=1_800_000)
|
s = CronSchedule(kind="every", every_ms=1_800_000)
|
||||||
assert CronTool._format_timing(s) == "every 30m"
|
assert tool._format_timing(s) == "every 30m"
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_every_seconds() -> None:
|
def test_format_timing_every_seconds(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
s = CronSchedule(kind="every", every_ms=30_000)
|
s = CronSchedule(kind="every", every_ms=30_000)
|
||||||
assert CronTool._format_timing(s) == "every 30s"
|
assert tool._format_timing(s) == "every 30s"
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_every_non_minute_seconds() -> None:
|
def test_format_timing_every_non_minute_seconds(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
s = CronSchedule(kind="every", every_ms=90_000)
|
s = CronSchedule(kind="every", every_ms=90_000)
|
||||||
assert CronTool._format_timing(s) == "every 90s"
|
assert tool._format_timing(s) == "every 90s"
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_every_milliseconds() -> None:
|
def test_format_timing_every_milliseconds(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
s = CronSchedule(kind="every", every_ms=200)
|
s = CronSchedule(kind="every", every_ms=200)
|
||||||
assert CronTool._format_timing(s) == "every 200ms"
|
assert tool._format_timing(s) == "every 200ms"
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_at() -> None:
|
def test_format_timing_at(tmp_path) -> None:
|
||||||
|
tool = _make_tool_with_tz(tmp_path, "Asia/Shanghai")
|
||||||
s = CronSchedule(kind="at", at_ms=1773684000000)
|
s = CronSchedule(kind="at", at_ms=1773684000000)
|
||||||
result = CronTool._format_timing(s)
|
result = tool._format_timing(s)
|
||||||
|
assert "Asia/Shanghai" in result
|
||||||
assert result.startswith("at 2026-")
|
assert result.startswith("at 2026-")
|
||||||
|
|
||||||
|
|
||||||
def test_format_timing_fallback() -> None:
|
def test_format_timing_fallback(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
s = CronSchedule(kind="every") # no every_ms
|
s = CronSchedule(kind="every") # no every_ms
|
||||||
assert CronTool._format_timing(s) == "every"
|
assert tool._format_timing(s) == "every"
|
||||||
|
|
||||||
|
|
||||||
# -- _format_state tests --
|
# -- _format_state tests --
|
||||||
|
|
||||||
|
|
||||||
def test_format_state_empty() -> None:
|
def test_format_state_empty(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
state = CronJobState()
|
state = CronJobState()
|
||||||
assert CronTool._format_state(state) == []
|
assert tool._format_state(state, CronSchedule(kind="every")) == []
|
||||||
|
|
||||||
|
|
||||||
def test_format_state_last_run_ok() -> None:
|
def test_format_state_last_run_ok(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
state = CronJobState(last_run_at_ms=1773673200000, last_status="ok")
|
state = CronJobState(last_run_at_ms=1773673200000, last_status="ok")
|
||||||
lines = CronTool._format_state(state)
|
lines = tool._format_state(state, CronSchedule(kind="cron", expr="0 9 * * *", tz="UTC"))
|
||||||
assert len(lines) == 1
|
assert len(lines) == 1
|
||||||
assert "Last run:" in lines[0]
|
assert "Last run:" in lines[0]
|
||||||
assert "ok" in lines[0]
|
assert "ok" in lines[0]
|
||||||
|
|
||||||
|
|
||||||
def test_format_state_last_run_with_error() -> None:
|
def test_format_state_last_run_with_error(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
state = CronJobState(last_run_at_ms=1773673200000, last_status="error", last_error="timeout")
|
state = CronJobState(last_run_at_ms=1773673200000, last_status="error", last_error="timeout")
|
||||||
lines = CronTool._format_state(state)
|
lines = tool._format_state(state, CronSchedule(kind="cron", expr="0 9 * * *", tz="UTC"))
|
||||||
assert len(lines) == 1
|
assert len(lines) == 1
|
||||||
assert "error" in lines[0]
|
assert "error" in lines[0]
|
||||||
assert "timeout" in lines[0]
|
assert "timeout" in lines[0]
|
||||||
|
|
||||||
|
|
||||||
def test_format_state_next_run_only() -> None:
|
def test_format_state_next_run_only(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
state = CronJobState(next_run_at_ms=1773684000000)
|
state = CronJobState(next_run_at_ms=1773684000000)
|
||||||
lines = CronTool._format_state(state)
|
lines = tool._format_state(state, CronSchedule(kind="cron", expr="0 9 * * *", tz="UTC"))
|
||||||
assert len(lines) == 1
|
assert len(lines) == 1
|
||||||
assert "Next run:" in lines[0]
|
assert "Next run:" in lines[0]
|
||||||
|
|
||||||
|
|
||||||
def test_format_state_both() -> None:
|
def test_format_state_both(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
state = CronJobState(
|
state = CronJobState(
|
||||||
last_run_at_ms=1773673200000, last_status="ok", next_run_at_ms=1773684000000
|
last_run_at_ms=1773673200000, last_status="ok", next_run_at_ms=1773684000000
|
||||||
)
|
)
|
||||||
lines = CronTool._format_state(state)
|
lines = tool._format_state(state, CronSchedule(kind="cron", expr="0 9 * * *", tz="UTC"))
|
||||||
assert len(lines) == 2
|
assert len(lines) == 2
|
||||||
assert "Last run:" in lines[0]
|
assert "Last run:" in lines[0]
|
||||||
assert "Next run:" in lines[1]
|
assert "Next run:" in lines[1]
|
||||||
|
|
||||||
|
|
||||||
def test_format_state_unknown_status() -> None:
|
def test_format_state_unknown_status(tmp_path) -> None:
|
||||||
|
tool = _make_tool(tmp_path)
|
||||||
state = CronJobState(last_run_at_ms=1773673200000, last_status=None)
|
state = CronJobState(last_run_at_ms=1773673200000, last_status=None)
|
||||||
lines = CronTool._format_state(state)
|
lines = tool._format_state(state, CronSchedule(kind="cron", expr="0 9 * * *", tz="UTC"))
|
||||||
assert "unknown" in lines[0]
|
assert "unknown" in lines[0]
|
||||||
|
|
||||||
|
|
||||||
@@ -181,7 +204,7 @@ def test_list_every_job_milliseconds(tmp_path) -> None:
|
|||||||
|
|
||||||
|
|
||||||
def test_list_at_job_shows_iso_timestamp(tmp_path) -> None:
|
def test_list_at_job_shows_iso_timestamp(tmp_path) -> None:
|
||||||
tool = _make_tool(tmp_path)
|
tool = _make_tool_with_tz(tmp_path, "Asia/Shanghai")
|
||||||
tool._cron.add_job(
|
tool._cron.add_job(
|
||||||
name="One-shot",
|
name="One-shot",
|
||||||
schedule=CronSchedule(kind="at", at_ms=1773684000000),
|
schedule=CronSchedule(kind="at", at_ms=1773684000000),
|
||||||
@@ -189,6 +212,7 @@ def test_list_at_job_shows_iso_timestamp(tmp_path) -> None:
|
|||||||
)
|
)
|
||||||
result = tool._list_jobs()
|
result = tool._list_jobs()
|
||||||
assert "at 2026-" in result
|
assert "at 2026-" in result
|
||||||
|
assert "Asia/Shanghai" in result
|
||||||
|
|
||||||
|
|
||||||
def test_list_shows_last_run_state(tmp_path) -> None:
|
def test_list_shows_last_run_state(tmp_path) -> None:
|
||||||
@@ -206,6 +230,7 @@ def test_list_shows_last_run_state(tmp_path) -> None:
|
|||||||
result = tool._list_jobs()
|
result = tool._list_jobs()
|
||||||
assert "Last run:" in result
|
assert "Last run:" in result
|
||||||
assert "ok" in result
|
assert "ok" in result
|
||||||
|
assert "(UTC)" in result
|
||||||
|
|
||||||
|
|
||||||
def test_list_shows_error_message(tmp_path) -> None:
|
def test_list_shows_error_message(tmp_path) -> None:
|
||||||
@@ -234,6 +259,30 @@ def test_list_shows_next_run(tmp_path) -> None:
|
|||||||
)
|
)
|
||||||
result = tool._list_jobs()
|
result = tool._list_jobs()
|
||||||
assert "Next run:" in result
|
assert "Next run:" in result
|
||||||
|
assert "(UTC)" in result
|
||||||
|
|
||||||
|
|
||||||
|
def test_add_cron_job_defaults_to_tool_timezone(tmp_path) -> None:
|
||||||
|
tool = _make_tool_with_tz(tmp_path, "Asia/Shanghai")
|
||||||
|
tool.set_context("telegram", "chat-1")
|
||||||
|
|
||||||
|
result = tool._add_job("Morning standup", None, "0 8 * * *", None, None)
|
||||||
|
|
||||||
|
assert result.startswith("Created job")
|
||||||
|
job = tool._cron.list_jobs()[0]
|
||||||
|
assert job.schedule.tz == "Asia/Shanghai"
|
||||||
|
|
||||||
|
|
||||||
|
def test_add_at_job_uses_default_timezone_for_naive_datetime(tmp_path) -> None:
|
||||||
|
tool = _make_tool_with_tz(tmp_path, "Asia/Shanghai")
|
||||||
|
tool.set_context("telegram", "chat-1")
|
||||||
|
|
||||||
|
result = tool._add_job("Morning reminder", None, None, None, "2026-03-25T08:00:00")
|
||||||
|
|
||||||
|
assert result.startswith("Created job")
|
||||||
|
job = tool._cron.list_jobs()[0]
|
||||||
|
expected = int(datetime(2026, 3, 25, 0, 0, 0, tzinfo=timezone.utc).timestamp() * 1000)
|
||||||
|
assert job.schedule.at_ms == expected
|
||||||
|
|
||||||
|
|
||||||
def test_list_excludes_disabled_jobs(tmp_path) -> None:
|
def test_list_excludes_disabled_jobs(tmp_path) -> None:
|
||||||
|
|||||||
@@ -60,6 +60,45 @@ 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_openrouter_sets_default_attribution_headers() -> None:
|
||||||
|
spec = find_by_name("openrouter")
|
||||||
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient:
|
||||||
|
OpenAICompatProvider(
|
||||||
|
api_key="sk-or-test-key",
|
||||||
|
api_base="https://openrouter.ai/api/v1",
|
||||||
|
default_model="anthropic/claude-sonnet-4-5",
|
||||||
|
spec=spec,
|
||||||
|
)
|
||||||
|
|
||||||
|
headers = MockClient.call_args.kwargs["default_headers"]
|
||||||
|
assert headers["HTTP-Referer"] == "https://github.com/HKUDS/nanobot"
|
||||||
|
assert headers["X-OpenRouter-Title"] == "nanobot"
|
||||||
|
assert headers["X-OpenRouter-Categories"] == "cli-agent,personal-agent"
|
||||||
|
assert "x-session-affinity" in headers
|
||||||
|
|
||||||
|
|
||||||
|
def test_openrouter_user_headers_override_default_attribution() -> None:
|
||||||
|
spec = find_by_name("openrouter")
|
||||||
|
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient:
|
||||||
|
OpenAICompatProvider(
|
||||||
|
api_key="sk-or-test-key",
|
||||||
|
api_base="https://openrouter.ai/api/v1",
|
||||||
|
default_model="anthropic/claude-sonnet-4-5",
|
||||||
|
extra_headers={
|
||||||
|
"HTTP-Referer": "https://nanobot.ai",
|
||||||
|
"X-OpenRouter-Title": "Nanobot Pro",
|
||||||
|
"X-Custom-App": "enabled",
|
||||||
|
},
|
||||||
|
spec=spec,
|
||||||
|
)
|
||||||
|
|
||||||
|
headers = MockClient.call_args.kwargs["default_headers"]
|
||||||
|
assert headers["HTTP-Referer"] == "https://nanobot.ai"
|
||||||
|
assert headers["X-OpenRouter-Title"] == "Nanobot Pro"
|
||||||
|
assert headers["X-OpenRouter-Categories"] == "cli-agent,personal-agent"
|
||||||
|
assert headers["X-Custom-App"] == "enabled"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_openrouter_keeps_model_name_intact() -> None:
|
async def test_openrouter_keeps_model_name_intact() -> None:
|
||||||
"""OpenRouter gateway keeps the full model name (gateway does its own routing)."""
|
"""OpenRouter gateway keeps the full model name (gateway does its own routing)."""
|
||||||
|
|||||||
@@ -0,0 +1,147 @@
|
|||||||
|
"""Tests for the Nanobot programmatic facade."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.nanobot import Nanobot, RunResult
|
||||||
|
|
||||||
|
|
||||||
|
def _write_config(tmp_path: Path, overrides: dict | None = None) -> Path:
|
||||||
|
data = {
|
||||||
|
"providers": {"openrouter": {"apiKey": "sk-test-key"}},
|
||||||
|
"agents": {"defaults": {"model": "openai/gpt-4.1"}},
|
||||||
|
}
|
||||||
|
if overrides:
|
||||||
|
data.update(overrides)
|
||||||
|
config_path = tmp_path / "config.json"
|
||||||
|
config_path.write_text(json.dumps(data))
|
||||||
|
return config_path
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_config_missing_file():
|
||||||
|
with pytest.raises(FileNotFoundError):
|
||||||
|
Nanobot.from_config("/nonexistent/config.json")
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_config_creates_instance(tmp_path):
|
||||||
|
config_path = _write_config(tmp_path)
|
||||||
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
|
assert bot._loop is not None
|
||||||
|
assert bot._loop.workspace == tmp_path
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_config_default_path():
|
||||||
|
from nanobot.config.schema import Config
|
||||||
|
|
||||||
|
with patch("nanobot.config.loader.load_config") as mock_load, \
|
||||||
|
patch("nanobot.nanobot._make_provider") as mock_prov:
|
||||||
|
mock_load.return_value = Config()
|
||||||
|
mock_prov.return_value = MagicMock()
|
||||||
|
mock_prov.return_value.get_default_model.return_value = "test"
|
||||||
|
mock_prov.return_value.generation.max_tokens = 4096
|
||||||
|
Nanobot.from_config()
|
||||||
|
mock_load.assert_called_once_with(None)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_returns_result(tmp_path):
|
||||||
|
config_path = _write_config(tmp_path)
|
||||||
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
|
|
||||||
|
from nanobot.bus.events import OutboundMessage
|
||||||
|
|
||||||
|
mock_response = OutboundMessage(
|
||||||
|
channel="cli", chat_id="direct", content="Hello back!"
|
||||||
|
)
|
||||||
|
bot._loop.process_direct = AsyncMock(return_value=mock_response)
|
||||||
|
|
||||||
|
result = await bot.run("hi")
|
||||||
|
|
||||||
|
assert isinstance(result, RunResult)
|
||||||
|
assert result.content == "Hello back!"
|
||||||
|
bot._loop.process_direct.assert_awaited_once_with("hi", session_key="sdk:default")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_with_hooks(tmp_path):
|
||||||
|
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||||
|
from nanobot.bus.events import OutboundMessage
|
||||||
|
|
||||||
|
config_path = _write_config(tmp_path)
|
||||||
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
|
|
||||||
|
class TestHook(AgentHook):
|
||||||
|
async def before_iteration(self, context: AgentHookContext) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
mock_response = OutboundMessage(
|
||||||
|
channel="cli", chat_id="direct", content="done"
|
||||||
|
)
|
||||||
|
bot._loop.process_direct = AsyncMock(return_value=mock_response)
|
||||||
|
|
||||||
|
result = await bot.run("hi", hooks=[TestHook()])
|
||||||
|
|
||||||
|
assert result.content == "done"
|
||||||
|
assert bot._loop._extra_hooks == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_hooks_restored_on_error(tmp_path):
|
||||||
|
config_path = _write_config(tmp_path)
|
||||||
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
|
|
||||||
|
from nanobot.agent.hook import AgentHook
|
||||||
|
|
||||||
|
bot._loop.process_direct = AsyncMock(side_effect=RuntimeError("boom"))
|
||||||
|
original_hooks = bot._loop._extra_hooks
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError):
|
||||||
|
await bot.run("hi", hooks=[AgentHook()])
|
||||||
|
|
||||||
|
assert bot._loop._extra_hooks is original_hooks
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_none_response(tmp_path):
|
||||||
|
config_path = _write_config(tmp_path)
|
||||||
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
|
bot._loop.process_direct = AsyncMock(return_value=None)
|
||||||
|
|
||||||
|
result = await bot.run("hi")
|
||||||
|
assert result.content == ""
|
||||||
|
|
||||||
|
|
||||||
|
def test_workspace_override(tmp_path):
|
||||||
|
config_path = _write_config(tmp_path)
|
||||||
|
custom_ws = tmp_path / "custom_workspace"
|
||||||
|
custom_ws.mkdir()
|
||||||
|
|
||||||
|
bot = Nanobot.from_config(config_path, workspace=custom_ws)
|
||||||
|
assert bot._loop.workspace == custom_ws
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_custom_session_key(tmp_path):
|
||||||
|
from nanobot.bus.events import OutboundMessage
|
||||||
|
|
||||||
|
config_path = _write_config(tmp_path)
|
||||||
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
|
|
||||||
|
mock_response = OutboundMessage(
|
||||||
|
channel="cli", chat_id="direct", content="ok"
|
||||||
|
)
|
||||||
|
bot._loop.process_direct = AsyncMock(return_value=mock_response)
|
||||||
|
|
||||||
|
await bot.run("hi", session_key="user-alice")
|
||||||
|
bot._loop.process_direct.assert_awaited_once_with("hi", session_key="user-alice")
|
||||||
|
|
||||||
|
|
||||||
|
def test_import_from_top_level():
|
||||||
|
from nanobot import Nanobot as N, RunResult as R
|
||||||
|
assert N is Nanobot
|
||||||
|
assert R is RunResult
|
||||||
@@ -0,0 +1,372 @@
|
|||||||
|
"""Focused tests for the fixed-session OpenAI-compatible API."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import pytest_asyncio
|
||||||
|
|
||||||
|
from nanobot.api.server import (
|
||||||
|
API_CHAT_ID,
|
||||||
|
API_SESSION_KEY,
|
||||||
|
_chat_completion_response,
|
||||||
|
_error_json,
|
||||||
|
create_app,
|
||||||
|
handle_chat_completions,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
from aiohttp.test_utils import TestClient, TestServer
|
||||||
|
|
||||||
|
HAS_AIOHTTP = True
|
||||||
|
except ImportError:
|
||||||
|
HAS_AIOHTTP = False
|
||||||
|
|
||||||
|
pytest_plugins = ("pytest_asyncio",)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_mock_agent(response_text: str = "mock response") -> MagicMock:
|
||||||
|
agent = MagicMock()
|
||||||
|
agent.process_direct = AsyncMock(return_value=response_text)
|
||||||
|
agent._connect_mcp = AsyncMock()
|
||||||
|
agent.close_mcp = AsyncMock()
|
||||||
|
return agent
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def mock_agent():
|
||||||
|
return _make_mock_agent()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def app(mock_agent):
|
||||||
|
return create_app(mock_agent, model_name="test-model", request_timeout=10.0)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
async def aiohttp_client():
|
||||||
|
clients: list[TestClient] = []
|
||||||
|
|
||||||
|
async def _make_client(app):
|
||||||
|
client = TestClient(TestServer(app))
|
||||||
|
await client.start_server()
|
||||||
|
clients.append(client)
|
||||||
|
return client
|
||||||
|
|
||||||
|
try:
|
||||||
|
yield _make_client
|
||||||
|
finally:
|
||||||
|
for client in clients:
|
||||||
|
await client.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_error_json() -> None:
|
||||||
|
resp = _error_json(400, "bad request")
|
||||||
|
assert resp.status == 400
|
||||||
|
body = json.loads(resp.body)
|
||||||
|
assert body["error"]["message"] == "bad request"
|
||||||
|
assert body["error"]["code"] == 400
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_completion_response() -> None:
|
||||||
|
result = _chat_completion_response("hello world", "test-model")
|
||||||
|
assert result["object"] == "chat.completion"
|
||||||
|
assert result["model"] == "test-model"
|
||||||
|
assert result["choices"][0]["message"]["content"] == "hello world"
|
||||||
|
assert result["choices"][0]["finish_reason"] == "stop"
|
||||||
|
assert result["id"].startswith("chatcmpl-")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_missing_messages_returns_400(aiohttp_client, app) -> None:
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.post("/v1/chat/completions", json={"model": "test"})
|
||||||
|
assert resp.status == 400
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_no_user_message_returns_400(aiohttp_client, app) -> None:
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={"messages": [{"role": "system", "content": "you are a bot"}]},
|
||||||
|
)
|
||||||
|
assert resp.status == 400
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stream_true_returns_400(aiohttp_client, app) -> None:
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={"messages": [{"role": "user", "content": "hello"}], "stream": True},
|
||||||
|
)
|
||||||
|
assert resp.status == 400
|
||||||
|
body = await resp.json()
|
||||||
|
assert "stream" in body["error"]["message"].lower()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_model_mismatch_returns_400() -> None:
|
||||||
|
request = MagicMock()
|
||||||
|
request.json = AsyncMock(
|
||||||
|
return_value={
|
||||||
|
"model": "other-model",
|
||||||
|
"messages": [{"role": "user", "content": "hello"}],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
request.app = {
|
||||||
|
"agent_loop": _make_mock_agent(),
|
||||||
|
"model_name": "test-model",
|
||||||
|
"request_timeout": 10.0,
|
||||||
|
"session_lock": asyncio.Lock(),
|
||||||
|
}
|
||||||
|
|
||||||
|
resp = await handle_chat_completions(request)
|
||||||
|
assert resp.status == 400
|
||||||
|
body = json.loads(resp.body)
|
||||||
|
assert "test-model" in body["error"]["message"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_single_user_message_required() -> None:
|
||||||
|
request = MagicMock()
|
||||||
|
request.json = AsyncMock(
|
||||||
|
return_value={
|
||||||
|
"messages": [
|
||||||
|
{"role": "user", "content": "hello"},
|
||||||
|
{"role": "assistant", "content": "previous reply"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
request.app = {
|
||||||
|
"agent_loop": _make_mock_agent(),
|
||||||
|
"model_name": "test-model",
|
||||||
|
"request_timeout": 10.0,
|
||||||
|
"session_lock": asyncio.Lock(),
|
||||||
|
}
|
||||||
|
|
||||||
|
resp = await handle_chat_completions(request)
|
||||||
|
assert resp.status == 400
|
||||||
|
body = json.loads(resp.body)
|
||||||
|
assert "single user message" in body["error"]["message"].lower()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_single_user_message_must_have_user_role() -> None:
|
||||||
|
request = MagicMock()
|
||||||
|
request.json = AsyncMock(
|
||||||
|
return_value={
|
||||||
|
"messages": [{"role": "system", "content": "you are a bot"}],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
request.app = {
|
||||||
|
"agent_loop": _make_mock_agent(),
|
||||||
|
"model_name": "test-model",
|
||||||
|
"request_timeout": 10.0,
|
||||||
|
"session_lock": asyncio.Lock(),
|
||||||
|
}
|
||||||
|
|
||||||
|
resp = await handle_chat_completions(request)
|
||||||
|
assert resp.status == 400
|
||||||
|
body = json.loads(resp.body)
|
||||||
|
assert "single user message" in body["error"]["message"].lower()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_successful_request_uses_fixed_api_session(aiohttp_client, mock_agent) -> None:
|
||||||
|
app = create_app(mock_agent, model_name="test-model")
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={"messages": [{"role": "user", "content": "hello"}]},
|
||||||
|
)
|
||||||
|
assert resp.status == 200
|
||||||
|
body = await resp.json()
|
||||||
|
assert body["choices"][0]["message"]["content"] == "mock response"
|
||||||
|
assert body["model"] == "test-model"
|
||||||
|
mock_agent.process_direct.assert_called_once_with(
|
||||||
|
content="hello",
|
||||||
|
session_key=API_SESSION_KEY,
|
||||||
|
channel="api",
|
||||||
|
chat_id=API_CHAT_ID,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_followup_requests_share_same_session_key(aiohttp_client) -> None:
|
||||||
|
call_log: list[str] = []
|
||||||
|
|
||||||
|
async def fake_process(content, session_key="", channel="", chat_id=""):
|
||||||
|
call_log.append(session_key)
|
||||||
|
return f"reply to {content}"
|
||||||
|
|
||||||
|
agent = MagicMock()
|
||||||
|
agent.process_direct = fake_process
|
||||||
|
agent._connect_mcp = AsyncMock()
|
||||||
|
agent.close_mcp = AsyncMock()
|
||||||
|
|
||||||
|
app = create_app(agent, model_name="m")
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
|
||||||
|
r1 = await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={"messages": [{"role": "user", "content": "first"}]},
|
||||||
|
)
|
||||||
|
r2 = await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={"messages": [{"role": "user", "content": "second"}]},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert r1.status == 200
|
||||||
|
assert r2.status == 200
|
||||||
|
assert call_log == [API_SESSION_KEY, API_SESSION_KEY]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_fixed_session_requests_are_serialized(aiohttp_client) -> None:
|
||||||
|
order: list[str] = []
|
||||||
|
barrier = asyncio.Event()
|
||||||
|
|
||||||
|
async def slow_process(content, session_key="", channel="", chat_id=""):
|
||||||
|
order.append(f"start:{content}")
|
||||||
|
if content == "first":
|
||||||
|
barrier.set()
|
||||||
|
await asyncio.sleep(0.1)
|
||||||
|
else:
|
||||||
|
await barrier.wait()
|
||||||
|
order.append(f"end:{content}")
|
||||||
|
return content
|
||||||
|
|
||||||
|
agent = MagicMock()
|
||||||
|
agent.process_direct = slow_process
|
||||||
|
agent._connect_mcp = AsyncMock()
|
||||||
|
agent.close_mcp = AsyncMock()
|
||||||
|
|
||||||
|
app = create_app(agent, model_name="m")
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
|
||||||
|
async def send(msg: str):
|
||||||
|
return await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={"messages": [{"role": "user", "content": msg}]},
|
||||||
|
)
|
||||||
|
|
||||||
|
r1, r2 = await asyncio.gather(send("first"), send("second"))
|
||||||
|
assert r1.status == 200
|
||||||
|
assert r2.status == 200
|
||||||
|
assert order.index("end:first") < order.index("start:second")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_models_endpoint(aiohttp_client, app) -> None:
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.get("/v1/models")
|
||||||
|
assert resp.status == 200
|
||||||
|
body = await resp.json()
|
||||||
|
assert body["object"] == "list"
|
||||||
|
assert body["data"][0]["id"] == "test-model"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_health_endpoint(aiohttp_client, app) -> None:
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.get("/health")
|
||||||
|
assert resp.status == 200
|
||||||
|
body = await resp.json()
|
||||||
|
assert body["status"] == "ok"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_multimodal_content_extracts_text(aiohttp_client, mock_agent) -> None:
|
||||||
|
app = create_app(mock_agent, model_name="m")
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "text", "text": "describe this"},
|
||||||
|
{"type": "image_url", "image_url": {"url": "data:image/png;base64,abc"}},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert resp.status == 200
|
||||||
|
mock_agent.process_direct.assert_called_once_with(
|
||||||
|
content="describe this",
|
||||||
|
session_key=API_SESSION_KEY,
|
||||||
|
channel="api",
|
||||||
|
chat_id=API_CHAT_ID,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_empty_response_retry_then_success(aiohttp_client) -> None:
|
||||||
|
call_count = 0
|
||||||
|
|
||||||
|
async def sometimes_empty(content, session_key="", channel="", chat_id=""):
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
if call_count == 1:
|
||||||
|
return ""
|
||||||
|
return "recovered response"
|
||||||
|
|
||||||
|
agent = MagicMock()
|
||||||
|
agent.process_direct = sometimes_empty
|
||||||
|
agent._connect_mcp = AsyncMock()
|
||||||
|
agent.close_mcp = AsyncMock()
|
||||||
|
|
||||||
|
app = create_app(agent, model_name="m")
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={"messages": [{"role": "user", "content": "hello"}]},
|
||||||
|
)
|
||||||
|
assert resp.status == 200
|
||||||
|
body = await resp.json()
|
||||||
|
assert body["choices"][0]["message"]["content"] == "recovered response"
|
||||||
|
assert call_count == 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_empty_response_falls_back(aiohttp_client) -> None:
|
||||||
|
call_count = 0
|
||||||
|
|
||||||
|
async def always_empty(content, session_key="", channel="", chat_id=""):
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
return ""
|
||||||
|
|
||||||
|
agent = MagicMock()
|
||||||
|
agent.process_direct = always_empty
|
||||||
|
agent._connect_mcp = AsyncMock()
|
||||||
|
agent.close_mcp = AsyncMock()
|
||||||
|
|
||||||
|
app = create_app(agent, model_name="m")
|
||||||
|
client = await aiohttp_client(app)
|
||||||
|
resp = await client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={"messages": [{"role": "user", "content": "hello"}]},
|
||||||
|
)
|
||||||
|
assert resp.status == 200
|
||||||
|
body = await resp.json()
|
||||||
|
assert body["choices"][0]["message"]["content"] == "I've completed processing but have no response to give."
|
||||||
|
assert call_count == 2
|
||||||
@@ -408,6 +408,56 @@ async def test_exec_timeout_capped_at_max() -> None:
|
|||||||
assert "Exit code: 0" in result
|
assert "Exit code: 0" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_applied() -> None:
|
||||||
|
"""command_wrapper should wrap the original command."""
|
||||||
|
tool = ExecTool(command_wrapper="echo WRAPPED: {command}")
|
||||||
|
result = await tool.execute(command="hello")
|
||||||
|
assert "WRAPPED: hello" in result
|
||||||
|
assert "Exit code: 0" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_with_cwd(tmp_path) -> None:
|
||||||
|
"""command_wrapper should substitute {cwd} with the absolute working directory."""
|
||||||
|
tool = ExecTool(command_wrapper="echo CWD:{cwd} CMD:{command}")
|
||||||
|
result = await tool.execute(command="hi", working_dir=str(tmp_path))
|
||||||
|
assert str(tmp_path) in result
|
||||||
|
assert "CMD:hi" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_empty_noop() -> None:
|
||||||
|
"""Empty command_wrapper should leave the command unchanged."""
|
||||||
|
tool = ExecTool(command_wrapper="")
|
||||||
|
result = await tool.execute(command="echo direct")
|
||||||
|
assert "direct" in result
|
||||||
|
assert "Exit code: 0" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_guard_runs_before_wrapper() -> None:
|
||||||
|
"""Safety guard should run before wrapper substitution."""
|
||||||
|
tool = ExecTool(command_wrapper="echo WRAPPED:{command}")
|
||||||
|
result = await tool.execute(command="rm -rf /")
|
||||||
|
assert "blocked by safety guard" in result
|
||||||
|
assert "WRAPPED:" not in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_ignores_unknown_placeholders() -> None:
|
||||||
|
"""Unknown {placeholders} in the wrapper should be left as-is, not raise KeyError."""
|
||||||
|
tool = ExecTool(command_wrapper="echo {command} {unknown}")
|
||||||
|
result = await tool.execute(command="hello")
|
||||||
|
assert "Exit code: 0" in result
|
||||||
|
assert "{unknown}" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_does_not_leak_attributes() -> None:
|
||||||
|
"""Wrapper should not expose Python internals via attribute access."""
|
||||||
|
tool = ExecTool(command_wrapper="echo {command.__class__}")
|
||||||
|
result = await tool.execute(command="hello")
|
||||||
|
assert "Exit code: 0" in result
|
||||||
|
# {command.__class__} is not a valid placeholder; {command} gets replaced
|
||||||
|
# leaving {.__class__} as a literal string — no Python object is leaked.
|
||||||
|
assert "<class" not in result
|
||||||
|
|
||||||
|
|
||||||
# --- _resolve_type and nullable param tests ---
|
# --- _resolve_type and nullable param tests ---
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user