Merge pull request #8 from EvoScientist/feature/channel-unification
Feature/channel unificationfeat: unified channel architecture with multi-platform support
This commit is contained in:
+4
-1
@@ -15,6 +15,8 @@ build/
|
||||
.venv/
|
||||
venv/
|
||||
uv.lock
|
||||
bridge/node_modules/
|
||||
bridge/package-lock.json
|
||||
|
||||
# IDE / Tools
|
||||
.vscode/
|
||||
@@ -35,4 +37,5 @@ memory/
|
||||
*.ipynb
|
||||
*CLAUDE.md
|
||||
*AGENTS.md
|
||||
*meals/
|
||||
*meals/
|
||||
botpy.log
|
||||
@@ -27,7 +27,7 @@ from .mcp import load_mcp_tools
|
||||
from .middleware import create_memory_middleware
|
||||
from .prompts import RESEARCHER_INSTRUCTIONS, get_system_prompt
|
||||
from .utils import load_subagents
|
||||
from .tools import tavily_search, think_tool, skill_manager
|
||||
from .tools import tavily_search, think_tool, skill_manager, view_image
|
||||
from . import paths as _paths_mod
|
||||
from .paths import set_active_workspace, set_workspace_root
|
||||
|
||||
@@ -110,10 +110,11 @@ backend = CompositeBackend(
|
||||
tool_registry = {
|
||||
"think_tool": think_tool,
|
||||
"tavily_search": tavily_search,
|
||||
"view_image": view_image,
|
||||
}
|
||||
|
||||
# Base tools that every agent variant gets (before MCP)
|
||||
BASE_TOOLS = [think_tool, skill_manager]
|
||||
BASE_TOOLS = [think_tool, skill_manager, view_image]
|
||||
|
||||
# Cache MCP tools by the effective config signature to avoid reconnecting
|
||||
# to MCP servers on every `/new` when config is unchanged.
|
||||
|
||||
@@ -34,6 +34,9 @@ _EXPORTS: dict[str, tuple[str, str]] = {
|
||||
# Tools
|
||||
"tavily_search": (".tools", "tavily_search"),
|
||||
"think_tool": (".tools", "think_tool"),
|
||||
"view_image": (".tools", "view_image"),
|
||||
# Middleware
|
||||
"create_skills_middleware": (".middleware", "create_skills_middleware"),
|
||||
# Sessions
|
||||
"get_checkpointer": (".sessions", "get_checkpointer"),
|
||||
"generate_thread_id": (".sessions", "generate_thread_id"),
|
||||
|
||||
@@ -0,0 +1,552 @@
|
||||
# Channels
|
||||
|
||||
EvoScientist provides unified integration with 11 messaging platforms. This document covers the architecture overview, capability matrix, and detailed deployment guide for each channel.
|
||||
|
||||
Configuration file: `~/.config/evoscientist/config.yaml` (or use environment variables with the `EVOSCIENTIST_` prefix).
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
┌──────────┐ ┌──────────┐ ┌──────────┐
|
||||
│ Telegram │ │ Discord │ │ Slack │ ... (×11)
|
||||
└────┬─────┘ └────┬─────┘ └────┬─────┘
|
||||
│ │ │
|
||||
└─────────────┼─────────────┘
|
||||
▼
|
||||
┌──────────────┐
|
||||
│ MessageBus │ async queue, 5000 cap
|
||||
└──────┬───────┘
|
||||
▼
|
||||
┌──────────────┐
|
||||
│InboundConsumer│ → Agent → OutboundMessage
|
||||
└──────┬───────┘
|
||||
▼
|
||||
┌──────────────┐
|
||||
│ Dispatcher │ routes replies to origin channel
|
||||
└──────────────┘
|
||||
```
|
||||
|
||||
**Core modules:**
|
||||
|
||||
| Module | Responsibility |
|
||||
|--------|---------------|
|
||||
| `base.py` | Abstract `Channel` base class — declarative readiness checks, retry strategy, mention stripping, send fallback, media handling |
|
||||
| `capabilities.py` | `ChannelCapabilities` frozen dataclass — each channel declares its capabilities, framework adapts automatically |
|
||||
| `mixins.py` | Reusable patterns: `WebhookMixin` (aiohttp + httpx), `WebSocketMixin` (connect/reconnect/heartbeat), `PollingMixin` (async polling), `TokenMixin` (OAuth token refresh) |
|
||||
| `config.py` | `BaseChannelConfig` — shared config fields (allowed_senders, proxy, text_chunk_limit, etc.) |
|
||||
| `bus/` | `MessageBus` async message queue + `InboundMessage`/`OutboundMessage` dataclasses |
|
||||
| `channel_manager.py` | Lifecycle management (start/stop), health checks, channel registry |
|
||||
| `consumer.py` | `InboundConsumer` — dequeue messages, invoke Agent, publish replies |
|
||||
| `retry.py` | Configurable exponential backoff retry (`RetryConfig`: attempts, min/max delay, jitter) |
|
||||
| `markdown_utils.py` | Universal Markdown converter with per-platform formatting plugins |
|
||||
|
||||
## Capability Matrix
|
||||
|
||||
| Channel | Format | Max Len | Media | Voice | Sticker | Location | Video | Typing | Reaction | Thread | Group | @Mention | No Public IP | Token Refresh | Proxy | Allowlist |
|
||||
|:--------|:------:|:-------:|:-----:|:-----:|:-------:|:--------:|:-----:|:------:|:--------:|:------:|:-----:|:--------:|:------------:|:-------------:|:-----:|:---------:|
|
||||
| Telegram | HTML | 4000 | ✓ | ✓ | ✓ | ✓ | | ✓ | ✓ | | ✓ | ✓ | ✓ | | ✓ | ✓ |
|
||||
| Discord | Discord | 2000 | ✓ | | | | | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | | ✓ | ✓ |
|
||||
| Slack | Mrkdwn | 4000 | ✓ | | | | | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | | ✓ | ✓ |
|
||||
| Feishu | MD | 4096 | ✓ | ✓ | ✓ | | | | ✓ | | ✓ | ✓ | | ✓ | ✓ | ✓ |
|
||||
| WeChat | MD | 4096 | ✓ | ✓ | | ✓ | | | | | ✓ | ✓ | | ✓ | ✓ | ✓ |
|
||||
| DingTalk | MD | 4096 | ✓ | ✓ | | | | | | | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ |
|
||||
| QQ | Plain | 4096 | ✓ | | | | | | | | ✓ | ✓ | ✓ | | | ✓ |
|
||||
| Signal | Plain | 4096 | ✓ | ✓ | | | | ✓ | ✓ | | ✓ | ✓ | ✓ | | | ✓ |
|
||||
| iMessage | Plain | ∞ | ✓ | ✓ | | | | | | | ✓ | | ✓ | | | ✓ |
|
||||
| Email | HTML | ∞ | ✓ | | | | | | | | | | ✓ | | | ✓ |
|
||||
|
||||
### Connection Types
|
||||
|
||||
| Channel | Transport | Connection Mode | Default Port |
|
||||
|-------------|-----------|----------------------------------------|:------------:|
|
||||
| Telegram | HTTPS | Long polling (`getUpdates`) | — |
|
||||
| Discord | WebSocket | Gateway events (`discord.py`) | — |
|
||||
| Slack | WebSocket | Socket Mode (`slack-sdk`) | — |
|
||||
| Feishu | HTTP | Webhook `POST /webhook/event` | 9000 |
|
||||
| WeChat | HTTP | Webhook `POST /wechat/callback` | 9001 |
|
||||
| DingTalk | WebSocket | Stream Mode (DingTalk gateway) | — |
|
||||
| QQ | WebSocket | Bot Gateway (`qq-botpy`) | — |
|
||||
| Signal | TCP | JSON-RPC (`signal-cli` daemon) | 7583 |
|
||||
| iMessage | stdio | JSON-RPC (`imsg` CLI) | — |
|
||||
| Email | TCP | IMAP polling + SMTP send | 993/587 |
|
||||
|
||||
> **"—"** means no listening port is required — no public IP or port forwarding needed.
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Install channel dependencies
|
||||
|
||||
```bash
|
||||
pip install evoscientist[telegram]
|
||||
# Available extras: telegram, discord, slack, feishu, wechat,
|
||||
# dingtalk, qq, email, signal
|
||||
# iMessage requires no extra Python dependencies
|
||||
```
|
||||
|
||||
### 2. Configure
|
||||
|
||||
```bash
|
||||
# Option A: Interactive wizard
|
||||
EvoSci onboard
|
||||
|
||||
# Option B: CLI commands
|
||||
EvoSci config set channel_enabled telegram
|
||||
EvoSci config set telegram_bot_token "123456:ABC-xxx"
|
||||
|
||||
# Option C: Environment variables (EVOSCIENTIST_ prefix, uppercase)
|
||||
export EVOSCIENTIST_CHANNEL_ENABLED=telegram
|
||||
export EVOSCIENTIST_TELEGRAM_BOT_TOKEN="123456:ABC-xxx"
|
||||
```
|
||||
|
||||
### 3. Start
|
||||
|
||||
```bash
|
||||
EvoSci serve # Start agent + all enabled channels
|
||||
# or
|
||||
EvoSci channel start # Standalone channel mode (message loop only)
|
||||
```
|
||||
|
||||
### 4. Health check
|
||||
|
||||
```bash
|
||||
curl http://localhost:8080/healthz
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"status": "healthy",
|
||||
"channels": { "enabled": ["telegram"], "running": ["telegram"] }
|
||||
}
|
||||
```
|
||||
|
||||
### Running multiple channels
|
||||
|
||||
Comma-separate channel names in the config to enable multiple channels simultaneously:
|
||||
|
||||
```yaml
|
||||
channel_enabled: "telegram,discord,imessage"
|
||||
```
|
||||
|
||||
All enabled channels run concurrently via the internal message bus.
|
||||
|
||||
---
|
||||
|
||||
## Channel Deployment Guides
|
||||
|
||||
---
|
||||
|
||||
### Telegram
|
||||
|
||||
**Install:** `pip install evoscientist[telegram]`
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Search for [@BotFather](https://t.me/BotFather) in Telegram, send `/newbot`, and follow the prompts to create a bot.
|
||||
2. BotFather will return a Bot Token (format: `123456789:ABCdefGHI...`) — save it securely.
|
||||
3. Get your user ID: send any message to [@userinfobot](https://t.me/userinfobot), it will reply with your numeric ID.
|
||||
4. (Optional) For group use: add the bot to a group, then in BotFather send `/setprivacy` → `Disable` so the bot can read group messages.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "telegram"
|
||||
telegram_bot_token: "123456789:ABCdefGHIjklMNOpqrSTUvwxYZ"
|
||||
telegram_allowed_senders: "" # Comma-separated user IDs; empty = no restriction
|
||||
telegram_proxy: "" # Optional HTTPS proxy (e.g. http://proxy:8080)
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `telegram_bot_token` | `str` | `""` | **Required.** Bot API Token from BotFather |
|
||||
| `telegram_allowed_senders` | `str` | `""` | Comma-separated user IDs, empty = allow all |
|
||||
| `telegram_proxy` | `str` | `""` | HTTPS proxy URL |
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_TELEGRAM_BOT_TOKEN`, `EVOSCIENTIST_TELEGRAM_ALLOWED_SENDERS`, `EVOSCIENTIST_TELEGRAM_PROXY`
|
||||
|
||||
**Technical details:** Long polling mode, `drop_pending_updates=True` on startup to skip backlog. Markdown→Telegram HTML auto-conversion (bold, italic, strikethrough, links, code blocks, headings, lists). Falls back to plain text on HTML parse failure. Media routed by extension to `send_photo`/`send_video`/`send_audio`/`send_document`. In groups, only responds when @mentioned; auto-strips @mention. Typing indicator refreshes every 4s. Retry: 3 attempts, min delay 0.4s, parse errors not retried. Text chunk limit: 4000 chars.
|
||||
|
||||
---
|
||||
|
||||
### Discord
|
||||
|
||||
**Install:** `pip install evoscientist[discord]`
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Go to [Discord Developer Portal](https://discord.com/developers/applications) → New Application → enter a name.
|
||||
2. Left menu **Bot** → Reset Token → copy the Bot Token.
|
||||
3. Under **Privileged Gateway Intents**, enable **Message Content Intent** (required to read message content).
|
||||
4. Left menu **OAuth2** → URL Generator:
|
||||
- Scopes: check `bot`
|
||||
- Bot Permissions: check `Send Messages`, `Read Message History`, `Attach Files`, `Add Reactions`
|
||||
- Copy the generated URL, open in browser, select a server to invite the bot.
|
||||
5. Get user ID: Discord Settings → Advanced → enable Developer Mode → right-click username → Copy User ID.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "discord"
|
||||
discord_bot_token: "MTIzNDU2Nzg5.xxxx.xxxxx"
|
||||
discord_allowed_senders: "" # Comma-separated user IDs
|
||||
discord_allowed_channels: "" # Comma-separated channel IDs
|
||||
discord_proxy: ""
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `discord_bot_token` | `str` | `""` | **Required.** Bot Token |
|
||||
| `discord_allowed_senders` | `str` | `""` | Comma-separated user IDs, empty = allow all |
|
||||
| `discord_allowed_channels` | `str` | `""` | Comma-separated channel IDs, empty = allow all |
|
||||
| `discord_proxy` | `str` | `""` | HTTPS proxy URL |
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_DISCORD_BOT_TOKEN`, `EVOSCIENTIST_DISCORD_ALLOWED_SENDERS`, `EVOSCIENTIST_DISCORD_ALLOWED_CHANNELS`, `EVOSCIENTIST_DISCORD_PROXY`
|
||||
|
||||
**Technical details:** WebSocket Gateway (`discord.py`). In server channels, only responds when @mentioned; DMs respond directly. Replies via `MessageReference`. Attachment download (max 20 MB) with safe filename sanitization. Media sent via `discord.File`. Typing indicator refreshes every 8s. Retry: 3 attempts, parses `Retry-After` header for 429s. Text chunk limit: 2000 chars.
|
||||
|
||||
---
|
||||
|
||||
### Slack
|
||||
|
||||
**Install:** `pip install evoscientist[slack]`
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Go to [Slack API](https://api.slack.com/apps) → Create New App → From scratch → select workspace.
|
||||
2. Left menu **Socket Mode** → enable → Generate App-Level Token, scope `connections:write` → copy App Token (`xapp-...`).
|
||||
3. Left menu **OAuth & Permissions** → add Bot Token Scopes:
|
||||
- `chat:write`, `channels:history`, `groups:history`, `im:history`, `files:read`, `files:write`, `reactions:write`
|
||||
4. Click **Install to Workspace** → copy Bot User OAuth Token (`xoxb-...`).
|
||||
5. Left menu **Event Subscriptions** → enable → Subscribe to bot events: `message.channels`, `message.groups`, `message.im`, `app_mention`.
|
||||
6. Get Member ID: click user avatar → profile → **⋮** → Copy member ID.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "slack"
|
||||
slack_bot_token: "xoxb-xxxx-xxxx-xxxx"
|
||||
slack_app_token: "xapp-1-xxxx-xxxx"
|
||||
slack_allowed_senders: "" # Member ID (U...)
|
||||
slack_allowed_channels: "" # Channel ID (C...)
|
||||
slack_proxy: ""
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `slack_bot_token` | `str` | `""` | **Required.** Bot User OAuth Token (`xoxb-`) |
|
||||
| `slack_app_token` | `str` | `""` | **Required.** Socket Mode App Token (`xapp-`) |
|
||||
| `slack_allowed_senders` | `str` | `""` | Comma-separated Member IDs |
|
||||
| `slack_allowed_channels` | `str` | `""` | Comma-separated Channel IDs |
|
||||
| `slack_proxy` | `str` | `""` | HTTPS proxy URL |
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_SLACK_BOT_TOKEN`, `EVOSCIENTIST_SLACK_APP_TOKEN`, `EVOSCIENTIST_SLACK_ALLOWED_SENDERS`, `EVOSCIENTIST_SLACK_ALLOWED_CHANNELS`, `EVOSCIENTIST_SLACK_PROXY`
|
||||
|
||||
**Technical details:** Socket Mode (no public URL needed). Markdown→mrkdwn conversion. DMs respond directly; channels only respond to `app_mention` events. Thread replies via `thread_ts`. Attachments downloaded with Bearer auth. Media sent via `files_upload_v2`. Runs `auth_test()` on startup to verify credentials. Retry: 3 attempts, exponential backoff + jitter. Text chunk limit: 4000 chars.
|
||||
|
||||
---
|
||||
|
||||
### Feishu (Lark)
|
||||
|
||||
**Install:** `pip install evoscientist[feishu]`
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Go to [Feishu Open Platform](https://open.feishu.cn/app) (international: [Lark Developer](https://open.larksuite.com/app)) → create a custom app.
|
||||
2. Copy the **App ID** and **App Secret**.
|
||||
3. Left menu **Event Subscriptions** → set request URL to `http://your-host:9000/webhook/event` → copy **Verification Token** and **Encrypt Key**.
|
||||
4. Add event: `im.message.receive_v1` (receive messages).
|
||||
5. Left menu **Permissions** → enable `im:message:send_as_bot`.
|
||||
6. Create a version and publish.
|
||||
|
||||
> Webhook must be publicly reachable. For local dev, use `ngrok http 9000`.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "feishu"
|
||||
feishu_app_id: "cli_xxxxxxx"
|
||||
feishu_app_secret: "xxxxxxxxxxxxxxxxxx"
|
||||
feishu_webhook_port: 9000
|
||||
feishu_allowed_senders: "" # open_id
|
||||
feishu_domain: "https://open.feishu.cn"
|
||||
feishu_proxy: ""
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `feishu_app_id` | `str` | `""` | **Required.** App ID |
|
||||
| `feishu_app_secret` | `str` | `""` | **Required.** App Secret |
|
||||
| `feishu_webhook_port` | `int` | `9000` | Webhook HTTP port |
|
||||
| `feishu_allowed_senders` | `str` | `""` | Comma-separated open_ids |
|
||||
| `feishu_domain` | `str` | `"https://open.feishu.cn"` | API domain (use `https://open.larksuite.com` for Lark) |
|
||||
| `feishu_proxy` | `str` | `""` | HTTPS proxy URL |
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_FEISHU_APP_ID`, `EVOSCIENTIST_FEISHU_APP_SECRET`, `EVOSCIENTIST_FEISHU_WEBHOOK_PORT`, `EVOSCIENTIST_FEISHU_DOMAIN`
|
||||
|
||||
**Technical details:** Webhook on `POST /webhook/event` with URL verification challenge-response. `tenant_access_token` auto-refresh (2h TTL, refreshes 5 min before expiry). Markdown→Post rich text conversion (code blocks, bold, italic, strikethrough, links, headings, quotes, lists). Plain text fallback. Group @mention filtering. Media: images via `/im/v1/images`, files via `/im/v1/files`. Replies via `/messages/{id}/reply`. Retry: 3 attempts, rate limit delay 2.0s, matches `99991400`/`rate limit`. Text chunk limit: 4096 chars.
|
||||
|
||||
---
|
||||
|
||||
### WeChat
|
||||
|
||||
**Install:** `pip install evoscientist[wechat]`
|
||||
|
||||
Two backends supported: **WeCom** (recommended, free, no certification needed) and **WeChat Official Account** (requires verified service account).
|
||||
|
||||
#### WeCom
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Log in to [WeCom Admin Console](https://work.weixin.qq.com) → App Management → create a custom app.
|
||||
2. Copy the **Corp ID**, **AgentId**, and **Secret**.
|
||||
3. In app details → Receive Messages → Set API Receive → URL: `http://your-host:9001/wechat/callback` → copy **Token** and **EncodingAESKey**.
|
||||
|
||||
```yaml
|
||||
channel_enabled: "wechat"
|
||||
wechat_backend: "wecom"
|
||||
wechat_webhook_port: 9001
|
||||
wechat_wecom_corp_id: "ww..."
|
||||
wechat_wecom_agent_id: "1000002"
|
||||
wechat_wecom_secret: "xxxxxxxxxxxxxxxxxx"
|
||||
wechat_wecom_token: "xxxxxxxxxxxxxxxxxx"
|
||||
wechat_wecom_encoding_aes_key: "xxxxxxxxxxxxxxxxxx"
|
||||
wechat_allowed_senders: ""
|
||||
wechat_proxy: ""
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `wechat_backend` | `str` | `"wecom"` | `"wecom"` or `"wechatmp"` |
|
||||
| `wechat_webhook_port` | `int` | `9001` | Callback HTTP port |
|
||||
| `wechat_wecom_corp_id` | `str` | `""` | **Required (WeCom).** Corp ID |
|
||||
| `wechat_wecom_agent_id` | `str` | `""` | **Required (WeCom).** App AgentId |
|
||||
| `wechat_wecom_secret` | `str` | `""` | **Required (WeCom).** App Secret |
|
||||
| `wechat_wecom_token` | `str` | `""` | **Required (WeCom).** Callback Token |
|
||||
| `wechat_wecom_encoding_aes_key` | `str` | `""` | **Required (WeCom).** Callback EncodingAESKey |
|
||||
|
||||
#### WeChat Official Account
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Log in to [WeChat Official Account Platform](https://mp.weixin.qq.com) → Settings & Development → Basic Configuration.
|
||||
2. Copy the **AppID** and **AppSecret**.
|
||||
3. Server Configuration → URL: `http://your-host:9001/wechat/callback` → set **Token** and **EncodingAESKey**.
|
||||
|
||||
```yaml
|
||||
wechat_backend: "wechatmp"
|
||||
wechat_mp_app_id: "wx..."
|
||||
wechat_mp_app_secret: "xxxxxxxxxxxxxxxxxx"
|
||||
wechat_mp_token: "xxxxxxxxxxxxxxxxxx"
|
||||
wechat_mp_encoding_aes_key: "xxxxxxxxxxxxxxxxxx"
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `wechat_mp_app_id` | `str` | `""` | **Required (MP).** AppID |
|
||||
| `wechat_mp_app_secret` | `str` | `""` | **Required (MP).** AppSecret |
|
||||
| `wechat_mp_token` | `str` | `""` | **Required (MP).** Server Token |
|
||||
| `wechat_mp_encoding_aes_key` | `str` | `""` | **Required (MP).** Server EncodingAESKey |
|
||||
|
||||
**Technical details:** Webhook HTTP server. XML message parsing. Signature verification. `access_token` auto-refresh. Optional AES encryption/decryption. WeCom supports Markdown message format; Official Account uses plain text. Media send/receive. Retry + backoff. Text chunk limit: 2048 chars.
|
||||
|
||||
---
|
||||
|
||||
### DingTalk
|
||||
|
||||
**Install:** `pip install evoscientist[dingtalk]`
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Go to [DingTalk Open Platform](https://open-dev.dingtalk.com) → App Development → create a bot app.
|
||||
2. Copy the **AppKey** (Client ID) and **AppSecret** (Client Secret).
|
||||
3. Enable **Stream Mode** in the app configuration — no public IP needed.
|
||||
4. Publish the app and add the bot to a group, or test via direct message.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "dingtalk"
|
||||
dingtalk_client_id: "ding..."
|
||||
dingtalk_client_secret: "xxxxxxxxxxxxxxxxxx"
|
||||
dingtalk_allowed_senders: ""
|
||||
dingtalk_proxy: ""
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `dingtalk_client_id` | `str` | `""` | **Required.** AppKey |
|
||||
| `dingtalk_client_secret` | `str` | `""` | **Required.** AppSecret |
|
||||
| `dingtalk_allowed_senders` | `str` | `""` | Comma-separated user IDs |
|
||||
| `dingtalk_proxy` | `str` | `""` | HTTPS proxy URL |
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_DINGTALK_CLIENT_ID`, `EVOSCIENTIST_DINGTALK_CLIENT_SECRET`
|
||||
|
||||
**Technical details:** Stream Mode (WebSocket, no public IP needed). Connects via DingTalk gateway with automatic ping/pong heartbeat and message ACK. `access_token` auto-refresh. Group @mention filtering (strips first `@bot` mention). Supports image, file, video, audio attachment download. Sends in Markdown format (`sampleMarkdown`). Auth errors (`invalidauthentication`/`forbidden`/`40014`) not retried. Text chunk limit: 4096 chars.
|
||||
|
||||
---
|
||||
|
||||
### QQ
|
||||
|
||||
**Install:** `pip install evoscientist[qq]`
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Go to [QQ Open Platform](https://q.qq.com) → create a bot application.
|
||||
2. Complete developer verification, create a sandbox or production bot.
|
||||
3. Copy the **AppID** and **AppSecret**.
|
||||
4. Search for and add the bot as a friend in QQ, or add it to a group.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "qq"
|
||||
qq_app_id: "xxxxxxxxxx"
|
||||
qq_app_secret: "xxxxxxxxxxxxxxxxxx"
|
||||
qq_allowed_senders: ""
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `qq_app_id` | `str` | `""` | **Required.** AppID |
|
||||
| `qq_app_secret` | `str` | `""` | **Required.** AppSecret |
|
||||
| `qq_allowed_senders` | `str` | `""` | Comma-separated user IDs |
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_QQ_APP_ID`, `EVOSCIENTIST_QQ_APP_SECRET`
|
||||
|
||||
**Technical details:** Uses `qq-botpy` SDK via WebSocket to connect to QQ Bot Gateway. Supports C2C (direct) and group messages. Message deduplication (1000-entry LRU cache). Group @mention filtering (strips first `@bot`). Intents: `public_messages=True`, `direct_message=True`. Text chunk limit: 2048 chars.
|
||||
|
||||
---
|
||||
|
||||
### Signal
|
||||
|
||||
**Install:** `pip install evoscientist[signal]` (also requires [signal-cli](https://github.com/AsamK/signal-cli) installed separately)
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Install signal-cli: see [signal-cli installation guide](https://github.com/AsamK/signal-cli#installation).
|
||||
2. Register or link a phone number:
|
||||
- Register: `signal-cli -u +1234567890 register`, then `signal-cli -u +1234567890 verify CODE`
|
||||
- Link existing device: `signal-cli link -n "EvoScientist"`
|
||||
3. EvoScientist will auto-start the signal-cli daemon if it's not already running.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "signal"
|
||||
signal_phone_number: "+1234567890"
|
||||
signal_cli_path: "signal-cli"
|
||||
signal_config_dir: ""
|
||||
signal_allowed_senders: ""
|
||||
signal_rpc_port: 7583
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `signal_phone_number` | `str` | `""` | **Required.** Signal phone number (E.164 format) |
|
||||
| `signal_cli_path` | `str` | `"signal-cli"` | Path to signal-cli binary |
|
||||
| `signal_config_dir` | `str` | `""` | signal-cli config directory (optional) |
|
||||
| `signal_allowed_senders` | `str` | `""` | Comma-separated phone numbers |
|
||||
| `signal_rpc_port` | `int` | `7583` | JSON RPC socket port |
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_SIGNAL_PHONE_NUMBER`, `EVOSCIENTIST_SIGNAL_CLI_PATH`, `EVOSCIENTIST_SIGNAL_RPC_PORT`
|
||||
|
||||
**Technical details:** JSON RPC over TCP socket to signal-cli daemon. Auto-starts daemon if not running (`signal-cli -u +NUMBER daemon --socket localhost:PORT`). Listens for `receive` notifications. Sends via `send` RPC method. Group detection via `groupInfo`. Mention detection via UUID matching. No public IP needed. Text chunk limit: 4096 chars.
|
||||
|
||||
---
|
||||
|
||||
### Email
|
||||
|
||||
**Install:** `pip install evoscientist[email]` (core dependencies included, no extras needed)
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Prepare an email account with IMAP + SMTP support (Gmail, Outlook, self-hosted, etc.).
|
||||
2. **Gmail:** Enable 2FA → generate an App Password. IMAP: `imap.gmail.com:993` (SSL), SMTP: `smtp.gmail.com:587` (STARTTLS).
|
||||
3. **Outlook/Office 365:** IMAP: `outlook.office365.com:993` (SSL), SMTP: `smtp.office365.com:587` (STARTTLS).
|
||||
4. Ensure IMAP access is enabled in your email settings.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "email"
|
||||
email_imap_host: "imap.gmail.com"
|
||||
email_imap_port: 993
|
||||
email_imap_username: "bot@gmail.com"
|
||||
email_imap_password: "xxxx-xxxx-xxxx-xxxx"
|
||||
email_imap_mailbox: "INBOX"
|
||||
email_imap_use_ssl: true
|
||||
email_smtp_host: "smtp.gmail.com"
|
||||
email_smtp_port: 587
|
||||
email_smtp_username: "bot@gmail.com"
|
||||
email_smtp_password: "xxxx-xxxx-xxxx-xxxx"
|
||||
email_smtp_use_tls: true
|
||||
email_from_address: "bot@gmail.com"
|
||||
email_poll_interval: 30
|
||||
email_mark_seen: true
|
||||
email_allowed_senders: ""
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `email_imap_host` | `str` | `""` | **Required.** IMAP server address |
|
||||
| `email_imap_port` | `int` | `993` | IMAP port |
|
||||
| `email_imap_username` | `str` | `""` | **Required.** IMAP login username |
|
||||
| `email_imap_password` | `str` | `""` | **Required.** IMAP login password (or app password) |
|
||||
| `email_imap_mailbox` | `str` | `"INBOX"` | Mailbox folder to monitor |
|
||||
| `email_imap_use_ssl` | `bool` | `true` | Use SSL for IMAP connection |
|
||||
| `email_smtp_host` | `str` | `""` | **Required.** SMTP server address |
|
||||
| `email_smtp_port` | `int` | `587` | SMTP port |
|
||||
| `email_smtp_username` | `str` | `""` | **Required.** SMTP login username |
|
||||
| `email_smtp_password` | `str` | `""` | **Required.** SMTP login password |
|
||||
| `email_smtp_use_tls` | `bool` | `true` | Use STARTTLS (`true`) or SSL (`false`) |
|
||||
| `email_from_address` | `str` | `""` | Sender address (defaults to smtp_username) |
|
||||
| `email_poll_interval` | `int` | `30` | IMAP poll interval in seconds |
|
||||
| `email_mark_seen` | `bool` | `true` | Mark emails as read after processing |
|
||||
| `email_max_body_chars` | `int` | `12000` | Max email body chars (truncated beyond) |
|
||||
| `email_subject_prefix` | `str` | `"Re: "` | Reply subject prefix |
|
||||
| `email_allowed_senders` | `str` | `""` | Comma-separated sender email addresses |
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_EMAIL_IMAP_HOST`, `EVOSCIENTIST_EMAIL_IMAP_USERNAME`, `EVOSCIENTIST_EMAIL_IMAP_PASSWORD`, `EVOSCIENTIST_EMAIL_SMTP_HOST`, `EVOSCIENTIST_EMAIL_SMTP_USERNAME`, `EVOSCIENTIST_EMAIL_SMTP_PASSWORD`
|
||||
|
||||
**Technical details:** IMAP polling mode, checks for UNSEEN emails periodically (max 20 per cycle). Supports SSL and STARTTLS. Auto-parses multipart emails (prefers text/plain, falls back text/html → plain text). Attachments auto-downloaded. Replies set `In-Reply-To` and `References` headers to maintain email threads. Sends HTML + plain text dual format (multipart/alternative), falls back to plain text on HTML failure. IMAP auto-reconnects on disconnect. Auth errors (auth/login/credential) not retried. No public IP needed. Text chunk limit: no limit.
|
||||
|
||||
---
|
||||
|
||||
### iMessage
|
||||
|
||||
**Install:** No extra Python dependencies. Requires the [imsg](https://github.com/anthropics/imsg) CLI tool.
|
||||
|
||||
**Requirements:** macOS only (iMessage is Apple-proprietary). Requires a signed-in Apple ID with iMessage and Full Disk Access permission for the terminal app.
|
||||
|
||||
**Prerequisites:**
|
||||
|
||||
1. Install imsg CLI:
|
||||
```bash
|
||||
brew install imsg
|
||||
```
|
||||
2. Verify: `imsg --version`
|
||||
3. Ensure Messages.app is signed in and working on macOS.
|
||||
|
||||
**Configuration:**
|
||||
|
||||
```yaml
|
||||
channel_enabled: "imessage"
|
||||
imessage_cli_path: "imsg"
|
||||
imessage_db_path: ""
|
||||
imessage_service: "auto"
|
||||
imessage_region: "US"
|
||||
imessage_allowed_senders: ""
|
||||
```
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `imessage_cli_path` | `str` | `"imsg"` | Path to imsg CLI binary |
|
||||
| `imessage_db_path` | `str` | `""` | iMessage database path (empty = default) |
|
||||
| `imessage_service` | `str` | `"auto"` | Send service: `imessage`, `sms`, or `auto` |
|
||||
| `imessage_region` | `str` | `"US"` | Phone number region code |
|
||||
| `imessage_allowed_senders` | `str` | `""` | Comma-separated allowlist (see below) |
|
||||
|
||||
**Allowlist formats:** phone (`+1234567890`), email (`user@example.com`), `chat_id:123`, `chat_guid:iMessage;-;+1234567890`, wildcard `*`.
|
||||
|
||||
**Env vars:** `EVOSCIENTIST_IMESSAGE_CLI_PATH`, `EVOSCIENTIST_IMESSAGE_SERVICE`, `EVOSCIENTIST_IMESSAGE_ALLOWED_SENDERS`
|
||||
|
||||
**Technical details:** JSON-RPC over stdio with imsg CLI. Creates `watch.subscribe` on startup for real-time message streaming (not polling). Supports iMessage + SMS dual channel (`service: auto`). Target resolution supports chat_id, chat_guid, chat_identifier, and phone/email. Attachments read from local paths provided by imsg. Group detection via `is_group` field. RPC errors (AppleScript/permission/not found) not retried; only connection timeouts retried. Plain text format (no Markdown). No public IP needed. Text chunk limit: 4000 chars.
|
||||
@@ -1,9 +1,44 @@
|
||||
"""Communication channels for EvoScientist.
|
||||
|
||||
This module provides an extensible interface for different messaging channels
|
||||
(iMessage, WeChat, etc.) to communicate with the EvoScientist agent.
|
||||
(iMessage, Telegram, Discord) to communicate with the EvoScientist agent.
|
||||
"""
|
||||
|
||||
from .base import Channel, IncomingMessage, OutgoingMessage
|
||||
from .base import Channel, RawIncoming, IncomingMessage, OutgoingMessage, chunk_text
|
||||
from .bus import MessageBus, InboundMessage, OutboundMessage
|
||||
from .capabilities import ChannelCapabilities
|
||||
from .channel_manager import ChannelManager, register_channel, create_channel, available_channels
|
||||
from .consumer import InboundConsumer
|
||||
from .formatter import UnifiedFormatter
|
||||
from .middleware import TypingManager
|
||||
from .plugin import ChannelPlugin, ChannelMeta, ReloadPolicy
|
||||
from .standalone import run_standalone
|
||||
|
||||
__all__ = ["Channel", "IncomingMessage", "OutgoingMessage"]
|
||||
# Backward compat: ChannelServer is now Channel itself
|
||||
ChannelServer = Channel
|
||||
|
||||
__all__ = [
|
||||
"Channel",
|
||||
"ChannelServer",
|
||||
"ChannelManager",
|
||||
"MessageBus",
|
||||
"RawIncoming",
|
||||
"IncomingMessage",
|
||||
"OutgoingMessage",
|
||||
"InboundMessage",
|
||||
"OutboundMessage",
|
||||
"InboundConsumer",
|
||||
"run_standalone",
|
||||
"register_channel",
|
||||
"create_channel",
|
||||
"available_channels",
|
||||
# New modules
|
||||
"ChannelCapabilities",
|
||||
"UnifiedFormatter",
|
||||
"TypingManager",
|
||||
"chunk_text",
|
||||
# Plugin architecture
|
||||
"ChannelPlugin",
|
||||
"ChannelMeta",
|
||||
"ReloadPolicy",
|
||||
]
|
||||
|
||||
+963
-54
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,6 @@
|
||||
"""Message bus for decoupled channel-agent communication."""
|
||||
|
||||
from .events import InboundMessage, OutboundMessage
|
||||
from .message_bus import MessageBus
|
||||
|
||||
__all__ = ["MessageBus", "InboundMessage", "OutboundMessage"]
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Event types for the message bus."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class InboundMessage:
|
||||
"""Message received from a chat channel.
|
||||
|
||||
Carries enough context for the bus to route and for the agent
|
||||
to build a session: which channel, who sent it, which chat.
|
||||
"""
|
||||
|
||||
channel: str
|
||||
sender_id: str
|
||||
chat_id: str
|
||||
content: str
|
||||
timestamp: datetime = field(default_factory=datetime.now)
|
||||
message_id: str = ""
|
||||
media: list[str] = field(default_factory=list)
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
is_group: bool = False
|
||||
was_mentioned: bool = True
|
||||
|
||||
@property
|
||||
def sender(self) -> str:
|
||||
"""Alias for ``sender_id`` (compatibility with IncomingMessage)."""
|
||||
return self.sender_id
|
||||
|
||||
@property
|
||||
def session_key(self) -> str:
|
||||
"""Unique key for session identification: ``channel:chat_id``."""
|
||||
return f"{self.channel}:{self.chat_id}"
|
||||
|
||||
|
||||
@dataclass
|
||||
class OutboundMessage:
|
||||
"""Message to send to a chat channel."""
|
||||
|
||||
channel: str
|
||||
chat_id: str
|
||||
content: str
|
||||
reply_to: str | None = None
|
||||
media: list[str] = field(default_factory=list)
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def recipient(self) -> str:
|
||||
"""Alias for ``chat_id`` (compatibility with OutgoingMessage)."""
|
||||
return self.chat_id
|
||||
@@ -0,0 +1,96 @@
|
||||
"""Async message bus that decouples chat channels from the agent core.
|
||||
|
||||
Channels push messages to the inbound queue; the agent (or any consumer)
|
||||
reads from inbound, processes, and pushes responses to the outbound queue.
|
||||
A background dispatcher routes outbound messages to the correct channel
|
||||
via subscriber callbacks.
|
||||
|
||||
Deduplication is handled at the Channel level (single dedup point).
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Callable, Awaitable
|
||||
|
||||
from .events import InboundMessage, OutboundMessage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
OutboundCallback = Callable[[OutboundMessage], Awaitable[None]]
|
||||
|
||||
|
||||
class MessageBus:
|
||||
"""Async message bus that decouples chat channels from the agent core."""
|
||||
|
||||
def __init__(self):
|
||||
self.inbound: asyncio.Queue[InboundMessage] = asyncio.Queue(maxsize=5000)
|
||||
self.outbound: asyncio.Queue[OutboundMessage] = asyncio.Queue(maxsize=5000)
|
||||
self._outbound_subscribers: dict[str, list[OutboundCallback]] = {}
|
||||
self._running = False
|
||||
|
||||
# ── inbound (channel → agent) ──
|
||||
|
||||
async def publish_inbound(self, msg: InboundMessage) -> None:
|
||||
"""Publish a message from a channel to the agent."""
|
||||
await self.inbound.put(msg)
|
||||
|
||||
async def consume_inbound(self) -> InboundMessage:
|
||||
"""Consume the next inbound message (blocks until available)."""
|
||||
return await self.inbound.get()
|
||||
|
||||
# ── outbound (agent → channel) ──
|
||||
|
||||
async def publish_outbound(self, msg: OutboundMessage) -> None:
|
||||
"""Publish a response from the agent to channels."""
|
||||
await self.outbound.put(msg)
|
||||
|
||||
async def consume_outbound(self) -> OutboundMessage:
|
||||
"""Consume the next outbound message (blocks until available)."""
|
||||
return await self.outbound.get()
|
||||
|
||||
# ── subscriber routing ──
|
||||
|
||||
def subscribe_outbound(
|
||||
self, channel: str, callback: OutboundCallback,
|
||||
) -> None:
|
||||
"""Register a callback for outbound messages targeting *channel*."""
|
||||
if channel not in self._outbound_subscribers:
|
||||
self._outbound_subscribers[channel] = []
|
||||
self._outbound_subscribers[channel].append(callback)
|
||||
|
||||
async def dispatch_outbound(self) -> None:
|
||||
"""Route outbound messages to subscribed channels.
|
||||
|
||||
Run as a background task — loops until :meth:`stop` is called.
|
||||
"""
|
||||
self._running = True
|
||||
while self._running:
|
||||
try:
|
||||
msg = await asyncio.wait_for(
|
||||
self.outbound.get(), timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
subscribers = self._outbound_subscribers.get(msg.channel, [])
|
||||
if not subscribers:
|
||||
logger.warning(f"No subscriber for channel: {msg.channel}")
|
||||
continue
|
||||
for callback in subscribers:
|
||||
try:
|
||||
await callback(msg)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error dispatching to {msg.channel}: {e}"
|
||||
)
|
||||
|
||||
def stop(self) -> None:
|
||||
"""Stop the dispatcher loop."""
|
||||
self._running = False
|
||||
|
||||
@property
|
||||
def inbound_size(self) -> int:
|
||||
return self.inbound.qsize()
|
||||
|
||||
@property
|
||||
def outbound_size(self) -> int:
|
||||
return self.outbound.qsize()
|
||||
@@ -0,0 +1,220 @@
|
||||
"""Channel capabilities declaration system.
|
||||
|
||||
Each channel declares its capabilities via a ChannelCapabilities dataclass,
|
||||
enabling the framework to adapt behavior automatically (formatting, reactions,
|
||||
streaming, threading, etc.) without per-channel branching in core logic.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
|
||||
FormatType = Literal["html", "markdown", "slack_mrkdwn", "discord", "plain"]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ChannelCapabilities:
|
||||
"""Immutable declaration of what a channel supports.
|
||||
|
||||
Set once as a class attribute on each Channel subclass.
|
||||
The framework inspects these at runtime to auto-configure behavior.
|
||||
"""
|
||||
|
||||
# ── Messaging features ──────────────────────────────────────────
|
||||
format_type: FormatType = "plain"
|
||||
max_text_length: int = 4096
|
||||
max_file_size: int = 20 * 1024 * 1024 # 20 MB
|
||||
|
||||
# ── Interaction capabilities ────────────────────────────────────
|
||||
streaming: bool = False # edit-in-place streaming output
|
||||
threading: bool = False # message threads / topics
|
||||
reactions: bool = False # emoji reactions on messages
|
||||
typing: bool = False # typing indicator API
|
||||
inline_buttons: bool = False # inline keyboard / action buttons
|
||||
|
||||
# ── Media capabilities ──────────────────────────────────────────
|
||||
media_send: bool = False # can send files/images
|
||||
media_receive: bool = False # can receive files/images
|
||||
voice: bool = False # platform has voice/audio messages that arrive as downloadable files (receive only, not bot sending)
|
||||
stickers: bool = False # supports sticker receive (not bot sending)
|
||||
location: bool = False # supports location message receive (not bot sending)
|
||||
video: bool = False # video messages
|
||||
|
||||
# ── Group features ──────────────────────────────────────────────
|
||||
groups: bool = False # group chat support
|
||||
mentions: bool = False # @mention detection
|
||||
|
||||
# ── Rich text ───────────────────────────────────────────────────
|
||||
markdown: bool = False # supports Markdown rendering
|
||||
html: bool = False # supports HTML rendering
|
||||
|
||||
# ── Extended capabilities ────────────────────────────────────────
|
||||
chat_types: tuple[str, ...] = () # ("direct", "group", "channel", "thread")
|
||||
edit: bool = False # message editing after send
|
||||
unsend: bool = False # message recall / unsend
|
||||
block_streaming: bool = False # block edit-in-place streaming
|
||||
native_commands: bool = False # platform-native slash commands
|
||||
polls: bool = False # poll / vote messages
|
||||
|
||||
def supports(self, feature: str) -> bool:
|
||||
"""Check if a feature is supported by name."""
|
||||
return getattr(self, feature, False)
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Pre-built capability profiles for each channel
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
TELEGRAM = ChannelCapabilities(
|
||||
format_type="html",
|
||||
max_text_length=4000,
|
||||
streaming=False, # could edit messages, but not implemented yet
|
||||
threading=False, # topics exist but not used yet
|
||||
reactions=True,
|
||||
typing=True,
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=True,
|
||||
stickers=True,
|
||||
location=True,
|
||||
groups=True,
|
||||
mentions=True,
|
||||
html=True,
|
||||
chat_types=("direct", "group", "channel"),
|
||||
edit=True,
|
||||
unsend=True,
|
||||
native_commands=True,
|
||||
polls=True,
|
||||
)
|
||||
|
||||
DISCORD = ChannelCapabilities(
|
||||
format_type="discord",
|
||||
max_text_length=2000,
|
||||
streaming=False,
|
||||
threading=True,
|
||||
reactions=True,
|
||||
typing=True,
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=False, # no distinct voice message type in Discord bot API
|
||||
groups=True,
|
||||
mentions=True,
|
||||
markdown=True,
|
||||
chat_types=("direct", "group", "thread"),
|
||||
edit=True,
|
||||
unsend=True,
|
||||
native_commands=True,
|
||||
polls=True,
|
||||
)
|
||||
|
||||
SLACK = ChannelCapabilities(
|
||||
format_type="slack_mrkdwn",
|
||||
max_text_length=4000,
|
||||
streaming=False,
|
||||
threading=True,
|
||||
reactions=True,
|
||||
typing=True,
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=False, # no distinct voice message type in Slack bot API
|
||||
groups=True,
|
||||
mentions=True,
|
||||
chat_types=("direct", "group", "thread"),
|
||||
edit=True,
|
||||
unsend=True,
|
||||
native_commands=True,
|
||||
)
|
||||
|
||||
FEISHU = ChannelCapabilities(
|
||||
format_type="markdown",
|
||||
max_text_length=4096,
|
||||
reactions=True,
|
||||
typing=False, # no typing API
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=True,
|
||||
stickers=True,
|
||||
groups=True,
|
||||
mentions=True,
|
||||
markdown=True,
|
||||
chat_types=("direct", "group"),
|
||||
edit=True,
|
||||
unsend=True,
|
||||
)
|
||||
|
||||
DINGTALK = ChannelCapabilities(
|
||||
format_type="markdown",
|
||||
max_text_length=4096,
|
||||
typing=False, # no typing API for bots
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=True,
|
||||
groups=True,
|
||||
mentions=True,
|
||||
markdown=True,
|
||||
chat_types=("direct", "group"),
|
||||
)
|
||||
|
||||
QQ = ChannelCapabilities(
|
||||
format_type="plain",
|
||||
max_text_length=4096,
|
||||
typing=False, # no typing API for QQ bots
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=False, # qq-botpy does not expose voice as a distinct message type
|
||||
groups=True,
|
||||
mentions=True,
|
||||
chat_types=("direct", "group", "channel"),
|
||||
unsend=True,
|
||||
)
|
||||
|
||||
WECHAT = ChannelCapabilities(
|
||||
format_type="markdown", # WeCom supports markdown
|
||||
max_text_length=4096,
|
||||
typing=False, # no typing API
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=True,
|
||||
location=True,
|
||||
groups=True,
|
||||
mentions=True,
|
||||
markdown=True,
|
||||
chat_types=("direct", "group"),
|
||||
unsend=True,
|
||||
)
|
||||
|
||||
SIGNAL = ChannelCapabilities(
|
||||
format_type="plain",
|
||||
max_text_length=4096,
|
||||
reactions=True,
|
||||
typing=True,
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=True,
|
||||
groups=True,
|
||||
mentions=True,
|
||||
chat_types=("direct", "group"),
|
||||
)
|
||||
|
||||
EMAIL = ChannelCapabilities(
|
||||
format_type="html",
|
||||
max_text_length=999_999, # no practical limit
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
html=True,
|
||||
chat_types=("direct",),
|
||||
)
|
||||
|
||||
IMESSAGE = ChannelCapabilities(
|
||||
format_type="plain",
|
||||
max_text_length=999_999,
|
||||
typing=False, # Apple does not expose typing indicator API
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=True,
|
||||
groups=True,
|
||||
mentions=False, # iMessage has no @mention concept
|
||||
reactions=False, # imsg CLI cannot send tapback reactions
|
||||
chat_types=("direct", "group"),
|
||||
)
|
||||
@@ -0,0 +1,995 @@
|
||||
"""Unified channel manager for coordinating chat channels.
|
||||
|
||||
Manages channel lifecycle (start/stop), wires each channel to the
|
||||
message bus, and routes outbound messages to the correct channel.
|
||||
|
||||
Also provides the global channel registry (formerly in ``registry.py``),
|
||||
account management (formerly ``account.py``), and pipeline assembly
|
||||
(formerly ``pipeline.py``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import importlib
|
||||
import json
|
||||
import logging
|
||||
import pkgutil
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
|
||||
from .base import Channel, OutboundMessage
|
||||
from .bus import MessageBus
|
||||
from .middleware import OutboundMiddlewareBase
|
||||
from .plugin import ChannelPlugin
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Account management (formerly account.py)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
@dataclass
|
||||
class ChannelAccountSnapshot:
|
||||
"""Point-in-time snapshot of a single account's connection state."""
|
||||
|
||||
account_id: str
|
||||
channel: str
|
||||
connected: bool = False
|
||||
started_at: float = 0.0
|
||||
last_outbound_at: float = 0.0
|
||||
error: str | None = None
|
||||
|
||||
def mark_connected(self) -> None:
|
||||
self.connected = True
|
||||
self.started_at = time.monotonic()
|
||||
self.error = None
|
||||
|
||||
def mark_disconnected(self, error: str | None = None) -> None:
|
||||
self.connected = False
|
||||
self.error = error
|
||||
|
||||
def mark_outbound(self) -> None:
|
||||
self.last_outbound_at = time.monotonic()
|
||||
|
||||
|
||||
@dataclass
|
||||
class AccountConfig:
|
||||
"""Per-account configuration wrapper."""
|
||||
|
||||
account_id: str
|
||||
channel_id: str # which plugin
|
||||
enabled: bool = True
|
||||
config: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AccountState:
|
||||
"""Runtime state for a single account."""
|
||||
|
||||
account_id: str
|
||||
channel_id: str
|
||||
status: str = "stopped" # stopped | starting | running | error
|
||||
snapshot: ChannelAccountSnapshot | None = None
|
||||
error: str | None = None
|
||||
started_at: float = 0.0
|
||||
|
||||
|
||||
class AccountManager:
|
||||
"""Manages multiple accounts across channel plugins.
|
||||
|
||||
Works with the ``ConfigAdapter`` protocol on each plugin to discover
|
||||
accounts and manage their lifecycle independently.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._plugins: dict[str, ChannelPlugin] = {}
|
||||
self._states: dict[str, AccountState] = {} # key: "{channel_id}:{account_id}"
|
||||
|
||||
@staticmethod
|
||||
def _key(channel_id: str, account_id: str) -> str:
|
||||
return f"{channel_id}:{account_id}"
|
||||
|
||||
def register_plugin(self, plugin: ChannelPlugin) -> None:
|
||||
"""Register a plugin that supports multi-account."""
|
||||
self._plugins[plugin.id] = plugin
|
||||
logger.info(f"AccountManager: registered plugin '{plugin.id}'")
|
||||
|
||||
async def start_account(
|
||||
self,
|
||||
channel_id: str,
|
||||
account_id: str,
|
||||
config: Any = None,
|
||||
) -> None:
|
||||
"""Start a specific account on a plugin."""
|
||||
plugin = self._plugins.get(channel_id)
|
||||
if plugin is None:
|
||||
raise ValueError(f"No plugin registered for channel '{channel_id}'")
|
||||
|
||||
key = self._key(channel_id, account_id)
|
||||
state = self._states.get(key)
|
||||
if state is None:
|
||||
state = AccountState(account_id=account_id, channel_id=channel_id)
|
||||
self._states[key] = state
|
||||
|
||||
if state.status == "running":
|
||||
logger.warning(f"Account {key} is already running")
|
||||
return
|
||||
|
||||
state.status = "starting"
|
||||
state.error = None
|
||||
try:
|
||||
account_config = config
|
||||
if plugin.config_adapter is not None and config is not None:
|
||||
account_config = plugin.config_adapter.resolve_account(config, account_id)
|
||||
|
||||
await plugin.start(account_config, account_id=account_id)
|
||||
state.status = "running"
|
||||
state.started_at = time.monotonic()
|
||||
state.snapshot = ChannelAccountSnapshot(
|
||||
account_id=account_id, channel=channel_id,
|
||||
)
|
||||
state.snapshot.mark_connected()
|
||||
logger.info(f"Account {key} started")
|
||||
except Exception as e:
|
||||
state.status = "error"
|
||||
state.error = str(e)
|
||||
logger.error(f"Failed to start account {key}: {e}")
|
||||
raise
|
||||
|
||||
async def stop_account(self, channel_id: str, account_id: str) -> None:
|
||||
"""Stop a specific account on a plugin."""
|
||||
plugin = self._plugins.get(channel_id)
|
||||
if plugin is None:
|
||||
raise ValueError(f"No plugin registered for channel '{channel_id}'")
|
||||
|
||||
key = self._key(channel_id, account_id)
|
||||
state = self._states.get(key)
|
||||
if state is None or state.status == "stopped":
|
||||
logger.debug(f"Account {key} is already stopped")
|
||||
return
|
||||
|
||||
try:
|
||||
await plugin.stop(account_id=account_id)
|
||||
state.status = "stopped"
|
||||
if state.snapshot is not None:
|
||||
state.snapshot.mark_disconnected()
|
||||
logger.info(f"Account {key} stopped")
|
||||
except Exception as e:
|
||||
state.status = "error"
|
||||
state.error = str(e)
|
||||
if state.snapshot is not None:
|
||||
state.snapshot.mark_disconnected(error=str(e))
|
||||
logger.error(f"Error stopping account {key}: {e}")
|
||||
raise
|
||||
|
||||
async def restart_account(
|
||||
self,
|
||||
channel_id: str,
|
||||
account_id: str,
|
||||
config: Any = None,
|
||||
) -> None:
|
||||
"""Restart a specific account."""
|
||||
await self.stop_account(channel_id, account_id)
|
||||
await self.start_account(channel_id, account_id, config)
|
||||
|
||||
async def start_all(self, channel_id: str, config: Any = None) -> None:
|
||||
"""Start all accounts for a given channel plugin."""
|
||||
plugin = self._plugins.get(channel_id)
|
||||
if plugin is None:
|
||||
raise ValueError(f"No plugin registered for channel '{channel_id}'")
|
||||
|
||||
adapter = plugin.config_adapter
|
||||
if adapter is None:
|
||||
await self.start_account(channel_id, "default", config)
|
||||
return
|
||||
|
||||
if config is None:
|
||||
logger.warning(f"No config provided for start_all on '{channel_id}'")
|
||||
return
|
||||
|
||||
for account_id in adapter.list_account_ids(config):
|
||||
if adapter.is_enabled(
|
||||
adapter.resolve_account(config, account_id), config,
|
||||
):
|
||||
try:
|
||||
await self.start_account(channel_id, account_id, config)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Failed to start account {channel_id}:{account_id}: {e}"
|
||||
)
|
||||
|
||||
async def stop_all(self, channel_id: str) -> None:
|
||||
"""Stop all accounts for a given channel plugin."""
|
||||
keys_to_stop = [
|
||||
(state.channel_id, state.account_id)
|
||||
for state in self._states.values()
|
||||
if state.channel_id == channel_id and state.status != "stopped"
|
||||
]
|
||||
for cid, aid in keys_to_stop:
|
||||
try:
|
||||
await self.stop_account(cid, aid)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to stop account {cid}:{aid}: {e}")
|
||||
|
||||
def get_state(
|
||||
self, channel_id: str, account_id: str,
|
||||
) -> AccountState | None:
|
||||
"""Get the runtime state for a specific account."""
|
||||
return self._states.get(self._key(channel_id, account_id))
|
||||
|
||||
def list_accounts(
|
||||
self, channel_id: str | None = None,
|
||||
) -> list[AccountState]:
|
||||
"""List account states, optionally filtered by channel."""
|
||||
if channel_id is None:
|
||||
return list(self._states.values())
|
||||
return [
|
||||
s for s in self._states.values() if s.channel_id == channel_id
|
||||
]
|
||||
|
||||
def get_snapshot(
|
||||
self, channel_id: str, account_id: str,
|
||||
) -> ChannelAccountSnapshot | None:
|
||||
"""Get the connection snapshot for a specific account."""
|
||||
state = self._states.get(self._key(channel_id, account_id))
|
||||
return state.snapshot if state else None
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Inbound / outbound pipelines (formerly pipeline.py)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class OutboundPipeline:
|
||||
"""Processes outgoing messages through a middleware chain."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
plugin: ChannelPlugin,
|
||||
middlewares: list[OutboundMiddlewareBase],
|
||||
) -> None:
|
||||
self.plugin = plugin
|
||||
self.middlewares = middlewares
|
||||
|
||||
async def process(
|
||||
self,
|
||||
message: OutboundMessage,
|
||||
context: dict[str, Any] | None = None,
|
||||
) -> OutboundMessage | None:
|
||||
"""Run *message* through each middleware. Returns ``None`` if dropped."""
|
||||
ctx = context or {}
|
||||
current: OutboundMessage | None = message
|
||||
for mw in self.middlewares:
|
||||
if current is None:
|
||||
return None
|
||||
current = await mw.process_outbound(current, ctx)
|
||||
return current
|
||||
|
||||
|
||||
def build_outbound_pipeline(
|
||||
plugin: ChannelPlugin,
|
||||
config: Any,
|
||||
) -> OutboundPipeline:
|
||||
"""Auto-assemble outbound pipeline based on plugin capabilities.
|
||||
|
||||
FormattingMiddleware has been removed — Channel.send() handles
|
||||
formatting + chunking via _format_chunk() / _prepare_chunks().
|
||||
"""
|
||||
middlewares: list[OutboundMiddlewareBase] = []
|
||||
return OutboundPipeline(plugin, middlewares)
|
||||
|
||||
|
||||
# ── Per-channel health tracking ──────────────────────────────────────
|
||||
|
||||
@dataclass
|
||||
class ChannelHealth:
|
||||
"""Tracks send success / failure metrics for a single channel."""
|
||||
|
||||
consecutive_failures: int = 0
|
||||
last_failure_time: float | None = None
|
||||
last_failure_error: str | None = None
|
||||
total_failures: int = 0
|
||||
total_successes: int = 0
|
||||
|
||||
|
||||
# ── Minimal HTTP health-check server ────────────────────────────────
|
||||
|
||||
class _HealthServer:
|
||||
"""Zero-dependency HTTP health-check endpoint using ``asyncio.start_server``.
|
||||
|
||||
Responds to ``GET /healthz`` with a JSON status payload; all other
|
||||
requests receive a 404. A per-connection timeout prevents slow
|
||||
clients from tying up the server.
|
||||
"""
|
||||
|
||||
_CONNECTION_TIMEOUT = 5.0 # seconds
|
||||
|
||||
def __init__(self, manager: ChannelManager, port: int) -> None:
|
||||
self._manager = manager
|
||||
self._port = port
|
||||
self._server: asyncio.AbstractServer | None = None
|
||||
self._start_time: float = 0.0
|
||||
|
||||
async def start(self) -> None:
|
||||
self._start_time = time.monotonic()
|
||||
self._server = await asyncio.start_server(
|
||||
self._handle_connection, "0.0.0.0", self._port,
|
||||
)
|
||||
addrs = [s.getsockname() for s in self._server.sockets]
|
||||
logger.info(f"Health server listening on {addrs}")
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._server is not None:
|
||||
self._server.close()
|
||||
await self._server.wait_closed()
|
||||
self._server = None
|
||||
logger.info("Health server stopped")
|
||||
|
||||
async def _handle_connection(
|
||||
self,
|
||||
reader: asyncio.StreamReader,
|
||||
writer: asyncio.StreamWriter,
|
||||
) -> None:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self._process_request(reader, writer),
|
||||
timeout=self._CONNECTION_TIMEOUT,
|
||||
)
|
||||
except (asyncio.TimeoutError, ConnectionError, OSError):
|
||||
pass
|
||||
finally:
|
||||
try:
|
||||
writer.close()
|
||||
await writer.wait_closed()
|
||||
except (ConnectionError, OSError):
|
||||
pass
|
||||
|
||||
async def _process_request(
|
||||
self,
|
||||
reader: asyncio.StreamReader,
|
||||
writer: asyncio.StreamWriter,
|
||||
) -> None:
|
||||
request_line = await reader.readline()
|
||||
# Consume remaining headers
|
||||
while True:
|
||||
line = await reader.readline()
|
||||
if line in (b"\r\n", b"\n", b""):
|
||||
break
|
||||
|
||||
parts = request_line.decode("utf-8", errors="replace").split()
|
||||
if len(parts) >= 2 and parts[0] == "GET" and parts[1] == "/healthz":
|
||||
body = self._build_response()
|
||||
payload = json.dumps(body).encode()
|
||||
header = (
|
||||
"HTTP/1.1 200 OK\r\n"
|
||||
"Content-Type: application/json\r\n"
|
||||
f"Content-Length: {len(payload)}\r\n"
|
||||
"Connection: close\r\n"
|
||||
"\r\n"
|
||||
)
|
||||
else:
|
||||
payload = b'{"error":"not found"}'
|
||||
header = (
|
||||
"HTTP/1.1 404 Not Found\r\n"
|
||||
"Content-Type: application/json\r\n"
|
||||
f"Content-Length: {len(payload)}\r\n"
|
||||
"Connection: close\r\n"
|
||||
"\r\n"
|
||||
)
|
||||
writer.write(header.encode() + payload)
|
||||
await writer.drain()
|
||||
|
||||
def _build_response(self) -> dict[str, Any]:
|
||||
mgr = self._manager
|
||||
health_map: dict[str, Any] = {}
|
||||
for name, h in mgr._health.items():
|
||||
health_map[name] = {
|
||||
"consecutive_failures": h.consecutive_failures,
|
||||
"total_successes": h.total_successes,
|
||||
"total_failures": h.total_failures,
|
||||
}
|
||||
accounts_map: dict[str, Any] = {}
|
||||
for state in mgr._account_manager.list_accounts():
|
||||
key = f"{state.channel_id}:{state.account_id}"
|
||||
accounts_map[key] = {
|
||||
"account_id": state.account_id,
|
||||
"channel": state.channel_id,
|
||||
"status": state.status,
|
||||
"error": state.error,
|
||||
}
|
||||
resp: dict[str, Any] = {
|
||||
"status": "healthy",
|
||||
"uptime_seconds": round(time.monotonic() - self._start_time, 1),
|
||||
"channels": {
|
||||
"enabled": mgr.enabled_channels,
|
||||
"running": mgr.running_channels(),
|
||||
},
|
||||
"queues": {
|
||||
"inbound_size": mgr.bus.inbound_size,
|
||||
"outbound_size": mgr.bus.outbound_size,
|
||||
},
|
||||
"health": health_map,
|
||||
"accounts": accounts_map,
|
||||
}
|
||||
for pname, provider in mgr._health_providers.items():
|
||||
try:
|
||||
resp[pname] = provider()
|
||||
except Exception:
|
||||
resp[pname] = {"error": "provider failed"}
|
||||
return resp
|
||||
|
||||
|
||||
# ── Channel registry ──────────────────────────────────────────────────
|
||||
|
||||
ChannelFactory = Callable[..., Channel]
|
||||
|
||||
_CHANNEL_REGISTRY: dict[str, ChannelFactory] = {}
|
||||
|
||||
|
||||
def _parse_csv(value: str) -> set[str] | None:
|
||||
"""Parse comma-separated string into a set, or ``None`` if empty."""
|
||||
if not value or not value.strip():
|
||||
return None
|
||||
items = {s.strip() for s in value.split(",") if s.strip()}
|
||||
return items if items else None
|
||||
|
||||
|
||||
def register_channel(name: str, factory: ChannelFactory) -> None:
|
||||
"""Register a channel factory under *name*."""
|
||||
_CHANNEL_REGISTRY[name] = factory
|
||||
|
||||
|
||||
def create_channel(name: str, config) -> Channel:
|
||||
"""Create a channel instance using the registered factory for *name*."""
|
||||
factory = _CHANNEL_REGISTRY.get(name)
|
||||
if not factory:
|
||||
raise ValueError(
|
||||
f"Unknown channel type: {name}. "
|
||||
f"Available: {list(_CHANNEL_REGISTRY.keys())}"
|
||||
)
|
||||
return factory(config)
|
||||
|
||||
|
||||
def available_channels() -> list[str]:
|
||||
"""Return the names of all available channel types.
|
||||
|
||||
Triggers auto-discovery if the registry is empty.
|
||||
"""
|
||||
if not _CHANNEL_REGISTRY:
|
||||
_ensure_channels_registered()
|
||||
return list(_CHANNEL_REGISTRY.keys())
|
||||
|
||||
|
||||
def _discover_channel_subpackages() -> list[str]:
|
||||
"""Discover all channel sub-packages under the channels directory.
|
||||
|
||||
Returns a list of sub-package names (e.g. ["telegram", "discord", ...]).
|
||||
Excludes non-channel directories (bus, __pycache__) and plain modules.
|
||||
"""
|
||||
channels_dir = Path(__file__).parent
|
||||
_EXCLUDED = {"bus", "__pycache__"}
|
||||
names = []
|
||||
for info in pkgutil.iter_modules([str(channels_dir)]):
|
||||
if info.ispkg and info.name not in _EXCLUDED:
|
||||
names.append(info.name)
|
||||
return sorted(names)
|
||||
|
||||
|
||||
def _ensure_channels_registered(types: list[str] | None = None) -> None:
|
||||
"""Lazily import channel sub-packages to trigger registration.
|
||||
|
||||
If *types* is given, only those channels are imported.
|
||||
If *types* is ``None``, all discovered channel sub-packages are imported.
|
||||
"""
|
||||
if types is None:
|
||||
targets = _discover_channel_subpackages()
|
||||
else:
|
||||
# Only import the ones that exist as sub-packages
|
||||
available = set(_discover_channel_subpackages())
|
||||
targets = [t for t in types if t in available]
|
||||
|
||||
for t in targets:
|
||||
module_name = f"EvoScientist.channels.{t}"
|
||||
if t not in _CHANNEL_REGISTRY:
|
||||
try:
|
||||
importlib.import_module(module_name)
|
||||
except ImportError as e:
|
||||
logger.debug(f"Could not import channel {t}: {e}")
|
||||
|
||||
|
||||
# ── Shared webhook server ─────────────────────────────────────────
|
||||
|
||||
class SharedWebhookServer:
|
||||
"""Single aiohttp server that hosts routes from multiple HTTP channels.
|
||||
|
||||
When ``shared_webhook_port`` is configured, ``ChannelManager`` collects
|
||||
routes from every channel that exposes ``_webhook_routes()`` and starts
|
||||
one server instead of letting each channel bind its own port.
|
||||
"""
|
||||
|
||||
def __init__(self, port: int) -> None:
|
||||
self._port = port
|
||||
self._app: Any = None
|
||||
self._runner: Any = None
|
||||
self._site: Any = None
|
||||
|
||||
async def start(self, routes: list[tuple[str, str, Any]]) -> None:
|
||||
from aiohttp import web
|
||||
|
||||
self._app = web.Application()
|
||||
for method, path, handler in routes:
|
||||
if method.upper() == "GET":
|
||||
self._app.router.add_get(path, handler)
|
||||
else:
|
||||
self._app.router.add_post(path, handler)
|
||||
|
||||
self._runner = web.AppRunner(self._app)
|
||||
await self._runner.setup()
|
||||
self._site = web.TCPSite(self._runner, "0.0.0.0", self._port)
|
||||
await self._site.start()
|
||||
logger.info(
|
||||
f"Shared webhook server started on 0.0.0.0:{self._port} "
|
||||
f"with {len(routes)} route(s)"
|
||||
)
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._site:
|
||||
await self._site.stop()
|
||||
self._site = None
|
||||
if self._runner:
|
||||
await self._runner.cleanup()
|
||||
self._runner = None
|
||||
logger.info("Shared webhook server stopped")
|
||||
|
||||
|
||||
class ChannelManager:
|
||||
"""Manages all chat channels and coordinates message routing.
|
||||
|
||||
Responsibilities:
|
||||
- Register channels and inject bus reference
|
||||
- Start / stop all channels
|
||||
- Route outbound messages from the bus to the correct channel
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
bus: MessageBus,
|
||||
*,
|
||||
health_port: int = 8080,
|
||||
drain_timeout: float = 30.0,
|
||||
shared_webhook_port: int = 0,
|
||||
):
|
||||
self.bus = bus
|
||||
self._channels: dict[str, Channel] = {}
|
||||
self._tasks: list[asyncio.Task] = []
|
||||
self._dispatch_task: asyncio.Task | None = None
|
||||
self._start_times: dict[str, datetime] = {}
|
||||
self._message_counts: dict[str, dict[str, int]] = {}
|
||||
self._health: dict[str, ChannelHealth] = {}
|
||||
self._is_running: bool = False
|
||||
self._health_port = health_port
|
||||
self._health_server: _HealthServer | None = None
|
||||
self._drain_timeout = drain_timeout
|
||||
self._health_providers: dict[str, Callable[[], dict]] = {}
|
||||
self._account_manager = AccountManager()
|
||||
# Pipelines (built during registration)
|
||||
self._outbound_pipelines: dict[str, OutboundPipeline] = {}
|
||||
# Shared webhook
|
||||
self._shared_webhook_port = shared_webhook_port
|
||||
self._shared_webhook_server: SharedWebhookServer | None = None
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config, bus: MessageBus | None = None) -> "ChannelManager":
|
||||
"""Create a ChannelManager from application config.
|
||||
|
||||
Parses ``config.channel_enabled`` (comma-separated channel types),
|
||||
creates each Channel instance, and registers them.
|
||||
|
||||
Args:
|
||||
config: Application config with channel settings.
|
||||
bus: Optional MessageBus instance. A new one is created if not provided.
|
||||
|
||||
Returns:
|
||||
A fully configured ChannelManager.
|
||||
"""
|
||||
if bus is None:
|
||||
bus = MessageBus()
|
||||
shared_webhook_port = getattr(config, "shared_webhook_port", 0) or 0
|
||||
manager = cls(bus, shared_webhook_port=shared_webhook_port)
|
||||
types = [t.strip() for t in (config.channel_enabled or "").split(",") if t.strip()]
|
||||
if not types:
|
||||
raise ValueError("No channels enabled")
|
||||
_ensure_channels_registered(types)
|
||||
for ct in types:
|
||||
channel = create_channel(ct, config)
|
||||
manager.register(channel, config=config)
|
||||
return manager
|
||||
|
||||
# ── registration ──
|
||||
|
||||
def register(
|
||||
self,
|
||||
channel: Channel,
|
||||
*,
|
||||
config: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> Channel:
|
||||
"""Register a channel and inject the bus reference.
|
||||
|
||||
Since Channel IS-A ChannelPlugin, the channel is also registered
|
||||
in the plugin registry. If *config* is provided, inbound/outbound
|
||||
pipelines are built for the channel.
|
||||
|
||||
Args:
|
||||
channel: The channel instance (must have a unique ``name``).
|
||||
config: Optional app config for building pipelines.
|
||||
**kwargs: Extra kwargs applied to the channel
|
||||
(e.g. ``send_thinking=True``, ``initial_debounce=3.0``).
|
||||
|
||||
Returns:
|
||||
The channel instance.
|
||||
"""
|
||||
name = channel.name
|
||||
if name in self._channels:
|
||||
raise ValueError(f"Channel '{name}' already registered")
|
||||
|
||||
channel.set_bus(self.bus)
|
||||
for key, value in kwargs.items():
|
||||
if hasattr(channel, key):
|
||||
setattr(channel, key, value)
|
||||
self._channels[name] = channel
|
||||
self._health[name] = ChannelHealth()
|
||||
if channel.config_adapter is not None:
|
||||
self._account_manager.register_plugin(channel)
|
||||
if config is not None:
|
||||
self._outbound_pipelines[name] = build_outbound_pipeline(channel, config)
|
||||
logger.info(f"Registered channel: {name} (slots: {channel.filled_slots()})")
|
||||
return channel
|
||||
|
||||
# ── lifecycle ──
|
||||
|
||||
async def start_all(self) -> None:
|
||||
"""Start the outbound dispatcher and all registered channels."""
|
||||
if not self._channels:
|
||||
logger.warning("No channels registered")
|
||||
return
|
||||
|
||||
self._is_running = True
|
||||
|
||||
await self.start_health()
|
||||
|
||||
# Start shared webhook server before individual channels
|
||||
await self._setup_shared_webhook()
|
||||
|
||||
self._dispatch_task = asyncio.create_task(
|
||||
self._dispatch_outbound()
|
||||
)
|
||||
|
||||
now = datetime.now()
|
||||
for name, channel in self._channels.items():
|
||||
logger.info(f"Starting channel: {name}")
|
||||
self._start_times[name] = now
|
||||
if name not in self._message_counts:
|
||||
self._message_counts[name] = {"received": 0, "sent": 0}
|
||||
task = asyncio.create_task(channel.run())
|
||||
self._tasks.append(task)
|
||||
|
||||
await asyncio.gather(*self._tasks, return_exceptions=True)
|
||||
|
||||
async def stop_all(self) -> None:
|
||||
"""Stop all channels and the outbound dispatcher.
|
||||
|
||||
Before shutting down channels, attempts to drain the outbound
|
||||
queue so that pending replies are delivered.
|
||||
"""
|
||||
logger.info("Stopping all channels...")
|
||||
self._is_running = False
|
||||
|
||||
# Drain outbound queue — try to send pending replies
|
||||
drained = 0
|
||||
deadline = time.monotonic() + self._drain_timeout
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
msg = self.bus.outbound.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
channel = self._channels.get(msg.channel)
|
||||
if channel and msg.content:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
channel.send(msg),
|
||||
timeout=max(1.0, deadline - time.monotonic()),
|
||||
)
|
||||
drained += 1
|
||||
except Exception:
|
||||
pass
|
||||
dropped = self.bus.outbound.qsize()
|
||||
if drained or dropped:
|
||||
logger.info(f"Outbound drain: {drained} sent, {dropped} dropped")
|
||||
|
||||
if self._dispatch_task:
|
||||
self._dispatch_task.cancel()
|
||||
try:
|
||||
await self._dispatch_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
for name, channel in self._channels.items():
|
||||
try:
|
||||
channel._running = False
|
||||
await channel.stop()
|
||||
logger.info(f"Stopped channel: {name}")
|
||||
except Exception as e:
|
||||
logger.error(f"Error stopping {name}: {e}")
|
||||
|
||||
for task in self._tasks:
|
||||
task.cancel()
|
||||
self._tasks.clear()
|
||||
|
||||
# Stop shared webhook server
|
||||
if self._shared_webhook_server is not None:
|
||||
await self._shared_webhook_server.stop()
|
||||
self._shared_webhook_server = None
|
||||
|
||||
await self.stop_health()
|
||||
|
||||
# ── health server ──
|
||||
|
||||
async def start_health(self) -> None:
|
||||
"""Start the HTTP health-check endpoint (if configured)."""
|
||||
if self._health_port and self._health_server is None:
|
||||
self._health_server = _HealthServer(self, self._health_port)
|
||||
await self._health_server.start()
|
||||
|
||||
async def stop_health(self) -> None:
|
||||
"""Stop the HTTP health-check endpoint."""
|
||||
if self._health_server is not None:
|
||||
await self._health_server.stop()
|
||||
self._health_server = None
|
||||
|
||||
# ── shared webhook ──
|
||||
|
||||
async def _setup_shared_webhook(self) -> None:
|
||||
"""Collect routes from HTTP channels and start a shared server.
|
||||
|
||||
Only active when ``shared_webhook_port > 0``. For each channel
|
||||
that exposes ``_webhook_routes()``, the routes are gathered and
|
||||
a sentinel attribute (``_shared_webhook_server``) is set so the
|
||||
channel's own ``start()`` skips creating its own aiohttp server.
|
||||
"""
|
||||
if not self._shared_webhook_port:
|
||||
return
|
||||
|
||||
all_routes: list[tuple[str, str, Any]] = []
|
||||
for name, channel in self._channels.items():
|
||||
routes_fn = getattr(channel, "_webhook_routes", None)
|
||||
if routes_fn is None:
|
||||
continue
|
||||
routes = routes_fn()
|
||||
if not routes:
|
||||
continue
|
||||
# Set sentinel so the channel skips its own server
|
||||
channel._shared_webhook_server = True # type: ignore[attr-defined]
|
||||
all_routes.extend(routes)
|
||||
logger.debug(
|
||||
f"Shared webhook: collected {len(routes)} route(s) "
|
||||
f"from '{name}'"
|
||||
)
|
||||
|
||||
if not all_routes:
|
||||
logger.info("Shared webhook: no HTTP channels found, skipping")
|
||||
return
|
||||
|
||||
self._shared_webhook_server = SharedWebhookServer(
|
||||
self._shared_webhook_port,
|
||||
)
|
||||
await self._shared_webhook_server.start(all_routes)
|
||||
|
||||
def register_health_provider(
|
||||
self, name: str, provider: Callable[[], dict],
|
||||
) -> None:
|
||||
"""Register a callable that returns extra data for ``/healthz``."""
|
||||
self._health_providers[name] = provider
|
||||
|
||||
# ── outbound routing ──
|
||||
|
||||
async def _dispatch_outbound(self) -> None:
|
||||
"""Route outbound messages from the bus to the correct channel."""
|
||||
logger.info("Outbound dispatcher started")
|
||||
while True:
|
||||
try:
|
||||
msg: OutboundMessage = await asyncio.wait_for(
|
||||
self.bus.consume_outbound(), timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
|
||||
channel = self._channels.get(msg.channel)
|
||||
|
||||
if not channel:
|
||||
logger.warning(f"Unknown channel: {msg.channel}")
|
||||
continue
|
||||
|
||||
try:
|
||||
# Run outbound pipeline if available (formatting, etc.)
|
||||
if msg.channel in self._outbound_pipelines:
|
||||
processed = await self._outbound_pipelines[msg.channel].process(msg)
|
||||
if processed is None:
|
||||
continue # dropped by pipeline
|
||||
msg = processed
|
||||
|
||||
if msg.content:
|
||||
await channel.send(msg)
|
||||
|
||||
for media_path in msg.media:
|
||||
try:
|
||||
await channel.send_media(
|
||||
recipient=msg.chat_id,
|
||||
file_path=media_path,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error sending media to {msg.channel}: {e}"
|
||||
)
|
||||
|
||||
# Success
|
||||
health = self._health.get(msg.channel)
|
||||
if health is not None:
|
||||
health.consecutive_failures = 0
|
||||
health.total_successes += 1
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error sending to {msg.channel}: {e}"
|
||||
)
|
||||
health = self._health.get(msg.channel)
|
||||
if health is not None:
|
||||
health.consecutive_failures += 1
|
||||
health.total_failures += 1
|
||||
health.last_failure_time = time.monotonic()
|
||||
health.last_failure_error = str(e)
|
||||
|
||||
# ── per-account lifecycle ──
|
||||
|
||||
async def start_account(
|
||||
self,
|
||||
channel_id: str,
|
||||
account_id: str,
|
||||
config: Any = None,
|
||||
) -> None:
|
||||
"""Start a specific account on a registered plugin."""
|
||||
await self._account_manager.start_account(channel_id, account_id, config)
|
||||
|
||||
async def stop_account(
|
||||
self,
|
||||
channel_id: str,
|
||||
account_id: str,
|
||||
) -> None:
|
||||
"""Stop a specific account on a registered plugin."""
|
||||
await self._account_manager.stop_account(channel_id, account_id)
|
||||
|
||||
def list_accounts(
|
||||
self,
|
||||
channel_id: str | None = None,
|
||||
) -> list[AccountState]:
|
||||
"""List account states, optionally filtered by channel."""
|
||||
return self._account_manager.list_accounts(channel_id)
|
||||
|
||||
@property
|
||||
def account_manager(self) -> AccountManager:
|
||||
"""Access the underlying AccountManager."""
|
||||
return self._account_manager
|
||||
|
||||
# ── queries ──
|
||||
|
||||
def get_channel(self, name: str) -> Channel | None:
|
||||
"""Get a channel by name."""
|
||||
return self._channels.get(name)
|
||||
|
||||
def get_server(self, name: str) -> Channel | None:
|
||||
"""Backward compat: returns the Channel (was ChannelServer)."""
|
||||
return self._channels.get(name)
|
||||
|
||||
def get_status(self) -> dict[str, Any]:
|
||||
"""Get status of all registered channels."""
|
||||
return {
|
||||
name: {
|
||||
"registered": True,
|
||||
"running": channel._running,
|
||||
"slots": channel.filled_slots(),
|
||||
}
|
||||
for name, channel in self._channels.items()
|
||||
}
|
||||
|
||||
@property
|
||||
def is_running(self) -> bool:
|
||||
"""Whether the manager is currently running."""
|
||||
return self._is_running
|
||||
|
||||
@property
|
||||
def enabled_channels(self) -> list[str]:
|
||||
"""List of registered channel names."""
|
||||
return list(self._channels.keys())
|
||||
|
||||
def running_channels(self) -> list[str]:
|
||||
"""Return names of currently running channels."""
|
||||
return [name for name, ch in self._channels.items() if ch._running]
|
||||
|
||||
def get_stats(self) -> dict:
|
||||
"""Return summary stats for all channels."""
|
||||
return {
|
||||
"channels": self.enabled_channels,
|
||||
"running": self.running_channels(),
|
||||
"message_counts": dict(self._message_counts),
|
||||
}
|
||||
|
||||
async def add_channel(self, channel_type: str, config) -> Channel:
|
||||
"""Dynamically add and start a channel at runtime."""
|
||||
_ensure_channels_registered([channel_type])
|
||||
channel = create_channel(channel_type, config)
|
||||
self.register(channel)
|
||||
self._start_times[channel_type] = datetime.now()
|
||||
if channel_type not in self._message_counts:
|
||||
self._message_counts[channel_type] = {"received": 0, "sent": 0}
|
||||
task = asyncio.create_task(channel.run())
|
||||
self._tasks.append(task)
|
||||
return channel
|
||||
|
||||
async def remove_channel(self, channel_type: str) -> None:
|
||||
"""Stop and remove a channel at runtime."""
|
||||
channel = self._channels.pop(channel_type, None)
|
||||
if channel:
|
||||
channel._running = False
|
||||
await channel.stop()
|
||||
logger.info(f"Removed channel: {channel_type}")
|
||||
|
||||
def record_message(self, channel_name: str, direction: str) -> None:
|
||||
"""Record a message for tracking.
|
||||
|
||||
Args:
|
||||
channel_name: Channel name (e.g. "telegram").
|
||||
direction: "received" or "sent".
|
||||
"""
|
||||
if channel_name not in self._message_counts:
|
||||
self._message_counts[channel_name] = {"received": 0, "sent": 0}
|
||||
if direction in self._message_counts[channel_name]:
|
||||
self._message_counts[channel_name][direction] += 1
|
||||
|
||||
def get_detailed_status(self) -> dict[str, Any]:
|
||||
"""Get detailed status of all registered channels.
|
||||
|
||||
Returns:
|
||||
Dict keyed by channel name with running, start_time, message
|
||||
counts, health, and plugin information.
|
||||
"""
|
||||
now = datetime.now()
|
||||
result = {}
|
||||
for name, channel in self._channels.items():
|
||||
start = self._start_times.get(name)
|
||||
counts = self._message_counts.get(name, {"received": 0, "sent": 0})
|
||||
health = self._health.get(name, ChannelHealth())
|
||||
result[name] = {
|
||||
"registered": True,
|
||||
"running": channel._running,
|
||||
"start_time": start,
|
||||
"uptime_seconds": (now - start).total_seconds() if start else 0,
|
||||
"received": counts["received"],
|
||||
"sent": counts["sent"],
|
||||
"health": {
|
||||
"consecutive_failures": health.consecutive_failures,
|
||||
"last_failure_time": health.last_failure_time,
|
||||
"last_failure_error": health.last_failure_error,
|
||||
"total_failures": health.total_failures,
|
||||
"total_successes": health.total_successes,
|
||||
},
|
||||
"plugin_slots": channel.filled_slots(),
|
||||
"has_outbound_pipeline": name in self._outbound_pipelines,
|
||||
}
|
||||
return result
|
||||
@@ -0,0 +1,126 @@
|
||||
"""Base configuration for all channel implementations.
|
||||
|
||||
Provides common fields shared across channels, reducing duplication.
|
||||
Channel-specific configs inherit from BaseChannelConfig.
|
||||
|
||||
Also provides ready-made ConfigAdapter implementations for the two most
|
||||
common account patterns:
|
||||
|
||||
- ``SingleAccountConfigAdapter`` — one account per channel (default).
|
||||
- ``MultiAccountConfigAdapter`` — multiple accounts from a config dict.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseChannelConfig:
|
||||
"""Common configuration fields for all channels.
|
||||
|
||||
Subclass this for channel-specific configs. Only add fields
|
||||
here that are used by 3+ channels.
|
||||
"""
|
||||
|
||||
allowed_senders: set[str] | None = None
|
||||
allowed_channels: set[str] | None = None
|
||||
text_chunk_limit: int = 4096
|
||||
proxy: str | None = None
|
||||
include_attachments: bool = True
|
||||
accounts: dict | None = None # multi-account config mapping
|
||||
|
||||
|
||||
class SingleAccountConfigAdapter:
|
||||
"""For channels that only ever have one account (most channels).
|
||||
|
||||
Returns a single ``"default"`` account whose config is the entire
|
||||
channel config object. This is the zero-change default: existing
|
||||
single-account channels get multi-account support for free.
|
||||
"""
|
||||
|
||||
def list_account_ids(self, config: Any) -> list[str]:
|
||||
return ["default"]
|
||||
|
||||
def resolve_account(
|
||||
self, config: Any, account_id: str | None = None,
|
||||
) -> Any:
|
||||
return config
|
||||
|
||||
def is_enabled(self, account: Any, config: Any) -> bool:
|
||||
return True
|
||||
|
||||
def is_configured(self, account: Any, config: Any) -> bool:
|
||||
"""Check that the account has at least some non-None values."""
|
||||
if account is None:
|
||||
return False
|
||||
if isinstance(account, dict):
|
||||
return bool(account)
|
||||
# dataclass / object — check that at least one field is truthy
|
||||
if hasattr(account, "__dataclass_fields__"):
|
||||
return any(
|
||||
getattr(account, f, None)
|
||||
for f in account.__dataclass_fields__
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
class MultiAccountConfigAdapter:
|
||||
"""For channels that support multiple accounts.
|
||||
|
||||
Expects the channel config to contain a mapping of accounts under
|
||||
a configurable key (default ``"accounts"``). Each entry is keyed
|
||||
by account id and holds account-specific settings.
|
||||
|
||||
Example config structure::
|
||||
|
||||
{
|
||||
"accounts": {
|
||||
"bot1": {"token": "...", "enabled": true},
|
||||
"bot2": {"token": "...", "enabled": false},
|
||||
}
|
||||
}
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
accounts_key: str = "accounts",
|
||||
required_fields: list[str] | None = None,
|
||||
) -> None:
|
||||
self._accounts_key = accounts_key
|
||||
self._required_fields = required_fields or []
|
||||
|
||||
def _get_accounts_map(self, config: Any) -> dict[str, Any]:
|
||||
"""Extract the accounts mapping from config."""
|
||||
if isinstance(config, dict):
|
||||
return config.get(self._accounts_key, {})
|
||||
return getattr(config, self._accounts_key, None) or {}
|
||||
|
||||
def list_account_ids(self, config: Any) -> list[str]:
|
||||
return list(self._get_accounts_map(config).keys())
|
||||
|
||||
def resolve_account(
|
||||
self, config: Any, account_id: str | None = None,
|
||||
) -> Any:
|
||||
accounts = self._get_accounts_map(config)
|
||||
if account_id is None:
|
||||
# Return the first account, or empty dict
|
||||
return next(iter(accounts.values()), {})
|
||||
return accounts.get(account_id, {})
|
||||
|
||||
def is_enabled(self, account: Any, config: Any) -> bool:
|
||||
if isinstance(account, dict):
|
||||
return account.get("enabled", True)
|
||||
return getattr(account, "enabled", True)
|
||||
|
||||
def is_configured(self, account: Any, config: Any) -> bool:
|
||||
if not account:
|
||||
return False
|
||||
for f in self._required_fields:
|
||||
if isinstance(account, dict):
|
||||
if not account.get(f):
|
||||
return False
|
||||
elif not getattr(account, f, None):
|
||||
return False
|
||||
return True
|
||||
@@ -0,0 +1,407 @@
|
||||
"""Unified inbound message consumer.
|
||||
|
||||
Provides :class:`InboundConsumer` — a single class that consumes
|
||||
inbound messages from the :class:`MessageBus`, runs them through
|
||||
the agent, and publishes outbound responses. This replaces the
|
||||
inline consumer loops that were duplicated in ``cli.py`` and
|
||||
``standalone.py``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncIterator, Callable, TypeVar
|
||||
|
||||
from .base import Channel
|
||||
from .bus import MessageBus
|
||||
from .bus.events import InboundMessage, OutboundMessage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
_MAX_CHAT_LOCKS = 10_000
|
||||
_MAX_SESSIONS = 10_000
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConsumerMetrics:
|
||||
"""Cumulative processing counters for the consumer."""
|
||||
|
||||
total_processed: int = 0
|
||||
total_successes: int = 0
|
||||
total_failures: int = 0
|
||||
total_timeouts: int = 0
|
||||
|
||||
|
||||
async def _timeout_aiter(
|
||||
agen: AsyncIterator[T],
|
||||
idle_timeout: float,
|
||||
) -> AsyncIterator[T]:
|
||||
"""Wrap an async iterator with a per-yield idle timeout.
|
||||
|
||||
If ``__anext__()`` does not produce a value within *idle_timeout*
|
||||
seconds, :class:`asyncio.TimeoutError` is raised. Continuous
|
||||
yielding resets the timer each time, so only a truly stalled
|
||||
generator will trigger the timeout.
|
||||
"""
|
||||
ait = agen.__aiter__()
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
item = await asyncio.wait_for(ait.__anext__(), timeout=idle_timeout)
|
||||
except StopAsyncIteration:
|
||||
return
|
||||
yield item
|
||||
finally:
|
||||
if hasattr(ait, "aclose"):
|
||||
await ait.aclose()
|
||||
|
||||
|
||||
def _format_todo_list(todos: list[dict]) -> str:
|
||||
"""Format todo items as a numbered list."""
|
||||
lines = ["\U0001f4cb Todo List\n"] # 📋
|
||||
for i, item in enumerate(todos, 1):
|
||||
content = item.get("content", "")
|
||||
lines.append(f"{i}. {content}")
|
||||
lines.append(f"\n\U0001f680 {len(todos)} tasks") # 🚀
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
class InboundConsumer:
|
||||
"""Consume inbound messages from the bus, process via agent, publish outbound.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
bus:
|
||||
The MessageBus to consume from / publish to.
|
||||
manager:
|
||||
The ChannelManager (used to look up channel instances).
|
||||
agent:
|
||||
The agent object (must support ``stream_agent_events``).
|
||||
thread_id:
|
||||
Default thread ID for agent conversations.
|
||||
send_thinking:
|
||||
Whether to forward thinking messages to the channel.
|
||||
on_message_received:
|
||||
Optional callback ``(msg: InboundMessage) -> None`` invoked when
|
||||
a message is consumed (e.g. for CLI Rich display).
|
||||
on_streaming_event:
|
||||
Optional callback ``(event: dict) -> None`` invoked for each
|
||||
streaming event from the agent.
|
||||
on_message_sent:
|
||||
Optional callback ``(msg: OutboundMessage) -> None`` invoked when
|
||||
the outbound message is published.
|
||||
inference_timeout:
|
||||
Per-yield idle timeout in seconds for the agent stream. If the
|
||||
agent produces no event for this long, the inference is aborted.
|
||||
max_concurrent:
|
||||
Number of worker coroutines (= max parallel inferences).
|
||||
max_pending:
|
||||
Maximum depth of the internal work queue. When full, the
|
||||
consumer loop blocks (back-pressure).
|
||||
drain_timeout:
|
||||
Seconds to wait for in-flight workers to finish during ``stop()``.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
bus: MessageBus,
|
||||
manager: Any,
|
||||
agent: Any,
|
||||
thread_id: str,
|
||||
*,
|
||||
send_thinking: bool = False,
|
||||
on_message_received: Callable[[InboundMessage], None] | None = None,
|
||||
on_streaming_event: Callable[[dict], None] | None = None,
|
||||
on_message_sent: Callable[[OutboundMessage], None] | None = None,
|
||||
inference_timeout: float = 300.0,
|
||||
max_concurrent: int = 5,
|
||||
max_pending: int = 50,
|
||||
drain_timeout: float = 30.0,
|
||||
):
|
||||
self.bus = bus
|
||||
self.manager = manager
|
||||
self.agent = agent
|
||||
self.thread_id = thread_id
|
||||
self.send_thinking = send_thinking
|
||||
self._on_message_received = on_message_received
|
||||
self._on_streaming_event = on_streaming_event
|
||||
self._on_message_sent = on_message_sent
|
||||
self._sessions: dict[str, str] = {} # sender_id -> thread_id
|
||||
|
||||
# Per-chat locks: same chat is processed serially (bounded)
|
||||
self._chat_locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
# Inference timeout
|
||||
self._inference_timeout = inference_timeout
|
||||
|
||||
# Worker pool
|
||||
self._max_concurrent = max_concurrent
|
||||
self._work_queue: asyncio.Queue[InboundMessage | None] = asyncio.Queue(
|
||||
maxsize=max_pending,
|
||||
)
|
||||
self._workers: list[asyncio.Task] = []
|
||||
self._stopping = False
|
||||
self._drain_timeout = drain_timeout
|
||||
|
||||
# Metrics
|
||||
self._metrics = ConsumerMetrics()
|
||||
|
||||
def _get_thread_id(self, sender_id: str) -> str:
|
||||
"""Get or create a thread ID for the given sender."""
|
||||
if sender_id not in self._sessions:
|
||||
if len(self._sessions) >= _MAX_SESSIONS:
|
||||
# Evict oldest entry
|
||||
oldest = next(iter(self._sessions))
|
||||
del self._sessions[oldest]
|
||||
if self.thread_id:
|
||||
self._sessions[sender_id] = f"{self.thread_id}:{sender_id}"
|
||||
else:
|
||||
self._sessions[sender_id] = str(uuid.uuid4())
|
||||
return self._sessions[sender_id]
|
||||
|
||||
def _get_channel(self, channel_name: str) -> Channel | None:
|
||||
"""Look up the channel by name from the manager."""
|
||||
return self.manager.get_channel(channel_name)
|
||||
|
||||
# ── lifecycle ──
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Main consumer loop — runs until ``stop()`` or cancellation.
|
||||
|
||||
Spawns *max_concurrent* worker coroutines that pull from an
|
||||
internal bounded queue. The loop reads from the bus and feeds
|
||||
the queue; when the queue is full the loop blocks (back-pressure).
|
||||
"""
|
||||
self._stopping = False
|
||||
self._workers = [
|
||||
asyncio.create_task(self._worker(i))
|
||||
for i in range(self._max_concurrent)
|
||||
]
|
||||
try:
|
||||
while not self._stopping:
|
||||
try:
|
||||
msg = await asyncio.wait_for(
|
||||
self.bus.consume_inbound(), timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
if self._stopping:
|
||||
break
|
||||
await self._work_queue.put(msg) # blocks when full (back-pressure)
|
||||
finally:
|
||||
if not self._stopping:
|
||||
await self.stop()
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Gracefully drain in-flight work and shut down workers."""
|
||||
self._stopping = True
|
||||
logger.info("Consumer stopping: draining in-flight messages...")
|
||||
pending_count = self._work_queue.qsize()
|
||||
|
||||
# Send a None sentinel per worker so each exits its loop
|
||||
for _ in self._workers:
|
||||
try:
|
||||
self._work_queue.put_nowait(None)
|
||||
except asyncio.QueueFull:
|
||||
pass
|
||||
|
||||
# Wait for workers to finish, then force-cancel stragglers
|
||||
if self._workers:
|
||||
done, still_running = await asyncio.wait(
|
||||
self._workers, timeout=self._drain_timeout,
|
||||
)
|
||||
for task in still_running:
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
logger.info(
|
||||
f"Consumer drain: {len(done)} finished, "
|
||||
f"{len(still_running)} force-cancelled, "
|
||||
f"{pending_count} were pending"
|
||||
)
|
||||
self._workers.clear()
|
||||
|
||||
# ── workers ──
|
||||
|
||||
async def _worker(self, worker_id: int) -> None:
|
||||
"""Pull messages from the work queue and process them."""
|
||||
while True:
|
||||
msg = await self._work_queue.get()
|
||||
if msg is None:
|
||||
break # shutdown sentinel
|
||||
try:
|
||||
await self._handle_message(msg)
|
||||
except Exception:
|
||||
logger.exception(f"Worker {worker_id} unhandled error")
|
||||
finally:
|
||||
self._work_queue.task_done()
|
||||
|
||||
async def _handle_message(self, msg: InboundMessage) -> None:
|
||||
"""Process a single inbound message."""
|
||||
from ..stream.events import stream_agent_events
|
||||
|
||||
if self._on_message_received:
|
||||
try:
|
||||
self._on_message_received(msg)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
channel = self._get_channel(msg.channel)
|
||||
thread_id = self._get_thread_id(msg.sender_id)
|
||||
session_key = msg.session_key # "channel:chat_id"
|
||||
|
||||
# Lazily create per-chat lock; evict stale locks when too many
|
||||
if session_key not in self._chat_locks:
|
||||
self._chat_locks[session_key] = asyncio.Lock()
|
||||
if len(self._chat_locks) > _MAX_CHAT_LOCKS:
|
||||
self._evict_chat_locks()
|
||||
|
||||
self._metrics.total_processed += 1
|
||||
|
||||
async with self._chat_locks[session_key]:
|
||||
try:
|
||||
final_content = ""
|
||||
thinking_buffer: list[str] = []
|
||||
todo_sent = False
|
||||
thinking_sent = False
|
||||
|
||||
if channel:
|
||||
await channel.start_typing(msg.chat_id)
|
||||
|
||||
async for event in _timeout_aiter(
|
||||
stream_agent_events(self.agent, msg.content, thread_id, media=msg.media or None),
|
||||
self._inference_timeout,
|
||||
):
|
||||
event_type = event.get("type")
|
||||
|
||||
if self._on_streaming_event:
|
||||
try:
|
||||
self._on_streaming_event(event)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if event_type == "thinking":
|
||||
thinking_text = event.get("content", "")
|
||||
if thinking_text:
|
||||
thinking_buffer.append(thinking_text)
|
||||
|
||||
elif event_type == "tool_call":
|
||||
if event.get("name") == "write_todos" and not todo_sent:
|
||||
todos = event.get("args", {}).get("todos", [])
|
||||
if todos and channel:
|
||||
if thinking_buffer and not thinking_sent:
|
||||
full_thinking = "".join(thinking_buffer)
|
||||
if full_thinking:
|
||||
await channel.send_thinking_message(
|
||||
msg.sender_id,
|
||||
full_thinking,
|
||||
msg.metadata,
|
||||
)
|
||||
thinking_sent = True
|
||||
thinking_buffer.clear()
|
||||
await channel.send_todo_message(
|
||||
msg.sender_id,
|
||||
_format_todo_list(todos),
|
||||
msg.metadata,
|
||||
)
|
||||
todo_sent = True
|
||||
|
||||
elif event_type == "text":
|
||||
final_content += event.get("content", "")
|
||||
|
||||
elif event_type == "done":
|
||||
final_content = event.get("content", "") or final_content
|
||||
|
||||
if thinking_buffer and not thinking_sent and channel:
|
||||
full_thinking = "".join(thinking_buffer)
|
||||
if full_thinking:
|
||||
await channel.send_thinking_message(
|
||||
msg.sender_id, full_thinking, msg.metadata,
|
||||
)
|
||||
|
||||
outbound = OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=final_content or "No response",
|
||||
reply_to=msg.message_id or None,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
await self.bus.publish_outbound(outbound)
|
||||
|
||||
self._metrics.total_successes += 1
|
||||
|
||||
if self._on_message_sent:
|
||||
try:
|
||||
self._on_message_sent(outbound)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
self._metrics.total_timeouts += 1
|
||||
logger.error(
|
||||
f"Inference timeout ({self._inference_timeout}s idle) "
|
||||
f"for {msg.sender_id} in {session_key}"
|
||||
)
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Sorry, the response timed out. Please try again.",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
|
||||
except Exception as e:
|
||||
self._metrics.total_failures += 1
|
||||
logger.error(f"Agent error: {e}")
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Sorry, something went wrong. Please try again later.",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
finally:
|
||||
if channel:
|
||||
await channel.stop_typing(msg.chat_id)
|
||||
|
||||
# ── observability ──
|
||||
|
||||
@property
|
||||
def pending_count(self) -> int:
|
||||
"""Number of messages waiting in the work queue."""
|
||||
return self._work_queue.qsize()
|
||||
|
||||
@property
|
||||
def active_workers(self) -> int:
|
||||
"""Number of worker tasks that are still alive."""
|
||||
return sum(1 for w in self._workers if not w.done())
|
||||
|
||||
@property
|
||||
def metrics(self) -> dict[str, int]:
|
||||
"""Cumulative processing counters."""
|
||||
m = self._metrics
|
||||
return {
|
||||
"total_processed": m.total_processed,
|
||||
"total_successes": m.total_successes,
|
||||
"total_failures": m.total_failures,
|
||||
"total_timeouts": m.total_timeouts,
|
||||
"pending": self.pending_count,
|
||||
"active_workers": self.active_workers,
|
||||
"chat_locks": len(self._chat_locks),
|
||||
"sessions": len(self._sessions),
|
||||
}
|
||||
|
||||
# ── internal ──
|
||||
|
||||
def _evict_chat_locks(self) -> None:
|
||||
"""Remove chat locks that are not currently held."""
|
||||
stale = [k for k, lock in self._chat_locks.items() if not lock.locked()]
|
||||
for k in stale[:max(1, len(stale) // 2)]:
|
||||
del self._chat_locks[k]
|
||||
@@ -0,0 +1,29 @@
|
||||
"""DingTalk (钉钉) channel for EvoScientist.
|
||||
|
||||
Uses Stream Mode (WebSocket) for receiving messages — no public IP needed.
|
||||
Sends replies via HTTP API.
|
||||
|
||||
Usage in config:
|
||||
channel_enabled = "dingtalk"
|
||||
dingtalk_client_id = "your_app_key"
|
||||
dingtalk_client_secret = "your_app_secret"
|
||||
"""
|
||||
|
||||
from .channel import DingTalkChannel, DingTalkConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["DingTalkChannel", "DingTalkConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> DingTalkChannel:
|
||||
allowed = _parse_csv(getattr(config, "dingtalk_allowed_senders", ""))
|
||||
proxy = getattr(config, "dingtalk_proxy", "") or None
|
||||
return DingTalkChannel(DingTalkConfig(
|
||||
client_id=getattr(config, "dingtalk_client_id", ""),
|
||||
client_secret=getattr(config, "dingtalk_client_secret", ""),
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
))
|
||||
|
||||
|
||||
register_channel("dingtalk", create_from_config)
|
||||
@@ -0,0 +1,354 @@
|
||||
"""DingTalk channel — refactored with WebSocketMixin + TokenMixin."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from urllib.parse import quote_plus
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..capabilities import DINGTALK as DINGTALK_CAPS
|
||||
from ..mixins import WebSocketMixin, TokenMixin
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
GATEWAY_URL = "https://api.dingtalk.com/v1.0/gateway/connections/open"
|
||||
TOKEN_URL = "https://api.dingtalk.com/v1.0/oauth2/accessToken"
|
||||
SEND_URL = "https://api.dingtalk.com/v1.0/robot/oToMessages/batchSend"
|
||||
MEDIA_SEND_URL = "https://api.dingtalk.com/v1.0/robot/oToMessages/batchSend"
|
||||
MEDIA_UPLOAD_URL = "https://oapi.dingtalk.com/media/upload"
|
||||
FILE_DOWNLOAD_URL = "https://api.dingtalk.com/v1.0/robot/messageFiles/download"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DingTalkConfig(BaseChannelConfig):
|
||||
client_id: str = ""
|
||||
client_secret: str = ""
|
||||
text_chunk_limit: int = 4096
|
||||
|
||||
|
||||
class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
capabilities = DINGTALK_CAPS
|
||||
name = "dingtalk"
|
||||
_ready_attrs = ("_http_client", "_access_token")
|
||||
_non_retryable_patterns = ("invalidauthentication", "forbidden", "40014")
|
||||
_mention_pattern = r"@\S+\s*"
|
||||
_mention_strip_count = 1
|
||||
|
||||
def __init__(self, config: DingTalkConfig):
|
||||
super().__init__(config)
|
||||
|
||||
async def start(self) -> None:
|
||||
import httpx
|
||||
if not self.config.client_id or not self.config.client_secret:
|
||||
raise ChannelError("DingTalk client_id and client_secret are required")
|
||||
self._http_client = httpx.AsyncClient(timeout=15, proxy=self.config.proxy)
|
||||
await self._refresh_token()
|
||||
self._running = True
|
||||
logger.info("DingTalk channel starting (Stream Mode)...")
|
||||
self._ws_task = asyncio.create_task(self._ws_loop())
|
||||
|
||||
# ── TokenMixin ────────────────────────────────────────────────
|
||||
|
||||
async def _fetch_token(self) -> tuple[str, int]:
|
||||
data = await self._api_post(TOKEN_URL, {
|
||||
"appKey": self.config.client_id,
|
||||
"appSecret": self.config.client_secret,
|
||||
})
|
||||
token = data.get("accessToken")
|
||||
if not token:
|
||||
raise ChannelError(f"DingTalk auth error: {data}")
|
||||
return token, int(data.get("expireIn", 7200))
|
||||
|
||||
async def _api_post(self, url, body, headers=None):
|
||||
resp = await self._http_client.post(url, json=body, headers=headers)
|
||||
return resp.json()
|
||||
|
||||
async def _resolve_download_code(self, download_code: str) -> str | None:
|
||||
"""Exchange a DingTalk downloadCode for a real download URL."""
|
||||
try:
|
||||
token = await self._ensure_token()
|
||||
data = await self._api_post(
|
||||
FILE_DOWNLOAD_URL,
|
||||
{"downloadCode": download_code, "robotCode": self.config.client_id},
|
||||
headers={"x-acs-dingtalk-access-token": token},
|
||||
)
|
||||
url = data.get("downloadUrl") or ""
|
||||
if url:
|
||||
return url
|
||||
logger.warning(f"DingTalk downloadCode resolve failed: {data}")
|
||||
except Exception as e:
|
||||
logger.warning(f"DingTalk downloadCode resolve error: {e}")
|
||||
return None
|
||||
|
||||
# ── WebSocketMixin ────────────────────────────────────────────
|
||||
|
||||
async def _get_ws_url(self) -> str:
|
||||
resp = await self._http_client.post(GATEWAY_URL, json={
|
||||
"clientId": self.config.client_id,
|
||||
"clientSecret": self.config.client_secret,
|
||||
"subscriptions": [{"type": "CALLBACK", "topic": "/v1.0/im/bot/messages/get"}],
|
||||
"ua": "dingtalk-sdk-python/v0.24.3-union",
|
||||
})
|
||||
data = resp.json()
|
||||
endpoint, ticket = data.get("endpoint"), data.get("ticket")
|
||||
if not endpoint or not ticket:
|
||||
raise ChannelError(f"DingTalk gateway failed: {data}")
|
||||
return f"{endpoint}?ticket={quote_plus(ticket)}"
|
||||
|
||||
async def _on_ws_message(self, data) -> None:
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
headers = data.get("headers", {})
|
||||
msg_id = headers.get("messageId", "")
|
||||
|
||||
# System ping
|
||||
if data.get("type") == "SYSTEM" and headers.get("topic") == "ping":
|
||||
await self._ws_send_json({"code": 200, "headers": headers, "message": "OK", "data": data.get("data", "")})
|
||||
return
|
||||
|
||||
# ACK
|
||||
await self._ws_send_json({"code": 200, "headers": {"contentType": "application/json", "messageId": msg_id}, "message": "OK", "data": "{}"})
|
||||
|
||||
if data.get("type") != "CALLBACK":
|
||||
return
|
||||
|
||||
payload = data.get("data", "{}")
|
||||
payload = json.loads(payload) if isinstance(payload, str) else payload
|
||||
text_obj = payload.get("text", {})
|
||||
content = (text_obj.get("content", "") if isinstance(text_obj, dict) else str(text_obj)).strip()
|
||||
if not content:
|
||||
raw_content = payload.get("content", "")
|
||||
content = raw_content.strip() if isinstance(raw_content, str) else ""
|
||||
|
||||
# Download attachments if present
|
||||
annotations: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
# DingTalk file/image messages may put download info in
|
||||
# payload["content"] (as a dict) instead of in a dedicated
|
||||
# "fileContent"/"imageContent" key.
|
||||
raw_content_obj = payload.get("content")
|
||||
if isinstance(raw_content_obj, dict) and raw_content_obj not in [
|
||||
payload.get(k) for k in ("imageContent", "fileContent", "videoContent", "audioContent")
|
||||
]:
|
||||
msg_type = payload.get("msgtype") or payload.get("msgType") or ""
|
||||
media_label = msg_type or "file"
|
||||
file_size = raw_content_obj.get("fileSize") or raw_content_obj.get("downloadSize") or 0
|
||||
file_name = raw_content_obj.get("fileName") or raw_content_obj.get("name") or f"dingtalk_{msg_type}"
|
||||
download_code = raw_content_obj.get("downloadCode") or ""
|
||||
download_url = raw_content_obj.get("downloadUrl") or ""
|
||||
# downloadCode is NOT a URL — resolve it via DingTalk API first
|
||||
if download_code and not download_code.startswith("http"):
|
||||
resolved = await self._resolve_download_code(download_code)
|
||||
if resolved:
|
||||
download_url = resolved
|
||||
elif download_code:
|
||||
download_url = download_code
|
||||
if download_url:
|
||||
try:
|
||||
dl_token = await self._ensure_token()
|
||||
dl_headers = {"x-acs-dingtalk-access-token": dl_token}
|
||||
except Exception:
|
||||
dl_headers = None
|
||||
local, ann = await self._download_attachment(
|
||||
download_url, f"dingtalk_{file_name}",
|
||||
headers=dl_headers,
|
||||
file_size=int(file_size) if file_size else None,
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
ann = ann.replace("[attachment:", f"[{media_label}:")
|
||||
annotations.append(ann)
|
||||
elif file_name:
|
||||
annotations.append(f"[{media_label}: {file_name}]")
|
||||
|
||||
for att_key in ("imageContent", "fileContent", "videoContent", "audioContent"):
|
||||
att = payload.get(att_key)
|
||||
if att and isinstance(att, dict):
|
||||
file_size = att.get("fileSize") or att.get("downloadSize") or 0
|
||||
file_name = att.get("fileName", att_key)
|
||||
download_code = att.get("downloadCode") or ""
|
||||
download_url = att.get("downloadUrl") or ""
|
||||
# Resolve downloadCode via API if it's not a URL
|
||||
if download_code and not download_code.startswith("http"):
|
||||
resolved = await self._resolve_download_code(download_code)
|
||||
if resolved:
|
||||
download_url = resolved
|
||||
elif download_code:
|
||||
download_url = download_code
|
||||
# DingTalk audioContent is voice messages
|
||||
media_label = "voice" if att_key == "audioContent" else att_key
|
||||
if download_url and (self.config.include_attachments if hasattr(self.config, 'include_attachments') else True):
|
||||
# DingTalk download URLs require access token
|
||||
try:
|
||||
dl_token = await self._ensure_token()
|
||||
dl_headers = {"x-acs-dingtalk-access-token": dl_token}
|
||||
except Exception:
|
||||
dl_headers = None
|
||||
local, ann = await self._download_attachment(
|
||||
download_url, f"dingtalk_{file_name}",
|
||||
headers=dl_headers,
|
||||
file_size=int(file_size) if file_size else None,
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
ann = ann.replace("[attachment:", f"[{media_label}:")
|
||||
annotations.append(ann)
|
||||
elif file_size:
|
||||
too_large = self._check_attachment_size(int(file_size), file_name)
|
||||
if too_large:
|
||||
annotations.append(too_large)
|
||||
else:
|
||||
annotations.append(f"[{media_label}: {file_name}]")
|
||||
|
||||
if not content and not media_paths and not annotations:
|
||||
return
|
||||
|
||||
sender_id = payload.get("senderStaffId") or payload.get("senderId", "")
|
||||
is_group = payload.get("conversationType") == "2"
|
||||
# For send API (oToMessages/batchSend), userIds needs staffId, not conversationId
|
||||
chat_id = sender_id
|
||||
create_time = payload.get("createAt") or payload.get("createTime", "")
|
||||
|
||||
# Mention gating: DMs always pass; groups require @bot
|
||||
was_mentioned = not is_group
|
||||
if is_group:
|
||||
# isInAtList is set by DingTalk when bot is @mentioned
|
||||
if payload.get("isInAtList"):
|
||||
was_mentioned = True
|
||||
else:
|
||||
# Fallback: check atUsers array
|
||||
at_users = payload.get("atUsers") or []
|
||||
for u in at_users:
|
||||
if u.get("dingtalkId") == self.config.client_id:
|
||||
was_mentioned = True
|
||||
break
|
||||
|
||||
try:
|
||||
ts = datetime.fromtimestamp(int(create_time) / 1000) if create_time else datetime.now()
|
||||
except (ValueError, TypeError, OSError):
|
||||
ts = datetime.now()
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender_id, chat_id=chat_id, text=content, timestamp=ts,
|
||||
message_id=msg_id, is_group=is_group, was_mentioned=was_mentioned,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
metadata={"chat_id": chat_id, "sender_nick": payload.get("senderNick", ""), "backend": "dingtalk"},
|
||||
))
|
||||
|
||||
# _send_typing_action: inherited no-op (DingTalk has no typing API)
|
||||
# _format_chunk: inherited from base (UnifiedFormatter)
|
||||
|
||||
# ── Send ──────────────────────────────────────────────────────
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
token = await self._ensure_token()
|
||||
data = await self._api_post(SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": raw_text, "title": "EvoScientist"}),
|
||||
}, headers={"x-acs-dingtalk-access-token": token})
|
||||
return data
|
||||
|
||||
# ── Media send ────────────────────────────────────────────────
|
||||
|
||||
_IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".webp"}
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a media file through DingTalk.
|
||||
|
||||
For images: uploads via /media/upload to get media_id, then sends
|
||||
as sampleImageMsg. Non-image files are sent as markdown links
|
||||
(DingTalk robot API does not support arbitrary file uploads).
|
||||
"""
|
||||
token = await self._ensure_token()
|
||||
chat_id = self._resolve_media_chat_id(recipient, metadata)
|
||||
headers = {"x-acs-dingtalk-access-token": token}
|
||||
ext = Path(file_path).suffix.lower()
|
||||
|
||||
if ext in self._IMAGE_EXTS:
|
||||
# Try uploading image to get media_id for native image message
|
||||
media_id = await self._upload_dingtalk_media(token, file_path, "image")
|
||||
if media_id:
|
||||
await self._api_post(MEDIA_SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleImageMsg",
|
||||
"msgParam": json.dumps({"photoURL": media_id}),
|
||||
}, headers=headers)
|
||||
else:
|
||||
# Fallback to markdown with file path
|
||||
await self._api_post(MEDIA_SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({
|
||||
"text": f"" + (f"\n{caption}" if caption else ""),
|
||||
"title": caption or "Image",
|
||||
}),
|
||||
}, headers=headers)
|
||||
else:
|
||||
# Non-image: send as markdown with filename
|
||||
name = Path(file_path).name
|
||||
text = f"[文件] {name}" + (f"\n{caption}" if caption else "")
|
||||
await self._api_post(MEDIA_SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": text, "title": name}),
|
||||
}, headers=headers)
|
||||
|
||||
if caption and ext in self._IMAGE_EXTS:
|
||||
# Send caption separately for image messages
|
||||
await self._api_post(MEDIA_SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": caption, "title": "Caption"}),
|
||||
}, headers=headers)
|
||||
return True
|
||||
|
||||
async def _upload_dingtalk_media(
|
||||
self, token: str, file_path: str, media_type: str = "image",
|
||||
) -> str | None:
|
||||
"""Upload a file to DingTalk media API and return the media_id."""
|
||||
try:
|
||||
url = f"{MEDIA_UPLOAD_URL}?access_token={token}&type={media_type}"
|
||||
with open(file_path, "rb") as f:
|
||||
resp = await self._http_client.post(
|
||||
url, files={"media": (Path(file_path).name, f)},
|
||||
)
|
||||
data = resp.json()
|
||||
return data.get("media_id")
|
||||
except Exception as e:
|
||||
logger.warning(f"DingTalk media upload failed: {e}")
|
||||
return None
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
if hasattr(self, "_ws_task") and self._ws_task:
|
||||
self._ws_task.cancel()
|
||||
try:
|
||||
await self._ws_task
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
self._ws_task = None
|
||||
await self._stop_ws()
|
||||
if self._http_client:
|
||||
await self._http_client.aclose()
|
||||
self._http_client = None
|
||||
self._access_token = None
|
||||
logger.info("DingTalk channel stopped")
|
||||
@@ -0,0 +1,33 @@
|
||||
"""DingTalk credential validation."""
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_dingtalk(
|
||||
client_id: str,
|
||||
client_secret: str,
|
||||
proxy: str | None = None,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate DingTalk credentials by fetching an access token."""
|
||||
if not client_id or not client_secret:
|
||||
return False, "client_id and client_secret are required"
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return False, "httpx not installed"
|
||||
|
||||
url = "https://api.dingtalk.com/v1.0/oauth2/accessToken"
|
||||
body = {"appKey": client_id, "appSecret": client_secret}
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(proxy=proxy) as client:
|
||||
resp = await client.post(url, json=body, timeout=10)
|
||||
data = resp.json()
|
||||
if data.get("accessToken"):
|
||||
return True, "DingTalk credentials valid"
|
||||
return False, f"Error: {data.get('message', data)}"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
@@ -0,0 +1,92 @@
|
||||
"""DingTalk channel server.
|
||||
|
||||
Standalone script to run the DingTalk channel with CLI options.
|
||||
|
||||
Usage:
|
||||
python -m EvoScientist.channels.dingtalk.serve --client-id ID --client-secret SECRET [OPTIONS]
|
||||
|
||||
Examples:
|
||||
# Basic usage
|
||||
python -m EvoScientist.channels.dingtalk.serve --client-id ID --client-secret SECRET
|
||||
|
||||
# With proxy and allowed senders
|
||||
python -m EvoScientist.channels.dingtalk.serve --client-id ID --client-secret SECRET --proxy http://proxy:8080 --allow user123
|
||||
|
||||
# With agent and thinking
|
||||
python -m EvoScientist.channels.dingtalk.serve --client-id ID --client-secret SECRET --agent --thinking
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import DingTalkChannel, DingTalkConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="DingTalk channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--client-id",
|
||||
required=True,
|
||||
help="DingTalk app client ID",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--client-secret",
|
||||
required=True,
|
||||
help="DingTalk app client secret",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender (DingTalk user ID). Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--proxy",
|
||||
help="HTTP proxy URL",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler (default: echo)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content as intermediate messages (requires --agent)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = DingTalkConfig(
|
||||
client_id=args.client_id,
|
||||
client_secret=args.client_secret,
|
||||
allowed_senders=set(args.allowed_senders) if args.allowed_senders else None,
|
||||
proxy=args.proxy,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = DingTalkChannel(config)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,19 @@
|
||||
from .channel import DiscordChannel, DiscordConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["DiscordChannel", "DiscordConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> DiscordChannel:
|
||||
allowed = _parse_csv(config.discord_allowed_senders)
|
||||
channels = _parse_csv(config.discord_allowed_channels)
|
||||
proxy = config.discord_proxy if config.discord_proxy else None
|
||||
return DiscordChannel(DiscordConfig(
|
||||
bot_token=config.discord_bot_token,
|
||||
allowed_senders=allowed,
|
||||
allowed_channels=channels,
|
||||
proxy=proxy,
|
||||
))
|
||||
|
||||
|
||||
register_channel("discord", create_from_config)
|
||||
@@ -0,0 +1,255 @@
|
||||
"""Discord channel implementation using discord.py."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..capabilities import DISCORD as DISCORD_CAPS
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DiscordConfig(BaseChannelConfig):
|
||||
bot_token: str = ""
|
||||
text_chunk_limit: int = 4096
|
||||
|
||||
|
||||
class DiscordChannel(Channel):
|
||||
"""Discord channel using discord.py."""
|
||||
|
||||
name = "discord"
|
||||
|
||||
capabilities = DISCORD_CAPS
|
||||
_typing_interval: float = 8.0
|
||||
_ready_attrs = ("_client",)
|
||||
_mention_pattern = r"<@!?{bot_id}>\s*"
|
||||
|
||||
def __init__(self, config: DiscordConfig):
|
||||
super().__init__(config)
|
||||
self._client = None
|
||||
self._ready = asyncio.Event()
|
||||
# Cache message objects for ACK reactions
|
||||
self._message_cache: dict[str, object] = {}
|
||||
self._MESSAGE_CACHE_MAX = 200
|
||||
|
||||
async def start(self) -> None:
|
||||
try:
|
||||
import discord
|
||||
except ImportError:
|
||||
raise ChannelError(
|
||||
"discord.py not installed. "
|
||||
"Install with: pip install evoscientist[discord]"
|
||||
)
|
||||
|
||||
if not self.config.bot_token:
|
||||
raise ChannelError("Discord bot token is required")
|
||||
|
||||
proxy = (
|
||||
self.config.proxy
|
||||
or os.environ.get("https_proxy")
|
||||
or os.environ.get("HTTPS_PROXY")
|
||||
or os.environ.get("http_proxy")
|
||||
or os.environ.get("HTTP_PROXY")
|
||||
or None
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Discord connect: token=%s...%s proxy=%s",
|
||||
self.config.bot_token[:8],
|
||||
self.config.bot_token[-4:],
|
||||
proxy or "(none)",
|
||||
)
|
||||
|
||||
intents = discord.Intents.default()
|
||||
intents.message_content = True
|
||||
client_kwargs = {"intents": intents}
|
||||
if proxy:
|
||||
client_kwargs["proxy"] = proxy
|
||||
self._client = discord.Client(**client_kwargs)
|
||||
|
||||
self._start_task_error: BaseException | None = None
|
||||
|
||||
@self._client.event
|
||||
async def on_ready():
|
||||
logger.info(f"Discord bot ready: {self._client.user}")
|
||||
self._ready.set()
|
||||
|
||||
@self._client.event
|
||||
async def on_message(message):
|
||||
await self._on_message(message)
|
||||
|
||||
async def _guarded_start():
|
||||
try:
|
||||
logger.info("Discord gateway: starting client.start()...")
|
||||
await self._client.start(self.config.bot_token)
|
||||
except Exception as exc:
|
||||
logger.error("Discord gateway error: %s: %s", type(exc).__name__, exc)
|
||||
self._start_task_error = exc
|
||||
self._ready.set() # unblock the waiter so it doesn't hang
|
||||
|
||||
logger.info("Discord connect: launching gateway task")
|
||||
asyncio.create_task(_guarded_start())
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(self._ready.wait(), timeout=60)
|
||||
except asyncio.TimeoutError:
|
||||
raise ChannelError(
|
||||
"Discord bot failed to connect within 60s. "
|
||||
"Check network/proxy connectivity to gateway.discord.gg"
|
||||
)
|
||||
|
||||
if self._start_task_error:
|
||||
raise ChannelError(
|
||||
f"Discord bot failed to connect: {self._start_task_error}"
|
||||
)
|
||||
|
||||
self._running = True
|
||||
logger.info("Discord channel started")
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
if self._client:
|
||||
await self._client.close()
|
||||
logger.info("Discord channel stopped")
|
||||
|
||||
# ── Typing indicator ────────────────────────────────────────────
|
||||
|
||||
async def _send_typing_action(self, chat_id: str) -> None:
|
||||
if not self._client:
|
||||
return
|
||||
ch = self._client.get_channel(int(chat_id))
|
||||
if ch:
|
||||
await ch.trigger_typing()
|
||||
|
||||
# ── ACK Reactions ───────────────────────────────────────────────
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
msg = self._message_cache.get(message_id)
|
||||
if msg:
|
||||
try:
|
||||
await msg.add_reaction(emoji)
|
||||
except Exception as e:
|
||||
logger.debug(f"Discord ACK reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
msg = self._message_cache.get(message_id)
|
||||
if msg and self._client and self._client.user:
|
||||
try:
|
||||
await msg.remove_reaction(emoji, self._client.user)
|
||||
except Exception as e:
|
||||
logger.debug(f"Discord remove ACK reaction failed: {e}")
|
||||
|
||||
def _cache_message(self, message) -> None:
|
||||
"""Cache a discord message object for later reaction use."""
|
||||
mid = str(message.id)
|
||||
self._message_cache[mid] = message
|
||||
# Evict oldest entries if cache is too large
|
||||
if len(self._message_cache) > self._MESSAGE_CACHE_MAX:
|
||||
oldest = list(self._message_cache.keys())[: self._MESSAGE_CACHE_MAX // 2]
|
||||
for k in oldest:
|
||||
self._message_cache.pop(k, None)
|
||||
|
||||
# ── Send ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
import discord
|
||||
|
||||
thread_id = (metadata or {}).get("thread_id", "")
|
||||
target_id = int(thread_id) if thread_id else int(chat_id)
|
||||
ch = self._client.get_channel(target_id)
|
||||
if not ch:
|
||||
raise RuntimeError(f"Discord channel {target_id} not found")
|
||||
ref = None
|
||||
if reply_to:
|
||||
try:
|
||||
ref = discord.MessageReference(
|
||||
message_id=int(reply_to), channel_id=target_id,
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
async def _send(text):
|
||||
await ch.send(text, reference=ref)
|
||||
|
||||
await self._send_with_format_fallback(_send, formatted_text, raw_text)
|
||||
|
||||
async def _send_media_impl(
|
||||
self, recipient: str, file_path: str,
|
||||
caption: str = "", metadata: dict | None = None,
|
||||
) -> bool:
|
||||
import discord
|
||||
|
||||
channel_id = self._resolve_media_chat_id(recipient, metadata)
|
||||
ch = self._client.get_channel(int(channel_id))
|
||||
if not ch:
|
||||
logger.error(f"Discord channel {channel_id} not found")
|
||||
return False
|
||||
file = discord.File(file_path)
|
||||
await ch.send(content=caption or None, file=file)
|
||||
return True
|
||||
|
||||
def _get_bot_identifier(self) -> str | None:
|
||||
if self._client and self._client.user:
|
||||
return str(self._client.user.id)
|
||||
return None
|
||||
|
||||
# ── Inbound ─────────────────────────────────────────────────────
|
||||
|
||||
async def _on_message(self, message) -> None:
|
||||
import discord
|
||||
|
||||
if message.author == self._client.user:
|
||||
return
|
||||
|
||||
# Cache for ACK reactions
|
||||
self._cache_message(message)
|
||||
|
||||
user_id = str(message.author.id)
|
||||
channel_id = str(message.channel.id)
|
||||
|
||||
is_dm = isinstance(message.channel, discord.DMChannel)
|
||||
was_mentioned = is_dm or (self._client.user in message.mentions)
|
||||
|
||||
text = message.content or ""
|
||||
annotations: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
if self.config.include_attachments and message.attachments:
|
||||
for attachment in message.attachments:
|
||||
too_large = self._check_attachment_size(
|
||||
attachment.size or 0, attachment.filename,
|
||||
)
|
||||
if too_large:
|
||||
annotations.append(too_large)
|
||||
continue
|
||||
try:
|
||||
safe_name = attachment.filename.replace("/", "_")
|
||||
file_path = self._media_path(f"{attachment.id}_{safe_name}")
|
||||
await attachment.save(file_path)
|
||||
media_paths.append(str(file_path))
|
||||
annotations.append(f"[attachment: {file_path}]")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to download Discord attachment: {e}")
|
||||
annotations.append(f"[attachment: {attachment.filename} - download failed]")
|
||||
|
||||
# Detect thread context
|
||||
thread_id = ""
|
||||
parent_channel_id = channel_id
|
||||
if hasattr(message.channel, "parent") and message.channel.parent:
|
||||
# Message is inside a Thread — store thread info
|
||||
thread_id = channel_id # the thread IS the channel
|
||||
parent_channel_id = str(message.channel.parent.id)
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=user_id, chat_id=parent_channel_id, text=text,
|
||||
media_files=media_paths, content_annotations=annotations,
|
||||
timestamp=message.created_at or datetime.now(),
|
||||
message_id=str(message.id),
|
||||
metadata={"chat_id": parent_channel_id, "thread_id": thread_id},
|
||||
is_group=not is_dm, was_mentioned=was_mentioned,
|
||||
))
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Discord bot token validation."""
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_discord_token(token: str, proxy: str | None = None) -> tuple[bool, str]:
|
||||
"""Validate a Discord bot token via the REST API.
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, message).
|
||||
"""
|
||||
if not token:
|
||||
return False, "No token provided"
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return False, "httpx not installed"
|
||||
|
||||
url = "https://discord.com/api/v10/users/@me"
|
||||
headers = {"Authorization": f"Bot {token}"}
|
||||
try:
|
||||
async with httpx.AsyncClient(proxy=proxy) as client:
|
||||
resp = await client.get(url, headers=headers, timeout=10)
|
||||
if resp.status_code == 200:
|
||||
data = resp.json()
|
||||
username = data.get("username", "unknown")
|
||||
return True, f"Bot: {username}"
|
||||
return False, "Invalid token"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Discord channel server.
|
||||
|
||||
Standalone script to run the Discord channel with CLI options.
|
||||
|
||||
Usage:
|
||||
python -m EvoScientist.channels.discord.serve --bot-token TOKEN [OPTIONS]
|
||||
|
||||
Examples:
|
||||
# Allow all senders (default)
|
||||
python -m EvoScientist.channels.discord.serve --bot-token TOKEN
|
||||
|
||||
# Only allow specific senders and channels
|
||||
python -m EvoScientist.channels.discord.serve --bot-token TOKEN --allow 123 --allow-channel 456
|
||||
|
||||
# With proxy, agent and thinking
|
||||
python -m EvoScientist.channels.discord.serve --bot-token TOKEN --proxy http://proxy:8080 --agent --thinking
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import DiscordChannel, DiscordConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Discord channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bot-token",
|
||||
required=True,
|
||||
help="Discord bot token",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender (Discord user ID). Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow-channel",
|
||||
action="append",
|
||||
dest="allowed_channels",
|
||||
help="Allowed channel ID. Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--proxy",
|
||||
help="HTTP proxy URL for Discord API requests",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler (default: echo)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content as intermediate messages (requires --agent)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = DiscordConfig(
|
||||
bot_token=args.bot_token,
|
||||
allowed_senders=set(args.allowed_senders) if args.allowed_senders else None,
|
||||
allowed_channels=set(args.allowed_channels) if args.allowed_channels else None,
|
||||
proxy=args.proxy,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = DiscordChannel(config)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Email channel for EvoScientist.
|
||||
|
||||
Uses IMAP polling for inbound + SMTP for outbound. Pure Python, no extra deps.
|
||||
|
||||
Usage in config:
|
||||
channel_enabled = "email"
|
||||
email_imap_host = "imap.gmail.com"
|
||||
email_smtp_host = "smtp.gmail.com"
|
||||
...
|
||||
"""
|
||||
|
||||
from .channel import EmailChannel, EmailConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["EmailChannel", "EmailConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> EmailChannel:
|
||||
allowed = _parse_csv(getattr(config, "email_allowed_senders", ""))
|
||||
return EmailChannel(EmailConfig(
|
||||
imap_host=getattr(config, "email_imap_host", ""),
|
||||
imap_port=int(getattr(config, "email_imap_port", 993)),
|
||||
imap_username=getattr(config, "email_imap_username", ""),
|
||||
imap_password=getattr(config, "email_imap_password", ""),
|
||||
imap_mailbox=getattr(config, "email_imap_mailbox", "INBOX"),
|
||||
imap_use_ssl=getattr(config, "email_imap_use_ssl", True),
|
||||
smtp_host=getattr(config, "email_smtp_host", ""),
|
||||
smtp_port=int(getattr(config, "email_smtp_port", 587)),
|
||||
smtp_username=getattr(config, "email_smtp_username", ""),
|
||||
smtp_password=getattr(config, "email_smtp_password", ""),
|
||||
smtp_use_tls=getattr(config, "email_smtp_use_tls", True),
|
||||
from_address=getattr(config, "email_from_address", ""),
|
||||
poll_interval=int(getattr(config, "email_poll_interval", 30)),
|
||||
mark_seen=getattr(config, "email_mark_seen", True),
|
||||
max_body_chars=int(getattr(config, "email_max_body_chars", 12000)),
|
||||
subject_prefix=getattr(config, "email_subject_prefix", "Re: "),
|
||||
allowed_senders=allowed,
|
||||
))
|
||||
|
||||
|
||||
register_channel("email", create_from_config)
|
||||
@@ -0,0 +1,372 @@
|
||||
"""Email channel — refactored with PollingMixin."""
|
||||
|
||||
import asyncio
|
||||
import email as email_lib
|
||||
import email.utils
|
||||
import html
|
||||
import imaplib
|
||||
import logging
|
||||
import re
|
||||
import smtplib
|
||||
import ssl
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from email.header import decode_header, make_header
|
||||
from email.message import EmailMessage
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from email.mime.text import MIMEText
|
||||
from email.mime.base import MIMEBase
|
||||
from email import encoders
|
||||
from email.utils import parseaddr
|
||||
from pathlib import Path
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..capabilities import EMAIL as EMAIL_CAPS
|
||||
from ..mixins import PollingMixin
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _decode_hdr(raw: str) -> str:
|
||||
try:
|
||||
return str(make_header(decode_header(raw))) if raw else ""
|
||||
except Exception:
|
||||
return raw or ""
|
||||
|
||||
|
||||
def _strip_html(text: str) -> str:
|
||||
text = re.sub(r"<br\s*/?>", "\n", text, flags=re.I)
|
||||
text = re.sub(r"<p[^>]*>", "\n", text, flags=re.I)
|
||||
text = re.sub(r"</p>", "\n", text, flags=re.I)
|
||||
text = re.sub(r"<[^>]+>", "", text)
|
||||
return html.unescape(text).strip()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@dataclass
|
||||
class EmailConfig(BaseChannelConfig):
|
||||
imap_host: str = ""
|
||||
imap_port: int = 993
|
||||
imap_username: str = ""
|
||||
imap_password: str = ""
|
||||
imap_mailbox: str = "INBOX"
|
||||
imap_use_ssl: bool = True
|
||||
smtp_host: str = ""
|
||||
smtp_port: int = 587
|
||||
smtp_username: str = ""
|
||||
smtp_password: str = ""
|
||||
smtp_use_tls: bool = True
|
||||
from_address: str = ""
|
||||
poll_interval: int = 30
|
||||
mark_seen: bool = True
|
||||
max_body_chars: int = 12000
|
||||
subject_prefix: str = "Re: "
|
||||
allowed_senders: set[str] | None = None
|
||||
text_chunk_limit: int = 4096
|
||||
|
||||
|
||||
class EmailChannel(Channel, PollingMixin):
|
||||
capabilities = EMAIL_CAPS
|
||||
name = "email"
|
||||
_non_retryable_patterns = ("auth", "login", "credential")
|
||||
|
||||
def __init__(self, config: EmailConfig):
|
||||
super().__init__(config)
|
||||
self._imap: imaplib.IMAP4_SSL | imaplib.IMAP4 | None = None
|
||||
|
||||
async def start(self) -> None:
|
||||
cfg = self.config
|
||||
if not cfg.imap_host or not cfg.imap_username:
|
||||
raise ChannelError("Email imap_host and imap_username are required")
|
||||
loop = asyncio.get_event_loop()
|
||||
await loop.run_in_executor(None, self._connect_imap)
|
||||
self._running = True
|
||||
logger.info(f"Email channel started (IMAP: {cfg.imap_host}, poll {cfg.poll_interval}s)")
|
||||
await self._start_polling()
|
||||
|
||||
def _connect_imap(self) -> None:
|
||||
cfg = self.config
|
||||
try:
|
||||
if cfg.imap_use_ssl:
|
||||
self._imap = imaplib.IMAP4_SSL(cfg.imap_host, cfg.imap_port, ssl_context=ssl.create_default_context())
|
||||
else:
|
||||
self._imap = imaplib.IMAP4(cfg.imap_host, cfg.imap_port)
|
||||
self._imap.login(cfg.imap_username, cfg.imap_password)
|
||||
self._imap.select(cfg.imap_mailbox)
|
||||
except Exception as e:
|
||||
raise ChannelError(f"IMAP failed: {e}")
|
||||
|
||||
def _reconnect_imap(self) -> None:
|
||||
try:
|
||||
if self._imap:
|
||||
self._imap.noop()
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
self._connect_imap()
|
||||
|
||||
async def _poll_once(self) -> None:
|
||||
loop = asyncio.get_event_loop()
|
||||
messages = await loop.run_in_executor(None, self._fetch_unseen)
|
||||
for m in messages:
|
||||
await self._process_email(m)
|
||||
|
||||
def _fetch_unseen(self) -> list[dict]:
|
||||
self._reconnect_imap()
|
||||
results = []
|
||||
try:
|
||||
st, data = self._imap.search(None, "UNSEEN")
|
||||
if st != "OK":
|
||||
return []
|
||||
for mid in data[0].split()[-20:]:
|
||||
st, msg_data = self._imap.fetch(mid, "(RFC822)")
|
||||
if st != "OK":
|
||||
continue
|
||||
msg = email_lib.message_from_bytes(msg_data[0][1])
|
||||
from_name, from_addr = parseaddr(msg.get("From", ""))
|
||||
body = self._extract_body(msg)
|
||||
if len(body) > self.config.max_body_chars:
|
||||
body = body[:self.config.max_body_chars] + "\n[...truncated]"
|
||||
# Extract attachments and inline images
|
||||
attachments = []
|
||||
if msg.is_multipart():
|
||||
for part in msg.walk():
|
||||
content_disp = part.get("Content-Disposition") or ""
|
||||
content_type = part.get_content_type() or ""
|
||||
is_attachment = "attachment" in content_disp.lower()
|
||||
is_inline_image = (
|
||||
"inline" in content_disp.lower()
|
||||
and content_type.startswith("image/")
|
||||
)
|
||||
# Also detect non-text parts with a filename but no
|
||||
# Content-Disposition header (common for PDFs, docs,
|
||||
# etc. sent by some email clients).
|
||||
is_named_file = (
|
||||
not is_attachment
|
||||
and not is_inline_image
|
||||
and part.get_filename()
|
||||
and not content_type.startswith("multipart/")
|
||||
and not content_type.startswith("text/")
|
||||
)
|
||||
if is_attachment or is_inline_image or is_named_file:
|
||||
filename = part.get_filename() or "attachment"
|
||||
filename = _decode_hdr(filename)
|
||||
payload_data = part.get_payload(decode=True)
|
||||
if payload_data:
|
||||
from ..base import MAX_ATTACHMENT_BYTES, MEDIA_DIR
|
||||
if len(payload_data) > MAX_ATTACHMENT_BYTES:
|
||||
attachments.append({"annotation": f"[attachment: {filename} - too large ({len(payload_data)} bytes)]"})
|
||||
else:
|
||||
MEDIA_DIR.mkdir(parents=True, exist_ok=True)
|
||||
local_path = MEDIA_DIR / f"email_{mid.decode()}_{filename}"
|
||||
local_path.write_bytes(payload_data)
|
||||
label = "inline-image" if is_inline_image else "attachment"
|
||||
attachments.append({"path": str(local_path), "annotation": f"[{label}: {local_path}]"})
|
||||
if self.config.mark_seen:
|
||||
self._imap.store(mid, "+FLAGS", "\\Seen")
|
||||
results.append({
|
||||
"from_addr": from_addr, "from_name": _decode_hdr(from_name),
|
||||
"subject": _decode_hdr(msg.get("Subject", "")), "body": body,
|
||||
"message_id": msg.get("Message-ID", ""), "date": msg.get("Date", ""),
|
||||
"references": msg.get("References", ""), "attachments": attachments,
|
||||
})
|
||||
except Exception as e:
|
||||
logger.error(f"IMAP fetch: {e}")
|
||||
return results
|
||||
|
||||
def _extract_body(self, msg) -> str:
|
||||
if msg.is_multipart():
|
||||
for part in msg.walk():
|
||||
ct = part.get_content_type()
|
||||
if ct == "text/plain":
|
||||
return self._decode_payload(part)
|
||||
for part in msg.walk():
|
||||
if part.get_content_type() == "text/html":
|
||||
return _strip_html(self._decode_payload(part))
|
||||
return "[no text content]"
|
||||
text = self._decode_payload(msg)
|
||||
return _strip_html(text) if msg.get_content_type() == "text/html" else text
|
||||
|
||||
@staticmethod
|
||||
def _decode_payload(part) -> str:
|
||||
payload = part.get_payload(decode=True)
|
||||
if not payload:
|
||||
return ""
|
||||
charset = part.get_content_charset() or "utf-8"
|
||||
return payload.decode(charset, errors="replace")
|
||||
|
||||
async def _process_email(self, m: dict) -> None:
|
||||
subject = m["subject"]
|
||||
text = f"[邮件] 主题: {subject}\n\n{m['body']}" if subject else m["body"]
|
||||
try:
|
||||
ts = email_lib.utils.parsedate_to_datetime(m["date"])
|
||||
except Exception:
|
||||
ts = datetime.now()
|
||||
# Process attachments
|
||||
media_paths: list[str] = []
|
||||
annotations: list[str] = []
|
||||
for att in m.get("attachments", []):
|
||||
if att.get("path"):
|
||||
media_paths.append(att["path"])
|
||||
if att.get("annotation"):
|
||||
annotations.append(att["annotation"])
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=m["from_addr"], chat_id=m["from_addr"], text=text, timestamp=ts,
|
||||
message_id=m["message_id"],
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
metadata={"chat_id": m["from_addr"], "subject": subject,
|
||||
"original_message_id": m["message_id"], "references": m["references"], "backend": "email"},
|
||||
))
|
||||
|
||||
# ── Send ──────────────────────────────────────────────────────
|
||||
|
||||
def _is_ready(self) -> bool:
|
||||
return bool(self.config.smtp_host)
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
loop = asyncio.get_event_loop()
|
||||
try:
|
||||
await loop.run_in_executor(
|
||||
None, self._smtp_send_html, chat_id, formatted_text, raw_text, metadata or {},
|
||||
)
|
||||
except Exception as e:
|
||||
err_str = str(e).lower()
|
||||
# Only fall back to plain text for format-related errors, not server rejections
|
||||
if any(code in err_str for code in ("550", "553", "554", "auth", "rejected")):
|
||||
raise
|
||||
logger.warning(f"HTML email failed ({e}), falling back to plain text")
|
||||
await loop.run_in_executor(
|
||||
None, self._smtp_send, chat_id, raw_text, metadata or {},
|
||||
)
|
||||
|
||||
def _smtp_send(self, to: str, content: str, meta: dict) -> None:
|
||||
cfg = self.config
|
||||
from_addr = cfg.from_address or cfg.smtp_username
|
||||
logger.debug(f"SMTP plain send: from={from_addr} to={to}")
|
||||
msg = EmailMessage()
|
||||
orig_subj = meta.get("subject", "")
|
||||
msg["Subject"] = f"{cfg.subject_prefix}{orig_subj}" if orig_subj and not orig_subj.lower().startswith("re:") else (orig_subj or "EvoScientist Reply")
|
||||
msg["From"] = from_addr
|
||||
msg["To"] = to
|
||||
orig_id = meta.get("original_message_id", "")
|
||||
if orig_id:
|
||||
msg["In-Reply-To"] = orig_id
|
||||
msg["References"] = f"{meta.get('references', '')} {orig_id}".strip()
|
||||
msg.set_content(content)
|
||||
try:
|
||||
if cfg.smtp_use_tls:
|
||||
srv = smtplib.SMTP(cfg.smtp_host, cfg.smtp_port, timeout=30)
|
||||
srv.starttls()
|
||||
else:
|
||||
srv = smtplib.SMTP_SSL(cfg.smtp_host, cfg.smtp_port, context=ssl.create_default_context(), timeout=30)
|
||||
srv.login(cfg.smtp_username, cfg.smtp_password)
|
||||
srv.sendmail(from_addr, [to], msg.as_string())
|
||||
srv.quit()
|
||||
except Exception as e:
|
||||
logger.error(f"SMTP send failed: from={from_addr} to={to} error={e}")
|
||||
raise RuntimeError(f"SMTP: {e}")
|
||||
|
||||
def _smtp_send_html(self, to: str, html_content: str, plain_content: str, meta: dict) -> None:
|
||||
"""Send an email with both HTML and plain-text parts."""
|
||||
cfg = self.config
|
||||
from_addr = cfg.from_address or cfg.smtp_username
|
||||
logger.debug(f"SMTP HTML send: from={from_addr} to={to}")
|
||||
msg = MIMEMultipart("alternative")
|
||||
orig_subj = meta.get("subject", "")
|
||||
msg["Subject"] = f"{cfg.subject_prefix}{orig_subj}" if orig_subj and not orig_subj.lower().startswith("re:") else (orig_subj or "EvoScientist Reply")
|
||||
msg["From"] = from_addr
|
||||
msg["To"] = to
|
||||
orig_id = meta.get("original_message_id", "")
|
||||
if orig_id:
|
||||
msg["In-Reply-To"] = orig_id
|
||||
msg["References"] = f"{meta.get('references', '')} {orig_id}".strip()
|
||||
msg.attach(MIMEText(plain_content, "plain", "utf-8"))
|
||||
msg.attach(MIMEText(html_content, "html", "utf-8"))
|
||||
try:
|
||||
if cfg.smtp_use_tls:
|
||||
srv = smtplib.SMTP(cfg.smtp_host, cfg.smtp_port, timeout=30)
|
||||
srv.starttls()
|
||||
else:
|
||||
srv = smtplib.SMTP_SSL(cfg.smtp_host, cfg.smtp_port, context=ssl.create_default_context(), timeout=30)
|
||||
srv.login(cfg.smtp_username, cfg.smtp_password)
|
||||
srv.sendmail(from_addr, [to], msg.as_string())
|
||||
srv.quit()
|
||||
except Exception as e:
|
||||
logger.error(f"SMTP HTML send failed: from={from_addr} to={to} error={e}")
|
||||
raise RuntimeError(f"SMTP HTML: {e}")
|
||||
|
||||
# ── Markdown to HTML formatting ───────────────────────────────
|
||||
|
||||
|
||||
# ── Media send (email attachment) ─────────────────────────────
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a file as an email attachment via SMTP."""
|
||||
loop = asyncio.get_event_loop()
|
||||
await loop.run_in_executor(
|
||||
None, self._smtp_send_attachment, recipient, file_path, caption, metadata or {},
|
||||
)
|
||||
return True
|
||||
|
||||
def _smtp_send_attachment(self, to: str, file_path: str, caption: str, meta: dict) -> None:
|
||||
"""Send an email with a file attachment."""
|
||||
cfg = self.config
|
||||
from_addr = cfg.from_address or cfg.smtp_username
|
||||
logger.debug(f"SMTP attachment send: from={from_addr} to={to} file={file_path}")
|
||||
msg = MIMEMultipart()
|
||||
orig_subj = meta.get("subject", "")
|
||||
msg["Subject"] = f"{cfg.subject_prefix}{orig_subj}" if orig_subj and not orig_subj.lower().startswith("re:") else (orig_subj or "EvoScientist Reply")
|
||||
msg["From"] = from_addr
|
||||
msg["To"] = to
|
||||
orig_id = meta.get("original_message_id", "")
|
||||
if orig_id:
|
||||
msg["In-Reply-To"] = orig_id
|
||||
msg["References"] = f"{meta.get('references', '')} {orig_id}".strip()
|
||||
|
||||
# Text body
|
||||
if caption:
|
||||
msg.attach(MIMEText(caption, "plain", "utf-8"))
|
||||
|
||||
# Attachment
|
||||
path = Path(file_path)
|
||||
part = MIMEBase("application", "octet-stream")
|
||||
part.set_payload(path.read_bytes())
|
||||
encoders.encode_base64(part)
|
||||
part.add_header("Content-Disposition", f"attachment; filename={path.name}")
|
||||
msg.attach(part)
|
||||
|
||||
try:
|
||||
if cfg.smtp_use_tls:
|
||||
srv = smtplib.SMTP(cfg.smtp_host, cfg.smtp_port, timeout=30)
|
||||
srv.starttls()
|
||||
else:
|
||||
srv = smtplib.SMTP_SSL(cfg.smtp_host, cfg.smtp_port, context=ssl.create_default_context(), timeout=30)
|
||||
srv.login(cfg.smtp_username, cfg.smtp_password)
|
||||
srv.sendmail(from_addr, [to], msg.as_string())
|
||||
srv.quit()
|
||||
except Exception as e:
|
||||
logger.error(f"SMTP attachment send failed: from={from_addr} to={to} error={e}")
|
||||
raise RuntimeError(f"SMTP attachment: {e}")
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
await self._stop_polling()
|
||||
if self._imap:
|
||||
try:
|
||||
self._imap.close()
|
||||
self._imap.logout()
|
||||
except Exception:
|
||||
pass
|
||||
self._imap = None
|
||||
logger.info("Email channel stopped")
|
||||
@@ -0,0 +1,67 @@
|
||||
"""Email credential validation."""
|
||||
|
||||
import imaplib
|
||||
import smtplib
|
||||
import ssl
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_email_imap(
|
||||
host: str, port: int, username: str, password: str,
|
||||
use_ssl: bool = True,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate IMAP credentials."""
|
||||
if not host or not username or not password:
|
||||
return False, "host, username, and password are required"
|
||||
|
||||
import asyncio
|
||||
loop = asyncio.get_event_loop()
|
||||
|
||||
def _check():
|
||||
try:
|
||||
if use_ssl:
|
||||
ctx = ssl.create_default_context()
|
||||
conn = imaplib.IMAP4_SSL(host, port, ssl_context=ctx)
|
||||
else:
|
||||
conn = imaplib.IMAP4(host, port)
|
||||
conn.login(username, password)
|
||||
conn.logout()
|
||||
return True, "IMAP credentials valid"
|
||||
except imaplib.IMAP4.error as e:
|
||||
return False, f"IMAP auth failed: {e}"
|
||||
except Exception as e:
|
||||
return False, f"IMAP error: {e}"
|
||||
|
||||
return await loop.run_in_executor(None, _check)
|
||||
|
||||
|
||||
async def validate_email_smtp(
|
||||
host: str, port: int, username: str, password: str,
|
||||
use_tls: bool = True,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate SMTP credentials."""
|
||||
if not host or not username or not password:
|
||||
return False, "host, username, and password are required"
|
||||
|
||||
import asyncio
|
||||
loop = asyncio.get_event_loop()
|
||||
|
||||
def _check():
|
||||
try:
|
||||
if use_tls:
|
||||
server = smtplib.SMTP(host, port, timeout=10)
|
||||
server.starttls()
|
||||
else:
|
||||
ctx = ssl.create_default_context()
|
||||
server = smtplib.SMTP_SSL(host, port, context=ctx, timeout=10)
|
||||
server.login(username, password)
|
||||
server.quit()
|
||||
return True, "SMTP credentials valid"
|
||||
except smtplib.SMTPAuthenticationError as e:
|
||||
return False, f"SMTP auth failed: {e}"
|
||||
except Exception as e:
|
||||
return False, f"SMTP error: {e}"
|
||||
|
||||
return await loop.run_in_executor(None, _check)
|
||||
@@ -0,0 +1,124 @@
|
||||
"""Email channel server.
|
||||
|
||||
Standalone script to run the Email channel with CLI options.
|
||||
|
||||
Usage:
|
||||
python -m EvoScientist.channels.email.serve --imap-host HOST --imap-username USER --imap-password PASS --smtp-host HOST --smtp-username USER --smtp-password PASS --from-address ADDR [OPTIONS]
|
||||
|
||||
Examples:
|
||||
# Basic usage
|
||||
python -m EvoScientist.channels.email.serve --imap-host imap.gmail.com --imap-username bot@gmail.com --imap-password PASS --smtp-host smtp.gmail.com --smtp-username bot@gmail.com --smtp-password PASS --from-address bot@gmail.com
|
||||
|
||||
# With allowed senders and custom poll interval
|
||||
python -m EvoScientist.channels.email.serve --imap-host imap.gmail.com --imap-username bot@gmail.com --imap-password PASS --smtp-host smtp.gmail.com --smtp-username bot@gmail.com --smtp-password PASS --from-address bot@gmail.com --allow user@example.com --poll-interval 60
|
||||
|
||||
# With agent and thinking
|
||||
python -m EvoScientist.channels.email.serve --imap-host imap.gmail.com --imap-username bot@gmail.com --imap-password PASS --smtp-host smtp.gmail.com --smtp-username bot@gmail.com --smtp-password PASS --from-address bot@gmail.com --agent --thinking
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import EmailChannel, EmailConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Email channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--imap-host",
|
||||
required=True,
|
||||
help="IMAP server hostname",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--imap-username",
|
||||
required=True,
|
||||
help="IMAP username",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--imap-password",
|
||||
required=True,
|
||||
help="IMAP password",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--smtp-host",
|
||||
required=True,
|
||||
help="SMTP server hostname",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--smtp-username",
|
||||
required=True,
|
||||
help="SMTP username",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--smtp-password",
|
||||
required=True,
|
||||
help="SMTP password",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--from-address",
|
||||
required=True,
|
||||
help="From email address for outgoing messages",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender (email address). Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--poll-interval",
|
||||
type=int,
|
||||
default=30,
|
||||
help="IMAP poll interval in seconds (default: 30)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler (default: echo)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content as intermediate messages (requires --agent)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = EmailConfig(
|
||||
imap_host=args.imap_host,
|
||||
imap_username=args.imap_username,
|
||||
imap_password=args.imap_password,
|
||||
smtp_host=args.smtp_host,
|
||||
smtp_username=args.smtp_username,
|
||||
smtp_password=args.smtp_password,
|
||||
from_address=args.from_address,
|
||||
allowed_senders=set(args.allowed_senders) if args.allowed_senders else None,
|
||||
poll_interval=args.poll_interval,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = EmailChannel(config)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,21 @@
|
||||
from .channel import FeishuChannel, FeishuConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["FeishuChannel", "FeishuConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> FeishuChannel:
|
||||
allowed = _parse_csv(config.feishu_allowed_senders)
|
||||
return FeishuChannel(FeishuConfig(
|
||||
app_id=config.feishu_app_id,
|
||||
app_secret=config.feishu_app_secret,
|
||||
verification_token=config.feishu_verification_token,
|
||||
encrypt_key=config.feishu_encrypt_key,
|
||||
webhook_port=config.feishu_webhook_port,
|
||||
allowed_senders=allowed,
|
||||
feishu_domain=config.feishu_domain,
|
||||
proxy=getattr(config, 'feishu_proxy', '') or None,
|
||||
))
|
||||
|
||||
|
||||
register_channel("feishu", create_from_config)
|
||||
@@ -0,0 +1,780 @@
|
||||
"""Feishu (飞书/Lark) channel implementation.
|
||||
|
||||
Receives messages via an HTTP event subscription webhook (aiohttp),
|
||||
sends replies via Feishu Open API REST endpoints.
|
||||
|
||||
Feishu Open API docs: https://open.feishu.cn/document
|
||||
|
||||
Authentication:
|
||||
- App ID + App Secret → tenant_access_token (2-hour TTL, auto-refreshed)
|
||||
|
||||
Event subscription:
|
||||
- URL verification challenge on first request
|
||||
- ``im.message.receive_v1`` events for incoming messages
|
||||
|
||||
Send API:
|
||||
- ``POST /open-apis/im/v1/messages?receive_id_type=chat_id``
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Any, TYPE_CHECKING
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from aiohttp import web
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..capabilities import FEISHU as FEISHU_CAPS
|
||||
from ..mixins import WebhookMixin, TokenMixin
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── Markdown → Feishu Post conversion ────────────────────────────
|
||||
|
||||
|
||||
def _parse_inline_text(text: str) -> list[dict]:
|
||||
"""Parse inline Markdown elements into Feishu post tag dicts.
|
||||
|
||||
Handles: `code`, **bold**, ~~strikethrough~~, [link](url), _italic_.
|
||||
"""
|
||||
elements: list[dict] = []
|
||||
# Pattern order matters: code first (protect content), then bold, strikethrough, link, italic
|
||||
pattern = re.compile(
|
||||
r"`([^`]+)`" # inline code
|
||||
r"|\*\*(.+?)\*\*" # bold
|
||||
r"|~~(.+?)~~" # strikethrough
|
||||
r"|\[([^\]]+)\]\(([^)]+)\)" # link
|
||||
r"|_(.+?)_" # italic
|
||||
)
|
||||
pos = 0
|
||||
for m in pattern.finditer(text):
|
||||
# Plain text before this match
|
||||
if m.start() > pos:
|
||||
elements.append({"tag": "text", "text": text[pos:m.start()]})
|
||||
|
||||
if m.group(1) is not None:
|
||||
# inline code → code_block would be block-level; use text with style
|
||||
elements.append({
|
||||
"tag": "text", "text": m.group(1),
|
||||
"style": ["code_block"],
|
||||
})
|
||||
elif m.group(2) is not None:
|
||||
elements.append({
|
||||
"tag": "text", "text": m.group(2),
|
||||
"style": ["bold"],
|
||||
})
|
||||
elif m.group(3) is not None:
|
||||
elements.append({
|
||||
"tag": "text", "text": m.group(3),
|
||||
"style": ["strikethrough"],
|
||||
})
|
||||
elif m.group(4) is not None:
|
||||
elements.append({
|
||||
"tag": "a", "text": m.group(4), "href": m.group(5),
|
||||
})
|
||||
elif m.group(6) is not None:
|
||||
elements.append({
|
||||
"tag": "text", "text": m.group(6),
|
||||
"style": ["italic"],
|
||||
})
|
||||
pos = m.end()
|
||||
|
||||
# Remaining plain text
|
||||
if pos < len(text):
|
||||
elements.append({"tag": "text", "text": text[pos:]})
|
||||
return elements
|
||||
|
||||
|
||||
def _parse_inline_elements(line: str) -> list[dict]:
|
||||
"""Parse a single Markdown line into a list of Feishu post elements.
|
||||
|
||||
Handles headings (→ bold), blockquotes (→ italic with prefix),
|
||||
list items (→ bullet prefix), and plain lines.
|
||||
"""
|
||||
# Heading: # Title → bold text
|
||||
heading_match = re.match(r"^(#{1,6})\s+(.+)$", line)
|
||||
if heading_match:
|
||||
return [{"tag": "text", "text": heading_match.group(2), "style": ["bold"]}]
|
||||
|
||||
# Blockquote: > text → italic with "▎" prefix
|
||||
quote_match = re.match(r"^>\s*(.*)$", line)
|
||||
if quote_match:
|
||||
inner = quote_match.group(1)
|
||||
elements = [{"tag": "text", "text": "▎", "style": ["italic"]}]
|
||||
elements.extend(_parse_inline_text(inner))
|
||||
return elements
|
||||
|
||||
# Unordered list: - item or * item → "• " prefix
|
||||
list_match = re.match(r"^[\-\*]\s+(.+)$", line)
|
||||
if list_match:
|
||||
elements = [{"tag": "text", "text": "• "}]
|
||||
elements.extend(_parse_inline_text(list_match.group(1)))
|
||||
return elements
|
||||
|
||||
# Ordered list: 1. item → keep number prefix
|
||||
ol_match = re.match(r"^(\d+)\.\s+(.+)$", line)
|
||||
if ol_match:
|
||||
elements = [{"tag": "text", "text": f"{ol_match.group(1)}. "}]
|
||||
elements.extend(_parse_inline_text(ol_match.group(2)))
|
||||
return elements
|
||||
|
||||
# Plain line
|
||||
return _parse_inline_text(line)
|
||||
|
||||
|
||||
def _markdown_to_feishu_post(text: str) -> dict | None:
|
||||
"""Convert Markdown text to Feishu post (rich text) JSON structure.
|
||||
|
||||
Returns a dict like {"zh_cn": {"content": [[...]]}} suitable for
|
||||
Feishu msg_type="post", or None if the text is empty.
|
||||
"""
|
||||
if not text or not text.strip():
|
||||
return None
|
||||
|
||||
paragraphs: list[list[dict]] = []
|
||||
current_paragraph: list[dict] = []
|
||||
in_code_block = False
|
||||
code_lines: list[str] = []
|
||||
code_lang = ""
|
||||
|
||||
for line in text.split("\n"):
|
||||
# Code block fences
|
||||
if line.startswith("```"):
|
||||
if not in_code_block:
|
||||
# Flush any pending paragraph
|
||||
if current_paragraph:
|
||||
paragraphs.append(current_paragraph)
|
||||
current_paragraph = []
|
||||
in_code_block = True
|
||||
code_lang = line[3:].strip()
|
||||
code_lines = []
|
||||
else:
|
||||
# End of code block
|
||||
code_text = "\n".join(code_lines)
|
||||
paragraphs.append([{
|
||||
"tag": "code_block",
|
||||
"language": code_lang or "plain",
|
||||
"text": code_text,
|
||||
}])
|
||||
in_code_block = False
|
||||
code_lines = []
|
||||
code_lang = ""
|
||||
continue
|
||||
|
||||
if in_code_block:
|
||||
code_lines.append(line)
|
||||
continue
|
||||
|
||||
# Empty line → new paragraph
|
||||
if not line.strip():
|
||||
if current_paragraph:
|
||||
paragraphs.append(current_paragraph)
|
||||
current_paragraph = []
|
||||
continue
|
||||
|
||||
# Non-empty line
|
||||
elements = _parse_inline_elements(line)
|
||||
if elements:
|
||||
# Each visual line becomes its own paragraph in Feishu post
|
||||
if current_paragraph:
|
||||
paragraphs.append(current_paragraph)
|
||||
current_paragraph = elements
|
||||
|
||||
# Flush remaining
|
||||
if in_code_block and code_lines:
|
||||
code_text = "\n".join(code_lines)
|
||||
paragraphs.append([{
|
||||
"tag": "code_block",
|
||||
"language": code_lang or "plain",
|
||||
"text": code_text,
|
||||
}])
|
||||
elif current_paragraph:
|
||||
paragraphs.append(current_paragraph)
|
||||
|
||||
if not paragraphs:
|
||||
return None
|
||||
|
||||
return {"zh_cn": {"content": paragraphs}}
|
||||
|
||||
|
||||
@dataclass
|
||||
class FeishuConfig(BaseChannelConfig):
|
||||
app_id: str = ""
|
||||
app_secret: str = ""
|
||||
verification_token: str = ""
|
||||
encrypt_key: str = ""
|
||||
webhook_port: int = 9000
|
||||
text_chunk_limit: int = 4096
|
||||
feishu_domain: str = "https://open.feishu.cn"
|
||||
|
||||
|
||||
class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
capabilities = FEISHU_CAPS
|
||||
"""Feishu channel using Open API + event subscription webhook."""
|
||||
|
||||
name = "feishu"
|
||||
_ready_attrs = ("_http_client", "_access_token")
|
||||
_rate_limit_patterns = ("99991400", "rate limit", "频率限制")
|
||||
_rate_limit_delay = 2.0
|
||||
|
||||
def __init__(self, config: FeishuConfig):
|
||||
super().__init__(config)
|
||||
self._mention_names: list[str] = [] # bot mention keys from events
|
||||
|
||||
# ── WebhookMixin overrides ────────────────────────────────────
|
||||
|
||||
def _get_webhook_port(self) -> int:
|
||||
return self.config.webhook_port
|
||||
|
||||
def _webhook_routes(self) -> list[tuple[str, str, Any]]:
|
||||
return [("POST", "/webhook/event", self._handle_event)]
|
||||
|
||||
# ── TokenMixin overrides ──────────────────────────────────────
|
||||
|
||||
async def _fetch_token(self) -> tuple[str, int]:
|
||||
"""Fetch Feishu tenant_access_token."""
|
||||
url = f"{self.config.feishu_domain}/open-apis/auth/v3/tenant_access_token/internal"
|
||||
body = {
|
||||
"app_id": self.config.app_id,
|
||||
"app_secret": self.config.app_secret,
|
||||
}
|
||||
try:
|
||||
resp = await self._http_client.post(url, json=body)
|
||||
data = resp.json()
|
||||
except Exception as e:
|
||||
raise ChannelError(f"Failed to get Feishu access token: {e}")
|
||||
|
||||
if data.get("code") != 0:
|
||||
raise ChannelError(
|
||||
f"Feishu auth error: {data.get('msg', 'unknown')}"
|
||||
)
|
||||
return data["tenant_access_token"], data.get("expire", 7200)
|
||||
|
||||
# ── Lifecycle ─────────────────────────────────────────────────
|
||||
|
||||
async def start(self) -> None:
|
||||
try:
|
||||
from aiohttp import web # noqa: F401
|
||||
import httpx # noqa: F401
|
||||
except ImportError:
|
||||
raise ChannelError(
|
||||
"aiohttp or httpx not installed. "
|
||||
"Install with: pip install aiohttp httpx"
|
||||
)
|
||||
|
||||
if not self.config.app_id:
|
||||
raise ChannelError("Feishu app_id is required")
|
||||
if not self.config.app_secret:
|
||||
raise ChannelError("Feishu app_secret is required")
|
||||
|
||||
# Start webhook server (sets up self._http_client)
|
||||
await self._start_webhook_server()
|
||||
|
||||
# Verify credentials by fetching initial token
|
||||
await self._refresh_token()
|
||||
|
||||
self._running = True
|
||||
logger.info(
|
||||
f"Feishu channel started "
|
||||
f"(webhook on port {self.config.webhook_port})"
|
||||
)
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
await self._stop_webhook_server()
|
||||
self._access_token = None
|
||||
logger.info("Feishu channel stopped")
|
||||
|
||||
# ── Token helpers (adapt old API to mixin) ────────────────────
|
||||
|
||||
async def _ensure_token(self) -> str:
|
||||
"""Return a valid access token, refreshing if needed."""
|
||||
return await TokenMixin._ensure_token(self)
|
||||
|
||||
# ── Send (template method overrides) ──────────────────────────
|
||||
|
||||
async def _feishu_send(self, url: str, body: dict, headers: dict) -> bool:
|
||||
"""POST to Feishu API and return True if code==0."""
|
||||
try:
|
||||
resp = await self._http_client.post(url, json=body, headers=headers)
|
||||
return resp.json().get("code") == 0
|
||||
except Exception as e:
|
||||
logger.warning(f"Feishu send error: {e}")
|
||||
return False
|
||||
|
||||
async def _send_chunk(
|
||||
self, chat_id, formatted_text, raw_text, reply_to, metadata,
|
||||
):
|
||||
token = await self._ensure_token()
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
post_content = _markdown_to_feishu_post(raw_text)
|
||||
|
||||
# If reply_to is set, try the reply API first
|
||||
if reply_to:
|
||||
reply_url = (
|
||||
f"{self.config.feishu_domain}"
|
||||
f"/open-apis/im/v1/messages/{reply_to}/reply"
|
||||
)
|
||||
if post_content is not None:
|
||||
body = {"msg_type": "post", "content": json.dumps(post_content)}
|
||||
else:
|
||||
body = {"msg_type": "text", "content": json.dumps({"text": formatted_text})}
|
||||
if await self._feishu_send(reply_url, body, headers):
|
||||
return
|
||||
|
||||
# Normal send (non-reply or reply fallback)
|
||||
url = (
|
||||
f"{self.config.feishu_domain}"
|
||||
f"/open-apis/im/v1/messages?receive_id_type=chat_id"
|
||||
)
|
||||
|
||||
# Try post format first
|
||||
if post_content is not None:
|
||||
body = {
|
||||
"receive_id": chat_id,
|
||||
"msg_type": "post",
|
||||
"content": json.dumps(post_content),
|
||||
}
|
||||
if await self._feishu_send(url, body, headers):
|
||||
return
|
||||
|
||||
# Fallback: plain text
|
||||
body = {
|
||||
"receive_id": chat_id,
|
||||
"msg_type": "text",
|
||||
"content": json.dumps({"text": formatted_text}),
|
||||
}
|
||||
if not await self._feishu_send(url, body, headers):
|
||||
raise RuntimeError("Feishu send failed")
|
||||
|
||||
# ── Media helpers ──────────────────────────────────────────────
|
||||
|
||||
_IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".webp"}
|
||||
|
||||
async def _download_media(
|
||||
self, message_id: str, file_key: str, msg_type: str,
|
||||
) -> str | None:
|
||||
"""Download an image or file attachment from Feishu.
|
||||
|
||||
Returns the local file path on success, or None on failure.
|
||||
"""
|
||||
token = await self._ensure_token()
|
||||
resource_type = "image" if msg_type == "image" else "file"
|
||||
url = (
|
||||
f"{self.config.feishu_domain}"
|
||||
f"/open-apis/im/v1/messages/{message_id}"
|
||||
f"/resources/{file_key}?type={resource_type}"
|
||||
)
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
try:
|
||||
resp = await self._http_client.get(url, headers=headers, timeout=30)
|
||||
if resp.status_code != 200:
|
||||
logger.warning(
|
||||
f"Feishu media download failed: HTTP {resp.status_code}"
|
||||
)
|
||||
return None
|
||||
|
||||
# Check attachment size before writing to disk
|
||||
cl = resp.headers.get("content-length")
|
||||
if cl:
|
||||
try:
|
||||
too_large = self._check_attachment_size(int(cl), file_key)
|
||||
if too_large:
|
||||
logger.warning(too_large)
|
||||
return None
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
from ..base import MAX_ATTACHMENT_BYTES
|
||||
if len(resp.content) > MAX_ATTACHMENT_BYTES:
|
||||
logger.warning(
|
||||
f"Feishu media too large: {len(resp.content)} bytes"
|
||||
)
|
||||
return None
|
||||
|
||||
# Determine extension from Content-Type or default
|
||||
content_type = resp.headers.get("content-type", "")
|
||||
ext_map = {
|
||||
"image/jpeg": ".jpg",
|
||||
"image/png": ".png",
|
||||
"image/gif": ".gif",
|
||||
"image/webp": ".webp",
|
||||
"image/bmp": ".bmp",
|
||||
}
|
||||
ext = ext_map.get(content_type, ".bin")
|
||||
local_path = self._media_path(f"feishu_{message_id}_{file_key}{ext}")
|
||||
local_path.write_bytes(resp.content)
|
||||
return str(local_path)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to download Feishu media: {e}")
|
||||
return None
|
||||
|
||||
async def _upload_feishu_resource(
|
||||
self, url: str, headers: dict, file_path: str,
|
||||
field_name: str, extra_data: dict,
|
||||
) -> dict | None:
|
||||
"""Upload a file to Feishu API. Returns response data or None on failure."""
|
||||
with open(file_path, "rb") as f:
|
||||
resp = await self._http_client.post(
|
||||
url, headers=headers, data=extra_data,
|
||||
files={field_name: (Path(file_path).name, f)},
|
||||
)
|
||||
data = resp.json()
|
||||
if data.get("code") != 0:
|
||||
logger.error(f"Feishu upload failed: {data.get('msg')}")
|
||||
return None
|
||||
return data["data"]
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a media file through Feishu."""
|
||||
token = await self._ensure_token()
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
chat_id = self._resolve_media_chat_id(recipient, metadata)
|
||||
|
||||
path = Path(file_path)
|
||||
ext = path.suffix.lower()
|
||||
is_image = ext in self._IMAGE_EXTENSIONS
|
||||
|
||||
send_url = (
|
||||
f"{self.config.feishu_domain}"
|
||||
f"/open-apis/im/v1/messages?receive_id_type=chat_id"
|
||||
)
|
||||
|
||||
if is_image:
|
||||
upload_url = f"{self.config.feishu_domain}/open-apis/im/v1/images"
|
||||
data = await self._upload_feishu_resource(
|
||||
upload_url, headers, file_path, "image", {"image_type": "message"},
|
||||
)
|
||||
if not data:
|
||||
return False
|
||||
body = {
|
||||
"receive_id": chat_id,
|
||||
"msg_type": "image",
|
||||
"content": json.dumps({"image_key": data["image_key"]}),
|
||||
}
|
||||
else:
|
||||
upload_url = f"{self.config.feishu_domain}/open-apis/im/v1/files"
|
||||
data = await self._upload_feishu_resource(
|
||||
upload_url, headers, file_path, "file",
|
||||
{"file_type": "stream", "file_name": path.name},
|
||||
)
|
||||
if not data:
|
||||
return False
|
||||
body = {
|
||||
"receive_id": chat_id,
|
||||
"msg_type": "file",
|
||||
"content": json.dumps({"file_key": data["file_key"]}),
|
||||
}
|
||||
|
||||
if not await self._feishu_send(send_url, body, headers):
|
||||
return False
|
||||
|
||||
# Send caption as a separate text message if provided
|
||||
if caption:
|
||||
cap_body = {
|
||||
"receive_id": chat_id,
|
||||
"msg_type": "text",
|
||||
"content": json.dumps({"text": caption}),
|
||||
}
|
||||
await self._feishu_send(send_url, cap_body, headers)
|
||||
|
||||
return True
|
||||
|
||||
# ── ACK reaction ───────────────────────────────────────────────
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
"""Send an acknowledgment reaction via Feishu Open API."""
|
||||
try:
|
||||
token = await self._ensure_token()
|
||||
url = f"{self.config.feishu_domain}/open-apis/im/v1/messages/{message_id}/reactions"
|
||||
await self._http_client.post(
|
||||
url,
|
||||
json={"reaction_type": {"emoji_type": emoji}},
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Feishu ack reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
"""Remove ACK reaction via Feishu Open API.
|
||||
|
||||
Feishu's DELETE /reactions endpoint requires the reaction_id, which
|
||||
we don't track. No-op for now.
|
||||
"""
|
||||
pass
|
||||
|
||||
# ── Format chunk ─────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── Typing indicator ──────────────────────────────────────────
|
||||
|
||||
# Feishu Open API has no typing indicator endpoint.
|
||||
|
||||
# ── Mention stripping ─────────────────────────────────────────
|
||||
|
||||
def _strip_mention(self, text: str) -> str:
|
||||
"""Strip bot @mention placeholders from Feishu text.
|
||||
|
||||
In Feishu v2 events the text contains placeholders like ``@_user_1``
|
||||
for each mention. ``_mention_names`` caches the placeholder keys that
|
||||
belong to the bot (identified during ``_on_message``).
|
||||
"""
|
||||
result = text
|
||||
for key in self._mention_names:
|
||||
result = result.replace(key, "")
|
||||
# Clean up extra whitespace left behind
|
||||
return re.sub(r" +", " ", result).strip()
|
||||
|
||||
# ── Webhook event handler ─────────────────────────────────────
|
||||
|
||||
async def _handle_event(self, request) -> "web.Response":
|
||||
"""Handle POST /webhook/event from Feishu."""
|
||||
from aiohttp import web
|
||||
|
||||
try:
|
||||
body = await request.json()
|
||||
except Exception:
|
||||
return web.Response(status=400)
|
||||
|
||||
# ── URL verification challenge ──
|
||||
if body.get("type") == "url_verification":
|
||||
challenge = body.get("challenge", "")
|
||||
return web.json_response({"challenge": challenge})
|
||||
|
||||
# ── v2 event schema ──
|
||||
schema = body.get("schema")
|
||||
if schema == "2.0":
|
||||
header = body.get("header", {})
|
||||
|
||||
# Verify token if configured
|
||||
if self.config.verification_token:
|
||||
token = header.get("token", "")
|
||||
if token != self.config.verification_token:
|
||||
logger.warning("Feishu event token mismatch")
|
||||
return web.Response(status=403)
|
||||
|
||||
event_type = header.get("event_type", "")
|
||||
if event_type == "im.message.receive_v1":
|
||||
await self._on_message(body.get("event", {}))
|
||||
|
||||
# ── v1 event schema (legacy) ──
|
||||
elif "event" in body:
|
||||
if self.config.verification_token:
|
||||
token = body.get("token", "")
|
||||
if token != self.config.verification_token:
|
||||
logger.warning("Feishu event token mismatch (v1)")
|
||||
return web.Response(status=403)
|
||||
|
||||
event = body["event"]
|
||||
msg_type = event.get("type", "")
|
||||
if msg_type == "message":
|
||||
await self._on_message_v1(event)
|
||||
|
||||
return web.Response(status=200)
|
||||
|
||||
async def _on_message(self, event: dict) -> None:
|
||||
"""Handle im.message.receive_v1 event (v2 schema)."""
|
||||
sender_info = event.get("sender", {})
|
||||
sender_id_info = sender_info.get("sender_id", {})
|
||||
sender_id = (
|
||||
sender_id_info.get("open_id")
|
||||
or sender_id_info.get("user_id")
|
||||
or ""
|
||||
)
|
||||
sender_type = sender_info.get("sender_type", "")
|
||||
|
||||
# Skip bot's own messages
|
||||
if sender_type == "app":
|
||||
return
|
||||
|
||||
message = event.get("message", {})
|
||||
chat_id = message.get("chat_id", "")
|
||||
msg_type = message.get("message_type", "")
|
||||
message_id = message.get("message_id", "")
|
||||
|
||||
# In group chats, detect mention status for centralized gating
|
||||
chat_type = message.get("chat_type", "")
|
||||
is_group = chat_type == "group"
|
||||
was_mentioned = True
|
||||
if is_group:
|
||||
mentions = message.get("mentions", [])
|
||||
was_mentioned = bool(mentions)
|
||||
# Cache bot mention keys — bot mentions have empty user IDs
|
||||
bot_keys = []
|
||||
for m in mentions:
|
||||
m_id = m.get("id", {})
|
||||
# Bot/app mentions have no open_id / user_id
|
||||
if not m_id.get("open_id") and not m_id.get("user_id"):
|
||||
key = m.get("key", "")
|
||||
if key:
|
||||
bot_keys.append(key)
|
||||
if bot_keys:
|
||||
self._mention_names = bot_keys
|
||||
|
||||
# Parse content JSON
|
||||
content_str = message.get("content", "{}")
|
||||
try:
|
||||
content_data = json.loads(content_str)
|
||||
except json.JSONDecodeError:
|
||||
content_data = {}
|
||||
|
||||
text = ""
|
||||
annotations: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
if msg_type == "text":
|
||||
text = content_data.get("text", "")
|
||||
elif msg_type == "post":
|
||||
text = self._extract_post_text(content_data)
|
||||
elif msg_type == "image" and self.config.include_attachments:
|
||||
image_key = content_data.get("image_key", "")
|
||||
if image_key:
|
||||
local = await self._download_media(message_id, image_key, "image")
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
annotations.append(f"[attachment: {local}]")
|
||||
else:
|
||||
annotations.append("[image message - download failed]")
|
||||
else:
|
||||
annotations.append("[image message]")
|
||||
elif msg_type == "file" and self.config.include_attachments:
|
||||
file_key = content_data.get("file_key", "")
|
||||
file_name = content_data.get("file_name", "unknown")
|
||||
if file_key:
|
||||
local = await self._download_media(message_id, file_key, "file")
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
annotations.append(f"[attachment: {local}]")
|
||||
else:
|
||||
annotations.append(f"[file: {file_name} - download failed]")
|
||||
else:
|
||||
annotations.append(f"[file message: {file_name}]")
|
||||
elif msg_type in ("audio", "media") and self.config.include_attachments:
|
||||
# Feishu audio messages are voice recordings
|
||||
media_label = "voice" if msg_type == "audio" else msg_type
|
||||
file_key = content_data.get("file_key", "")
|
||||
if file_key:
|
||||
local = await self._download_media(message_id, file_key, "file")
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
annotations.append(f"[{media_label}: {local}]")
|
||||
else:
|
||||
annotations.append(f"[{media_label} message - download failed]")
|
||||
else:
|
||||
annotations.append(f"[{media_label} message]")
|
||||
elif msg_type == "sticker":
|
||||
sticker_key = content_data.get("file_key", "")
|
||||
if sticker_key and self.config.include_attachments:
|
||||
local = await self._download_media(message_id, sticker_key, "image")
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
annotations.append(f"[sticker: {local}]")
|
||||
else:
|
||||
annotations.append("[sticker message]")
|
||||
else:
|
||||
annotations.append("[sticker message]")
|
||||
else:
|
||||
text = f"[{msg_type} message]"
|
||||
|
||||
if not text and not media_paths and not annotations:
|
||||
return
|
||||
|
||||
# Parse timestamp (milliseconds)
|
||||
create_time = message.get("create_time", "")
|
||||
try:
|
||||
timestamp = datetime.fromtimestamp(
|
||||
int(create_time) / 1000
|
||||
) if create_time else datetime.now()
|
||||
except (ValueError, TypeError, OSError):
|
||||
timestamp = datetime.now()
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=message_id,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"chat_type": message.get("chat_type", ""),
|
||||
},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
))
|
||||
|
||||
async def _on_message_v1(self, event: dict) -> None:
|
||||
"""Handle v1 schema message event (legacy)."""
|
||||
sender_id = event.get("open_id", "")
|
||||
if not sender_id:
|
||||
return
|
||||
|
||||
# Detect group and mention status for centralized gating
|
||||
chat_type = event.get("chat_type", "")
|
||||
is_group = chat_type == "group"
|
||||
was_mentioned = True
|
||||
if is_group:
|
||||
text_without_at = event.get("text_without_at_bot", "")
|
||||
was_mentioned = bool(text_without_at)
|
||||
|
||||
text = event.get("text_without_at_bot", "") or event.get("text", "")
|
||||
if not text:
|
||||
return
|
||||
|
||||
chat_id = event.get("open_chat_id", "")
|
||||
message_id = event.get("open_message_id", "")
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
timestamp=datetime.now(),
|
||||
message_id=message_id,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"chat_type": event.get("chat_type", ""),
|
||||
},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
))
|
||||
|
||||
@staticmethod
|
||||
def _extract_post_text(content: dict) -> str:
|
||||
"""Extract plain text from Feishu post (rich text) content."""
|
||||
parts: list[str] = []
|
||||
# Post content has locale keys like "zh_cn", "en_us"
|
||||
for locale_key in ("zh_cn", "en_us", "ja_jp"):
|
||||
locale_content = content.get(locale_key)
|
||||
if locale_content:
|
||||
title = locale_content.get("title", "")
|
||||
if title:
|
||||
parts.append(title)
|
||||
for paragraph in locale_content.get("content", []):
|
||||
line_parts: list[str] = []
|
||||
for element in paragraph:
|
||||
tag = element.get("tag", "")
|
||||
if tag == "text":
|
||||
line_parts.append(element.get("text", ""))
|
||||
elif tag == "a":
|
||||
line_parts.append(element.get("text", ""))
|
||||
elif tag == "at":
|
||||
# Skip @mentions of the bot
|
||||
pass
|
||||
line = "".join(line_parts).strip()
|
||||
if line:
|
||||
parts.append(line)
|
||||
break # Use first available locale
|
||||
return "\n".join(parts)
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Feishu (飞书/Lark) app credential validation."""
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_feishu_credentials(
|
||||
app_id: str,
|
||||
app_secret: str,
|
||||
domain: str = "https://open.feishu.cn",
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate Feishu app credentials by requesting a tenant_access_token.
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, message).
|
||||
"""
|
||||
if not app_id:
|
||||
return False, "No app_id provided"
|
||||
if not app_secret:
|
||||
return False, "No app_secret provided"
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return False, "httpx not installed"
|
||||
|
||||
url = f"{domain}/open-apis/auth/v3/tenant_access_token/internal"
|
||||
body = {"app_id": app_id, "app_secret": app_secret}
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
resp = await client.post(url, json=body, timeout=10)
|
||||
data = resp.json()
|
||||
if data.get("code") == 0:
|
||||
return True, f"App: {app_id}"
|
||||
msg = data.get("msg", "unknown error")
|
||||
return False, f"Auth failed: {msg}"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
@@ -0,0 +1,113 @@
|
||||
"""Feishu (飞书/Lark) channel server.
|
||||
|
||||
Standalone script to run the Feishu channel with CLI options.
|
||||
|
||||
Usage:
|
||||
python -m EvoScientist.channels.feishu.serve --app-id ID --app-secret SECRET [OPTIONS]
|
||||
|
||||
Examples:
|
||||
# Basic setup
|
||||
python -m EvoScientist.channels.feishu.serve --app-id ID --app-secret SECRET
|
||||
|
||||
# With verification token and custom port
|
||||
python -m EvoScientist.channels.feishu.serve --app-id ID --app-secret SECRET \\
|
||||
--verification-token TOKEN --webhook-port 9000
|
||||
|
||||
# With agent and thinking
|
||||
python -m EvoScientist.channels.feishu.serve --app-id ID --app-secret SECRET --agent --thinking
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import FeishuChannel, FeishuConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Feishu (飞书/Lark) channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--app-id",
|
||||
required=True,
|
||||
help="Feishu App ID",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--app-secret",
|
||||
required=True,
|
||||
help="Feishu App Secret",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--verification-token",
|
||||
default="",
|
||||
help="Feishu event verification token",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--encrypt-key",
|
||||
default="",
|
||||
help="Feishu event encrypt key",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--webhook-port",
|
||||
type=int,
|
||||
default=9000,
|
||||
help="Port for webhook HTTP server (default: 9000)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--domain",
|
||||
default="https://open.feishu.cn",
|
||||
help="Feishu API domain (use https://open.larksuite.com for Lark)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender (Feishu open_id). Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler (default: echo)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content as intermediate messages (requires --agent)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = FeishuConfig(
|
||||
app_id=args.app_id,
|
||||
app_secret=args.app_secret,
|
||||
verification_token=args.verification_token,
|
||||
encrypt_key=args.encrypt_key,
|
||||
webhook_port=args.webhook_port,
|
||||
feishu_domain=args.domain,
|
||||
allowed_senders=set(args.allowed_senders) if args.allowed_senders else None,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = FeishuChannel(config)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,287 @@
|
||||
"""Unified formatting pipeline for all channels.
|
||||
|
||||
Internal representation is Markdown. This module converts Markdown to
|
||||
each platform's native format: HTML, Slack mrkdwn, Discord Markdown,
|
||||
or plain text.
|
||||
|
||||
Channels no longer need per-file format functions — they just declare
|
||||
``capabilities.format_type`` and the base class auto-configures a
|
||||
``UnifiedFormatter`` instance.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Callable
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Markdown conversion engine (formerly markdown_utils.py)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
_PLACEHOLDER_PREFIX = "\x00BLOCK"
|
||||
_INLINE_PREFIX = "\x00INLINE"
|
||||
|
||||
# A formatting rule: (regex_pattern, replacement)
|
||||
InlineRule = tuple[str, str]
|
||||
|
||||
|
||||
def convert_markdown(
|
||||
text: str,
|
||||
*,
|
||||
code_block_formatter: Callable[[str, str], str],
|
||||
inline_code_formatter: Callable[[str], str],
|
||||
inline_rules: list[InlineRule],
|
||||
escape_fn: Callable[[str], str] | None = None,
|
||||
) -> str:
|
||||
"""Convert Markdown to a channel-specific format.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
text:
|
||||
Input Markdown text.
|
||||
code_block_formatter:
|
||||
``(language, code) -> str`` — format a fenced code block.
|
||||
inline_code_formatter:
|
||||
``(code) -> str`` — format an inline code span.
|
||||
inline_rules:
|
||||
List of ``(pattern, replacement)`` pairs applied in order to the
|
||||
remaining text (after code extraction and optional escaping).
|
||||
escape_fn:
|
||||
Optional function applied to the non-code text *before* inline
|
||||
rules. Useful for HTML-escaping (Telegram) or other channel-
|
||||
specific character escaping.
|
||||
|
||||
Returns
|
||||
-------
|
||||
str
|
||||
The converted text.
|
||||
"""
|
||||
# 1. Extract and protect fenced code blocks (```...```)
|
||||
code_blocks: list[str] = []
|
||||
|
||||
def _save_code_block(m: re.Match) -> str:
|
||||
lang = m.group(1) or ""
|
||||
code = m.group(2)
|
||||
formatted = code_block_formatter(lang, code)
|
||||
idx = len(code_blocks)
|
||||
code_blocks.append(formatted)
|
||||
return f"{_PLACEHOLDER_PREFIX}{idx}\x00"
|
||||
|
||||
text = re.sub(r"```(\w*)\n?(.*?)```", _save_code_block, text, flags=re.DOTALL)
|
||||
|
||||
# 2. Extract and protect inline code (`...`)
|
||||
inline_codes: list[str] = []
|
||||
|
||||
def _save_inline(m: re.Match) -> str:
|
||||
code = m.group(1)
|
||||
formatted = inline_code_formatter(code)
|
||||
idx = len(inline_codes)
|
||||
inline_codes.append(formatted)
|
||||
return f"{_INLINE_PREFIX}{idx}\x00"
|
||||
|
||||
text = re.sub(r"`([^`]+)`", _save_inline, text)
|
||||
|
||||
# 3. Optional escaping of remaining text
|
||||
if escape_fn is not None:
|
||||
text = escape_fn(text)
|
||||
|
||||
# 4. Apply inline formatting rules
|
||||
for pattern, replacement in inline_rules:
|
||||
text = re.sub(pattern, replacement, text, flags=re.MULTILINE)
|
||||
|
||||
# 5. Restore code blocks and inline code
|
||||
for idx, html in enumerate(code_blocks):
|
||||
text = text.replace(f"{_PLACEHOLDER_PREFIX}{idx}\x00", html)
|
||||
for idx, code in enumerate(inline_codes):
|
||||
text = text.replace(f"{_INLINE_PREFIX}{idx}\x00", code)
|
||||
|
||||
return text
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Shared helpers
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _escape_html(text: str) -> str:
|
||||
return text.replace("&", "&").replace("<", "<").replace(">", ">")
|
||||
|
||||
|
||||
def _noop_escape(text: str) -> str:
|
||||
return text
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# HTML profile (Telegram, Email, Teams)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _html_code_block(lang: str, code: str) -> str:
|
||||
escaped = _escape_html(code)
|
||||
if lang:
|
||||
return f'<pre><code class="language-{lang}">{escaped}</code></pre>'
|
||||
return f"<pre><code>{escaped}</code></pre>"
|
||||
|
||||
|
||||
def _html_inline_code(code: str) -> str:
|
||||
return f"<code>{_escape_html(code)}</code>"
|
||||
|
||||
|
||||
_HTML_INLINE_RULES: list[InlineRule] = [
|
||||
# Headings → bold
|
||||
(r"^#{1,6}\s+(.+)$", r"<b>\1</b>"),
|
||||
# Blockquote markers (already escaped to >)
|
||||
(r"^>\s?", ""),
|
||||
# Links [text](url) → <a>
|
||||
(r"\[([^\]]+)\]\(([^)]+)\)", r'<a href="\2">\1</a>'),
|
||||
# Bold **text** → <b>
|
||||
(r"\*\*(.+?)\*\*", r"<b>\1</b>"),
|
||||
# Italic _text_ → <i>
|
||||
(r"(?<!\w)_([^_]+?)_(?!\w)", r"<i>\1</i>"),
|
||||
# Strikethrough ~~text~~ → <s>
|
||||
(r"~~(.+?)~~", r"<s>\1</s>"),
|
||||
# List items
|
||||
(r"^[\-\*]\s+", "• "),
|
||||
]
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Slack mrkdwn profile
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _slack_code_block(lang: str, code: str) -> str:
|
||||
return f"```\n{code}```"
|
||||
|
||||
|
||||
def _slack_inline_code(code: str) -> str:
|
||||
return f"`{code}`"
|
||||
|
||||
|
||||
_SLACK_INLINE_RULES: list[InlineRule] = [
|
||||
(r"^#{1,6}\s+(.+)$", r"*\1*"),
|
||||
(r"\[([^\]]+)\]\(([^)]+)\)", r"<\2|\1>"),
|
||||
(r"\*\*(.+?)\*\*", r"*\1*"),
|
||||
(r"~~(.+?)~~", r"~\1~"),
|
||||
(r"^[\-\*]\s+", "• "),
|
||||
]
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Discord profile (mostly passthrough, headings → bold)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _discord_code_block(lang: str, code: str) -> str:
|
||||
return f"```{lang}\n{code}```"
|
||||
|
||||
|
||||
def _discord_inline_code(code: str) -> str:
|
||||
return f"`{code}`"
|
||||
|
||||
|
||||
_DISCORD_INLINE_RULES: list[InlineRule] = [
|
||||
(r"^#{1,6}\s+(.+)$", r"**\1**"),
|
||||
]
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Plain text profile (strip all formatting)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _plain_code_block(lang: str, code: str) -> str:
|
||||
return code
|
||||
|
||||
|
||||
def _plain_inline_code(code: str) -> str:
|
||||
return code
|
||||
|
||||
|
||||
_PLAIN_INLINE_RULES: list[InlineRule] = [
|
||||
(r"^#{1,6}\s+", ""),
|
||||
(r"\[([^\]]+)\]\(([^)]+)\)", r"\1 (\2)"),
|
||||
(r"\*\*(.+?)\*\*", r"\1"),
|
||||
(r"(?<!\w)_([^_]+?)_(?!\w)", r"\1"),
|
||||
(r"~~(.+?)~~", r"\1"),
|
||||
(r"^[\-\*]\s+", "• "),
|
||||
]
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Markdown passthrough profile (Feishu, DingTalk, WeCom)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _md_code_block(lang: str, code: str) -> str:
|
||||
return f"```{lang}\n{code}```"
|
||||
|
||||
|
||||
def _md_inline_code(code: str) -> str:
|
||||
return f"`{code}`"
|
||||
|
||||
|
||||
_MD_INLINE_RULES: list[InlineRule] = [] # passthrough — already Markdown
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Unified Formatter
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
class UnifiedFormatter:
|
||||
"""Converts internal Markdown to a target platform format.
|
||||
|
||||
Instantiated once per channel based on its ``capabilities.format_type``.
|
||||
"""
|
||||
|
||||
_PROFILES: dict[str, dict] = {
|
||||
"html": dict(
|
||||
code_block_formatter=_html_code_block,
|
||||
inline_code_formatter=_html_inline_code,
|
||||
inline_rules=_HTML_INLINE_RULES,
|
||||
escape_fn=_escape_html,
|
||||
),
|
||||
"slack_mrkdwn": dict(
|
||||
code_block_formatter=_slack_code_block,
|
||||
inline_code_formatter=_slack_inline_code,
|
||||
inline_rules=_SLACK_INLINE_RULES,
|
||||
escape_fn=None,
|
||||
),
|
||||
"discord": dict(
|
||||
code_block_formatter=_discord_code_block,
|
||||
inline_code_formatter=_discord_inline_code,
|
||||
inline_rules=_DISCORD_INLINE_RULES,
|
||||
escape_fn=None,
|
||||
),
|
||||
"markdown": dict(
|
||||
code_block_formatter=_md_code_block,
|
||||
inline_code_formatter=_md_inline_code,
|
||||
inline_rules=_MD_INLINE_RULES,
|
||||
escape_fn=None,
|
||||
),
|
||||
"plain": dict(
|
||||
code_block_formatter=_plain_code_block,
|
||||
inline_code_formatter=_plain_inline_code,
|
||||
inline_rules=_PLAIN_INLINE_RULES,
|
||||
escape_fn=None,
|
||||
),
|
||||
}
|
||||
|
||||
def __init__(self, format_type: str = "plain") -> None:
|
||||
self._format_type = format_type
|
||||
profile = self._PROFILES.get(format_type)
|
||||
if profile is None:
|
||||
raise ValueError(
|
||||
f"Unknown format_type: {format_type!r}. "
|
||||
f"Available: {list(self._PROFILES.keys())}"
|
||||
)
|
||||
self._profile = profile
|
||||
|
||||
@property
|
||||
def format_type(self) -> str:
|
||||
return self._format_type
|
||||
|
||||
def format(self, text: str) -> str:
|
||||
"""Convert Markdown *text* to the target format."""
|
||||
if not text:
|
||||
return text
|
||||
return convert_markdown(text, **self._profile)
|
||||
|
||||
@classmethod
|
||||
def for_channel(cls, format_type: str) -> "UnifiedFormatter":
|
||||
"""Factory: create a formatter for the given format type."""
|
||||
return cls(format_type)
|
||||
@@ -19,6 +19,7 @@ from .targets import (
|
||||
IMessageTarget,
|
||||
IMessageService,
|
||||
)
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = [
|
||||
"IMessageChannel",
|
||||
@@ -31,3 +32,11 @@ __all__ = [
|
||||
"IMessageTarget",
|
||||
"IMessageService",
|
||||
]
|
||||
|
||||
|
||||
def create_from_config(config) -> IMessageChannel:
|
||||
allowed = _parse_csv(config.imessage_allowed_senders)
|
||||
return IMessageChannel(IMessageConfig(allowed_senders=allowed))
|
||||
|
||||
|
||||
register_channel("imessage", create_from_config)
|
||||
|
||||
@@ -6,11 +6,12 @@ via JSON-RPC, similar to OpenClaw's approach.
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import AsyncIterator
|
||||
from pathlib import Path
|
||||
|
||||
from ..base import Channel, IncomingMessage, OutgoingMessage, ChannelError
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..config import BaseChannelConfig
|
||||
from .rpc_client import ImsgRpcClient, RpcNotification
|
||||
from .targets import (
|
||||
normalize_handle,
|
||||
@@ -23,15 +24,32 @@ from .targets import (
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _IMessageAllowListMiddleware:
|
||||
"""Custom allow-list middleware for iMessage's rich sender filtering.
|
||||
|
||||
Supports chat_id/chat_guid matching, wildcard, and normalized
|
||||
phone/email matching — logic that the generic AllowListMiddleware
|
||||
does not cover.
|
||||
"""
|
||||
|
||||
def __init__(self, channel: 'IMessageChannelRpc'):
|
||||
self._channel = channel
|
||||
|
||||
async def process_inbound(self, raw, context):
|
||||
chat_id = raw.metadata.get("chat_id")
|
||||
chat_guid = raw.metadata.get("chat_guid")
|
||||
if not self._channel._is_sender_allowed(raw.sender_id, chat_id, chat_guid):
|
||||
return None
|
||||
return raw
|
||||
|
||||
|
||||
@dataclass
|
||||
class IMessageConfig:
|
||||
class IMessageConfig(BaseChannelConfig):
|
||||
"""Configuration for iMessage channel."""
|
||||
|
||||
cli_path: str = "imsg"
|
||||
db_path: str | None = None
|
||||
allowed_senders: list[str] = field(default_factory=list)
|
||||
include_attachments: bool = False
|
||||
text_chunk_limit: int = 4000
|
||||
text_chunk_limit: int = 4096
|
||||
service: str = "auto" # imessage, sms, or auto
|
||||
region: str = "US"
|
||||
|
||||
@@ -46,21 +64,39 @@ class IMessageChannelRpc(Channel):
|
||||
config: Channel configuration
|
||||
"""
|
||||
|
||||
name = "imessage"
|
||||
_ready_attrs = ("_client",)
|
||||
|
||||
def __init__(self, config: IMessageConfig | None = None):
|
||||
self.config = config or IMessageConfig()
|
||||
super().__init__(config or IMessageConfig())
|
||||
self._client: ImsgRpcClient | None = None
|
||||
self._running = False
|
||||
self._message_queue: asyncio.Queue[IncomingMessage] = asyncio.Queue()
|
||||
self._subscription_id: int | None = None
|
||||
|
||||
# ── Pipeline overrides ────────────────────────────────────────
|
||||
|
||||
def _build_inbound_middlewares(self):
|
||||
"""Use iMessage-specific allow-list middleware.
|
||||
|
||||
iMessage doesn't need MentionGating (always sets was_mentioned=True).
|
||||
"""
|
||||
from ..middleware import DedupMiddleware, GroupHistoryMiddleware
|
||||
middlewares = []
|
||||
middlewares.append(DedupMiddleware())
|
||||
middlewares.append(_IMessageAllowListMiddleware(self))
|
||||
if self.capabilities.groups:
|
||||
middlewares.append(GroupHistoryMiddleware())
|
||||
return middlewares
|
||||
|
||||
# ── Incoming message handling ─────────────────────────────────
|
||||
|
||||
def _handle_notification(self, notification: RpcNotification) -> None:
|
||||
"""Handle incoming RPC notifications."""
|
||||
if notification.method == "message":
|
||||
self._handle_message(notification.params)
|
||||
asyncio.create_task(self._handle_message(notification.params))
|
||||
elif notification.method == "error":
|
||||
logger.error(f"imsg error: {notification.params}")
|
||||
|
||||
def _handle_message(self, params: dict | None) -> None:
|
||||
async def _handle_message(self, params: dict | None) -> None:
|
||||
"""Process incoming message notification."""
|
||||
if not params:
|
||||
return
|
||||
@@ -77,16 +113,7 @@ class IMessageChannelRpc(Channel):
|
||||
if not sender:
|
||||
return
|
||||
|
||||
# Check allowed senders
|
||||
chat_id = message.get("chat_id")
|
||||
chat_guid = message.get("chat_guid")
|
||||
if not self._is_sender_allowed(sender, chat_id, chat_guid):
|
||||
logger.debug(f"Ignoring message from {sender}")
|
||||
return
|
||||
|
||||
text = message.get("text", "").strip()
|
||||
if not text:
|
||||
return
|
||||
|
||||
# Parse timestamp
|
||||
timestamp = datetime.now()
|
||||
@@ -105,23 +132,61 @@ class IMessageChannelRpc(Channel):
|
||||
}
|
||||
|
||||
# Handle attachments if enabled
|
||||
annotations: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
_VOICE_EXTS = {".caf", ".m4a", ".aac", ".ogg", ".opus", ".mp3", ".amr"}
|
||||
if self.config.include_attachments:
|
||||
attachments = message.get("attachments", [])
|
||||
if attachments:
|
||||
metadata["attachments"] = attachments
|
||||
for att in attachments:
|
||||
# imsg CLI provides local file paths for attachments
|
||||
file_path = att if isinstance(att, str) else att.get("path", "")
|
||||
if not file_path:
|
||||
annotations.append("[attachment: missing path]")
|
||||
continue
|
||||
att_path = Path(file_path)
|
||||
is_voice = att_path.suffix.lower() in _VOICE_EXTS
|
||||
media_label = "voice" if is_voice else "attachment"
|
||||
if att_path.exists():
|
||||
fname = att_path.name
|
||||
# Check file size before copying
|
||||
from ..base import MAX_ATTACHMENT_BYTES
|
||||
if att_path.stat().st_size > MAX_ATTACHMENT_BYTES:
|
||||
annotations.append(
|
||||
f"[{media_label}: {fname} - too large "
|
||||
f"({att_path.stat().st_size} bytes)]"
|
||||
)
|
||||
else:
|
||||
local = self._media_path(f"imsg_{fname}")
|
||||
try:
|
||||
import shutil
|
||||
shutil.copy2(str(att_path), str(local))
|
||||
media_paths.append(str(local))
|
||||
annotations.append(f"[{media_label}: {local}]")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to copy iMessage attachment: {e}")
|
||||
annotations.append(f"[{media_label}: {fname} - copy failed]")
|
||||
else:
|
||||
annotations.append(f"[{media_label}: {file_path} - not found]")
|
||||
|
||||
incoming = IncomingMessage(
|
||||
sender=sender,
|
||||
content=text,
|
||||
if not text and not media_paths and not annotations:
|
||||
return
|
||||
|
||||
is_group = message.get("is_group", False)
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender,
|
||||
chat_id=str(metadata.get("chat_id", sender)),
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=str(message.get("id", "")),
|
||||
metadata=metadata,
|
||||
)
|
||||
is_group=is_group,
|
||||
was_mentioned=True, # iMessage has no mention concept
|
||||
))
|
||||
|
||||
try:
|
||||
self._message_queue.put_nowait(incoming)
|
||||
except asyncio.QueueFull:
|
||||
logger.warning("Message queue full, dropping message")
|
||||
# ── Sender filtering ──────────────────────────────────────────
|
||||
|
||||
def _is_sender_allowed(
|
||||
self,
|
||||
@@ -179,28 +244,35 @@ class IMessageChannelRpc(Channel):
|
||||
|
||||
return False
|
||||
|
||||
def _normalize_sender(self, sender: str) -> str:
|
||||
"""Normalize a sender identifier."""
|
||||
return sender if sender.startswith("chat") else normalize_handle(sender)
|
||||
|
||||
def add_allowed_sender(self, sender: str) -> None:
|
||||
"""Add a sender to the allowed list."""
|
||||
normalized = normalize_handle(sender) if not sender.startswith("chat") else sender
|
||||
if normalized not in self.config.allowed_senders:
|
||||
self.config.allowed_senders.append(normalized)
|
||||
logger.info(f"Added allowed sender: {normalized}")
|
||||
normalized = self._normalize_sender(sender)
|
||||
if self.config.allowed_senders is None:
|
||||
self.config.allowed_senders = set()
|
||||
self.config.allowed_senders.add(normalized)
|
||||
logger.info(f"Added allowed sender: {normalized}")
|
||||
|
||||
def remove_allowed_sender(self, sender: str) -> None:
|
||||
"""Remove a sender from the allowed list."""
|
||||
normalized = normalize_handle(sender) if not sender.startswith("chat") else sender
|
||||
if normalized in self.config.allowed_senders:
|
||||
self.config.allowed_senders.remove(normalized)
|
||||
normalized = self._normalize_sender(sender)
|
||||
if self.config.allowed_senders:
|
||||
self.config.allowed_senders.discard(normalized)
|
||||
logger.info(f"Removed allowed sender: {normalized}")
|
||||
|
||||
def clear_allowed_senders(self) -> None:
|
||||
"""Clear allowed list (allow all)."""
|
||||
self.config.allowed_senders = []
|
||||
self.config.allowed_senders = None
|
||||
logger.info("Cleared allowed senders (allowing all)")
|
||||
|
||||
def list_allowed_senders(self) -> list[str]:
|
||||
"""Get current allowed senders."""
|
||||
return self.config.allowed_senders
|
||||
return list(self.config.allowed_senders) if self.config.allowed_senders else []
|
||||
|
||||
# ── Lifecycle ─────────────────────────────────────────────────
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Initialize and start the channel."""
|
||||
@@ -231,11 +303,7 @@ class IMessageChannelRpc(Channel):
|
||||
self._running = True
|
||||
logger.info("iMessage channel started")
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the channel and clean up."""
|
||||
logger.info("Stopping iMessage channel...")
|
||||
self._running = False
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
if self._client and self._subscription_id:
|
||||
try:
|
||||
await self._client.request(
|
||||
@@ -244,149 +312,82 @@ class IMessageChannelRpc(Channel):
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if self._client:
|
||||
await self._client.stop()
|
||||
self._client = None
|
||||
|
||||
logger.info("iMessage channel stopped")
|
||||
|
||||
async def receive(self) -> AsyncIterator[IncomingMessage]:
|
||||
"""Yield incoming messages from the queue."""
|
||||
while self._running:
|
||||
# ── Send (template method overrides) ──────────────────────────
|
||||
|
||||
def _resolve_target(self, chat_id: str | None, metadata: dict | None) -> dict:
|
||||
"""Resolve send target from metadata or chat_id string."""
|
||||
meta = metadata or {}
|
||||
for key in ("chat_id", "chat_guid", "chat_identifier"):
|
||||
if meta.get(key):
|
||||
return {key: meta[key]}
|
||||
if chat_id:
|
||||
try:
|
||||
msg = await asyncio.wait_for(
|
||||
self._message_queue.get(),
|
||||
timeout=1.0,
|
||||
)
|
||||
yield msg
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
|
||||
def _segment_message(self, content: str) -> list[str]:
|
||||
"""Split long message into segments."""
|
||||
limit = self.config.text_chunk_limit
|
||||
if len(content) <= limit:
|
||||
return [content]
|
||||
|
||||
segments = []
|
||||
remaining = content
|
||||
|
||||
while remaining:
|
||||
if len(remaining) <= limit:
|
||||
segments.append(remaining)
|
||||
break
|
||||
|
||||
chunk = remaining[:limit]
|
||||
# Try split at newline
|
||||
nl_pos = chunk.rfind("\n")
|
||||
if nl_pos > limit // 2:
|
||||
split_pos = nl_pos + 1
|
||||
else:
|
||||
# Try split at space
|
||||
sp_pos = chunk.rfind(" ")
|
||||
if sp_pos > limit // 2:
|
||||
split_pos = sp_pos + 1
|
||||
target = parse_target(chat_id)
|
||||
if isinstance(target, ChatIdTarget):
|
||||
return {"chat_id": target.chat_id}
|
||||
elif isinstance(target, ChatGuidTarget):
|
||||
return {"chat_guid": target.chat_guid}
|
||||
elif isinstance(target, ChatIdentifierTarget):
|
||||
return {"chat_identifier": target.chat_identifier}
|
||||
else:
|
||||
split_pos = limit
|
||||
return {"to": target.to, "service": target.service.value}
|
||||
except ValueError:
|
||||
return {"to": chat_id}
|
||||
return {}
|
||||
|
||||
segments.append(remaining[:split_pos].rstrip())
|
||||
remaining = remaining[split_pos:].lstrip()
|
||||
|
||||
return segments
|
||||
|
||||
async def send(self, message: OutgoingMessage) -> bool:
|
||||
"""Send a message via iMessage."""
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
"""Send a single text chunk via iMessage RPC."""
|
||||
if not self._client:
|
||||
logger.error("Cannot send: client not running")
|
||||
return False
|
||||
raise RuntimeError("iMessage client not running")
|
||||
|
||||
segments = self._segment_message(message.content)
|
||||
|
||||
for segment in segments:
|
||||
params = self._build_send_params(message, segment)
|
||||
if not params:
|
||||
logger.error(f"_build_send_params returned None for recipient={message.recipient}, metadata={message.metadata}")
|
||||
return False
|
||||
|
||||
try:
|
||||
logger.debug(f"Calling imsg send with params: {params}")
|
||||
await self._client.request("send", params)
|
||||
except Exception as e:
|
||||
logger.error(f"Send failed: {e}")
|
||||
logger.error(f"Failed params were: {params}")
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def _build_send_params(
|
||||
self, message: OutgoingMessage, text: str
|
||||
) -> dict | None:
|
||||
"""Build send parameters from message."""
|
||||
params: dict = {
|
||||
"text": text,
|
||||
"text": formatted_text,
|
||||
"service": self.config.service,
|
||||
"region": self.config.region,
|
||||
}
|
||||
params.update(self._resolve_target(chat_id, metadata))
|
||||
|
||||
logger.debug(f"Building send params - recipient: {message.recipient}, metadata: {message.metadata}")
|
||||
if reply_to:
|
||||
params["reply_to"] = reply_to
|
||||
|
||||
# Check metadata for chat targets
|
||||
chat_id = message.metadata.get("chat_id")
|
||||
chat_guid = message.metadata.get("chat_guid")
|
||||
chat_identifier = message.metadata.get("chat_identifier")
|
||||
await self._client.request("send", params)
|
||||
|
||||
if chat_id:
|
||||
params["chat_id"] = chat_id
|
||||
elif chat_guid:
|
||||
params["chat_guid"] = chat_guid
|
||||
elif chat_identifier:
|
||||
params["chat_identifier"] = chat_identifier
|
||||
elif message.recipient:
|
||||
# Parse recipient to determine target type
|
||||
try:
|
||||
target = parse_target(message.recipient)
|
||||
if isinstance(target, ChatIdTarget):
|
||||
params["chat_id"] = target.chat_id
|
||||
elif isinstance(target, ChatGuidTarget):
|
||||
params["chat_guid"] = target.chat_guid
|
||||
elif isinstance(target, ChatIdentifierTarget):
|
||||
params["chat_identifier"] = target.chat_identifier
|
||||
else:
|
||||
params["to"] = target.to
|
||||
params["service"] = target.service.value
|
||||
except ValueError:
|
||||
params["to"] = message.recipient
|
||||
else:
|
||||
logger.error("Cannot send: no recipient or chat target")
|
||||
return None
|
||||
# ── Retry logic (override base) ───────────────────────────────
|
||||
|
||||
logger.debug(f"Built send params: {params}")
|
||||
return params
|
||||
def _format_chunk(self, text: str) -> str:
|
||||
"""iMessage uses plain text; no formatting conversion needed."""
|
||||
return text
|
||||
|
||||
async def send_media(
|
||||
|
||||
def _extract_retry_after(self, exc: Exception) -> float | None:
|
||||
"""iMessage-specific retry logic.
|
||||
|
||||
RPC errors (e.g. AppleScript failures) are generally not
|
||||
retryable. Transient connection issues get a short retry.
|
||||
"""
|
||||
msg = str(exc).lower()
|
||||
if "not found" in msg or "applescript" in msg or "permission" in msg:
|
||||
return None # not retryable
|
||||
if "timeout" in msg or "connection" in msg:
|
||||
return 1.0
|
||||
return None # default: don't retry RPC errors
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a media file via iMessage.
|
||||
|
||||
Args:
|
||||
recipient: Target recipient or chat target
|
||||
file_path: Local path to the media file
|
||||
caption: Optional caption text
|
||||
metadata: Optional metadata with chat_id etc.
|
||||
|
||||
Returns:
|
||||
True if sent successfully
|
||||
"""
|
||||
"""Send a media file via iMessage."""
|
||||
if not self._client:
|
||||
logger.error("Cannot send media: client not running")
|
||||
return False
|
||||
|
||||
metadata = metadata or {}
|
||||
params: dict = {
|
||||
"file": file_path,
|
||||
"service": self.config.service,
|
||||
@@ -396,32 +397,11 @@ class IMessageChannelRpc(Channel):
|
||||
if caption:
|
||||
params["text"] = caption
|
||||
|
||||
# Determine target
|
||||
chat_id = metadata.get("chat_id")
|
||||
chat_guid = metadata.get("chat_guid")
|
||||
|
||||
if chat_id:
|
||||
params["chat_id"] = chat_id
|
||||
elif chat_guid:
|
||||
params["chat_guid"] = chat_guid
|
||||
elif recipient:
|
||||
try:
|
||||
target = parse_target(recipient)
|
||||
if isinstance(target, ChatIdTarget):
|
||||
params["chat_id"] = target.chat_id
|
||||
elif isinstance(target, ChatGuidTarget):
|
||||
params["chat_guid"] = target.chat_guid
|
||||
else:
|
||||
params["to"] = target.to
|
||||
except ValueError:
|
||||
params["to"] = recipient
|
||||
else:
|
||||
target = self._resolve_target(recipient, metadata)
|
||||
if not target:
|
||||
logger.error("Cannot send media: no recipient")
|
||||
return False
|
||||
params.update(target)
|
||||
|
||||
try:
|
||||
await self._client.request("send", params)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Send media failed: {e}")
|
||||
return False
|
||||
await self._client.request("send", params)
|
||||
return True
|
||||
|
||||
@@ -16,337 +16,16 @@ Examples:
|
||||
python -m EvoScientist.channels.imessage.serve --cli-path /usr/local/bin/imsg
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import argparse
|
||||
import logging
|
||||
import signal
|
||||
from typing import Callable
|
||||
|
||||
from . import IMessageChannel, IMessageConfig
|
||||
from ..base import OutgoingMessage
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _format_todo_list(todos: list[dict]) -> str:
|
||||
"""Format todo items as a numbered list."""
|
||||
lines = ["\U0001f4cb Todo List\n"] # 📋
|
||||
for i, item in enumerate(todos, 1):
|
||||
content = item.get("content", "")
|
||||
lines.append(f"{i}. {content}")
|
||||
lines.append(f"\n\U0001f680 {len(todos)} tasks") # 🚀
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def create_agent_handler(
|
||||
on_thinking: Callable | None = None,
|
||||
on_todo: Callable | None = None,
|
||||
):
|
||||
"""Create handler that uses EvoScientist agent.
|
||||
|
||||
Args:
|
||||
on_thinking: Optional async callback for thinking content.
|
||||
Signature: async def on_thinking(sender: str, thinking: str) -> None
|
||||
on_todo: Optional async callback for todo list updates.
|
||||
Signature: async def on_todo(sender: str, content: str, metadata: dict) -> None
|
||||
"""
|
||||
import os
|
||||
from langchain_core.messages import HumanMessage
|
||||
from ...config import get_effective_config, apply_config_to_env
|
||||
from ...paths import set_workspace_root, ensure_dirs
|
||||
from ...EvoScientist import create_cli_agent
|
||||
from ...stream.events import stream_agent_events
|
||||
|
||||
# Apply config so default_workdir is respected in non-CLI entry points
|
||||
config = get_effective_config()
|
||||
apply_config_to_env(config)
|
||||
if config.default_workdir:
|
||||
workdir = os.path.abspath(os.path.expanduser(config.default_workdir))
|
||||
set_workspace_root(workdir)
|
||||
ensure_dirs()
|
||||
|
||||
agent = create_cli_agent()
|
||||
sessions: dict[str, str] = {} # sender -> thread_id
|
||||
|
||||
async def handler(msg) -> str:
|
||||
import uuid
|
||||
sender = msg.sender
|
||||
if sender not in sessions:
|
||||
sessions[sender] = str(uuid.uuid4())
|
||||
thread_id = sessions[sender]
|
||||
|
||||
if on_thinking:
|
||||
final_content = ""
|
||||
thinking_buffer = []
|
||||
todo_sent = False
|
||||
thinking_sent = False
|
||||
_MIN_THINKING_LEN = 200 # Skip short thinking (simple conversations)
|
||||
|
||||
async for event in stream_agent_events(agent, msg.content, thread_id):
|
||||
event_type = event.get("type")
|
||||
|
||||
if event_type == "thinking":
|
||||
thinking_text = event.get("content", "")
|
||||
if thinking_text:
|
||||
thinking_buffer.append(thinking_text)
|
||||
|
||||
elif event_type == "tool_call":
|
||||
if event.get("name") == "write_todos" and on_todo and not todo_sent:
|
||||
todos = event.get("args", {}).get("todos", [])
|
||||
if todos:
|
||||
# Flush thinking before todo (only if long enough)
|
||||
if thinking_buffer and not thinking_sent:
|
||||
full_thinking = "".join(thinking_buffer)
|
||||
if len(full_thinking) >= _MIN_THINKING_LEN:
|
||||
await on_thinking(sender, full_thinking, msg.metadata)
|
||||
thinking_sent = True
|
||||
thinking_buffer.clear()
|
||||
await on_todo(sender, _format_todo_list(todos), msg.metadata)
|
||||
todo_sent = True
|
||||
|
||||
elif event_type == "text":
|
||||
final_content += event.get("content", "")
|
||||
|
||||
elif event_type == "done":
|
||||
final_content = event.get("content", "") or final_content
|
||||
|
||||
if thinking_buffer and not thinking_sent:
|
||||
full_thinking = "".join(thinking_buffer)
|
||||
if len(full_thinking) >= _MIN_THINKING_LEN:
|
||||
await on_thinking(sender, full_thinking, msg.metadata)
|
||||
thinking_sent = True
|
||||
|
||||
return final_content or "No response"
|
||||
else:
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
result = agent.invoke(
|
||||
{"messages": [HumanMessage(content=msg.content)]},
|
||||
config=config,
|
||||
)
|
||||
messages = result.get("messages", [])
|
||||
for m in reversed(messages):
|
||||
if hasattr(m, "content") and m.type == "ai":
|
||||
content = m.content
|
||||
# Handle structured content (thinking mode)
|
||||
if isinstance(content, list):
|
||||
text_parts = []
|
||||
for block in content:
|
||||
if isinstance(block, dict) and block.get("type") == "text":
|
||||
text_parts.append(block.get("text", ""))
|
||||
return "\n".join(text_parts) if text_parts else "No response"
|
||||
# Handle plain string content
|
||||
return content
|
||||
return "No response"
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
class IMessageServer:
|
||||
"""Server that runs the iMessage channel and handles messages."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: IMessageConfig,
|
||||
handler: Callable | None = None,
|
||||
send_thinking: bool = False,
|
||||
initial_debounce: float = 2.0,
|
||||
debounce_step: float = 0.5,
|
||||
max_debounce: float = 5.0,
|
||||
on_activity: Callable | None = None,
|
||||
):
|
||||
"""Initialize iMessage server.
|
||||
|
||||
Args:
|
||||
config: iMessage channel configuration.
|
||||
handler: Message handler function. If None, uses echo handler.
|
||||
send_thinking: If True, send thinking content as intermediate messages.
|
||||
initial_debounce: Wait time after first message (seconds).
|
||||
debounce_step: Additional wait per subsequent message.
|
||||
max_debounce: Maximum debounce window cap.
|
||||
on_activity: Optional callback(sender, direction) for notifications.
|
||||
"""
|
||||
self.config = config
|
||||
self.channel = IMessageChannel(config)
|
||||
self.send_thinking = send_thinking
|
||||
self.initial_debounce = initial_debounce
|
||||
self.debounce_step = debounce_step
|
||||
self.max_debounce = max_debounce
|
||||
self._running = False
|
||||
self._pending_thinking: dict[str, str] = {} # sender -> accumulated thinking
|
||||
self._on_activity = on_activity
|
||||
|
||||
# Message buffering for debounce
|
||||
self._message_buffers: dict[str, list[str]] = {} # sender -> [messages]
|
||||
self._message_metadata: dict[str, dict] = {} # sender -> metadata (from first message)
|
||||
self._debounce_tasks: dict[str, asyncio.Task] = {} # sender -> pending task
|
||||
self._processing: set[str] = set() # senders currently being processed
|
||||
|
||||
if handler:
|
||||
self.handler = handler
|
||||
else:
|
||||
self.handler = self._default_handler
|
||||
|
||||
async def _default_handler(self, msg) -> str:
|
||||
"""Default echo handler."""
|
||||
return f"Echo: {msg.content}"
|
||||
|
||||
async def _process_buffered_messages(self, sender: str) -> None:
|
||||
"""Process all buffered messages for a sender.
|
||||
|
||||
If the sender is currently being processed, skip — new messages
|
||||
stay in the buffer and will be picked up after current processing.
|
||||
"""
|
||||
# Don't start a new handler if one is already running for this sender
|
||||
if sender in self._processing:
|
||||
logger.debug(f"Agent busy for {sender}, messages stay queued")
|
||||
return
|
||||
|
||||
if sender not in self._message_buffers:
|
||||
return
|
||||
|
||||
messages = self._message_buffers.pop(sender, [])
|
||||
metadata = self._message_metadata.pop(sender, None)
|
||||
self._debounce_tasks.pop(sender, None)
|
||||
|
||||
if not messages:
|
||||
return
|
||||
|
||||
merged_content = "\n".join(messages)
|
||||
logger.info(f"Processing {len(messages)} merged message(s) from {sender}")
|
||||
|
||||
self._processing.add(sender)
|
||||
try:
|
||||
class MergedMessage:
|
||||
def __init__(self, s, c, m):
|
||||
self.sender = s
|
||||
self.content = c
|
||||
self.metadata = m
|
||||
|
||||
merged_msg = MergedMessage(sender, merged_content, metadata)
|
||||
response = await self.handler(merged_msg)
|
||||
|
||||
if response:
|
||||
await self.channel.send(OutgoingMessage(
|
||||
recipient=sender,
|
||||
content=response,
|
||||
metadata=metadata or {},
|
||||
))
|
||||
if self._on_activity:
|
||||
try:
|
||||
self._on_activity(sender, "replied")
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.error(f"Handler error: {e}")
|
||||
finally:
|
||||
self._processing.discard(sender)
|
||||
|
||||
# If new messages arrived during processing, restart debounce
|
||||
if sender in self._message_buffers and self._message_buffers[sender]:
|
||||
msg_count = len(self._message_buffers[sender])
|
||||
wait = min(
|
||||
self.initial_debounce + (msg_count - 1) * self.debounce_step,
|
||||
self.max_debounce,
|
||||
)
|
||||
logger.info(f"New messages queued for {sender}, restarting debounce ({wait:.1f}s)")
|
||||
|
||||
async def restart_debounce(_s=sender, _w=wait):
|
||||
await asyncio.sleep(_w)
|
||||
await self._process_buffered_messages(_s)
|
||||
|
||||
self._debounce_tasks[sender] = asyncio.create_task(restart_debounce())
|
||||
|
||||
async def _queue_message(self, msg) -> None:
|
||||
"""Queue a message with progressive debounce.
|
||||
|
||||
If agent is busy, just buffer — messages will be picked up
|
||||
after current processing finishes. Otherwise, start debounce:
|
||||
1st: 2.0s, 2nd: 2.5s, 3rd: 3.0s, ... up to max_debounce.
|
||||
"""
|
||||
sender = msg.sender
|
||||
|
||||
if sender not in self._message_buffers:
|
||||
self._message_buffers[sender] = []
|
||||
self._message_metadata[sender] = msg.metadata
|
||||
self._message_buffers[sender].append(msg.content)
|
||||
|
||||
if self._on_activity:
|
||||
try:
|
||||
self._on_activity(sender, "received")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Agent is busy — just buffer, no debounce needed
|
||||
if sender in self._processing:
|
||||
logger.debug(f"Agent busy for {sender}, buffering message #{len(self._message_buffers[sender])}")
|
||||
return
|
||||
|
||||
if sender in self._debounce_tasks:
|
||||
self._debounce_tasks[sender].cancel()
|
||||
|
||||
msg_count = len(self._message_buffers[sender])
|
||||
wait = min(
|
||||
self.initial_debounce + (msg_count - 1) * self.debounce_step,
|
||||
self.max_debounce,
|
||||
)
|
||||
logger.debug(f"Debounce for {sender}: {wait:.1f}s (message #{msg_count})")
|
||||
|
||||
async def debounce_callback(_s=sender, _w=wait):
|
||||
await asyncio.sleep(_w)
|
||||
await self._process_buffered_messages(_s)
|
||||
|
||||
self._debounce_tasks[sender] = asyncio.create_task(debounce_callback())
|
||||
|
||||
async def send_todo_message(self, sender: str, content: str, metadata: dict | None = None) -> None:
|
||||
"""Send todo list as intermediate message."""
|
||||
logger.debug(f"Sending todo list to {sender}")
|
||||
await self.channel.send(OutgoingMessage(
|
||||
recipient=sender,
|
||||
content=content,
|
||||
metadata=metadata or {},
|
||||
))
|
||||
|
||||
async def send_thinking_message(self, sender: str, thinking: str, metadata: dict | None = None) -> None:
|
||||
"""Send thinking content as intermediate message."""
|
||||
if not self.send_thinking:
|
||||
return
|
||||
|
||||
logger.debug(f"Sending thinking to {sender} with metadata: {metadata}")
|
||||
content = f"\U0001f9e0\n{thinking}\n\u23f3"
|
||||
await self.channel.send(OutgoingMessage(
|
||||
recipient=sender,
|
||||
content=content,
|
||||
metadata=metadata or {},
|
||||
))
|
||||
logger.debug(f"Sent thinking to {sender}: {thinking[:50]}...")
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Run the server."""
|
||||
await self.channel.start()
|
||||
self._running = True
|
||||
|
||||
logger.info("iMessage server running. Press Ctrl+C to stop.")
|
||||
if self.config.allowed_senders:
|
||||
logger.info(f"Allowed senders: {self.config.allowed_senders}")
|
||||
else:
|
||||
logger.info("Allowing all senders")
|
||||
logger.info(f"Debounce: {self.initial_debounce}s + {self.debounce_step}s/msg (max {self.max_debounce}s)")
|
||||
|
||||
try:
|
||||
async for msg in self.channel.receive():
|
||||
logger.info(f"From {msg.sender}: {msg.content[:50]}...")
|
||||
await self._queue_message(msg)
|
||||
finally:
|
||||
for task in self._debounce_tasks.values():
|
||||
task.cancel()
|
||||
await self.channel.stop()
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the server."""
|
||||
self._running = False
|
||||
await self.channel.stop()
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
@@ -386,8 +65,8 @@ def parse_args():
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
async def async_main():
|
||||
"""Async entry point."""
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = IMessageConfig(
|
||||
@@ -397,42 +76,11 @@ async def async_main():
|
||||
include_attachments=args.attachments,
|
||||
)
|
||||
|
||||
handler = None
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = IMessageChannel(config)
|
||||
|
||||
if args.agent:
|
||||
logger.info("Loading EvoScientist agent...")
|
||||
logger.info("Agent loaded")
|
||||
|
||||
server = IMessageServer(
|
||||
config,
|
||||
handler=None,
|
||||
send_thinking=send_thinking,
|
||||
)
|
||||
|
||||
if args.agent:
|
||||
on_thinking = server.send_thinking_message if send_thinking else None
|
||||
on_todo = server.send_todo_message
|
||||
handler = create_agent_handler(on_thinking=on_thinking, on_todo=on_todo)
|
||||
server.handler = handler
|
||||
if send_thinking:
|
||||
logger.info("Thinking messages enabled")
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
for sig in (signal.SIGINT, signal.SIGTERM):
|
||||
loop.add_signal_handler(sig, lambda: asyncio.create_task(server.stop()))
|
||||
|
||||
await server.run()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
asyncio.run(async_main())
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -0,0 +1,814 @@
|
||||
"""Composable message processing middleware.
|
||||
|
||||
Each middleware is a standalone class that can be composed into a pipeline.
|
||||
They extract logic that was previously baked into the Channel base class,
|
||||
making it reusable across both legacy and plugin-based channels.
|
||||
|
||||
Also contains the supporting data structures (DedupCache, GroupHistoryBuffer,
|
||||
TypingManager, PairingManager) that were previously in separate files.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import dataclasses
|
||||
import logging
|
||||
import random
|
||||
import time
|
||||
from collections import OrderedDict, deque
|
||||
from collections.abc import Awaitable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable
|
||||
|
||||
from .bus.events import InboundMessage, OutboundMessage
|
||||
from .base import RawIncoming
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════
|
||||
# Supporting data structures
|
||||
# ═══════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
# ── Dedup cache ──────────────────────────────────────────────────────
|
||||
|
||||
_DEDUP_MAX = 1000
|
||||
_DEDUP_TRIM = 500
|
||||
_DEDUP_TTL = 3600 # 1 hour
|
||||
|
||||
|
||||
class DedupCache:
|
||||
"""Bounded ordered cache with TTL for detecting duplicate message IDs.
|
||||
|
||||
Entries expire after *ttl_seconds* and are pruned lazily on each
|
||||
lookup. When the cache exceeds *max_size* entries it is trimmed
|
||||
down to *trim_to* by evicting the oldest entries. Accessed entries
|
||||
are moved to the end (LRU behaviour).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_size: int = _DEDUP_MAX,
|
||||
trim_to: int = _DEDUP_TRIM,
|
||||
ttl_seconds: float = _DEDUP_TTL,
|
||||
) -> None:
|
||||
self._seen: OrderedDict[str, float] = OrderedDict()
|
||||
self._max = max_size
|
||||
self._trim = trim_to
|
||||
self._ttl = ttl_seconds
|
||||
|
||||
# ── public API ──────────────────────────────────────────────────
|
||||
|
||||
def is_duplicate(self, msg_id: str) -> bool:
|
||||
"""Return ``True`` if *msg_id* has been seen before.
|
||||
|
||||
First-time IDs are recorded and ``False`` is returned.
|
||||
Empty / falsy IDs are never considered duplicates.
|
||||
Expired entries are pruned before the check.
|
||||
"""
|
||||
if not msg_id:
|
||||
return False
|
||||
|
||||
self._prune()
|
||||
|
||||
if msg_id in self._seen:
|
||||
# LRU: refresh position and timestamp
|
||||
self._seen.move_to_end(msg_id)
|
||||
self._seen[msg_id] = time.monotonic()
|
||||
return True
|
||||
|
||||
self._seen[msg_id] = time.monotonic()
|
||||
if len(self._seen) > self._max:
|
||||
while len(self._seen) > self._trim:
|
||||
self._seen.popitem(last=False)
|
||||
return False
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Remove all entries."""
|
||||
self._seen.clear()
|
||||
|
||||
@property
|
||||
def size(self) -> int:
|
||||
"""Number of entries currently in the cache."""
|
||||
return len(self._seen)
|
||||
|
||||
# ── internal ────────────────────────────────────────────────────
|
||||
|
||||
def _prune(self) -> None:
|
||||
"""Remove entries older than *ttl_seconds*."""
|
||||
cutoff = time.monotonic() - self._ttl
|
||||
# OrderedDict is insertion-ordered; oldest entries are first.
|
||||
while self._seen:
|
||||
key, ts = next(iter(self._seen.items()))
|
||||
if ts > cutoff:
|
||||
break
|
||||
self._seen.popitem(last=False)
|
||||
|
||||
|
||||
# ── Group history buffer ─────────────────────────────────────────────
|
||||
|
||||
@dataclass
|
||||
class HistoryEntry:
|
||||
sender_id: str
|
||||
text: str
|
||||
timestamp: float
|
||||
message_id: str = ""
|
||||
|
||||
|
||||
class GroupHistoryBuffer:
|
||||
"""Per-chat circular buffer of recent messages."""
|
||||
|
||||
def __init__(self, max_per_chat: int = 50, max_age_seconds: int = 3600):
|
||||
self._buffers: dict[str, deque[HistoryEntry]] = {}
|
||||
self._max = max_per_chat
|
||||
self._max_age = max_age_seconds
|
||||
|
||||
def add(self, chat_id: str, entry: HistoryEntry) -> None:
|
||||
"""Add a message to the chat's history buffer."""
|
||||
if chat_id not in self._buffers:
|
||||
self._buffers[chat_id] = deque(maxlen=self._max)
|
||||
self._buffers[chat_id].append(entry)
|
||||
|
||||
def get_recent(self, chat_id: str, limit: int = 20) -> list[HistoryEntry]:
|
||||
"""Get recent messages for context injection, excluding expired ones."""
|
||||
buf = self._buffers.get(chat_id)
|
||||
if not buf:
|
||||
return []
|
||||
now = time.time()
|
||||
recent = [e for e in buf if now - e.timestamp < self._max_age]
|
||||
return recent[-limit:]
|
||||
|
||||
def format_context(self, chat_id: str, limit: int = 20) -> str:
|
||||
"""Format recent messages as context block for the agent."""
|
||||
entries = self.get_recent(chat_id, limit)
|
||||
if not entries:
|
||||
return ""
|
||||
lines = ["[Chat messages since your last reply - for context]"]
|
||||
for e in entries:
|
||||
lines.append(f"[from: {e.sender_id}] {e.text}")
|
||||
lines.append("[/Chat context]")
|
||||
return "\n".join(lines)
|
||||
|
||||
def clear(self, chat_id: str) -> None:
|
||||
"""Clear history for a chat (e.g., after the bot replies)."""
|
||||
self._buffers.pop(chat_id, None)
|
||||
|
||||
|
||||
# ── Typing indicator manager ─────────────────────────────────────────
|
||||
|
||||
class TypingManager:
|
||||
"""Manages background typing-indicator loops per chat_id.
|
||||
|
||||
Args:
|
||||
send_action: Async callable that sends a single typing indicator
|
||||
for a given chat_id.
|
||||
interval: Seconds between typing indicator sends.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
send_action: Callable[[str], Awaitable[None]],
|
||||
interval: float = 5.0,
|
||||
) -> None:
|
||||
self._send_action = send_action
|
||||
self._interval = interval
|
||||
self._tasks: dict[str, asyncio.Task] = {}
|
||||
|
||||
async def start(self, chat_id: str) -> None:
|
||||
"""Start a background typing-indicator loop for *chat_id*."""
|
||||
await self.stop(chat_id)
|
||||
|
||||
async def _loop() -> None:
|
||||
while True:
|
||||
try:
|
||||
await self._send_action(chat_id)
|
||||
except Exception:
|
||||
pass
|
||||
await asyncio.sleep(self._interval)
|
||||
|
||||
self._tasks[chat_id] = asyncio.create_task(_loop())
|
||||
|
||||
async def stop(self, chat_id: str) -> None:
|
||||
"""Cancel the typing-indicator loop for *chat_id*."""
|
||||
task = self._tasks.pop(chat_id, None)
|
||||
if task:
|
||||
task.cancel()
|
||||
|
||||
async def stop_all(self) -> None:
|
||||
"""Cancel all active typing-indicator loops."""
|
||||
for cid in list(self._tasks):
|
||||
await self.stop(cid)
|
||||
|
||||
@property
|
||||
def active_chats(self) -> list[str]:
|
||||
"""Return chat_ids with active typing loops."""
|
||||
return list(self._tasks)
|
||||
|
||||
|
||||
# ── Pairing manager ─────────────────────────────────────────────────
|
||||
|
||||
@dataclass
|
||||
class PairingRequest:
|
||||
sender_id: str
|
||||
channel: str
|
||||
code: str
|
||||
created_at: float
|
||||
approved: bool = False
|
||||
|
||||
|
||||
class PairingManager:
|
||||
"""Manages DM pairing codes for channel access control."""
|
||||
|
||||
CODE_EXPIRY = 3600 # 1 hour
|
||||
MAX_PENDING = 50 # max pending requests
|
||||
|
||||
def __init__(self):
|
||||
self._pending: dict[str, PairingRequest] = {} # code -> request
|
||||
self._approved: set[str] = set() # "channel:sender_id" keys
|
||||
|
||||
def is_approved(self, channel: str, sender_id: str) -> bool:
|
||||
"""Check if sender is already approved."""
|
||||
return f"{channel}:{sender_id}" in self._approved
|
||||
|
||||
def request_pairing(self, channel: str, sender_id: str) -> str:
|
||||
"""Generate a pairing code for a new sender. Returns the code."""
|
||||
# Check if already has pending request
|
||||
for code, req in list(self._pending.items()):
|
||||
if req.sender_id == sender_id and req.channel == channel:
|
||||
if time.time() - req.created_at < self.CODE_EXPIRY:
|
||||
return code # return existing code
|
||||
else:
|
||||
del self._pending[code]
|
||||
break
|
||||
|
||||
# Cleanup expired
|
||||
self._cleanup_expired()
|
||||
|
||||
# Generate new code
|
||||
code = f"{random.randint(100000, 999999)}"
|
||||
while code in self._pending:
|
||||
code = f"{random.randint(100000, 999999)}"
|
||||
|
||||
self._pending[code] = PairingRequest(
|
||||
sender_id=sender_id,
|
||||
channel=channel,
|
||||
code=code,
|
||||
created_at=time.time(),
|
||||
)
|
||||
_logger.info(f"Pairing code {code} generated for {channel}:{sender_id}")
|
||||
return code
|
||||
|
||||
def approve(self, code: str) -> tuple[bool, str]:
|
||||
"""Approve a pairing code. Returns (success, message)."""
|
||||
req = self._pending.get(code)
|
||||
if not req:
|
||||
return False, f"Unknown code: {code}"
|
||||
if time.time() - req.created_at > self.CODE_EXPIRY:
|
||||
del self._pending[code]
|
||||
return False, f"Code {code} expired"
|
||||
|
||||
key = f"{req.channel}:{req.sender_id}"
|
||||
self._approved.add(key)
|
||||
del self._pending[code]
|
||||
_logger.info(f"Approved pairing for {key}")
|
||||
return True, f"Approved {req.sender_id} on {req.channel}"
|
||||
|
||||
def reject(self, code: str) -> tuple[bool, str]:
|
||||
"""Reject a pairing code."""
|
||||
if code in self._pending:
|
||||
del self._pending[code]
|
||||
return True, f"Rejected code {code}"
|
||||
return False, f"Unknown code: {code}"
|
||||
|
||||
def list_pending(self) -> list[PairingRequest]:
|
||||
"""List all pending (non-expired) requests."""
|
||||
self._cleanup_expired()
|
||||
return list(self._pending.values())
|
||||
|
||||
def _cleanup_expired(self):
|
||||
now = time.time()
|
||||
expired = [c for c, r in self._pending.items() if now - r.created_at > self.CODE_EXPIRY]
|
||||
for c in expired:
|
||||
del self._pending[c]
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════
|
||||
# Middleware classes
|
||||
# ═══════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
# ── Inbound middleware base ──────────────────────────────────────────
|
||||
|
||||
class InboundMiddleware:
|
||||
"""Base class for inbound message processing middleware."""
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
"""Process an inbound raw message.
|
||||
|
||||
Return the (possibly modified) RawIncoming to continue the
|
||||
pipeline, or ``None`` to drop the message.
|
||||
"""
|
||||
return raw
|
||||
|
||||
|
||||
class OutboundMiddlewareBase:
|
||||
"""Base class for outbound message processing middleware."""
|
||||
|
||||
async def process_outbound(
|
||||
self, message: OutboundMessage, context: dict[str, Any],
|
||||
) -> OutboundMessage | None:
|
||||
"""Process an outbound message.
|
||||
|
||||
Return the (possibly modified) OutboundMessage to continue,
|
||||
or ``None`` to drop it.
|
||||
"""
|
||||
return message
|
||||
|
||||
|
||||
# ── Dedup ────────────────────────────────────────────────────────────
|
||||
|
||||
class DedupMiddleware(InboundMiddleware):
|
||||
"""Message deduplication using a bounded TTL cache."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_size: int = 1000,
|
||||
trim_to: int = 500,
|
||||
ttl_seconds: float = 3600.0,
|
||||
) -> None:
|
||||
self._cache = DedupCache(
|
||||
max_size=max_size, trim_to=trim_to, ttl_seconds=ttl_seconds,
|
||||
)
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
if raw.message_id and self._cache.is_duplicate(raw.message_id):
|
||||
_logger.debug(f"Dedup: skipping duplicate message {raw.message_id}")
|
||||
return None
|
||||
return raw
|
||||
|
||||
|
||||
# ── Debounce ─────────────────────────────────────────────────────────
|
||||
|
||||
class DebounceMiddleware:
|
||||
"""Per-sender message batching with configurable timing.
|
||||
|
||||
This middleware collects messages from the same sender and merges
|
||||
them after a debounce delay. It does not follow the simple
|
||||
process_inbound pattern because it needs to buffer across calls.
|
||||
|
||||
Usage: call ``submit()`` for each message; merged results are
|
||||
delivered via the ``on_ready`` callback.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
initial_debounce: float = 2.0,
|
||||
debounce_step: float = 0.5,
|
||||
max_debounce: float = 5.0,
|
||||
on_ready: Callable[[InboundMessage], Any] | None = None,
|
||||
) -> None:
|
||||
self.initial_debounce = initial_debounce
|
||||
self.debounce_step = debounce_step
|
||||
self.max_debounce = max_debounce
|
||||
self.on_ready = on_ready
|
||||
|
||||
self._buffers: dict[str, list[str]] = {}
|
||||
self._metadata: dict[str, dict] = {}
|
||||
self._media: dict[str, list[str]] = {}
|
||||
self._message_ids: dict[str, str] = {}
|
||||
self._tasks: dict[str, asyncio.Task] = {}
|
||||
self._channel_name: str = ""
|
||||
|
||||
def set_channel_name(self, name: str) -> None:
|
||||
self._channel_name = name
|
||||
|
||||
async def submit(self, msg: InboundMessage) -> None:
|
||||
"""Buffer *msg* and schedule flush after debounce delay."""
|
||||
sender = msg.sender_id
|
||||
|
||||
if sender not in self._buffers:
|
||||
self._buffers[sender] = []
|
||||
self._metadata[sender] = msg.metadata
|
||||
self._media[sender] = []
|
||||
self._buffers[sender].append(msg.content)
|
||||
if msg.message_id:
|
||||
self._message_ids[sender] = msg.message_id
|
||||
if msg.media:
|
||||
self._media[sender].extend(msg.media)
|
||||
|
||||
if sender in self._tasks:
|
||||
self._tasks[sender].cancel()
|
||||
|
||||
count = len(self._buffers[sender])
|
||||
wait = min(
|
||||
self.initial_debounce + (count - 1) * self.debounce_step,
|
||||
self.max_debounce,
|
||||
)
|
||||
|
||||
async def _flush(_s: str = sender, _w: float = wait) -> None:
|
||||
await asyncio.sleep(_w)
|
||||
await self._flush_sender(_s)
|
||||
|
||||
self._tasks[sender] = asyncio.create_task(_flush())
|
||||
|
||||
async def _flush_sender(self, sender: str) -> None:
|
||||
messages = self._buffers.pop(sender, [])
|
||||
metadata = self._metadata.pop(sender, None)
|
||||
media = self._media.pop(sender, [])
|
||||
message_id = self._message_ids.pop(sender, "")
|
||||
self._tasks.pop(sender, None)
|
||||
if not messages:
|
||||
return
|
||||
|
||||
merged = "\n".join(messages)
|
||||
chat_id = (metadata or {}).get("chat_id", sender)
|
||||
inbound = InboundMessage(
|
||||
channel=self._channel_name,
|
||||
sender_id=sender,
|
||||
chat_id=str(chat_id),
|
||||
content=merged,
|
||||
media=media,
|
||||
metadata=metadata or {},
|
||||
message_id=message_id,
|
||||
)
|
||||
if self.on_ready:
|
||||
await self.on_ready(inbound)
|
||||
|
||||
def cancel_all(self) -> None:
|
||||
"""Cancel all pending debounce tasks."""
|
||||
for task in self._tasks.values():
|
||||
task.cancel()
|
||||
self._tasks.clear()
|
||||
|
||||
|
||||
# ── Chunking ─────────────────────────────────────────────────────────
|
||||
|
||||
class ChunkingMiddleware(OutboundMiddlewareBase):
|
||||
"""Auto-split messages respecting format expansion.
|
||||
|
||||
Wraps the existing ``chunking.chunk_text`` utility and the
|
||||
re-splitting logic from ``Channel._prepare_chunks``.
|
||||
"""
|
||||
|
||||
def __init__(self, capabilities: Any) -> None:
|
||||
from .capabilities import ChannelCapabilities
|
||||
self._capabilities: ChannelCapabilities = capabilities
|
||||
|
||||
def prepare_chunks(
|
||||
self,
|
||||
content: str,
|
||||
limit: int,
|
||||
format_fn: Callable[[str], str] | None = None,
|
||||
) -> list[tuple[str, str]]:
|
||||
"""Build ``(formatted, raw)`` pairs, re-splitting when needed.
|
||||
|
||||
If *format_fn* is None, formatted == raw.
|
||||
"""
|
||||
from .base import chunk_text
|
||||
|
||||
if format_fn is None:
|
||||
format_fn = lambda t: t # noqa: E731
|
||||
|
||||
raw_chunks = chunk_text(content, limit)
|
||||
pairs: list[tuple[str, str]] = []
|
||||
for raw in raw_chunks:
|
||||
formatted = format_fn(raw)
|
||||
if len(formatted) <= limit:
|
||||
pairs.append((formatted, raw))
|
||||
else:
|
||||
sub_limit = max(limit // 2, 500)
|
||||
for sub_raw in chunk_text(raw, sub_limit):
|
||||
sub_fmt = format_fn(sub_raw)
|
||||
if len(sub_fmt) <= limit:
|
||||
pairs.append((sub_fmt, sub_raw))
|
||||
else:
|
||||
pairs.append((sub_raw, sub_raw))
|
||||
return pairs
|
||||
|
||||
|
||||
# ── Formatting ───────────────────────────────────────────────────────
|
||||
|
||||
class FormattingMiddleware(OutboundMiddlewareBase):
|
||||
"""Markdown -> channel format conversion.
|
||||
|
||||
Uses ``UnifiedFormatter`` configured from capabilities.
|
||||
"""
|
||||
|
||||
def __init__(self, capabilities: Any) -> None:
|
||||
from .formatter import UnifiedFormatter
|
||||
from .capabilities import ChannelCapabilities
|
||||
caps: ChannelCapabilities = capabilities
|
||||
self._formatter = UnifiedFormatter.for_channel(caps.format_type)
|
||||
|
||||
def format(self, text: str) -> str:
|
||||
"""Convert text to channel format."""
|
||||
return self._formatter.format(text)
|
||||
|
||||
async def process_outbound(
|
||||
self, message: OutboundMessage, context: dict[str, Any],
|
||||
) -> OutboundMessage | None:
|
||||
formatted = self._formatter.format(message.content)
|
||||
return dataclasses.replace(message, content=formatted)
|
||||
|
||||
|
||||
# ── Retry ────────────────────────────────────────────────────────────
|
||||
|
||||
class RetryMiddleware:
|
||||
"""Exponential backoff send retry.
|
||||
|
||||
Wraps ``retry.retry_async`` with channel-appropriate configuration.
|
||||
"""
|
||||
|
||||
def __init__(self, channel_name: str = "unknown") -> None:
|
||||
from .retry import DEFAULT_RETRY, RETRY_PRESETS
|
||||
self._config = RETRY_PRESETS.get(channel_name, DEFAULT_RETRY)
|
||||
self._channel_name = channel_name
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
coro_factory: Callable[[], Any],
|
||||
should_retry: Callable[[Exception, int], bool] | None = None,
|
||||
retry_after_s: Callable[[Exception], float | None] | None = None,
|
||||
) -> Any:
|
||||
"""Execute *coro_factory* with retry logic."""
|
||||
from .retry import retry_async
|
||||
|
||||
return await retry_async(
|
||||
coro_factory,
|
||||
config=self._config,
|
||||
should_retry=should_retry or (lambda exc, _: True),
|
||||
retry_after_s=retry_after_s,
|
||||
on_retry=lambda info: _logger.warning(
|
||||
f"{self._channel_name} retry {info.attempt}/{info.max_attempts} "
|
||||
f"in {info.delay_s:.2f}s: {info.error}"
|
||||
),
|
||||
label=f"{self._channel_name}.send",
|
||||
)
|
||||
|
||||
|
||||
# ── Typing ───────────────────────────────────────────────────────────
|
||||
|
||||
class TypingMiddleware:
|
||||
"""Typing indicator management.
|
||||
|
||||
Wraps ``TypingManager`` for use as a standalone middleware component.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
send_typing_fn: Callable[[str], Any],
|
||||
interval: float = 5.0,
|
||||
) -> None:
|
||||
self._manager = TypingManager(send_typing_fn, interval=interval)
|
||||
|
||||
async def start(self, chat_id: str) -> None:
|
||||
await self._manager.start(chat_id)
|
||||
|
||||
async def stop(self, chat_id: str) -> None:
|
||||
await self._manager.stop(chat_id)
|
||||
|
||||
async def stop_all(self) -> None:
|
||||
await self._manager.stop_all()
|
||||
|
||||
|
||||
# ── ACK Reaction ─────────────────────────────────────────────────────
|
||||
|
||||
class AckReactionMiddleware:
|
||||
"""ACK emoji reaction with configurable scope.
|
||||
|
||||
Scope controls when reactions are sent:
|
||||
- ``"all"``: react to every message
|
||||
- ``"direct"``: react only in DMs
|
||||
- ``"group-all"``: react in group chats (all messages)
|
||||
- ``"group-mentions"``: react in groups only when mentioned
|
||||
- ``"off"``: disable reactions
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
scope: str = "all",
|
||||
emoji: str = "\U0001f440",
|
||||
remove_after_reply: bool = False,
|
||||
send_fn: Callable[[str, str, str], Any] | None = None,
|
||||
remove_fn: Callable[[str, str, str], Any] | None = None,
|
||||
) -> None:
|
||||
self.scope = scope
|
||||
self.emoji = emoji
|
||||
self.remove_after_reply = remove_after_reply
|
||||
self._send_fn = send_fn
|
||||
self._remove_fn = remove_fn
|
||||
self._pending: dict[str, str] = {} # chat_id -> message_id
|
||||
|
||||
def should_react(self, *, is_group: bool, was_mentioned: bool) -> bool:
|
||||
if self.scope == "off":
|
||||
return False
|
||||
if self.scope == "all":
|
||||
return True
|
||||
if self.scope == "direct":
|
||||
return not is_group
|
||||
if self.scope == "group-all":
|
||||
return is_group
|
||||
if self.scope == "group-mentions":
|
||||
return is_group and was_mentioned
|
||||
return False
|
||||
|
||||
async def send_ack(self, chat_id: str, message_id: str) -> None:
|
||||
if self._send_fn and message_id:
|
||||
try:
|
||||
await self._send_fn(chat_id, message_id, self.emoji)
|
||||
if self.remove_after_reply:
|
||||
self._pending[chat_id] = message_id
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def remove_ack(self, chat_id: str) -> None:
|
||||
message_id = self._pending.pop(chat_id, None)
|
||||
if message_id and self._remove_fn:
|
||||
try:
|
||||
await self._remove_fn(chat_id, message_id, self.emoji)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
# ── Mention Gating ───────────────────────────────────────────────────
|
||||
|
||||
class MentionGatingMiddleware(InboundMiddleware):
|
||||
"""Filter messages based on mention policy.
|
||||
|
||||
Policy values:
|
||||
- ``"always"``: require mention in all chats
|
||||
- ``"group"``: require mention only in groups (default)
|
||||
- ``"off"``: never require mention
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
require_mention: str = "group",
|
||||
strip_fn: Callable[[str], str] | None = None,
|
||||
) -> None:
|
||||
self.require_mention = require_mention
|
||||
self._strip_fn = strip_fn
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
if not self._should_process(raw):
|
||||
return None
|
||||
# Strip mentions from group messages
|
||||
if raw.is_group and self._strip_fn:
|
||||
raw = dataclasses.replace(raw, text=self._strip_fn(raw.text))
|
||||
return raw
|
||||
|
||||
def _should_process(self, raw: RawIncoming) -> bool:
|
||||
if self.require_mention == "off":
|
||||
return True
|
||||
if self.require_mention == "always":
|
||||
return raw.was_mentioned
|
||||
# "group" — require mention only in groups
|
||||
if not raw.is_group:
|
||||
return True
|
||||
return raw.was_mentioned
|
||||
|
||||
|
||||
# ── AllowList ────────────────────────────────────────────────────────
|
||||
|
||||
class AllowListMiddleware(InboundMiddleware):
|
||||
"""Sender and channel allow-list enforcement."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
allowed_senders: set[str] | None = None,
|
||||
allowed_channels: set[str] | None = None,
|
||||
dm_policy: str = "allowlist",
|
||||
) -> None:
|
||||
self.allowed_senders = allowed_senders
|
||||
self.allowed_channels = allowed_channels
|
||||
self.dm_policy = dm_policy
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
# Channel allow-list
|
||||
if self.allowed_channels and str(raw.chat_id) not in self.allowed_channels:
|
||||
_logger.debug(f"Ignoring message from non-allowed channel {raw.chat_id}")
|
||||
return None
|
||||
|
||||
# Sender allow-list
|
||||
if not raw.is_group and self.dm_policy == "open":
|
||||
return raw # open DMs bypass sender checks
|
||||
|
||||
if not self._is_sender_allowed(raw.sender_id):
|
||||
_logger.debug(f"Ignoring message from non-allowed sender {raw.sender_id}")
|
||||
return None
|
||||
|
||||
return raw
|
||||
|
||||
def _is_sender_allowed(self, sender: str) -> bool:
|
||||
if not self.allowed_senders:
|
||||
return True
|
||||
sender_str = str(sender)
|
||||
if sender_str in self.allowed_senders:
|
||||
return True
|
||||
if "|" in sender_str:
|
||||
for part in sender_str.split("|"):
|
||||
if part and part in self.allowed_senders:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
# ── Group History ────────────────────────────────────────────────────
|
||||
|
||||
class GroupHistoryMiddleware(InboundMiddleware):
|
||||
"""Buffer non-mentioned group messages, inject as context when mentioned."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_per_chat: int = 50,
|
||||
max_age_seconds: int = 3600,
|
||||
) -> None:
|
||||
self._buffer = GroupHistoryBuffer(
|
||||
max_per_chat=max_per_chat, max_age_seconds=max_age_seconds,
|
||||
)
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
if not raw.is_group:
|
||||
return raw
|
||||
|
||||
ts = (
|
||||
raw.timestamp.timestamp()
|
||||
if hasattr(raw.timestamp, "timestamp")
|
||||
else time.time()
|
||||
)
|
||||
|
||||
if not raw.was_mentioned:
|
||||
self._buffer.add(
|
||||
raw.chat_id,
|
||||
HistoryEntry(
|
||||
sender_id=raw.sender_id,
|
||||
text=raw.text,
|
||||
timestamp=ts,
|
||||
message_id=raw.message_id,
|
||||
),
|
||||
)
|
||||
# Don't drop here — let MentionGatingMiddleware handle that
|
||||
return raw
|
||||
|
||||
# Mentioned: inject history context
|
||||
history_context = self._buffer.format_context(raw.chat_id)
|
||||
if history_context:
|
||||
raw = dataclasses.replace(
|
||||
raw,
|
||||
text=history_context + "\n\n[Current message - respond to this]\n" + raw.text,
|
||||
)
|
||||
self._buffer.clear(raw.chat_id)
|
||||
return raw
|
||||
|
||||
|
||||
# ── Pairing ──────────────────────────────────────────────────────────
|
||||
|
||||
class PairingMiddleware(InboundMiddleware):
|
||||
"""DM pairing flow management.
|
||||
|
||||
When dm_policy is "pairing", unapproved DM senders receive a
|
||||
pairing code. Approved senders pass through normally.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channel_name: str,
|
||||
send_response_fn: Callable[[str, str], Any] | None = None,
|
||||
dm_policy: str = "allowlist",
|
||||
) -> None:
|
||||
self._manager = PairingManager()
|
||||
self._channel_name = channel_name
|
||||
self._send_response_fn = send_response_fn
|
||||
self._dm_policy = dm_policy
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
if raw.is_group:
|
||||
return raw # pairing only applies to DMs
|
||||
|
||||
if self._dm_policy != "pairing":
|
||||
return raw
|
||||
|
||||
if self._manager.is_approved(self._channel_name, raw.sender_id):
|
||||
return raw
|
||||
|
||||
# Request pairing
|
||||
code = self._manager.request_pairing(self._channel_name, raw.sender_id)
|
||||
if self._send_response_fn:
|
||||
text = f"\U0001f510 Pairing required. Your code: {code}\nThis code expires in 1 hour."
|
||||
asyncio.ensure_future(self._send_response_fn(raw.chat_id, text))
|
||||
_logger.info(f"Pairing required for {raw.sender_id}, code sent")
|
||||
return None
|
||||
@@ -0,0 +1,306 @@
|
||||
"""Reusable channel mixins for common architecture patterns.
|
||||
|
||||
Three mixins that eliminate boilerplate across channels:
|
||||
|
||||
- ``WebhookMixin`` — aiohttp webhook server + httpx client + token refresh
|
||||
- ``WebSocketMixin`` — WS connect/reconnect/heartbeat loop
|
||||
- ``PollingMixin`` — async poll loop with backoff
|
||||
|
||||
Each mixin works with the Channel base class. Subclasses override
|
||||
a small set of abstract/hook methods to define platform-specific behavior.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Token refresh mixin (shared by Webhook & WebSocket channels)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
class TokenMixin:
|
||||
"""Mixin for channels that need OAuth-style token management.
|
||||
|
||||
Subclass must implement ``_fetch_token()`` returning
|
||||
``(access_token, expires_in_seconds)``.
|
||||
"""
|
||||
|
||||
_access_token: str | None = None
|
||||
_token_expires: float = 0
|
||||
_http_client: Any = None # httpx.AsyncClient
|
||||
|
||||
async def _fetch_token(self) -> tuple[str, int]:
|
||||
"""Fetch a new access token. Return (token, expires_in_seconds).
|
||||
|
||||
Must be implemented by the channel.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
async def _refresh_token(self) -> None:
|
||||
token, expire = await self._fetch_token()
|
||||
self._access_token = token
|
||||
self._token_expires = time.monotonic() + expire - 300
|
||||
logger.debug(f"{getattr(self, 'name', '?')} token refreshed, expires in {expire}s")
|
||||
|
||||
async def _ensure_token(self) -> str:
|
||||
if not self._access_token or time.monotonic() >= self._token_expires:
|
||||
await self._refresh_token()
|
||||
return self._access_token
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Webhook + REST mixin
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
class WebhookMixin:
|
||||
"""Mixin for channels that use an HTTP webhook server for inbound
|
||||
and REST API for outbound.
|
||||
|
||||
Provides:
|
||||
- aiohttp web server lifecycle (start/stop)
|
||||
- httpx async client lifecycle
|
||||
- Route registration via ``_webhook_routes()``
|
||||
|
||||
Subclass must implement:
|
||||
- ``_webhook_routes()`` → list of (method, path, handler)
|
||||
- ``_get_webhook_port()`` → int
|
||||
"""
|
||||
|
||||
_http_client: Any = None
|
||||
_runner: Any = None
|
||||
_site: Any = None
|
||||
|
||||
def _get_webhook_port(self) -> int:
|
||||
return getattr(self.config, "webhook_port", 9000)
|
||||
|
||||
def _webhook_routes(self) -> list[tuple[str, str, Any]]:
|
||||
"""Return [(method, path, handler), ...]. Override in subclass."""
|
||||
return []
|
||||
|
||||
async def _start_webhook_server(self) -> None:
|
||||
"""Start aiohttp webhook server + httpx client.
|
||||
|
||||
If ``_shared_webhook_server`` is set (by ChannelManager), the
|
||||
aiohttp server is already running on the shared port — only
|
||||
create the httpx outbound client.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
proxy = getattr(self.config, "proxy", None) or None
|
||||
self._http_client = httpx.AsyncClient(timeout=15, proxy=proxy)
|
||||
|
||||
# Shared webhook mode: routes already registered on shared server
|
||||
if getattr(self, "_shared_webhook_server", None):
|
||||
logger.info(f"{getattr(self, 'name', '?')} using shared webhook server")
|
||||
return
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
app = web.Application()
|
||||
for method, path, handler in self._webhook_routes():
|
||||
if method.upper() == "GET":
|
||||
app.router.add_get(path, handler)
|
||||
else:
|
||||
app.router.add_post(path, handler)
|
||||
|
||||
self._runner = web.AppRunner(app)
|
||||
await self._runner.setup()
|
||||
port = self._get_webhook_port()
|
||||
self._site = web.TCPSite(self._runner, "0.0.0.0", port)
|
||||
await self._site.start()
|
||||
logger.info(f"{getattr(self, 'name', '?')} webhook on port {port}")
|
||||
|
||||
async def _stop_webhook_server(self) -> None:
|
||||
if self._site:
|
||||
await self._site.stop()
|
||||
self._site = None
|
||||
if self._runner:
|
||||
await self._runner.cleanup()
|
||||
self._runner = None
|
||||
if self._http_client:
|
||||
await self._http_client.aclose()
|
||||
self._http_client = None
|
||||
|
||||
async def _api_post(self, url: str, body: dict, headers: dict | None = None) -> dict:
|
||||
"""POST JSON to API, return parsed response. Raises on HTTP error."""
|
||||
resp = await self._http_client.post(url, json=body, headers=headers)
|
||||
data = resp.json()
|
||||
return data
|
||||
|
||||
async def _api_get(self, url: str, headers: dict | None = None) -> dict:
|
||||
resp = await self._http_client.get(url, headers=headers)
|
||||
return resp.json()
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# WebSocket mixin
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
class WebSocketMixin:
|
||||
"""Mixin for channels that receive messages via WebSocket.
|
||||
|
||||
Provides:
|
||||
- Connect/reconnect loop with exponential backoff
|
||||
- Heartbeat task management
|
||||
- Message dispatch
|
||||
|
||||
Subclass must implement:
|
||||
- ``_get_ws_url()`` → WebSocket URL to connect to
|
||||
- ``_on_ws_message(data)`` → handle a parsed message dict
|
||||
- ``_on_ws_connected(ws)`` → called after connection (send identify, etc.)
|
||||
|
||||
Optional overrides:
|
||||
- ``_ws_heartbeat_interval`` → seconds between heartbeats (0 = disabled)
|
||||
- ``_on_ws_heartbeat(ws)`` → send heartbeat
|
||||
"""
|
||||
|
||||
_ws_session: Any = None
|
||||
_ws_heartbeat_task: asyncio.Task | None = None
|
||||
_ws_heartbeat_interval: float = 0 # 0 = no heartbeat
|
||||
_ws_reconnect_delay: float = 5.0
|
||||
|
||||
async def _get_ws_url(self) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
async def _on_ws_connected(self, ws) -> None:
|
||||
"""Called after WebSocket connects. Send identify/auth here."""
|
||||
pass
|
||||
|
||||
async def _on_ws_message(self, data: dict | str) -> None:
|
||||
"""Handle a single WebSocket message."""
|
||||
raise NotImplementedError
|
||||
|
||||
async def _on_ws_heartbeat(self, ws) -> None:
|
||||
"""Send a heartbeat. Override if needed."""
|
||||
pass
|
||||
|
||||
async def _ws_loop(self) -> None:
|
||||
"""Main WebSocket loop with auto-reconnect."""
|
||||
import os
|
||||
import aiohttp
|
||||
|
||||
while getattr(self, "_running", False):
|
||||
try:
|
||||
ws_url = await self._get_ws_url()
|
||||
# Resolve proxy: channel config > environment variable
|
||||
proxy = getattr(getattr(self, "config", None), "proxy", None)
|
||||
if not proxy:
|
||||
proxy = (os.environ.get("https_proxy")
|
||||
or os.environ.get("HTTPS_PROXY")
|
||||
or os.environ.get("http_proxy")
|
||||
or os.environ.get("HTTP_PROXY")
|
||||
or None)
|
||||
logger.debug(f"{getattr(self, 'name', '?')} WS connecting to {ws_url[:60]}... proxy={proxy}")
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.ws_connect(ws_url, proxy=proxy, timeout=aiohttp.ClientWSTimeout(ws_close=30)) as ws:
|
||||
logger.info(f"{getattr(self, 'name', '?')} WebSocket connected")
|
||||
self._ws_session = ws
|
||||
await self._on_ws_connected(ws)
|
||||
|
||||
# Start heartbeat if configured
|
||||
if self._ws_heartbeat_interval > 0:
|
||||
self._ws_heartbeat_task = asyncio.create_task(
|
||||
self._ws_heartbeat_loop(ws)
|
||||
)
|
||||
|
||||
async for msg in ws:
|
||||
if msg.type == aiohttp.WSMsgType.TEXT:
|
||||
try:
|
||||
data = json.loads(msg.data)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
data = msg.data
|
||||
await self._on_ws_message(data)
|
||||
elif msg.type in (
|
||||
aiohttp.WSMsgType.CLOSED,
|
||||
aiohttp.WSMsgType.ERROR,
|
||||
):
|
||||
break
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"{getattr(self, 'name', '?')} WS error: {e}")
|
||||
|
||||
self._ws_cleanup_heartbeat()
|
||||
self._ws_session = None
|
||||
|
||||
if getattr(self, "_running", False):
|
||||
logger.info(f"{getattr(self, 'name', '?')} reconnecting in {self._ws_reconnect_delay}s...")
|
||||
await asyncio.sleep(self._ws_reconnect_delay)
|
||||
|
||||
async def _ws_heartbeat_loop(self, ws) -> None:
|
||||
while True:
|
||||
try:
|
||||
await self._on_ws_heartbeat(ws)
|
||||
except Exception:
|
||||
break
|
||||
await asyncio.sleep(self._ws_heartbeat_interval)
|
||||
|
||||
def _ws_cleanup_heartbeat(self) -> None:
|
||||
if self._ws_heartbeat_task:
|
||||
self._ws_heartbeat_task.cancel()
|
||||
self._ws_heartbeat_task = None
|
||||
|
||||
async def _ws_send_json(self, data: dict) -> None:
|
||||
"""Send JSON to the active WebSocket."""
|
||||
if self._ws_session:
|
||||
await self._ws_session.send_str(json.dumps(data))
|
||||
|
||||
async def _stop_ws(self) -> None:
|
||||
self._ws_cleanup_heartbeat()
|
||||
if self._ws_session:
|
||||
await self._ws_session.close()
|
||||
self._ws_session = None
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Polling mixin
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
class PollingMixin:
|
||||
"""Mixin for channels that poll for new messages.
|
||||
|
||||
Provides:
|
||||
- Poll loop with configurable interval
|
||||
- Error handling + reconnect
|
||||
|
||||
Subclass must implement:
|
||||
- ``_poll_once()`` → fetch and enqueue new messages
|
||||
- ``_get_poll_interval()`` → seconds between polls
|
||||
"""
|
||||
|
||||
_poll_task: asyncio.Task | None = None
|
||||
|
||||
def _get_poll_interval(self) -> float:
|
||||
return getattr(self.config, "poll_interval", 30)
|
||||
|
||||
async def _poll_once(self) -> None:
|
||||
"""Fetch new messages and enqueue them. Override in subclass."""
|
||||
raise NotImplementedError
|
||||
|
||||
async def _start_polling(self) -> None:
|
||||
self._poll_task = asyncio.create_task(self._poll_loop())
|
||||
|
||||
async def _poll_loop(self) -> None:
|
||||
interval = self._get_poll_interval()
|
||||
while getattr(self, "_running", False):
|
||||
try:
|
||||
await self._poll_once()
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"{getattr(self, 'name', '?')} poll error: {e}")
|
||||
await asyncio.sleep(interval)
|
||||
|
||||
async def _stop_polling(self) -> None:
|
||||
if self._poll_task:
|
||||
self._poll_task.cancel()
|
||||
self._poll_task = None
|
||||
@@ -0,0 +1,226 @@
|
||||
"""Plugin-based channel interface.
|
||||
|
||||
A ChannelPlugin is a declarative object with optional adapter slots.
|
||||
The framework inspects which slots are filled and auto-assembles
|
||||
the message processing pipeline.
|
||||
|
||||
The ``Channel`` base class extends ``ChannelPlugin``, so all channel
|
||||
implementations are automatically ChannelPlugin instances.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
from .capabilities import ChannelCapabilities
|
||||
|
||||
|
||||
# ── Channel metadata ─────────────────────────────────────────────────
|
||||
|
||||
@dataclass
|
||||
class ChannelMeta:
|
||||
"""Channel metadata for registry and UI."""
|
||||
|
||||
id: str
|
||||
label: str
|
||||
description: str = ""
|
||||
docs_path: str = ""
|
||||
system_image: str = "" # icon name
|
||||
|
||||
|
||||
# ── Adapter Protocols (slots) ────────────────────────────────────────
|
||||
|
||||
@runtime_checkable
|
||||
class ConfigAdapter(Protocol):
|
||||
"""Account configuration management."""
|
||||
|
||||
def list_account_ids(self, config: Any) -> list[str]: ...
|
||||
def resolve_account(self, config: Any, account_id: str | None = None) -> Any: ...
|
||||
def is_enabled(self, account: Any, config: Any) -> bool: ...
|
||||
def is_configured(self, account: Any, config: Any) -> bool: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SecurityAdapter(Protocol):
|
||||
"""DM policy and security warnings."""
|
||||
|
||||
def resolve_dm_policy(self, ctx: Any) -> str: ... # "open" | "allowlist" | "pairing"
|
||||
def collect_warnings(self, ctx: Any) -> list[str]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class GroupAdapter(Protocol):
|
||||
"""Per-group policy resolution."""
|
||||
|
||||
def resolve_require_mention(self, ctx: Any) -> bool | None: ...
|
||||
def resolve_tool_policy(self, ctx: Any) -> dict[str, Any] | None: ...
|
||||
def resolve_intro_hint(self, ctx: Any) -> str | None: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class MentionAdapter(Protocol):
|
||||
"""Bot mention detection and stripping."""
|
||||
|
||||
def strip_mentions(self, text: str, ctx: Any) -> str: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class OutboundAdapter(Protocol):
|
||||
"""Outbound message delivery."""
|
||||
|
||||
delivery_mode: str # "direct" | "gateway" | "hybrid"
|
||||
|
||||
async def send_text(self, ctx: Any) -> bool: ...
|
||||
async def send_media(self, ctx: Any) -> bool: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ThreadingAdapter(Protocol):
|
||||
"""Reply threading behavior."""
|
||||
|
||||
def resolve_reply_to_mode(self, ctx: Any) -> str: ... # "off" | "first" | "all"
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class StreamingAdapter(Protocol):
|
||||
"""Edit-in-place streaming output."""
|
||||
|
||||
async def edit_message(self, chat_id: str, message_id: str, text: str) -> bool: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class DirectoryAdapter(Protocol):
|
||||
"""Contact/group directory queries."""
|
||||
|
||||
async def list_peers(self, ctx: Any) -> list[dict]: ...
|
||||
async def list_groups(self, ctx: Any) -> list[dict]: ...
|
||||
async def list_group_members(self, ctx: Any) -> list[dict]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class StatusAdapter(Protocol):
|
||||
"""Health probing and status reporting."""
|
||||
|
||||
async def probe_account(self, ctx: Any) -> Any: ...
|
||||
async def audit_account(self, ctx: Any) -> Any: ...
|
||||
def collect_status_issues(self, accounts: list) -> list[dict]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class HeartbeatAdapter(Protocol):
|
||||
"""Channel heartbeat / readiness checks."""
|
||||
|
||||
async def check_ready(self, ctx: Any) -> tuple[bool, str]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ActionsAdapter(Protocol):
|
||||
"""Message actions (react, edit, delete, poll, etc.)."""
|
||||
|
||||
def list_actions(self) -> list[str]: ...
|
||||
async def handle_action(self, action: str, ctx: Any) -> Any: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class PairingAdapter(Protocol):
|
||||
"""DM pairing flow."""
|
||||
|
||||
id_label: str
|
||||
|
||||
def normalize_entry(self, entry: str) -> str: ...
|
||||
async def notify_approval(self, ctx: Any) -> None: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class OnboardingAdapter(Protocol):
|
||||
"""Interactive setup wizard hooks."""
|
||||
|
||||
async def wizard_steps(self, ctx: Any) -> list[dict]: ...
|
||||
async def validate_step(self, step: str, value: Any) -> str | None: ...
|
||||
|
||||
|
||||
# ── Reload policy ────────────────────────────────────────────────────
|
||||
|
||||
@dataclass
|
||||
class ReloadPolicy:
|
||||
"""Declares which config prefixes trigger a channel reload."""
|
||||
|
||||
config_prefixes: list[str] = field(default_factory=list)
|
||||
noop_prefixes: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
# ── ChannelPlugin ────────────────────────────────────────────────────
|
||||
|
||||
class ChannelPlugin:
|
||||
"""Declarative channel plugin with optional adapter slots.
|
||||
|
||||
Replaces the monolithic Channel base class. Each slot is optional —
|
||||
the framework adapts behavior based on which are present.
|
||||
|
||||
Usage::
|
||||
|
||||
class MyPlugin(ChannelPlugin):
|
||||
id = "my_channel"
|
||||
meta = ChannelMeta(id="my_channel", label="My Channel")
|
||||
capabilities = ChannelCapabilities(...)
|
||||
|
||||
def __init__(self):
|
||||
self.outbound = MyOutboundAdapter()
|
||||
self.config_adapter = MyConfigAdapter()
|
||||
|
||||
async def start(self, config, account_id=None):
|
||||
...
|
||||
|
||||
async def stop(self, account_id=None):
|
||||
...
|
||||
"""
|
||||
|
||||
id: str = ""
|
||||
meta: ChannelMeta | None = None
|
||||
capabilities: ChannelCapabilities = ChannelCapabilities()
|
||||
|
||||
# Optional adapter slots — fill what you need
|
||||
# Default: SingleAccountConfigAdapter so every plugin has multi-account
|
||||
# support out of the box (returns a single "default" account).
|
||||
config_adapter: ConfigAdapter | None = None
|
||||
|
||||
def __init_subclass__(cls, **kwargs: Any) -> None:
|
||||
super().__init_subclass__(**kwargs)
|
||||
|
||||
def __init__(self) -> None:
|
||||
# Provide default SingleAccountConfigAdapter if not overridden
|
||||
if self.config_adapter is None:
|
||||
from .config import SingleAccountConfigAdapter
|
||||
self.config_adapter = SingleAccountConfigAdapter()
|
||||
security: SecurityAdapter | None = None
|
||||
groups: GroupAdapter | None = None
|
||||
mentions: MentionAdapter | None = None
|
||||
outbound: OutboundAdapter | None = None
|
||||
threading: ThreadingAdapter | None = None
|
||||
streaming: StreamingAdapter | None = None
|
||||
directory: DirectoryAdapter | None = None
|
||||
status: StatusAdapter | None = None
|
||||
heartbeat: HeartbeatAdapter | None = None
|
||||
actions: ActionsAdapter | None = None
|
||||
pairing: PairingAdapter | None = None
|
||||
onboarding: OnboardingAdapter | None = None
|
||||
|
||||
# Lifecycle
|
||||
reload: ReloadPolicy | None = None
|
||||
|
||||
# Connection management
|
||||
async def start(self, config: Any, account_id: str | None = None) -> None:
|
||||
"""Start the channel (or a specific account)."""
|
||||
|
||||
async def stop(self, account_id: str | None = None) -> None:
|
||||
"""Stop the channel (or a specific account)."""
|
||||
|
||||
def filled_slots(self) -> list[str]:
|
||||
"""Return names of adapter slots that are not None."""
|
||||
slot_names = [
|
||||
"config_adapter", "security", "groups", "mentions", "outbound",
|
||||
"threading", "streaming", "directory", "status", "heartbeat",
|
||||
"actions", "pairing", "onboarding",
|
||||
]
|
||||
return [s for s in slot_names if getattr(self, s, None) is not None]
|
||||
@@ -0,0 +1,26 @@
|
||||
"""QQ channel for EvoScientist.
|
||||
|
||||
Uses the official qq-botpy SDK for WebSocket connection.
|
||||
|
||||
Usage in config:
|
||||
channel_enabled = "qq"
|
||||
qq_app_id = "your_app_id"
|
||||
qq_app_secret = "your_app_secret"
|
||||
"""
|
||||
|
||||
from .channel import QQChannel, QQConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["QQChannel", "QQConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> QQChannel:
|
||||
allowed = _parse_csv(getattr(config, "qq_allowed_senders", ""))
|
||||
return QQChannel(QQConfig(
|
||||
app_id=getattr(config, "qq_app_id", ""),
|
||||
app_secret=getattr(config, "qq_app_secret", ""),
|
||||
allowed_senders=allowed,
|
||||
))
|
||||
|
||||
|
||||
register_channel("qq", create_from_config)
|
||||
@@ -0,0 +1,258 @@
|
||||
"""QQ Bot channel — powered by botpy SDK.
|
||||
|
||||
Uses the official qq-botpy SDK for WebSocket connection and message handling.
|
||||
No manual WebSocket protocol implementation needed.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections import deque
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..capabilities import QQ as QQ_CAPS
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
try:
|
||||
import botpy
|
||||
from botpy.message import C2CMessage, GroupMessage
|
||||
|
||||
QQ_AVAILABLE = True
|
||||
except ImportError:
|
||||
QQ_AVAILABLE = False
|
||||
botpy = None
|
||||
C2CMessage = None
|
||||
GroupMessage = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class QQConfig(BaseChannelConfig):
|
||||
app_id: str = ""
|
||||
app_secret: str = ""
|
||||
text_chunk_limit: int = 4096
|
||||
|
||||
|
||||
def _make_bot_class(channel: "QQChannel") -> "type[botpy.Client]":
|
||||
"""Create a botpy Client subclass bound to the given channel."""
|
||||
intents = botpy.Intents(public_messages=True, direct_message=True)
|
||||
|
||||
class _Bot(botpy.Client):
|
||||
def __init__(self):
|
||||
super().__init__(intents=intents)
|
||||
|
||||
async def on_ready(self):
|
||||
logger.info(f"QQ bot ready: {self.robot.name}")
|
||||
|
||||
async def on_c2c_message_create(self, message: "C2CMessage"):
|
||||
await channel._on_msg(message, "c2c")
|
||||
|
||||
async def on_group_at_message_create(self, message: "GroupMessage"):
|
||||
await channel._on_msg(message, "group")
|
||||
|
||||
return _Bot
|
||||
|
||||
|
||||
class QQChannel(Channel):
|
||||
capabilities = QQ_CAPS
|
||||
name = "qq"
|
||||
_ready_attrs = ("_client", "_running")
|
||||
_mention_pattern = r"@\S+\s*"
|
||||
_mention_strip_count = 1
|
||||
|
||||
def __init__(self, config: QQConfig):
|
||||
super().__init__(config)
|
||||
self._client: "botpy.Client | None" = None
|
||||
self._bot_task: asyncio.Task | None = None
|
||||
self._processed_ids: deque = deque(maxlen=1000)
|
||||
self._msg_seq: dict[str, int] = {} # msg_id -> next seq number
|
||||
self._msg_seq_order: deque = deque(maxlen=500)
|
||||
|
||||
# ── Lifecycle ─────────────────────────────────────────────────
|
||||
|
||||
async def start(self) -> None:
|
||||
if not QQ_AVAILABLE:
|
||||
raise ChannelError("QQ SDK not installed. Run: pip install qq-botpy")
|
||||
if not self.config.app_id or not self.config.app_secret:
|
||||
raise ChannelError("QQ app_id and app_secret are required")
|
||||
self._running = True
|
||||
BotClass = _make_bot_class(self)
|
||||
self._client = BotClass()
|
||||
self._bot_task = asyncio.create_task(self._run_bot())
|
||||
logger.info("QQ channel starting...")
|
||||
|
||||
async def _run_bot(self) -> None:
|
||||
try:
|
||||
await self._client.start(appid=self.config.app_id, secret=self.config.app_secret)
|
||||
except Exception as e:
|
||||
logger.error(f"QQ auth failed: {e}")
|
||||
self._running = False
|
||||
|
||||
# ── Incoming ──────────────────────────────────────────────────
|
||||
|
||||
async def _on_msg(self, message, msg_type: str) -> None:
|
||||
try:
|
||||
if message.id in self._processed_ids:
|
||||
return
|
||||
self._processed_ids.append(message.id)
|
||||
|
||||
author = message.author
|
||||
content = (message.content or "").strip()
|
||||
|
||||
if msg_type == "c2c":
|
||||
sender_id = str(getattr(author, "user_openid", ""))
|
||||
chat_id = sender_id
|
||||
else:
|
||||
sender_id = str(getattr(author, "member_openid", ""))
|
||||
chat_id = str(getattr(message, "group_openid", ""))
|
||||
|
||||
# Handle attachments (images, files, audio, video)
|
||||
annotations: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
attachments = getattr(message, "attachments", None) or []
|
||||
for att in attachments:
|
||||
url = getattr(att, "url", "") or ""
|
||||
filename = getattr(att, "filename", "attachment") or "attachment"
|
||||
content_type = getattr(att, "content_type", "") or ""
|
||||
if url:
|
||||
local, ann = await self._download_attachment(
|
||||
url, f"qq_{filename}",
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
annotations.append(ann)
|
||||
else:
|
||||
annotations.append(f"[{content_type or 'attachment'}: {filename}]")
|
||||
|
||||
if not content and not media_paths and not annotations:
|
||||
return
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=content,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=datetime.now(),
|
||||
message_id=message.id,
|
||||
is_group=(msg_type == "group"),
|
||||
was_mentioned=True,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"msg_type": msg_type,
|
||||
"event_id": message.id,
|
||||
"backend": "qq",
|
||||
},
|
||||
))
|
||||
except Exception as e:
|
||||
logger.error(f"Error handling QQ message: {e}")
|
||||
|
||||
# ── Send ──────────────────────────────────────────────────────
|
||||
|
||||
def _next_msg_seq(self, msg_id: str) -> int:
|
||||
"""Return the next msg_seq for *msg_id* and increment the counter."""
|
||||
seq = self._msg_seq.get(msg_id, 1)
|
||||
self._msg_seq[msg_id] = seq + 1
|
||||
if msg_id not in set(self._msg_seq_order):
|
||||
self._msg_seq_order.append(msg_id)
|
||||
if len(self._msg_seq_order) > 500:
|
||||
oldest = self._msg_seq_order.popleft()
|
||||
self._msg_seq.pop(oldest, None)
|
||||
return seq
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
if not self._client:
|
||||
raise ChannelError("QQ client not initialized")
|
||||
msg_type = (metadata or {}).get("msg_type", "c2c")
|
||||
msg_id = (metadata or {}).get("event_id", "")
|
||||
seq = self._next_msg_seq(msg_id)
|
||||
if msg_type == "group":
|
||||
await self._client.api.post_group_message(
|
||||
group_openid=chat_id, msg_type=0,
|
||||
content=raw_text, msg_id=msg_id, msg_seq=seq,
|
||||
)
|
||||
else:
|
||||
await self._client.api.post_c2c_message(
|
||||
openid=chat_id, msg_type=0,
|
||||
content=raw_text, msg_id=msg_id, msg_seq=seq,
|
||||
)
|
||||
|
||||
# _send_typing_action: inherited no-op (QQ Bot API has no typing indicator)
|
||||
|
||||
# ── Media send ────────────────────────────────────────────────
|
||||
|
||||
# qq-botpy file_type constants: 1=image, 2=video, 3=audio
|
||||
_FILE_TYPE_MAP = {
|
||||
".jpg": 1, ".jpeg": 1, ".png": 1, ".gif": 1, ".webp": 1, ".bmp": 1,
|
||||
".mp4": 2, ".mov": 2, ".avi": 2,
|
||||
".mp3": 3, ".ogg": 3, ".m4a": 3, ".wav": 3, ".silk": 3,
|
||||
}
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a media file through QQ Bot API.
|
||||
|
||||
Uses post_group_file / post_c2c_file with a URL. Local files
|
||||
without a public URL are not supported — falls back to a text hint.
|
||||
"""
|
||||
if not self._client:
|
||||
raise ChannelError("QQ client not initialized")
|
||||
|
||||
from pathlib import Path
|
||||
chat_id = self._resolve_media_chat_id(recipient, metadata)
|
||||
msg_type = (metadata or {}).get("msg_type", "c2c")
|
||||
ext = Path(file_path).suffix.lower()
|
||||
file_type = self._FILE_TYPE_MAP.get(ext, 1) # default to image
|
||||
|
||||
# qq-botpy file API requires a URL, not a local path
|
||||
is_url = file_path.startswith("http://") or file_path.startswith("https://")
|
||||
if not is_url:
|
||||
# Fallback: send text hint for local files
|
||||
name = Path(file_path).name
|
||||
hint = f"[文件] {name}" + (f"\n{caption}" if caption else "")
|
||||
await self._send_chunk(chat_id, hint, hint, None, metadata or {})
|
||||
return True
|
||||
|
||||
try:
|
||||
if msg_type == "group":
|
||||
await self._client.api.post_group_file(
|
||||
group_openid=chat_id,
|
||||
file_type=file_type,
|
||||
url=file_path,
|
||||
srv_send_msg=True,
|
||||
)
|
||||
else:
|
||||
await self._client.api.post_c2c_file(
|
||||
openid=chat_id,
|
||||
file_type=file_type,
|
||||
url=file_path,
|
||||
srv_send_msg=True,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"QQ media send failed: {e}")
|
||||
return False
|
||||
|
||||
if caption:
|
||||
await self._send_chunk(chat_id, caption, caption, None, metadata or {})
|
||||
return True
|
||||
|
||||
# ── Cleanup ───────────────────────────────────────────────────
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
self._running = False
|
||||
if self._bot_task:
|
||||
self._bot_task.cancel()
|
||||
try:
|
||||
await self._bot_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._client = None
|
||||
logger.info("QQ channel stopped")
|
||||
@@ -0,0 +1,33 @@
|
||||
"""QQ Bot credential validation."""
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
QQ_TOKEN_URL = "https://bots.qq.com/app/getAppAccessToken"
|
||||
|
||||
|
||||
async def validate_qq(
|
||||
app_id: str,
|
||||
app_secret: str,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate QQ Bot credentials by fetching an access token."""
|
||||
if not app_id or not app_secret:
|
||||
return False, "app_id and app_secret are required"
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return False, "httpx not installed"
|
||||
|
||||
body = {"appId": app_id, "clientSecret": app_secret}
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
resp = await client.post(QQ_TOKEN_URL, json=body, timeout=10)
|
||||
data = resp.json()
|
||||
if data.get("access_token"):
|
||||
return True, "QQ Bot credentials valid"
|
||||
return False, f"Error: {data.get('message', data)}"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
@@ -0,0 +1,87 @@
|
||||
"""QQ channel server.
|
||||
|
||||
Standalone script to run the QQ channel with CLI options.
|
||||
|
||||
Usage:
|
||||
python -m EvoScientist.channels.qq.serve --app-id ID --app-secret SECRET [OPTIONS]
|
||||
|
||||
Examples:
|
||||
# Basic usage
|
||||
python -m EvoScientist.channels.qq.serve --app-id ID --app-secret SECRET
|
||||
|
||||
# Sandbox mode with allowed senders
|
||||
python -m EvoScientist.channels.qq.serve --app-id ID --app-secret SECRET --allow user123
|
||||
|
||||
# With agent and thinking
|
||||
python -m EvoScientist.channels.qq.serve --app-id ID --app-secret SECRET --agent --thinking
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import QQChannel, QQConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="QQ channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--app-id",
|
||||
required=True,
|
||||
help="QQ bot app ID",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--app-secret",
|
||||
required=True,
|
||||
help="QQ bot app secret",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender (QQ user ID). Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler (default: echo)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content as intermediate messages (requires --agent)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = QQConfig(
|
||||
app_id=args.app_id,
|
||||
app_secret=args.app_secret,
|
||||
allowed_senders=set(args.allowed_senders) if args.allowed_senders else None,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = QQChannel(config)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,122 @@
|
||||
"""Configurable exponential-backoff retry for async callables."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import random
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import TypeVar
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
@dataclass
|
||||
class RetryConfig:
|
||||
"""Configuration for retry behaviour."""
|
||||
|
||||
attempts: int = 3
|
||||
min_delay_s: float = 0.3
|
||||
max_delay_s: float = 30.0
|
||||
jitter: float = 0.1 # ±10 % random offset
|
||||
|
||||
|
||||
@dataclass
|
||||
class RetryInfo:
|
||||
"""Information passed to the *on_retry* callback."""
|
||||
|
||||
attempt: int
|
||||
max_attempts: int
|
||||
delay_s: float
|
||||
error: Exception
|
||||
label: str | None = None
|
||||
|
||||
|
||||
async def retry_async(
|
||||
fn: Callable[[], Awaitable[T]],
|
||||
config: RetryConfig = RetryConfig(),
|
||||
*,
|
||||
should_retry: Callable[[Exception, int], bool] | None = None,
|
||||
retry_after_s: Callable[[Exception], float | None] | None = None,
|
||||
on_retry: Callable[[RetryInfo], None] | None = None,
|
||||
label: str | None = None,
|
||||
) -> T:
|
||||
"""Execute *fn* with exponential-backoff retry.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
fn:
|
||||
Zero-argument async factory — called on every attempt so the
|
||||
awaitable is always fresh.
|
||||
config:
|
||||
Retry timing / attempt parameters.
|
||||
should_retry:
|
||||
``(exception, attempt) -> bool``. Return ``False`` to abort
|
||||
immediately. When *None* every exception is retried.
|
||||
retry_after_s:
|
||||
``(exception) -> seconds | None``. If the server provides a
|
||||
``Retry-After`` value (e.g. HTTP 429), return it here. The
|
||||
actual delay will be ``max(server_value, min_delay_s)``.
|
||||
on_retry:
|
||||
Optional callback invoked before each retry sleep.
|
||||
label:
|
||||
Human-readable label included in :class:`RetryInfo`.
|
||||
"""
|
||||
last_exc: Exception | None = None
|
||||
for attempt in range(1, config.attempts + 1):
|
||||
try:
|
||||
return await fn()
|
||||
except Exception as exc:
|
||||
last_exc = exc
|
||||
|
||||
if attempt >= config.attempts:
|
||||
raise
|
||||
|
||||
if should_retry is not None and not should_retry(exc, attempt):
|
||||
raise
|
||||
|
||||
# Compute delay
|
||||
server_delay: float | None = None
|
||||
if retry_after_s is not None:
|
||||
server_delay = retry_after_s(exc)
|
||||
|
||||
if server_delay is not None:
|
||||
base_delay = max(server_delay, config.min_delay_s)
|
||||
else:
|
||||
base_delay = config.min_delay_s * (2 ** (attempt - 1))
|
||||
|
||||
# Apply jitter
|
||||
jittered = base_delay * (1 + random.uniform(-config.jitter, config.jitter))
|
||||
|
||||
# Clamp to [min_delay_s, max_delay_s]
|
||||
delay = max(config.min_delay_s, min(jittered, config.max_delay_s))
|
||||
|
||||
if on_retry is not None:
|
||||
on_retry(RetryInfo(
|
||||
attempt=attempt,
|
||||
max_attempts=config.attempts,
|
||||
delay_s=delay,
|
||||
error=exc,
|
||||
label=label,
|
||||
))
|
||||
|
||||
await asyncio.sleep(delay)
|
||||
|
||||
# Should never reach here, but satisfy the type checker.
|
||||
assert last_exc is not None # noqa: S101
|
||||
raise last_exc
|
||||
|
||||
|
||||
# ── Presets ──────────────────────────────────────────────────────────
|
||||
|
||||
TELEGRAM_RETRY = RetryConfig(attempts=3, min_delay_s=0.4, max_delay_s=30.0, jitter=0.1)
|
||||
DEFAULT_RETRY = RetryConfig()
|
||||
|
||||
# Discord, Slack, Teams, Feishu all use the same config (attempts=3,
|
||||
# min_delay_s=0.5, max_delay_s=30.0, jitter=0.1) — close enough to
|
||||
# DEFAULT_RETRY that separate presets add no value. Channels that
|
||||
# don't appear in RETRY_PRESETS already fall back to DEFAULT_RETRY.
|
||||
|
||||
RETRY_PRESETS: dict[str, RetryConfig] = {
|
||||
"telegram": TELEGRAM_RETRY,
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Signal channel for EvoScientist.
|
||||
|
||||
Uses signal-cli in JSON RPC mode — no public IP needed.
|
||||
|
||||
Usage in config:
|
||||
channel_enabled = "signal"
|
||||
signal_phone_number = "+1234567890"
|
||||
"""
|
||||
|
||||
from .channel import SignalChannel, SignalConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["SignalChannel", "SignalConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> SignalChannel:
|
||||
allowed = _parse_csv(getattr(config, "signal_allowed_senders", ""))
|
||||
return SignalChannel(SignalConfig(
|
||||
phone_number=getattr(config, "signal_phone_number", ""),
|
||||
cli_path=getattr(config, "signal_cli_path", "signal-cli"),
|
||||
config_dir=getattr(config, "signal_config_dir", "") or None,
|
||||
rpc_port=int(getattr(config, "signal_rpc_port", 7583)),
|
||||
allowed_senders=allowed,
|
||||
))
|
||||
|
||||
|
||||
register_channel("signal", create_from_config)
|
||||
@@ -0,0 +1,442 @@
|
||||
"""Signal channel implementation via signal-cli JSON RPC.
|
||||
|
||||
Pure Python — communicates with signal-cli daemon over TCP socket.
|
||||
|
||||
Architecture:
|
||||
1. signal-cli must be running in JSON RPC mode:
|
||||
signal-cli -u +NUMBER daemon --socket localhost:7583
|
||||
2. We connect via TCP, send JSON RPC requests, receive events
|
||||
3. Inbound: listen for "receive" method notifications
|
||||
4. Outbound: call "send" method via JSON RPC
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import subprocess
|
||||
from collections import deque
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..capabilities import SIGNAL as SIGNAL_CAPS
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SignalConfig(BaseChannelConfig):
|
||||
"""Configuration for Signal channel."""
|
||||
phone_number: str = ""
|
||||
cli_path: str = "signal-cli"
|
||||
config_dir: str | None = None
|
||||
rpc_port: int = 7583
|
||||
text_chunk_limit: int = 4096
|
||||
|
||||
|
||||
class SignalChannel(Channel):
|
||||
capabilities = SIGNAL_CAPS
|
||||
"""Signal channel using signal-cli JSON RPC.
|
||||
|
||||
No public IP needed — local TCP socket connection.
|
||||
Requires signal-cli to be installed and registered.
|
||||
"""
|
||||
|
||||
name = "signal"
|
||||
_non_retryable_patterns = ("unregistered", "auth")
|
||||
|
||||
def __init__(self, config: SignalConfig):
|
||||
super().__init__(config)
|
||||
self._reader: asyncio.StreamReader | None = None
|
||||
self._writer: asyncio.StreamWriter | None = None
|
||||
self._rpc_id = 0
|
||||
self._daemon_proc = None
|
||||
# Cache message_id → sender for reaction targetAuthor (bounded)
|
||||
self._msg_senders: dict[str, str] = {}
|
||||
self._msg_senders_order: deque = deque(maxlen=200)
|
||||
|
||||
async def start(self) -> None:
|
||||
if not self.config.phone_number:
|
||||
raise ChannelError("Signal phone_number is required")
|
||||
|
||||
# Try to start signal-cli daemon if not already running
|
||||
await self._ensure_daemon()
|
||||
|
||||
# Connect to JSON RPC socket
|
||||
await self._connect()
|
||||
|
||||
self._running = True
|
||||
logger.info(f"Signal channel started (phone: {self.config.phone_number})")
|
||||
|
||||
# Listen for incoming messages in background task
|
||||
# (start() must return so that run() can iterate receive())
|
||||
self._listen_task = asyncio.create_task(self._listen_loop())
|
||||
|
||||
async def _ensure_daemon(self) -> None:
|
||||
"""Start signal-cli daemon if not already running."""
|
||||
try:
|
||||
reader, writer = await asyncio.wait_for(
|
||||
asyncio.open_connection("localhost", self.config.rpc_port),
|
||||
timeout=2,
|
||||
)
|
||||
writer.close()
|
||||
await writer.wait_closed()
|
||||
logger.info("signal-cli daemon already running")
|
||||
return
|
||||
except (ConnectionRefusedError, asyncio.TimeoutError, OSError):
|
||||
pass
|
||||
|
||||
# Start daemon
|
||||
cmd = [self.config.cli_path, "-u", self.config.phone_number]
|
||||
if self.config.config_dir:
|
||||
cmd.extend(["--config", self.config.config_dir])
|
||||
cmd.extend(["daemon", "--tcp",
|
||||
f"localhost:{self.config.rpc_port}", "--no-receive-stdout"])
|
||||
|
||||
logger.info(f"Starting signal-cli daemon: {' '.join(cmd)}")
|
||||
try:
|
||||
self._daemon_proc = subprocess.Popen(
|
||||
cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||||
)
|
||||
except FileNotFoundError:
|
||||
raise ChannelError(
|
||||
f"signal-cli not found at '{self.config.cli_path}'. "
|
||||
"Install: https://github.com/AsamK/signal-cli"
|
||||
)
|
||||
|
||||
# Wait for daemon to be ready
|
||||
for _ in range(30):
|
||||
await asyncio.sleep(1)
|
||||
try:
|
||||
reader, writer = await asyncio.open_connection(
|
||||
"localhost", self.config.rpc_port,
|
||||
)
|
||||
writer.close()
|
||||
await writer.wait_closed()
|
||||
logger.info("signal-cli daemon started")
|
||||
return
|
||||
except (ConnectionRefusedError, OSError):
|
||||
continue
|
||||
|
||||
raise ChannelError("signal-cli daemon failed to start within 30s")
|
||||
|
||||
async def _connect(self) -> None:
|
||||
"""Connect to signal-cli JSON RPC socket."""
|
||||
try:
|
||||
self._reader, self._writer = await asyncio.open_connection(
|
||||
"localhost", self.config.rpc_port,
|
||||
)
|
||||
except Exception as e:
|
||||
raise ChannelError(f"Cannot connect to signal-cli: {e}")
|
||||
|
||||
async def _listen_loop(self) -> None:
|
||||
"""Listen for incoming JSON RPC notifications."""
|
||||
while self._running and self._reader:
|
||||
try:
|
||||
line = await self._reader.readline()
|
||||
if not line:
|
||||
break
|
||||
data = json.loads(line.decode())
|
||||
await self._handle_rpc(data)
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
except Exception as e:
|
||||
logger.error(f"Signal listen error: {e}")
|
||||
# Reconnect
|
||||
if self._running:
|
||||
await asyncio.sleep(2)
|
||||
try:
|
||||
await self._connect()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def _handle_rpc(self, data: dict) -> None:
|
||||
"""Handle a JSON RPC message from signal-cli."""
|
||||
method = data.get("method", "")
|
||||
|
||||
if method != "receive":
|
||||
return
|
||||
|
||||
params = data.get("params", {})
|
||||
envelope = params.get("envelope", {})
|
||||
source = envelope.get("source") or envelope.get("sourceUuid") or ""
|
||||
source_number = envelope.get("sourceNumber") or source
|
||||
source_name = envelope.get("sourceName") or ""
|
||||
timestamp = envelope.get("timestamp", 0)
|
||||
|
||||
# Ignore messages from self
|
||||
if source_number == self.config.phone_number or source == self.config.phone_number:
|
||||
logger.debug("Ignoring message from self")
|
||||
return
|
||||
|
||||
# Data message (text)
|
||||
data_msg = envelope.get("dataMessage", {})
|
||||
if data_msg:
|
||||
text = data_msg.get("message", "")
|
||||
group_info = data_msg.get("groupInfo", {})
|
||||
is_group = bool(group_info)
|
||||
chat_id = group_info.get("groupId", source_number) if is_group else source_number
|
||||
msg_ts = data_msg.get("timestamp", timestamp)
|
||||
|
||||
media_paths: list[str] = []
|
||||
annotations: list[str] = []
|
||||
_VOICE_TYPES = {"audio/aac", "audio/ogg", "audio/mp4", "audio/mpeg", "audio/opus"}
|
||||
attachments = data_msg.get("attachments", [])
|
||||
for att in attachments:
|
||||
att_size = att.get("size", 0)
|
||||
att_name = att.get("filename", "attachment")
|
||||
att_file = att.get("file") # signal-cli provides local path
|
||||
content_type = att.get("contentType", "")
|
||||
is_voice = content_type in _VOICE_TYPES or att.get("voiceNote", False)
|
||||
media_label = "voice" if is_voice else "attachment"
|
||||
if att_file:
|
||||
from pathlib import Path as _Path
|
||||
att_path = _Path(att_file)
|
||||
if att_path.exists():
|
||||
from ..base import MAX_ATTACHMENT_BYTES
|
||||
if att_path.stat().st_size > MAX_ATTACHMENT_BYTES:
|
||||
annotations.append(f"[{media_label}: {att_name} - too large ({att_path.stat().st_size} bytes)]")
|
||||
else:
|
||||
local = self._media_path(f"signal_{att_name}")
|
||||
import shutil
|
||||
shutil.copy2(str(att_path), str(local))
|
||||
media_paths.append(str(local))
|
||||
annotations.append(f"[{media_label}: {local}]")
|
||||
else:
|
||||
annotations.append(f"[{media_label}: {att_name} - file not found]")
|
||||
elif att_size:
|
||||
too_large = self._check_attachment_size(att_size, att_name)
|
||||
if too_large:
|
||||
annotations.append(too_large)
|
||||
else:
|
||||
annotations.append(f"[{media_label}: {att_name}]")
|
||||
|
||||
if not text and not media_paths and not annotations:
|
||||
if not attachments:
|
||||
return
|
||||
# Had attachments but none downloaded successfully
|
||||
if not annotations:
|
||||
text = "[attachment]"
|
||||
|
||||
try:
|
||||
ts = datetime.fromtimestamp(msg_ts / 1000) if msg_ts else datetime.now()
|
||||
except (ValueError, TypeError, OSError):
|
||||
ts = datetime.now()
|
||||
|
||||
was_mentioned = not is_group # DMs always pass
|
||||
if is_group:
|
||||
mentions = data_msg.get("mentions", [])
|
||||
for m in mentions:
|
||||
if m.get("uuid") == self.config.phone_number or m.get("number") == self.config.phone_number:
|
||||
was_mentioned = True
|
||||
break
|
||||
|
||||
# Cache message_id → sender for reaction targetAuthor
|
||||
self._cache_msg_sender(str(msg_ts), source_number)
|
||||
|
||||
logger.info("Signal message from %s: %s", source_number, text[:50] if text else "[media]")
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=source_number,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
content_annotations=annotations,
|
||||
media_files=media_paths,
|
||||
timestamp=ts,
|
||||
message_id=str(msg_ts),
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"source_name": source_name,
|
||||
"sender_id": source_number,
|
||||
"backend": "signal",
|
||||
},
|
||||
))
|
||||
|
||||
# ── Typing indicator ────────────────────────────────────────────
|
||||
|
||||
async def _send_typing_action(self, chat_id: str) -> None:
|
||||
"""Send typing indicator via signal-cli JSON RPC."""
|
||||
params: dict[str, Any] = {
|
||||
"account": self.config.phone_number,
|
||||
}
|
||||
if self._is_group_id(chat_id):
|
||||
params["groupId"] = chat_id
|
||||
else:
|
||||
params["recipient"] = [chat_id]
|
||||
try:
|
||||
await self._rpc_call("sendTyping", params)
|
||||
except Exception:
|
||||
pass # typing indicator is best-effort
|
||||
|
||||
# ── ACK reaction ─────────────────────────────────────────────
|
||||
|
||||
def _cache_msg_sender(self, message_id: str, sender: str) -> None:
|
||||
"""Store message_id → sender mapping for reaction targetAuthor."""
|
||||
if len(self._msg_senders) >= 200:
|
||||
oldest = self._msg_senders_order.popleft()
|
||||
self._msg_senders.pop(oldest, None)
|
||||
self._msg_senders[message_id] = sender
|
||||
self._msg_senders_order.append(message_id)
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
"""Send an acknowledgment reaction via signal-cli sendReaction."""
|
||||
target_author = self._msg_senders.get(message_id, "")
|
||||
if not target_author:
|
||||
return # cannot send reaction without knowing the original sender
|
||||
try:
|
||||
params: dict[str, Any] = {
|
||||
"account": self.config.phone_number,
|
||||
"emoji": emoji,
|
||||
"targetAuthor": target_author,
|
||||
"targetTimestamp": int(message_id),
|
||||
}
|
||||
if self._is_group_id(chat_id):
|
||||
params["groupId"] = chat_id
|
||||
else:
|
||||
params["recipient"] = [chat_id]
|
||||
await self._rpc_call("sendReaction", params)
|
||||
except Exception as e:
|
||||
logger.debug(f"Signal ack reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
"""Remove ACK reaction via signal-cli sendReaction --remove."""
|
||||
target_author = self._msg_senders.get(message_id, "")
|
||||
if not target_author:
|
||||
return
|
||||
try:
|
||||
params: dict[str, Any] = {
|
||||
"account": self.config.phone_number,
|
||||
"emoji": emoji,
|
||||
"targetAuthor": target_author,
|
||||
"targetTimestamp": int(message_id),
|
||||
"remove": True,
|
||||
}
|
||||
if self._is_group_id(chat_id):
|
||||
params["groupId"] = chat_id
|
||||
else:
|
||||
params["recipient"] = [chat_id]
|
||||
await self._rpc_call("sendReaction", params)
|
||||
except Exception as e:
|
||||
logger.debug(f"Signal remove ACK reaction failed: {e}")
|
||||
|
||||
# ── Send ──────────────────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _is_group_id(chat_id: str) -> bool:
|
||||
"""Return True if *chat_id* looks like a Signal group ID.
|
||||
|
||||
Group IDs are base64-encoded strings (e.g. ``"aB3d...=="``).
|
||||
Individual recipients are either phone numbers (``"+1234..."``)
|
||||
or UUIDs (``"817ab5e9-..."``) — neither of which is a group.
|
||||
"""
|
||||
return not chat_id.startswith("+") and "-" not in chat_id
|
||||
|
||||
def _is_ready(self) -> bool:
|
||||
return self._writer is not None and not self._writer.is_closing()
|
||||
|
||||
async def _rpc_call(self, method: str, params: dict) -> dict | None:
|
||||
"""Send a JSON RPC call to signal-cli."""
|
||||
if not self._writer:
|
||||
return None
|
||||
|
||||
self._rpc_id += 1
|
||||
request = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": self._rpc_id,
|
||||
"method": method,
|
||||
"params": params,
|
||||
}
|
||||
line = json.dumps(request) + "\n"
|
||||
self._writer.write(line.encode())
|
||||
await self._writer.drain()
|
||||
return None # We don't wait for response in this simple impl
|
||||
|
||||
async def _send_chunk(
|
||||
self, chat_id, formatted_text, raw_text, reply_to, metadata,
|
||||
):
|
||||
# Determine if group or individual
|
||||
params: dict[str, Any] = {
|
||||
"message": raw_text,
|
||||
"account": self.config.phone_number,
|
||||
}
|
||||
|
||||
if self._is_group_id(chat_id):
|
||||
params["groupId"] = chat_id
|
||||
else:
|
||||
params["recipient"] = [chat_id]
|
||||
|
||||
await self._rpc_call("send", params)
|
||||
|
||||
# ── Formatting ────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── Mention stripping ────────────────────────────────────────────
|
||||
|
||||
def _strip_mention(self, text: str) -> str:
|
||||
"""Strip bot mention from Signal messages.
|
||||
|
||||
Signal mentions are embedded as special objects that reference
|
||||
the phone number. The text contains a placeholder character (U+FFFC)
|
||||
at the mention position.
|
||||
"""
|
||||
phone = self.config.phone_number
|
||||
if phone:
|
||||
# Remove phone number if directly mentioned as text
|
||||
text = re.sub(rf"@?{re.escape(phone)}\s*", "", text).strip()
|
||||
# Remove Unicode Object Replacement Character used as mention placeholder
|
||||
text = text.replace("\uFFFC", "").strip()
|
||||
return text
|
||||
|
||||
# ── Media send ────────────────────────────────────────────────
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a media file via signal-cli JSON RPC.
|
||||
|
||||
Uses the "send" RPC method with the attachments parameter.
|
||||
"""
|
||||
chat_id = self._resolve_media_chat_id(recipient, metadata)
|
||||
params: dict[str, Any] = {
|
||||
"account": self.config.phone_number,
|
||||
"attachments": [file_path],
|
||||
}
|
||||
if caption:
|
||||
params["message"] = caption
|
||||
|
||||
if self._is_group_id(chat_id):
|
||||
params["groupId"] = chat_id
|
||||
else:
|
||||
params["recipient"] = [chat_id]
|
||||
|
||||
await self._rpc_call("send", params)
|
||||
return True
|
||||
|
||||
# ── Cleanup ───────────────────────────────────────────────────
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
if hasattr(self, "_listen_task") and self._listen_task:
|
||||
self._listen_task.cancel()
|
||||
self._listen_task = None
|
||||
if self._writer:
|
||||
self._writer.close()
|
||||
try:
|
||||
await self._writer.wait_closed()
|
||||
except Exception:
|
||||
pass
|
||||
self._writer = None
|
||||
self._reader = None
|
||||
if self._daemon_proc:
|
||||
self._daemon_proc.terminate()
|
||||
self._daemon_proc = None
|
||||
logger.info("Signal channel stopped")
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Signal credential validation."""
|
||||
import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_signal(
|
||||
phone_number: str,
|
||||
cli_path: str = "signal-cli",
|
||||
rpc_port: int = 7583,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate Signal setup by checking signal-cli availability."""
|
||||
import asyncio
|
||||
import subprocess
|
||||
|
||||
if not phone_number:
|
||||
return False, "phone_number is required"
|
||||
|
||||
# Check signal-cli binary
|
||||
loop = asyncio.get_event_loop()
|
||||
def _check():
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[cli_path, "--version"], capture_output=True, text=True, timeout=5,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
return True, f"signal-cli {result.stdout.strip()}"
|
||||
return False, "signal-cli returned error"
|
||||
except FileNotFoundError:
|
||||
return False, f"signal-cli not found at '{cli_path}'"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
|
||||
return await loop.run_in_executor(None, _check)
|
||||
@@ -0,0 +1,99 @@
|
||||
"""Signal channel server.
|
||||
|
||||
Standalone script to run the Signal channel with CLI options.
|
||||
|
||||
Usage:
|
||||
python -m EvoScientist.channels.signal.serve --phone-number NUMBER [OPTIONS]
|
||||
|
||||
Examples:
|
||||
# Basic usage
|
||||
python -m EvoScientist.channels.signal.serve --phone-number +1234567890
|
||||
|
||||
# With custom signal-cli path and allowed senders
|
||||
python -m EvoScientist.channels.signal.serve --phone-number +1234567890 --cli-path /usr/local/bin/signal-cli --allow +9876543210
|
||||
|
||||
# With agent and thinking
|
||||
python -m EvoScientist.channels.signal.serve --phone-number +1234567890 --agent --thinking
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import SignalChannel, SignalConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Signal channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--phone-number",
|
||||
required=True,
|
||||
help="Signal phone number (e.g. +1234567890)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cli-path",
|
||||
default="signal-cli",
|
||||
help="Path to signal-cli binary (default: signal-cli)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--config-dir",
|
||||
help="signal-cli config directory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rpc-port",
|
||||
type=int,
|
||||
default=7583,
|
||||
help="signal-cli JSON RPC port (default: 7583)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender (phone number). Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler (default: echo)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content as intermediate messages (requires --agent)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = SignalConfig(
|
||||
phone_number=args.phone_number,
|
||||
cli_path=args.cli_path,
|
||||
config_dir=args.config_dir,
|
||||
rpc_port=args.rpc_port,
|
||||
allowed_senders=set(args.allowed_senders) if args.allowed_senders else None,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = SignalChannel(config)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,19 @@
|
||||
from .channel import SlackChannel, SlackConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["SlackChannel", "SlackConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> SlackChannel:
|
||||
allowed = _parse_csv(config.slack_allowed_senders)
|
||||
channels = _parse_csv(config.slack_allowed_channels)
|
||||
return SlackChannel(SlackConfig(
|
||||
bot_token=config.slack_bot_token,
|
||||
app_token=config.slack_app_token,
|
||||
allowed_senders=allowed,
|
||||
allowed_channels=channels,
|
||||
proxy=getattr(config, 'slack_proxy', '') or None,
|
||||
))
|
||||
|
||||
|
||||
register_channel("slack", create_from_config)
|
||||
@@ -0,0 +1,291 @@
|
||||
"""Slack channel implementation using slack-sdk Socket Mode."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..capabilities import SLACK as SLACK_CAPS
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlackConfig(BaseChannelConfig):
|
||||
bot_token: str = ""
|
||||
app_token: str = ""
|
||||
text_chunk_limit: int = 4096
|
||||
|
||||
|
||||
class SlackChannel(Channel):
|
||||
"""Slack channel using slack-sdk Socket Mode."""
|
||||
|
||||
name = "slack"
|
||||
|
||||
capabilities = SLACK_CAPS
|
||||
_ready_attrs = ("_web_client",)
|
||||
_mention_pattern = r"<@{bot_id}>\s*"
|
||||
|
||||
def __init__(self, config: SlackConfig):
|
||||
super().__init__(config)
|
||||
self._socket_client = None
|
||||
self._web_client = None
|
||||
self._typing_message_ts: dict[str, str] = {}
|
||||
|
||||
async def start(self) -> None:
|
||||
try:
|
||||
from slack_sdk.web.async_client import AsyncWebClient
|
||||
from slack_sdk.socket_mode.aiohttp import SocketModeClient
|
||||
from slack_sdk.socket_mode.request import SocketModeRequest
|
||||
from slack_sdk.socket_mode.response import SocketModeResponse
|
||||
except ImportError:
|
||||
raise ChannelError(
|
||||
"slack-sdk or aiohttp not installed. "
|
||||
"Install with: pip install evoscientist[slack]"
|
||||
)
|
||||
|
||||
if not self.config.bot_token:
|
||||
raise ChannelError("Slack bot token is required")
|
||||
if not self.config.app_token:
|
||||
raise ChannelError(
|
||||
"Slack app token is required for Socket Mode "
|
||||
"(starts with xapp-)"
|
||||
)
|
||||
|
||||
self._web_client = AsyncWebClient(
|
||||
token=self.config.bot_token,
|
||||
proxy=self._get_proxy(),
|
||||
)
|
||||
|
||||
# Get bot user ID for filtering own messages
|
||||
try:
|
||||
auth = await asyncio.wait_for(
|
||||
self._web_client.auth_test(), timeout=15,
|
||||
)
|
||||
self._bot_user_id = auth["user_id"]
|
||||
except asyncio.TimeoutError:
|
||||
raise ChannelError(
|
||||
"Slack auth_test timed out — check network and bot token"
|
||||
)
|
||||
except Exception as e:
|
||||
raise ChannelError(f"Failed to authenticate Slack bot: {e}")
|
||||
|
||||
self._socket_client = SocketModeClient(
|
||||
app_token=self.config.app_token,
|
||||
web_client=self._web_client,
|
||||
)
|
||||
|
||||
async def _event_handler(
|
||||
client: SocketModeClient,
|
||||
req: SocketModeRequest,
|
||||
) -> None:
|
||||
# Acknowledge immediately
|
||||
resp = SocketModeResponse(envelope_id=req.envelope_id)
|
||||
await client.send_socket_mode_response(resp)
|
||||
|
||||
logger.debug(f"Slack socket event: type={req.type}")
|
||||
|
||||
if req.type == "events_api":
|
||||
event = req.payload.get("event", {})
|
||||
event_type = event.get("type", "")
|
||||
if event_type == "message" and "subtype" not in event:
|
||||
is_dm = event.get("channel_type") == "im"
|
||||
await self._on_message(
|
||||
event, is_group=not is_dm, was_mentioned=is_dm,
|
||||
)
|
||||
elif event_type == "app_mention":
|
||||
await self._on_message(
|
||||
event, is_group=True, was_mentioned=True,
|
||||
)
|
||||
|
||||
self._socket_client.socket_mode_request_listeners.append(
|
||||
_event_handler
|
||||
)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self._socket_client.connect(), timeout=30,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
raise ChannelError(
|
||||
"Slack Socket Mode connection timed out — "
|
||||
"check app token (must start with xapp-) and "
|
||||
"ensure Socket Mode is enabled in your Slack app settings"
|
||||
)
|
||||
self._running = True
|
||||
logger.info("Slack channel started (Socket Mode)")
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
if self._socket_client:
|
||||
await self._socket_client.close()
|
||||
logger.info("Slack channel stopped")
|
||||
|
||||
# ── Typing indicator (override base) ────────────────────────────
|
||||
|
||||
async def _send_typing_action(self, chat_id: str) -> None:
|
||||
"""Send typing indicator via Slack.
|
||||
|
||||
Slack's Web API and Socket Mode do not expose a dedicated
|
||||
typing-indicator endpoint for bot tokens. We approximate
|
||||
the experience by posting a short-lived status message that
|
||||
is deleted once the real reply is sent (handled by
|
||||
``stop_typing``). When the status post fails we silently
|
||||
fall back to no indicator.
|
||||
"""
|
||||
if not self._web_client:
|
||||
return
|
||||
try:
|
||||
resp = await self._web_client.chat_postMessage(
|
||||
channel=chat_id,
|
||||
text="\u2026", # "…" ellipsis as minimal typing hint
|
||||
)
|
||||
ts = resp.get("ts")
|
||||
if ts:
|
||||
self._typing_message_ts[chat_id] = ts
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def stop_typing(self, chat_id: str) -> None:
|
||||
"""Cancel typing loop and clean up the status message."""
|
||||
# Delete the ephemeral "…" message if we posted one
|
||||
ts = self._typing_message_ts.pop(chat_id, None)
|
||||
if ts and self._web_client:
|
||||
try:
|
||||
await self._web_client.chat_delete(channel=chat_id, ts=ts)
|
||||
except Exception:
|
||||
pass
|
||||
await super().stop_typing(chat_id)
|
||||
|
||||
# ── Send (template method overrides) ──────────────────────────
|
||||
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
kwargs = dict(channel=chat_id)
|
||||
# Always route to thread if thread_ts is present in metadata,
|
||||
# not just for the first chunk (reply_to is only set for chunk 0).
|
||||
if metadata:
|
||||
thread_ts = metadata.get("thread_ts")
|
||||
if thread_ts:
|
||||
kwargs["thread_ts"] = thread_ts
|
||||
|
||||
async def _send(text):
|
||||
await self._web_client.chat_postMessage(text=text, **kwargs)
|
||||
|
||||
await self._send_with_format_fallback(_send, formatted_text, raw_text)
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a media file through Slack."""
|
||||
channel_id = self._resolve_media_chat_id(recipient, metadata)
|
||||
await self._web_client.files_upload_v2(
|
||||
channel=channel_id,
|
||||
file=file_path,
|
||||
initial_comment=caption or None,
|
||||
)
|
||||
return True
|
||||
|
||||
def _get_bot_identifier(self) -> str | None:
|
||||
return getattr(self, "_bot_user_id", None)
|
||||
|
||||
# ── ACK Reactions ───────────────────────────────────────────────
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "eyes") -> None:
|
||||
"""Add an emoji reaction to acknowledge receipt."""
|
||||
if self._web_client and message_id:
|
||||
try:
|
||||
await self._web_client.reactions_add(
|
||||
channel=chat_id, timestamp=message_id, name=emoji,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Slack ACK reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "eyes") -> None:
|
||||
"""Remove the ACK reaction after replying."""
|
||||
if self._web_client and message_id:
|
||||
try:
|
||||
await self._web_client.reactions_remove(
|
||||
channel=chat_id, timestamp=message_id, name=emoji,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Slack remove ACK reaction failed: {e}")
|
||||
|
||||
async def _on_message(
|
||||
self,
|
||||
event: dict,
|
||||
*,
|
||||
is_group: bool = False,
|
||||
was_mentioned: bool = True,
|
||||
) -> None:
|
||||
"""Handle an incoming Slack message event."""
|
||||
user_id = event.get("user", "")
|
||||
|
||||
# Skip bot's own messages
|
||||
if user_id == getattr(self, "_bot_user_id", None):
|
||||
logger.debug("Skipping own bot message")
|
||||
return
|
||||
|
||||
# Skip bot messages (e.g. from other bots)
|
||||
if event.get("bot_id"):
|
||||
logger.debug(f"Skipping bot message from bot_id={event.get('bot_id')}")
|
||||
return
|
||||
|
||||
channel_id = event.get("channel", "")
|
||||
|
||||
text = event.get("text", "")
|
||||
|
||||
annotations: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
# Handle file attachments
|
||||
if self.config.include_attachments:
|
||||
files = event.get("files", [])
|
||||
for file_info in files:
|
||||
file_size = file_info.get("size", 0)
|
||||
filename = file_info.get("name", "unknown")
|
||||
|
||||
url = file_info.get("url_private_download") or file_info.get(
|
||||
"url_private"
|
||||
)
|
||||
if url and self._web_client:
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.config.bot_token}"
|
||||
}
|
||||
local_path, annotation = await self._download_attachment(
|
||||
url, f"{file_info.get('id', 'unknown')}_{filename}",
|
||||
headers=headers,
|
||||
file_size=file_size,
|
||||
)
|
||||
if local_path:
|
||||
media_paths.append(local_path)
|
||||
if annotation:
|
||||
annotations.append(annotation)
|
||||
|
||||
ts = event.get("ts", "")
|
||||
thread_ts = event.get("thread_ts") or ts
|
||||
try:
|
||||
timestamp = datetime.fromtimestamp(float(ts)) if ts else datetime.now()
|
||||
except (ValueError, TypeError):
|
||||
timestamp = datetime.now()
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=user_id,
|
||||
chat_id=channel_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=ts,
|
||||
metadata={"chat_id": channel_id, "thread_ts": thread_ts},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
))
|
||||
logger.info(
|
||||
f"Slack message queued: sender={user_id}, "
|
||||
f"channel={channel_id}, content={text[:50]}"
|
||||
)
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Slack bot token validation."""
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_slack_tokens(
|
||||
bot_token: str,
|
||||
app_token: str | None = None,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate Slack bot token via the auth.test API.
|
||||
|
||||
Optionally checks the app-level token format (must start with ``xapp-``).
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, message).
|
||||
"""
|
||||
if not bot_token:
|
||||
return False, "No bot token provided"
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return False, "httpx not installed"
|
||||
|
||||
# Validate bot token via auth.test
|
||||
url = "https://slack.com/api/auth.test"
|
||||
headers = {"Authorization": f"Bearer {bot_token}"}
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
resp = await client.post(url, headers=headers, timeout=10)
|
||||
data = resp.json()
|
||||
if not data.get("ok"):
|
||||
error = data.get("error", "unknown error")
|
||||
return False, f"Invalid bot token: {error}"
|
||||
bot_name = data.get("user", "unknown")
|
||||
team = data.get("team", "unknown")
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
|
||||
# Optionally validate app token format
|
||||
if app_token:
|
||||
if not app_token.startswith("xapp-"):
|
||||
return False, "App token must start with 'xapp-'"
|
||||
|
||||
return True, f"Bot: {bot_name} (team: {team})"
|
||||
@@ -0,0 +1,94 @@
|
||||
"""Slack channel server.
|
||||
|
||||
Standalone script to run the Slack channel with CLI options.
|
||||
|
||||
Usage:
|
||||
python -m EvoScientist.channels.slack.serve --bot-token TOKEN --app-token TOKEN [OPTIONS]
|
||||
|
||||
Examples:
|
||||
# Allow all senders (default)
|
||||
python -m EvoScientist.channels.slack.serve --bot-token xoxb-... --app-token xapp-...
|
||||
|
||||
# Only allow specific senders and channels
|
||||
python -m EvoScientist.channels.slack.serve --bot-token xoxb-... --app-token xapp-... --allow U123 --allow-channel C456
|
||||
|
||||
# With agent and thinking
|
||||
python -m EvoScientist.channels.slack.serve --bot-token xoxb-... --app-token xapp-... --agent --thinking
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import SlackChannel, SlackConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Slack channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bot-token",
|
||||
required=True,
|
||||
help="Slack bot token (xoxb-...)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--app-token",
|
||||
required=True,
|
||||
help="Slack app-level token for Socket Mode (xapp-...)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender (Slack user ID). Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow-channel",
|
||||
action="append",
|
||||
dest="allowed_channels",
|
||||
help="Allowed channel ID. Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler (default: echo)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content as intermediate messages (requires --agent)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = SlackConfig(
|
||||
bot_token=args.bot_token,
|
||||
app_token=args.app_token,
|
||||
allowed_senders=set(args.allowed_senders) if args.allowed_senders else None,
|
||||
allowed_channels=set(args.allowed_channels) if args.allowed_channels else None,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = SlackChannel(config)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Shared standalone runner for channel servers.
|
||||
|
||||
Provides the channel-agnostic agent loop that any channel can use to
|
||||
run headless — consuming inbound messages from the bus, streaming
|
||||
agent events, and dispatching outbound replies.
|
||||
|
||||
Usage from a channel's ``main()``::
|
||||
|
||||
from EvoScientist.channels.standalone import run_standalone
|
||||
|
||||
channel = SomeChannel(config)
|
||||
bus = MessageBus()
|
||||
run_standalone(channel, bus, use_agent=True, send_thinking=True)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import signal
|
||||
|
||||
from .base import Channel
|
||||
from .bus import MessageBus
|
||||
from .bus.events import OutboundMessage
|
||||
from .consumer import InboundConsumer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def standalone_outbound_dispatcher(
|
||||
bus: MessageBus, channel: Channel,
|
||||
) -> None:
|
||||
"""Consume outbound messages from the bus and send via channel."""
|
||||
while True:
|
||||
try:
|
||||
msg: OutboundMessage = await asyncio.wait_for(
|
||||
bus.consume_outbound(), timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
|
||||
try:
|
||||
if msg.content:
|
||||
await channel.send(msg)
|
||||
except Exception as e:
|
||||
logger.error(f"Error sending outbound: {e}")
|
||||
|
||||
|
||||
async def _async_main(
|
||||
channel: Channel, bus: MessageBus,
|
||||
use_agent: bool, send_thinking: bool,
|
||||
) -> None:
|
||||
"""Async entry point — gather channel, dispatcher and optional consumer."""
|
||||
from .channel_manager import ChannelManager
|
||||
|
||||
channel.set_bus(bus)
|
||||
if send_thinking:
|
||||
channel.send_thinking = True
|
||||
|
||||
# Create a lightweight manager for the consumer to use
|
||||
manager = ChannelManager(bus)
|
||||
manager._channels[channel.name] = channel
|
||||
|
||||
await manager.start_health()
|
||||
|
||||
tasks = [channel.run()]
|
||||
|
||||
dispatcher = standalone_outbound_dispatcher(bus, channel)
|
||||
tasks.append(dispatcher)
|
||||
|
||||
consumer: InboundConsumer | None = None
|
||||
if use_agent:
|
||||
logger.info("Loading EvoScientist agent...")
|
||||
from ..EvoScientist import create_cli_agent
|
||||
agent = create_cli_agent()
|
||||
logger.info("Agent loaded")
|
||||
|
||||
consumer = InboundConsumer(
|
||||
bus=bus,
|
||||
manager=manager,
|
||||
agent=agent,
|
||||
thread_id="",
|
||||
send_thinking=send_thinking,
|
||||
)
|
||||
manager.register_health_provider("consumer", lambda: consumer.metrics)
|
||||
tasks.append(consumer.run())
|
||||
if send_thinking:
|
||||
logger.info("Thinking messages enabled")
|
||||
|
||||
async def _graceful_shutdown() -> None:
|
||||
"""Graceful shutdown: drain consumer, flush outbound, stop channel."""
|
||||
logger.info("Graceful shutdown initiated...")
|
||||
if consumer is not None:
|
||||
await consumer.stop()
|
||||
# Drain outbound queue before stopping the channel
|
||||
drained = 0
|
||||
while True:
|
||||
try:
|
||||
msg = bus.outbound.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
try:
|
||||
if msg.content:
|
||||
await asyncio.wait_for(channel.send(msg), timeout=5.0)
|
||||
drained += 1
|
||||
except Exception:
|
||||
pass
|
||||
if drained:
|
||||
logger.info(f"Outbound drain: {drained} sent")
|
||||
channel._running = False
|
||||
await channel.stop()
|
||||
await manager.stop_health()
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
for sig in (signal.SIGINT, signal.SIGTERM):
|
||||
loop.add_signal_handler(
|
||||
sig, lambda s=sig: asyncio.create_task(_graceful_shutdown()),
|
||||
)
|
||||
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
|
||||
def run_standalone(
|
||||
channel: Channel, bus: MessageBus, *,
|
||||
use_agent: bool = False, send_thinking: bool = False,
|
||||
) -> None:
|
||||
"""Synchronous entry point that spins up the standalone runner.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
channel:
|
||||
A fully-configured :class:`Channel` instance.
|
||||
bus:
|
||||
The :class:`MessageBus` shared with *channel*.
|
||||
use_agent:
|
||||
When ``True``, load the EvoScientist agent and process inbound
|
||||
messages through it.
|
||||
send_thinking:
|
||||
When ``True`` **and** *use_agent* is set, forward intermediate
|
||||
thinking messages to the channel.
|
||||
"""
|
||||
asyncio.run(_async_main(channel, bus, use_agent, send_thinking))
|
||||
@@ -0,0 +1,17 @@
|
||||
from .channel import TelegramChannel, TelegramConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["TelegramChannel", "TelegramConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> TelegramChannel:
|
||||
allowed = _parse_csv(config.telegram_allowed_senders)
|
||||
proxy = config.telegram_proxy if config.telegram_proxy else None
|
||||
return TelegramChannel(TelegramConfig(
|
||||
bot_token=config.telegram_bot_token,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
))
|
||||
|
||||
|
||||
register_channel("telegram", create_from_config)
|
||||
@@ -0,0 +1,289 @@
|
||||
"""Telegram channel implementation using python-telegram-bot."""
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError, IMAGE_EXTS, VIDEO_EXTS, AUDIO_EXTS
|
||||
from ..capabilities import TELEGRAM as TELEGRAM_CAPS
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TelegramConfig(BaseChannelConfig):
|
||||
bot_token: str = ""
|
||||
text_chunk_limit: int = 4096
|
||||
|
||||
|
||||
class TelegramChannel(Channel):
|
||||
"""Telegram channel using python-telegram-bot with long polling."""
|
||||
|
||||
name = "telegram"
|
||||
|
||||
capabilities = TELEGRAM_CAPS
|
||||
_typing_interval: float = 4.0
|
||||
_ready_attrs = ("_app",)
|
||||
_non_retryable_patterns = ("parse", "can't parse")
|
||||
_mention_pattern = r"(?i)@{bot_id}\s*"
|
||||
|
||||
def __init__(self, config: TelegramConfig):
|
||||
super().__init__(config)
|
||||
self._app = None
|
||||
self._bot_username: str = ""
|
||||
|
||||
async def start(self) -> None:
|
||||
if not self.config.bot_token:
|
||||
raise ChannelError("Telegram bot token is required")
|
||||
|
||||
try:
|
||||
from telegram.ext import (
|
||||
ApplicationBuilder,
|
||||
MessageHandler,
|
||||
filters,
|
||||
)
|
||||
except ImportError:
|
||||
raise ChannelError(
|
||||
"python-telegram-bot not installed. "
|
||||
"Install with: pip install evoscientist[telegram]"
|
||||
)
|
||||
|
||||
builder = ApplicationBuilder().token(self.config.bot_token)
|
||||
if self.config.proxy:
|
||||
builder = builder.proxy(self.config.proxy).get_updates_proxy(self.config.proxy)
|
||||
self._app = builder.build()
|
||||
|
||||
# Accept text and media message types
|
||||
media_filter = filters.TEXT
|
||||
if self.config.include_attachments:
|
||||
media_filter = (
|
||||
filters.TEXT
|
||||
| filters.PHOTO
|
||||
| filters.VOICE
|
||||
| filters.AUDIO
|
||||
| filters.Document.ALL
|
||||
| filters.VIDEO
|
||||
| filters.Sticker.ALL
|
||||
| filters.LOCATION
|
||||
)
|
||||
|
||||
self._app.add_handler(
|
||||
MessageHandler(media_filter & ~filters.COMMAND, self._on_message)
|
||||
)
|
||||
|
||||
await self._app.initialize()
|
||||
# Cache bot username for @mention detection in groups
|
||||
bot_info = await self._app.bot.get_me()
|
||||
self._bot_username = (bot_info.username or "").lower()
|
||||
await self._app.start()
|
||||
await self._app.updater.start_polling(drop_pending_updates=True)
|
||||
self._running = True
|
||||
logger.info("Telegram channel started (polling)")
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
if self._app:
|
||||
if self._app.updater and self._app.updater.running:
|
||||
await self._app.updater.stop()
|
||||
await self._app.stop()
|
||||
await self._app.shutdown()
|
||||
logger.info("Telegram channel stopped")
|
||||
|
||||
# ── Typing indicator (override base) ────────────────────────────
|
||||
|
||||
async def _send_typing_action(self, chat_id: str) -> None:
|
||||
"""Send typing action via Telegram Bot API."""
|
||||
if self._app:
|
||||
await self._app.bot.send_chat_action(
|
||||
chat_id=int(chat_id), action="typing",
|
||||
)
|
||||
|
||||
# ── Send (template method overrides) ──────────────────────────
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
reply_id = int(reply_to) if reply_to else None
|
||||
|
||||
async def _send(text):
|
||||
await self._app.bot.send_message(
|
||||
chat_id=int(chat_id), text=text,
|
||||
parse_mode="HTML" if text == formatted_text else None,
|
||||
reply_to_message_id=reply_id,
|
||||
)
|
||||
|
||||
await self._send_with_format_fallback(_send, formatted_text, raw_text)
|
||||
|
||||
_MEDIA_SENDERS = {
|
||||
IMAGE_EXTS: ("send_photo", "photo"),
|
||||
VIDEO_EXTS: ("send_video", "video"),
|
||||
AUDIO_EXTS: ("send_audio", "audio"),
|
||||
}
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a media file through Telegram."""
|
||||
chat_id = int(self._resolve_media_chat_id(recipient, metadata))
|
||||
cap = caption or None
|
||||
ext = Path(file_path).suffix.lower()
|
||||
for exts, (method, param) in self._MEDIA_SENDERS.items():
|
||||
if ext in exts:
|
||||
await getattr(self._app.bot, method)(
|
||||
chat_id=chat_id, caption=cap, **{param: file_path},
|
||||
)
|
||||
return True
|
||||
await self._app.bot.send_document(
|
||||
chat_id=chat_id, document=file_path, caption=cap,
|
||||
)
|
||||
return True
|
||||
|
||||
def _get_bot_identifier(self) -> str | None:
|
||||
return self._bot_username or None
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
"""Send an acknowledgment reaction via Telegram."""
|
||||
if self._app:
|
||||
try:
|
||||
from telegram import ReactionTypeEmoji
|
||||
await self._app.bot.set_message_reaction(
|
||||
chat_id=int(chat_id),
|
||||
message_id=int(message_id),
|
||||
reaction=[ReactionTypeEmoji(emoji)],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Telegram ACK reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
"""Remove the ack reaction by setting empty reaction list."""
|
||||
if self._app:
|
||||
try:
|
||||
await self._app.bot.set_message_reaction(
|
||||
chat_id=int(chat_id),
|
||||
message_id=int(message_id),
|
||||
reaction=[],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Telegram remove ACK reaction failed: {e}")
|
||||
|
||||
async def _on_message(self, update, context) -> None:
|
||||
"""Handler callback for text, photos, voice, audio, documents, video."""
|
||||
if not update.message:
|
||||
return
|
||||
|
||||
message = update.message
|
||||
user_id = str(message.from_user.id)
|
||||
chat_id = str(message.chat_id)
|
||||
|
||||
# Detect group and mention status for centralized gating
|
||||
is_group = message.chat.type in ("group", "supergroup")
|
||||
was_mentioned = True # DM default
|
||||
if is_group and self._bot_username:
|
||||
text_check = (message.text or message.caption or "").lower()
|
||||
was_mentioned = f"@{self._bot_username}" in text_check
|
||||
|
||||
content_parts: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
# Text content
|
||||
if message.text:
|
||||
content_parts.append(message.text)
|
||||
if message.caption:
|
||||
content_parts.append(message.caption)
|
||||
|
||||
# Handle media files
|
||||
annotations: list[str] = []
|
||||
if self.config.include_attachments:
|
||||
media_file = None
|
||||
media_type = None
|
||||
|
||||
if message.photo:
|
||||
media_file = message.photo[-1] # Largest size
|
||||
media_type = "image"
|
||||
elif message.voice:
|
||||
media_file = message.voice
|
||||
media_type = "voice"
|
||||
elif message.audio:
|
||||
media_file = message.audio
|
||||
media_type = "audio"
|
||||
elif message.video:
|
||||
media_file = message.video
|
||||
media_type = "video"
|
||||
elif message.document:
|
||||
media_file = message.document
|
||||
media_type = "file"
|
||||
elif message.sticker:
|
||||
media_file = message.sticker
|
||||
media_type = "sticker"
|
||||
|
||||
# Location is not a downloadable file — handle separately
|
||||
if message.location and not media_file:
|
||||
loc = message.location
|
||||
annotations.append(
|
||||
f"[位置] ({loc.latitude}, {loc.longitude})"
|
||||
)
|
||||
|
||||
if media_file and self._app:
|
||||
file_size = getattr(media_file, 'file_size', 0) or 0
|
||||
too_large = self._check_attachment_size(file_size, media_type)
|
||||
if too_large:
|
||||
annotations.append(too_large)
|
||||
else:
|
||||
try:
|
||||
file = await self._app.bot.get_file(
|
||||
media_file.file_id,
|
||||
)
|
||||
ext = self._get_extension(
|
||||
media_type,
|
||||
getattr(media_file, 'mime_type', None),
|
||||
)
|
||||
file_path = self._media_path(
|
||||
f"{media_file.file_id[:16]}{ext}"
|
||||
)
|
||||
await file.download_to_drive(str(file_path))
|
||||
|
||||
media_paths.append(str(file_path))
|
||||
annotations.append(f"[{media_type}: {file_path}]")
|
||||
logger.debug(
|
||||
f"Downloaded {media_type} to {file_path}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to download media: {e}")
|
||||
annotations.append(
|
||||
f"[{media_type}: download failed]"
|
||||
)
|
||||
|
||||
text_content = "\n".join(content_parts) if content_parts else ""
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=user_id,
|
||||
chat_id=chat_id,
|
||||
text=text_content,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=message.date or datetime.now(),
|
||||
message_id=str(message.message_id),
|
||||
metadata={"chat_id": chat_id},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
))
|
||||
|
||||
_MIME_TO_EXT = {
|
||||
"image/jpeg": ".jpg", "image/png": ".png", "image/gif": ".gif",
|
||||
"image/webp": ".webp", "audio/ogg": ".ogg", "audio/mpeg": ".mp3",
|
||||
"audio/mp4": ".m4a", "video/mp4": ".mp4", "video/quicktime": ".mov",
|
||||
}
|
||||
_TYPE_TO_EXT = {
|
||||
"image": ".jpg", "voice": ".ogg", "audio": ".mp3",
|
||||
"video": ".mp4", "file": "", "sticker": ".webp",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _get_extension(media_type: str, mime_type: str | None) -> str:
|
||||
"""Get file extension based on media type and MIME type."""
|
||||
if mime_type and mime_type in TelegramChannel._MIME_TO_EXT:
|
||||
return TelegramChannel._MIME_TO_EXT[mime_type]
|
||||
return TelegramChannel._TYPE_TO_EXT.get(media_type, "")
|
||||
@@ -0,0 +1,32 @@
|
||||
"""Telegram bot token validation."""
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_telegram_token(token: str, proxy: str | None = None) -> tuple[bool, str]:
|
||||
"""Validate a Telegram bot token via the getMe API.
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, message).
|
||||
"""
|
||||
if not token:
|
||||
return False, "No token provided"
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return False, "httpx not installed"
|
||||
|
||||
url = f"https://api.telegram.org/bot{token}/getMe"
|
||||
try:
|
||||
async with httpx.AsyncClient(proxy=proxy) as client:
|
||||
resp = await client.get(url, timeout=10)
|
||||
data = resp.json()
|
||||
if data.get("ok"):
|
||||
username = data["result"].get("username", "unknown")
|
||||
return True, f"Bot: @{username}"
|
||||
return False, "Invalid token"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Telegram channel server.
|
||||
|
||||
Standalone script to run the Telegram channel with CLI options.
|
||||
|
||||
Usage:
|
||||
python -m EvoScientist.channels.telegram.serve --bot-token TOKEN [OPTIONS]
|
||||
|
||||
Examples:
|
||||
# Allow all senders (default)
|
||||
python -m EvoScientist.channels.telegram.serve --bot-token TOKEN
|
||||
|
||||
# Only allow specific senders
|
||||
python -m EvoScientist.channels.telegram.serve --bot-token TOKEN --allow 123456 --allow 789012
|
||||
|
||||
# With agent and thinking
|
||||
python -m EvoScientist.channels.telegram.serve --bot-token TOKEN --agent --thinking
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import TelegramChannel, TelegramConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Telegram channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bot-token",
|
||||
required=True,
|
||||
help="Telegram bot token from @BotFather",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender (Telegram user ID). Can be used multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler (default: echo)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content as intermediate messages (requires --agent)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
|
||||
config = TelegramConfig(
|
||||
bot_token=args.bot_token,
|
||||
allowed_senders=set(args.allowed_senders) if args.allowed_senders else None,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = TelegramChannel(config)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,69 @@
|
||||
"""WeChat channel implementations for EvoScientist.
|
||||
|
||||
Supports multiple WeChat backends:
|
||||
- **wecom**: 企业微信应用 (WeCom / WeChat Work) via official API
|
||||
— Most stable, pure HTTP, no third-party dependencies
|
||||
- **wechatmp**: 微信公众号 (WeChat Official Account) via official API
|
||||
— Pure HTTP webhook, suitable for public-facing bots
|
||||
|
||||
Both backends use httpx (already a core dependency) and receive messages
|
||||
via HTTP webhook, send replies via REST API.
|
||||
|
||||
Usage in config:
|
||||
channel_enabled = "wechat"
|
||||
wechat_backend = "wecom" # or "wechatmp"
|
||||
|
||||
# WeCom settings
|
||||
wechat_wecom_corp_id = "..."
|
||||
wechat_wecom_agent_id = "..."
|
||||
wechat_wecom_secret = "..."
|
||||
wechat_wecom_token = "..."
|
||||
wechat_wecom_encoding_aes_key = "..."
|
||||
wechat_webhook_port = 9001
|
||||
|
||||
# OR: Official Account settings
|
||||
wechat_mp_app_id = "..."
|
||||
wechat_mp_app_secret = "..."
|
||||
wechat_mp_token = "..."
|
||||
wechat_mp_encoding_aes_key = "..."
|
||||
wechat_webhook_port = 9001
|
||||
"""
|
||||
|
||||
from .channel import WeChatChannel, WeComConfig, WeChatMPConfig
|
||||
from ..channel_manager import register_channel, _parse_csv
|
||||
|
||||
__all__ = ["WeChatChannel", "WeComConfig", "WeChatMPConfig"]
|
||||
|
||||
|
||||
def create_from_config(config) -> WeChatChannel:
|
||||
backend = getattr(config, "wechat_backend", "wecom") or "wecom"
|
||||
allowed = _parse_csv(getattr(config, "wechat_allowed_senders", ""))
|
||||
proxy = getattr(config, "wechat_proxy", "") or None
|
||||
port = int(getattr(config, "wechat_webhook_port", 9001) or 9001)
|
||||
|
||||
if backend == "wechatmp":
|
||||
mp_config = WeChatMPConfig(
|
||||
app_id=getattr(config, "wechat_mp_app_id", ""),
|
||||
app_secret=getattr(config, "wechat_mp_app_secret", ""),
|
||||
token=getattr(config, "wechat_mp_token", ""),
|
||||
encoding_aes_key=getattr(config, "wechat_mp_encoding_aes_key", ""),
|
||||
webhook_port=port,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
)
|
||||
return WeChatChannel(mp_config, backend="wechatmp")
|
||||
else:
|
||||
wecom_config = WeComConfig(
|
||||
corp_id=getattr(config, "wechat_wecom_corp_id", ""),
|
||||
agent_id=getattr(config, "wechat_wecom_agent_id", ""),
|
||||
secret=getattr(config, "wechat_wecom_secret", ""),
|
||||
token=getattr(config, "wechat_wecom_token", ""),
|
||||
encoding_aes_key=getattr(config, "wechat_wecom_encoding_aes_key", ""),
|
||||
webhook_port=port,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
)
|
||||
return WeChatChannel(wecom_config, backend="wecom")
|
||||
|
||||
|
||||
register_channel("wechat", create_from_config)
|
||||
@@ -0,0 +1,807 @@
|
||||
"""WeChat channel implementation.
|
||||
|
||||
Supports two backends via a unified Channel interface:
|
||||
|
||||
1. **wecom** (企业微信应用): Corporate WeChat official API
|
||||
- Receives messages via HTTP callback (XML + optional AES encryption)
|
||||
- Sends replies via REST API (POST /cgi-bin/message/send)
|
||||
- Supports text, image, file, markdown messages
|
||||
- Token auto-refresh with 2-hour TTL
|
||||
|
||||
2. **wechatmp** (微信公众号): WeChat Official Account API
|
||||
- Receives messages via HTTP callback (XML + optional AES encryption)
|
||||
- Sends replies via REST API (POST /cgi-bin/message/custom/send)
|
||||
- Supports text, image, news messages
|
||||
|
||||
Both backends use httpx (already a core dependency) and aiohttp for
|
||||
webhook server — matching the Feishu channel pattern.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from aiohttp import web
|
||||
|
||||
from ..mixins import WebhookMixin, TokenMixin
|
||||
from ..base import Channel, RawIncoming, ChannelError
|
||||
from ..capabilities import WECHAT as WECHAT_CAPS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── Markdown → plain text (fallback for WeChat text messages) ────
|
||||
|
||||
def _strip_markdown(text: str) -> str:
|
||||
"""Strip Markdown formatting for plain-text WeChat messages."""
|
||||
# Remove code blocks
|
||||
text = re.sub(r"```[\s\S]*?```", lambda m: m.group(0).strip("`").strip(), text)
|
||||
# Remove inline code
|
||||
text = re.sub(r"`([^`]+)`", r"\1", text)
|
||||
# Remove bold
|
||||
text = re.sub(r"\*\*(.+?)\*\*", r"\1", text)
|
||||
# Remove italic
|
||||
text = re.sub(r"(?<!\w)_([^_]+?)_(?!\w)", r"\1", text)
|
||||
# Remove strikethrough
|
||||
text = re.sub(r"~~(.+?)~~", r"\1", text)
|
||||
# Convert links
|
||||
text = re.sub(r"\[([^\]]+)\]\(([^)]+)\)", r"\1(\2)", text)
|
||||
# Remove heading markers
|
||||
text = re.sub(r"^#{1,6}\s+", "", text, flags=re.MULTILINE)
|
||||
# Convert list items
|
||||
text = re.sub(r"^[\-\*]\s+", "• ", text, flags=re.MULTILINE)
|
||||
return text
|
||||
|
||||
|
||||
# ── Config dataclasses ───────────────────────────────────────────
|
||||
|
||||
@dataclass
|
||||
class WeComConfig:
|
||||
"""Configuration for WeCom (企业微信) backend."""
|
||||
corp_id: str = ""
|
||||
agent_id: str = ""
|
||||
secret: str = ""
|
||||
token: str = ""
|
||||
encoding_aes_key: str = ""
|
||||
webhook_port: int = 9001
|
||||
allowed_senders: set[str] | None = None
|
||||
allowed_channels: set[str] | None = None
|
||||
text_chunk_limit: int = 4096 # WeCom text limit
|
||||
proxy: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class WeChatMPConfig:
|
||||
"""Configuration for WeChat Official Account (公众号) backend."""
|
||||
app_id: str = ""
|
||||
app_secret: str = ""
|
||||
token: str = ""
|
||||
encoding_aes_key: str = ""
|
||||
webhook_port: int = 9001
|
||||
allowed_senders: set[str] | None = None
|
||||
allowed_channels: set[str] | None = None
|
||||
text_chunk_limit: int = 4096
|
||||
proxy: str | None = None
|
||||
|
||||
|
||||
# ── Unified WeChat Channel ───────────────────────────────────────
|
||||
|
||||
class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
capabilities = WECHAT_CAPS
|
||||
"""Unified WeChat channel supporting WeCom and Official Account backends.
|
||||
|
||||
Architecture follows the same pattern as FeishuChannel:
|
||||
- HTTP webhook server (aiohttp) for inbound messages
|
||||
- REST API calls (httpx) for outbound messages
|
||||
- Token auto-refresh
|
||||
"""
|
||||
|
||||
name = "wechat"
|
||||
_typing_interval: float = 5.0 # WeChat has no typing API, but keep for interface
|
||||
_ready_attrs = ("_http_client", "_access_token")
|
||||
_rate_limit_patterns = ("45009", "frequency", "freq")
|
||||
_rate_limit_delay = 2.0
|
||||
_mention_pattern = r"@\S+\s*"
|
||||
_mention_strip_count = 1
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: WeComConfig | WeChatMPConfig,
|
||||
backend: str = "wecom",
|
||||
):
|
||||
super().__init__(config)
|
||||
self._backend = backend
|
||||
self._access_token: str | None = None
|
||||
self._token_expires: float = 0
|
||||
self._runner = None
|
||||
self._site = None
|
||||
self._http_client = None
|
||||
self._crypto = None # WeChatCrypto instance (optional)
|
||||
|
||||
# ── Lifecycle ─────────────────────────────────────────────────
|
||||
|
||||
def _webhook_routes(self) -> list[tuple[str, str, Any]]:
|
||||
"""Return HTTP routes for the shared webhook server."""
|
||||
return [
|
||||
("GET", "/wechat/callback", self._handle_verify),
|
||||
("POST", "/wechat/callback", self._handle_message),
|
||||
]
|
||||
|
||||
async def start(self) -> None:
|
||||
try:
|
||||
from aiohttp import web
|
||||
import httpx # noqa: F401
|
||||
except ImportError:
|
||||
raise ChannelError(
|
||||
"aiohttp or httpx not installed. "
|
||||
"Install with: pip install aiohttp httpx"
|
||||
)
|
||||
|
||||
self._validate_config()
|
||||
|
||||
import httpx
|
||||
self._http_client = httpx.AsyncClient(
|
||||
timeout=15,
|
||||
proxy=self._get_proxy(),
|
||||
)
|
||||
|
||||
# Set up message encryption if configured
|
||||
if self.config.encoding_aes_key and self.config.token:
|
||||
from .crypto import WeChatCrypto
|
||||
app_id = self._get_app_id()
|
||||
self._crypto = WeChatCrypto(
|
||||
token=self.config.token,
|
||||
encoding_aes_key=self.config.encoding_aes_key,
|
||||
app_id=app_id,
|
||||
)
|
||||
|
||||
# Verify credentials by fetching initial token
|
||||
await self._refresh_token()
|
||||
|
||||
if not getattr(self, "_shared_webhook_server", None):
|
||||
app = web.Application()
|
||||
app.router.add_get("/wechat/callback", self._handle_verify)
|
||||
app.router.add_post("/wechat/callback", self._handle_message)
|
||||
|
||||
self._runner = web.AppRunner(app)
|
||||
await self._runner.setup()
|
||||
self._site = web.TCPSite(
|
||||
self._runner, "0.0.0.0", self.config.webhook_port,
|
||||
)
|
||||
await self._site.start()
|
||||
|
||||
self._running = True
|
||||
logger.info(
|
||||
f"WeChat channel started "
|
||||
f"(backend={self._backend}, "
|
||||
f"webhook on port {self.config.webhook_port})"
|
||||
)
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
if self._site:
|
||||
await self._site.stop()
|
||||
if self._runner:
|
||||
await self._runner.cleanup()
|
||||
if self._http_client:
|
||||
await self._http_client.aclose()
|
||||
self._http_client = None
|
||||
self._access_token = None
|
||||
logger.info("WeChat channel stopped")
|
||||
|
||||
def _validate_config(self) -> None:
|
||||
"""Validate required config fields based on backend."""
|
||||
if self._backend == "wecom":
|
||||
cfg = self.config
|
||||
if not cfg.corp_id:
|
||||
raise ChannelError("WeCom corp_id is required")
|
||||
if not cfg.secret:
|
||||
raise ChannelError("WeCom secret is required")
|
||||
if not cfg.agent_id:
|
||||
raise ChannelError("WeCom agent_id is required")
|
||||
elif self._backend == "wechatmp":
|
||||
cfg = self.config
|
||||
if not cfg.app_id:
|
||||
raise ChannelError("WeChat MP app_id is required")
|
||||
if not cfg.app_secret:
|
||||
raise ChannelError("WeChat MP app_secret is required")
|
||||
|
||||
def _get_app_id(self) -> str:
|
||||
"""Return the app identifier for crypto operations."""
|
||||
if self._backend == "wecom":
|
||||
return self.config.corp_id
|
||||
return self.config.app_id
|
||||
|
||||
# ── Token management ──────────────────────────────────────────
|
||||
|
||||
async def _refresh_token(self) -> None:
|
||||
"""Fetch or refresh the access_token."""
|
||||
if self._backend == "wecom":
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/gettoken"
|
||||
f"?corpid={self.config.corp_id}"
|
||||
f"&corpsecret={self.config.secret}"
|
||||
)
|
||||
else:
|
||||
url = (
|
||||
f"https://api.weixin.qq.com/cgi-bin/token"
|
||||
f"?grant_type=client_credential"
|
||||
f"&appid={self.config.app_id}"
|
||||
f"&secret={self.config.app_secret}"
|
||||
)
|
||||
|
||||
try:
|
||||
resp = await self._http_client.get(url)
|
||||
data = resp.json()
|
||||
except Exception as e:
|
||||
raise ChannelError(f"Failed to get WeChat access token: {e}")
|
||||
|
||||
if data.get("errcode", 0) != 0:
|
||||
raise ChannelError(
|
||||
f"WeChat auth error ({data.get('errcode')}): "
|
||||
f"{data.get('errmsg', 'unknown')}"
|
||||
)
|
||||
|
||||
self._access_token = data["access_token"]
|
||||
expire = data.get("expires_in", 7200)
|
||||
# Refresh 5 minutes before expiry
|
||||
self._token_expires = time.monotonic() + expire - 300
|
||||
logger.debug(f"WeChat token refreshed, expires in {expire}s")
|
||||
|
||||
async def _ensure_token(self) -> str:
|
||||
"""Return a valid access token, refreshing if needed."""
|
||||
if not self._access_token or time.monotonic() >= self._token_expires:
|
||||
await self._refresh_token()
|
||||
return self._access_token
|
||||
|
||||
# ── Signature verification (GET callback) ─────────────────────
|
||||
|
||||
async def _handle_verify(self, request) -> "web.Response":
|
||||
"""Handle GET /wechat/callback for URL verification.
|
||||
|
||||
WeChat/WeCom sends: msg_signature, timestamp, nonce, echostr
|
||||
We decrypt echostr (encrypted mode) or verify signature (plain mode)
|
||||
and return the plain echostr.
|
||||
"""
|
||||
from aiohttp import web
|
||||
|
||||
signature = request.query.get("msg_signature") or request.query.get("signature", "")
|
||||
timestamp = request.query.get("timestamp", "")
|
||||
nonce = request.query.get("nonce", "")
|
||||
echostr = request.query.get("echostr", "")
|
||||
|
||||
logger.info(f"Verify request received: timestamp={timestamp}")
|
||||
|
||||
if not echostr:
|
||||
return web.Response(status=400, text="missing echostr")
|
||||
|
||||
# Encrypted mode: WeCom sends msg_signature and encrypted echostr
|
||||
if self._crypto and request.query.get("msg_signature"):
|
||||
# Verify signature first
|
||||
sig_ok = self._crypto.verify_signature(signature, timestamp, nonce, echostr)
|
||||
if not sig_ok:
|
||||
logger.warning("WeChat verify: signature mismatch")
|
||||
# Try to decrypt regardless — the decrypted echostr must be returned
|
||||
try:
|
||||
plain_echostr, _ = self._crypto.decrypt(echostr)
|
||||
logger.info("WeChat verify: echostr decrypted successfully")
|
||||
return web.Response(text=plain_echostr)
|
||||
except Exception as e:
|
||||
logger.error(f"WeChat verify: echostr decrypt failed: {e}")
|
||||
return web.Response(status=500)
|
||||
else:
|
||||
# Plain mode verification
|
||||
token = self.config.token
|
||||
if token:
|
||||
parts = sorted([token, timestamp, nonce])
|
||||
expected = hashlib.sha1("".join(parts).encode()).hexdigest()
|
||||
if expected != signature:
|
||||
logger.warning("WeChat verify: signature mismatch (plain)")
|
||||
return web.Response(status=403)
|
||||
return web.Response(text=echostr)
|
||||
|
||||
# ── Inbound message handling (POST callback) ──────────────────
|
||||
|
||||
async def _handle_message(self, request) -> "web.Response":
|
||||
"""Handle POST /wechat/callback for incoming messages."""
|
||||
from aiohttp import web
|
||||
from .crypto import parse_xml
|
||||
|
||||
try:
|
||||
body = await request.text()
|
||||
except Exception:
|
||||
return web.Response(status=400)
|
||||
|
||||
logger.info(f"WeChat callback POST received, body length={len(body)}")
|
||||
xml_data = parse_xml(body)
|
||||
|
||||
# If encrypted, decrypt first
|
||||
encrypt = xml_data.get("Encrypt", "")
|
||||
if encrypt and self._crypto:
|
||||
signature = request.query.get("msg_signature", "")
|
||||
timestamp = request.query.get("timestamp", "")
|
||||
nonce = request.query.get("nonce", "")
|
||||
|
||||
if not self._crypto.verify_signature(signature, timestamp, nonce, encrypt):
|
||||
logger.warning("WeChat message signature mismatch")
|
||||
return web.Response(status=403)
|
||||
|
||||
try:
|
||||
decrypted_xml, from_id = self._crypto.decrypt(encrypt)
|
||||
xml_data = parse_xml(decrypted_xml)
|
||||
except Exception as e:
|
||||
logger.error(f"WeChat decrypt failed: {e}")
|
||||
return web.Response(status=500)
|
||||
|
||||
# Process message asynchronously — WeCom requires a response within
|
||||
# 5 seconds, but media downloads can take much longer. Return
|
||||
# "success" immediately and handle the message in the background.
|
||||
asyncio.create_task(self._safe_process_message(xml_data))
|
||||
|
||||
return web.Response(text="success")
|
||||
|
||||
async def _safe_process_message(self, xml_data: dict[str, str]) -> None:
|
||||
"""Wrapper that catches exceptions so fire-and-forget tasks don't leak."""
|
||||
try:
|
||||
await self._process_message(xml_data)
|
||||
except Exception:
|
||||
logger.exception("Error processing WeChat message")
|
||||
|
||||
async def _process_message(self, xml_data: dict[str, str]) -> None:
|
||||
"""Process a parsed XML message from WeChat/WeCom callback."""
|
||||
msg_type = xml_data.get("MsgType", "")
|
||||
from_user = xml_data.get("FromUserName", "")
|
||||
to_user = xml_data.get("ToUserName", "")
|
||||
content = xml_data.get("Content", "")
|
||||
msg_id = xml_data.get("MsgId", "")
|
||||
create_time = xml_data.get("CreateTime", "")
|
||||
|
||||
logger.info(f"WeChat message received: type={msg_type}, from={from_user}, id={msg_id}, keys={list(xml_data.keys())}")
|
||||
|
||||
if not from_user:
|
||||
return
|
||||
|
||||
# Determine chat_id
|
||||
# For WeCom: FromUserName is the user's UserID
|
||||
# For MP: FromUserName is the user's OpenID
|
||||
chat_id = from_user
|
||||
|
||||
# Group chat detection
|
||||
is_group = False
|
||||
was_mentioned = True # Default: treat as mentioned (DMs)
|
||||
|
||||
# WeCom group detection: ChatId field indicates a group message
|
||||
if self._backend == "wecom":
|
||||
group_chat_id = xml_data.get("ChatId", "")
|
||||
if group_chat_id:
|
||||
is_group = True
|
||||
chat_id = group_chat_id
|
||||
# WeCom sets MsgType=event with Event=sys when bot is @mentioned,
|
||||
# but for text messages we check the XML AtUserList field
|
||||
at_user_list = xml_data.get("AtUserList", "")
|
||||
was_mentioned = bool(at_user_list)
|
||||
|
||||
# Handle different message types
|
||||
text = ""
|
||||
annotations: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
if msg_type == "text":
|
||||
text = content
|
||||
elif msg_type == "image":
|
||||
pic_url = xml_data.get("PicUrl", "")
|
||||
media_id = xml_data.get("MediaId", "")
|
||||
if pic_url:
|
||||
local, ann = await self._download_attachment(
|
||||
pic_url, f"wechat_{msg_id}.jpg",
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
annotations.append(ann)
|
||||
elif media_id:
|
||||
local, ann = await self._download_wechat_media(
|
||||
media_id, f"wechat_image_{msg_id}",
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
annotations.append(ann)
|
||||
else:
|
||||
annotations.append("[image: no download source]")
|
||||
elif msg_type == "voice":
|
||||
recognition = xml_data.get("Recognition", "")
|
||||
media_id = xml_data.get("MediaId", "")
|
||||
if media_id:
|
||||
local, ann = await self._download_wechat_media(media_id, f"wechat_voice_{msg_id}")
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
ann = ann.replace("[attachment:", "[voice:")
|
||||
annotations.append(ann)
|
||||
if recognition:
|
||||
text = f"[语音识别] {recognition}"
|
||||
elif not media_paths:
|
||||
annotations.append("[voice message]")
|
||||
elif msg_type in ("video", "shortvideo"):
|
||||
media_id = xml_data.get("MediaId", "")
|
||||
if media_id:
|
||||
local, ann = await self._download_wechat_media(media_id, f"wechat_{msg_type}_{msg_id}")
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
annotations.append(ann)
|
||||
if not media_paths:
|
||||
annotations.append(f"[{msg_type} message]")
|
||||
elif msg_type == "location":
|
||||
label = xml_data.get("Label", "")
|
||||
lat = xml_data.get("Location_X", "")
|
||||
lon = xml_data.get("Location_Y", "")
|
||||
text = f"[位置] {label} ({lat}, {lon})"
|
||||
elif msg_type == "file":
|
||||
media_id = xml_data.get("MediaId", "")
|
||||
file_name = xml_data.get("FileName", "") or xml_data.get("Title", f"wechat_file_{msg_id}")
|
||||
logger.info(f"WeChat file message: name={file_name}, media_id={media_id!r}, keys={list(xml_data.keys())}")
|
||||
if media_id:
|
||||
local, ann = await self._download_wechat_media(
|
||||
media_id, f"wechat_file_{msg_id}_{file_name}",
|
||||
)
|
||||
logger.info(f"WeChat file download result: local={local}, ann={ann}")
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
annotations.append(ann)
|
||||
if not media_paths:
|
||||
annotations.append(f"[file: {file_name}]")
|
||||
elif msg_type == "link":
|
||||
title = xml_data.get("Title", "")
|
||||
description = xml_data.get("Description", "")
|
||||
url = xml_data.get("Url", "")
|
||||
text = f"[链接] {title}\n{description}\n{url}"
|
||||
elif msg_type == "event":
|
||||
event_type = xml_data.get("Event", "")
|
||||
if event_type == "subscribe":
|
||||
text = "[用户关注]"
|
||||
elif event_type == "unsubscribe":
|
||||
logger.info(f"User {from_user} unsubscribed")
|
||||
return # Don't process
|
||||
elif event_type == "CLICK":
|
||||
event_key = xml_data.get("EventKey", "")
|
||||
text = f"[菜单点击] {event_key}"
|
||||
elif event_type in ("LOCATION", "VIEW"):
|
||||
# Periodic location reports and menu-link clicks — ignore
|
||||
return
|
||||
else:
|
||||
logger.debug(f"Ignoring WeChat event: {event_type}")
|
||||
return
|
||||
else:
|
||||
text = f"[{msg_type} message]"
|
||||
|
||||
if not text and not media_paths and not annotations:
|
||||
return
|
||||
|
||||
# Parse timestamp
|
||||
try:
|
||||
timestamp = datetime.fromtimestamp(
|
||||
int(create_time)
|
||||
) if create_time else datetime.now()
|
||||
except (ValueError, TypeError, OSError):
|
||||
timestamp = datetime.now()
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=from_user,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=msg_id,
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"to_user": to_user,
|
||||
"backend": self._backend,
|
||||
},
|
||||
))
|
||||
|
||||
# ── Send (template method overrides) ──────────────────────────
|
||||
|
||||
def _format_chunk(self, text: str) -> str:
|
||||
"""WeCom uses markdown formatter; MP uses plain text."""
|
||||
if self._backend == "wecom":
|
||||
return self._formatter.format(text) # markdown profile
|
||||
return _strip_markdown(text)
|
||||
|
||||
async def _send_chunk(
|
||||
self, chat_id, formatted_text, raw_text, reply_to, metadata,
|
||||
):
|
||||
token = await self._ensure_token()
|
||||
|
||||
if self._backend == "wecom":
|
||||
# Group chat: use appchat/send endpoint
|
||||
if chat_id.startswith("wr"):
|
||||
try:
|
||||
await self._wecom_send_group_markdown(token, chat_id, raw_text)
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
await self._wecom_send_group_text(token, chat_id, raw_text)
|
||||
else:
|
||||
# DM: Try markdown first, fall back to plain text
|
||||
try:
|
||||
await self._wecom_send_markdown(token, chat_id, raw_text)
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
await self._wecom_send_text(token, chat_id, raw_text)
|
||||
else:
|
||||
await self._mp_send_text(token, chat_id, raw_text)
|
||||
|
||||
# ── WeCom send ────────────────────────────────────────────────
|
||||
|
||||
async def _wecom_send_text(
|
||||
self, token: str, user_id: str, text: str,
|
||||
) -> None:
|
||||
"""Send a text message via WeCom API."""
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/message/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
body = {
|
||||
"touser": user_id,
|
||||
"msgtype": "text",
|
||||
"agentid": int(self.config.agent_id),
|
||||
"text": {"content": _strip_markdown(text)},
|
||||
}
|
||||
await self._post_api(url, body)
|
||||
|
||||
async def _wecom_send_markdown(
|
||||
self, token: str, user_id: str, text: str,
|
||||
) -> None:
|
||||
"""Send a markdown message via WeCom API.
|
||||
|
||||
Note: WeCom markdown only supports a subset of Markdown
|
||||
(no code blocks, no images). Falls back to text if the
|
||||
message is too complex.
|
||||
"""
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/message/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
body = {
|
||||
"touser": user_id,
|
||||
"msgtype": "markdown",
|
||||
"agentid": int(self.config.agent_id),
|
||||
"markdown": {"content": text},
|
||||
}
|
||||
await self._post_api(url, body)
|
||||
|
||||
# ── WeCom group send ────────────────────────────────────────────
|
||||
|
||||
async def _wecom_send_group_text(
|
||||
self, token: str, chatid: str, text: str,
|
||||
) -> None:
|
||||
"""Send a text message to a WeCom group chat."""
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/appchat/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
body = {
|
||||
"chatid": chatid,
|
||||
"msgtype": "text",
|
||||
"text": {"content": _strip_markdown(text)},
|
||||
}
|
||||
await self._post_api(url, body)
|
||||
|
||||
async def _wecom_send_group_markdown(
|
||||
self, token: str, chatid: str, text: str,
|
||||
) -> None:
|
||||
"""Send a markdown message to a WeCom group chat."""
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/appchat/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
body = {
|
||||
"chatid": chatid,
|
||||
"msgtype": "markdown",
|
||||
"markdown": {"content": text},
|
||||
}
|
||||
await self._post_api(url, body)
|
||||
|
||||
# ── MP send ───────────────────────────────────────────────────
|
||||
|
||||
async def _mp_send_text(
|
||||
self, token: str, openid: str, text: str,
|
||||
) -> None:
|
||||
"""Send a text message via WeChat MP customer service API."""
|
||||
url = (
|
||||
f"https://api.weixin.qq.com/cgi-bin/message/custom/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
body = {
|
||||
"touser": openid,
|
||||
"msgtype": "text",
|
||||
"text": {"content": _strip_markdown(text)},
|
||||
}
|
||||
await self._post_api(url, body)
|
||||
|
||||
# ── Media send ────────────────────────────────────────────────
|
||||
|
||||
async def _send_media_impl(
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
"""Send a media file via WeChat/WeCom."""
|
||||
token = await self._ensure_token()
|
||||
chat_id = self._resolve_media_chat_id(recipient, metadata)
|
||||
|
||||
# Upload media to get media_id
|
||||
media_id = await self._upload_media(token, file_path)
|
||||
if not media_id:
|
||||
return False
|
||||
|
||||
path = Path(file_path)
|
||||
ext = path.suffix.lower()
|
||||
is_image = ext in {".jpg", ".jpeg", ".png", ".gif", ".bmp"}
|
||||
|
||||
if self._backend == "wecom":
|
||||
msg_type = "image" if is_image else "file"
|
||||
# Group chat: use appchat/send endpoint
|
||||
if chat_id.startswith("wr"):
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/appchat/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
body = {
|
||||
"chatid": chat_id,
|
||||
"msgtype": msg_type,
|
||||
msg_type: {"media_id": media_id},
|
||||
}
|
||||
else:
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/message/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
body = {
|
||||
"touser": chat_id,
|
||||
"msgtype": msg_type,
|
||||
"agentid": int(self.config.agent_id),
|
||||
msg_type: {"media_id": media_id},
|
||||
}
|
||||
else:
|
||||
url = (
|
||||
f"https://api.weixin.qq.com/cgi-bin/message/custom/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
msg_type = "image" if is_image else "file" # MP only supports image
|
||||
if not is_image:
|
||||
# MP doesn't support file via customer service API;
|
||||
# send caption as text instead
|
||||
if caption:
|
||||
await self._mp_send_text(token, chat_id, f"[文件] {path.name}\n{caption}")
|
||||
return True
|
||||
body = {
|
||||
"touser": chat_id,
|
||||
"msgtype": "image",
|
||||
"image": {"media_id": media_id},
|
||||
}
|
||||
|
||||
await self._post_api(url, body)
|
||||
|
||||
# Send caption separately if provided
|
||||
if caption:
|
||||
if self._backend == "wecom":
|
||||
if chat_id.startswith("wr"):
|
||||
await self._wecom_send_group_text(token, chat_id, caption)
|
||||
else:
|
||||
await self._wecom_send_text(token, chat_id, caption)
|
||||
else:
|
||||
await self._mp_send_text(token, chat_id, caption)
|
||||
|
||||
return True
|
||||
|
||||
async def _upload_media(
|
||||
self, token: str, file_path: str,
|
||||
) -> str | None:
|
||||
"""Upload a media file and return the media_id."""
|
||||
path = Path(file_path)
|
||||
ext = path.suffix.lower()
|
||||
is_image = ext in {".jpg", ".jpeg", ".png", ".gif", ".bmp"}
|
||||
media_type = "image" if is_image else "file"
|
||||
|
||||
if self._backend == "wecom":
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/media/upload"
|
||||
f"?access_token={token}&type={media_type}"
|
||||
)
|
||||
else:
|
||||
url = (
|
||||
f"https://api.weixin.qq.com/cgi-bin/media/upload"
|
||||
f"?access_token={token}&type={media_type}"
|
||||
)
|
||||
|
||||
try:
|
||||
with open(file_path, "rb") as f:
|
||||
resp = await self._http_client.post(
|
||||
url,
|
||||
files={"media": (path.name, f)},
|
||||
)
|
||||
data = resp.json()
|
||||
if data.get("errcode", 0) != 0 and "media_id" not in data:
|
||||
logger.error(
|
||||
f"WeChat media upload failed: {data.get('errmsg')}"
|
||||
)
|
||||
return None
|
||||
return data.get("media_id")
|
||||
except Exception as e:
|
||||
logger.error(f"WeChat media upload error: {e}")
|
||||
return None
|
||||
|
||||
# ── Media download helper ────────────────────────────────────
|
||||
|
||||
async def _download_wechat_media(
|
||||
self, media_id: str, filename: str,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""Download media by media_id via WeChat/WeCom media API."""
|
||||
token = await self._ensure_token()
|
||||
if self._backend == "wecom":
|
||||
url = f"https://qyapi.weixin.qq.com/cgi-bin/media/get?access_token={token}&media_id={media_id}"
|
||||
else:
|
||||
url = f"https://api.weixin.qq.com/cgi-bin/media/get?access_token={token}&media_id={media_id}"
|
||||
return await self._download_attachment(url, filename)
|
||||
|
||||
# ── Shared API helper ─────────────────────────────────────────
|
||||
|
||||
async def _post_api(self, url: str, body: dict) -> dict:
|
||||
"""POST to WeChat/WeCom API, check errcode, return response."""
|
||||
try:
|
||||
resp = await self._http_client.post(url, json=body)
|
||||
data = resp.json()
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"WeChat API error: {e}")
|
||||
|
||||
errcode = data.get("errcode", 0)
|
||||
if errcode != 0:
|
||||
errmsg = data.get("errmsg", "unknown")
|
||||
# Token expired — refresh and retry once
|
||||
if errcode in (40014, 42001):
|
||||
logger.warning("WeChat token expired, refreshing...")
|
||||
await self._refresh_token()
|
||||
token = self._access_token
|
||||
# Replace token in URL
|
||||
if "access_token=" in url:
|
||||
url = re.sub(
|
||||
r"access_token=[^&]+",
|
||||
f"access_token={token}",
|
||||
url,
|
||||
)
|
||||
resp = await self._http_client.post(url, json=body)
|
||||
data = resp.json()
|
||||
if data.get("errcode", 0) != 0:
|
||||
raise RuntimeError(
|
||||
f"WeChat API error after retry: "
|
||||
f"{data.get('errmsg')}"
|
||||
)
|
||||
return data
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"WeChat API error ({errcode}): {errmsg}"
|
||||
)
|
||||
|
||||
return data
|
||||
|
||||
# _send_typing_action: inherited no-op (WeChat has no typing API)
|
||||
|
||||
|
||||
@@ -0,0 +1,187 @@
|
||||
"""WeChat / WeCom crypto helpers.
|
||||
|
||||
Implements the message encryption/decryption protocol used by both
|
||||
WeCom (企业微信) and WeChat Official Account (公众号) callback APIs.
|
||||
|
||||
The protocol uses AES-256-CBC with a key derived from the EncodingAESKey
|
||||
(base64-encoded 43-char string → 32-byte AES key).
|
||||
|
||||
References:
|
||||
- WeCom: https://developer.work.weixin.qq.com/document/path/90930
|
||||
- MP: https://developers.weixin.qq.com/doc/offiaccount/Message_Management/Message_Encryption_and_Decryption_Instructions.html
|
||||
"""
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import struct
|
||||
import time
|
||||
import xml.etree.ElementTree as ET
|
||||
|
||||
# Crypto imports — all from the Python standard library + pycryptodome
|
||||
# (but we'll use a pure-Python fallback if not available)
|
||||
try:
|
||||
from Crypto.Cipher import AES
|
||||
_HAS_PYCRYPTO = True
|
||||
except ImportError:
|
||||
_HAS_PYCRYPTO = False
|
||||
|
||||
|
||||
def _pkcs7_pad(data: bytes, block_size: int = 32) -> bytes:
|
||||
"""PKCS#7 padding."""
|
||||
pad_len = block_size - (len(data) % block_size)
|
||||
return data + bytes([pad_len]) * pad_len
|
||||
|
||||
|
||||
def _pkcs7_unpad(data: bytes) -> bytes:
|
||||
"""PKCS#7 unpadding."""
|
||||
pad_len = data[-1]
|
||||
if pad_len < 1 or pad_len > 32:
|
||||
return data
|
||||
return data[:-pad_len]
|
||||
|
||||
|
||||
def _aes_decrypt(key: bytes, iv: bytes, ciphertext: bytes) -> bytes:
|
||||
"""AES-256-CBC decryption."""
|
||||
if _HAS_PYCRYPTO:
|
||||
cipher = AES.new(key, AES.MODE_CBC, iv)
|
||||
return cipher.decrypt(ciphertext)
|
||||
else:
|
||||
# Pure-Python AES fallback (slower but no C deps)
|
||||
# We'll try pyaes as a fallback
|
||||
try:
|
||||
import pyaes
|
||||
decrypter = pyaes.Decrypter(
|
||||
pyaes.AESModeOfOperationCBC(key, iv=iv)
|
||||
)
|
||||
decrypted = decrypter.feed(ciphertext)
|
||||
decrypted += decrypter.feed()
|
||||
return decrypted
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"WeChat message decryption requires pycryptodome or pyaes. "
|
||||
"Install with: pip install pycryptodome"
|
||||
)
|
||||
|
||||
|
||||
def _aes_encrypt(key: bytes, iv: bytes, plaintext: bytes) -> bytes:
|
||||
"""AES-256-CBC encryption."""
|
||||
if _HAS_PYCRYPTO:
|
||||
cipher = AES.new(key, AES.MODE_CBC, iv)
|
||||
return cipher.encrypt(plaintext)
|
||||
else:
|
||||
try:
|
||||
import pyaes
|
||||
encrypter = pyaes.Encrypter(
|
||||
pyaes.AESModeOfOperationCBC(key, iv=iv)
|
||||
)
|
||||
encrypted = encrypter.feed(plaintext)
|
||||
encrypted += encrypter.feed()
|
||||
return encrypted
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"WeChat message encryption requires pycryptodome or pyaes. "
|
||||
"Install with: pip install pycryptodome"
|
||||
)
|
||||
|
||||
|
||||
class WeChatCrypto:
|
||||
"""Handles WeChat/WeCom message encryption and decryption.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
token:
|
||||
The Token configured in the WeChat/WeCom callback URL settings.
|
||||
encoding_aes_key:
|
||||
The 43-character EncodingAESKey (base64-encoded).
|
||||
app_id:
|
||||
The AppID (for MP) or CorpID (for WeCom).
|
||||
"""
|
||||
|
||||
def __init__(self, token: str, encoding_aes_key: str, app_id: str):
|
||||
self.token = token
|
||||
self.app_id = app_id
|
||||
# Decode the AES key: EncodingAESKey + "=" → base64 decode → 32 bytes
|
||||
self.aes_key = base64.b64decode(encoding_aes_key + "=")
|
||||
# IV is the first 16 bytes of the key
|
||||
self.iv = self.aes_key[:16]
|
||||
|
||||
def verify_signature(
|
||||
self, signature: str, timestamp: str, nonce: str,
|
||||
encrypt: str = "",
|
||||
) -> bool:
|
||||
"""Verify the callback signature.
|
||||
|
||||
For plain-mode verification (no encryption), *encrypt* can be empty.
|
||||
"""
|
||||
parts = sorted([self.token, timestamp, nonce] + ([encrypt] if encrypt else []))
|
||||
sha1 = hashlib.sha1("".join(parts).encode()).hexdigest()
|
||||
return sha1 == signature
|
||||
|
||||
def decrypt(self, encrypt: str) -> tuple[str, str]:
|
||||
"""Decrypt an encrypted message.
|
||||
|
||||
Returns ``(xml_content, from_app_id)`` tuple.
|
||||
"""
|
||||
ciphertext = base64.b64decode(encrypt)
|
||||
plaintext = _aes_decrypt(self.aes_key, self.iv, ciphertext)
|
||||
plaintext = _pkcs7_unpad(plaintext)
|
||||
|
||||
# plaintext layout:
|
||||
# 16 bytes random + 4 bytes msg_len (big-endian) + msg + app_id
|
||||
msg_len = struct.unpack("!I", plaintext[16:20])[0]
|
||||
msg = plaintext[20:20 + msg_len].decode("utf-8")
|
||||
from_app_id = plaintext[20 + msg_len:].decode("utf-8")
|
||||
return msg, from_app_id
|
||||
|
||||
def encrypt(self, reply_msg: str) -> str:
|
||||
"""Encrypt a reply message.
|
||||
|
||||
Returns the base64-encoded ciphertext.
|
||||
"""
|
||||
msg_bytes = reply_msg.encode("utf-8")
|
||||
app_id_bytes = self.app_id.encode("utf-8")
|
||||
|
||||
# Random 16 bytes + msg_len (4 bytes big-endian) + msg + app_id
|
||||
import os
|
||||
random_bytes = os.urandom(16)
|
||||
msg_len = struct.pack("!I", len(msg_bytes))
|
||||
plaintext = random_bytes + msg_len + msg_bytes + app_id_bytes
|
||||
plaintext = _pkcs7_pad(plaintext)
|
||||
|
||||
ciphertext = _aes_encrypt(self.aes_key, self.iv, plaintext)
|
||||
return base64.b64encode(ciphertext).decode("utf-8")
|
||||
|
||||
def generate_signature(
|
||||
self, encrypt: str, timestamp: str, nonce: str,
|
||||
) -> str:
|
||||
"""Generate the msg_signature for an encrypted reply."""
|
||||
parts = sorted([self.token, timestamp, nonce, encrypt])
|
||||
return hashlib.sha1("".join(parts).encode()).hexdigest()
|
||||
|
||||
def wrap_encrypted_reply(self, reply_msg: str) -> str:
|
||||
"""Encrypt a reply and wrap it in the XML envelope.
|
||||
|
||||
Returns the full XML string to return in the HTTP response.
|
||||
"""
|
||||
encrypt = self.encrypt(reply_msg)
|
||||
timestamp = str(int(time.time()))
|
||||
nonce = hashlib.md5(str(time.time()).encode()).hexdigest()[:10]
|
||||
signature = self.generate_signature(encrypt, timestamp, nonce)
|
||||
|
||||
return (
|
||||
f"<xml>"
|
||||
f"<Encrypt><![CDATA[{encrypt}]]></Encrypt>"
|
||||
f"<MsgSignature><![CDATA[{signature}]]></MsgSignature>"
|
||||
f"<TimeStamp>{timestamp}</TimeStamp>"
|
||||
f"<Nonce><![CDATA[{nonce}]]></Nonce>"
|
||||
f"</xml>"
|
||||
)
|
||||
|
||||
|
||||
def parse_xml(xml_str: str) -> dict[str, str]:
|
||||
"""Parse a WeChat callback XML into a flat dict."""
|
||||
root = ET.fromstring(xml_str)
|
||||
result = {}
|
||||
for child in root:
|
||||
result[child.tag] = child.text or ""
|
||||
return result
|
||||
@@ -0,0 +1,72 @@
|
||||
"""WeChat/WeCom credential validation."""
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_wecom(
|
||||
corp_id: str,
|
||||
secret: str,
|
||||
proxy: str | None = None,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate WeCom credentials by fetching an access token.
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, message).
|
||||
"""
|
||||
if not corp_id or not secret:
|
||||
return False, "corp_id and secret are required"
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return False, "httpx not installed"
|
||||
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/gettoken"
|
||||
f"?corpid={corp_id}&corpsecret={secret}"
|
||||
)
|
||||
try:
|
||||
async with httpx.AsyncClient(proxy=proxy) as client:
|
||||
resp = await client.get(url, timeout=10)
|
||||
data = resp.json()
|
||||
if data.get("errcode", 0) == 0:
|
||||
return True, "WeCom credentials valid"
|
||||
return False, f"Error ({data.get('errcode')}): {data.get('errmsg')}"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
|
||||
|
||||
async def validate_wechat_mp(
|
||||
app_id: str,
|
||||
app_secret: str,
|
||||
proxy: str | None = None,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate WeChat Official Account credentials.
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, message).
|
||||
"""
|
||||
if not app_id or not app_secret:
|
||||
return False, "app_id and app_secret are required"
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return False, "httpx not installed"
|
||||
|
||||
url = (
|
||||
f"https://api.weixin.qq.com/cgi-bin/token"
|
||||
f"?grant_type=client_credential"
|
||||
f"&appid={app_id}&secret={app_secret}"
|
||||
)
|
||||
try:
|
||||
async with httpx.AsyncClient(proxy=proxy) as client:
|
||||
resp = await client.get(url, timeout=10)
|
||||
data = resp.json()
|
||||
if "access_token" in data:
|
||||
return True, "WeChat MP credentials valid"
|
||||
return False, f"Error ({data.get('errcode')}): {data.get('errmsg')}"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
@@ -0,0 +1,130 @@
|
||||
"""WeChat channel server.
|
||||
|
||||
Standalone script to run the WeChat channel with CLI options.
|
||||
|
||||
Usage:
|
||||
# WeCom (企业微信应用)
|
||||
python -m EvoScientist.channels.wechat.serve \\
|
||||
--backend wecom \\
|
||||
--corp-id CORP_ID \\
|
||||
--agent-id AGENT_ID \\
|
||||
--secret SECRET \\
|
||||
--token TOKEN \\
|
||||
--aes-key AES_KEY
|
||||
|
||||
# WeChat Official Account (公众号)
|
||||
python -m EvoScientist.channels.wechat.serve \\
|
||||
--backend wechatmp \\
|
||||
--app-id APP_ID \\
|
||||
--app-secret APP_SECRET \\
|
||||
--token TOKEN \\
|
||||
--aes-key AES_KEY
|
||||
|
||||
Options:
|
||||
--port PORT Webhook listen port (default: 9001)
|
||||
--allow USER_ID Allowed sender (repeatable)
|
||||
--agent Use EvoScientist agent as handler
|
||||
--thinking Send thinking content as intermediate messages
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from .channel import WeChatChannel, WeComConfig, WeChatMPConfig
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_args():
|
||||
"""Parse command line arguments."""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="WeChat channel server",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
choices=["wecom", "wechatmp"],
|
||||
default="wecom",
|
||||
help="WeChat backend type (default: wecom)",
|
||||
)
|
||||
parser.add_argument("--port", type=int, default=9001, help="Webhook port")
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
action="append",
|
||||
dest="allowed_senders",
|
||||
help="Allowed sender ID (repeatable)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--agent",
|
||||
action="store_true",
|
||||
help="Use EvoScientist agent as handler",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--thinking",
|
||||
action="store_true",
|
||||
help="Send thinking content (requires --agent)",
|
||||
)
|
||||
|
||||
# WeCom settings
|
||||
wecom = parser.add_argument_group("WeCom (企业微信)")
|
||||
wecom.add_argument("--corp-id", default="", help="WeCom Corp ID")
|
||||
wecom.add_argument("--agent-id", default="", help="WeCom Agent ID")
|
||||
wecom.add_argument("--secret", default="", help="WeCom Secret")
|
||||
|
||||
# MP settings
|
||||
mp = parser.add_argument_group("WeChat Official Account (公众号)")
|
||||
mp.add_argument("--app-id", default="", help="MP App ID")
|
||||
mp.add_argument("--app-secret", default="", help="MP App Secret")
|
||||
|
||||
# Shared settings
|
||||
parser.add_argument("--token", default="", help="Callback verification token")
|
||||
parser.add_argument("--aes-key", default="", help="EncodingAESKey")
|
||||
parser.add_argument("--proxy", default="", help="HTTP proxy URL")
|
||||
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
args = parse_args()
|
||||
allowed = set(args.allowed_senders) if args.allowed_senders else None
|
||||
proxy = args.proxy or None
|
||||
|
||||
if args.backend == "wecom":
|
||||
config = WeComConfig(
|
||||
corp_id=args.corp_id,
|
||||
agent_id=args.agent_id,
|
||||
secret=args.secret,
|
||||
token=args.token,
|
||||
encoding_aes_key=args.aes_key,
|
||||
webhook_port=args.port,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
)
|
||||
else:
|
||||
config = WeChatMPConfig(
|
||||
app_id=args.app_id,
|
||||
app_secret=args.app_secret,
|
||||
token=args.token,
|
||||
encoding_aes_key=args.aes_key,
|
||||
webhook_port=args.port,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = WeChatChannel(config, backend=args.backend)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,175 @@
|
||||
"""WeChat callback verification server.
|
||||
|
||||
Provides a lightweight temporary HTTP server that handles the WeChat/WeCom
|
||||
URL verification handshake during onboarding. This solves the chicken-and-egg
|
||||
problem: WeChat requires a live server to verify the callback URL before
|
||||
saving, but the main EvoScientist service isn't running during onboard.
|
||||
|
||||
Usage:
|
||||
server = VerifyServer(port, token, encoding_aes_key, corp_id)
|
||||
await server.start()
|
||||
# ... user clicks "Save" in WeCom admin console ...
|
||||
# ... server auto-responds to the verification GET request ...
|
||||
await server.wait_for_verify(timeout=120)
|
||||
await server.stop()
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from aiohttp import web
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class VerifyServer:
|
||||
"""Temporary HTTP server for WeChat/WeCom callback URL verification.
|
||||
|
||||
Handles the GET verification request (signature + echostr) and
|
||||
signals when verification succeeds.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
port: int,
|
||||
token: str,
|
||||
encoding_aes_key: str = "",
|
||||
app_id: str = "",
|
||||
):
|
||||
self.port = port
|
||||
self.token = token
|
||||
self._crypto = None
|
||||
self._runner = None
|
||||
self._site = None
|
||||
self._verified = asyncio.Event()
|
||||
|
||||
if encoding_aes_key and token and app_id:
|
||||
from .crypto import WeChatCrypto
|
||||
self._crypto = WeChatCrypto(
|
||||
token=token,
|
||||
encoding_aes_key=encoding_aes_key,
|
||||
app_id=app_id,
|
||||
)
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Start the verification server."""
|
||||
from aiohttp import web
|
||||
|
||||
app = web.Application()
|
||||
app.router.add_get("/wechat/callback", self._handle)
|
||||
# Also handle POST in case WeCom sends a POST for some reason
|
||||
app.router.add_post("/wechat/callback", self._handle_post)
|
||||
|
||||
self._runner = web.AppRunner(app)
|
||||
await self._runner.setup()
|
||||
self._site = web.TCPSite(self._runner, "0.0.0.0", self.port)
|
||||
await self._site.start()
|
||||
logger.info(f"Verify server listening on port {self.port}")
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the verification server."""
|
||||
if self._site:
|
||||
await self._site.stop()
|
||||
if self._runner:
|
||||
await self._runner.cleanup()
|
||||
self._site = None
|
||||
self._runner = None
|
||||
|
||||
async def wait_for_verify(self, timeout: float = 120) -> bool:
|
||||
"""Wait for verification to succeed.
|
||||
|
||||
Returns True if verified within timeout, False otherwise.
|
||||
"""
|
||||
try:
|
||||
await asyncio.wait_for(self._verified.wait(), timeout=timeout)
|
||||
return True
|
||||
except asyncio.TimeoutError:
|
||||
return False
|
||||
|
||||
@property
|
||||
def is_verified(self) -> bool:
|
||||
return self._verified.is_set()
|
||||
|
||||
async def _handle(self, request) -> web.Response:
|
||||
"""Handle GET verification request.
|
||||
|
||||
During onboarding we use a lenient approach:
|
||||
1. Try strict crypto verification (encrypted mode)
|
||||
2. Try strict plain-mode signature check
|
||||
3. If both fail, fall back to decrypting echostr without
|
||||
signature check (WeCom requires the decrypted echostr)
|
||||
4. Last resort: echo back raw echostr
|
||||
|
||||
This ensures the callback URL can be saved even if Token/AESKey
|
||||
have minor issues, while still attempting proper verification.
|
||||
"""
|
||||
from aiohttp import web
|
||||
|
||||
signature = (
|
||||
request.query.get("msg_signature")
|
||||
or request.query.get("signature", "")
|
||||
)
|
||||
timestamp = request.query.get("timestamp", "")
|
||||
nonce = request.query.get("nonce", "")
|
||||
echostr = request.query.get("echostr", "")
|
||||
|
||||
logger.info(
|
||||
f"Verify request: msg_signature={signature[:16]}... "
|
||||
f"timestamp={timestamp} nonce={nonce} "
|
||||
f"echostr={echostr[:32]}..."
|
||||
)
|
||||
|
||||
if not echostr:
|
||||
return web.Response(status=400, text="missing echostr")
|
||||
|
||||
# Attempt 1: Encrypted mode with full signature verification
|
||||
if self._crypto and request.query.get("msg_signature"):
|
||||
sig_ok = self._crypto.verify_signature(
|
||||
signature, timestamp, nonce, echostr,
|
||||
)
|
||||
if sig_ok:
|
||||
try:
|
||||
plain_echostr, _ = self._crypto.decrypt(echostr)
|
||||
self._verified.set()
|
||||
logger.info("✓ Verified (encrypted, signature OK)")
|
||||
return web.Response(text=plain_echostr)
|
||||
except Exception as e:
|
||||
logger.warning(f"Signature OK but decrypt failed: {e}")
|
||||
else:
|
||||
logger.warning("Signature mismatch, trying decrypt anyway...")
|
||||
|
||||
# Attempt 2: Try decrypt without signature check
|
||||
# (WeCom requires the decrypted echostr to be returned)
|
||||
try:
|
||||
plain_echostr, _ = self._crypto.decrypt(echostr)
|
||||
self._verified.set()
|
||||
logger.info("✓ Verified (decrypted, signature skipped)")
|
||||
return web.Response(text=plain_echostr)
|
||||
except Exception as e:
|
||||
logger.warning(f"Decrypt also failed: {e}")
|
||||
|
||||
# Attempt 3: Plain mode signature check
|
||||
if self.token:
|
||||
parts = sorted([self.token, timestamp, nonce])
|
||||
expected = hashlib.sha1("".join(parts).encode()).hexdigest()
|
||||
if expected == signature:
|
||||
self._verified.set()
|
||||
logger.info("✓ Verified (plain mode)")
|
||||
return web.Response(text=echostr)
|
||||
|
||||
# Attempt 4: Last resort — just echo back the echostr
|
||||
# This won't work for encrypted mode (WeCom expects decrypted),
|
||||
# but works for plain mode with wrong token.
|
||||
logger.warning("All verification methods failed, echoing raw echostr")
|
||||
self._verified.set()
|
||||
return web.Response(text=echostr)
|
||||
|
||||
async def _handle_post(self, request) -> web.Response:
|
||||
"""Handle POST — just acknowledge during verification phase."""
|
||||
from aiohttp import web
|
||||
return web.Response(text="success")
|
||||
@@ -7,7 +7,7 @@ from ..stream.state import ( # noqa: F401
|
||||
_parse_todo_items,
|
||||
_build_todo_stats,
|
||||
)
|
||||
from .channel import ChannelMessage, _ChannelState # noqa: F401
|
||||
from .channel import _channels_is_running, _channels_stop # noqa: F401
|
||||
from .agent import _deduplicate_run_name # noqa: F401
|
||||
|
||||
from ._app import app # noqa: F401
|
||||
|
||||
@@ -45,3 +45,7 @@ Sub-agents (-e): planner-agent | research-agent | code-agent | debug-agent | dat
|
||||
"""
|
||||
mcp_app = typer.Typer(help=_MCP_HELP, invoke_without_command=True)
|
||||
app.add_typer(mcp_app, name="mcp")
|
||||
|
||||
# Channel subcommand group
|
||||
channel_app = typer.Typer(help="Channel management commands")
|
||||
app.add_typer(channel_app, name="channel")
|
||||
|
||||
+353
-215
@@ -1,14 +1,12 @@
|
||||
"""Background iMessage channel — state management, thread lifecycle, handlers."""
|
||||
"""Background channel management — bus mode with ChannelManager."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, Optional
|
||||
|
||||
from rich.panel import Panel
|
||||
from rich.table import Table
|
||||
from rich.text import Text
|
||||
|
||||
from ..stream.display import console
|
||||
@@ -16,213 +14,267 @@ from ..stream.display import console
|
||||
_channel_logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelMessage:
|
||||
"""Message from a channel (iMessage, Email, etc.)."""
|
||||
msg_id: str
|
||||
content: str
|
||||
sender: str
|
||||
channel_type: str # "iMessage", "Email", "Slack"
|
||||
metadata: Any = None
|
||||
# Module-level channel state (bus mode)
|
||||
_manager: Optional[Any] = None # ChannelManager
|
||||
_bus_loop: Optional[asyncio.AbstractEventLoop] = None
|
||||
_bus_thread: Optional[threading.Thread] = None
|
||||
_cli_agent: Any = None # shared agent reference (same as CLI)
|
||||
_cli_thread_id: Optional[str] = None # shared thread_id (same conversation)
|
||||
|
||||
|
||||
class _ChannelState:
|
||||
"""Singleton tracking background iMessage channel and message queue."""
|
||||
def _channels_is_running(channel_type: str | None = None) -> bool:
|
||||
"""Check whether channels are running."""
|
||||
if _manager is None:
|
||||
return False
|
||||
if channel_type:
|
||||
ch = _manager.get_channel(channel_type)
|
||||
return ch is not None and ch._running
|
||||
return _manager.is_running and bool(_manager.running_channels())
|
||||
|
||||
server = None # IMessageServer | None
|
||||
thread = None # threading.Thread | None
|
||||
loop = None # asyncio.AbstractEventLoop | None
|
||||
agent = None # shared agent reference (same as CLI)
|
||||
thread_id = None # shared thread_id (same conversation as CLI)
|
||||
|
||||
# Queue-based communication between channel thread and main CLI thread
|
||||
message_queue: queue.Queue = queue.Queue()
|
||||
pending_responses: dict = {} # msg_id -> {"event": Event, "response": str | None}
|
||||
_response_lock = threading.Lock()
|
||||
def _channels_running_list() -> list[str]:
|
||||
"""Return names of running channels."""
|
||||
return _manager.running_channels() if _manager else []
|
||||
|
||||
@classmethod
|
||||
def is_running(cls) -> bool:
|
||||
return cls.thread is not None and cls.thread.is_alive()
|
||||
|
||||
@classmethod
|
||||
def stop(cls):
|
||||
if cls.loop and cls.server:
|
||||
cls.loop.call_soon_threadsafe(
|
||||
lambda: asyncio.ensure_future(cls.server.stop())
|
||||
def _channels_stop(channel_type: str | None = None) -> None:
|
||||
"""Stop channel(s) and clean up module-level state."""
|
||||
global _manager, _bus_loop, _bus_thread, _cli_agent, _cli_thread_id
|
||||
|
||||
if channel_type is None:
|
||||
# Stop everything
|
||||
if _bus_loop and _manager:
|
||||
try:
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
_manager.stop_all(), _bus_loop,
|
||||
)
|
||||
future.result(timeout=10)
|
||||
except Exception:
|
||||
pass
|
||||
if _manager:
|
||||
_manager.bus.stop()
|
||||
if _bus_thread:
|
||||
_bus_thread.join(timeout=5)
|
||||
_manager = None
|
||||
_bus_loop = None
|
||||
_bus_thread = None
|
||||
_cli_agent = None
|
||||
_cli_thread_id = None
|
||||
return
|
||||
|
||||
# Stop a specific channel
|
||||
if _manager and _bus_loop:
|
||||
try:
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
_manager.remove_channel(channel_type), _bus_loop,
|
||||
)
|
||||
if cls.thread:
|
||||
cls.thread.join(timeout=5)
|
||||
cls.server = None
|
||||
cls.thread = None
|
||||
cls.loop = None
|
||||
cls.agent = None
|
||||
cls.thread_id = None
|
||||
# Clear pending responses
|
||||
with cls._response_lock:
|
||||
for slot in cls.pending_responses.values():
|
||||
slot["event"].set() # Unblock any waiting handlers
|
||||
cls.pending_responses.clear()
|
||||
future.result(timeout=5)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def enqueue(
|
||||
cls,
|
||||
content: str,
|
||||
sender: str,
|
||||
channel_type: str,
|
||||
metadata: Any = None,
|
||||
) -> tuple[str, threading.Event]:
|
||||
"""Enqueue a message from any channel for main thread processing.
|
||||
|
||||
Returns:
|
||||
Tuple of (msg_id, event) - caller can wait on event for response.
|
||||
"""
|
||||
msg_id = str(uuid.uuid4())
|
||||
event = threading.Event()
|
||||
with cls._response_lock:
|
||||
cls.pending_responses[msg_id] = {"event": event, "response": None}
|
||||
cls.message_queue.put(ChannelMessage(msg_id, content, sender, channel_type, metadata))
|
||||
return msg_id, event
|
||||
|
||||
@classmethod
|
||||
def set_response(cls, msg_id: str, response: str) -> None:
|
||||
"""Set response and signal completion."""
|
||||
with cls._response_lock:
|
||||
if msg_id in cls.pending_responses:
|
||||
cls.pending_responses[msg_id]["response"] = response
|
||||
cls.pending_responses[msg_id]["event"].set()
|
||||
|
||||
@classmethod
|
||||
def get_response(cls, msg_id: str, timeout: float = 300) -> str | None:
|
||||
"""Wait for and retrieve response.
|
||||
|
||||
Args:
|
||||
msg_id: The message ID to get response for.
|
||||
timeout: Maximum seconds to wait (default 300 = 5 minutes).
|
||||
|
||||
Returns:
|
||||
The response text, or None if timed out or not found.
|
||||
"""
|
||||
with cls._response_lock:
|
||||
slot = cls.pending_responses.get(msg_id)
|
||||
if not slot:
|
||||
return None
|
||||
if slot["event"].wait(timeout=timeout):
|
||||
with cls._response_lock:
|
||||
return cls.pending_responses.pop(msg_id, {}).get("response")
|
||||
return None
|
||||
if _manager and not _manager.running_channels():
|
||||
_cli_agent = None
|
||||
_cli_thread_id = None
|
||||
|
||||
|
||||
def _run_channel_thread(server):
|
||||
"""Entry point for background channel thread."""
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
_ChannelState.loop = loop
|
||||
try:
|
||||
loop.run_until_complete(server.run())
|
||||
except Exception as e:
|
||||
_channel_logger.error(f"Channel error: {e}")
|
||||
finally:
|
||||
loop.close()
|
||||
def _start_channels_bus_mode(config, agent, thread_id: str, show_thinking: bool = True) -> None:
|
||||
"""Start all channels in bus mode with MessageBus + ChannelManager.
|
||||
|
||||
|
||||
def _create_channel_handler():
|
||||
"""Create iMessage handler that enqueues messages for main thread processing.
|
||||
|
||||
The handler enqueues messages to the shared queue and waits for the main
|
||||
CLI thread to process them with full Rich Live streaming. This ensures
|
||||
channel messages get the same display quality as direct CLI input.
|
||||
|
||||
Returns:
|
||||
Async handler function: (msg) -> str
|
||||
Creates a single event loop in a daemon thread running the bus,
|
||||
ChannelManager, and the inbound consumer.
|
||||
"""
|
||||
global _manager, _bus_loop, _bus_thread
|
||||
|
||||
async def handler(msg) -> str:
|
||||
# Enqueue for main thread to process with full Live streaming
|
||||
msg_id, event = _ChannelState.enqueue(
|
||||
content=msg.content,
|
||||
sender=msg.sender,
|
||||
channel_type="iMessage",
|
||||
metadata=msg.metadata,
|
||||
from ..channels.channel_manager import ChannelManager
|
||||
|
||||
mgr = ChannelManager.from_config(config)
|
||||
|
||||
if show_thinking:
|
||||
for channel in mgr._channels.values():
|
||||
channel.send_thinking = True
|
||||
|
||||
_manager = mgr
|
||||
|
||||
def _bus_thread_entry():
|
||||
global _bus_loop
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
_bus_loop = loop
|
||||
|
||||
async def _run():
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(mgr.bus, mgr, agent, thread_id, show_thinking)
|
||||
)
|
||||
try:
|
||||
await mgr.start_all()
|
||||
finally:
|
||||
consumer.cancel()
|
||||
try:
|
||||
await consumer
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
try:
|
||||
loop.run_until_complete(_run())
|
||||
except Exception as e:
|
||||
_channel_logger.error(f"Bus thread error: {e}")
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
thread = threading.Thread(target=_bus_thread_entry, daemon=True)
|
||||
_bus_thread = thread
|
||||
thread.start()
|
||||
|
||||
# Wait briefly for the loop to start
|
||||
import time
|
||||
for _ in range(20):
|
||||
if _bus_loop is not None:
|
||||
break
|
||||
time.sleep(0.1)
|
||||
|
||||
|
||||
def _add_channel_to_running_bus(channel_type: str, config) -> None:
|
||||
"""Dynamically add a single channel to the already-running bus.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the bus loop or manager is not initialised.
|
||||
ValueError: If the channel type is unknown or already registered.
|
||||
"""
|
||||
if not _manager or not _bus_loop:
|
||||
raise RuntimeError("Bus not initialised")
|
||||
|
||||
async def _do_add():
|
||||
channel = await _manager.add_channel(channel_type, config)
|
||||
channel.send_thinking = True
|
||||
|
||||
future = asyncio.run_coroutine_threadsafe(_do_add(), _bus_loop)
|
||||
future.result(timeout=10)
|
||||
|
||||
|
||||
async def _bus_inbound_consumer(
|
||||
bus, manager, agent, thread_id: str, show_thinking: bool = True,
|
||||
) -> None:
|
||||
"""Core bridge: consume inbound messages from bus and run agent.
|
||||
|
||||
Streams agent events on the bus loop with Rich Live real-time display
|
||||
(identical to interactive CLI) and sends thinking / todo / answer to
|
||||
the originating channel via direct ``await`` calls.
|
||||
"""
|
||||
from ..stream import events as _stream_events_mod
|
||||
from ..stream.display import (
|
||||
console, create_streaming_display,
|
||||
)
|
||||
from ..stream.state import StreamState
|
||||
from ..channels.consumer import _format_todo_list
|
||||
from ..channels.bus.events import OutboundMessage
|
||||
from rich.live import Live
|
||||
from rich.text import Text as _Text
|
||||
|
||||
def _print_separator():
|
||||
width = console.size.width
|
||||
console.print(_Text("\u2500" * width, style="dim"))
|
||||
|
||||
while True:
|
||||
try:
|
||||
msg = await asyncio.wait_for(bus.consume_inbound(), timeout=1.0)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
|
||||
_channel_logger.info(
|
||||
f"[bus] Processing from {msg.channel}:{msg.sender_id}: "
|
||||
f"{msg.content[:60]}..."
|
||||
)
|
||||
manager.record_message(msg.channel, "received")
|
||||
|
||||
# Wait indefinitely for main thread to process and set response
|
||||
# (no timeout - let the agent work as long as needed)
|
||||
await asyncio.to_thread(event.wait)
|
||||
# CLI: show query from channel (mirrors interactive prompt)
|
||||
source_label = _Text()
|
||||
source_label.append(f"[{msg.channel}] ", style="cyan bold")
|
||||
source_label.append(msg.content)
|
||||
console.print(source_label)
|
||||
|
||||
# Get the response
|
||||
with _ChannelState._response_lock:
|
||||
response = _ChannelState.pending_responses.pop(msg_id, {}).get("response", "")
|
||||
channel = manager.get_channel(msg.channel)
|
||||
state = StreamState()
|
||||
thinking_sent = False
|
||||
todo_sent = False
|
||||
|
||||
return response if response else "(empty response)"
|
||||
if channel:
|
||||
await channel.start_typing(msg.chat_id)
|
||||
|
||||
return handler
|
||||
try:
|
||||
with Live(console=console, refresh_per_second=10, transient=False) as live:
|
||||
live.update(create_streaming_display(is_waiting=True))
|
||||
|
||||
async for event in _stream_events_mod.stream_agent_events(
|
||||
agent, msg.content, thread_id,
|
||||
):
|
||||
etype = state.handle_event(event)
|
||||
|
||||
def _cmd_channel(args: str, agent: Any, thread_id: str) -> None:
|
||||
"""Start iMessage channel in background thread using the shared agent.
|
||||
# Channel: send thinking on transition
|
||||
if (etype != "thinking"
|
||||
and not thinking_sent
|
||||
and state.thinking_text):
|
||||
if channel and show_thinking:
|
||||
await channel.send_thinking_message(
|
||||
msg.sender_id, state.thinking_text, msg.metadata,
|
||||
)
|
||||
thinking_sent = True
|
||||
|
||||
CLI and iMessage share the same agent + thread_id (same conversation).
|
||||
When an iMessage arrives, the main CLI thread processes it with full
|
||||
Rich Live streaming — same experience as direct CLI input.
|
||||
# Channel: send todo list
|
||||
if (etype == "tool_call"
|
||||
and event.get("name") == "write_todos"
|
||||
and not todo_sent
|
||||
and state.todo_items):
|
||||
if channel:
|
||||
await channel.send_todo_message(
|
||||
msg.sender_id,
|
||||
_format_todo_list(state.todo_items),
|
||||
msg.metadata,
|
||||
)
|
||||
todo_sent = True
|
||||
|
||||
Usage: /channel [--allow SENDER]
|
||||
"""
|
||||
from ..channels.imessage import IMessageConfig
|
||||
from ..channels.imessage.serve import IMessageServer
|
||||
# CLI: Live update
|
||||
live.update(create_streaming_display(
|
||||
**state.get_display_args(),
|
||||
show_thinking=show_thinking,
|
||||
))
|
||||
if etype in (
|
||||
"tool_call", "tool_result",
|
||||
"subagent_start", "subagent_tool_call",
|
||||
"subagent_tool_result", "subagent_end",
|
||||
):
|
||||
live.refresh()
|
||||
|
||||
if _ChannelState.is_running():
|
||||
console.print("[dim]iMessage channel already running[/dim]")
|
||||
console.print("[dim]Use[/dim] /channel stop [dim]to disconnect[/dim]\n")
|
||||
return
|
||||
# Flush remaining thinking
|
||||
if (not thinking_sent
|
||||
and state.thinking_text):
|
||||
if channel and show_thinking:
|
||||
await channel.send_thinking_message(
|
||||
msg.sender_id, state.thinking_text, msg.metadata,
|
||||
)
|
||||
|
||||
parts = args.split() if args else []
|
||||
allowed = set()
|
||||
|
||||
for i, p in enumerate(parts):
|
||||
if p == "--allow" and i + 1 < len(parts):
|
||||
allowed.add(parts[i + 1])
|
||||
|
||||
config = IMessageConfig(
|
||||
allowed_senders=list(allowed) if allowed else [],
|
||||
)
|
||||
|
||||
# Store shared agent reference — no separate agent creation
|
||||
_ChannelState.agent = agent
|
||||
_ChannelState.thread_id = thread_id
|
||||
|
||||
# Read send_thinking preference from config
|
||||
from ..config import load_config as _load_config
|
||||
send_thinking = _load_config().imessage_send_thinking
|
||||
|
||||
server = IMessageServer(
|
||||
config,
|
||||
handler=_create_channel_handler(),
|
||||
send_thinking=send_thinking,
|
||||
)
|
||||
|
||||
_ChannelState.server = server
|
||||
_ChannelState.thread = threading.Thread(
|
||||
target=_run_channel_thread,
|
||||
args=(server,),
|
||||
daemon=True,
|
||||
)
|
||||
_ChannelState.thread.start()
|
||||
|
||||
console.print("[green]iMessage channel running in background[/green]")
|
||||
if allowed:
|
||||
console.print(f"[dim]Allowed:[/dim] {allowed}")
|
||||
else:
|
||||
console.print("[dim]Allowed: all senders[/dim]")
|
||||
console.print("[dim]Use[/dim] /channel stop [dim]to disconnect[/dim]\n")
|
||||
|
||||
|
||||
def _cmd_channel_stop() -> None:
|
||||
"""Stop background iMessage channel."""
|
||||
if not _ChannelState.is_running():
|
||||
console.print("[dim]No channel running[/dim]\n")
|
||||
return
|
||||
_ChannelState.stop()
|
||||
console.print("[dim]iMessage channel stopped[/dim]\n")
|
||||
# Channel: publish answer
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=state.response_text or "No response",
|
||||
reply_to=msg.message_id or None,
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
manager.record_message(msg.channel, "sent")
|
||||
console.print(_Text("> ", style="blue bold"), end="")
|
||||
except Exception as e:
|
||||
_channel_logger.error(f"[bus] Agent error: {e}")
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=f"Error processing message: {e}",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
finally:
|
||||
if channel:
|
||||
await channel.stop_typing(msg.chat_id)
|
||||
|
||||
|
||||
def _print_channel_panel(channels: list[tuple[str, bool, str]]) -> None:
|
||||
@@ -252,38 +304,124 @@ def _print_channel_panel(channels: list[tuple[str, bool, str]]) -> None:
|
||||
console.print()
|
||||
|
||||
|
||||
def _auto_start_channel(agent: Any, thread_id: str, allowed_senders_csv: str, send_thinking: bool = True) -> None:
|
||||
"""Start iMessage channel automatically from config.
|
||||
def _cmd_channel(args: str, agent: Any, thread_id: str) -> None:
|
||||
"""Start a channel in background using bus mode.
|
||||
|
||||
Usage:
|
||||
/channel [telegram|discord|imessage] -- start channel (default from config)
|
||||
/channel status -- show current channel status
|
||||
/channel stop -- stop running channel
|
||||
"""
|
||||
global _cli_agent, _cli_thread_id
|
||||
|
||||
from ..config import load_config
|
||||
app_config = load_config()
|
||||
|
||||
channel_type = args.strip().lower() if args and args.strip() else ""
|
||||
if channel_type == "status":
|
||||
running = _channels_running_list()
|
||||
if running and _manager:
|
||||
detailed = _manager.get_detailed_status()
|
||||
table = Table(title="Channel Status", show_header=True, expand=False)
|
||||
table.add_column("Channel", style="cyan")
|
||||
table.add_column("Status")
|
||||
table.add_column("Uptime", style="dim")
|
||||
table.add_column("Rx", justify="right")
|
||||
table.add_column("Tx", justify="right")
|
||||
for ch_name in running:
|
||||
info = detailed.get(ch_name, {})
|
||||
secs = info.get("uptime_seconds", 0)
|
||||
mins, s = divmod(int(secs), 60)
|
||||
hours, mins = divmod(mins, 60)
|
||||
uptime = f"{hours}h{mins:02d}m" if hours else f"{mins}m{s:02d}s"
|
||||
rx = str(info.get("received", 0))
|
||||
tx = str(info.get("sent", 0))
|
||||
table.add_row(ch_name, "[green]running[/green]", uptime, rx, tx)
|
||||
console.print(table)
|
||||
console.print()
|
||||
else:
|
||||
console.print("[dim]No channel running[/dim]\n")
|
||||
return
|
||||
|
||||
if not channel_type:
|
||||
channel_type = app_config.channel_enabled
|
||||
if not channel_type:
|
||||
console.print("[yellow]No channel configured.[/yellow]")
|
||||
console.print("[dim]Run[/dim] evosci onboard [dim]or specify:[/dim] /channel telegram\n")
|
||||
return
|
||||
|
||||
requested = [t.strip() for t in channel_type.split(",") if t.strip()]
|
||||
|
||||
if _channels_is_running():
|
||||
running = _channels_running_list()
|
||||
results: list[tuple[str, bool, str]] = []
|
||||
for ct in requested:
|
||||
if ct in running:
|
||||
results.append((ct, True, "already running"))
|
||||
else:
|
||||
try:
|
||||
_add_channel_to_running_bus(ct, app_config)
|
||||
results.append((ct, True, "connected (bus)"))
|
||||
except Exception as e:
|
||||
results.append((ct, False, str(e)))
|
||||
_print_channel_panel(results)
|
||||
return
|
||||
|
||||
_cli_agent = agent
|
||||
_cli_thread_id = thread_id
|
||||
|
||||
# Override channel_enabled for this invocation
|
||||
original = app_config.channel_enabled
|
||||
app_config.channel_enabled = channel_type
|
||||
try:
|
||||
_start_channels_bus_mode(app_config, agent, thread_id)
|
||||
results = [(ct, True, "connected (bus)") for ct in requested]
|
||||
except Exception as e:
|
||||
results = [(ct, False, str(e)) for ct in requested]
|
||||
finally:
|
||||
app_config.channel_enabled = original
|
||||
|
||||
_print_channel_panel(results)
|
||||
|
||||
|
||||
def _cmd_channel_stop(channel_type: str | None = None) -> None:
|
||||
"""Stop background channel(s).
|
||||
|
||||
Args:
|
||||
channel_type: Specific channel to stop, or None to stop all.
|
||||
"""
|
||||
if not _channels_is_running():
|
||||
console.print("[dim]No channel running[/dim]\n")
|
||||
return
|
||||
if channel_type:
|
||||
if not _channels_is_running(channel_type):
|
||||
console.print(f"[dim]{channel_type} is not running[/dim]\n")
|
||||
return
|
||||
_channels_stop(channel_type)
|
||||
console.print(f"[dim]{channel_type} stopped[/dim]\n")
|
||||
else:
|
||||
running = _channels_running_list()
|
||||
_channels_stop()
|
||||
console.print(f"[dim]{', '.join(running)} stopped[/dim]\n")
|
||||
|
||||
|
||||
def _auto_start_channel(agent: Any, thread_id: str, config) -> None:
|
||||
"""Start channels automatically from config (bus mode).
|
||||
|
||||
Args:
|
||||
agent: Compiled agent graph.
|
||||
thread_id: Current thread ID.
|
||||
allowed_senders_csv: Comma-separated allowed senders (empty = all).
|
||||
send_thinking: Whether to forward thinking content to channel.
|
||||
config: EvoScientistConfig with channel settings.
|
||||
"""
|
||||
try:
|
||||
from ..channels.imessage import IMessageConfig
|
||||
from ..channels.imessage.serve import IMessageServer
|
||||
global _cli_agent, _cli_thread_id
|
||||
|
||||
allowed: set[str] | None = None
|
||||
if allowed_senders_csv.strip():
|
||||
allowed = {s.strip() for s in allowed_senders_csv.split(",") if s.strip()}
|
||||
if not config.channel_enabled:
|
||||
return
|
||||
|
||||
config = IMessageConfig(allowed_senders=list(allowed) if allowed else [])
|
||||
_cli_agent = agent
|
||||
_cli_thread_id = thread_id
|
||||
|
||||
_ChannelState.agent = agent
|
||||
_ChannelState.thread_id = thread_id
|
||||
|
||||
server = IMessageServer(config, handler=_create_channel_handler(), send_thinking=send_thinking)
|
||||
_ChannelState.server = server
|
||||
_ChannelState.thread = threading.Thread(
|
||||
target=_run_channel_thread,
|
||||
args=(server,),
|
||||
daemon=True,
|
||||
)
|
||||
_ChannelState.thread.start()
|
||||
|
||||
detail = ", ".join(sorted(allowed)) if allowed else "all senders"
|
||||
_print_channel_panel([("iMessage", True, detail)])
|
||||
except Exception as e:
|
||||
_print_channel_panel([("iMessage", False, str(e))])
|
||||
_start_channels_bus_mode(config, agent, thread_id)
|
||||
types = [t.strip() for t in config.channel_enabled.split(",") if t.strip()]
|
||||
results = [(ct, True, "connected (bus)") for ct in types]
|
||||
_print_channel_panel(results)
|
||||
|
||||
@@ -11,9 +11,10 @@ import typer # type: ignore[import-untyped]
|
||||
from rich.table import Table
|
||||
|
||||
from ..stream.display import console
|
||||
from ..paths import ensure_dirs, set_workspace_root
|
||||
from ._app import app, config_app, mcp_app
|
||||
from .agent import _deduplicate_run_name, _create_session_workspace, _load_agent
|
||||
from ..paths import ensure_dirs, default_workspace_dir, set_workspace_root
|
||||
from ._app import app, config_app, mcp_app, channel_app
|
||||
from .agent import _deduplicate_run_name, _create_session_workspace, _load_agent, _shorten_path
|
||||
from .channel import _channels_stop, _start_channels_bus_mode
|
||||
from .mcp_ui import (
|
||||
_mcp_list_servers,
|
||||
_mcp_add_server_from_kwargs,
|
||||
@@ -45,6 +46,100 @@ def onboard(
|
||||
run_onboard(skip_validation=skip_validation)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Channel setup command
|
||||
# =============================================================================
|
||||
|
||||
@channel_app.command("setup")
|
||||
def channel_setup():
|
||||
"""Interactive channel configuration wizard.
|
||||
|
||||
Guides you through selecting and configuring a messaging channel
|
||||
(Telegram, Discord, or iMessage).
|
||||
"""
|
||||
import asyncio
|
||||
try:
|
||||
asyncio.get_event_loop()
|
||||
except RuntimeError:
|
||||
asyncio.set_event_loop(asyncio.new_event_loop())
|
||||
|
||||
from ..config import load_config, save_config
|
||||
from ..config.onboard import _step_channels
|
||||
|
||||
config = load_config()
|
||||
updates = _step_channels(config)
|
||||
if updates:
|
||||
for key, value in updates.items():
|
||||
setattr(config, key, value)
|
||||
save_config(config)
|
||||
console.print("[green]Channel configuration saved.[/green]")
|
||||
else:
|
||||
console.print("[dim]No changes made.[/dim]")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Serve command (headless mode)
|
||||
# =============================================================================
|
||||
|
||||
@app.command()
|
||||
def serve(
|
||||
no_thinking: bool = typer.Option(False, "--no-thinking", help="Disable thinking relay to channels"),
|
||||
workdir: Optional[str] = typer.Option(None, "--workdir", help="Override workspace directory"),
|
||||
):
|
||||
"""Run EvoScientist in headless mode -- channels only, no interactive prompt.
|
||||
|
||||
Starts all configured channels and processes messages via the agent.
|
||||
Press Ctrl+C to shut down.
|
||||
"""
|
||||
import nest_asyncio # type: ignore[import-untyped]
|
||||
import uuid
|
||||
nest_asyncio.apply()
|
||||
|
||||
from dotenv import load_dotenv, find_dotenv # type: ignore[import-untyped]
|
||||
load_dotenv(find_dotenv(), override=True)
|
||||
|
||||
from ..config import get_effective_config, apply_config_to_env
|
||||
|
||||
config = get_effective_config()
|
||||
apply_config_to_env(config)
|
||||
|
||||
if not config.channel_enabled:
|
||||
console.print("[red]No channels configured.[/red]")
|
||||
console.print("[dim]Run [bold]evosci channel setup[/bold] first.[/dim]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
show_thinking = not no_thinking
|
||||
ensure_dirs()
|
||||
|
||||
if workdir:
|
||||
ws = os.path.abspath(os.path.expanduser(workdir))
|
||||
os.makedirs(ws, exist_ok=True)
|
||||
else:
|
||||
ws = str(default_workspace_dir())
|
||||
os.makedirs(ws, exist_ok=True)
|
||||
|
||||
console.print("[dim]Loading agent...[/dim]")
|
||||
agent = _load_agent(workspace_dir=ws)
|
||||
tid = str(uuid.uuid4())
|
||||
|
||||
_start_channels_bus_mode(config, agent, tid, show_thinking)
|
||||
console.print("[green]Serve mode started (bus mode).[/green]")
|
||||
|
||||
console.print(f"[dim]Thread: {tid}[/dim]")
|
||||
console.print(f"[dim]Workspace: {_shorten_path(ws)}[/dim]")
|
||||
console.print("[dim]Press Ctrl+C to stop.[/dim]\n")
|
||||
|
||||
import time
|
||||
try:
|
||||
while True:
|
||||
time.sleep(1)
|
||||
except KeyboardInterrupt:
|
||||
console.print("\n[dim]Shutting down...[/dim]")
|
||||
finally:
|
||||
_channels_stop()
|
||||
console.print("[dim]Stopped.[/dim]")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Config commands
|
||||
# =============================================================================
|
||||
@@ -435,9 +530,6 @@ def _main_callback(
|
||||
mode=effective_mode,
|
||||
model=config.model,
|
||||
provider=config.provider,
|
||||
imessage_enabled=config.imessage_enabled,
|
||||
imessage_allowed_senders=config.imessage_allowed_senders,
|
||||
imessage_send_thinking=config.imessage_send_thinking,
|
||||
run_name=name,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
|
||||
+15
-127
@@ -2,7 +2,6 @@
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import queue
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
@@ -33,12 +32,12 @@ from ..sessions import (
|
||||
from ..stream.display import console, _run_streaming
|
||||
from .agent import _shorten_path, _create_session_workspace, _load_agent
|
||||
from .channel import (
|
||||
ChannelMessage,
|
||||
_ChannelState,
|
||||
_channels_is_running,
|
||||
_cmd_channel,
|
||||
_cmd_channel_stop,
|
||||
_auto_start_channel,
|
||||
)
|
||||
import EvoScientist.cli.channel as _ch_mod
|
||||
from .mcp_ui import _cmd_mcp
|
||||
from .skills_cmd import _cmd_list_skills, _cmd_install_skill, _cmd_uninstall_skill
|
||||
|
||||
@@ -177,9 +176,6 @@ def cmd_interactive(
|
||||
mode: str | None = None,
|
||||
model: str | None = None,
|
||||
provider: str | None = None,
|
||||
imessage_enabled: bool = False,
|
||||
imessage_allowed_senders: str = "",
|
||||
imessage_send_thinking: bool = True,
|
||||
run_name: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
) -> None:
|
||||
@@ -195,9 +191,6 @@ def cmd_interactive(
|
||||
mode: Workspace mode ('daemon' or 'run'), displayed in banner
|
||||
model: Model name to display in banner
|
||||
provider: LLM provider name to display in banner
|
||||
imessage_enabled: Whether to auto-start iMessage channel
|
||||
imessage_allowed_senders: Comma-separated allowed senders
|
||||
imessage_send_thinking: Whether to forward thinking to channel
|
||||
run_name: Optional run name for /new session deduplication
|
||||
thread_id: Optional thread ID to resume a previous session
|
||||
"""
|
||||
@@ -231,106 +224,6 @@ def cmd_interactive(
|
||||
"resumed": False,
|
||||
}
|
||||
|
||||
def _process_channel_message(msg: ChannelMessage) -> None:
|
||||
"""Process a message from a channel with full Live streaming."""
|
||||
# Move past the current prompt line to avoid interference with prompt_toolkit
|
||||
# Then move back up and clear that line
|
||||
sys.stdout.write("\n\033[A\033[2K\r")
|
||||
sys.stdout.flush()
|
||||
# Display prompt with channel source on second line
|
||||
console.print(f"[bold blue]>[/bold blue] {msg.content}")
|
||||
console.print(Text.assemble(
|
||||
("[", "dim"),
|
||||
(f"{msg.channel_type}: Received from ", "dim"),
|
||||
(msg.sender, "cyan"),
|
||||
("]", "dim"),
|
||||
))
|
||||
_print_separator()
|
||||
console.print()
|
||||
|
||||
# Build channel callbacks for intermediate messages (thinking + todo + files)
|
||||
on_thinking = None
|
||||
on_todo = None
|
||||
on_file_write = None
|
||||
if _ChannelState.is_running() and _ChannelState.server and _ChannelState.loop:
|
||||
def _send_thinking(thinking_text: str) -> None:
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
_ChannelState.server.send_thinking_message(
|
||||
msg.sender, thinking_text, msg.metadata,
|
||||
),
|
||||
_ChannelState.loop,
|
||||
)
|
||||
except Exception:
|
||||
pass # Non-critical — don't break main flow
|
||||
|
||||
def _send_todo(todo_items: list) -> None:
|
||||
try:
|
||||
lines = [f"\U0001f4cb {len(todo_items)} tasks ongoing"] # 📋
|
||||
for i, item in enumerate(todo_items, 1):
|
||||
content = item.get("content", "")
|
||||
lines.append(f"{i}. {content}")
|
||||
lines.append("\U0001f680") # 🚀
|
||||
formatted = "\n".join(lines)
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
_ChannelState.server.send_todo_message(
|
||||
msg.sender, formatted, msg.metadata,
|
||||
),
|
||||
_ChannelState.loop,
|
||||
)
|
||||
except Exception:
|
||||
pass # Non-critical — don't break main flow
|
||||
|
||||
def _send_file(real_path: str) -> None:
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
_ChannelState.server.channel.send_media(
|
||||
recipient=msg.sender, file_path=real_path,
|
||||
metadata=msg.metadata,
|
||||
),
|
||||
_ChannelState.loop,
|
||||
)
|
||||
except Exception:
|
||||
pass # Non-critical — don't break main flow
|
||||
|
||||
on_thinking = _send_thinking
|
||||
on_todo = _send_todo
|
||||
on_file_write = _send_file
|
||||
|
||||
try:
|
||||
meta = _build_metadata(state["workspace_dir"], model)
|
||||
# Use SAME _run_streaming as CLI input — full Live experience
|
||||
response_text = _run_streaming(
|
||||
state["agent"], msg.content, state["thread_id"], show_thinking,
|
||||
interactive=True, on_thinking=on_thinking, on_todo=on_todo,
|
||||
on_file_write=on_file_write, metadata=meta,
|
||||
)
|
||||
|
||||
# Set response for channel handler to retrieve
|
||||
_ChannelState.set_response(msg.msg_id, response_text or "")
|
||||
# Show replied indicator
|
||||
console.print(Text.assemble(
|
||||
("[", "dim"),
|
||||
(f"{msg.channel_type}: Replied to ", "dim"),
|
||||
(msg.sender, "cyan"),
|
||||
("]", "dim"),
|
||||
))
|
||||
except Exception as e:
|
||||
console.print(f"[red]Channel processing error: {e}[/red]")
|
||||
_ChannelState.set_response(msg.msg_id, f"Error: {e}")
|
||||
|
||||
_print_separator()
|
||||
|
||||
async def _check_channel_queue():
|
||||
"""Background task to check channel queue periodically."""
|
||||
while state["running"]:
|
||||
try:
|
||||
msg = _ChannelState.message_queue.get_nowait()
|
||||
_process_channel_message(msg)
|
||||
except queue.Empty:
|
||||
pass
|
||||
await asyncio.sleep(0.1) # Check every 100ms
|
||||
|
||||
async def _resolve_thread_id(tid: str) -> str | None:
|
||||
"""Resolve a (possibly partial) thread ID. Returns full ID or None."""
|
||||
if await thread_exists(tid):
|
||||
@@ -486,9 +379,9 @@ def cmd_interactive(
|
||||
console.print("[dim]Loading session...[/dim]")
|
||||
state["agent"] = _load_agent(workspace_dir=state["workspace_dir"], checkpointer=checkpointer)
|
||||
# Sync shared refs if channel is running
|
||||
if _ChannelState.is_running():
|
||||
_ChannelState.agent = state["agent"]
|
||||
_ChannelState.thread_id = state["thread_id"]
|
||||
if _channels_is_running():
|
||||
_ch_mod._cli_agent = state["agent"]
|
||||
_ch_mod._cli_thread_id = state["thread_id"]
|
||||
console.print(f"[green]Resumed session:[/green] [yellow]{resolved}[/yellow]")
|
||||
if state["workspace_dir"]:
|
||||
console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]")
|
||||
@@ -537,11 +430,13 @@ def cmd_interactive(
|
||||
print_banner(state["thread_id"], state["workspace_dir"], memory_dir, mode, model, provider)
|
||||
|
||||
# Start background queue checker
|
||||
queue_task = asyncio.create_task(_check_channel_queue())
|
||||
# (no longer needed — bus mode handles messages internally)
|
||||
|
||||
# Auto-start iMessage channel if enabled in config
|
||||
if imessage_enabled and not _ChannelState.is_running():
|
||||
_auto_start_channel(state["agent"], state["thread_id"], imessage_allowed_senders, imessage_send_thinking)
|
||||
# Auto-start channel if enabled in config
|
||||
from ..config import load_config
|
||||
config = load_config()
|
||||
if config and config.channel_enabled and not _channels_is_running():
|
||||
_auto_start_channel(state["agent"], state["thread_id"], config)
|
||||
|
||||
try:
|
||||
_print_separator()
|
||||
@@ -588,10 +483,6 @@ def cmd_interactive(
|
||||
state["agent"] = _load_agent(workspace_dir=state["workspace_dir"], checkpointer=checkpointer)
|
||||
state["thread_id"] = generate_thread_id()
|
||||
state["resumed"] = False
|
||||
# Sync shared refs if channel is running
|
||||
if _ChannelState.is_running():
|
||||
_ChannelState.agent = state["agent"]
|
||||
_ChannelState.thread_id = state["thread_id"]
|
||||
console.print(f"[green]New session:[/green] [yellow]{state['thread_id']}[/yellow]")
|
||||
if state["workspace_dir"]:
|
||||
console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]\n")
|
||||
@@ -626,8 +517,9 @@ def cmd_interactive(
|
||||
|
||||
if user_input.lower().startswith("/channel"):
|
||||
args = user_input[len("/channel"):].strip()
|
||||
if args.lower() == "stop":
|
||||
_cmd_channel_stop()
|
||||
if args.lower().startswith("stop"):
|
||||
stop_arg = args[len("stop"):].strip()
|
||||
_cmd_channel_stop(stop_arg or None)
|
||||
else:
|
||||
_cmd_channel(args, state["agent"], state["thread_id"])
|
||||
continue
|
||||
@@ -660,11 +552,7 @@ def cmd_interactive(
|
||||
else:
|
||||
console.print(f"[red]Error: {e}[/red]")
|
||||
finally:
|
||||
queue_task.cancel()
|
||||
try:
|
||||
await queue_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
pass
|
||||
|
||||
# Run the async main loop
|
||||
try:
|
||||
|
||||
@@ -85,9 +85,101 @@ class EvoScientistConfig:
|
||||
show_thinking: bool = True
|
||||
|
||||
# Channel Settings
|
||||
imessage_enabled: bool = False
|
||||
imessage_allowed_senders: str = "" # comma-separated, empty = allow all
|
||||
imessage_send_thinking: bool = True # forward thinking to channel
|
||||
channel_enabled: str = "" # "imessage" | "telegram" | "discord" | "" (comma-separated for multiple)
|
||||
require_mention: str = "group" # "always" | "group" | "off"
|
||||
text_chunk_limit: int = 0 # 0 = use capability default
|
||||
allowed_channels: str = "" # comma-separated channel IDs, empty = allow all
|
||||
|
||||
# iMessage Settings
|
||||
imessage_enabled: bool = False # legacy compat
|
||||
imessage_allowed_senders: str = ""
|
||||
imessage_send_thinking: bool = True
|
||||
|
||||
# Telegram Settings
|
||||
telegram_bot_token: str = ""
|
||||
telegram_allowed_senders: str = ""
|
||||
telegram_proxy: str = ""
|
||||
|
||||
# Discord Settings
|
||||
discord_bot_token: str = ""
|
||||
discord_allowed_senders: str = ""
|
||||
discord_allowed_channels: str = ""
|
||||
discord_proxy: str = ""
|
||||
|
||||
# Slack Settings
|
||||
slack_bot_token: str = ""
|
||||
slack_app_token: str = ""
|
||||
slack_allowed_senders: str = ""
|
||||
slack_allowed_channels: str = ""
|
||||
slack_proxy: str = ""
|
||||
|
||||
# Feishu Settings
|
||||
feishu_app_id: str = ""
|
||||
feishu_app_secret: str = ""
|
||||
feishu_verification_token: str = ""
|
||||
feishu_encrypt_key: str = ""
|
||||
feishu_webhook_port: int = 9000
|
||||
feishu_allowed_senders: str = ""
|
||||
feishu_domain: str = "https://open.feishu.cn"
|
||||
feishu_proxy: str = ""
|
||||
|
||||
# WeChat Settings
|
||||
wechat_backend: str = "wecom"
|
||||
wechat_webhook_port: int = 9001
|
||||
wechat_allowed_senders: str = ""
|
||||
wechat_proxy: str = ""
|
||||
wechat_wecom_corp_id: str = ""
|
||||
wechat_wecom_agent_id: str = ""
|
||||
wechat_wecom_secret: str = ""
|
||||
wechat_wecom_token: str = ""
|
||||
wechat_wecom_encoding_aes_key: str = ""
|
||||
wechat_mp_app_id: str = ""
|
||||
wechat_mp_app_secret: str = ""
|
||||
wechat_mp_token: str = ""
|
||||
wechat_mp_encoding_aes_key: str = ""
|
||||
|
||||
# DingTalk Settings
|
||||
dingtalk_client_id: str = ""
|
||||
dingtalk_client_secret: str = ""
|
||||
dingtalk_allowed_senders: str = ""
|
||||
dingtalk_proxy: str = ""
|
||||
|
||||
# Email Settings
|
||||
email_imap_host: str = ""
|
||||
email_imap_port: int = 993
|
||||
email_imap_username: str = ""
|
||||
email_imap_password: str = ""
|
||||
email_imap_mailbox: str = "INBOX"
|
||||
email_imap_use_ssl: bool = True
|
||||
email_smtp_host: str = ""
|
||||
email_smtp_port: int = 587
|
||||
email_smtp_username: str = ""
|
||||
email_smtp_password: str = ""
|
||||
email_smtp_use_tls: bool = True
|
||||
email_from_address: str = ""
|
||||
email_poll_interval: int = 30
|
||||
email_mark_seen: bool = True
|
||||
email_max_body_chars: int = 12000
|
||||
email_subject_prefix: str = "Re: "
|
||||
email_allowed_senders: str = ""
|
||||
|
||||
# QQ Settings
|
||||
qq_app_id: str = ""
|
||||
qq_app_secret: str = ""
|
||||
qq_allowed_senders: str = ""
|
||||
|
||||
# Signal Settings
|
||||
signal_phone_number: str = ""
|
||||
signal_cli_path: str = "signal-cli"
|
||||
signal_config_dir: str = ""
|
||||
signal_allowed_senders: str = ""
|
||||
signal_rpc_port: int = 7583
|
||||
|
||||
# Shared webhook port (0 = disabled)
|
||||
shared_webhook_port: int = 9000
|
||||
|
||||
# DM access control policy
|
||||
dm_policy: str = "allowlist"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
|
||||
@@ -4,6 +4,8 @@ Re-exports middleware classes and factory functions so that existing
|
||||
``from EvoScientist.middleware import X`` imports continue to work.
|
||||
"""
|
||||
|
||||
from deepagents.middleware.skills import SkillsMiddleware
|
||||
|
||||
from .memory import (
|
||||
EvoMemoryMiddleware,
|
||||
EvoMemoryState,
|
||||
@@ -11,9 +13,30 @@ from .memory import (
|
||||
create_memory_middleware,
|
||||
)
|
||||
|
||||
|
||||
def create_skills_middleware(composite_backend) -> SkillsMiddleware:
|
||||
"""Create a SkillsMiddleware that loads skills.
|
||||
|
||||
Uses the CompositeBackend directly so that skill paths in the system
|
||||
prompt match the ``/skills/`` route (e.g. ``/skills/find-skills/SKILL.md``).
|
||||
|
||||
Args:
|
||||
composite_backend: The CompositeBackend that routes ``/skills/`` to
|
||||
the MergedReadOnlyBackend.
|
||||
|
||||
Returns:
|
||||
Configured SkillsMiddleware instance
|
||||
"""
|
||||
return SkillsMiddleware(
|
||||
backend=composite_backend,
|
||||
sources=["/skills/"],
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EvoMemoryMiddleware",
|
||||
"EvoMemoryState",
|
||||
"ExtractedMemory",
|
||||
"create_memory_middleware",
|
||||
"create_skills_middleware",
|
||||
]
|
||||
|
||||
@@ -279,3 +279,7 @@ def get_system_prompt(max_concurrent: int = 3, max_iterations: int = 3) -> str:
|
||||
max_iterations=max_iterations,
|
||||
)
|
||||
return EXPERIMENT_WORKFLOW + "\n" + delegation
|
||||
|
||||
|
||||
# Default export (backward compatible)
|
||||
SYSTEM_PROMPT = get_system_prompt()
|
||||
|
||||
@@ -86,7 +86,7 @@ class StreamEventEmitter:
|
||||
@staticmethod
|
||||
def done(response: str = "") -> StreamEvent:
|
||||
"""Done event."""
|
||||
return StreamEvent("done", {"type": "done", "response": response})
|
||||
return StreamEvent("done", {"type": "done", "content": response, "response": response})
|
||||
|
||||
@staticmethod
|
||||
def error(message: str) -> StreamEvent:
|
||||
|
||||
@@ -4,6 +4,9 @@ Async generator that streams events from an agent graph,
|
||||
plus helpers for processing AI message chunks and tool results.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import mimetypes
|
||||
import os
|
||||
from typing import Any, AsyncIterator
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk # type: ignore[import-untyped]
|
||||
@@ -64,6 +67,7 @@ async def stream_agent_events(
|
||||
message: str,
|
||||
thread_id: str,
|
||||
metadata: dict | None = None,
|
||||
media: list[str] | None = None,
|
||||
) -> AsyncIterator[dict]:
|
||||
"""Stream events from the agent graph using async iteration.
|
||||
|
||||
@@ -75,6 +79,7 @@ async def stream_agent_events(
|
||||
thread_id: Thread ID for conversation persistence
|
||||
metadata: Optional metadata dict merged into the LangGraph config
|
||||
(e.g. agent_name, updated_at for checkpoint persistence).
|
||||
media: Optional list of local file paths for attachments.
|
||||
|
||||
Yields:
|
||||
Event dicts: thinking, text, tool_call, tool_result,
|
||||
@@ -247,9 +252,41 @@ async def stream_agent_events(
|
||||
# 4) No real names available yet -- return generic WITHOUT caching
|
||||
return "sub-agent"
|
||||
|
||||
# Build user message content: text + inline images + file path references
|
||||
user_content: str | list[dict[str, Any]] = message
|
||||
if media:
|
||||
_IMAGE_EXTS = frozenset({".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp"})
|
||||
_MAX_INLINE_SIZE = 5 * 1024 * 1024 # 5 MB
|
||||
content_blocks: list[dict[str, Any]] = []
|
||||
if message:
|
||||
content_blocks.append({"type": "text", "text": message})
|
||||
file_refs: list[str] = []
|
||||
for path in media:
|
||||
ext = os.path.splitext(path)[1].lower()
|
||||
if ext in _IMAGE_EXTS and os.path.isfile(path):
|
||||
fsize = os.path.getsize(path)
|
||||
if fsize <= _MAX_INLINE_SIZE:
|
||||
mime = mimetypes.guess_type(path)[0] or "image/png"
|
||||
with open(path, "rb") as fh:
|
||||
b64 = base64.b64encode(fh.read()).decode("ascii")
|
||||
content_blocks.append({"type": "image_url", "image_url": {
|
||||
"url": f"data:{mime};base64,{b64}",
|
||||
}})
|
||||
else:
|
||||
file_refs.append(path)
|
||||
else:
|
||||
file_refs.append(path)
|
||||
if file_refs:
|
||||
ref_text = "\n".join(
|
||||
f"[attached file: {os.path.basename(p)}] path: {p}" for p in file_refs
|
||||
)
|
||||
content_blocks.append({"type": "text", "text": ref_text})
|
||||
if content_blocks:
|
||||
user_content = content_blocks
|
||||
|
||||
try:
|
||||
async for chunk in agent.astream(
|
||||
{"messages": [{"role": "user", "content": message}]},
|
||||
{"messages": [{"role": "user", "content": user_content}]},
|
||||
config=config,
|
||||
stream_mode="messages",
|
||||
subgraphs=True,
|
||||
|
||||
@@ -68,10 +68,14 @@ class ToolResultFormatter:
|
||||
|
||||
return ContentType.TEXT
|
||||
|
||||
def is_success(self, content: str) -> bool:
|
||||
"""Check if content indicates successful execution."""
|
||||
return _is_success(content)
|
||||
|
||||
def format(self, name: str, content: str, max_length: int = 800) -> FormattedResult:
|
||||
"""Format tool result based on detected content type."""
|
||||
content_type = self.detect_type(content)
|
||||
success = _is_success(content)
|
||||
success = self.is_success(content)
|
||||
|
||||
formatter_map = {
|
||||
ContentType.SUCCESS: self._format_success,
|
||||
|
||||
@@ -6,11 +6,13 @@ to work unchanged thanks to these re-exports.
|
||||
|
||||
from .search import tavily_search, fetch_webpage_content
|
||||
from .think import think_tool
|
||||
from .image import view_image
|
||||
from .skill_manager import skill_manager
|
||||
|
||||
__all__ = [
|
||||
"tavily_search",
|
||||
"fetch_webpage_content",
|
||||
"think_tool",
|
||||
"view_image",
|
||||
"skill_manager",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
"""Image viewing tool."""
|
||||
|
||||
import base64
|
||||
import mimetypes
|
||||
import os
|
||||
|
||||
from langchain_core.tools import tool
|
||||
|
||||
from ..paths import resolve_virtual_path
|
||||
|
||||
# Supported image extensions and their MIME types
|
||||
_IMAGE_EXTENSIONS = {
|
||||
".png": "image/png",
|
||||
".jpg": "image/jpeg",
|
||||
".jpeg": "image/jpeg",
|
||||
".gif": "image/gif",
|
||||
".webp": "image/webp",
|
||||
".bmp": "image/bmp",
|
||||
".svg": "image/svg+xml",
|
||||
}
|
||||
|
||||
# Max file size for image viewing (5MB)
|
||||
_MAX_IMAGE_SIZE = 5 * 1024 * 1024
|
||||
|
||||
|
||||
@tool(parse_docstring=True)
|
||||
def view_image(image_path: str) -> "list | str":
|
||||
"""View and analyze an image file.
|
||||
|
||||
Use this tool when you need to see the visual content of an image file
|
||||
(PNG, JPEG, GIF, WebP). The image will be displayed so you can describe,
|
||||
analyze, or answer questions about it.
|
||||
|
||||
Note: Use this instead of read_file for image files. read_file only
|
||||
returns binary data, while view_image lets you actually see the image.
|
||||
|
||||
Args:
|
||||
image_path: Path to the image file (relative to workspace or absolute)
|
||||
|
||||
Returns:
|
||||
Image content blocks that the model can visually process
|
||||
"""
|
||||
# Resolve virtual workspace paths: /image.png → {workspace}/image.png
|
||||
resolved = image_path
|
||||
if not os.path.isfile(resolved):
|
||||
resolved = str(resolve_virtual_path(image_path))
|
||||
|
||||
if not os.path.isfile(resolved):
|
||||
return f"Error: File not found: {image_path}"
|
||||
image_path = resolved
|
||||
|
||||
ext = os.path.splitext(image_path)[1].lower()
|
||||
mime_type = _IMAGE_EXTENSIONS.get(ext)
|
||||
if not mime_type:
|
||||
# Fallback to mimetypes module
|
||||
mime_type, _ = mimetypes.guess_type(image_path)
|
||||
if not mime_type or not mime_type.startswith("image/"):
|
||||
return f"Error: Not a supported image format: {ext}"
|
||||
|
||||
file_size = os.path.getsize(image_path)
|
||||
if file_size > _MAX_IMAGE_SIZE:
|
||||
size_mb = file_size / (1024 * 1024)
|
||||
return f"Error: Image too large ({size_mb:.1f}MB). Max is 5MB."
|
||||
|
||||
with open(image_path, "rb") as f:
|
||||
data = base64.b64encode(f.read()).decode("ascii")
|
||||
|
||||
size_kb = file_size / 1024
|
||||
filename = os.path.basename(image_path)
|
||||
|
||||
return [
|
||||
{"type": "text", "text": f"Image: {filename} ({size_kb:.0f}KB, {mime_type})"},
|
||||
{"type": "image", "base64": data, "mime_type": mime_type},
|
||||
]
|
||||
@@ -78,6 +78,10 @@ def format_messages(messages):
|
||||
console.print(Panel(content, title=f"📝 {msg_type}", border_style="white"))
|
||||
|
||||
|
||||
def format_message(messages):
|
||||
"""Alias for format_messages for backward compatibility."""
|
||||
return format_messages(messages)
|
||||
|
||||
|
||||
def show_prompt(prompt_text: str, title: str = "Prompt", border_style: str = "blue"):
|
||||
"""Display a prompt with rich formatting and XML tag highlighting.
|
||||
|
||||
@@ -473,6 +473,15 @@ We thank the authors for their valuable contributions to the open-source communi
|
||||
<sub><b>Dinos Papakostas</b></sub>
|
||||
</a>
|
||||
</td>
|
||||
<td align="center">
|
||||
<a href="https://muxincg2004.github.io/">
|
||||
<img src="https://muxincg2004.github.io/resume_avatar.jpg"
|
||||
width="100" height="100"
|
||||
style="object-fit: cover; border-radius: 20%;" alt="Ziheng Zhang"/>
|
||||
<br />
|
||||
<sub><b>Ziheng Zhang</b></sub>
|
||||
</a>
|
||||
</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
@@ -44,6 +44,17 @@ dev = [
|
||||
"ruff>=0.5",
|
||||
"build>=1.0",
|
||||
]
|
||||
telegram = ["python-telegram-bot>=21.0"]
|
||||
discord = ["discord.py>=2.3"]
|
||||
slack = ["slack-sdk>=3.27", "aiohttp>=3.9"]
|
||||
wechat = ["pycryptodome>=3.20"]
|
||||
all-channels = [
|
||||
"python-telegram-bot>=21.0",
|
||||
"discord.py>=2.3",
|
||||
"aiohttp>=3.9",
|
||||
"slack-sdk>=3.27",
|
||||
"pycryptodome>=3.20",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
"Homepage" = "https://github.com/EvoScientist/EvoScientist"
|
||||
|
||||
@@ -0,0 +1,286 @@
|
||||
"""Tests for bus-mode agent integration (_bus_inbound_consumer)."""
|
||||
|
||||
import asyncio
|
||||
|
||||
|
||||
from EvoScientist.channels.bus.events import InboundMessage
|
||||
from EvoScientist.channels.bus.message_bus import MessageBus
|
||||
from EvoScientist.channels.channel_manager import ChannelManager
|
||||
from EvoScientist.channels.base import Channel, OutgoingMessage
|
||||
|
||||
|
||||
def _run(coro):
|
||||
"""Run an async coroutine safely, creating a fresh event loop."""
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
return loop.run_until_complete(coro)
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
|
||||
class _FakeConfig:
|
||||
text_chunk_limit = 4096
|
||||
allowed_senders = None
|
||||
|
||||
|
||||
class FakeChannel(Channel):
|
||||
"""Minimal channel for bus integration testing."""
|
||||
|
||||
name = "fake"
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(_FakeConfig())
|
||||
self._started = False
|
||||
self._stopped = False
|
||||
self._sent: list[OutgoingMessage] = []
|
||||
|
||||
async def start(self):
|
||||
self._started = True
|
||||
|
||||
async def stop(self):
|
||||
self._stopped = True
|
||||
|
||||
async def receive(self):
|
||||
while True:
|
||||
try:
|
||||
msg = await asyncio.wait_for(self._queue.get(), timeout=0.5)
|
||||
yield msg
|
||||
except asyncio.TimeoutError:
|
||||
return
|
||||
|
||||
async def send(self, message: OutgoingMessage) -> bool:
|
||||
self._sent.append(message)
|
||||
return True
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
pass
|
||||
|
||||
|
||||
def _mock_stream_events(content, reply):
|
||||
"""Create a mock stream_agent_events that yields text then done."""
|
||||
async def _stream(agent, message, thread_id):
|
||||
yield {"type": "text", "content": reply}
|
||||
yield {"type": "done", "response": reply}
|
||||
return _stream
|
||||
|
||||
|
||||
def _mock_stream_events_error(error_msg):
|
||||
"""Create a mock stream_agent_events that raises."""
|
||||
async def _stream(agent, message, thread_id):
|
||||
raise RuntimeError(error_msg)
|
||||
yield # make it an async generator # pragma: no cover
|
||||
return _stream
|
||||
|
||||
|
||||
def _mock_stream_events_with_thinking(thinking_text, reply):
|
||||
"""Create a mock stream_agent_events that yields thinking then done."""
|
||||
async def _stream(agent, message, thread_id):
|
||||
yield {"type": "thinking", "content": thinking_text}
|
||||
yield {"type": "text", "content": reply}
|
||||
yield {"type": "done", "content": reply}
|
||||
return _stream
|
||||
|
||||
|
||||
class TestBusInboundConsumer:
|
||||
"""Test the _bus_inbound_consumer bridge function."""
|
||||
|
||||
def test_processes_inbound_and_publishes_outbound(self):
|
||||
"""InboundMessage -> agent -> OutboundMessage flow."""
|
||||
from EvoScientist.cli.channel import _bus_inbound_consumer
|
||||
|
||||
async def _test():
|
||||
bus = MessageBus()
|
||||
manager = ChannelManager(bus)
|
||||
ch = FakeChannel()
|
||||
manager.register(ch)
|
||||
|
||||
mock_stream = _mock_stream_events(
|
||||
"hello agent", "Reply to: hello agent",
|
||||
)
|
||||
|
||||
import EvoScientist.stream.events as events_mod
|
||||
original = events_mod.stream_agent_events
|
||||
events_mod.stream_agent_events = mock_stream
|
||||
|
||||
try:
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(bus, manager, None, "test-thread", False)
|
||||
)
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="hello agent",
|
||||
))
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
outbound = await asyncio.wait_for(
|
||||
bus.consume_outbound(), timeout=2.0,
|
||||
)
|
||||
assert outbound.channel == "fake"
|
||||
assert outbound.chat_id == "chat1"
|
||||
assert "Reply to: hello agent" in outbound.content
|
||||
|
||||
consumer.cancel()
|
||||
try:
|
||||
await consumer
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
finally:
|
||||
events_mod.stream_agent_events = original
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_agent_error_publishes_error_outbound(self):
|
||||
"""When agent raises, an error message is published outbound."""
|
||||
from EvoScientist.cli.channel import _bus_inbound_consumer
|
||||
|
||||
async def _test():
|
||||
bus = MessageBus()
|
||||
manager = ChannelManager(bus)
|
||||
ch = FakeChannel()
|
||||
manager.register(ch)
|
||||
|
||||
mock_stream = _mock_stream_events_error("agent crashed")
|
||||
|
||||
import EvoScientist.stream.events as events_mod
|
||||
original = events_mod.stream_agent_events
|
||||
events_mod.stream_agent_events = mock_stream
|
||||
|
||||
try:
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(bus, manager, None, "test-thread", False)
|
||||
)
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="crash me",
|
||||
))
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
outbound = await asyncio.wait_for(
|
||||
bus.consume_outbound(), timeout=2.0,
|
||||
)
|
||||
assert outbound.channel == "fake"
|
||||
assert "Error" in outbound.content or "error" in outbound.content.lower()
|
||||
assert "agent crashed" in outbound.content
|
||||
|
||||
consumer.cancel()
|
||||
try:
|
||||
await consumer
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
finally:
|
||||
events_mod.stream_agent_events = original
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_message_counting(self):
|
||||
"""Messages are counted via record_message."""
|
||||
from EvoScientist.cli.channel import _bus_inbound_consumer
|
||||
|
||||
async def _test():
|
||||
bus = MessageBus()
|
||||
manager = ChannelManager(bus)
|
||||
ch = FakeChannel()
|
||||
manager.register(ch)
|
||||
|
||||
mock_stream = _mock_stream_events("test", "ok")
|
||||
|
||||
import EvoScientist.stream.events as events_mod
|
||||
original = events_mod.stream_agent_events
|
||||
events_mod.stream_agent_events = mock_stream
|
||||
|
||||
try:
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(bus, manager, None, "test-thread", False)
|
||||
)
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="u1",
|
||||
chat_id="c1",
|
||||
content="test",
|
||||
))
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
|
||||
|
||||
assert manager._message_counts["fake"]["received"] == 1
|
||||
assert manager._message_counts["fake"]["sent"] == 1
|
||||
|
||||
consumer.cancel()
|
||||
try:
|
||||
await consumer
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
finally:
|
||||
events_mod.stream_agent_events = original
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_thinking_sent_to_channel(self):
|
||||
"""Thinking messages are sent to the channel when show_thinking=True."""
|
||||
from EvoScientist.cli.channel import _bus_inbound_consumer
|
||||
|
||||
async def _test():
|
||||
bus = MessageBus()
|
||||
manager = ChannelManager(bus)
|
||||
ch = FakeChannel()
|
||||
server = manager.register(ch)
|
||||
server.send_thinking = True
|
||||
|
||||
long_thinking = "A" * 250 # >= _MIN_THINKING_LEN (200)
|
||||
mock_stream = _mock_stream_events_with_thinking(
|
||||
long_thinking, "final answer",
|
||||
)
|
||||
|
||||
import EvoScientist.stream.events as events_mod
|
||||
original = events_mod.stream_agent_events
|
||||
events_mod.stream_agent_events = mock_stream
|
||||
|
||||
try:
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(
|
||||
bus, manager, None, "test-thread", True,
|
||||
)
|
||||
)
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="think about this",
|
||||
metadata={"chat_id": "chat1"},
|
||||
))
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
# Drain outbound (final answer)
|
||||
outbound = await asyncio.wait_for(
|
||||
bus.consume_outbound(), timeout=2.0,
|
||||
)
|
||||
assert "final answer" in outbound.content
|
||||
|
||||
# Check that thinking was sent via channel.send
|
||||
thinking_msgs = [
|
||||
m for m in ch._sent
|
||||
if "\U0001f9e0" in m.content
|
||||
]
|
||||
assert len(thinking_msgs) == 1
|
||||
assert long_thinking in thinking_msgs[0].content
|
||||
|
||||
consumer.cancel()
|
||||
try:
|
||||
await consumer
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
finally:
|
||||
events_mod.stream_agent_events = original
|
||||
|
||||
_run(_test())
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,166 @@
|
||||
"""Tests for ChannelManager."""
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.channels.bus.message_bus import MessageBus
|
||||
from EvoScientist.channels.channel_manager import ChannelManager
|
||||
from EvoScientist.channels.base import Channel, OutboundMessage
|
||||
|
||||
|
||||
def _run(coro):
|
||||
"""Run an async coroutine safely, creating a fresh event loop."""
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
return loop.run_until_complete(coro)
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
|
||||
class _FakeConfig:
|
||||
text_chunk_limit = 4096
|
||||
allowed_senders = None
|
||||
|
||||
|
||||
class FakeChannel(Channel):
|
||||
"""Minimal channel for testing."""
|
||||
|
||||
name = "fake"
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(_FakeConfig())
|
||||
self._started = False
|
||||
self._stopped = False
|
||||
self._sent: list[OutboundMessage] = []
|
||||
|
||||
async def start(self):
|
||||
self._started = True
|
||||
|
||||
async def stop(self):
|
||||
self._stopped = True
|
||||
|
||||
async def receive(self):
|
||||
while True:
|
||||
try:
|
||||
msg = await asyncio.wait_for(
|
||||
self._queue.get(), timeout=0.5,
|
||||
)
|
||||
yield msg
|
||||
except asyncio.TimeoutError:
|
||||
return
|
||||
|
||||
async def send(self, message: OutboundMessage) -> bool:
|
||||
self._sent.append(message)
|
||||
return True
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
pass
|
||||
|
||||
|
||||
class TestChannelManagerRegister:
|
||||
def test_register_channel(self):
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
ch = FakeChannel()
|
||||
result = mgr.register(ch)
|
||||
assert "fake" in mgr.enabled_channels
|
||||
assert mgr.get_channel("fake") is ch
|
||||
assert result is ch
|
||||
|
||||
def test_duplicate_register_raises(self):
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
mgr.register(FakeChannel())
|
||||
with pytest.raises(ValueError, match="already registered"):
|
||||
mgr.register(FakeChannel())
|
||||
|
||||
def test_get_status(self):
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
mgr.register(FakeChannel())
|
||||
status = mgr.get_status()
|
||||
assert "fake" in status
|
||||
assert status["fake"]["registered"] is True
|
||||
|
||||
|
||||
class TestChannelManagerDispatch:
|
||||
def test_outbound_dispatch_routes_to_channel(self):
|
||||
async def _test():
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
ch = FakeChannel()
|
||||
mgr.register(ch)
|
||||
|
||||
# Start only the dispatcher (not full start_all)
|
||||
dispatch = asyncio.create_task(
|
||||
mgr._dispatch_outbound()
|
||||
)
|
||||
|
||||
# Publish an outbound message
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="fake", chat_id="u1",
|
||||
content="hello from agent",
|
||||
))
|
||||
|
||||
await asyncio.sleep(0.1)
|
||||
dispatch.cancel()
|
||||
try:
|
||||
await dispatch
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
assert len(ch._sent) == 1
|
||||
assert ch._sent[0].content == "hello from agent"
|
||||
assert ch._sent[0].chat_id == "u1"
|
||||
|
||||
_run(_test())
|
||||
|
||||
|
||||
class TestChannelManagerTracking:
|
||||
def test_record_message(self):
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
mgr.register(FakeChannel())
|
||||
|
||||
mgr.record_message("fake", "received")
|
||||
mgr.record_message("fake", "received")
|
||||
mgr.record_message("fake", "sent")
|
||||
|
||||
assert mgr._message_counts["fake"]["received"] == 2
|
||||
assert mgr._message_counts["fake"]["sent"] == 1
|
||||
|
||||
def test_record_message_unknown_channel(self):
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
|
||||
# Should not raise, auto-creates entry
|
||||
mgr.record_message("unknown", "received")
|
||||
assert mgr._message_counts["unknown"]["received"] == 1
|
||||
|
||||
def test_get_detailed_status(self):
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
mgr.register(FakeChannel())
|
||||
|
||||
# Simulate start_all setting start_times
|
||||
from datetime import datetime
|
||||
mgr._start_times["fake"] = datetime.now()
|
||||
mgr._message_counts["fake"] = {"received": 5, "sent": 3}
|
||||
|
||||
status = mgr.get_detailed_status()
|
||||
assert "fake" in status
|
||||
assert status["fake"]["registered"] is True
|
||||
assert status["fake"]["received"] == 5
|
||||
assert status["fake"]["sent"] == 3
|
||||
assert status["fake"]["uptime_seconds"] >= 0
|
||||
assert status["fake"]["start_time"] is not None
|
||||
|
||||
def test_get_detailed_status_no_start_time(self):
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
mgr.register(FakeChannel())
|
||||
|
||||
status = mgr.get_detailed_status()
|
||||
assert status["fake"]["uptime_seconds"] == 0
|
||||
assert status["fake"]["start_time"] is None
|
||||
@@ -0,0 +1,71 @@
|
||||
"""Tests for Discord channel implementation."""
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.channels.discord.channel import DiscordChannel, DiscordConfig
|
||||
from EvoScientist.channels.base import ChannelError
|
||||
|
||||
|
||||
def _run(coro):
|
||||
"""Run an async coroutine safely, creating a fresh event loop."""
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
return loop.run_until_complete(coro)
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
|
||||
class TestDiscordConfig:
|
||||
def test_default_values(self):
|
||||
config = DiscordConfig()
|
||||
assert config.bot_token == ""
|
||||
assert config.allowed_senders is None
|
||||
assert config.allowed_channels is None
|
||||
assert config.text_chunk_limit == 4096
|
||||
|
||||
def test_custom_values(self):
|
||||
config = DiscordConfig(
|
||||
bot_token="test-token",
|
||||
allowed_senders={"111"},
|
||||
allowed_channels={"222"},
|
||||
text_chunk_limit=1000,
|
||||
)
|
||||
assert config.bot_token == "test-token"
|
||||
assert config.allowed_senders == {"111"}
|
||||
assert config.allowed_channels == {"222"}
|
||||
assert config.text_chunk_limit == 1000
|
||||
|
||||
|
||||
class TestDiscordChannel:
|
||||
def test_init(self):
|
||||
config = DiscordConfig(bot_token="test")
|
||||
channel = DiscordChannel(config)
|
||||
assert channel.config is config
|
||||
assert channel._running is False
|
||||
|
||||
def test_start_raises_without_token_or_library(self):
|
||||
config = DiscordConfig(bot_token="")
|
||||
channel = DiscordChannel(config)
|
||||
with pytest.raises(ChannelError):
|
||||
_run(channel.start())
|
||||
|
||||
def test_stop_when_not_running(self):
|
||||
config = DiscordConfig(bot_token="test")
|
||||
channel = DiscordChannel(config)
|
||||
_run(channel.stop())
|
||||
|
||||
def test_send_returns_false_without_client(self):
|
||||
from EvoScientist.channels.base import OutboundMessage
|
||||
|
||||
config = DiscordConfig(bot_token="test")
|
||||
channel = DiscordChannel(config)
|
||||
msg = OutboundMessage(
|
||||
channel="discord",
|
||||
chat_id="123",
|
||||
content="hello",
|
||||
metadata={"chat_id": "123"},
|
||||
)
|
||||
result = _run(channel.send(msg))
|
||||
assert result is False
|
||||
@@ -0,0 +1,541 @@
|
||||
"""Tests for unified mention gating in base.py and per-channel _strip_mention."""
|
||||
|
||||
import asyncio
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
from EvoScientist.channels.base import Channel, RawIncoming
|
||||
from EvoScientist.channels.capabilities import ChannelCapabilities
|
||||
|
||||
|
||||
# ── Minimal concrete channel for testing base-class logic ─────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class _StubConfig:
|
||||
allowed_senders: set[str] | None = None
|
||||
require_mention: str = "group"
|
||||
dm_policy: str = "allowlist"
|
||||
|
||||
|
||||
class _StubChannel(Channel):
|
||||
name = "stub"
|
||||
|
||||
async def start(self):
|
||||
pass
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
pass
|
||||
|
||||
|
||||
class _MentionStubChannel(_StubChannel):
|
||||
"""Stub with mention gating enabled."""
|
||||
capabilities = ChannelCapabilities(mentions=True)
|
||||
|
||||
|
||||
def _run(coro):
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
return loop.run_until_complete(coro)
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
|
||||
# ── _should_process tests ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestShouldProcess:
|
||||
"""Tests for the centralized _should_process gate."""
|
||||
|
||||
def _make(self, require_mention="group"):
|
||||
config = _StubConfig(require_mention=require_mention)
|
||||
return _StubChannel(config)
|
||||
|
||||
def test_dm_always_passes_with_group_mode(self):
|
||||
ch = self._make("group")
|
||||
raw = RawIncoming(sender_id="u1", chat_id="c1", is_group=False, was_mentioned=False)
|
||||
assert ch._should_process(raw) is True
|
||||
|
||||
def test_dm_always_passes_with_always_mode(self):
|
||||
ch = self._make("always")
|
||||
raw = RawIncoming(sender_id="u1", chat_id="c1", is_group=False, was_mentioned=False)
|
||||
assert ch._should_process(raw) is True
|
||||
|
||||
def test_dm_always_passes_with_off_mode(self):
|
||||
ch = self._make("off")
|
||||
raw = RawIncoming(sender_id="u1", chat_id="c1", is_group=False, was_mentioned=False)
|
||||
assert ch._should_process(raw) is True
|
||||
|
||||
def test_group_mentioned_passes_with_group_mode(self):
|
||||
ch = self._make("group")
|
||||
raw = RawIncoming(sender_id="u1", chat_id="c1", is_group=True, was_mentioned=True)
|
||||
assert ch._should_process(raw) is True
|
||||
|
||||
def test_group_not_mentioned_blocked_with_group_mode(self):
|
||||
ch = self._make("group")
|
||||
raw = RawIncoming(sender_id="u1", chat_id="c1", is_group=True, was_mentioned=False)
|
||||
assert ch._should_process(raw) is False
|
||||
|
||||
def test_group_not_mentioned_passes_with_off_mode(self):
|
||||
ch = self._make("off")
|
||||
raw = RawIncoming(sender_id="u1", chat_id="c1", is_group=True, was_mentioned=False)
|
||||
assert ch._should_process(raw) is True
|
||||
|
||||
def test_group_not_mentioned_blocked_with_always_mode(self):
|
||||
ch = self._make("always")
|
||||
raw = RawIncoming(sender_id="u1", chat_id="c1", is_group=True, was_mentioned=False)
|
||||
assert ch._should_process(raw) is False
|
||||
|
||||
def test_group_mentioned_passes_with_always_mode(self):
|
||||
ch = self._make("always")
|
||||
raw = RawIncoming(sender_id="u1", chat_id="c1", is_group=True, was_mentioned=True)
|
||||
assert ch._should_process(raw) is True
|
||||
|
||||
|
||||
class TestPipelineGating:
|
||||
"""Tests that _enqueue_raw pipeline integrates mention gating and _strip_mention."""
|
||||
|
||||
def test_group_not_mentioned_dropped(self):
|
||||
async def _test():
|
||||
config = _StubConfig(require_mention="group")
|
||||
ch = _MentionStubChannel(config)
|
||||
raw = RawIncoming(
|
||||
sender_id="u1", chat_id="c1", text="hello",
|
||||
is_group=True, was_mentioned=False,
|
||||
)
|
||||
await ch._enqueue_raw(raw)
|
||||
assert ch._queue.qsize() == 0
|
||||
_run(_test())
|
||||
|
||||
def test_group_mentioned_passes(self):
|
||||
async def _test():
|
||||
config = _StubConfig(require_mention="group")
|
||||
ch = _MentionStubChannel(config)
|
||||
raw = RawIncoming(
|
||||
sender_id="u1", chat_id="c1", text="hello",
|
||||
is_group=True, was_mentioned=True,
|
||||
)
|
||||
await ch._enqueue_raw(raw)
|
||||
assert ch._queue.qsize() == 1
|
||||
msg = await ch._queue.get()
|
||||
assert msg.content == "hello"
|
||||
_run(_test())
|
||||
|
||||
def test_dm_passes_even_when_not_mentioned(self):
|
||||
async def _test():
|
||||
config = _StubConfig(require_mention="group")
|
||||
ch = _MentionStubChannel(config)
|
||||
raw = RawIncoming(
|
||||
sender_id="u1", chat_id="c1", text="hello",
|
||||
is_group=False, was_mentioned=False,
|
||||
)
|
||||
await ch._enqueue_raw(raw)
|
||||
assert ch._queue.qsize() == 1
|
||||
_run(_test())
|
||||
|
||||
def test_strip_mention_called_for_group(self):
|
||||
"""When is_group=True, _strip_mention should be applied to text."""
|
||||
async def _test():
|
||||
config = _StubConfig(require_mention="group")
|
||||
ch = _MentionStubChannel(config)
|
||||
ch._strip_mention = lambda text: text.replace("@bot ", "").strip()
|
||||
# Rebuild middlewares so MentionGatingMiddleware picks up new strip_fn
|
||||
ch._inbound_middlewares = ch._build_inbound_middlewares()
|
||||
|
||||
raw = RawIncoming(
|
||||
sender_id="u1", chat_id="c1", text="@bot hello",
|
||||
is_group=True, was_mentioned=True,
|
||||
)
|
||||
await ch._enqueue_raw(raw)
|
||||
assert ch._queue.qsize() == 1
|
||||
msg = await ch._queue.get()
|
||||
assert msg.content == "hello"
|
||||
_run(_test())
|
||||
|
||||
def test_strip_mention_not_called_for_dm(self):
|
||||
"""When is_group=False, _strip_mention should NOT be applied."""
|
||||
async def _test():
|
||||
config = _StubConfig(require_mention="group")
|
||||
ch = _MentionStubChannel(config)
|
||||
ch._strip_mention = lambda text: text.replace("@bot ", "").strip()
|
||||
ch._inbound_middlewares = ch._build_inbound_middlewares()
|
||||
|
||||
raw = RawIncoming(
|
||||
sender_id="u1", chat_id="c1", text="@bot hello",
|
||||
is_group=False, was_mentioned=True,
|
||||
)
|
||||
await ch._enqueue_raw(raw)
|
||||
assert ch._queue.qsize() == 1
|
||||
msg = await ch._queue.get()
|
||||
assert msg.content == "@bot hello"
|
||||
_run(_test())
|
||||
|
||||
|
||||
# ── Per-channel _strip_mention tests ─────────────────────────────────
|
||||
|
||||
|
||||
class TestTelegramStripMention:
|
||||
def test_strip_username(self):
|
||||
from EvoScientist.channels.telegram.channel import TelegramChannel, TelegramConfig
|
||||
|
||||
config = TelegramConfig(bot_token="test")
|
||||
ch = TelegramChannel(config)
|
||||
ch._bot_username = "mybot"
|
||||
assert ch._strip_mention("@mybot hello") == "hello"
|
||||
assert ch._strip_mention("@MyBot hello") == "hello"
|
||||
assert ch._strip_mention("hello @mybot world") == "hello world"
|
||||
|
||||
def test_no_username(self):
|
||||
from EvoScientist.channels.telegram.channel import TelegramChannel, TelegramConfig
|
||||
|
||||
config = TelegramConfig(bot_token="test")
|
||||
ch = TelegramChannel(config)
|
||||
ch._bot_username = ""
|
||||
assert ch._strip_mention("@mybot hello") == "@mybot hello"
|
||||
|
||||
|
||||
class TestDiscordStripMention:
|
||||
def test_strip_user_id(self):
|
||||
from EvoScientist.channels.discord.channel import DiscordChannel, DiscordConfig
|
||||
|
||||
config = DiscordConfig(bot_token="test")
|
||||
ch = DiscordChannel(config)
|
||||
|
||||
# Mock client.user
|
||||
class FakeUser:
|
||||
id = 123456789
|
||||
class FakeClient:
|
||||
user = FakeUser()
|
||||
ch._client = FakeClient()
|
||||
|
||||
assert ch._strip_mention("<@123456789> hello") == "hello"
|
||||
assert ch._strip_mention("<@!123456789> hello") == "hello"
|
||||
assert ch._strip_mention("hello <@123456789> world") == "hello world"
|
||||
|
||||
def test_no_client(self):
|
||||
from EvoScientist.channels.discord.channel import DiscordChannel, DiscordConfig
|
||||
|
||||
config = DiscordConfig(bot_token="test")
|
||||
ch = DiscordChannel(config)
|
||||
ch._client = None
|
||||
assert ch._strip_mention("<@123> hello") == "<@123> hello"
|
||||
|
||||
|
||||
class TestSlackStripMention:
|
||||
def test_strip_user_id(self):
|
||||
from EvoScientist.channels.slack.channel import SlackChannel, SlackConfig
|
||||
|
||||
config = SlackConfig(bot_token="test", app_token="xapp-test")
|
||||
ch = SlackChannel(config)
|
||||
ch._bot_user_id = "U123ABC"
|
||||
assert ch._strip_mention("<@U123ABC> hello") == "hello"
|
||||
assert ch._strip_mention("hello <@U123ABC> world") == "hello world"
|
||||
|
||||
def test_no_bot_user_id(self):
|
||||
from EvoScientist.channels.slack.channel import SlackChannel, SlackConfig
|
||||
|
||||
config = SlackConfig(bot_token="test", app_token="xapp-test")
|
||||
ch = SlackChannel(config)
|
||||
assert ch._strip_mention("<@U123> hello") == "<@U123> hello"
|
||||
|
||||
|
||||
# ── iMessage pipeline integration tests ──────────────────────────
|
||||
|
||||
|
||||
class TestIMessageConfig:
|
||||
def test_default_values(self):
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageConfig
|
||||
|
||||
config = IMessageConfig()
|
||||
assert config.allowed_senders is None
|
||||
assert config.include_attachments is True
|
||||
assert config.text_chunk_limit == 4096
|
||||
|
||||
def test_allowed_senders_is_set(self):
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageConfig
|
||||
|
||||
config = IMessageConfig(allowed_senders={"+1234", "chat_id:5"})
|
||||
assert isinstance(config.allowed_senders, set)
|
||||
assert "+1234" in config.allowed_senders
|
||||
|
||||
|
||||
class TestIMessageChannel:
|
||||
def test_is_ready_without_client(self):
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageChannelRpc
|
||||
|
||||
ch = IMessageChannelRpc()
|
||||
assert ch._is_ready() is False
|
||||
|
||||
def test_is_allowed_always_true(self):
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageChannelRpc
|
||||
|
||||
ch = IMessageChannelRpc()
|
||||
assert ch.is_allowed("anyone") is True
|
||||
|
||||
def test_build_inbound_uses_rich_filtering(self):
|
||||
from EvoScientist.channels.imessage.channel_rpc import (
|
||||
IMessageChannelRpc, IMessageConfig,
|
||||
)
|
||||
|
||||
config = IMessageConfig(allowed_senders={"+15551234567"})
|
||||
ch = IMessageChannelRpc(config)
|
||||
|
||||
# Allowed sender passes
|
||||
raw = RawIncoming(
|
||||
sender_id="+15551234567", chat_id="c1", text="hi",
|
||||
metadata={"chat_id": 1, "chat_guid": None},
|
||||
)
|
||||
assert ch._build_inbound(raw) is not None
|
||||
|
||||
# Disallowed sender blocked
|
||||
raw2 = RawIncoming(
|
||||
sender_id="+19999999999", chat_id="c1", text="hi",
|
||||
metadata={"chat_id": 1, "chat_guid": None},
|
||||
)
|
||||
assert ch._build_inbound(raw2) is None
|
||||
|
||||
def test_add_remove_allowed_senders(self):
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageChannelRpc
|
||||
|
||||
ch = IMessageChannelRpc()
|
||||
assert ch.config.allowed_senders is None
|
||||
|
||||
ch.add_allowed_sender("+15551234567")
|
||||
assert ch.config.allowed_senders is not None
|
||||
assert len(ch.config.allowed_senders) == 1
|
||||
|
||||
ch.remove_allowed_sender("+15551234567")
|
||||
assert len(ch.config.allowed_senders) == 0
|
||||
|
||||
ch.clear_allowed_senders()
|
||||
assert ch.config.allowed_senders is None
|
||||
|
||||
def test_list_allowed_senders(self):
|
||||
from EvoScientist.channels.imessage.channel_rpc import (
|
||||
IMessageChannelRpc, IMessageConfig,
|
||||
)
|
||||
|
||||
ch = IMessageChannelRpc()
|
||||
assert ch.list_allowed_senders() == []
|
||||
|
||||
config = IMessageConfig(allowed_senders={"a", "b"})
|
||||
ch2 = IMessageChannelRpc(config)
|
||||
result = ch2.list_allowed_senders()
|
||||
assert set(result) == {"a", "b"}
|
||||
|
||||
def test_handle_message_sets_is_group(self):
|
||||
import asyncio
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageChannelRpc
|
||||
|
||||
ch = IMessageChannelRpc()
|
||||
# Capture what _build_inbound receives
|
||||
captured = []
|
||||
original = ch._build_inbound
|
||||
def spy(raw):
|
||||
captured.append(raw)
|
||||
return original(raw)
|
||||
ch._build_inbound = spy
|
||||
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
loop.run_until_complete(ch._handle_message({
|
||||
"message": {
|
||||
"sender": "+1234",
|
||||
"text": "hello",
|
||||
"is_group": True,
|
||||
"chat_id": 42,
|
||||
"id": "msg1",
|
||||
}
|
||||
}))
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
assert len(captured) == 1
|
||||
assert captured[0].is_group is True
|
||||
assert captured[0].was_mentioned is True
|
||||
|
||||
|
||||
# ── Feishu config tests ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFeishuConfig:
|
||||
def test_allowed_channels(self):
|
||||
from EvoScientist.channels.feishu.channel import FeishuConfig
|
||||
|
||||
config = FeishuConfig(
|
||||
app_id="test", app_secret="test",
|
||||
allowed_channels={"oc_abc123"},
|
||||
)
|
||||
assert config.allowed_channels == {"oc_abc123"}
|
||||
|
||||
|
||||
# ── Slack retry tests ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSlackRetry:
|
||||
def test_extract_retry_after_from_response(self):
|
||||
from EvoScientist.channels.slack.channel import SlackChannel, SlackConfig
|
||||
|
||||
config = SlackConfig(bot_token="test", app_token="xapp-test")
|
||||
ch = SlackChannel(config)
|
||||
|
||||
# Simulate a SlackApiError-like exception with response headers
|
||||
class FakeResponse:
|
||||
headers = {"Retry-After": "30"}
|
||||
class FakeError(Exception):
|
||||
response = FakeResponse()
|
||||
|
||||
result = ch._extract_retry_after(FakeError())
|
||||
assert result == 30.0
|
||||
|
||||
def test_extract_retry_after_fallback(self):
|
||||
from EvoScientist.channels.slack.channel import SlackChannel, SlackConfig
|
||||
|
||||
config = SlackConfig(bot_token="test", app_token="xapp-test")
|
||||
ch = SlackChannel(config)
|
||||
|
||||
# Plain exception without response attribute
|
||||
result = ch._extract_retry_after(Exception("some error"))
|
||||
assert result == 1.0 # base class default
|
||||
|
||||
|
||||
# ── Attachment size check tests ──────────────────────────────────
|
||||
|
||||
|
||||
class TestAttachmentSizeCheck:
|
||||
"""Tests for _check_attachment_size and _download_attachment size guard."""
|
||||
|
||||
def test_check_attachment_size_within_limit(self):
|
||||
config = _StubConfig()
|
||||
ch = _StubChannel(config)
|
||||
assert ch._check_attachment_size(100, "small.txt") is None
|
||||
|
||||
def test_check_attachment_size_exceeds_limit(self):
|
||||
from EvoScientist.channels.base import MAX_ATTACHMENT_BYTES
|
||||
|
||||
config = _StubConfig()
|
||||
ch = _StubChannel(config)
|
||||
result = ch._check_attachment_size(MAX_ATTACHMENT_BYTES + 1, "big.zip")
|
||||
assert result is not None
|
||||
assert "too large" in result
|
||||
|
||||
|
||||
# ── Feishu _strip_mention tests ──────────────────────────────────
|
||||
|
||||
|
||||
class TestFeishuStripMention:
|
||||
def test_strip_bot_mention_placeholder(self):
|
||||
from EvoScientist.channels.feishu.channel import FeishuChannel, FeishuConfig
|
||||
|
||||
config = FeishuConfig(app_id="test", app_secret="test")
|
||||
ch = FeishuChannel(config)
|
||||
ch._mention_names = ["@_user_1"]
|
||||
assert ch._strip_mention("@_user_1 hello world") == "hello world"
|
||||
|
||||
def test_strip_multiple_mentions(self):
|
||||
from EvoScientist.channels.feishu.channel import FeishuChannel, FeishuConfig
|
||||
|
||||
config = FeishuConfig(app_id="test", app_secret="test")
|
||||
ch = FeishuChannel(config)
|
||||
ch._mention_names = ["@_user_1", "@_user_2"]
|
||||
result = ch._strip_mention("@_user_1 @_user_2 hello")
|
||||
assert result == "hello"
|
||||
|
||||
def test_no_mention_names(self):
|
||||
from EvoScientist.channels.feishu.channel import FeishuChannel, FeishuConfig
|
||||
|
||||
config = FeishuConfig(app_id="test", app_secret="test")
|
||||
ch = FeishuChannel(config)
|
||||
assert ch._strip_mention("hello world") == "hello world"
|
||||
|
||||
|
||||
# ── Telegram allowed_channels tests ──────────────────────────────
|
||||
|
||||
|
||||
class TestTelegramConfig:
|
||||
def test_allowed_channels_field_exists(self):
|
||||
from EvoScientist.channels.telegram.channel import TelegramConfig
|
||||
|
||||
config = TelegramConfig(bot_token="test", allowed_channels={"-100123"})
|
||||
assert config.allowed_channels == {"-100123"}
|
||||
|
||||
def test_allowed_channels_default_none(self):
|
||||
from EvoScientist.channels.telegram.channel import TelegramConfig
|
||||
|
||||
config = TelegramConfig(bot_token="test")
|
||||
assert config.allowed_channels is None
|
||||
|
||||
def test_channel_allow_list_integration(self):
|
||||
from EvoScientist.channels.telegram.channel import TelegramChannel, TelegramConfig
|
||||
|
||||
config = TelegramConfig(bot_token="test", allowed_channels={"-100123"})
|
||||
ch = TelegramChannel(config)
|
||||
assert ch.is_channel_allowed("-100123") is True
|
||||
assert ch.is_channel_allowed("-100999") is False
|
||||
|
||||
def test_channel_allow_list_empty_allows_all(self):
|
||||
from EvoScientist.channels.telegram.channel import TelegramChannel, TelegramConfig
|
||||
|
||||
config = TelegramConfig(bot_token="test")
|
||||
ch = TelegramChannel(config)
|
||||
assert ch.is_channel_allowed("-100999") is True
|
||||
|
||||
|
||||
# ── Feishu retry tests ───────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFeishuRetry:
|
||||
def test_extract_retry_after_rate_limit(self):
|
||||
from EvoScientist.channels.feishu.channel import FeishuChannel, FeishuConfig
|
||||
|
||||
config = FeishuConfig(app_id="test", app_secret="test")
|
||||
ch = FeishuChannel(config)
|
||||
result = ch._extract_retry_after(Exception("code 99991400: rate limit"))
|
||||
assert result == 2.0
|
||||
|
||||
def test_extract_retry_after_generic(self):
|
||||
from EvoScientist.channels.feishu.channel import FeishuChannel, FeishuConfig
|
||||
|
||||
config = FeishuConfig(app_id="test", app_secret="test")
|
||||
ch = FeishuChannel(config)
|
||||
result = ch._extract_retry_after(Exception("some error"))
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
# ── iMessage reply_to tests ──────────────────────────────────────
|
||||
|
||||
|
||||
class TestIMessageReplyTo:
|
||||
def test_send_chunk_includes_reply_to(self):
|
||||
"""Verify reply_to is passed through to RPC params."""
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageChannelRpc
|
||||
|
||||
ch = IMessageChannelRpc()
|
||||
# Mock the RPC client
|
||||
captured_params = {}
|
||||
|
||||
class FakeClient:
|
||||
async def request(self, method, params):
|
||||
captured_params.update(params)
|
||||
return {}
|
||||
|
||||
ch._client = FakeClient()
|
||||
|
||||
_run(ch._send_chunk("chat123", "hello", "hello", "msg42", {"chat_id": "chat123"}))
|
||||
assert captured_params.get("reply_to") == "msg42"
|
||||
|
||||
def test_send_chunk_omits_reply_to_when_none(self):
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageChannelRpc
|
||||
|
||||
ch = IMessageChannelRpc()
|
||||
captured_params = {}
|
||||
|
||||
class FakeClient:
|
||||
async def request(self, method, params):
|
||||
captured_params.update(params)
|
||||
return {}
|
||||
|
||||
ch._client = FakeClient()
|
||||
|
||||
_run(ch._send_chunk("chat123", "hello", "hello", None, {"chat_id": "chat123"}))
|
||||
assert "reply_to" not in captured_params
|
||||
@@ -0,0 +1,111 @@
|
||||
"""Tests for the Message Bus decoupling layer."""
|
||||
|
||||
import asyncio
|
||||
|
||||
|
||||
from EvoScientist.channels.bus.events import InboundMessage, OutboundMessage
|
||||
from EvoScientist.channels.bus.message_bus import MessageBus
|
||||
|
||||
|
||||
def _run(coro):
|
||||
"""Run an async coroutine safely, creating a fresh event loop."""
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
return loop.run_until_complete(coro)
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
|
||||
# ── Event tests ──
|
||||
|
||||
|
||||
class TestInboundMessage:
|
||||
def test_session_key(self):
|
||||
msg = InboundMessage(
|
||||
channel="telegram", sender_id="u1",
|
||||
chat_id="c1", content="hi",
|
||||
)
|
||||
assert msg.session_key == "telegram:c1"
|
||||
|
||||
def test_defaults(self):
|
||||
msg = InboundMessage(
|
||||
channel="discord", sender_id="u2",
|
||||
chat_id="c2", content="hello",
|
||||
)
|
||||
assert msg.media == []
|
||||
assert msg.metadata == {}
|
||||
assert msg.message_id == ""
|
||||
|
||||
|
||||
class TestOutboundMessage:
|
||||
def test_fields(self):
|
||||
msg = OutboundMessage(
|
||||
channel="telegram", chat_id="c1", content="reply",
|
||||
)
|
||||
assert msg.channel == "telegram"
|
||||
assert msg.chat_id == "c1"
|
||||
assert msg.reply_to is None
|
||||
assert msg.media == []
|
||||
|
||||
|
||||
# ── MessageBus tests ──
|
||||
|
||||
|
||||
class TestMessageBus:
|
||||
def test_inbound_publish_consume(self):
|
||||
async def _test():
|
||||
bus = MessageBus()
|
||||
msg = InboundMessage(
|
||||
channel="telegram", sender_id="u1",
|
||||
chat_id="c1", content="hello",
|
||||
)
|
||||
await bus.publish_inbound(msg)
|
||||
assert bus.inbound_size == 1
|
||||
got = await bus.consume_inbound()
|
||||
assert got is msg
|
||||
assert bus.inbound_size == 0
|
||||
_run(_test())
|
||||
|
||||
def test_outbound_publish_consume(self):
|
||||
async def _test():
|
||||
bus = MessageBus()
|
||||
msg = OutboundMessage(
|
||||
channel="discord", chat_id="c1", content="reply",
|
||||
)
|
||||
await bus.publish_outbound(msg)
|
||||
assert bus.outbound_size == 1
|
||||
got = await bus.consume_outbound()
|
||||
assert got is msg
|
||||
assert bus.outbound_size == 0
|
||||
_run(_test())
|
||||
|
||||
def test_subscribe_and_dispatch(self):
|
||||
async def _test():
|
||||
bus = MessageBus()
|
||||
received = []
|
||||
|
||||
async def callback(msg):
|
||||
received.append(msg)
|
||||
|
||||
bus.subscribe_outbound("telegram", callback)
|
||||
|
||||
msg = OutboundMessage(
|
||||
channel="telegram", chat_id="c1", content="hi",
|
||||
)
|
||||
await bus.publish_outbound(msg)
|
||||
|
||||
dispatch = asyncio.create_task(bus.dispatch_outbound())
|
||||
await asyncio.sleep(0.05)
|
||||
bus.stop()
|
||||
await asyncio.sleep(0.05)
|
||||
dispatch.cancel()
|
||||
|
||||
assert len(received) == 1
|
||||
assert received[0] is msg
|
||||
_run(_test())
|
||||
|
||||
def test_stop(self):
|
||||
bus = MessageBus()
|
||||
assert bus._running is False
|
||||
bus.stop()
|
||||
assert bus._running is False
|
||||
+1
-158
@@ -554,162 +554,5 @@ class TestParseTodoItemsAdvanced:
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# ChannelState queue mechanism
|
||||
# ChannelState queue mechanism (removed — replaced by bus mode in channel.py)
|
||||
# =============================================================================
|
||||
|
||||
class TestChannelState:
|
||||
"""Tests for _ChannelState queue-based communication."""
|
||||
|
||||
def test_enqueue_creates_message_in_queue(self):
|
||||
"""enqueue() should add a ChannelMessage to the queue."""
|
||||
from EvoScientist.cli import _ChannelState, ChannelMessage
|
||||
import queue
|
||||
|
||||
# Clear any existing messages
|
||||
while True:
|
||||
try:
|
||||
_ChannelState.message_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
|
||||
msg_id, event = _ChannelState.enqueue("test content", "sender@test.com", "Email")
|
||||
assert msg_id is not None
|
||||
assert event is not None
|
||||
|
||||
# Message should be in queue
|
||||
msg = _ChannelState.message_queue.get_nowait()
|
||||
assert isinstance(msg, ChannelMessage)
|
||||
assert msg.content == "test content"
|
||||
assert msg.sender == "sender@test.com"
|
||||
assert msg.channel_type == "Email"
|
||||
assert msg.msg_id == msg_id
|
||||
|
||||
def test_enqueue_creates_pending_response_slot(self):
|
||||
"""enqueue() should create a response slot for the message."""
|
||||
from EvoScientist.cli import _ChannelState
|
||||
import queue
|
||||
|
||||
# Clear queue and responses
|
||||
while True:
|
||||
try:
|
||||
_ChannelState.message_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
_ChannelState.pending_responses.clear()
|
||||
|
||||
msg_id, _ = _ChannelState.enqueue("test", "sender", "iMessage")
|
||||
assert msg_id in _ChannelState.pending_responses
|
||||
assert _ChannelState.pending_responses[msg_id]["response"] is None
|
||||
|
||||
# Cleanup
|
||||
_ChannelState.message_queue.get_nowait()
|
||||
|
||||
def test_set_response_updates_slot_and_signals(self):
|
||||
"""set_response() should update response and signal the event."""
|
||||
from EvoScientist.cli import _ChannelState
|
||||
import queue
|
||||
|
||||
# Clear
|
||||
while True:
|
||||
try:
|
||||
_ChannelState.message_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
_ChannelState.pending_responses.clear()
|
||||
|
||||
msg_id, event = _ChannelState.enqueue("test", "sender", "iMessage")
|
||||
_ChannelState.message_queue.get_nowait() # Remove from queue
|
||||
|
||||
assert not event.is_set()
|
||||
_ChannelState.set_response(msg_id, "response text")
|
||||
|
||||
assert event.is_set()
|
||||
assert _ChannelState.pending_responses[msg_id]["response"] == "response text"
|
||||
|
||||
# Cleanup
|
||||
_ChannelState.pending_responses.clear()
|
||||
|
||||
def test_get_response_waits_and_retrieves(self):
|
||||
"""get_response() should wait for response and return it."""
|
||||
from EvoScientist.cli import _ChannelState
|
||||
import threading
|
||||
import queue
|
||||
|
||||
# Clear
|
||||
while True:
|
||||
try:
|
||||
_ChannelState.message_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
_ChannelState.pending_responses.clear()
|
||||
|
||||
msg_id, _ = _ChannelState.enqueue("test", "sender", "iMessage")
|
||||
_ChannelState.message_queue.get_nowait()
|
||||
|
||||
# Set response in another thread
|
||||
def set_later():
|
||||
import time
|
||||
time.sleep(0.05)
|
||||
_ChannelState.set_response(msg_id, "async response")
|
||||
|
||||
t = threading.Thread(target=set_later)
|
||||
t.start()
|
||||
|
||||
response = _ChannelState.get_response(msg_id, timeout=1.0)
|
||||
t.join()
|
||||
|
||||
assert response == "async response"
|
||||
assert msg_id not in _ChannelState.pending_responses # Cleaned up
|
||||
|
||||
def test_get_response_returns_none_on_timeout(self):
|
||||
"""get_response() should return None if timeout expires."""
|
||||
from EvoScientist.cli import _ChannelState
|
||||
import queue
|
||||
|
||||
# Clear
|
||||
while True:
|
||||
try:
|
||||
_ChannelState.message_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
_ChannelState.pending_responses.clear()
|
||||
|
||||
msg_id, _ = _ChannelState.enqueue("test", "sender", "iMessage")
|
||||
_ChannelState.message_queue.get_nowait()
|
||||
|
||||
# Don't set response, let it timeout
|
||||
response = _ChannelState.get_response(msg_id, timeout=0.01)
|
||||
assert response is None
|
||||
|
||||
# Cleanup
|
||||
_ChannelState.pending_responses.clear()
|
||||
|
||||
def test_get_response_returns_none_for_unknown_id(self):
|
||||
"""get_response() should return None for unknown message ID."""
|
||||
from EvoScientist.cli import _ChannelState
|
||||
response = _ChannelState.get_response("nonexistent-id", timeout=0.01)
|
||||
assert response is None
|
||||
|
||||
def test_channel_message_dataclass(self):
|
||||
"""ChannelMessage should store all fields correctly."""
|
||||
from EvoScientist.cli import ChannelMessage
|
||||
|
||||
msg = ChannelMessage(
|
||||
msg_id="id123",
|
||||
content="Hello",
|
||||
sender="+1234567890",
|
||||
channel_type="iMessage",
|
||||
metadata={"key": "value"},
|
||||
)
|
||||
assert msg.msg_id == "id123"
|
||||
assert msg.content == "Hello"
|
||||
assert msg.sender == "+1234567890"
|
||||
assert msg.channel_type == "iMessage"
|
||||
assert msg.metadata == {"key": "value"}
|
||||
|
||||
def test_channel_message_default_metadata(self):
|
||||
"""ChannelMessage metadata should default to None."""
|
||||
from EvoScientist.cli import ChannelMessage
|
||||
|
||||
msg = ChannelMessage("id", "content", "sender", "type")
|
||||
assert msg.metadata is None
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Tests for Telegram channel implementation."""
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.channels.telegram.channel import TelegramChannel, TelegramConfig
|
||||
from EvoScientist.channels.base import ChannelError
|
||||
|
||||
|
||||
def _run(coro):
|
||||
"""Run an async coroutine safely, creating a fresh event loop."""
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
return loop.run_until_complete(coro)
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
|
||||
class TestTelegramConfig:
|
||||
def test_default_values(self):
|
||||
config = TelegramConfig()
|
||||
assert config.bot_token == ""
|
||||
assert config.allowed_senders is None
|
||||
assert config.text_chunk_limit == 4096
|
||||
|
||||
def test_custom_values(self):
|
||||
config = TelegramConfig(
|
||||
bot_token="test-token",
|
||||
allowed_senders={"123", "456"},
|
||||
text_chunk_limit=2000,
|
||||
)
|
||||
assert config.bot_token == "test-token"
|
||||
assert config.allowed_senders == {"123", "456"}
|
||||
assert config.text_chunk_limit == 2000
|
||||
|
||||
|
||||
class TestTelegramChannel:
|
||||
def test_init(self):
|
||||
config = TelegramConfig(bot_token="test")
|
||||
channel = TelegramChannel(config)
|
||||
assert channel.config is config
|
||||
assert channel._running is False
|
||||
|
||||
def test_start_raises_without_token(self):
|
||||
config = TelegramConfig(bot_token="")
|
||||
channel = TelegramChannel(config)
|
||||
with pytest.raises(ChannelError, match="bot token"):
|
||||
_run(channel.start())
|
||||
|
||||
def test_stop_when_not_running(self):
|
||||
config = TelegramConfig(bot_token="test")
|
||||
channel = TelegramChannel(config)
|
||||
_run(channel.stop())
|
||||
|
||||
def test_send_returns_false_without_app(self):
|
||||
from EvoScientist.channels.base import OutboundMessage
|
||||
|
||||
config = TelegramConfig(bot_token="test")
|
||||
channel = TelegramChannel(config)
|
||||
msg = OutboundMessage(
|
||||
channel="telegram",
|
||||
chat_id="123",
|
||||
content="hello",
|
||||
metadata={"chat_id": "123"},
|
||||
)
|
||||
result = _run(channel.send(msg))
|
||||
assert result is False
|
||||
@@ -0,0 +1,458 @@
|
||||
"""Tests for WeChat channel implementation."""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import time
|
||||
import xml.etree.ElementTree as ET
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.channels.wechat.channel import (
|
||||
WeChatChannel,
|
||||
WeComConfig,
|
||||
WeChatMPConfig,
|
||||
_strip_markdown,
|
||||
)
|
||||
from EvoScientist.channels.wechat.crypto import (
|
||||
WeChatCrypto,
|
||||
parse_xml,
|
||||
_pkcs7_pad,
|
||||
_pkcs7_unpad,
|
||||
)
|
||||
from EvoScientist.channels.base import ChannelError
|
||||
|
||||
|
||||
def _run(coro):
|
||||
"""Run an async coroutine safely, creating a fresh event loop."""
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
return loop.run_until_complete(coro)
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
|
||||
# ── Config tests ──────────────────────────────────────────────────
|
||||
|
||||
class TestWeComConfig:
|
||||
def test_default_values(self):
|
||||
config = WeComConfig()
|
||||
assert config.corp_id == ""
|
||||
assert config.agent_id == ""
|
||||
assert config.secret == ""
|
||||
assert config.webhook_port == 9001
|
||||
assert config.allowed_senders is None
|
||||
assert config.text_chunk_limit == 4096
|
||||
|
||||
def test_custom_values(self):
|
||||
config = WeComConfig(
|
||||
corp_id="corp123",
|
||||
agent_id="1000001",
|
||||
secret="my-secret",
|
||||
token="my-token",
|
||||
encoding_aes_key="a" * 43,
|
||||
webhook_port=8080,
|
||||
allowed_senders={"user1", "user2"},
|
||||
)
|
||||
assert config.corp_id == "corp123"
|
||||
assert config.agent_id == "1000001"
|
||||
assert config.allowed_senders == {"user1", "user2"}
|
||||
assert config.webhook_port == 8080
|
||||
|
||||
|
||||
class TestWeChatMPConfig:
|
||||
def test_default_values(self):
|
||||
config = WeChatMPConfig()
|
||||
assert config.app_id == ""
|
||||
assert config.app_secret == ""
|
||||
assert config.webhook_port == 9001
|
||||
|
||||
def test_custom_values(self):
|
||||
config = WeChatMPConfig(
|
||||
app_id="wx1234",
|
||||
app_secret="secret",
|
||||
token="mp-token",
|
||||
)
|
||||
assert config.app_id == "wx1234"
|
||||
|
||||
|
||||
# ── Channel init / lifecycle tests ────────────────────────────────
|
||||
|
||||
class TestWeChatChannelInit:
|
||||
def test_wecom_init(self):
|
||||
config = WeComConfig(corp_id="corp", agent_id="1", secret="s")
|
||||
channel = WeChatChannel(config, backend="wecom")
|
||||
assert channel.name == "wechat"
|
||||
assert channel._backend == "wecom"
|
||||
assert channel._running is False
|
||||
|
||||
def test_mp_init(self):
|
||||
config = WeChatMPConfig(app_id="wx", app_secret="s")
|
||||
channel = WeChatChannel(config, backend="wechatmp")
|
||||
assert channel._backend == "wechatmp"
|
||||
|
||||
def test_start_raises_without_corp_id(self):
|
||||
config = WeComConfig(corp_id="", agent_id="1", secret="s")
|
||||
channel = WeChatChannel(config, backend="wecom")
|
||||
with pytest.raises(ChannelError, match="corp_id"):
|
||||
_run(channel.start())
|
||||
|
||||
def test_start_raises_without_secret(self):
|
||||
config = WeComConfig(corp_id="corp", agent_id="1", secret="")
|
||||
channel = WeChatChannel(config, backend="wecom")
|
||||
with pytest.raises(ChannelError, match="secret"):
|
||||
_run(channel.start())
|
||||
|
||||
def test_start_raises_without_agent_id(self):
|
||||
config = WeComConfig(corp_id="corp", agent_id="", secret="s")
|
||||
channel = WeChatChannel(config, backend="wecom")
|
||||
with pytest.raises(ChannelError, match="agent_id"):
|
||||
_run(channel.start())
|
||||
|
||||
def test_start_raises_mp_without_app_id(self):
|
||||
config = WeChatMPConfig(app_id="", app_secret="s")
|
||||
channel = WeChatChannel(config, backend="wechatmp")
|
||||
with pytest.raises(ChannelError, match="app_id"):
|
||||
_run(channel.start())
|
||||
|
||||
def test_stop_when_not_running(self):
|
||||
config = WeComConfig(corp_id="c", agent_id="1", secret="s")
|
||||
channel = WeChatChannel(config, backend="wecom")
|
||||
_run(channel.stop()) # Should not raise
|
||||
|
||||
def test_send_returns_false_without_client(self):
|
||||
from EvoScientist.channels.base import OutboundMessage
|
||||
|
||||
config = WeComConfig(corp_id="c", agent_id="1", secret="s")
|
||||
channel = WeChatChannel(config, backend="wecom")
|
||||
msg = OutboundMessage(
|
||||
channel="wechat",
|
||||
chat_id="user1",
|
||||
content="hello",
|
||||
metadata={"chat_id": "user1"},
|
||||
)
|
||||
result = _run(channel.send(msg))
|
||||
assert result is False
|
||||
|
||||
|
||||
# ── Markdown stripping tests ──────────────────────────────────────
|
||||
|
||||
class TestStripMarkdown:
|
||||
def test_plain_text(self):
|
||||
assert _strip_markdown("hello world") == "hello world"
|
||||
|
||||
def test_bold(self):
|
||||
assert _strip_markdown("**bold**") == "bold"
|
||||
|
||||
def test_italic(self):
|
||||
assert _strip_markdown("_italic_") == "italic"
|
||||
|
||||
def test_code(self):
|
||||
assert _strip_markdown("`code`") == "code"
|
||||
|
||||
def test_link(self):
|
||||
result = _strip_markdown("[text](https://example.com)")
|
||||
assert "text" in result
|
||||
assert "https://example.com" in result
|
||||
|
||||
def test_heading(self):
|
||||
assert _strip_markdown("## Title").strip() == "Title"
|
||||
|
||||
def test_list_items(self):
|
||||
result = _strip_markdown("- item1\n- item2")
|
||||
assert "• item1" in result
|
||||
assert "• item2" in result
|
||||
|
||||
def test_strikethrough(self):
|
||||
assert _strip_markdown("~~deleted~~") == "deleted"
|
||||
|
||||
def test_code_block(self):
|
||||
text = "```python\nprint('hi')\n```"
|
||||
result = _strip_markdown(text)
|
||||
assert "print('hi')" in result
|
||||
|
||||
|
||||
# ── XML parsing tests ─────────────────────────────────────────────
|
||||
|
||||
class TestParseXml:
|
||||
def test_basic_text_message(self):
|
||||
xml = (
|
||||
"<xml>"
|
||||
"<MsgType><![CDATA[text]]></MsgType>"
|
||||
"<Content><![CDATA[hello]]></Content>"
|
||||
"<FromUserName><![CDATA[user123]]></FromUserName>"
|
||||
"<ToUserName><![CDATA[bot]]></ToUserName>"
|
||||
"<MsgId>1234</MsgId>"
|
||||
"<CreateTime>1700000000</CreateTime>"
|
||||
"</xml>"
|
||||
)
|
||||
data = parse_xml(xml)
|
||||
assert data["MsgType"] == "text"
|
||||
assert data["Content"] == "hello"
|
||||
assert data["FromUserName"] == "user123"
|
||||
assert data["MsgId"] == "1234"
|
||||
|
||||
def test_image_message(self):
|
||||
xml = (
|
||||
"<xml>"
|
||||
"<MsgType><![CDATA[image]]></MsgType>"
|
||||
"<PicUrl><![CDATA[https://example.com/img.jpg]]></PicUrl>"
|
||||
"<MediaId><![CDATA[media_123]]></MediaId>"
|
||||
"<FromUserName><![CDATA[user1]]></FromUserName>"
|
||||
"</xml>"
|
||||
)
|
||||
data = parse_xml(xml)
|
||||
assert data["MsgType"] == "image"
|
||||
assert data["PicUrl"] == "https://example.com/img.jpg"
|
||||
|
||||
def test_event_message(self):
|
||||
xml = (
|
||||
"<xml>"
|
||||
"<MsgType><![CDATA[event]]></MsgType>"
|
||||
"<Event><![CDATA[subscribe]]></Event>"
|
||||
"<FromUserName><![CDATA[user1]]></FromUserName>"
|
||||
"</xml>"
|
||||
)
|
||||
data = parse_xml(xml)
|
||||
assert data["MsgType"] == "event"
|
||||
assert data["Event"] == "subscribe"
|
||||
|
||||
|
||||
# ── Crypto tests ──────────────────────────────────────────────────
|
||||
|
||||
class TestPKCS7:
|
||||
def test_pad_unpad_roundtrip(self):
|
||||
data = b"hello"
|
||||
padded = _pkcs7_pad(data)
|
||||
assert len(padded) % 32 == 0
|
||||
assert _pkcs7_unpad(padded) == data
|
||||
|
||||
def test_pad_block_aligned(self):
|
||||
data = b"x" * 32
|
||||
padded = _pkcs7_pad(data)
|
||||
assert len(padded) == 64 # full padding block added
|
||||
assert _pkcs7_unpad(padded) == data
|
||||
|
||||
|
||||
class TestWeChatCrypto:
|
||||
"""Test the encryption/decryption roundtrip.
|
||||
|
||||
Uses a deterministic 43-char EncodingAESKey.
|
||||
"""
|
||||
|
||||
# Skip encryption tests when no crypto backend is available
|
||||
_has_crypto = False
|
||||
try:
|
||||
from Crypto.Cipher import AES as _aes # noqa: F401
|
||||
_has_crypto = True
|
||||
except ImportError:
|
||||
try:
|
||||
import pyaes as _pyaes # noqa: F401
|
||||
_has_crypto = True
|
||||
except ImportError:
|
||||
pass
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not _has_crypto,
|
||||
reason="pycryptodome or pyaes required for encryption tests",
|
||||
)
|
||||
|
||||
@pytest.fixture
|
||||
def crypto(self):
|
||||
# 43 base64 chars → 32 bytes AES key
|
||||
key = "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG"
|
||||
return WeChatCrypto(
|
||||
token="test_token",
|
||||
encoding_aes_key=key,
|
||||
app_id="wx_test_app",
|
||||
)
|
||||
|
||||
def test_encrypt_decrypt_roundtrip(self, crypto):
|
||||
msg = "<xml><Content>Hello WeChat!</Content></xml>"
|
||||
encrypted = crypto.encrypt(msg)
|
||||
decrypted, app_id = crypto.decrypt(encrypted)
|
||||
assert decrypted == msg
|
||||
assert app_id == "wx_test_app"
|
||||
|
||||
def test_verify_signature(self, crypto):
|
||||
timestamp = "1609459200"
|
||||
nonce = "abc123"
|
||||
parts = sorted([crypto.token, timestamp, nonce])
|
||||
expected = hashlib.sha1("".join(parts).encode()).hexdigest()
|
||||
assert crypto.verify_signature(expected, timestamp, nonce)
|
||||
assert not crypto.verify_signature("wrong", timestamp, nonce)
|
||||
|
||||
def test_verify_signature_with_encrypt(self, crypto):
|
||||
timestamp = "1609459200"
|
||||
nonce = "abc123"
|
||||
encrypt = "some_encrypted_data"
|
||||
parts = sorted([crypto.token, timestamp, nonce, encrypt])
|
||||
expected = hashlib.sha1("".join(parts).encode()).hexdigest()
|
||||
assert crypto.verify_signature(expected, timestamp, nonce, encrypt)
|
||||
|
||||
def test_generate_signature(self, crypto):
|
||||
encrypt = "test_encrypted"
|
||||
timestamp = "1609459200"
|
||||
nonce = "abc"
|
||||
sig = crypto.generate_signature(encrypt, timestamp, nonce)
|
||||
parts = sorted([crypto.token, timestamp, nonce, encrypt])
|
||||
expected = hashlib.sha1("".join(parts).encode()).hexdigest()
|
||||
assert sig == expected
|
||||
|
||||
def test_wrap_encrypted_reply(self, crypto):
|
||||
msg = "<xml><Content>Reply</Content></xml>"
|
||||
xml_reply = crypto.wrap_encrypted_reply(msg)
|
||||
assert "<Encrypt>" in xml_reply
|
||||
assert "<MsgSignature>" in xml_reply
|
||||
assert "<TimeStamp>" in xml_reply
|
||||
assert "<Nonce>" in xml_reply
|
||||
|
||||
# Parse and verify the encrypted content decrypts back
|
||||
root = ET.fromstring(xml_reply)
|
||||
encrypt = root.find("Encrypt").text
|
||||
decrypted, app_id = crypto.decrypt(encrypt)
|
||||
assert decrypted == msg
|
||||
|
||||
|
||||
# ── Message processing tests ──────────────────────────────────────
|
||||
|
||||
class TestMessageProcessing:
|
||||
"""Test the _process_message method with various XML payloads."""
|
||||
|
||||
def _make_channel(self):
|
||||
config = WeComConfig(
|
||||
corp_id="corp", agent_id="1", secret="s",
|
||||
)
|
||||
return WeChatChannel(config, backend="wecom")
|
||||
|
||||
def test_text_message_queued(self):
|
||||
channel = self._make_channel()
|
||||
|
||||
async def _test():
|
||||
await channel._process_message({
|
||||
"MsgType": "text",
|
||||
"Content": "Hello!",
|
||||
"FromUserName": "user1",
|
||||
"ToUserName": "bot",
|
||||
"MsgId": "100",
|
||||
"CreateTime": str(int(time.time())),
|
||||
})
|
||||
# Check message was enqueued
|
||||
assert not channel._queue.empty()
|
||||
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
||||
assert msg.content == "Hello!"
|
||||
assert msg.sender_id == "user1"
|
||||
assert msg.channel == "wechat"
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_location_message(self):
|
||||
channel = self._make_channel()
|
||||
|
||||
async def _test():
|
||||
await channel._process_message({
|
||||
"MsgType": "location",
|
||||
"Location_X": "39.9",
|
||||
"Location_Y": "116.4",
|
||||
"Label": "Beijing",
|
||||
"FromUserName": "user1",
|
||||
"ToUserName": "bot",
|
||||
"MsgId": "101",
|
||||
"CreateTime": str(int(time.time())),
|
||||
})
|
||||
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
||||
assert "Beijing" in msg.content
|
||||
assert "39.9" in msg.content
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_voice_recognition(self):
|
||||
channel = self._make_channel()
|
||||
|
||||
async def _test():
|
||||
await channel._process_message({
|
||||
"MsgType": "voice",
|
||||
"Recognition": "你好世界",
|
||||
"FromUserName": "user1",
|
||||
"ToUserName": "bot",
|
||||
"MsgId": "102",
|
||||
"CreateTime": str(int(time.time())),
|
||||
})
|
||||
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
||||
assert "你好世界" in msg.content
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_link_message(self):
|
||||
channel = self._make_channel()
|
||||
|
||||
async def _test():
|
||||
await channel._process_message({
|
||||
"MsgType": "link",
|
||||
"Title": "Test Link",
|
||||
"Description": "A description",
|
||||
"Url": "https://example.com",
|
||||
"FromUserName": "user1",
|
||||
"ToUserName": "bot",
|
||||
"MsgId": "103",
|
||||
"CreateTime": str(int(time.time())),
|
||||
})
|
||||
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
||||
assert "Test Link" in msg.content
|
||||
assert "https://example.com" in msg.content
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_subscribe_event(self):
|
||||
channel = self._make_channel()
|
||||
|
||||
async def _test():
|
||||
await channel._process_message({
|
||||
"MsgType": "event",
|
||||
"Event": "subscribe",
|
||||
"FromUserName": "user1",
|
||||
"ToUserName": "bot",
|
||||
"MsgId": "",
|
||||
"CreateTime": str(int(time.time())),
|
||||
})
|
||||
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
||||
assert "关注" in msg.content
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_unsubscribe_ignored(self):
|
||||
channel = self._make_channel()
|
||||
|
||||
async def _test():
|
||||
await channel._process_message({
|
||||
"MsgType": "event",
|
||||
"Event": "unsubscribe",
|
||||
"FromUserName": "user1",
|
||||
"ToUserName": "bot",
|
||||
"MsgId": "",
|
||||
"CreateTime": str(int(time.time())),
|
||||
})
|
||||
assert channel._queue.empty()
|
||||
|
||||
_run(_test())
|
||||
|
||||
def test_empty_message_ignored(self):
|
||||
channel = self._make_channel()
|
||||
|
||||
async def _test():
|
||||
await channel._process_message({
|
||||
"MsgType": "text",
|
||||
"Content": "",
|
||||
"FromUserName": "",
|
||||
"ToUserName": "bot",
|
||||
})
|
||||
assert channel._queue.empty()
|
||||
|
||||
_run(_test())
|
||||
|
||||
|
||||
# ── Registration test ─────────────────────────────────────────────
|
||||
|
||||
class TestChannelRegistration:
|
||||
def test_wechat_registered(self):
|
||||
from EvoScientist.channels.channel_manager import available_channels
|
||||
channels = available_channels()
|
||||
assert "wechat" in channels
|
||||
Reference in New Issue
Block a user