mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-09 22:08:38 +03:00
Compare commits
27
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5257453c4c | ||
|
|
a4dfbdf996 | ||
|
|
949a10f536 | ||
|
|
2a6c616080 | ||
|
|
1bcd5f9742 | ||
|
|
26947db479 | ||
|
|
0514233217 | ||
|
|
345c393e53 | ||
|
|
faf2b07923 | ||
|
|
efd42cc236 | ||
|
|
3823042290 | ||
|
|
5bdb7a90b1 | ||
|
|
bc8fbd1ce4 | ||
|
|
6aad945719 | ||
|
|
f450c6ef6c | ||
|
|
8956df3668 | ||
|
|
0506e6c1c1 | ||
|
|
b94d4c0509 | ||
|
|
d0c68157b1 | ||
|
|
2dce5e07c1 | ||
|
|
1a4ad67628 | ||
|
|
ed2ca759e7 | ||
|
|
79a915307c | ||
|
|
2abd990b89 | ||
|
|
0207b541df | ||
|
|
b1d5475681 | ||
|
|
e04e1c24ff |
@@ -1387,6 +1387,8 @@ MCP tools are automatically discovered and registered on startup. The LLM can us
|
|||||||
| `tools.restrictToWorkspace` | `false` | When `true`, restricts **all** agent tools (shell, file read/write/edit, list) to the workspace directory. Prevents path traversal and out-of-scope access. |
|
| `tools.restrictToWorkspace` | `false` | When `true`, restricts **all** agent tools (shell, file read/write/edit, list) to the workspace directory. Prevents path traversal and out-of-scope access. |
|
||||||
| `tools.exec.enable` | `true` | When `false`, the shell `exec` tool is not registered at all. Use this to completely disable shell command execution. |
|
| `tools.exec.enable` | `true` | When `false`, the shell `exec` tool is not registered at all. Use this to completely disable shell command execution. |
|
||||||
| `tools.exec.pathAppend` | `""` | Extra directories to append to `PATH` when running shell commands (e.g. `/usr/sbin` for `ufw`). |
|
| `tools.exec.pathAppend` | `""` | Extra directories to append to `PATH` when running shell commands (e.g. `/usr/sbin` for `ufw`). |
|
||||||
|
| `tools.exec.commandWrapper` | `""` | Sandbox wrapper command template. See [Exec Tool Sandbox](docs/COMMAND_WRAPPER.md) for details and examples. |
|
||||||
|
|
||||||
| `channels.*.allowFrom` | `[]` (deny all) | Whitelist of user IDs. Empty denies all; use `["*"]` to allow everyone. |
|
| `channels.*.allowFrom` | `[]` (deny all) | Whitelist of user IDs. Empty denies all; use `["*"]` to allow everyone. |
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,86 @@
|
|||||||
|
# Exec Tool Sandbox (`commandWrapper`)
|
||||||
|
|
||||||
|
The `tools.exec.commandWrapper` config option wraps every shell command in a user-defined template before execution. This allows you to add a sandbox layer (e.g. bubblewrap, firejail, nsjail) without any code changes to nanobot.
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"commandWrapper": "<template>"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Leave empty (the default) to run commands directly with no wrapper.
|
||||||
|
|
||||||
|
## Placeholders
|
||||||
|
|
||||||
|
Two placeholders are available in the template:
|
||||||
|
|
||||||
|
| Placeholder | Value |
|
||||||
|
|---|---|
|
||||||
|
| `{command}` | The original shell command generated by the LLM |
|
||||||
|
| `{cwd}` | Absolute path of the working directory |
|
||||||
|
|
||||||
|
nanobot performs plain string replacement — it does not parse, validate, or shell-escape the values. The wrapper template is trusted configuration.
|
||||||
|
|
||||||
|
## Examples
|
||||||
|
|
||||||
|
### bubblewrap
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"commandWrapper": "bwrap --ro-bind /usr /usr --ro-bind-try /bin /bin --ro-bind-try /lib /lib --ro-bind-try /lib64 /lib64 --proc /proc --dev /dev --tmpfs /tmp --bind {cwd} {cwd} --chdir {cwd} -- sh -c \"{command}\""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Requires: `apt install bubblewrap` (or equivalent for your distro).
|
||||||
|
|
||||||
|
### firejail
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"commandWrapper": "firejail --noprofile --private={cwd} -- {command}"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### nsjail
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"exec": {
|
||||||
|
"commandWrapper": "nsjail -Mo --chroot /sandbox --cwd {cwd} -- {command}"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Caveats
|
||||||
|
|
||||||
|
> [!WARNING]
|
||||||
|
> **Do not wrap `{command}` in shell quotes.** If the original command contains the same quote character, the shell will break the quoting context. For example, `sh -c '{command}'` will fail on any command that contains single quotes.
|
||||||
|
|
||||||
|
This is an inherent limitation of the template approach — nanobot substitutes `{command}` as a raw string and cannot safely shell-quote it (the command may contain compound syntax like `&&`, `|`, `;` that must be preserved for the inner shell).
|
||||||
|
|
||||||
|
### Interaction with `create_subprocess_shell`
|
||||||
|
|
||||||
|
nanobot executes the wrapped command via `create_subprocess_shell`, which adds an outer shell layer. Keep this in mind when designing your template:
|
||||||
|
|
||||||
|
- **Without `sh -c`** (e.g. `firejail ... -- {command}`): The outer shell parses `{command}` directly. Compound commands with `&&` and `|` work as expected because they are parsed by the outer shell before the sandbox tool receives them.
|
||||||
|
- **With `sh -c`** (e.g. `bwrap ... -- sh -c "{command}"`): The command is passed through two shell layers. This is only needed if the sandbox tool requires a single command argument but you want to support compound syntax.
|
||||||
|
|
||||||
|
### `restrict_to_workspace` is independent
|
||||||
|
|
||||||
|
The `tools.restrictToWorkspace` setting and `commandWrapper` are orthogonal features. The workspace restriction guards against path traversal in the original command (before wrapping). The sandbox wrapper provides OS-level isolation. You can use either or both — they address different threat models.
|
||||||
+10
-9
@@ -248,6 +248,7 @@ class AgentLoop:
|
|||||||
timeout=self.exec_config.timeout,
|
timeout=self.exec_config.timeout,
|
||||||
restrict_to_workspace=self.restrict_to_workspace,
|
restrict_to_workspace=self.restrict_to_workspace,
|
||||||
path_append=self.exec_config.path_append,
|
path_append=self.exec_config.path_append,
|
||||||
|
command_wrapper=self.exec_config.command_wrapper,
|
||||||
))
|
))
|
||||||
self.tools.register(WebSearchTool(config=self.web_search_config, proxy=self.web_proxy))
|
self.tools.register(WebSearchTool(config=self.web_search_config, proxy=self.web_proxy))
|
||||||
self.tools.register(WebFetchTool(proxy=self.web_proxy))
|
self.tools.register(WebFetchTool(proxy=self.web_proxy))
|
||||||
@@ -403,25 +404,25 @@ class AgentLoop:
|
|||||||
return f"{stream_base_id}:{stream_segment}"
|
return f"{stream_base_id}:{stream_segment}"
|
||||||
|
|
||||||
async def on_stream(delta: str) -> None:
|
async def on_stream(delta: str) -> None:
|
||||||
|
meta = dict(msg.metadata or {})
|
||||||
|
meta["_stream_delta"] = True
|
||||||
|
meta["_stream_id"] = _current_stream_id()
|
||||||
await self.bus.publish_outbound(OutboundMessage(
|
await self.bus.publish_outbound(OutboundMessage(
|
||||||
channel=msg.channel, chat_id=msg.chat_id,
|
channel=msg.channel, chat_id=msg.chat_id,
|
||||||
content=delta,
|
content=delta,
|
||||||
metadata={
|
metadata=meta,
|
||||||
"_stream_delta": True,
|
|
||||||
"_stream_id": _current_stream_id(),
|
|
||||||
},
|
|
||||||
))
|
))
|
||||||
|
|
||||||
async def on_stream_end(*, resuming: bool = False) -> None:
|
async def on_stream_end(*, resuming: bool = False) -> None:
|
||||||
nonlocal stream_segment
|
nonlocal stream_segment
|
||||||
|
meta = dict(msg.metadata or {})
|
||||||
|
meta["_stream_end"] = True
|
||||||
|
meta["_resuming"] = resuming
|
||||||
|
meta["_stream_id"] = _current_stream_id()
|
||||||
await self.bus.publish_outbound(OutboundMessage(
|
await self.bus.publish_outbound(OutboundMessage(
|
||||||
channel=msg.channel, chat_id=msg.chat_id,
|
channel=msg.channel, chat_id=msg.chat_id,
|
||||||
content="",
|
content="",
|
||||||
metadata={
|
metadata=meta,
|
||||||
"_stream_end": True,
|
|
||||||
"_resuming": resuming,
|
|
||||||
"_stream_id": _current_stream_id(),
|
|
||||||
},
|
|
||||||
))
|
))
|
||||||
stream_segment += 1
|
stream_segment += 1
|
||||||
|
|
||||||
|
|||||||
@@ -121,6 +121,7 @@ class SubagentManager:
|
|||||||
timeout=self.exec_config.timeout,
|
timeout=self.exec_config.timeout,
|
||||||
restrict_to_workspace=self.restrict_to_workspace,
|
restrict_to_workspace=self.restrict_to_workspace,
|
||||||
path_append=self.exec_config.path_append,
|
path_append=self.exec_config.path_append,
|
||||||
|
command_wrapper=self.exec_config.command_wrapper,
|
||||||
))
|
))
|
||||||
tools.register(WebSearchTool(config=self.web_search_config, proxy=self.web_proxy))
|
tools.register(WebSearchTool(config=self.web_search_config, proxy=self.web_proxy))
|
||||||
tools.register(WebFetchTool(proxy=self.web_proxy))
|
tools.register(WebFetchTool(proxy=self.web_proxy))
|
||||||
|
|||||||
@@ -23,9 +23,11 @@ class ExecTool(Tool):
|
|||||||
allow_patterns: list[str] | None = None,
|
allow_patterns: list[str] | None = None,
|
||||||
restrict_to_workspace: bool = False,
|
restrict_to_workspace: bool = False,
|
||||||
path_append: str = "",
|
path_append: str = "",
|
||||||
|
command_wrapper: str = "",
|
||||||
):
|
):
|
||||||
self.timeout = timeout
|
self.timeout = timeout
|
||||||
self.working_dir = working_dir
|
self.working_dir = working_dir
|
||||||
|
self.command_wrapper = command_wrapper
|
||||||
self.deny_patterns = deny_patterns or [
|
self.deny_patterns = deny_patterns or [
|
||||||
r"\brm\s+-[rf]{1,2}\b", # rm -r, rm -rf, rm -fr
|
r"\brm\s+-[rf]{1,2}\b", # rm -r, rm -rf, rm -fr
|
||||||
r"\bdel\s+/[fq]\b", # del /f, del /q
|
r"\bdel\s+/[fq]\b", # del /f, del /q
|
||||||
@@ -82,11 +84,16 @@ class ExecTool(Tool):
|
|||||||
self, command: str, working_dir: str | None = None,
|
self, command: str, working_dir: str | None = None,
|
||||||
timeout: int | None = None, **kwargs: Any,
|
timeout: int | None = None, **kwargs: Any,
|
||||||
) -> str:
|
) -> str:
|
||||||
cwd = working_dir or self.working_dir or os.getcwd()
|
cwd = os.path.abspath(working_dir or self.working_dir or os.getcwd())
|
||||||
guard_error = self._guard_command(command, cwd)
|
guard_error = self._guard_command(command, cwd)
|
||||||
if guard_error:
|
if guard_error:
|
||||||
return guard_error
|
return guard_error
|
||||||
|
|
||||||
|
if self.command_wrapper:
|
||||||
|
original_command = command
|
||||||
|
command = self.command_wrapper.replace("{cwd}", cwd).replace("{command}", command)
|
||||||
|
logger.debug("command_wrapper applied: {} -> {}", original_command, command)
|
||||||
|
|
||||||
effective_timeout = min(timeout or self.timeout, self._MAX_TIMEOUT)
|
effective_timeout = min(timeout or self.timeout, self._MAX_TIMEOUT)
|
||||||
|
|
||||||
env = os.environ.copy()
|
env = os.environ.copy()
|
||||||
|
|||||||
+404
-283
@@ -1,25 +1,37 @@
|
|||||||
"""Discord channel implementation using Discord Gateway websocket."""
|
"""Discord channel implementation using discord.py."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import importlib.util
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Literal
|
from typing import TYPE_CHECKING, Any, Literal
|
||||||
|
|
||||||
import httpx
|
|
||||||
from pydantic import Field
|
|
||||||
import websockets
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
from pydantic import Field
|
||||||
|
|
||||||
from nanobot.bus.events import OutboundMessage
|
from nanobot.bus.events import OutboundMessage
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.channels.base import BaseChannel
|
from nanobot.channels.base import BaseChannel
|
||||||
|
from nanobot.command.builtin import build_help_text
|
||||||
from nanobot.config.paths import get_media_dir
|
from nanobot.config.paths import get_media_dir
|
||||||
from nanobot.config.schema import Base
|
from nanobot.config.schema import Base
|
||||||
from nanobot.utils.helpers import split_message
|
from nanobot.utils.helpers import safe_filename, split_message
|
||||||
|
|
||||||
|
DISCORD_AVAILABLE = importlib.util.find_spec("discord") is not None
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import discord
|
||||||
|
from discord import app_commands
|
||||||
|
from discord.abc import Messageable
|
||||||
|
|
||||||
|
if DISCORD_AVAILABLE:
|
||||||
|
import discord
|
||||||
|
from discord import app_commands
|
||||||
|
from discord.abc import Messageable
|
||||||
|
|
||||||
DISCORD_API_BASE = "https://discord.com/api/v10"
|
|
||||||
MAX_ATTACHMENT_BYTES = 20 * 1024 * 1024 # 20MB
|
MAX_ATTACHMENT_BYTES = 20 * 1024 * 1024 # 20MB
|
||||||
MAX_MESSAGE_LEN = 2000 # Discord message character limit
|
MAX_MESSAGE_LEN = 2000 # Discord message character limit
|
||||||
|
TYPING_INTERVAL_S = 8
|
||||||
|
|
||||||
|
|
||||||
class DiscordConfig(Base):
|
class DiscordConfig(Base):
|
||||||
@@ -28,145 +40,155 @@ class DiscordConfig(Base):
|
|||||||
enabled: bool = False
|
enabled: bool = False
|
||||||
token: str = ""
|
token: str = ""
|
||||||
allow_from: list[str] = Field(default_factory=list)
|
allow_from: list[str] = Field(default_factory=list)
|
||||||
gateway_url: str = "wss://gateway.discord.gg/?v=10&encoding=json"
|
|
||||||
intents: int = 37377
|
intents: int = 37377
|
||||||
group_policy: Literal["mention", "open"] = "mention"
|
group_policy: Literal["mention", "open"] = "mention"
|
||||||
|
read_receipt_emoji: str = "👀"
|
||||||
|
working_emoji: str = "🔧"
|
||||||
|
working_emoji_delay: float = 2.0
|
||||||
|
|
||||||
|
|
||||||
class DiscordChannel(BaseChannel):
|
if DISCORD_AVAILABLE:
|
||||||
"""Discord channel using Gateway websocket."""
|
|
||||||
|
|
||||||
name = "discord"
|
class DiscordBotClient(discord.Client):
|
||||||
display_name = "Discord"
|
"""discord.py client that forwards events to the channel."""
|
||||||
|
|
||||||
@classmethod
|
def __init__(self, channel: DiscordChannel, *, intents: discord.Intents) -> None:
|
||||||
def default_config(cls) -> dict[str, Any]:
|
super().__init__(intents=intents)
|
||||||
return DiscordConfig().model_dump(by_alias=True)
|
self._channel = channel
|
||||||
|
self.tree = app_commands.CommandTree(self)
|
||||||
|
self._register_app_commands()
|
||||||
|
|
||||||
def __init__(self, config: Any, bus: MessageBus):
|
async def on_ready(self) -> None:
|
||||||
if isinstance(config, dict):
|
self._channel._bot_user_id = str(self.user.id) if self.user else None
|
||||||
config = DiscordConfig.model_validate(config)
|
logger.info("Discord bot connected as user {}", self._channel._bot_user_id)
|
||||||
super().__init__(config, bus)
|
|
||||||
self.config: DiscordConfig = config
|
|
||||||
self._ws: websockets.WebSocketClientProtocol | None = None
|
|
||||||
self._seq: int | None = None
|
|
||||||
self._heartbeat_task: asyncio.Task | None = None
|
|
||||||
self._typing_tasks: dict[str, asyncio.Task] = {}
|
|
||||||
self._http: httpx.AsyncClient | None = None
|
|
||||||
self._bot_user_id: str | None = None
|
|
||||||
|
|
||||||
async def start(self) -> None:
|
|
||||||
"""Start the Discord gateway connection."""
|
|
||||||
if not self.config.token:
|
|
||||||
logger.error("Discord bot token not configured")
|
|
||||||
return
|
|
||||||
|
|
||||||
self._running = True
|
|
||||||
self._http = httpx.AsyncClient(timeout=30.0)
|
|
||||||
|
|
||||||
while self._running:
|
|
||||||
try:
|
try:
|
||||||
logger.info("Connecting to Discord gateway...")
|
synced = await self.tree.sync()
|
||||||
async with websockets.connect(self.config.gateway_url) as ws:
|
logger.info("Discord app commands synced: {}", len(synced))
|
||||||
self._ws = ws
|
|
||||||
await self._gateway_loop()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("Discord gateway error: {}", e)
|
logger.warning("Discord app command sync failed: {}", e)
|
||||||
if self._running:
|
|
||||||
logger.info("Reconnecting to Discord gateway in 5 seconds...")
|
|
||||||
await asyncio.sleep(5)
|
|
||||||
|
|
||||||
async def stop(self) -> None:
|
async def on_message(self, message: discord.Message) -> None:
|
||||||
"""Stop the Discord channel."""
|
await self._channel._handle_discord_message(message)
|
||||||
self._running = False
|
|
||||||
if self._heartbeat_task:
|
|
||||||
self._heartbeat_task.cancel()
|
|
||||||
self._heartbeat_task = None
|
|
||||||
for task in self._typing_tasks.values():
|
|
||||||
task.cancel()
|
|
||||||
self._typing_tasks.clear()
|
|
||||||
if self._ws:
|
|
||||||
await self._ws.close()
|
|
||||||
self._ws = None
|
|
||||||
if self._http:
|
|
||||||
await self._http.aclose()
|
|
||||||
self._http = None
|
|
||||||
|
|
||||||
async def send(self, msg: OutboundMessage) -> None:
|
async def _reply_ephemeral(self, interaction: discord.Interaction, text: str) -> bool:
|
||||||
"""Send a message through Discord REST API, including file attachments."""
|
"""Send an ephemeral interaction response and report success."""
|
||||||
if not self._http:
|
try:
|
||||||
logger.warning("Discord HTTP client not initialized")
|
await interaction.response.send_message(text, ephemeral=True)
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Discord interaction response failed: {}", e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _forward_slash_command(
|
||||||
|
self,
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
command_text: str,
|
||||||
|
) -> None:
|
||||||
|
sender_id = str(interaction.user.id)
|
||||||
|
channel_id = interaction.channel_id
|
||||||
|
|
||||||
|
if channel_id is None:
|
||||||
|
logger.warning("Discord slash command missing channel_id: {}", command_text)
|
||||||
return
|
return
|
||||||
|
|
||||||
url = f"{DISCORD_API_BASE}/channels/{msg.chat_id}/messages"
|
if not self._channel.is_allowed(sender_id):
|
||||||
headers = {"Authorization": f"Bot {self.config.token}"}
|
await self._reply_ephemeral(interaction, "You are not allowed to use this bot.")
|
||||||
|
return
|
||||||
|
|
||||||
|
await self._reply_ephemeral(interaction, f"Processing {command_text}...")
|
||||||
|
|
||||||
|
await self._channel._handle_message(
|
||||||
|
sender_id=sender_id,
|
||||||
|
chat_id=str(channel_id),
|
||||||
|
content=command_text,
|
||||||
|
metadata={
|
||||||
|
"interaction_id": str(interaction.id),
|
||||||
|
"guild_id": str(interaction.guild_id) if interaction.guild_id else None,
|
||||||
|
"is_slash_command": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
def _register_app_commands(self) -> None:
|
||||||
|
commands = (
|
||||||
|
("new", "Start a new conversation", "/new"),
|
||||||
|
("stop", "Stop the current task", "/stop"),
|
||||||
|
("restart", "Restart the bot", "/restart"),
|
||||||
|
("status", "Show bot status", "/status"),
|
||||||
|
)
|
||||||
|
|
||||||
|
for name, description, command_text in commands:
|
||||||
|
@self.tree.command(name=name, description=description)
|
||||||
|
async def command_handler(
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
_command_text: str = command_text,
|
||||||
|
) -> None:
|
||||||
|
await self._forward_slash_command(interaction, _command_text)
|
||||||
|
|
||||||
|
@self.tree.command(name="help", description="Show available commands")
|
||||||
|
async def help_command(interaction: discord.Interaction) -> None:
|
||||||
|
sender_id = str(interaction.user.id)
|
||||||
|
if not self._channel.is_allowed(sender_id):
|
||||||
|
await self._reply_ephemeral(interaction, "You are not allowed to use this bot.")
|
||||||
|
return
|
||||||
|
await self._reply_ephemeral(interaction, build_help_text())
|
||||||
|
|
||||||
|
@self.tree.error
|
||||||
|
async def on_app_command_error(
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
error: app_commands.AppCommandError,
|
||||||
|
) -> None:
|
||||||
|
command_name = interaction.command.qualified_name if interaction.command else "?"
|
||||||
|
logger.warning(
|
||||||
|
"Discord app command failed user={} channel={} cmd={} error={}",
|
||||||
|
interaction.user.id,
|
||||||
|
interaction.channel_id,
|
||||||
|
command_name,
|
||||||
|
error,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_outbound(self, msg: OutboundMessage) -> None:
|
||||||
|
"""Send a nanobot outbound message using Discord transport rules."""
|
||||||
|
channel_id = int(msg.chat_id)
|
||||||
|
|
||||||
|
channel = self.get_channel(channel_id)
|
||||||
|
if channel is None:
|
||||||
try:
|
try:
|
||||||
|
channel = await self.fetch_channel(channel_id)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Discord channel {} unavailable: {}", msg.chat_id, e)
|
||||||
|
return
|
||||||
|
|
||||||
|
reference, mention_settings = self._build_reply_context(channel, msg.reply_to)
|
||||||
sent_media = False
|
sent_media = False
|
||||||
failed_media: list[str] = []
|
failed_media: list[str] = []
|
||||||
|
|
||||||
# Send file attachments first
|
for index, media_path in enumerate(msg.media or []):
|
||||||
for media_path in msg.media or []:
|
if await self._send_file(
|
||||||
if await self._send_file(url, headers, media_path, reply_to=msg.reply_to):
|
channel,
|
||||||
|
media_path,
|
||||||
|
reference=reference if index == 0 else None,
|
||||||
|
mention_settings=mention_settings,
|
||||||
|
):
|
||||||
sent_media = True
|
sent_media = True
|
||||||
else:
|
else:
|
||||||
failed_media.append(Path(media_path).name)
|
failed_media.append(Path(media_path).name)
|
||||||
|
|
||||||
# Send text content
|
for index, chunk in enumerate(self._build_chunks(msg.content or "", failed_media, sent_media)):
|
||||||
chunks = split_message(msg.content or "", MAX_MESSAGE_LEN)
|
kwargs: dict[str, Any] = {"content": chunk}
|
||||||
if not chunks and failed_media and not sent_media:
|
if index == 0 and reference is not None and not sent_media:
|
||||||
chunks = split_message(
|
kwargs["reference"] = reference
|
||||||
"\n".join(f"[attachment: {name} - send failed]" for name in failed_media),
|
kwargs["allowed_mentions"] = mention_settings
|
||||||
MAX_MESSAGE_LEN,
|
await channel.send(**kwargs)
|
||||||
)
|
|
||||||
if not chunks:
|
|
||||||
return
|
|
||||||
|
|
||||||
for i, chunk in enumerate(chunks):
|
|
||||||
payload: dict[str, Any] = {"content": chunk}
|
|
||||||
|
|
||||||
# Let the first successful attachment carry the reply if present.
|
|
||||||
if i == 0 and msg.reply_to and not sent_media:
|
|
||||||
payload["message_reference"] = {"message_id": msg.reply_to}
|
|
||||||
payload["allowed_mentions"] = {"replied_user": False}
|
|
||||||
|
|
||||||
if not await self._send_payload(url, headers, payload):
|
|
||||||
break # Abort remaining chunks on failure
|
|
||||||
finally:
|
|
||||||
await self._stop_typing(msg.chat_id)
|
|
||||||
|
|
||||||
async def _send_payload(
|
|
||||||
self, url: str, headers: dict[str, str], payload: dict[str, Any]
|
|
||||||
) -> bool:
|
|
||||||
"""Send a single Discord API payload with retry on rate-limit. Returns True on success."""
|
|
||||||
for attempt in range(3):
|
|
||||||
try:
|
|
||||||
response = await self._http.post(url, headers=headers, json=payload)
|
|
||||||
if response.status_code == 429:
|
|
||||||
data = response.json()
|
|
||||||
retry_after = float(data.get("retry_after", 1.0))
|
|
||||||
logger.warning("Discord rate limited, retrying in {}s", retry_after)
|
|
||||||
await asyncio.sleep(retry_after)
|
|
||||||
continue
|
|
||||||
response.raise_for_status()
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
if attempt == 2:
|
|
||||||
logger.error("Error sending Discord message: {}", e)
|
|
||||||
else:
|
|
||||||
await asyncio.sleep(1)
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def _send_file(
|
async def _send_file(
|
||||||
self,
|
self,
|
||||||
url: str,
|
channel: Messageable,
|
||||||
headers: dict[str, str],
|
|
||||||
file_path: str,
|
file_path: str,
|
||||||
reply_to: str | None = None,
|
*,
|
||||||
|
reference: discord.PartialMessage | None,
|
||||||
|
mention_settings: discord.AllowedMentions,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Send a file attachment via Discord REST API using multipart/form-data."""
|
"""Send a file attachment via discord.py."""
|
||||||
path = Path(file_path)
|
path = Path(file_path)
|
||||||
if not path.is_file():
|
if not path.is_file():
|
||||||
logger.warning("Discord file not found, skipping: {}", file_path)
|
logger.warning("Discord file not found, skipping: {}", file_path)
|
||||||
@@ -176,220 +198,319 @@ class DiscordChannel(BaseChannel):
|
|||||||
logger.warning("Discord file too large (>20MB), skipping: {}", path.name)
|
logger.warning("Discord file too large (>20MB), skipping: {}", path.name)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
payload_json: dict[str, Any] = {}
|
|
||||||
if reply_to:
|
|
||||||
payload_json["message_reference"] = {"message_id": reply_to}
|
|
||||||
payload_json["allowed_mentions"] = {"replied_user": False}
|
|
||||||
|
|
||||||
for attempt in range(3):
|
|
||||||
try:
|
try:
|
||||||
with open(path, "rb") as f:
|
kwargs: dict[str, Any] = {"file": discord.File(path)}
|
||||||
files = {"files[0]": (path.name, f, "application/octet-stream")}
|
if reference is not None:
|
||||||
data: dict[str, Any] = {}
|
kwargs["reference"] = reference
|
||||||
if payload_json:
|
kwargs["allowed_mentions"] = mention_settings
|
||||||
data["payload_json"] = json.dumps(payload_json)
|
await channel.send(**kwargs)
|
||||||
response = await self._http.post(
|
|
||||||
url, headers=headers, files=files, data=data
|
|
||||||
)
|
|
||||||
if response.status_code == 429:
|
|
||||||
resp_data = response.json()
|
|
||||||
retry_after = float(resp_data.get("retry_after", 1.0))
|
|
||||||
logger.warning("Discord rate limited, retrying in {}s", retry_after)
|
|
||||||
await asyncio.sleep(retry_after)
|
|
||||||
continue
|
|
||||||
response.raise_for_status()
|
|
||||||
logger.info("Discord file sent: {}", path.name)
|
logger.info("Discord file sent: {}", path.name)
|
||||||
return True
|
return True
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
if attempt == 2:
|
|
||||||
logger.error("Error sending Discord file {}: {}", path.name, e)
|
logger.error("Error sending Discord file {}: {}", path.name, e)
|
||||||
else:
|
|
||||||
await asyncio.sleep(1)
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
async def _gateway_loop(self) -> None:
|
@staticmethod
|
||||||
"""Main gateway loop: identify, heartbeat, dispatch events."""
|
def _build_chunks(content: str, failed_media: list[str], sent_media: bool) -> list[str]:
|
||||||
if not self._ws:
|
"""Build outbound text chunks, including attachment-failure fallback text."""
|
||||||
|
chunks = split_message(content, MAX_MESSAGE_LEN)
|
||||||
|
if chunks or not failed_media or sent_media:
|
||||||
|
return chunks
|
||||||
|
fallback = "\n".join(f"[attachment: {name} - send failed]" for name in failed_media)
|
||||||
|
return split_message(fallback, MAX_MESSAGE_LEN)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_reply_context(
|
||||||
|
channel: Messageable,
|
||||||
|
reply_to: str | None,
|
||||||
|
) -> tuple[discord.PartialMessage | None, discord.AllowedMentions]:
|
||||||
|
"""Build reply context for outbound messages."""
|
||||||
|
mention_settings = discord.AllowedMentions(replied_user=False)
|
||||||
|
if not reply_to:
|
||||||
|
return None, mention_settings
|
||||||
|
try:
|
||||||
|
message_id = int(reply_to)
|
||||||
|
except ValueError:
|
||||||
|
logger.warning("Invalid Discord reply target: {}", reply_to)
|
||||||
|
return None, mention_settings
|
||||||
|
|
||||||
|
return channel.get_partial_message(message_id), mention_settings
|
||||||
|
|
||||||
|
|
||||||
|
class DiscordChannel(BaseChannel):
|
||||||
|
"""Discord channel using discord.py."""
|
||||||
|
|
||||||
|
name = "discord"
|
||||||
|
display_name = "Discord"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def default_config(cls) -> dict[str, Any]:
|
||||||
|
return DiscordConfig().model_dump(by_alias=True)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _channel_key(channel_or_id: Any) -> str:
|
||||||
|
"""Normalize channel-like objects and ids to a stable string key."""
|
||||||
|
channel_id = getattr(channel_or_id, "id", channel_or_id)
|
||||||
|
return str(channel_id)
|
||||||
|
|
||||||
|
def __init__(self, config: Any, bus: MessageBus):
|
||||||
|
if isinstance(config, dict):
|
||||||
|
config = DiscordConfig.model_validate(config)
|
||||||
|
super().__init__(config, bus)
|
||||||
|
self.config: DiscordConfig = config
|
||||||
|
self._client: DiscordBotClient | None = None
|
||||||
|
self._typing_tasks: dict[str, asyncio.Task[None]] = {}
|
||||||
|
self._bot_user_id: str | None = None
|
||||||
|
self._pending_reactions: dict[str, Any] = {} # chat_id -> message object
|
||||||
|
self._working_emoji_tasks: dict[str, asyncio.Task[None]] = {}
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
"""Start the Discord client."""
|
||||||
|
if not DISCORD_AVAILABLE:
|
||||||
|
logger.error("discord.py not installed. Run: pip install nanobot-ai[discord]")
|
||||||
return
|
return
|
||||||
|
|
||||||
async for raw in self._ws:
|
if not self.config.token:
|
||||||
try:
|
logger.error("Discord bot token not configured")
|
||||||
data = json.loads(raw)
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
logger.warning("Invalid JSON from Discord gateway: {}", raw[:100])
|
|
||||||
continue
|
|
||||||
|
|
||||||
op = data.get("op")
|
|
||||||
event_type = data.get("t")
|
|
||||||
seq = data.get("s")
|
|
||||||
payload = data.get("d")
|
|
||||||
|
|
||||||
if seq is not None:
|
|
||||||
self._seq = seq
|
|
||||||
|
|
||||||
if op == 10:
|
|
||||||
# HELLO: start heartbeat and identify
|
|
||||||
interval_ms = payload.get("heartbeat_interval", 45000)
|
|
||||||
await self._start_heartbeat(interval_ms / 1000)
|
|
||||||
await self._identify()
|
|
||||||
elif op == 0 and event_type == "READY":
|
|
||||||
logger.info("Discord gateway READY")
|
|
||||||
# Capture bot user ID for mention detection
|
|
||||||
user_data = payload.get("user") or {}
|
|
||||||
self._bot_user_id = user_data.get("id")
|
|
||||||
logger.info("Discord bot connected as user {}", self._bot_user_id)
|
|
||||||
elif op == 0 and event_type == "MESSAGE_CREATE":
|
|
||||||
await self._handle_message_create(payload)
|
|
||||||
elif op == 7:
|
|
||||||
# RECONNECT: exit loop to reconnect
|
|
||||||
logger.info("Discord gateway requested reconnect")
|
|
||||||
break
|
|
||||||
elif op == 9:
|
|
||||||
# INVALID_SESSION: reconnect
|
|
||||||
logger.warning("Discord gateway invalid session")
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _identify(self) -> None:
|
|
||||||
"""Send IDENTIFY payload."""
|
|
||||||
if not self._ws:
|
|
||||||
return
|
return
|
||||||
|
|
||||||
identify = {
|
|
||||||
"op": 2,
|
|
||||||
"d": {
|
|
||||||
"token": self.config.token,
|
|
||||||
"intents": self.config.intents,
|
|
||||||
"properties": {
|
|
||||||
"os": "nanobot",
|
|
||||||
"browser": "nanobot",
|
|
||||||
"device": "nanobot",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
await self._ws.send(json.dumps(identify))
|
|
||||||
|
|
||||||
async def _start_heartbeat(self, interval_s: float) -> None:
|
|
||||||
"""Start or restart the heartbeat loop."""
|
|
||||||
if self._heartbeat_task:
|
|
||||||
self._heartbeat_task.cancel()
|
|
||||||
|
|
||||||
async def heartbeat_loop() -> None:
|
|
||||||
while self._running and self._ws:
|
|
||||||
payload = {"op": 1, "d": self._seq}
|
|
||||||
try:
|
try:
|
||||||
await self._ws.send(json.dumps(payload))
|
intents = discord.Intents.none()
|
||||||
|
intents.value = self.config.intents
|
||||||
|
self._client = DiscordBotClient(self, intents=intents)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("Discord heartbeat failed: {}", e)
|
logger.error("Failed to initialize Discord client: {}", e)
|
||||||
break
|
self._client = None
|
||||||
await asyncio.sleep(interval_s)
|
self._running = False
|
||||||
|
|
||||||
self._heartbeat_task = asyncio.create_task(heartbeat_loop())
|
|
||||||
|
|
||||||
async def _handle_message_create(self, payload: dict[str, Any]) -> None:
|
|
||||||
"""Handle incoming Discord messages."""
|
|
||||||
author = payload.get("author") or {}
|
|
||||||
if author.get("bot"):
|
|
||||||
return
|
return
|
||||||
|
|
||||||
sender_id = str(author.get("id", ""))
|
self._running = True
|
||||||
channel_id = str(payload.get("channel_id", ""))
|
logger.info("Starting Discord client via discord.py...")
|
||||||
content = payload.get("content") or ""
|
|
||||||
guild_id = payload.get("guild_id")
|
|
||||||
|
|
||||||
if not sender_id or not channel_id:
|
|
||||||
return
|
|
||||||
|
|
||||||
if not self.is_allowed(sender_id):
|
|
||||||
return
|
|
||||||
|
|
||||||
# Check group channel policy (DMs always respond if is_allowed passes)
|
|
||||||
if guild_id is not None:
|
|
||||||
if not self._should_respond_in_group(payload, content):
|
|
||||||
return
|
|
||||||
|
|
||||||
content_parts = [content] if content else []
|
|
||||||
media_paths: list[str] = []
|
|
||||||
media_dir = get_media_dir("discord")
|
|
||||||
|
|
||||||
for attachment in payload.get("attachments") or []:
|
|
||||||
url = attachment.get("url")
|
|
||||||
filename = attachment.get("filename") or "attachment"
|
|
||||||
size = attachment.get("size") or 0
|
|
||||||
if not url or not self._http:
|
|
||||||
continue
|
|
||||||
if size and size > MAX_ATTACHMENT_BYTES:
|
|
||||||
content_parts.append(f"[attachment: {filename} - too large]")
|
|
||||||
continue
|
|
||||||
try:
|
try:
|
||||||
media_dir.mkdir(parents=True, exist_ok=True)
|
await self._client.start(self.config.token)
|
||||||
file_path = media_dir / f"{attachment.get('id', 'file')}_{filename.replace('/', '_')}"
|
except asyncio.CancelledError:
|
||||||
resp = await self._http.get(url)
|
raise
|
||||||
resp.raise_for_status()
|
|
||||||
file_path.write_bytes(resp.content)
|
|
||||||
media_paths.append(str(file_path))
|
|
||||||
content_parts.append(f"[attachment: {file_path}]")
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("Failed to download Discord attachment: {}", e)
|
logger.error("Discord client startup failed: {}", e)
|
||||||
content_parts.append(f"[attachment: {filename} - download failed]")
|
finally:
|
||||||
|
self._running = False
|
||||||
|
await self._reset_runtime_state(close_client=True)
|
||||||
|
|
||||||
reply_to = (payload.get("referenced_message") or {}).get("id")
|
async def stop(self) -> None:
|
||||||
|
"""Stop the Discord channel."""
|
||||||
|
self._running = False
|
||||||
|
await self._reset_runtime_state(close_client=True)
|
||||||
|
|
||||||
await self._start_typing(channel_id)
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
|
"""Send a message through Discord using discord.py."""
|
||||||
|
client = self._client
|
||||||
|
if client is None or not client.is_ready():
|
||||||
|
logger.warning("Discord client not ready; dropping outbound message")
|
||||||
|
return
|
||||||
|
|
||||||
|
is_progress = bool((msg.metadata or {}).get("_progress"))
|
||||||
|
|
||||||
|
try:
|
||||||
|
await client.send_outbound(msg)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Error sending Discord message: {}", e)
|
||||||
|
finally:
|
||||||
|
if not is_progress:
|
||||||
|
await self._stop_typing(msg.chat_id)
|
||||||
|
await self._clear_reactions(msg.chat_id)
|
||||||
|
|
||||||
|
async def _handle_discord_message(self, message: discord.Message) -> None:
|
||||||
|
"""Handle incoming Discord messages from discord.py."""
|
||||||
|
if message.author.bot:
|
||||||
|
return
|
||||||
|
|
||||||
|
sender_id = str(message.author.id)
|
||||||
|
channel_id = self._channel_key(message.channel)
|
||||||
|
content = message.content or ""
|
||||||
|
|
||||||
|
if not self._should_accept_inbound(message, sender_id, content):
|
||||||
|
return
|
||||||
|
|
||||||
|
media_paths, attachment_markers = await self._download_attachments(message.attachments)
|
||||||
|
full_content = self._compose_inbound_content(content, attachment_markers)
|
||||||
|
metadata = self._build_inbound_metadata(message)
|
||||||
|
|
||||||
|
await self._start_typing(message.channel)
|
||||||
|
|
||||||
|
# Add read receipt reaction immediately, working emoji after delay
|
||||||
|
channel_id = self._channel_key(message.channel)
|
||||||
|
try:
|
||||||
|
await message.add_reaction(self.config.read_receipt_emoji)
|
||||||
|
self._pending_reactions[channel_id] = message
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Failed to add read receipt reaction: {}", e)
|
||||||
|
|
||||||
|
# Delayed working indicator (cosmetic — not tied to subagent lifecycle)
|
||||||
|
async def _delayed_working_emoji() -> None:
|
||||||
|
await asyncio.sleep(self.config.working_emoji_delay)
|
||||||
|
try:
|
||||||
|
await message.add_reaction(self.config.working_emoji)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
self._working_emoji_tasks[channel_id] = asyncio.create_task(_delayed_working_emoji())
|
||||||
|
|
||||||
|
try:
|
||||||
await self._handle_message(
|
await self._handle_message(
|
||||||
sender_id=sender_id,
|
sender_id=sender_id,
|
||||||
chat_id=channel_id,
|
chat_id=channel_id,
|
||||||
content="\n".join(p for p in content_parts if p) or "[empty message]",
|
content=full_content,
|
||||||
media=media_paths,
|
media=media_paths,
|
||||||
metadata={
|
metadata=metadata,
|
||||||
"message_id": str(payload.get("id", "")),
|
|
||||||
"guild_id": guild_id,
|
|
||||||
"reply_to": reply_to,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
|
except Exception:
|
||||||
|
await self._clear_reactions(channel_id)
|
||||||
|
await self._stop_typing(channel_id)
|
||||||
|
raise
|
||||||
|
|
||||||
def _should_respond_in_group(self, payload: dict[str, Any], content: str) -> bool:
|
async def _on_message(self, message: discord.Message) -> None:
|
||||||
"""Check if bot should respond in a group channel based on policy."""
|
"""Backward-compatible alias for legacy tests/callers."""
|
||||||
|
await self._handle_discord_message(message)
|
||||||
|
|
||||||
|
def _should_accept_inbound(
|
||||||
|
self,
|
||||||
|
message: discord.Message,
|
||||||
|
sender_id: str,
|
||||||
|
content: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Check if inbound Discord message should be processed."""
|
||||||
|
if not self.is_allowed(sender_id):
|
||||||
|
return False
|
||||||
|
if message.guild is not None and not self._should_respond_in_group(message, content):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def _download_attachments(
|
||||||
|
self,
|
||||||
|
attachments: list[discord.Attachment],
|
||||||
|
) -> tuple[list[str], list[str]]:
|
||||||
|
"""Download supported attachments and return paths + display markers."""
|
||||||
|
media_paths: list[str] = []
|
||||||
|
markers: list[str] = []
|
||||||
|
media_dir = get_media_dir("discord")
|
||||||
|
|
||||||
|
for attachment in attachments:
|
||||||
|
filename = attachment.filename or "attachment"
|
||||||
|
if attachment.size and attachment.size > MAX_ATTACHMENT_BYTES:
|
||||||
|
markers.append(f"[attachment: {filename} - too large]")
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
media_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
safe_name = safe_filename(filename)
|
||||||
|
file_path = media_dir / f"{attachment.id}_{safe_name}"
|
||||||
|
await attachment.save(file_path)
|
||||||
|
media_paths.append(str(file_path))
|
||||||
|
markers.append(f"[attachment: {file_path.name}]")
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Failed to download Discord attachment: {}", e)
|
||||||
|
markers.append(f"[attachment: {filename} - download failed]")
|
||||||
|
|
||||||
|
return media_paths, markers
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _compose_inbound_content(content: str, attachment_markers: list[str]) -> str:
|
||||||
|
"""Combine message text with attachment markers."""
|
||||||
|
content_parts = [content] if content else []
|
||||||
|
content_parts.extend(attachment_markers)
|
||||||
|
return "\n".join(part for part in content_parts if part) or "[empty message]"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_inbound_metadata(message: discord.Message) -> dict[str, str | None]:
|
||||||
|
"""Build metadata for inbound Discord messages."""
|
||||||
|
reply_to = str(message.reference.message_id) if message.reference and message.reference.message_id else None
|
||||||
|
return {
|
||||||
|
"message_id": str(message.id),
|
||||||
|
"guild_id": str(message.guild.id) if message.guild else None,
|
||||||
|
"reply_to": reply_to,
|
||||||
|
}
|
||||||
|
|
||||||
|
def _should_respond_in_group(self, message: discord.Message, content: str) -> bool:
|
||||||
|
"""Check if the bot should respond in a guild channel based on policy."""
|
||||||
if self.config.group_policy == "open":
|
if self.config.group_policy == "open":
|
||||||
return True
|
return True
|
||||||
|
|
||||||
if self.config.group_policy == "mention":
|
if self.config.group_policy == "mention":
|
||||||
# Check if bot was mentioned in the message
|
bot_user_id = self._bot_user_id
|
||||||
if self._bot_user_id:
|
if bot_user_id is None:
|
||||||
# Check mentions array
|
logger.debug("Discord message in {} ignored (bot identity unavailable)", message.channel.id)
|
||||||
mentions = payload.get("mentions") or []
|
return False
|
||||||
for mention in mentions:
|
|
||||||
if str(mention.get("id")) == self._bot_user_id:
|
if any(str(user.id) == bot_user_id for user in message.mentions):
|
||||||
return True
|
return True
|
||||||
# Also check content for mention format <@USER_ID>
|
if f"<@{bot_user_id}>" in content or f"<@!{bot_user_id}>" in content:
|
||||||
if f"<@{self._bot_user_id}>" in content or f"<@!{self._bot_user_id}>" in content:
|
|
||||||
return True
|
return True
|
||||||
logger.debug("Discord message in {} ignored (bot not mentioned)", payload.get("channel_id"))
|
|
||||||
|
logger.debug("Discord message in {} ignored (bot not mentioned)", message.channel.id)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
async def _start_typing(self, channel_id: str) -> None:
|
async def _start_typing(self, channel: Messageable) -> None:
|
||||||
"""Start periodic typing indicator for a channel."""
|
"""Start periodic typing indicator for a channel."""
|
||||||
|
channel_id = self._channel_key(channel)
|
||||||
await self._stop_typing(channel_id)
|
await self._stop_typing(channel_id)
|
||||||
|
|
||||||
async def typing_loop() -> None:
|
async def typing_loop() -> None:
|
||||||
url = f"{DISCORD_API_BASE}/channels/{channel_id}/typing"
|
|
||||||
headers = {"Authorization": f"Bot {self.config.token}"}
|
|
||||||
while self._running:
|
while self._running:
|
||||||
try:
|
try:
|
||||||
await self._http.post(url, headers=headers)
|
async with channel.typing():
|
||||||
|
await asyncio.sleep(TYPING_INTERVAL_S)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
return
|
return
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug("Discord typing indicator failed for {}: {}", channel_id, e)
|
logger.debug("Discord typing indicator failed for {}: {}", channel_id, e)
|
||||||
return
|
return
|
||||||
await asyncio.sleep(8)
|
|
||||||
|
|
||||||
self._typing_tasks[channel_id] = asyncio.create_task(typing_loop())
|
self._typing_tasks[channel_id] = asyncio.create_task(typing_loop())
|
||||||
|
|
||||||
async def _stop_typing(self, channel_id: str) -> None:
|
async def _stop_typing(self, channel_id: str) -> None:
|
||||||
"""Stop typing indicator for a channel."""
|
"""Stop typing indicator for a channel."""
|
||||||
task = self._typing_tasks.pop(channel_id, None)
|
task = self._typing_tasks.pop(self._channel_key(channel_id), None)
|
||||||
if task:
|
if task is None:
|
||||||
|
return
|
||||||
task.cancel()
|
task.cancel()
|
||||||
|
try:
|
||||||
|
await task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
async def _clear_reactions(self, chat_id: str) -> None:
|
||||||
|
"""Remove all pending reactions after bot replies."""
|
||||||
|
# Cancel delayed working emoji if it hasn't fired yet
|
||||||
|
task = self._working_emoji_tasks.pop(chat_id, None)
|
||||||
|
if task and not task.done():
|
||||||
|
task.cancel()
|
||||||
|
|
||||||
|
msg_obj = self._pending_reactions.pop(chat_id, None)
|
||||||
|
if msg_obj is None:
|
||||||
|
return
|
||||||
|
bot_user = self._client.user if self._client else None
|
||||||
|
for emoji in (self.config.read_receipt_emoji, self.config.working_emoji):
|
||||||
|
try:
|
||||||
|
await msg_obj.remove_reaction(emoji, bot_user)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def _cancel_all_typing(self) -> None:
|
||||||
|
"""Stop all typing tasks."""
|
||||||
|
channel_ids = list(self._typing_tasks)
|
||||||
|
for channel_id in channel_ids:
|
||||||
|
await self._stop_typing(channel_id)
|
||||||
|
|
||||||
|
async def _reset_runtime_state(self, close_client: bool) -> None:
|
||||||
|
"""Reset client and typing state."""
|
||||||
|
await self._cancel_all_typing()
|
||||||
|
if close_client and self._client is not None and not self._client.is_closed():
|
||||||
|
try:
|
||||||
|
await self._client.close()
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Discord client close failed: {}", e)
|
||||||
|
self._client = None
|
||||||
|
self._bot_user_id = None
|
||||||
|
|||||||
+115
-7
@@ -3,6 +3,8 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import mimetypes
|
import mimetypes
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Literal, TypeAlias
|
from typing import Any, Literal, TypeAlias
|
||||||
|
|
||||||
@@ -28,7 +30,7 @@ try:
|
|||||||
RoomSendError,
|
RoomSendError,
|
||||||
RoomTypingError,
|
RoomTypingError,
|
||||||
SyncError,
|
SyncError,
|
||||||
UploadError,
|
UploadError, RoomSendResponse,
|
||||||
)
|
)
|
||||||
from nio.crypto.attachments import decrypt_attachment
|
from nio.crypto.attachments import decrypt_attachment
|
||||||
from nio.exceptions import EncryptionError
|
from nio.exceptions import EncryptionError
|
||||||
@@ -97,6 +99,22 @@ MATRIX_HTML_CLEANER = nh3.Cleaner(
|
|||||||
link_rel="noopener noreferrer",
|
link_rel="noopener noreferrer",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _StreamBuf:
|
||||||
|
"""
|
||||||
|
Represents a buffer for managing LLM response stream data.
|
||||||
|
|
||||||
|
:ivar text: Stores the text content of the buffer.
|
||||||
|
:type text: str
|
||||||
|
:ivar event_id: Identifier for the associated event. None indicates no
|
||||||
|
specific event association.
|
||||||
|
:type event_id: str | None
|
||||||
|
:ivar last_edit: Timestamp of the most recent edit to the buffer.
|
||||||
|
:type last_edit: float
|
||||||
|
"""
|
||||||
|
text: str = ""
|
||||||
|
event_id: str | None = None
|
||||||
|
last_edit: float = 0.0
|
||||||
|
|
||||||
def _render_markdown_html(text: str) -> str | None:
|
def _render_markdown_html(text: str) -> str | None:
|
||||||
"""Render markdown to sanitized HTML; returns None for plain text."""
|
"""Render markdown to sanitized HTML; returns None for plain text."""
|
||||||
@@ -114,12 +132,47 @@ def _render_markdown_html(text: str) -> str | None:
|
|||||||
return formatted
|
return formatted
|
||||||
|
|
||||||
|
|
||||||
def _build_matrix_text_content(text: str) -> dict[str, object]:
|
def _build_matrix_text_content(
|
||||||
"""Build Matrix m.text payload with optional HTML formatted_body."""
|
text: str,
|
||||||
|
event_id: str | None = None,
|
||||||
|
thread_relates_to: dict[str, object] | None = None,
|
||||||
|
) -> dict[str, object]:
|
||||||
|
"""
|
||||||
|
Constructs and returns a dictionary representing the matrix text content with optional
|
||||||
|
HTML formatting and reference to an existing event for replacement. This function is
|
||||||
|
primarily used to create content payloads compatible with the Matrix messaging protocol.
|
||||||
|
|
||||||
|
:param text: The plain text content to include in the message.
|
||||||
|
:type text: str
|
||||||
|
:param event_id: Optional ID of the event to replace. If provided, the function will
|
||||||
|
include information indicating that the message is a replacement of the specified
|
||||||
|
event.
|
||||||
|
:type event_id: str | None
|
||||||
|
:param thread_relates_to: Optional Matrix thread relation metadata. For edits this is
|
||||||
|
stored in ``m.new_content`` so the replacement remains in the same thread.
|
||||||
|
:type thread_relates_to: dict[str, object] | None
|
||||||
|
:return: A dictionary containing the matrix text content, potentially enriched with
|
||||||
|
HTML formatting and replacement metadata if applicable.
|
||||||
|
:rtype: dict[str, object]
|
||||||
|
"""
|
||||||
content: dict[str, object] = {"msgtype": "m.text", "body": text, "m.mentions": {}}
|
content: dict[str, object] = {"msgtype": "m.text", "body": text, "m.mentions": {}}
|
||||||
if html := _render_markdown_html(text):
|
if html := _render_markdown_html(text):
|
||||||
content["format"] = MATRIX_HTML_FORMAT
|
content["format"] = MATRIX_HTML_FORMAT
|
||||||
content["formatted_body"] = html
|
content["formatted_body"] = html
|
||||||
|
if event_id:
|
||||||
|
content["m.new_content"] = {
|
||||||
|
"body": text,
|
||||||
|
"msgtype": "m.text",
|
||||||
|
}
|
||||||
|
content["m.relates_to"] = {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": event_id,
|
||||||
|
}
|
||||||
|
if thread_relates_to:
|
||||||
|
content["m.new_content"]["m.relates_to"] = thread_relates_to
|
||||||
|
elif thread_relates_to:
|
||||||
|
content["m.relates_to"] = thread_relates_to
|
||||||
|
|
||||||
return content
|
return content
|
||||||
|
|
||||||
|
|
||||||
@@ -159,7 +212,8 @@ class MatrixConfig(Base):
|
|||||||
allow_from: list[str] = Field(default_factory=list)
|
allow_from: list[str] = Field(default_factory=list)
|
||||||
group_policy: Literal["open", "mention", "allowlist"] = "open"
|
group_policy: Literal["open", "mention", "allowlist"] = "open"
|
||||||
group_allow_from: list[str] = Field(default_factory=list)
|
group_allow_from: list[str] = Field(default_factory=list)
|
||||||
allow_room_mentions: bool = False
|
allow_room_mentions: bool = False,
|
||||||
|
streaming: bool = False
|
||||||
|
|
||||||
|
|
||||||
class MatrixChannel(BaseChannel):
|
class MatrixChannel(BaseChannel):
|
||||||
@@ -167,6 +221,8 @@ class MatrixChannel(BaseChannel):
|
|||||||
|
|
||||||
name = "matrix"
|
name = "matrix"
|
||||||
display_name = "Matrix"
|
display_name = "Matrix"
|
||||||
|
_STREAM_EDIT_INTERVAL = 2 # min seconds between edit_message_text calls
|
||||||
|
monotonic_time = time.monotonic
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def default_config(cls) -> dict[str, Any]:
|
def default_config(cls) -> dict[str, Any]:
|
||||||
@@ -192,6 +248,8 @@ class MatrixChannel(BaseChannel):
|
|||||||
)
|
)
|
||||||
self._server_upload_limit_bytes: int | None = None
|
self._server_upload_limit_bytes: int | None = None
|
||||||
self._server_upload_limit_checked = False
|
self._server_upload_limit_checked = False
|
||||||
|
self._stream_bufs: dict[str, _StreamBuf] = {}
|
||||||
|
|
||||||
|
|
||||||
async def start(self) -> None:
|
async def start(self) -> None:
|
||||||
"""Start Matrix client and begin sync loop."""
|
"""Start Matrix client and begin sync loop."""
|
||||||
@@ -297,14 +355,17 @@ class MatrixChannel(BaseChannel):
|
|||||||
room = getattr(self.client, "rooms", {}).get(room_id)
|
room = getattr(self.client, "rooms", {}).get(room_id)
|
||||||
return bool(getattr(room, "encrypted", False))
|
return bool(getattr(room, "encrypted", False))
|
||||||
|
|
||||||
async def _send_room_content(self, room_id: str, content: dict[str, Any]) -> None:
|
async def _send_room_content(self, room_id: str,
|
||||||
|
content: dict[str, Any]) -> None | RoomSendResponse | RoomSendError:
|
||||||
"""Send m.room.message with E2EE options."""
|
"""Send m.room.message with E2EE options."""
|
||||||
if not self.client:
|
if not self.client:
|
||||||
return
|
return None
|
||||||
kwargs: dict[str, Any] = {"room_id": room_id, "message_type": "m.room.message", "content": content}
|
kwargs: dict[str, Any] = {"room_id": room_id, "message_type": "m.room.message", "content": content}
|
||||||
|
|
||||||
if self.config.e2ee_enabled:
|
if self.config.e2ee_enabled:
|
||||||
kwargs["ignore_unverified_devices"] = True
|
kwargs["ignore_unverified_devices"] = True
|
||||||
await self.client.room_send(**kwargs)
|
response = await self.client.room_send(**kwargs)
|
||||||
|
return response
|
||||||
|
|
||||||
async def _resolve_server_upload_limit_bytes(self) -> int | None:
|
async def _resolve_server_upload_limit_bytes(self) -> int | None:
|
||||||
"""Query homeserver upload limit once per channel lifecycle."""
|
"""Query homeserver upload limit once per channel lifecycle."""
|
||||||
@@ -414,6 +475,53 @@ class MatrixChannel(BaseChannel):
|
|||||||
if not is_progress:
|
if not is_progress:
|
||||||
await self._stop_typing_keepalive(msg.chat_id, clear_typing=True)
|
await self._stop_typing_keepalive(msg.chat_id, clear_typing=True)
|
||||||
|
|
||||||
|
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||||
|
meta = metadata or {}
|
||||||
|
relates_to = self._build_thread_relates_to(metadata)
|
||||||
|
|
||||||
|
if meta.get("_stream_end"):
|
||||||
|
buf = self._stream_bufs.pop(chat_id, None)
|
||||||
|
if not buf or not buf.event_id or not buf.text:
|
||||||
|
return
|
||||||
|
|
||||||
|
await self._stop_typing_keepalive(chat_id, clear_typing=True)
|
||||||
|
|
||||||
|
content = _build_matrix_text_content(
|
||||||
|
buf.text,
|
||||||
|
buf.event_id,
|
||||||
|
thread_relates_to=relates_to,
|
||||||
|
)
|
||||||
|
await self._send_room_content(chat_id, content)
|
||||||
|
return
|
||||||
|
|
||||||
|
buf = self._stream_bufs.get(chat_id)
|
||||||
|
if buf is None:
|
||||||
|
buf = _StreamBuf()
|
||||||
|
self._stream_bufs[chat_id] = buf
|
||||||
|
buf.text += delta
|
||||||
|
|
||||||
|
if not buf.text.strip():
|
||||||
|
return
|
||||||
|
|
||||||
|
now = self.monotonic_time()
|
||||||
|
|
||||||
|
if not buf.last_edit or (now - buf.last_edit) >= self._STREAM_EDIT_INTERVAL:
|
||||||
|
try:
|
||||||
|
content = _build_matrix_text_content(
|
||||||
|
buf.text,
|
||||||
|
buf.event_id,
|
||||||
|
thread_relates_to=relates_to,
|
||||||
|
)
|
||||||
|
response = await self._send_room_content(chat_id, content)
|
||||||
|
buf.last_edit = now
|
||||||
|
if not buf.event_id:
|
||||||
|
# we are editing the same message all the time, so only the first time the event id needs to be set
|
||||||
|
buf.event_id = response.event_id
|
||||||
|
except Exception:
|
||||||
|
await self._stop_typing_keepalive(chat_id, clear_typing=True)
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
def _register_event_callbacks(self) -> None:
|
def _register_event_callbacks(self) -> None:
|
||||||
self.client.add_event_callback(self._on_message, RoomMessageText)
|
self.client.add_event_callback(self._on_message, RoomMessageText)
|
||||||
self.client.add_event_callback(self._on_media_message, MATRIX_MEDIA_EVENT_FILTER)
|
self.client.add_event_callback(self._on_media_message, MATRIX_MEDIA_EVENT_FILTER)
|
||||||
|
|||||||
+332
-66
@@ -15,6 +15,7 @@ import hashlib
|
|||||||
import json
|
import json
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import os
|
import os
|
||||||
|
import random
|
||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
@@ -53,7 +54,26 @@ MESSAGE_TYPE_BOT = 2
|
|||||||
MESSAGE_STATE_FINISH = 2
|
MESSAGE_STATE_FINISH = 2
|
||||||
|
|
||||||
WEIXIN_MAX_MESSAGE_LEN = 4000
|
WEIXIN_MAX_MESSAGE_LEN = 4000
|
||||||
WEIXIN_CHANNEL_VERSION = "1.0.3"
|
WEIXIN_CHANNEL_VERSION = "2.1.1"
|
||||||
|
ILINK_APP_ID = "bot"
|
||||||
|
|
||||||
|
|
||||||
|
def _build_client_version(version: str) -> int:
|
||||||
|
"""Encode semantic version as 0x00MMNNPP (major/minor/patch in one uint32)."""
|
||||||
|
parts = version.split(".")
|
||||||
|
|
||||||
|
def _as_int(idx: int) -> int:
|
||||||
|
try:
|
||||||
|
return int(parts[idx])
|
||||||
|
except Exception:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
major = _as_int(0)
|
||||||
|
minor = _as_int(1)
|
||||||
|
patch = _as_int(2)
|
||||||
|
return ((major & 0xFF) << 16) | ((minor & 0xFF) << 8) | (patch & 0xFF)
|
||||||
|
|
||||||
|
ILINK_APP_CLIENT_VERSION = _build_client_version(WEIXIN_CHANNEL_VERSION)
|
||||||
BASE_INFO: dict[str, str] = {"channel_version": WEIXIN_CHANNEL_VERSION}
|
BASE_INFO: dict[str, str] = {"channel_version": WEIXIN_CHANNEL_VERSION}
|
||||||
|
|
||||||
# Session-expired error code
|
# Session-expired error code
|
||||||
@@ -65,18 +85,32 @@ MAX_CONSECUTIVE_FAILURES = 3
|
|||||||
BACKOFF_DELAY_S = 30
|
BACKOFF_DELAY_S = 30
|
||||||
RETRY_DELAY_S = 2
|
RETRY_DELAY_S = 2
|
||||||
MAX_QR_REFRESH_COUNT = 3
|
MAX_QR_REFRESH_COUNT = 3
|
||||||
|
TYPING_STATUS_TYPING = 1
|
||||||
|
TYPING_STATUS_CANCEL = 2
|
||||||
|
TYPING_TICKET_TTL_S = 24 * 60 * 60
|
||||||
|
TYPING_KEEPALIVE_INTERVAL_S = 5
|
||||||
|
CONFIG_CACHE_INITIAL_RETRY_S = 2
|
||||||
|
CONFIG_CACHE_MAX_RETRY_S = 60 * 60
|
||||||
|
|
||||||
# Default long-poll timeout; overridden by server via longpolling_timeout_ms.
|
# Default long-poll timeout; overridden by server via longpolling_timeout_ms.
|
||||||
DEFAULT_LONG_POLL_TIMEOUT_S = 35
|
DEFAULT_LONG_POLL_TIMEOUT_S = 35
|
||||||
|
|
||||||
# Media-type codes for getuploadurl (1=image, 2=video, 3=file)
|
# Media-type codes for getuploadurl (1=image, 2=video, 3=file, 4=voice)
|
||||||
UPLOAD_MEDIA_IMAGE = 1
|
UPLOAD_MEDIA_IMAGE = 1
|
||||||
UPLOAD_MEDIA_VIDEO = 2
|
UPLOAD_MEDIA_VIDEO = 2
|
||||||
UPLOAD_MEDIA_FILE = 3
|
UPLOAD_MEDIA_FILE = 3
|
||||||
|
UPLOAD_MEDIA_VOICE = 4
|
||||||
|
|
||||||
# File extensions considered as images / videos for outbound media
|
# File extensions considered as images / videos for outbound media
|
||||||
_IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".webp", ".tiff", ".ico", ".svg"}
|
_IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".webp", ".tiff", ".ico", ".svg"}
|
||||||
_VIDEO_EXTS = {".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv"}
|
_VIDEO_EXTS = {".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv"}
|
||||||
|
_VOICE_EXTS = {".mp3", ".wav", ".amr", ".silk", ".ogg", ".m4a", ".aac", ".flac"}
|
||||||
|
|
||||||
|
|
||||||
|
def _has_downloadable_media_locator(media: dict[str, Any] | None) -> bool:
|
||||||
|
if not isinstance(media, dict):
|
||||||
|
return False
|
||||||
|
return bool(str(media.get("encrypt_query_param", "") or "") or str(media.get("full_url", "") or "").strip())
|
||||||
|
|
||||||
|
|
||||||
class WeixinConfig(Base):
|
class WeixinConfig(Base):
|
||||||
@@ -124,6 +158,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
self._poll_task: asyncio.Task | None = None
|
self._poll_task: asyncio.Task | None = None
|
||||||
self._next_poll_timeout_s: int = DEFAULT_LONG_POLL_TIMEOUT_S
|
self._next_poll_timeout_s: int = DEFAULT_LONG_POLL_TIMEOUT_S
|
||||||
self._session_pause_until: float = 0.0
|
self._session_pause_until: float = 0.0
|
||||||
|
self._typing_tickets: dict[str, dict[str, Any]] = {}
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# State persistence
|
# State persistence
|
||||||
@@ -162,8 +197,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
if base_url:
|
if base_url:
|
||||||
self.config.base_url = base_url
|
self.config.base_url = base_url
|
||||||
return bool(self._token)
|
return bool(self._token)
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.warning("Failed to load WeChat state: {}", e)
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def _save_state(self) -> None:
|
def _save_state(self) -> None:
|
||||||
@@ -176,8 +210,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
"base_url": self.config.base_url,
|
"base_url": self.config.base_url,
|
||||||
}
|
}
|
||||||
state_file.write_text(json.dumps(data, ensure_ascii=False))
|
state_file.write_text(json.dumps(data, ensure_ascii=False))
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.warning("Failed to save WeChat state: {}", e)
|
pass
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# HTTP helpers (matches api.ts buildHeaders / apiFetch)
|
# HTTP helpers (matches api.ts buildHeaders / apiFetch)
|
||||||
@@ -199,6 +233,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
"X-WECHAT-UIN": self._random_wechat_uin(),
|
"X-WECHAT-UIN": self._random_wechat_uin(),
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
"AuthorizationType": "ilink_bot_token",
|
"AuthorizationType": "ilink_bot_token",
|
||||||
|
"iLink-App-Id": ILINK_APP_ID,
|
||||||
|
"iLink-App-ClientVersion": str(ILINK_APP_CLIENT_VERSION),
|
||||||
}
|
}
|
||||||
if auth and self._token:
|
if auth and self._token:
|
||||||
headers["Authorization"] = f"Bearer {self._token}"
|
headers["Authorization"] = f"Bearer {self._token}"
|
||||||
@@ -206,6 +242,15 @@ class WeixinChannel(BaseChannel):
|
|||||||
headers["SKRouteTag"] = str(self.config.route_tag).strip()
|
headers["SKRouteTag"] = str(self.config.route_tag).strip()
|
||||||
return headers
|
return headers
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_retryable_media_download_error(err: Exception) -> bool:
|
||||||
|
if isinstance(err, httpx.TimeoutException | httpx.TransportError):
|
||||||
|
return True
|
||||||
|
if isinstance(err, httpx.HTTPStatusError):
|
||||||
|
status_code = err.response.status_code if err.response is not None else 0
|
||||||
|
return status_code >= 500
|
||||||
|
return False
|
||||||
|
|
||||||
async def _api_get(
|
async def _api_get(
|
||||||
self,
|
self,
|
||||||
endpoint: str,
|
endpoint: str,
|
||||||
@@ -223,6 +268,25 @@ class WeixinChannel(BaseChannel):
|
|||||||
resp.raise_for_status()
|
resp.raise_for_status()
|
||||||
return resp.json()
|
return resp.json()
|
||||||
|
|
||||||
|
async def _api_get_with_base(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
base_url: str,
|
||||||
|
endpoint: str,
|
||||||
|
params: dict | None = None,
|
||||||
|
auth: bool = True,
|
||||||
|
extra_headers: dict[str, str] | None = None,
|
||||||
|
) -> dict:
|
||||||
|
"""GET helper that allows overriding base_url for QR redirect polling."""
|
||||||
|
assert self._client is not None
|
||||||
|
url = f"{base_url.rstrip('/')}/{endpoint}"
|
||||||
|
hdrs = self._make_headers(auth=auth)
|
||||||
|
if extra_headers:
|
||||||
|
hdrs.update(extra_headers)
|
||||||
|
resp = await self._client.get(url, params=params, headers=hdrs)
|
||||||
|
resp.raise_for_status()
|
||||||
|
return resp.json()
|
||||||
|
|
||||||
async def _api_post(
|
async def _api_post(
|
||||||
self,
|
self,
|
||||||
endpoint: str,
|
endpoint: str,
|
||||||
@@ -259,23 +323,27 @@ class WeixinChannel(BaseChannel):
|
|||||||
async def _qr_login(self) -> bool:
|
async def _qr_login(self) -> bool:
|
||||||
"""Perform QR code login flow. Returns True on success."""
|
"""Perform QR code login flow. Returns True on success."""
|
||||||
try:
|
try:
|
||||||
logger.info("Starting WeChat QR code login...")
|
|
||||||
refresh_count = 0
|
refresh_count = 0
|
||||||
qrcode_id, scan_url = await self._fetch_qr_code()
|
qrcode_id, scan_url = await self._fetch_qr_code()
|
||||||
self._print_qr_code(scan_url)
|
self._print_qr_code(scan_url)
|
||||||
|
current_poll_base_url = self.config.base_url
|
||||||
|
|
||||||
logger.info("Waiting for QR code scan...")
|
|
||||||
while self._running:
|
while self._running:
|
||||||
try:
|
try:
|
||||||
# Reference plugin sends iLink-App-ClientVersion header for
|
status_data = await self._api_get_with_base(
|
||||||
# QR status polling (login-qr.ts:81).
|
base_url=current_poll_base_url,
|
||||||
status_data = await self._api_get(
|
endpoint="ilink/bot/get_qrcode_status",
|
||||||
"ilink/bot/get_qrcode_status",
|
|
||||||
params={"qrcode": qrcode_id},
|
params={"qrcode": qrcode_id},
|
||||||
auth=False,
|
auth=False,
|
||||||
extra_headers={"iLink-App-ClientVersion": "1"},
|
|
||||||
)
|
)
|
||||||
except httpx.TimeoutException:
|
except Exception as e:
|
||||||
|
if self._is_retryable_qr_poll_error(e):
|
||||||
|
await asyncio.sleep(1)
|
||||||
|
continue
|
||||||
|
raise
|
||||||
|
|
||||||
|
if not isinstance(status_data, dict):
|
||||||
|
await asyncio.sleep(1)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
status = status_data.get("status", "")
|
status = status_data.get("status", "")
|
||||||
@@ -298,8 +366,15 @@ class WeixinChannel(BaseChannel):
|
|||||||
else:
|
else:
|
||||||
logger.error("Login confirmed but no bot_token in response")
|
logger.error("Login confirmed but no bot_token in response")
|
||||||
return False
|
return False
|
||||||
elif status == "scaned":
|
elif status == "scaned_but_redirect":
|
||||||
logger.info("QR code scanned, waiting for confirmation...")
|
redirect_host = str(status_data.get("redirect_host", "") or "").strip()
|
||||||
|
if redirect_host:
|
||||||
|
if redirect_host.startswith("http://") or redirect_host.startswith("https://"):
|
||||||
|
redirected_base = redirect_host
|
||||||
|
else:
|
||||||
|
redirected_base = f"https://{redirect_host}"
|
||||||
|
if redirected_base != current_poll_base_url:
|
||||||
|
current_poll_base_url = redirected_base
|
||||||
elif status == "expired":
|
elif status == "expired":
|
||||||
refresh_count += 1
|
refresh_count += 1
|
||||||
if refresh_count > MAX_QR_REFRESH_COUNT:
|
if refresh_count > MAX_QR_REFRESH_COUNT:
|
||||||
@@ -309,14 +384,9 @@ class WeixinChannel(BaseChannel):
|
|||||||
MAX_QR_REFRESH_COUNT,
|
MAX_QR_REFRESH_COUNT,
|
||||||
)
|
)
|
||||||
return False
|
return False
|
||||||
logger.warning(
|
|
||||||
"QR code expired, refreshing... ({}/{})",
|
|
||||||
refresh_count,
|
|
||||||
MAX_QR_REFRESH_COUNT,
|
|
||||||
)
|
|
||||||
qrcode_id, scan_url = await self._fetch_qr_code()
|
qrcode_id, scan_url = await self._fetch_qr_code()
|
||||||
|
current_poll_base_url = self.config.base_url
|
||||||
self._print_qr_code(scan_url)
|
self._print_qr_code(scan_url)
|
||||||
logger.info("New QR code generated, waiting for scan...")
|
|
||||||
continue
|
continue
|
||||||
# status == "wait" — keep polling
|
# status == "wait" — keep polling
|
||||||
|
|
||||||
@@ -327,6 +397,16 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_retryable_qr_poll_error(err: Exception) -> bool:
|
||||||
|
if isinstance(err, httpx.TimeoutException | httpx.TransportError):
|
||||||
|
return True
|
||||||
|
if isinstance(err, httpx.HTTPStatusError):
|
||||||
|
status_code = err.response.status_code if err.response is not None else 0
|
||||||
|
if status_code >= 500:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _print_qr_code(url: str) -> None:
|
def _print_qr_code(url: str) -> None:
|
||||||
try:
|
try:
|
||||||
@@ -337,7 +417,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
qr.make(fit=True)
|
qr.make(fit=True)
|
||||||
qr.print_ascii(invert=True)
|
qr.print_ascii(invert=True)
|
||||||
except ImportError:
|
except ImportError:
|
||||||
logger.info("QR code URL (install 'qrcode' for terminal display): {}", url)
|
|
||||||
print(f"\nLogin URL: {url}\n")
|
print(f"\nLogin URL: {url}\n")
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@@ -399,12 +478,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
if not self._running:
|
if not self._running:
|
||||||
break
|
break
|
||||||
consecutive_failures += 1
|
consecutive_failures += 1
|
||||||
logger.error(
|
|
||||||
"WeChat poll error ({}/{}): {}",
|
|
||||||
consecutive_failures,
|
|
||||||
MAX_CONSECUTIVE_FAILURES,
|
|
||||||
e,
|
|
||||||
)
|
|
||||||
if consecutive_failures >= MAX_CONSECUTIVE_FAILURES:
|
if consecutive_failures >= MAX_CONSECUTIVE_FAILURES:
|
||||||
consecutive_failures = 0
|
consecutive_failures = 0
|
||||||
await asyncio.sleep(BACKOFF_DELAY_S)
|
await asyncio.sleep(BACKOFF_DELAY_S)
|
||||||
@@ -419,8 +492,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
await self._client.aclose()
|
await self._client.aclose()
|
||||||
self._client = None
|
self._client = None
|
||||||
self._save_state()
|
self._save_state()
|
||||||
logger.info("WeChat channel stopped")
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Polling (matches monitor.ts monitorWeixinProvider)
|
# Polling (matches monitor.ts monitorWeixinProvider)
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@@ -446,10 +517,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
async def _poll_once(self) -> None:
|
async def _poll_once(self) -> None:
|
||||||
remaining = self._session_pause_remaining_s()
|
remaining = self._session_pause_remaining_s()
|
||||||
if remaining > 0:
|
if remaining > 0:
|
||||||
logger.warning(
|
|
||||||
"WeChat session paused, waiting {} min before next poll.",
|
|
||||||
max((remaining + 59) // 60, 1),
|
|
||||||
)
|
|
||||||
await asyncio.sleep(remaining)
|
await asyncio.sleep(remaining)
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -499,8 +566,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
for msg in msgs:
|
for msg in msgs:
|
||||||
try:
|
try:
|
||||||
await self._process_message(msg)
|
await self._process_message(msg)
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.error("Error processing WeChat message: {}", e)
|
pass
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Inbound message processing (matches inbound.ts + process-message.ts)
|
# Inbound message processing (matches inbound.ts + process-message.ts)
|
||||||
@@ -536,6 +603,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
item_list: list[dict] = msg.get("item_list") or []
|
item_list: list[dict] = msg.get("item_list") or []
|
||||||
content_parts: list[str] = []
|
content_parts: list[str] = []
|
||||||
media_paths: list[str] = []
|
media_paths: list[str] = []
|
||||||
|
has_top_level_downloadable_media = False
|
||||||
|
|
||||||
for item in item_list:
|
for item in item_list:
|
||||||
item_type = item.get("type", 0)
|
item_type = item.get("type", 0)
|
||||||
@@ -572,6 +640,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
elif item_type == ITEM_IMAGE:
|
elif item_type == ITEM_IMAGE:
|
||||||
image_item = item.get("image_item") or {}
|
image_item = item.get("image_item") or {}
|
||||||
|
if _has_downloadable_media_locator(image_item.get("media")):
|
||||||
|
has_top_level_downloadable_media = True
|
||||||
file_path = await self._download_media_item(image_item, "image")
|
file_path = await self._download_media_item(image_item, "image")
|
||||||
if file_path:
|
if file_path:
|
||||||
content_parts.append(f"[image]\n[Image: source: {file_path}]")
|
content_parts.append(f"[image]\n[Image: source: {file_path}]")
|
||||||
@@ -586,6 +656,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
if voice_text:
|
if voice_text:
|
||||||
content_parts.append(f"[voice] {voice_text}")
|
content_parts.append(f"[voice] {voice_text}")
|
||||||
else:
|
else:
|
||||||
|
if _has_downloadable_media_locator(voice_item.get("media")):
|
||||||
|
has_top_level_downloadable_media = True
|
||||||
file_path = await self._download_media_item(voice_item, "voice")
|
file_path = await self._download_media_item(voice_item, "voice")
|
||||||
if file_path:
|
if file_path:
|
||||||
transcription = await self.transcribe_audio(file_path)
|
transcription = await self.transcribe_audio(file_path)
|
||||||
@@ -599,6 +671,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
elif item_type == ITEM_FILE:
|
elif item_type == ITEM_FILE:
|
||||||
file_item = item.get("file_item") or {}
|
file_item = item.get("file_item") or {}
|
||||||
|
if _has_downloadable_media_locator(file_item.get("media")):
|
||||||
|
has_top_level_downloadable_media = True
|
||||||
file_name = file_item.get("file_name", "unknown")
|
file_name = file_item.get("file_name", "unknown")
|
||||||
file_path = await self._download_media_item(
|
file_path = await self._download_media_item(
|
||||||
file_item,
|
file_item,
|
||||||
@@ -613,6 +687,8 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
elif item_type == ITEM_VIDEO:
|
elif item_type == ITEM_VIDEO:
|
||||||
video_item = item.get("video_item") or {}
|
video_item = item.get("video_item") or {}
|
||||||
|
if _has_downloadable_media_locator(video_item.get("media")):
|
||||||
|
has_top_level_downloadable_media = True
|
||||||
file_path = await self._download_media_item(video_item, "video")
|
file_path = await self._download_media_item(video_item, "video")
|
||||||
if file_path:
|
if file_path:
|
||||||
content_parts.append(f"[video]\n[Video: source: {file_path}]")
|
content_parts.append(f"[video]\n[Video: source: {file_path}]")
|
||||||
@@ -620,17 +696,56 @@ class WeixinChannel(BaseChannel):
|
|||||||
else:
|
else:
|
||||||
content_parts.append("[video]")
|
content_parts.append("[video]")
|
||||||
|
|
||||||
|
# Fallback: when no top-level media was downloaded, try quoted/referenced media.
|
||||||
|
# This aligns with the reference plugin behavior that checks ref_msg.message_item
|
||||||
|
# when main item_list has no downloadable media.
|
||||||
|
if not media_paths and not has_top_level_downloadable_media:
|
||||||
|
ref_media_item: dict[str, Any] | None = None
|
||||||
|
for item in item_list:
|
||||||
|
if item.get("type", 0) != ITEM_TEXT:
|
||||||
|
continue
|
||||||
|
ref = item.get("ref_msg") or {}
|
||||||
|
candidate = ref.get("message_item") or {}
|
||||||
|
if candidate.get("type", 0) in (ITEM_IMAGE, ITEM_VOICE, ITEM_FILE, ITEM_VIDEO):
|
||||||
|
ref_media_item = candidate
|
||||||
|
break
|
||||||
|
|
||||||
|
if ref_media_item:
|
||||||
|
ref_type = ref_media_item.get("type", 0)
|
||||||
|
if ref_type == ITEM_IMAGE:
|
||||||
|
image_item = ref_media_item.get("image_item") or {}
|
||||||
|
file_path = await self._download_media_item(image_item, "image")
|
||||||
|
if file_path:
|
||||||
|
content_parts.append(f"[image]\n[Image: source: {file_path}]")
|
||||||
|
media_paths.append(file_path)
|
||||||
|
elif ref_type == ITEM_VOICE:
|
||||||
|
voice_item = ref_media_item.get("voice_item") or {}
|
||||||
|
file_path = await self._download_media_item(voice_item, "voice")
|
||||||
|
if file_path:
|
||||||
|
transcription = await self.transcribe_audio(file_path)
|
||||||
|
if transcription:
|
||||||
|
content_parts.append(f"[voice] {transcription}")
|
||||||
|
else:
|
||||||
|
content_parts.append(f"[voice]\n[Audio: source: {file_path}]")
|
||||||
|
media_paths.append(file_path)
|
||||||
|
elif ref_type == ITEM_FILE:
|
||||||
|
file_item = ref_media_item.get("file_item") or {}
|
||||||
|
file_name = file_item.get("file_name", "unknown")
|
||||||
|
file_path = await self._download_media_item(file_item, "file", file_name)
|
||||||
|
if file_path:
|
||||||
|
content_parts.append(f"[file: {file_name}]\n[File: source: {file_path}]")
|
||||||
|
media_paths.append(file_path)
|
||||||
|
elif ref_type == ITEM_VIDEO:
|
||||||
|
video_item = ref_media_item.get("video_item") or {}
|
||||||
|
file_path = await self._download_media_item(video_item, "video")
|
||||||
|
if file_path:
|
||||||
|
content_parts.append(f"[video]\n[Video: source: {file_path}]")
|
||||||
|
media_paths.append(file_path)
|
||||||
|
|
||||||
content = "\n".join(content_parts)
|
content = "\n".join(content_parts)
|
||||||
if not content:
|
if not content:
|
||||||
return
|
return
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"WeChat inbound: from={} items={} bodyLen={}",
|
|
||||||
from_user_id,
|
|
||||||
",".join(str(i.get("type", 0)) for i in item_list),
|
|
||||||
len(content),
|
|
||||||
)
|
|
||||||
|
|
||||||
await self._handle_message(
|
await self._handle_message(
|
||||||
sender_id=from_user_id,
|
sender_id=from_user_id,
|
||||||
chat_id=from_user_id,
|
chat_id=from_user_id,
|
||||||
@@ -652,9 +767,10 @@ class WeixinChannel(BaseChannel):
|
|||||||
"""Download + AES-decrypt a media item. Returns local path or None."""
|
"""Download + AES-decrypt a media item. Returns local path or None."""
|
||||||
try:
|
try:
|
||||||
media = typed_item.get("media") or {}
|
media = typed_item.get("media") or {}
|
||||||
encrypt_query_param = media.get("encrypt_query_param", "")
|
encrypt_query_param = str(media.get("encrypt_query_param", "") or "")
|
||||||
|
full_url = str(media.get("full_url", "") or "").strip()
|
||||||
|
|
||||||
if not encrypt_query_param:
|
if not encrypt_query_param and not full_url:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Resolve AES key (media-download.ts:43-45, pic-decrypt.ts:40-52)
|
# Resolve AES key (media-download.ts:43-45, pic-decrypt.ts:40-52)
|
||||||
@@ -671,21 +787,50 @@ class WeixinChannel(BaseChannel):
|
|||||||
elif media_aes_key_b64:
|
elif media_aes_key_b64:
|
||||||
aes_key_b64 = media_aes_key_b64
|
aes_key_b64 = media_aes_key_b64
|
||||||
|
|
||||||
# Build CDN download URL with proper URL-encoding (cdn-url.ts:7)
|
# Reference protocol behavior: VOICE/FILE/VIDEO require aes_key;
|
||||||
cdn_url = (
|
# only IMAGE may be downloaded as plain bytes when key is missing.
|
||||||
|
if media_type != "image" and not aes_key_b64:
|
||||||
|
return None
|
||||||
|
|
||||||
|
assert self._client is not None
|
||||||
|
fallback_url = ""
|
||||||
|
if encrypt_query_param:
|
||||||
|
fallback_url = (
|
||||||
f"{self.config.cdn_base_url}/download"
|
f"{self.config.cdn_base_url}/download"
|
||||||
f"?encrypted_query_param={quote(encrypt_query_param)}"
|
f"?encrypted_query_param={quote(encrypt_query_param)}"
|
||||||
)
|
)
|
||||||
|
|
||||||
assert self._client is not None
|
download_candidates: list[tuple[str, str]] = []
|
||||||
|
if full_url:
|
||||||
|
download_candidates.append(("full_url", full_url))
|
||||||
|
if fallback_url and (not full_url or fallback_url != full_url):
|
||||||
|
download_candidates.append(("encrypt_query_param", fallback_url))
|
||||||
|
|
||||||
|
data = b""
|
||||||
|
for idx, (download_source, cdn_url) in enumerate(download_candidates):
|
||||||
|
try:
|
||||||
resp = await self._client.get(cdn_url)
|
resp = await self._client.get(cdn_url)
|
||||||
resp.raise_for_status()
|
resp.raise_for_status()
|
||||||
data = resp.content
|
data = resp.content
|
||||||
|
break
|
||||||
|
except Exception as e:
|
||||||
|
has_more_candidates = idx + 1 < len(download_candidates)
|
||||||
|
should_fallback = (
|
||||||
|
download_source == "full_url"
|
||||||
|
and has_more_candidates
|
||||||
|
and self._is_retryable_media_download_error(e)
|
||||||
|
)
|
||||||
|
if should_fallback:
|
||||||
|
logger.warning(
|
||||||
|
"WeChat media download failed via full_url, falling back to encrypt_query_param: type={} err={}",
|
||||||
|
media_type,
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
raise
|
||||||
|
|
||||||
if aes_key_b64 and data:
|
if aes_key_b64 and data:
|
||||||
data = _decrypt_aes_ecb(data, aes_key_b64)
|
data = _decrypt_aes_ecb(data, aes_key_b64)
|
||||||
elif not aes_key_b64:
|
|
||||||
logger.debug("No AES key for {} item, using raw bytes", media_type)
|
|
||||||
|
|
||||||
if not data:
|
if not data:
|
||||||
return None
|
return None
|
||||||
@@ -694,12 +839,12 @@ class WeixinChannel(BaseChannel):
|
|||||||
ext = _ext_for_type(media_type)
|
ext = _ext_for_type(media_type)
|
||||||
if not filename:
|
if not filename:
|
||||||
ts = int(time.time())
|
ts = int(time.time())
|
||||||
h = abs(hash(encrypt_query_param)) % 100000
|
hash_seed = encrypt_query_param or full_url
|
||||||
|
h = abs(hash(hash_seed)) % 100000
|
||||||
filename = f"{media_type}_{ts}_{h}{ext}"
|
filename = f"{media_type}_{ts}_{h}{ext}"
|
||||||
safe_name = os.path.basename(filename)
|
safe_name = os.path.basename(filename)
|
||||||
file_path = media_dir / safe_name
|
file_path = media_dir / safe_name
|
||||||
file_path.write_bytes(data)
|
file_path.write_bytes(data)
|
||||||
logger.debug("Downloaded WeChat {} to {}", media_type, file_path)
|
|
||||||
return str(file_path)
|
return str(file_path)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -710,14 +855,76 @@ class WeixinChannel(BaseChannel):
|
|||||||
# Outbound (matches send.ts buildTextMessageReq + sendMessageWeixin)
|
# Outbound (matches send.ts buildTextMessageReq + sendMessageWeixin)
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _get_typing_ticket(self, user_id: str, context_token: str = "") -> str:
|
||||||
|
"""Get typing ticket with per-user refresh + failure backoff cache."""
|
||||||
|
now = time.time()
|
||||||
|
entry = self._typing_tickets.get(user_id)
|
||||||
|
if entry and now < float(entry.get("next_fetch_at", 0)):
|
||||||
|
return str(entry.get("ticket", "") or "")
|
||||||
|
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"ilink_user_id": user_id,
|
||||||
|
"context_token": context_token or None,
|
||||||
|
"base_info": BASE_INFO,
|
||||||
|
}
|
||||||
|
data = await self._api_post("ilink/bot/getconfig", body)
|
||||||
|
if data.get("ret", 0) == 0:
|
||||||
|
ticket = str(data.get("typing_ticket", "") or "")
|
||||||
|
self._typing_tickets[user_id] = {
|
||||||
|
"ticket": ticket,
|
||||||
|
"ever_succeeded": True,
|
||||||
|
"next_fetch_at": now + (random.random() * TYPING_TICKET_TTL_S),
|
||||||
|
"retry_delay_s": CONFIG_CACHE_INITIAL_RETRY_S,
|
||||||
|
}
|
||||||
|
return ticket
|
||||||
|
|
||||||
|
prev_delay = float(entry.get("retry_delay_s", CONFIG_CACHE_INITIAL_RETRY_S)) if entry else CONFIG_CACHE_INITIAL_RETRY_S
|
||||||
|
next_delay = min(prev_delay * 2, CONFIG_CACHE_MAX_RETRY_S)
|
||||||
|
if entry:
|
||||||
|
entry["next_fetch_at"] = now + next_delay
|
||||||
|
entry["retry_delay_s"] = next_delay
|
||||||
|
return str(entry.get("ticket", "") or "")
|
||||||
|
|
||||||
|
self._typing_tickets[user_id] = {
|
||||||
|
"ticket": "",
|
||||||
|
"ever_succeeded": False,
|
||||||
|
"next_fetch_at": now + CONFIG_CACHE_INITIAL_RETRY_S,
|
||||||
|
"retry_delay_s": CONFIG_CACHE_INITIAL_RETRY_S,
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
|
||||||
|
async def _send_typing(self, user_id: str, typing_ticket: str, status: int) -> None:
|
||||||
|
"""Best-effort sendtyping wrapper."""
|
||||||
|
if not typing_ticket:
|
||||||
|
return
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"ilink_user_id": user_id,
|
||||||
|
"typing_ticket": typing_ticket,
|
||||||
|
"status": status,
|
||||||
|
"base_info": BASE_INFO,
|
||||||
|
}
|
||||||
|
await self._api_post("ilink/bot/sendtyping", body)
|
||||||
|
|
||||||
|
async def _typing_keepalive_loop(self, user_id: str, typing_ticket: str, stop_event: asyncio.Event) -> None:
|
||||||
|
try:
|
||||||
|
while not stop_event.is_set():
|
||||||
|
await asyncio.sleep(TYPING_KEEPALIVE_INTERVAL_S)
|
||||||
|
if stop_event.is_set():
|
||||||
|
break
|
||||||
|
try:
|
||||||
|
await self._send_typing(user_id, typing_ticket, TYPING_STATUS_TYPING)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
pass
|
||||||
|
|
||||||
async def send(self, msg: OutboundMessage) -> None:
|
async def send(self, msg: OutboundMessage) -> None:
|
||||||
if not self._client or not self._token:
|
if not self._client or not self._token:
|
||||||
logger.warning("WeChat client not initialized or not authenticated")
|
logger.warning("WeChat client not initialized or not authenticated")
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
self._assert_session_active()
|
self._assert_session_active()
|
||||||
except RuntimeError as e:
|
except RuntimeError:
|
||||||
logger.warning("WeChat send blocked: {}", e)
|
|
||||||
return
|
return
|
||||||
|
|
||||||
content = msg.content.strip()
|
content = msg.content.strip()
|
||||||
@@ -729,6 +936,26 @@ class WeixinChannel(BaseChannel):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
typing_ticket = ""
|
||||||
|
try:
|
||||||
|
typing_ticket = await self._get_typing_ticket(msg.chat_id, ctx_token)
|
||||||
|
except Exception:
|
||||||
|
typing_ticket = ""
|
||||||
|
|
||||||
|
if typing_ticket:
|
||||||
|
try:
|
||||||
|
await self._send_typing(msg.chat_id, typing_ticket, TYPING_STATUS_TYPING)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
typing_keepalive_stop = asyncio.Event()
|
||||||
|
typing_keepalive_task: asyncio.Task | None = None
|
||||||
|
if typing_ticket:
|
||||||
|
typing_keepalive_task = asyncio.create_task(
|
||||||
|
self._typing_keepalive_loop(msg.chat_id, typing_ticket, typing_keepalive_stop)
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
# --- Send media files first (following Telegram channel pattern) ---
|
# --- Send media files first (following Telegram channel pattern) ---
|
||||||
for media_path in (msg.media or []):
|
for media_path in (msg.media or []):
|
||||||
try:
|
try:
|
||||||
@@ -745,13 +972,26 @@ class WeixinChannel(BaseChannel):
|
|||||||
if not content:
|
if not content:
|
||||||
return
|
return
|
||||||
|
|
||||||
try:
|
|
||||||
chunks = split_message(content, WEIXIN_MAX_MESSAGE_LEN)
|
chunks = split_message(content, WEIXIN_MAX_MESSAGE_LEN)
|
||||||
for chunk in chunks:
|
for chunk in chunks:
|
||||||
await self._send_text(msg.chat_id, chunk, ctx_token)
|
await self._send_text(msg.chat_id, chunk, ctx_token)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error sending WeChat message: {}", e)
|
logger.error("Error sending WeChat message: {}", e)
|
||||||
raise
|
raise
|
||||||
|
finally:
|
||||||
|
if typing_keepalive_task:
|
||||||
|
typing_keepalive_stop.set()
|
||||||
|
typing_keepalive_task.cancel()
|
||||||
|
try:
|
||||||
|
await typing_keepalive_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if typing_ticket:
|
||||||
|
try:
|
||||||
|
await self._send_typing(msg.chat_id, typing_ticket, TYPING_STATUS_CANCEL)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
async def _send_text(
|
async def _send_text(
|
||||||
self,
|
self,
|
||||||
@@ -825,6 +1065,10 @@ class WeixinChannel(BaseChannel):
|
|||||||
upload_type = UPLOAD_MEDIA_VIDEO
|
upload_type = UPLOAD_MEDIA_VIDEO
|
||||||
item_type = ITEM_VIDEO
|
item_type = ITEM_VIDEO
|
||||||
item_key = "video_item"
|
item_key = "video_item"
|
||||||
|
elif ext in _VOICE_EXTS:
|
||||||
|
upload_type = UPLOAD_MEDIA_VOICE
|
||||||
|
item_type = ITEM_VOICE
|
||||||
|
item_key = "voice_item"
|
||||||
else:
|
else:
|
||||||
upload_type = UPLOAD_MEDIA_FILE
|
upload_type = UPLOAD_MEDIA_FILE
|
||||||
item_type = ITEM_FILE
|
item_type = ITEM_FILE
|
||||||
@@ -838,7 +1082,7 @@ class WeixinChannel(BaseChannel):
|
|||||||
# Matches aesEcbPaddedSize: Math.ceil((size + 1) / 16) * 16
|
# Matches aesEcbPaddedSize: Math.ceil((size + 1) / 16) * 16
|
||||||
padded_size = ((raw_size + 1 + 15) // 16) * 16
|
padded_size = ((raw_size + 1 + 15) // 16) * 16
|
||||||
|
|
||||||
# Step 1: Get upload URL (upload_param) from server
|
# Step 1: Get upload URL from server (prefer upload_full_url, fallback to upload_param)
|
||||||
file_key = os.urandom(16).hex()
|
file_key = os.urandom(16).hex()
|
||||||
upload_body: dict[str, Any] = {
|
upload_body: dict[str, Any] = {
|
||||||
"filekey": file_key,
|
"filekey": file_key,
|
||||||
@@ -853,22 +1097,27 @@ class WeixinChannel(BaseChannel):
|
|||||||
|
|
||||||
assert self._client is not None
|
assert self._client is not None
|
||||||
upload_resp = await self._api_post("ilink/bot/getuploadurl", upload_body)
|
upload_resp = await self._api_post("ilink/bot/getuploadurl", upload_body)
|
||||||
logger.debug("WeChat getuploadurl response: {}", upload_resp)
|
|
||||||
|
|
||||||
upload_param = upload_resp.get("upload_param", "")
|
upload_full_url = str(upload_resp.get("upload_full_url", "") or "").strip()
|
||||||
if not upload_param:
|
upload_param = str(upload_resp.get("upload_param", "") or "")
|
||||||
raise RuntimeError(f"getuploadurl returned no upload_param: {upload_resp}")
|
if not upload_full_url and not upload_param:
|
||||||
|
raise RuntimeError(
|
||||||
|
"getuploadurl returned no upload URL "
|
||||||
|
f"(need upload_full_url or upload_param): {upload_resp}"
|
||||||
|
)
|
||||||
|
|
||||||
# Step 2: AES-128-ECB encrypt and POST to CDN
|
# Step 2: AES-128-ECB encrypt and POST to CDN
|
||||||
aes_key_b64 = base64.b64encode(aes_key_raw).decode()
|
aes_key_b64 = base64.b64encode(aes_key_raw).decode()
|
||||||
encrypted_data = _encrypt_aes_ecb(raw_data, aes_key_b64)
|
encrypted_data = _encrypt_aes_ecb(raw_data, aes_key_b64)
|
||||||
|
|
||||||
|
if upload_full_url:
|
||||||
|
cdn_upload_url = upload_full_url
|
||||||
|
else:
|
||||||
cdn_upload_url = (
|
cdn_upload_url = (
|
||||||
f"{self.config.cdn_base_url}/upload"
|
f"{self.config.cdn_base_url}/upload"
|
||||||
f"?encrypted_query_param={quote(upload_param)}"
|
f"?encrypted_query_param={quote(upload_param)}"
|
||||||
f"&filekey={quote(file_key)}"
|
f"&filekey={quote(file_key)}"
|
||||||
)
|
)
|
||||||
logger.debug("WeChat CDN POST url={} ciphertextSize={}", cdn_upload_url[:80], len(encrypted_data))
|
|
||||||
|
|
||||||
cdn_resp = await self._client.post(
|
cdn_resp = await self._client.post(
|
||||||
cdn_upload_url,
|
cdn_upload_url,
|
||||||
@@ -884,7 +1133,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
"CDN upload response missing x-encrypted-param header; "
|
"CDN upload response missing x-encrypted-param header; "
|
||||||
f"status={cdn_resp.status_code} headers={dict(cdn_resp.headers)}"
|
f"status={cdn_resp.status_code} headers={dict(cdn_resp.headers)}"
|
||||||
)
|
)
|
||||||
logger.debug("WeChat CDN upload success for {}, got download_param", p.name)
|
|
||||||
|
|
||||||
# Step 3: Send message with the media item
|
# Step 3: Send message with the media item
|
||||||
# aes_key for CDNMedia is the hex key encoded as base64
|
# aes_key for CDNMedia is the hex key encoded as base64
|
||||||
@@ -933,7 +1181,6 @@ class WeixinChannel(BaseChannel):
|
|||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"WeChat send media error (code {errcode}): {data.get('errmsg', '')}"
|
f"WeChat send media error (code {errcode}): {data.get('errmsg', '')}"
|
||||||
)
|
)
|
||||||
logger.info("WeChat media sent: {} (type={})", p.name, item_key)
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -1005,24 +1252,43 @@ def _decrypt_aes_ecb(data: bytes, aes_key_b64: str) -> bytes:
|
|||||||
logger.warning("Failed to parse AES key, returning raw data: {}", e)
|
logger.warning("Failed to parse AES key, returning raw data: {}", e)
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
decrypted: bytes | None = None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from Crypto.Cipher import AES
|
from Crypto.Cipher import AES
|
||||||
|
|
||||||
cipher = AES.new(key, AES.MODE_ECB)
|
cipher = AES.new(key, AES.MODE_ECB)
|
||||||
return cipher.decrypt(data) # pycryptodome auto-strips PKCS7 with unpad
|
decrypted = cipher.decrypt(data)
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
if decrypted is None:
|
||||||
try:
|
try:
|
||||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||||
|
|
||||||
cipher_obj = Cipher(algorithms.AES(key), modes.ECB())
|
cipher_obj = Cipher(algorithms.AES(key), modes.ECB())
|
||||||
decryptor = cipher_obj.decryptor()
|
decryptor = cipher_obj.decryptor()
|
||||||
return decryptor.update(data) + decryptor.finalize()
|
decrypted = decryptor.update(data) + decryptor.finalize()
|
||||||
except ImportError:
|
except ImportError:
|
||||||
logger.warning("Cannot decrypt media: install 'pycryptodome' or 'cryptography'")
|
logger.warning("Cannot decrypt media: install 'pycryptodome' or 'cryptography'")
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
return _pkcs7_unpad_safe(decrypted)
|
||||||
|
|
||||||
|
|
||||||
|
def _pkcs7_unpad_safe(data: bytes, block_size: int = 16) -> bytes:
|
||||||
|
"""Safely remove PKCS7 padding when valid; otherwise return original bytes."""
|
||||||
|
if not data:
|
||||||
|
return data
|
||||||
|
if len(data) % block_size != 0:
|
||||||
|
return data
|
||||||
|
pad_len = data[-1]
|
||||||
|
if pad_len < 1 or pad_len > block_size:
|
||||||
|
return data
|
||||||
|
if data[-pad_len:] != bytes([pad_len]) * pad_len:
|
||||||
|
return data
|
||||||
|
return data[:-pad_len]
|
||||||
|
|
||||||
|
|
||||||
def _ext_for_type(media_type: str) -> str:
|
def _ext_for_type(media_type: str) -> str:
|
||||||
return {
|
return {
|
||||||
|
|||||||
@@ -84,6 +84,16 @@ async def cmd_new(ctx: CommandContext) -> OutboundMessage:
|
|||||||
|
|
||||||
async def cmd_help(ctx: CommandContext) -> OutboundMessage:
|
async def cmd_help(ctx: CommandContext) -> OutboundMessage:
|
||||||
"""Return available slash commands."""
|
"""Return available slash commands."""
|
||||||
|
return OutboundMessage(
|
||||||
|
channel=ctx.msg.channel,
|
||||||
|
chat_id=ctx.msg.chat_id,
|
||||||
|
content=build_help_text(),
|
||||||
|
metadata={"render_as": "text"},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_help_text() -> str:
|
||||||
|
"""Build canonical help text shared across channels."""
|
||||||
lines = [
|
lines = [
|
||||||
"🐈 nanobot commands:",
|
"🐈 nanobot commands:",
|
||||||
"/new — Start a new conversation",
|
"/new — Start a new conversation",
|
||||||
@@ -92,12 +102,7 @@ async def cmd_help(ctx: CommandContext) -> OutboundMessage:
|
|||||||
"/status — Show bot status",
|
"/status — Show bot status",
|
||||||
"/help — Show available commands",
|
"/help — Show available commands",
|
||||||
]
|
]
|
||||||
return OutboundMessage(
|
return "\n".join(lines)
|
||||||
channel=ctx.msg.channel,
|
|
||||||
chat_id=ctx.msg.chat_id,
|
|
||||||
content="\n".join(lines),
|
|
||||||
metadata={"render_as": "text"},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def register_builtin_commands(router: CommandRouter) -> None:
|
def register_builtin_commands(router: CommandRouter) -> None:
|
||||||
|
|||||||
@@ -136,6 +136,7 @@ class ExecToolConfig(Base):
|
|||||||
enable: bool = True
|
enable: bool = True
|
||||||
timeout: int = 60
|
timeout: int = 60
|
||||||
path_append: str = ""
|
path_append: str = ""
|
||||||
|
command_wrapper: str = "" # sandbox wrapper command template; supports {command} and {cwd}
|
||||||
|
|
||||||
class MCPServerConfig(Base):
|
class MCPServerConfig(Base):
|
||||||
"""MCP server connection configuration (stdio or HTTP)."""
|
"""MCP server connection configuration (stdio or HTTP)."""
|
||||||
|
|||||||
@@ -67,6 +67,9 @@ matrix = [
|
|||||||
"mistune>=3.0.0,<4.0.0",
|
"mistune>=3.0.0,<4.0.0",
|
||||||
"nh3>=0.2.17,<1.0.0",
|
"nh3>=0.2.17,<1.0.0",
|
||||||
]
|
]
|
||||||
|
discord = [
|
||||||
|
"discord.py>=2.5.2,<3.0.0",
|
||||||
|
]
|
||||||
langsmith = [
|
langsmith = [
|
||||||
"langsmith>=0.1.0",
|
"langsmith>=0.1.0",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -117,6 +117,43 @@ class TestDispatch:
|
|||||||
out = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
out = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||||
assert out.content == "hi"
|
assert out.content == "hi"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_dispatch_streaming_preserves_message_metadata(self):
|
||||||
|
from nanobot.bus.events import InboundMessage
|
||||||
|
|
||||||
|
loop, bus = _make_loop()
|
||||||
|
msg = InboundMessage(
|
||||||
|
channel="matrix",
|
||||||
|
sender_id="u1",
|
||||||
|
chat_id="!room:matrix.org",
|
||||||
|
content="hello",
|
||||||
|
metadata={
|
||||||
|
"_wants_stream": True,
|
||||||
|
"thread_root_event_id": "$root1",
|
||||||
|
"thread_reply_to_event_id": "$reply1",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def fake_process(_msg, *, on_stream=None, on_stream_end=None, **kwargs):
|
||||||
|
assert on_stream is not None
|
||||||
|
assert on_stream_end is not None
|
||||||
|
await on_stream("hi")
|
||||||
|
await on_stream_end(resuming=False)
|
||||||
|
return None
|
||||||
|
|
||||||
|
loop._process_message = fake_process
|
||||||
|
|
||||||
|
await loop._dispatch(msg)
|
||||||
|
first = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||||
|
second = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||||
|
|
||||||
|
assert first.metadata["thread_root_event_id"] == "$root1"
|
||||||
|
assert first.metadata["thread_reply_to_event_id"] == "$reply1"
|
||||||
|
assert first.metadata["_stream_delta"] is True
|
||||||
|
assert second.metadata["thread_root_event_id"] == "$root1"
|
||||||
|
assert second.metadata["thread_reply_to_event_id"] == "$reply1"
|
||||||
|
assert second.metadata["_stream_end"] is True
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_processing_lock_serializes(self):
|
async def test_processing_lock_serializes(self):
|
||||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||||
|
|||||||
@@ -0,0 +1,676 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from pathlib import Path
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
discord = pytest.importorskip("discord")
|
||||||
|
|
||||||
|
from nanobot.bus.events import OutboundMessage
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.channels.discord import DiscordBotClient, DiscordChannel, DiscordConfig
|
||||||
|
from nanobot.command.builtin import build_help_text
|
||||||
|
|
||||||
|
|
||||||
|
# Minimal Discord client test double used to control startup/readiness behavior.
|
||||||
|
class _FakeDiscordClient:
|
||||||
|
instances: list["_FakeDiscordClient"] = []
|
||||||
|
start_error: Exception | None = None
|
||||||
|
|
||||||
|
def __init__(self, owner, *, intents) -> None:
|
||||||
|
self.owner = owner
|
||||||
|
self.intents = intents
|
||||||
|
self.closed = False
|
||||||
|
self.ready = True
|
||||||
|
self.channels: dict[int, object] = {}
|
||||||
|
self.user = SimpleNamespace(id=999)
|
||||||
|
self.__class__.instances.append(self)
|
||||||
|
|
||||||
|
async def start(self, token: str) -> None:
|
||||||
|
self.token = token
|
||||||
|
if self.__class__.start_error is not None:
|
||||||
|
raise self.__class__.start_error
|
||||||
|
|
||||||
|
async def close(self) -> None:
|
||||||
|
self.closed = True
|
||||||
|
|
||||||
|
def is_closed(self) -> bool:
|
||||||
|
return self.closed
|
||||||
|
|
||||||
|
def is_ready(self) -> bool:
|
||||||
|
return self.ready
|
||||||
|
|
||||||
|
def get_channel(self, channel_id: int):
|
||||||
|
return self.channels.get(channel_id)
|
||||||
|
|
||||||
|
async def send_outbound(self, msg: OutboundMessage) -> None:
|
||||||
|
channel = self.get_channel(int(msg.chat_id))
|
||||||
|
if channel is None:
|
||||||
|
return
|
||||||
|
await channel.send(content=msg.content)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeAttachment:
|
||||||
|
# Attachment double that can simulate successful or failing save() calls.
|
||||||
|
def __init__(self, attachment_id: int, filename: str, *, size: int = 1, fail: bool = False) -> None:
|
||||||
|
self.id = attachment_id
|
||||||
|
self.filename = filename
|
||||||
|
self.size = size
|
||||||
|
self._fail = fail
|
||||||
|
|
||||||
|
async def save(self, path: str | Path) -> None:
|
||||||
|
if self._fail:
|
||||||
|
raise RuntimeError("save failed")
|
||||||
|
Path(path).write_bytes(b"attachment")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakePartialMessage:
|
||||||
|
# Lightweight stand-in for Discord partial message references used in replies.
|
||||||
|
def __init__(self, message_id: int) -> None:
|
||||||
|
self.id = message_id
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeChannel:
|
||||||
|
# Channel double that records outbound payloads and typing activity.
|
||||||
|
def __init__(self, channel_id: int = 123) -> None:
|
||||||
|
self.id = channel_id
|
||||||
|
self.sent_payloads: list[dict] = []
|
||||||
|
self.trigger_typing_calls = 0
|
||||||
|
self.typing_enter_hook = None
|
||||||
|
|
||||||
|
async def send(self, **kwargs) -> None:
|
||||||
|
payload = dict(kwargs)
|
||||||
|
if "file" in payload:
|
||||||
|
payload["file_name"] = payload["file"].filename
|
||||||
|
del payload["file"]
|
||||||
|
self.sent_payloads.append(payload)
|
||||||
|
|
||||||
|
def get_partial_message(self, message_id: int) -> _FakePartialMessage:
|
||||||
|
return _FakePartialMessage(message_id)
|
||||||
|
|
||||||
|
def typing(self):
|
||||||
|
channel = self
|
||||||
|
|
||||||
|
class _TypingContext:
|
||||||
|
async def __aenter__(self):
|
||||||
|
channel.trigger_typing_calls += 1
|
||||||
|
if channel.typing_enter_hook is not None:
|
||||||
|
await channel.typing_enter_hook()
|
||||||
|
|
||||||
|
async def __aexit__(self, exc_type, exc, tb):
|
||||||
|
return False
|
||||||
|
|
||||||
|
return _TypingContext()
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeInteractionResponse:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.messages: list[dict] = []
|
||||||
|
self._done = False
|
||||||
|
|
||||||
|
async def send_message(self, content: str, *, ephemeral: bool = False) -> None:
|
||||||
|
self.messages.append({"content": content, "ephemeral": ephemeral})
|
||||||
|
self._done = True
|
||||||
|
|
||||||
|
def is_done(self) -> bool:
|
||||||
|
return self._done
|
||||||
|
|
||||||
|
|
||||||
|
def _make_interaction(
|
||||||
|
*,
|
||||||
|
user_id: int = 123,
|
||||||
|
channel_id: int | None = 456,
|
||||||
|
guild_id: int | None = None,
|
||||||
|
interaction_id: int = 999,
|
||||||
|
):
|
||||||
|
return SimpleNamespace(
|
||||||
|
user=SimpleNamespace(id=user_id),
|
||||||
|
channel_id=channel_id,
|
||||||
|
guild_id=guild_id,
|
||||||
|
id=interaction_id,
|
||||||
|
command=SimpleNamespace(qualified_name="new"),
|
||||||
|
response=_FakeInteractionResponse(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_message(
|
||||||
|
*,
|
||||||
|
author_id: int = 123,
|
||||||
|
author_bot: bool = False,
|
||||||
|
channel_id: int = 456,
|
||||||
|
message_id: int = 789,
|
||||||
|
content: str = "hello",
|
||||||
|
guild_id: int | None = None,
|
||||||
|
mentions: list[object] | None = None,
|
||||||
|
attachments: list[object] | None = None,
|
||||||
|
reply_to: int | None = None,
|
||||||
|
):
|
||||||
|
# Factory for incoming Discord message objects with optional guild/reply/attachments.
|
||||||
|
guild = SimpleNamespace(id=guild_id) if guild_id is not None else None
|
||||||
|
reference = SimpleNamespace(message_id=reply_to) if reply_to is not None else None
|
||||||
|
return SimpleNamespace(
|
||||||
|
author=SimpleNamespace(id=author_id, bot=author_bot),
|
||||||
|
channel=_FakeChannel(channel_id),
|
||||||
|
content=content,
|
||||||
|
guild=guild,
|
||||||
|
mentions=mentions or [],
|
||||||
|
attachments=attachments or [],
|
||||||
|
reference=reference,
|
||||||
|
id=message_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_returns_when_token_missing() -> None:
|
||||||
|
# If no token is configured, startup should no-op and leave channel stopped.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
|
||||||
|
await channel.start()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert channel._client is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_returns_when_discord_dependency_missing(monkeypatch) -> None:
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, token="token", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.DISCORD_AVAILABLE", False)
|
||||||
|
|
||||||
|
await channel.start()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert channel._client is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_handles_client_construction_failure(monkeypatch) -> None:
|
||||||
|
# Construction errors from the Discord client should be swallowed and keep state clean.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, token="token", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _boom(owner, *, intents):
|
||||||
|
raise RuntimeError("bad client")
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.DiscordBotClient", _boom)
|
||||||
|
|
||||||
|
await channel.start()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert channel._client is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_handles_client_start_failure(monkeypatch) -> None:
|
||||||
|
# If client.start fails, the partially created client should be closed and detached.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, token="token", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
|
||||||
|
_FakeDiscordClient.instances.clear()
|
||||||
|
_FakeDiscordClient.start_error = RuntimeError("connect failed")
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.DiscordBotClient", _FakeDiscordClient)
|
||||||
|
|
||||||
|
await channel.start()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert channel._client is None
|
||||||
|
assert _FakeDiscordClient.instances[0].intents.value == channel.config.intents
|
||||||
|
assert _FakeDiscordClient.instances[0].closed is True
|
||||||
|
|
||||||
|
_FakeDiscordClient.start_error = None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stop_is_safe_after_partial_start(monkeypatch) -> None:
|
||||||
|
# stop() should close/discard the client even when startup was only partially completed.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, token="token", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
client = _FakeDiscordClient(channel, intents=None)
|
||||||
|
channel._client = client
|
||||||
|
channel._running = True
|
||||||
|
|
||||||
|
await channel.stop()
|
||||||
|
|
||||||
|
assert channel.is_running is False
|
||||||
|
assert client.closed is True
|
||||||
|
assert channel._client is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_ignores_bot_messages() -> None:
|
||||||
|
# Incoming bot-authored messages must be ignored to prevent feedback loops.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
channel._handle_message = lambda **kwargs: handled.append(kwargs) # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await channel._on_message(_make_message(author_bot=True))
|
||||||
|
|
||||||
|
assert handled == []
|
||||||
|
|
||||||
|
# If inbound handling raises, typing should be stopped for that channel.
|
||||||
|
async def fail_handle(**kwargs) -> None:
|
||||||
|
raise RuntimeError("boom")
|
||||||
|
|
||||||
|
channel._handle_message = fail_handle # type: ignore[method-assign]
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="boom"):
|
||||||
|
await channel._on_message(_make_message(author_id=123, channel_id=456))
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_accepts_allowlisted_dm() -> None:
|
||||||
|
# Allowed direct messages should be forwarded with normalized metadata.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["123"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await channel._on_message(_make_message(author_id=123, channel_id=456, message_id=789))
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["chat_id"] == "456"
|
||||||
|
assert handled[0]["metadata"] == {"message_id": "789", "guild_id": None, "reply_to": None}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_ignores_unmentioned_guild_message() -> None:
|
||||||
|
# With mention-only group policy, guild messages without a bot mention are dropped.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, allow_from=["*"], group_policy="mention"),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
channel._bot_user_id = "999"
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await channel._on_message(_make_message(guild_id=1, content="hello everyone"))
|
||||||
|
|
||||||
|
assert handled == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_accepts_mentioned_guild_message() -> None:
|
||||||
|
# Mentioned guild messages should be accepted and preserve reply threading metadata.
|
||||||
|
channel = DiscordChannel(
|
||||||
|
DiscordConfig(enabled=True, allow_from=["*"], group_policy="mention"),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
channel._bot_user_id = "999"
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await channel._on_message(
|
||||||
|
_make_message(
|
||||||
|
guild_id=1,
|
||||||
|
content="<@999> hello",
|
||||||
|
mentions=[SimpleNamespace(id=999)],
|
||||||
|
reply_to=321,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["metadata"]["reply_to"] == "321"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_downloads_attachments(tmp_path, monkeypatch) -> None:
|
||||||
|
# Attachment downloads should be saved and referenced in forwarded content/media.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.get_media_dir", lambda _name: tmp_path)
|
||||||
|
|
||||||
|
await channel._on_message(
|
||||||
|
_make_message(
|
||||||
|
attachments=[_FakeAttachment(12, "photo.png")],
|
||||||
|
content="see file",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["media"] == [str(tmp_path / "12_photo.png")]
|
||||||
|
assert "[attachment:" in handled[0]["content"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_on_message_marks_failed_attachment_download(tmp_path, monkeypatch) -> None:
|
||||||
|
# Failed attachment downloads should emit a readable placeholder and no media path.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
monkeypatch.setattr("nanobot.channels.discord.get_media_dir", lambda _name: tmp_path)
|
||||||
|
|
||||||
|
await channel._on_message(
|
||||||
|
_make_message(
|
||||||
|
attachments=[_FakeAttachment(12, "photo.png", fail=True)],
|
||||||
|
content="",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["media"] == []
|
||||||
|
assert handled[0]["content"] == "[attachment: photo.png - download failed]"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_warns_when_client_not_ready() -> None:
|
||||||
|
# Sending without a running/ready client should be a safe no-op.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
|
||||||
|
await channel.send(OutboundMessage(channel="discord", chat_id="123", content="hello"))
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_skips_when_channel_not_cached() -> None:
|
||||||
|
# Outbound sends should be skipped when the destination channel is not resolvable.
|
||||||
|
owner = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = DiscordBotClient(owner, intents=discord.Intents.none())
|
||||||
|
fetch_calls: list[int] = []
|
||||||
|
|
||||||
|
async def fetch_channel(channel_id: int):
|
||||||
|
fetch_calls.append(channel_id)
|
||||||
|
raise RuntimeError("not found")
|
||||||
|
|
||||||
|
client.fetch_channel = fetch_channel # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await client.send_outbound(OutboundMessage(channel="discord", chat_id="123", content="hello"))
|
||||||
|
|
||||||
|
assert client.get_channel(123) is None
|
||||||
|
assert fetch_calls == [123]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_fetches_channel_when_not_cached() -> None:
|
||||||
|
owner = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = DiscordBotClient(owner, intents=discord.Intents.none())
|
||||||
|
target = _FakeChannel(channel_id=123)
|
||||||
|
|
||||||
|
async def fetch_channel(channel_id: int):
|
||||||
|
return target if channel_id == 123 else None
|
||||||
|
|
||||||
|
client.fetch_channel = fetch_channel # type: ignore[method-assign]
|
||||||
|
|
||||||
|
await client.send_outbound(OutboundMessage(channel="discord", chat_id="123", content="hello"))
|
||||||
|
|
||||||
|
assert target.sent_payloads == [{"content": "hello"}]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_slash_new_forwards_when_user_is_allowlisted() -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["123"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||||
|
interaction = _make_interaction(user_id=123, channel_id=456, interaction_id=321)
|
||||||
|
|
||||||
|
new_cmd = client.tree.get_command("new")
|
||||||
|
assert new_cmd is not None
|
||||||
|
await new_cmd.callback(interaction)
|
||||||
|
|
||||||
|
assert interaction.response.messages == [
|
||||||
|
{"content": "Processing /new...", "ephemeral": True}
|
||||||
|
]
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["content"] == "/new"
|
||||||
|
assert handled[0]["sender_id"] == "123"
|
||||||
|
assert handled[0]["chat_id"] == "456"
|
||||||
|
assert handled[0]["metadata"]["interaction_id"] == "321"
|
||||||
|
assert handled[0]["metadata"]["is_slash_command"] is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_slash_new_is_blocked_for_disallowed_user() -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["999"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||||
|
interaction = _make_interaction(user_id=123, channel_id=456)
|
||||||
|
|
||||||
|
new_cmd = client.tree.get_command("new")
|
||||||
|
assert new_cmd is not None
|
||||||
|
await new_cmd.callback(interaction)
|
||||||
|
|
||||||
|
assert interaction.response.messages == [
|
||||||
|
{"content": "You are not allowed to use this bot.", "ephemeral": True}
|
||||||
|
]
|
||||||
|
assert handled == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("slash_name", ["stop", "restart", "status"])
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_slash_commands_forward_via_handle_message(slash_name: str) -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||||
|
interaction = _make_interaction()
|
||||||
|
interaction.command.qualified_name = slash_name
|
||||||
|
|
||||||
|
cmd = client.tree.get_command(slash_name)
|
||||||
|
assert cmd is not None
|
||||||
|
await cmd.callback(interaction)
|
||||||
|
|
||||||
|
assert interaction.response.messages == [
|
||||||
|
{"content": f"Processing /{slash_name}...", "ephemeral": True}
|
||||||
|
]
|
||||||
|
assert len(handled) == 1
|
||||||
|
assert handled[0]["content"] == f"/{slash_name}"
|
||||||
|
assert handled[0]["metadata"]["is_slash_command"] is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_slash_help_returns_ephemeral_help_text() -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
handled: list[dict] = []
|
||||||
|
|
||||||
|
async def capture_handle(**kwargs) -> None:
|
||||||
|
handled.append(kwargs)
|
||||||
|
|
||||||
|
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||||
|
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||||
|
interaction = _make_interaction()
|
||||||
|
interaction.command.qualified_name = "help"
|
||||||
|
|
||||||
|
help_cmd = client.tree.get_command("help")
|
||||||
|
assert help_cmd is not None
|
||||||
|
await help_cmd.callback(interaction)
|
||||||
|
|
||||||
|
assert interaction.response.messages == [
|
||||||
|
{"content": build_help_text(), "ephemeral": True}
|
||||||
|
]
|
||||||
|
assert handled == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_client_send_outbound_chunks_text_replies_and_uploads_files(tmp_path) -> None:
|
||||||
|
# Outbound payloads should upload files, attach reply references, and chunk long text.
|
||||||
|
owner = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = DiscordBotClient(owner, intents=discord.Intents.none())
|
||||||
|
target = _FakeChannel(channel_id=123)
|
||||||
|
client.get_channel = lambda channel_id: target if channel_id == 123 else None # type: ignore[method-assign]
|
||||||
|
|
||||||
|
file_path = tmp_path / "demo.txt"
|
||||||
|
file_path.write_text("hi")
|
||||||
|
|
||||||
|
await client.send_outbound(
|
||||||
|
OutboundMessage(
|
||||||
|
channel="discord",
|
||||||
|
chat_id="123",
|
||||||
|
content="a" * 2100,
|
||||||
|
reply_to="55",
|
||||||
|
media=[str(file_path)],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(target.sent_payloads) == 3
|
||||||
|
assert target.sent_payloads[0]["file_name"] == "demo.txt"
|
||||||
|
assert target.sent_payloads[0]["reference"].id == 55
|
||||||
|
assert target.sent_payloads[1]["content"] == "a" * 2000
|
||||||
|
assert target.sent_payloads[2]["content"] == "a" * 100
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_client_send_outbound_reports_failed_attachments_when_no_text(tmp_path) -> None:
|
||||||
|
# If all attachment sends fail and no text exists, emit a failure placeholder message.
|
||||||
|
owner = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = DiscordBotClient(owner, intents=discord.Intents.none())
|
||||||
|
target = _FakeChannel(channel_id=123)
|
||||||
|
client.get_channel = lambda channel_id: target if channel_id == 123 else None # type: ignore[method-assign]
|
||||||
|
|
||||||
|
missing_file = tmp_path / "missing.txt"
|
||||||
|
|
||||||
|
await client.send_outbound(
|
||||||
|
OutboundMessage(
|
||||||
|
channel="discord",
|
||||||
|
chat_id="123",
|
||||||
|
content="",
|
||||||
|
media=[str(missing_file)],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert target.sent_payloads == [{"content": "[attachment: missing.txt - send failed]"}]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_stops_typing_after_send() -> None:
|
||||||
|
# Active typing indicators should be cancelled/cleared after a successful send.
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
client = _FakeDiscordClient(channel, intents=None)
|
||||||
|
channel._client = client
|
||||||
|
channel._running = True
|
||||||
|
|
||||||
|
start = asyncio.Event()
|
||||||
|
release = asyncio.Event()
|
||||||
|
|
||||||
|
async def slow_typing() -> None:
|
||||||
|
start.set()
|
||||||
|
await release.wait()
|
||||||
|
|
||||||
|
typing_channel = _FakeChannel(channel_id=123)
|
||||||
|
typing_channel.typing_enter_hook = slow_typing
|
||||||
|
|
||||||
|
await channel._start_typing(typing_channel)
|
||||||
|
await start.wait()
|
||||||
|
|
||||||
|
await channel.send(OutboundMessage(channel="discord", chat_id="123", content="hello"))
|
||||||
|
release.set()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
|
|
||||||
|
# Progress messages should keep typing active until a final (non-progress) send.
|
||||||
|
start = asyncio.Event()
|
||||||
|
release = asyncio.Event()
|
||||||
|
|
||||||
|
async def slow_typing_progress() -> None:
|
||||||
|
start.set()
|
||||||
|
await release.wait()
|
||||||
|
|
||||||
|
typing_channel = _FakeChannel(channel_id=123)
|
||||||
|
typing_channel.typing_enter_hook = slow_typing_progress
|
||||||
|
|
||||||
|
await channel._start_typing(typing_channel)
|
||||||
|
await start.wait()
|
||||||
|
|
||||||
|
await channel.send(
|
||||||
|
OutboundMessage(
|
||||||
|
channel="discord",
|
||||||
|
chat_id="123",
|
||||||
|
content="progress",
|
||||||
|
metadata={"_progress": True},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "123" in channel._typing_tasks
|
||||||
|
|
||||||
|
await channel.send(OutboundMessage(channel="discord", chat_id="123", content="final"))
|
||||||
|
release.set()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_start_typing_uses_typing_context_when_trigger_typing_missing() -> None:
|
||||||
|
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||||
|
channel._running = True
|
||||||
|
|
||||||
|
entered = asyncio.Event()
|
||||||
|
release = asyncio.Event()
|
||||||
|
|
||||||
|
class _TypingCtx:
|
||||||
|
async def __aenter__(self):
|
||||||
|
entered.set()
|
||||||
|
|
||||||
|
async def __aexit__(self, exc_type, exc, tb):
|
||||||
|
return False
|
||||||
|
|
||||||
|
class _NoTriggerChannel:
|
||||||
|
def __init__(self, channel_id: int = 123) -> None:
|
||||||
|
self.id = channel_id
|
||||||
|
|
||||||
|
def typing(self):
|
||||||
|
async def _waiter():
|
||||||
|
await release.wait()
|
||||||
|
# Hold the loop so task remains active until explicitly stopped.
|
||||||
|
class _Ctx(_TypingCtx):
|
||||||
|
async def __aenter__(self):
|
||||||
|
await super().__aenter__()
|
||||||
|
await _waiter()
|
||||||
|
return _Ctx()
|
||||||
|
|
||||||
|
typing_channel = _NoTriggerChannel(channel_id=123)
|
||||||
|
await channel._start_typing(typing_channel) # type: ignore[arg-type]
|
||||||
|
await entered.wait()
|
||||||
|
|
||||||
|
assert "123" in channel._typing_tasks
|
||||||
|
|
||||||
|
await channel._stop_typing("123")
|
||||||
|
release.set()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
assert channel._typing_tasks == {}
|
||||||
@@ -3,6 +3,9 @@ from pathlib import Path
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from nio import RoomSendResponse
|
||||||
|
|
||||||
|
from nanobot.channels.matrix import _build_matrix_text_content
|
||||||
|
|
||||||
# Check optional matrix dependencies before importing
|
# Check optional matrix dependencies before importing
|
||||||
try:
|
try:
|
||||||
@@ -65,6 +68,7 @@ class _FakeAsyncClient:
|
|||||||
self.raise_on_send = False
|
self.raise_on_send = False
|
||||||
self.raise_on_typing = False
|
self.raise_on_typing = False
|
||||||
self.raise_on_upload = False
|
self.raise_on_upload = False
|
||||||
|
self.room_send_response: RoomSendResponse | None = RoomSendResponse(event_id="", room_id="")
|
||||||
|
|
||||||
def add_event_callback(self, callback, event_type) -> None:
|
def add_event_callback(self, callback, event_type) -> None:
|
||||||
self.callbacks.append((callback, event_type))
|
self.callbacks.append((callback, event_type))
|
||||||
@@ -87,7 +91,7 @@ class _FakeAsyncClient:
|
|||||||
message_type: str,
|
message_type: str,
|
||||||
content: dict[str, object],
|
content: dict[str, object],
|
||||||
ignore_unverified_devices: object = _ROOM_SEND_UNSET,
|
ignore_unverified_devices: object = _ROOM_SEND_UNSET,
|
||||||
) -> None:
|
) -> RoomSendResponse:
|
||||||
call: dict[str, object] = {
|
call: dict[str, object] = {
|
||||||
"room_id": room_id,
|
"room_id": room_id,
|
||||||
"message_type": message_type,
|
"message_type": message_type,
|
||||||
@@ -98,6 +102,7 @@ class _FakeAsyncClient:
|
|||||||
self.room_send_calls.append(call)
|
self.room_send_calls.append(call)
|
||||||
if self.raise_on_send:
|
if self.raise_on_send:
|
||||||
raise RuntimeError("send failed")
|
raise RuntimeError("send failed")
|
||||||
|
return self.room_send_response
|
||||||
|
|
||||||
async def room_typing(
|
async def room_typing(
|
||||||
self,
|
self,
|
||||||
@@ -520,6 +525,7 @@ async def test_on_message_room_mention_requires_opt_in() -> None:
|
|||||||
source={"content": {"m.mentions": {"room": True}}},
|
source={"content": {"m.mentions": {"room": True}}},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
channel.config.allow_room_mentions = False
|
||||||
await channel._on_message(room, room_mention_event)
|
await channel._on_message(room, room_mention_event)
|
||||||
assert handled == []
|
assert handled == []
|
||||||
assert client.typing_calls == []
|
assert client.typing_calls == []
|
||||||
@@ -1322,3 +1328,302 @@ async def test_send_keeps_plaintext_only_for_plain_text() -> None:
|
|||||||
"body": text,
|
"body": text,
|
||||||
"m.mentions": {},
|
"m.mentions": {},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_basic_text() -> None:
|
||||||
|
"""Test basic text content without HTML formatting."""
|
||||||
|
result = _build_matrix_text_content("Hello, World!")
|
||||||
|
expected = {
|
||||||
|
"msgtype": "m.text",
|
||||||
|
"body": "Hello, World!",
|
||||||
|
"m.mentions": {}
|
||||||
|
}
|
||||||
|
assert expected == result
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_with_markdown() -> None:
|
||||||
|
"""Test text content with markdown that renders to HTML."""
|
||||||
|
text = "*Hello* **World**"
|
||||||
|
result = _build_matrix_text_content(text)
|
||||||
|
assert "msgtype" in result
|
||||||
|
assert "body" in result
|
||||||
|
assert result["body"] == text
|
||||||
|
assert "format" in result
|
||||||
|
assert result["format"] == "org.matrix.custom.html"
|
||||||
|
assert "formatted_body" in result
|
||||||
|
assert isinstance(result["formatted_body"], str)
|
||||||
|
assert len(result["formatted_body"]) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_with_event_id() -> None:
|
||||||
|
"""Test text content with event_id for message replacement."""
|
||||||
|
event_id = "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
result = _build_matrix_text_content("Updated message", event_id)
|
||||||
|
assert "msgtype" in result
|
||||||
|
assert "body" in result
|
||||||
|
assert result["m.new_content"]
|
||||||
|
assert result["m.new_content"]["body"] == "Updated message"
|
||||||
|
assert result["m.relates_to"]["rel_type"] == "m.replace"
|
||||||
|
assert result["m.relates_to"]["event_id"] == event_id
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_with_event_id_preserves_thread_relation() -> None:
|
||||||
|
"""Thread relations for edits should stay inside m.new_content."""
|
||||||
|
relates_to = {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": "$root1",
|
||||||
|
"m.in_reply_to": {"event_id": "$reply1"},
|
||||||
|
"is_falling_back": True,
|
||||||
|
}
|
||||||
|
result = _build_matrix_text_content("Updated message", "event-1", relates_to)
|
||||||
|
|
||||||
|
assert result["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "event-1",
|
||||||
|
}
|
||||||
|
assert result["m.new_content"]["m.relates_to"] == relates_to
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_no_event_id() -> None:
|
||||||
|
"""Test that when event_id is not provided, no extra properties are added."""
|
||||||
|
result = _build_matrix_text_content("Regular message")
|
||||||
|
|
||||||
|
# Basic required properties should be present
|
||||||
|
assert "msgtype" in result
|
||||||
|
assert "body" in result
|
||||||
|
assert result["body"] == "Regular message"
|
||||||
|
|
||||||
|
# Extra properties for replacement should NOT be present
|
||||||
|
assert "m.relates_to" not in result
|
||||||
|
assert "m.new_content" not in result
|
||||||
|
assert "format" not in result
|
||||||
|
assert "formatted_body" not in result
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_matrix_text_content_plain_text_no_html() -> None:
|
||||||
|
"""Test plain text that should not include HTML formatting."""
|
||||||
|
result = _build_matrix_text_content("Simple plain text")
|
||||||
|
assert "msgtype" in result
|
||||||
|
assert "body" in result
|
||||||
|
assert "format" not in result
|
||||||
|
assert "formatted_body" not in result
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_room_content_returns_room_send_response():
|
||||||
|
"""Test that _send_room_content returns the response from client.room_send."""
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
room_id = "!test_room:matrix.org"
|
||||||
|
content = {"msgtype": "m.text", "body": "Hello World"}
|
||||||
|
|
||||||
|
result = await channel._send_room_content(room_id, content)
|
||||||
|
|
||||||
|
assert result is client.room_send_response
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_creates_stream_buffer_and_sends_initial_message() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello")
|
||||||
|
|
||||||
|
assert "!room:matrix.org" in channel._stream_bufs
|
||||||
|
buf = channel._stream_bufs["!room:matrix.org"]
|
||||||
|
assert buf.text == "Hello"
|
||||||
|
assert buf.event_id == "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
assert client.room_send_calls[0]["content"]["body"] == "Hello"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_appends_without_sending_before_edit_interval(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
|
||||||
|
now = 100.0
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: now)
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello")
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", " world")
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
|
||||||
|
buf = channel._stream_bufs["!room:matrix.org"]
|
||||||
|
assert buf.text == "Hello world"
|
||||||
|
assert buf.event_id == "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_edits_again_after_interval(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo"
|
||||||
|
|
||||||
|
times = [100.0, 102.0, 104.0, 106.0, 108.0]
|
||||||
|
times.reverse()
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: times and times.pop())
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello")
|
||||||
|
await channel.send_delta("!room:matrix.org", " world")
|
||||||
|
|
||||||
|
assert len(client.room_send_calls) == 2
|
||||||
|
first_content = client.room_send_calls[0]["content"]
|
||||||
|
second_content = client.room_send_calls[1]["content"]
|
||||||
|
|
||||||
|
assert "body" in first_content
|
||||||
|
assert first_content["body"] == "Hello"
|
||||||
|
assert "m.relates_to" not in first_content
|
||||||
|
|
||||||
|
assert "body" in second_content
|
||||||
|
assert "m.relates_to" in second_content
|
||||||
|
assert second_content["body"] == "Hello world"
|
||||||
|
assert second_content["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "$8E2XVyINbEhcuAxvxd1d9JhQosNPzkVoU8TrbCAvyHo",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_stream_end_replaces_existing_message() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
channel._stream_bufs["!room:matrix.org"] = matrix_module._StreamBuf(
|
||||||
|
text="Final text",
|
||||||
|
event_id="event-1",
|
||||||
|
last_edit=100.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True})
|
||||||
|
|
||||||
|
assert "!room:matrix.org" not in channel._stream_bufs
|
||||||
|
assert client.typing_calls[-1] == ("!room:matrix.org", False, TYPING_NOTICE_TIMEOUT_MS)
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
assert client.room_send_calls[0]["content"]["body"] == "Final text"
|
||||||
|
assert client.room_send_calls[0]["content"]["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "event-1",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_starts_threaded_stream_inside_thread() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "event-1"
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
"thread_root_event_id": "$root1",
|
||||||
|
"thread_reply_to_event_id": "$reply1",
|
||||||
|
}
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello", metadata)
|
||||||
|
|
||||||
|
assert client.room_send_calls[0]["content"]["m.relates_to"] == {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": "$root1",
|
||||||
|
"m.in_reply_to": {"event_id": "$reply1"},
|
||||||
|
"is_falling_back": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_threaded_edit_keeps_replace_and_thread_relation(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
client.room_send_response.event_id = "event-1"
|
||||||
|
|
||||||
|
times = [100.0, 102.0, 104.0]
|
||||||
|
times.reverse()
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: times and times.pop())
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
"thread_root_event_id": "$root1",
|
||||||
|
"thread_reply_to_event_id": "$reply1",
|
||||||
|
}
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello", metadata)
|
||||||
|
await channel.send_delta("!room:matrix.org", " world", metadata)
|
||||||
|
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True, **metadata})
|
||||||
|
|
||||||
|
edit_content = client.room_send_calls[1]["content"]
|
||||||
|
final_content = client.room_send_calls[2]["content"]
|
||||||
|
|
||||||
|
assert edit_content["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "event-1",
|
||||||
|
}
|
||||||
|
assert edit_content["m.new_content"]["m.relates_to"] == {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": "$root1",
|
||||||
|
"m.in_reply_to": {"event_id": "$reply1"},
|
||||||
|
"is_falling_back": True,
|
||||||
|
}
|
||||||
|
assert final_content["m.relates_to"] == {
|
||||||
|
"rel_type": "m.replace",
|
||||||
|
"event_id": "event-1",
|
||||||
|
}
|
||||||
|
assert final_content["m.new_content"]["m.relates_to"] == {
|
||||||
|
"rel_type": "m.thread",
|
||||||
|
"event_id": "$root1",
|
||||||
|
"m.in_reply_to": {"event_id": "$reply1"},
|
||||||
|
"is_falling_back": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_stream_end_noop_when_buffer_missing() -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True})
|
||||||
|
|
||||||
|
assert client.room_send_calls == []
|
||||||
|
assert client.typing_calls == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_on_error_stops_typing(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
client.raise_on_send = True
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
now = 100.0
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: now)
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", "Hello", {"room_id": "!room:matrix.org"})
|
||||||
|
|
||||||
|
assert "!room:matrix.org" in channel._stream_bufs
|
||||||
|
assert channel._stream_bufs["!room:matrix.org"].text == "Hello"
|
||||||
|
assert len(client.room_send_calls) == 1
|
||||||
|
|
||||||
|
assert len(client.typing_calls) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_delta_ignores_whitespace_only_delta(monkeypatch) -> None:
|
||||||
|
channel = MatrixChannel(_make_config(), MessageBus())
|
||||||
|
client = _FakeAsyncClient("", "", "", None)
|
||||||
|
channel.client = client
|
||||||
|
|
||||||
|
now = 100.0
|
||||||
|
monkeypatch.setattr(channel, "monotonic_time", lambda: now)
|
||||||
|
|
||||||
|
await channel.send_delta("!room:matrix.org", " ")
|
||||||
|
|
||||||
|
assert "!room:matrix.org" in channel._stream_bufs
|
||||||
|
assert channel._stream_bufs["!room:matrix.org"].text == " "
|
||||||
|
assert client.room_send_calls == []
|
||||||
@@ -1,17 +1,22 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import tempfile
|
import tempfile
|
||||||
|
from pathlib import Path
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
import nanobot.channels.weixin as weixin_mod
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.channels.weixin import (
|
from nanobot.channels.weixin import (
|
||||||
ITEM_IMAGE,
|
ITEM_IMAGE,
|
||||||
ITEM_TEXT,
|
ITEM_TEXT,
|
||||||
MESSAGE_TYPE_BOT,
|
MESSAGE_TYPE_BOT,
|
||||||
WEIXIN_CHANNEL_VERSION,
|
WEIXIN_CHANNEL_VERSION,
|
||||||
|
_decrypt_aes_ecb,
|
||||||
|
_encrypt_aes_ecb,
|
||||||
WeixinChannel,
|
WeixinChannel,
|
||||||
WeixinConfig,
|
WeixinConfig,
|
||||||
)
|
)
|
||||||
@@ -42,10 +47,12 @@ def test_make_headers_includes_route_tag_when_configured() -> None:
|
|||||||
|
|
||||||
assert headers["Authorization"] == "Bearer token"
|
assert headers["Authorization"] == "Bearer token"
|
||||||
assert headers["SKRouteTag"] == "123"
|
assert headers["SKRouteTag"] == "123"
|
||||||
|
assert headers["iLink-App-Id"] == "bot"
|
||||||
|
assert headers["iLink-App-ClientVersion"] == str((2 << 16) | (1 << 8) | 1)
|
||||||
|
|
||||||
|
|
||||||
def test_channel_version_matches_reference_plugin_version() -> None:
|
def test_channel_version_matches_reference_plugin_version() -> None:
|
||||||
assert WEIXIN_CHANNEL_VERSION == "1.0.3"
|
assert WEIXIN_CHANNEL_VERSION == "2.1.1"
|
||||||
|
|
||||||
|
|
||||||
def test_save_and_load_state_persists_context_tokens(tmp_path) -> None:
|
def test_save_and_load_state_persists_context_tokens(tmp_path) -> None:
|
||||||
@@ -169,6 +176,120 @@ async def test_process_message_extracts_media_and_preserves_paths() -> None:
|
|||||||
assert inbound.media == ["/tmp/test.jpg"]
|
assert inbound.media == ["/tmp/test.jpg"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_process_message_falls_back_to_referenced_media_when_no_top_level_media() -> None:
|
||||||
|
channel, bus = _make_channel()
|
||||||
|
channel._download_media_item = AsyncMock(return_value="/tmp/ref.jpg")
|
||||||
|
|
||||||
|
await channel._process_message(
|
||||||
|
{
|
||||||
|
"message_type": 1,
|
||||||
|
"message_id": "m3-ref-fallback",
|
||||||
|
"from_user_id": "wx-user",
|
||||||
|
"context_token": "ctx-3-ref-fallback",
|
||||||
|
"item_list": [
|
||||||
|
{
|
||||||
|
"type": ITEM_TEXT,
|
||||||
|
"text_item": {"text": "reply to image"},
|
||||||
|
"ref_msg": {
|
||||||
|
"message_item": {
|
||||||
|
"type": ITEM_IMAGE,
|
||||||
|
"image_item": {"media": {"encrypt_query_param": "ref-enc"}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
inbound = await asyncio.wait_for(bus.consume_inbound(), timeout=1.0)
|
||||||
|
|
||||||
|
channel._download_media_item.assert_awaited_once_with(
|
||||||
|
{"media": {"encrypt_query_param": "ref-enc"}},
|
||||||
|
"image",
|
||||||
|
)
|
||||||
|
assert inbound.media == ["/tmp/ref.jpg"]
|
||||||
|
assert "reply to image" in inbound.content
|
||||||
|
assert "[image]" in inbound.content
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_process_message_does_not_use_referenced_fallback_when_top_level_media_exists() -> None:
|
||||||
|
channel, bus = _make_channel()
|
||||||
|
channel._download_media_item = AsyncMock(side_effect=["/tmp/top.jpg", "/tmp/ref.jpg"])
|
||||||
|
|
||||||
|
await channel._process_message(
|
||||||
|
{
|
||||||
|
"message_type": 1,
|
||||||
|
"message_id": "m3-ref-no-fallback",
|
||||||
|
"from_user_id": "wx-user",
|
||||||
|
"context_token": "ctx-3-ref-no-fallback",
|
||||||
|
"item_list": [
|
||||||
|
{"type": ITEM_IMAGE, "image_item": {"media": {"encrypt_query_param": "top-enc"}}},
|
||||||
|
{
|
||||||
|
"type": ITEM_TEXT,
|
||||||
|
"text_item": {"text": "has top-level media"},
|
||||||
|
"ref_msg": {
|
||||||
|
"message_item": {
|
||||||
|
"type": ITEM_IMAGE,
|
||||||
|
"image_item": {"media": {"encrypt_query_param": "ref-enc"}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
inbound = await asyncio.wait_for(bus.consume_inbound(), timeout=1.0)
|
||||||
|
|
||||||
|
channel._download_media_item.assert_awaited_once_with(
|
||||||
|
{"media": {"encrypt_query_param": "top-enc"}},
|
||||||
|
"image",
|
||||||
|
)
|
||||||
|
assert inbound.media == ["/tmp/top.jpg"]
|
||||||
|
assert "/tmp/ref.jpg" not in inbound.content
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_process_message_does_not_fallback_when_top_level_media_exists_but_download_fails() -> None:
|
||||||
|
channel, bus = _make_channel()
|
||||||
|
# Top-level image download fails (None), referenced image would succeed if fallback were triggered.
|
||||||
|
channel._download_media_item = AsyncMock(side_effect=[None, "/tmp/ref.jpg"])
|
||||||
|
|
||||||
|
await channel._process_message(
|
||||||
|
{
|
||||||
|
"message_type": 1,
|
||||||
|
"message_id": "m3-ref-no-fallback-on-failure",
|
||||||
|
"from_user_id": "wx-user",
|
||||||
|
"context_token": "ctx-3-ref-no-fallback-on-failure",
|
||||||
|
"item_list": [
|
||||||
|
{"type": ITEM_IMAGE, "image_item": {"media": {"encrypt_query_param": "top-enc"}}},
|
||||||
|
{
|
||||||
|
"type": ITEM_TEXT,
|
||||||
|
"text_item": {"text": "quoted has media"},
|
||||||
|
"ref_msg": {
|
||||||
|
"message_item": {
|
||||||
|
"type": ITEM_IMAGE,
|
||||||
|
"image_item": {"media": {"encrypt_query_param": "ref-enc"}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
inbound = await asyncio.wait_for(bus.consume_inbound(), timeout=1.0)
|
||||||
|
|
||||||
|
# Should only attempt top-level media item; reference fallback must not activate.
|
||||||
|
channel._download_media_item.assert_awaited_once_with(
|
||||||
|
{"media": {"encrypt_query_param": "top-enc"}},
|
||||||
|
"image",
|
||||||
|
)
|
||||||
|
assert inbound.media == []
|
||||||
|
assert "[image]" in inbound.content
|
||||||
|
assert "/tmp/ref.jpg" not in inbound.content
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_send_without_context_token_does_not_send_text() -> None:
|
async def test_send_without_context_token_does_not_send_text() -> None:
|
||||||
channel, _bus = _make_channel()
|
channel, _bus = _make_channel()
|
||||||
@@ -199,6 +320,70 @@ async def test_send_does_not_send_when_session_is_paused() -> None:
|
|||||||
channel._send_text.assert_not_awaited()
|
channel._send_text.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_typing_ticket_fetches_and_caches_per_user() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
channel._api_post = AsyncMock(return_value={"ret": 0, "typing_ticket": "ticket-1"})
|
||||||
|
|
||||||
|
first = await channel._get_typing_ticket("wx-user", "ctx-1")
|
||||||
|
second = await channel._get_typing_ticket("wx-user", "ctx-2")
|
||||||
|
|
||||||
|
assert first == "ticket-1"
|
||||||
|
assert second == "ticket-1"
|
||||||
|
channel._api_post.assert_awaited_once_with(
|
||||||
|
"ilink/bot/getconfig",
|
||||||
|
{"ilink_user_id": "wx-user", "context_token": "ctx-1", "base_info": weixin_mod.BASE_INFO},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_uses_typing_start_and_cancel_when_ticket_available() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
channel._context_tokens["wx-user"] = "ctx-typing"
|
||||||
|
channel._send_text = AsyncMock()
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"ret": 0, "typing_ticket": "ticket-typing"},
|
||||||
|
{"ret": 0},
|
||||||
|
{"ret": 0},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel.send(
|
||||||
|
type("Msg", (), {"chat_id": "wx-user", "content": "pong", "media": [], "metadata": {}})()
|
||||||
|
)
|
||||||
|
|
||||||
|
channel._send_text.assert_awaited_once_with("wx-user", "pong", "ctx-typing")
|
||||||
|
assert channel._api_post.await_count == 3
|
||||||
|
assert channel._api_post.await_args_list[0].args[0] == "ilink/bot/getconfig"
|
||||||
|
assert channel._api_post.await_args_list[1].args[0] == "ilink/bot/sendtyping"
|
||||||
|
assert channel._api_post.await_args_list[1].args[1]["status"] == 1
|
||||||
|
assert channel._api_post.await_args_list[2].args[0] == "ilink/bot/sendtyping"
|
||||||
|
assert channel._api_post.await_args_list[2].args[1]["status"] == 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_still_sends_text_when_typing_ticket_missing() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
channel._context_tokens["wx-user"] = "ctx-no-ticket"
|
||||||
|
channel._send_text = AsyncMock()
|
||||||
|
channel._api_post = AsyncMock(return_value={"ret": 1, "errmsg": "no config"})
|
||||||
|
|
||||||
|
await channel.send(
|
||||||
|
type("Msg", (), {"chat_id": "wx-user", "content": "pong", "media": [], "metadata": {}})()
|
||||||
|
)
|
||||||
|
|
||||||
|
channel._send_text.assert_awaited_once_with("wx-user", "pong", "ctx-no-ticket")
|
||||||
|
channel._api_post.assert_awaited_once()
|
||||||
|
assert channel._api_post.await_args_list[0].args[0] == "ilink/bot/getconfig"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_poll_once_pauses_session_on_expired_errcode() -> None:
|
async def test_poll_once_pauses_session_on_expired_errcode() -> None:
|
||||||
channel, _bus = _make_channel()
|
channel, _bus = _make_channel()
|
||||||
@@ -220,8 +405,12 @@ async def test_qr_login_refreshes_expired_qr_and_then_succeeds() -> None:
|
|||||||
channel._api_get = AsyncMock(
|
channel._api_get = AsyncMock(
|
||||||
side_effect=[
|
side_effect=[
|
||||||
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
||||||
{"status": "expired"},
|
|
||||||
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"status": "expired"},
|
||||||
{
|
{
|
||||||
"status": "confirmed",
|
"status": "confirmed",
|
||||||
"bot_token": "token-2",
|
"bot_token": "token-2",
|
||||||
@@ -247,12 +436,16 @@ async def test_qr_login_returns_false_after_too_many_expired_qr_codes() -> None:
|
|||||||
channel._api_get = AsyncMock(
|
channel._api_get = AsyncMock(
|
||||||
side_effect=[
|
side_effect=[
|
||||||
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
|
||||||
{"status": "expired"},
|
|
||||||
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
|
||||||
{"status": "expired"},
|
|
||||||
{"qrcode": "qr-3", "qrcode_img_content": "url-3"},
|
{"qrcode": "qr-3", "qrcode_img_content": "url-3"},
|
||||||
{"status": "expired"},
|
|
||||||
{"qrcode": "qr-4", "qrcode_img_content": "url-4"},
|
{"qrcode": "qr-4", "qrcode_img_content": "url-4"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"status": "expired"},
|
||||||
|
{"status": "expired"},
|
||||||
|
{"status": "expired"},
|
||||||
{"status": "expired"},
|
{"status": "expired"},
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
@@ -262,6 +455,105 @@ async def test_qr_login_returns_false_after_too_many_expired_qr_codes() -> None:
|
|||||||
assert ok is False
|
assert ok is False
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_switches_polling_base_url_on_redirect_status() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
|
||||||
|
|
||||||
|
status_side_effect = [
|
||||||
|
{"status": "scaned_but_redirect", "redirect_host": "idc.redirect.test"},
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-3",
|
||||||
|
"ilink_bot_id": "bot-3",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
channel._api_get = AsyncMock(side_effect=list(status_side_effect))
|
||||||
|
channel._api_get_with_base = AsyncMock(side_effect=list(status_side_effect))
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-3"
|
||||||
|
assert channel._api_get_with_base.await_count == 2
|
||||||
|
first_call = channel._api_get_with_base.await_args_list[0]
|
||||||
|
second_call = channel._api_get_with_base.await_args_list[1]
|
||||||
|
assert first_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
assert second_call.kwargs["base_url"] == "https://idc.redirect.test"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_redirect_without_host_keeps_current_polling_base_url() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
|
||||||
|
|
||||||
|
status_side_effect = [
|
||||||
|
{"status": "scaned_but_redirect"},
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-4",
|
||||||
|
"ilink_bot_id": "bot-4",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
channel._api_get = AsyncMock(side_effect=list(status_side_effect))
|
||||||
|
channel._api_get_with_base = AsyncMock(side_effect=list(status_side_effect))
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-4"
|
||||||
|
assert channel._api_get_with_base.await_count == 2
|
||||||
|
first_call = channel._api_get_with_base.await_args_list[0]
|
||||||
|
second_call = channel._api_get_with_base.await_args_list[1]
|
||||||
|
assert first_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
assert second_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_resets_redirect_base_url_after_qr_refresh() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(side_effect=[("qr-1", "url-1"), ("qr-2", "url-2")])
|
||||||
|
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"status": "scaned_but_redirect", "redirect_host": "idc.redirect.test"},
|
||||||
|
{"status": "expired"},
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-5",
|
||||||
|
"ilink_bot_id": "bot-5",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-5"
|
||||||
|
assert channel._api_get_with_base.await_count == 3
|
||||||
|
first_call = channel._api_get_with_base.await_args_list[0]
|
||||||
|
second_call = channel._api_get_with_base.await_args_list[1]
|
||||||
|
third_call = channel._api_get_with_base.await_args_list[2]
|
||||||
|
assert first_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
assert second_call.kwargs["base_url"] == "https://idc.redirect.test"
|
||||||
|
assert third_call.kwargs["base_url"] == "https://ilinkai.weixin.qq.com"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_process_message_skips_bot_messages() -> None:
|
async def test_process_message_skips_bot_messages() -> None:
|
||||||
channel, bus = _make_channel()
|
channel, bus = _make_channel()
|
||||||
@@ -278,3 +570,357 @@ async def test_process_message_skips_bot_messages() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert bus.inbound_size == 0
|
assert bus.inbound_size == 0
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyHttpResponse:
|
||||||
|
def __init__(self, *, headers: dict[str, str] | None = None, status_code: int = 200) -> None:
|
||||||
|
self.headers = headers or {}
|
||||||
|
self.status_code = status_code
|
||||||
|
|
||||||
|
def raise_for_status(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_media_uses_upload_full_url_when_present(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
|
||||||
|
media_file = tmp_path / "photo.jpg"
|
||||||
|
media_file.write_bytes(b"hello-weixin")
|
||||||
|
|
||||||
|
cdn_post = AsyncMock(return_value=_DummyHttpResponse(headers={"x-encrypted-param": "dl-param"}))
|
||||||
|
channel._client = SimpleNamespace(post=cdn_post)
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{
|
||||||
|
"upload_full_url": "https://upload-full.example.test/path?foo=bar",
|
||||||
|
"upload_param": "should-not-be-used",
|
||||||
|
},
|
||||||
|
{"ret": 0},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._send_media_file("wx-user", str(media_file), "ctx-1")
|
||||||
|
|
||||||
|
# first POST call is CDN upload
|
||||||
|
cdn_url = cdn_post.await_args_list[0].args[0]
|
||||||
|
assert cdn_url == "https://upload-full.example.test/path?foo=bar"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_media_falls_back_to_upload_param_url(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
|
||||||
|
media_file = tmp_path / "photo.jpg"
|
||||||
|
media_file.write_bytes(b"hello-weixin")
|
||||||
|
|
||||||
|
cdn_post = AsyncMock(return_value=_DummyHttpResponse(headers={"x-encrypted-param": "dl-param"}))
|
||||||
|
channel._client = SimpleNamespace(post=cdn_post)
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"upload_param": "enc-need-fallback"},
|
||||||
|
{"ret": 0},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._send_media_file("wx-user", str(media_file), "ctx-1")
|
||||||
|
|
||||||
|
cdn_url = cdn_post.await_args_list[0].args[0]
|
||||||
|
assert cdn_url.startswith(f"{channel.config.cdn_base_url}/upload?encrypted_query_param=enc-need-fallback")
|
||||||
|
assert "&filekey=" in cdn_url
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_media_voice_file_uses_voice_item_and_voice_upload_type(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
|
||||||
|
media_file = tmp_path / "voice.mp3"
|
||||||
|
media_file.write_bytes(b"voice-bytes")
|
||||||
|
|
||||||
|
cdn_post = AsyncMock(return_value=_DummyHttpResponse(headers={"x-encrypted-param": "voice-dl-param"}))
|
||||||
|
channel._client = SimpleNamespace(post=cdn_post)
|
||||||
|
channel._api_post = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
{"upload_full_url": "https://upload-full.example.test/voice?foo=bar"},
|
||||||
|
{"ret": 0},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
await channel._send_media_file("wx-user", str(media_file), "ctx-voice")
|
||||||
|
|
||||||
|
getupload_body = channel._api_post.await_args_list[0].args[1]
|
||||||
|
assert getupload_body["media_type"] == 4
|
||||||
|
|
||||||
|
sendmessage_body = channel._api_post.await_args_list[1].args[1]
|
||||||
|
item = sendmessage_body["msg"]["item_list"][0]
|
||||||
|
assert item["type"] == 3
|
||||||
|
assert "voice_item" in item
|
||||||
|
assert "file_item" not in item
|
||||||
|
assert item["voice_item"]["media"]["encrypt_query_param"] == "voice-dl-param"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_send_typing_uses_keepalive_until_send_finishes() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
channel._context_tokens["wx-user"] = "ctx-typing-loop"
|
||||||
|
async def _api_post_side_effect(endpoint: str, _body: dict | None = None, *, auth: bool = True):
|
||||||
|
if endpoint == "ilink/bot/getconfig":
|
||||||
|
return {"ret": 0, "typing_ticket": "ticket-keepalive"}
|
||||||
|
return {"ret": 0}
|
||||||
|
|
||||||
|
channel._api_post = AsyncMock(side_effect=_api_post_side_effect)
|
||||||
|
|
||||||
|
async def _slow_send_text(*_args, **_kwargs) -> None:
|
||||||
|
await asyncio.sleep(0.03)
|
||||||
|
|
||||||
|
channel._send_text = AsyncMock(side_effect=_slow_send_text)
|
||||||
|
|
||||||
|
old_interval = weixin_mod.TYPING_KEEPALIVE_INTERVAL_S
|
||||||
|
weixin_mod.TYPING_KEEPALIVE_INTERVAL_S = 0.01
|
||||||
|
try:
|
||||||
|
await channel.send(
|
||||||
|
type("Msg", (), {"chat_id": "wx-user", "content": "pong", "media": [], "metadata": {}})()
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
weixin_mod.TYPING_KEEPALIVE_INTERVAL_S = old_interval
|
||||||
|
|
||||||
|
status_calls = [
|
||||||
|
c.args[1]["status"]
|
||||||
|
for c in channel._api_post.await_args_list
|
||||||
|
if c.args and c.args[0] == "ilink/bot/sendtyping"
|
||||||
|
]
|
||||||
|
assert status_calls.count(1) >= 2
|
||||||
|
assert status_calls[-1] == 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_typing_ticket_failure_uses_backoff_and_cached_ticket(monkeypatch) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._client = object()
|
||||||
|
channel._token = "token"
|
||||||
|
|
||||||
|
now = {"value": 1000.0}
|
||||||
|
monkeypatch.setattr(weixin_mod.time, "time", lambda: now["value"])
|
||||||
|
monkeypatch.setattr(weixin_mod.random, "random", lambda: 0.5)
|
||||||
|
|
||||||
|
channel._api_post = AsyncMock(return_value={"ret": 0, "typing_ticket": "ticket-ok"})
|
||||||
|
first = await channel._get_typing_ticket("wx-user", "ctx-1")
|
||||||
|
assert first == "ticket-ok"
|
||||||
|
|
||||||
|
# force refresh window reached
|
||||||
|
now["value"] = now["value"] + (12 * 60 * 60) + 1
|
||||||
|
channel._api_post = AsyncMock(return_value={"ret": 1, "errmsg": "temporary failure"})
|
||||||
|
|
||||||
|
# On refresh failure, should still return cached ticket and apply backoff.
|
||||||
|
second = await channel._get_typing_ticket("wx-user", "ctx-2")
|
||||||
|
assert second == "ticket-ok"
|
||||||
|
assert channel._api_post.await_count == 1
|
||||||
|
|
||||||
|
# Before backoff expiry, no extra fetch should happen.
|
||||||
|
now["value"] += 1
|
||||||
|
third = await channel._get_typing_ticket("wx-user", "ctx-3")
|
||||||
|
assert third == "ticket-ok"
|
||||||
|
assert channel._api_post.await_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_treats_temporary_connect_error_as_wait_and_recovers() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
|
||||||
|
|
||||||
|
request = httpx.Request("GET", "https://ilinkai.weixin.qq.com/ilink/bot/get_qrcode_status")
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
httpx.ConnectError("temporary network", request=request),
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-net-ok",
|
||||||
|
"ilink_bot_id": "bot-id",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-net-ok"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_qr_login_treats_5xx_gateway_response_error_as_wait_and_recovers() -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
channel._running = True
|
||||||
|
channel._save_state = lambda: None
|
||||||
|
channel._print_qr_code = lambda url: None
|
||||||
|
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
|
||||||
|
|
||||||
|
request = httpx.Request("GET", "https://ilinkai.weixin.qq.com/ilink/bot/get_qrcode_status")
|
||||||
|
response = httpx.Response(status_code=524, request=request)
|
||||||
|
channel._api_get_with_base = AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
httpx.HTTPStatusError("gateway timeout", request=request, response=response),
|
||||||
|
{
|
||||||
|
"status": "confirmed",
|
||||||
|
"bot_token": "token-5xx-ok",
|
||||||
|
"ilink_bot_id": "bot-id",
|
||||||
|
"baseurl": "https://example.test",
|
||||||
|
"ilink_user_id": "wx-user",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
ok = await channel._qr_login()
|
||||||
|
|
||||||
|
assert ok is True
|
||||||
|
assert channel._token == "token-5xx-ok"
|
||||||
|
|
||||||
|
|
||||||
|
def test_decrypt_aes_ecb_strips_valid_pkcs7_padding() -> None:
|
||||||
|
key_b64 = "MDEyMzQ1Njc4OWFiY2RlZg==" # base64("0123456789abcdef")
|
||||||
|
plaintext = b"hello-weixin-padding"
|
||||||
|
|
||||||
|
ciphertext = _encrypt_aes_ecb(plaintext, key_b64)
|
||||||
|
decrypted = _decrypt_aes_ecb(ciphertext, key_b64)
|
||||||
|
|
||||||
|
assert decrypted == plaintext
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyDownloadResponse:
|
||||||
|
def __init__(self, content: bytes, status_code: int = 200) -> None:
|
||||||
|
self.content = content
|
||||||
|
self.status_code = status_code
|
||||||
|
|
||||||
|
def raise_for_status(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyErrorDownloadResponse(_DummyDownloadResponse):
|
||||||
|
def __init__(self, url: str, status_code: int) -> None:
|
||||||
|
super().__init__(content=b"", status_code=status_code)
|
||||||
|
self._url = url
|
||||||
|
|
||||||
|
def raise_for_status(self) -> None:
|
||||||
|
request = httpx.Request("GET", self._url)
|
||||||
|
response = httpx.Response(self.status_code, request=request)
|
||||||
|
raise httpx.HTTPStatusError(
|
||||||
|
f"download failed with status {self.status_code}",
|
||||||
|
request=request,
|
||||||
|
response=response,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_uses_full_url_when_present(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
full_url = "https://cdn.example.test/download/full"
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(return_value=_DummyDownloadResponse(content=b"raw-image-bytes"))
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {
|
||||||
|
"media": {
|
||||||
|
"full_url": full_url,
|
||||||
|
"encrypt_query_param": "enc-fallback-should-not-be-used",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
saved_path = await channel._download_media_item(item, "image")
|
||||||
|
|
||||||
|
assert saved_path is not None
|
||||||
|
assert Path(saved_path).read_bytes() == b"raw-image-bytes"
|
||||||
|
channel._client.get.assert_awaited_once_with(full_url)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_falls_back_when_full_url_returns_retryable_error(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
full_url = "https://cdn.example.test/download/full?taskid=123"
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(
|
||||||
|
side_effect=[
|
||||||
|
_DummyErrorDownloadResponse(full_url, 500),
|
||||||
|
_DummyDownloadResponse(content=b"fallback-bytes"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {
|
||||||
|
"media": {
|
||||||
|
"full_url": full_url,
|
||||||
|
"encrypt_query_param": "enc-fallback",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
saved_path = await channel._download_media_item(item, "image")
|
||||||
|
|
||||||
|
assert saved_path is not None
|
||||||
|
assert Path(saved_path).read_bytes() == b"fallback-bytes"
|
||||||
|
assert channel._client.get.await_count == 2
|
||||||
|
assert channel._client.get.await_args_list[0].args[0] == full_url
|
||||||
|
fallback_url = channel._client.get.await_args_list[1].args[0]
|
||||||
|
assert fallback_url.startswith(f"{channel.config.cdn_base_url}/download?encrypted_query_param=enc-fallback")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_falls_back_to_encrypt_query_param(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(return_value=_DummyDownloadResponse(content=b"fallback-bytes"))
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {"media": {"encrypt_query_param": "enc-fallback"}}
|
||||||
|
saved_path = await channel._download_media_item(item, "image")
|
||||||
|
|
||||||
|
assert saved_path is not None
|
||||||
|
assert Path(saved_path).read_bytes() == b"fallback-bytes"
|
||||||
|
called_url = channel._client.get.await_args_list[0].args[0]
|
||||||
|
assert called_url.startswith(f"{channel.config.cdn_base_url}/download?encrypted_query_param=enc-fallback")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_does_not_retry_when_full_url_fails_without_fallback(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
full_url = "https://cdn.example.test/download/full"
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(return_value=_DummyErrorDownloadResponse(full_url, 500))
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {"media": {"full_url": full_url}}
|
||||||
|
saved_path = await channel._download_media_item(item, "image")
|
||||||
|
|
||||||
|
assert saved_path is None
|
||||||
|
channel._client.get.assert_awaited_once_with(full_url)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_download_media_item_non_image_requires_aes_key_even_with_full_url(tmp_path) -> None:
|
||||||
|
channel, _bus = _make_channel()
|
||||||
|
weixin_mod.get_media_dir = lambda _name: tmp_path
|
||||||
|
|
||||||
|
full_url = "https://cdn.example.test/download/voice"
|
||||||
|
channel._client = SimpleNamespace(
|
||||||
|
get=AsyncMock(return_value=_DummyDownloadResponse(content=b"ciphertext-or-unknown"))
|
||||||
|
)
|
||||||
|
|
||||||
|
item = {
|
||||||
|
"media": {
|
||||||
|
"full_url": full_url,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
saved_path = await channel._download_media_item(item, "voice")
|
||||||
|
|
||||||
|
assert saved_path is None
|
||||||
|
channel._client.get.assert_not_awaited()
|
||||||
|
|||||||
@@ -408,6 +408,56 @@ async def test_exec_timeout_capped_at_max() -> None:
|
|||||||
assert "Exit code: 0" in result
|
assert "Exit code: 0" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_applied() -> None:
|
||||||
|
"""command_wrapper should wrap the original command."""
|
||||||
|
tool = ExecTool(command_wrapper="echo WRAPPED: {command}")
|
||||||
|
result = await tool.execute(command="hello")
|
||||||
|
assert "WRAPPED: hello" in result
|
||||||
|
assert "Exit code: 0" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_with_cwd(tmp_path) -> None:
|
||||||
|
"""command_wrapper should substitute {cwd} with the absolute working directory."""
|
||||||
|
tool = ExecTool(command_wrapper="echo CWD:{cwd} CMD:{command}")
|
||||||
|
result = await tool.execute(command="hi", working_dir=str(tmp_path))
|
||||||
|
assert str(tmp_path) in result
|
||||||
|
assert "CMD:hi" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_empty_noop() -> None:
|
||||||
|
"""Empty command_wrapper should leave the command unchanged."""
|
||||||
|
tool = ExecTool(command_wrapper="")
|
||||||
|
result = await tool.execute(command="echo direct")
|
||||||
|
assert "direct" in result
|
||||||
|
assert "Exit code: 0" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_guard_runs_before_wrapper() -> None:
|
||||||
|
"""Safety guard should run before wrapper substitution."""
|
||||||
|
tool = ExecTool(command_wrapper="echo WRAPPED:{command}")
|
||||||
|
result = await tool.execute(command="rm -rf /")
|
||||||
|
assert "blocked by safety guard" in result
|
||||||
|
assert "WRAPPED:" not in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_ignores_unknown_placeholders() -> None:
|
||||||
|
"""Unknown {placeholders} in the wrapper should be left as-is, not raise KeyError."""
|
||||||
|
tool = ExecTool(command_wrapper="echo {command} {unknown}")
|
||||||
|
result = await tool.execute(command="hello")
|
||||||
|
assert "Exit code: 0" in result
|
||||||
|
assert "{unknown}" in result
|
||||||
|
|
||||||
|
|
||||||
|
async def test_exec_command_wrapper_does_not_leak_attributes() -> None:
|
||||||
|
"""Wrapper should not expose Python internals via attribute access."""
|
||||||
|
tool = ExecTool(command_wrapper="echo {command.__class__}")
|
||||||
|
result = await tool.execute(command="hello")
|
||||||
|
assert "Exit code: 0" in result
|
||||||
|
# {command.__class__} is not a valid placeholder; {command} gets replaced
|
||||||
|
# leaving {.__class__} as a literal string — no Python object is leaked.
|
||||||
|
assert "<class" not in result
|
||||||
|
|
||||||
|
|
||||||
# --- _resolve_type and nullable param tests ---
|
# --- _resolve_type and nullable param tests ---
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user