da15b70535
Closes #392 (uncontroversial part). WeChat (`_handle_message`) and Feishu (`_handle_event`) gated their signature/decryption checks behind a condition the REQUEST controls: - WeChat: `if encrypt and self._crypto:` -- a POST with no `<Encrypt>` element took the false branch and reached `_safe_process_message` without any verification, even when `encoding_aes_key` + `token` were configured. - Feishu: `if self.config.encrypt_key and "encrypt" in body:` -- a plaintext body skipped decryption entirely and was processed directly. Since the webhook port is the channel's only inbound boundary, an attacker could POST forged plaintext and reach the agent, spoofing `sender_id` / `FromUserName` (and, with an empty allowlist, passing the sender gate). Fix: when encryption is configured, an inbound POST MUST carry the encrypted field (`<Encrypt>` / `encrypt`) -- otherwise it is rejected with 403 and never reaches the agent. Plaintext mode (no encryption configured) is unchanged, so existing plaintext deployments are not affected. The remaining fail-closed question (what to do when credentials are entirely unset) is left for the maintainers to decide as the policy part of the issue. Regression tests (9 new): - WeChat: plaintext rejected / missing Encrypt rejected / bad signature rejected / valid signature decrypts and processes / plaintext still accepted when no crypto. - Feishu: plaintext rejected / non-dict body rejected / encrypted body decrypts and processes / plaintext still accepted when no encrypt_key. 93 tests in the two channel files pass; full suite 3045 passed, 13 skipped; ruff clean. Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
575 lines
20 KiB
Python
575 lines
20 KiB
Python
"""Tests for WeChat channel implementation."""
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import time
|
|
import xml.etree.ElementTree as ET
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
import pytest
|
|
|
|
from EvoScientist.channels.base import ChannelError
|
|
from EvoScientist.channels.wechat.channel import (
|
|
WeChatChannel,
|
|
WeChatMPConfig,
|
|
WeComConfig,
|
|
_strip_markdown,
|
|
)
|
|
from EvoScientist.channels.wechat.crypto import (
|
|
WeChatCrypto,
|
|
_pkcs7_pad,
|
|
_pkcs7_unpad,
|
|
parse_xml,
|
|
)
|
|
|
|
# ── Config tests ──────────────────────────────────────────────────
|
|
|
|
|
|
class TestWeComConfig:
|
|
def test_default_values(self):
|
|
config = WeComConfig()
|
|
assert config.corp_id == ""
|
|
assert config.agent_id == ""
|
|
assert config.secret == ""
|
|
assert config.webhook_port == 9001
|
|
assert config.allowed_senders is None
|
|
assert config.text_chunk_limit == 4096
|
|
|
|
def test_custom_values(self):
|
|
config = WeComConfig(
|
|
corp_id="corp123",
|
|
agent_id="1000001",
|
|
secret="my-secret",
|
|
token="my-token",
|
|
encoding_aes_key="a" * 43,
|
|
webhook_port=8080,
|
|
allowed_senders={"user1", "user2"},
|
|
)
|
|
assert config.corp_id == "corp123"
|
|
assert config.agent_id == "1000001"
|
|
assert config.allowed_senders == {"user1", "user2"}
|
|
assert config.webhook_port == 8080
|
|
|
|
|
|
class TestWeChatMPConfig:
|
|
def test_default_values(self):
|
|
config = WeChatMPConfig()
|
|
assert config.app_id == ""
|
|
assert config.app_secret == ""
|
|
assert config.webhook_port == 9001
|
|
|
|
def test_custom_values(self):
|
|
config = WeChatMPConfig(
|
|
app_id="wx1234",
|
|
app_secret="secret",
|
|
token="mp-token",
|
|
)
|
|
assert config.app_id == "wx1234"
|
|
|
|
|
|
# ── Channel init / lifecycle tests ────────────────────────────────
|
|
|
|
|
|
class TestWeChatChannelInit:
|
|
def test_wecom_init(self):
|
|
config = WeComConfig(corp_id="corp", agent_id="1", secret="s")
|
|
channel = WeChatChannel(config, backend="wecom")
|
|
assert channel.name == "wechat"
|
|
assert channel._backend == "wecom"
|
|
assert channel._running is False
|
|
|
|
def test_mp_init(self):
|
|
config = WeChatMPConfig(app_id="wx", app_secret="s")
|
|
channel = WeChatChannel(config, backend="wechatmp")
|
|
assert channel._backend == "wechatmp"
|
|
|
|
async def test_start_raises_without_corp_id(self):
|
|
config = WeComConfig(corp_id="", agent_id="1", secret="s")
|
|
channel = WeChatChannel(config, backend="wecom")
|
|
with pytest.raises(ChannelError, match="corp_id"):
|
|
await channel.start()
|
|
|
|
async def test_start_raises_without_secret(self):
|
|
config = WeComConfig(corp_id="corp", agent_id="1", secret="")
|
|
channel = WeChatChannel(config, backend="wecom")
|
|
with pytest.raises(ChannelError, match="secret"):
|
|
await channel.start()
|
|
|
|
async def test_start_raises_without_agent_id(self):
|
|
config = WeComConfig(corp_id="corp", agent_id="", secret="s")
|
|
channel = WeChatChannel(config, backend="wecom")
|
|
with pytest.raises(ChannelError, match="agent_id"):
|
|
await channel.start()
|
|
|
|
async def test_start_raises_mp_without_app_id(self):
|
|
config = WeChatMPConfig(app_id="", app_secret="s")
|
|
channel = WeChatChannel(config, backend="wechatmp")
|
|
with pytest.raises(ChannelError, match="app_id"):
|
|
await channel.start()
|
|
|
|
async def test_stop_when_not_running(self):
|
|
config = WeComConfig(corp_id="c", agent_id="1", secret="s")
|
|
channel = WeChatChannel(config, backend="wecom")
|
|
await channel.stop() # Should not raise
|
|
|
|
async def test_send_returns_false_without_client(self):
|
|
from EvoScientist.channels.base import OutboundMessage
|
|
|
|
config = WeComConfig(corp_id="c", agent_id="1", secret="s")
|
|
channel = WeChatChannel(config, backend="wecom")
|
|
msg = OutboundMessage(
|
|
channel="wechat",
|
|
chat_id="user1",
|
|
content="hello",
|
|
metadata={"chat_id": "user1"},
|
|
)
|
|
result = await channel.send(msg)
|
|
assert result is False
|
|
|
|
|
|
# ── Markdown stripping tests ──────────────────────────────────────
|
|
|
|
|
|
class TestStripMarkdown:
|
|
def test_plain_text(self):
|
|
assert _strip_markdown("hello world") == "hello world"
|
|
|
|
def test_bold(self):
|
|
assert _strip_markdown("**bold**") == "bold"
|
|
|
|
def test_italic(self):
|
|
assert _strip_markdown("_italic_") == "italic"
|
|
|
|
def test_code(self):
|
|
assert _strip_markdown("`code`") == "code"
|
|
|
|
def test_link(self):
|
|
result = _strip_markdown("[text](https://example.com)")
|
|
assert "text" in result
|
|
assert "https://example.com" in result
|
|
|
|
def test_heading(self):
|
|
assert _strip_markdown("## Title").strip() == "Title"
|
|
|
|
def test_list_items(self):
|
|
result = _strip_markdown("- item1\n- item2")
|
|
assert "• item1" in result
|
|
assert "• item2" in result
|
|
|
|
def test_strikethrough(self):
|
|
assert _strip_markdown("~~deleted~~") == "deleted"
|
|
|
|
def test_code_block(self):
|
|
text = "```python\nprint('hi')\n```"
|
|
result = _strip_markdown(text)
|
|
assert "print('hi')" in result
|
|
|
|
|
|
# ── XML parsing tests ─────────────────────────────────────────────
|
|
|
|
|
|
class TestParseXml:
|
|
def test_basic_text_message(self):
|
|
xml = (
|
|
"<xml>"
|
|
"<MsgType><![CDATA[text]]></MsgType>"
|
|
"<Content><![CDATA[hello]]></Content>"
|
|
"<FromUserName><![CDATA[user123]]></FromUserName>"
|
|
"<ToUserName><![CDATA[bot]]></ToUserName>"
|
|
"<MsgId>1234</MsgId>"
|
|
"<CreateTime>1700000000</CreateTime>"
|
|
"</xml>"
|
|
)
|
|
data = parse_xml(xml)
|
|
assert data["MsgType"] == "text"
|
|
assert data["Content"] == "hello"
|
|
assert data["FromUserName"] == "user123"
|
|
assert data["MsgId"] == "1234"
|
|
|
|
def test_image_message(self):
|
|
xml = (
|
|
"<xml>"
|
|
"<MsgType><![CDATA[image]]></MsgType>"
|
|
"<PicUrl><![CDATA[https://example.com/img.jpg]]></PicUrl>"
|
|
"<MediaId><![CDATA[media_123]]></MediaId>"
|
|
"<FromUserName><![CDATA[user1]]></FromUserName>"
|
|
"</xml>"
|
|
)
|
|
data = parse_xml(xml)
|
|
assert data["MsgType"] == "image"
|
|
assert data["PicUrl"] == "https://example.com/img.jpg"
|
|
|
|
def test_event_message(self):
|
|
xml = (
|
|
"<xml>"
|
|
"<MsgType><![CDATA[event]]></MsgType>"
|
|
"<Event><![CDATA[subscribe]]></Event>"
|
|
"<FromUserName><![CDATA[user1]]></FromUserName>"
|
|
"</xml>"
|
|
)
|
|
data = parse_xml(xml)
|
|
assert data["MsgType"] == "event"
|
|
assert data["Event"] == "subscribe"
|
|
|
|
|
|
# ── Crypto tests ──────────────────────────────────────────────────
|
|
|
|
|
|
class TestPKCS7:
|
|
def test_pad_unpad_roundtrip(self):
|
|
data = b"hello"
|
|
padded = _pkcs7_pad(data)
|
|
assert len(padded) % 32 == 0
|
|
assert _pkcs7_unpad(padded) == data
|
|
|
|
def test_pad_block_aligned(self):
|
|
data = b"x" * 32
|
|
padded = _pkcs7_pad(data)
|
|
assert len(padded) == 64 # full padding block added
|
|
assert _pkcs7_unpad(padded) == data
|
|
|
|
|
|
class TestWeChatCrypto:
|
|
"""Test the encryption/decryption roundtrip.
|
|
|
|
Uses a deterministic 43-char EncodingAESKey.
|
|
"""
|
|
|
|
# Skip encryption tests when no crypto backend is available
|
|
_has_crypto = False
|
|
try:
|
|
from Crypto.Cipher import AES as _aes
|
|
|
|
_has_crypto = True
|
|
except ImportError:
|
|
try:
|
|
import pyaes as _pyaes
|
|
|
|
_has_crypto = True
|
|
except ImportError:
|
|
pass
|
|
pytestmark = pytest.mark.skipif(
|
|
not _has_crypto,
|
|
reason="pycryptodome or pyaes required for encryption tests",
|
|
)
|
|
|
|
@pytest.fixture
|
|
def crypto(self):
|
|
# 43 base64 chars → 32 bytes AES key
|
|
key = "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG"
|
|
return WeChatCrypto(
|
|
token="test_token",
|
|
encoding_aes_key=key,
|
|
app_id="wx_test_app",
|
|
)
|
|
|
|
def test_encrypt_decrypt_roundtrip(self, crypto):
|
|
msg = "<xml><Content>Hello WeChat!</Content></xml>"
|
|
encrypted = crypto.encrypt(msg)
|
|
decrypted, app_id = crypto.decrypt(encrypted)
|
|
assert decrypted == msg
|
|
assert app_id == "wx_test_app"
|
|
|
|
def test_verify_signature(self, crypto):
|
|
timestamp = "1609459200"
|
|
nonce = "abc123"
|
|
parts = sorted([crypto.token, timestamp, nonce])
|
|
expected = hashlib.sha1("".join(parts).encode()).hexdigest()
|
|
assert crypto.verify_signature(expected, timestamp, nonce)
|
|
assert not crypto.verify_signature("wrong", timestamp, nonce)
|
|
|
|
def test_verify_signature_with_encrypt(self, crypto):
|
|
timestamp = "1609459200"
|
|
nonce = "abc123"
|
|
encrypt = "some_encrypted_data"
|
|
parts = sorted([crypto.token, timestamp, nonce, encrypt])
|
|
expected = hashlib.sha1("".join(parts).encode()).hexdigest()
|
|
assert crypto.verify_signature(expected, timestamp, nonce, encrypt)
|
|
|
|
def test_generate_signature(self, crypto):
|
|
encrypt = "test_encrypted"
|
|
timestamp = "1609459200"
|
|
nonce = "abc"
|
|
sig = crypto.generate_signature(encrypt, timestamp, nonce)
|
|
parts = sorted([crypto.token, timestamp, nonce, encrypt])
|
|
expected = hashlib.sha1("".join(parts).encode()).hexdigest()
|
|
assert sig == expected
|
|
|
|
def test_wrap_encrypted_reply(self, crypto):
|
|
msg = "<xml><Content>Reply</Content></xml>"
|
|
xml_reply = crypto.wrap_encrypted_reply(msg)
|
|
assert "<Encrypt>" in xml_reply
|
|
assert "<MsgSignature>" in xml_reply
|
|
assert "<TimeStamp>" in xml_reply
|
|
assert "<Nonce>" in xml_reply
|
|
|
|
# Parse and verify the encrypted content decrypts back
|
|
root = ET.fromstring(xml_reply)
|
|
encrypt = root.find("Encrypt").text
|
|
decrypted, _app_id = crypto.decrypt(encrypt)
|
|
assert decrypted == msg
|
|
|
|
|
|
# ── Message processing tests ──────────────────────────────────────
|
|
|
|
|
|
class TestMessageProcessing:
|
|
"""Test the _process_message method with various XML payloads."""
|
|
|
|
def _make_channel(self):
|
|
config = WeComConfig(
|
|
corp_id="corp",
|
|
agent_id="1",
|
|
secret="s",
|
|
)
|
|
return WeChatChannel(config, backend="wecom")
|
|
|
|
async def test_text_message_queued(self):
|
|
channel = self._make_channel()
|
|
|
|
await channel._process_message(
|
|
{
|
|
"MsgType": "text",
|
|
"Content": "Hello!",
|
|
"FromUserName": "user1",
|
|
"ToUserName": "bot",
|
|
"MsgId": "100",
|
|
"CreateTime": str(int(time.time())),
|
|
}
|
|
)
|
|
# Check message was enqueued
|
|
assert not channel._queue.empty()
|
|
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
|
assert msg.content == "Hello!"
|
|
assert msg.sender_id == "user1"
|
|
assert msg.channel == "wechat"
|
|
|
|
async def test_location_message(self):
|
|
channel = self._make_channel()
|
|
|
|
await channel._process_message(
|
|
{
|
|
"MsgType": "location",
|
|
"Location_X": "39.9",
|
|
"Location_Y": "116.4",
|
|
"Label": "Beijing",
|
|
"FromUserName": "user1",
|
|
"ToUserName": "bot",
|
|
"MsgId": "101",
|
|
"CreateTime": str(int(time.time())),
|
|
}
|
|
)
|
|
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
|
assert "Beijing" in msg.content
|
|
assert "39.9" in msg.content
|
|
|
|
async def test_voice_recognition(self):
|
|
channel = self._make_channel()
|
|
|
|
await channel._process_message(
|
|
{
|
|
"MsgType": "voice",
|
|
"Recognition": "你好世界",
|
|
"FromUserName": "user1",
|
|
"ToUserName": "bot",
|
|
"MsgId": "102",
|
|
"CreateTime": str(int(time.time())),
|
|
}
|
|
)
|
|
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
|
assert "你好世界" in msg.content
|
|
|
|
async def test_link_message(self):
|
|
channel = self._make_channel()
|
|
|
|
await channel._process_message(
|
|
{
|
|
"MsgType": "link",
|
|
"Title": "Test Link",
|
|
"Description": "A description",
|
|
"Url": "https://example.com",
|
|
"FromUserName": "user1",
|
|
"ToUserName": "bot",
|
|
"MsgId": "103",
|
|
"CreateTime": str(int(time.time())),
|
|
}
|
|
)
|
|
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
|
assert "Test Link" in msg.content
|
|
assert "https://example.com" in msg.content
|
|
|
|
async def test_subscribe_event(self):
|
|
channel = self._make_channel()
|
|
|
|
await channel._process_message(
|
|
{
|
|
"MsgType": "event",
|
|
"Event": "subscribe",
|
|
"FromUserName": "user1",
|
|
"ToUserName": "bot",
|
|
"MsgId": "",
|
|
"CreateTime": str(int(time.time())),
|
|
}
|
|
)
|
|
msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0)
|
|
assert "关注" in msg.content
|
|
|
|
async def test_unsubscribe_ignored(self):
|
|
channel = self._make_channel()
|
|
|
|
await channel._process_message(
|
|
{
|
|
"MsgType": "event",
|
|
"Event": "unsubscribe",
|
|
"FromUserName": "user1",
|
|
"ToUserName": "bot",
|
|
"MsgId": "",
|
|
"CreateTime": str(int(time.time())),
|
|
}
|
|
)
|
|
assert channel._queue.empty()
|
|
|
|
async def test_empty_message_ignored(self):
|
|
channel = self._make_channel()
|
|
|
|
await channel._process_message(
|
|
{
|
|
"MsgType": "text",
|
|
"Content": "",
|
|
"FromUserName": "",
|
|
"ToUserName": "bot",
|
|
}
|
|
)
|
|
assert channel._queue.empty()
|
|
|
|
|
|
# ── Webhook signature bypass regression (issue #392) ──────────────
|
|
|
|
|
|
class _FakeWeChatRequest:
|
|
"""Minimal stand-in for aiohttp.web.Request for _handle_message tests."""
|
|
|
|
def __init__(self, text_body: str, query: dict | None = None):
|
|
self._text = text_body
|
|
self.query = query or {}
|
|
|
|
async def text(self) -> str:
|
|
return self._text
|
|
|
|
|
|
class TestWebhookSignatureBypass:
|
|
"""Regression tests for issue #392: when encryption is configured, an
|
|
unsigned POST must NOT reach the agent — verify before branching, not
|
|
inside the branch the request controls."""
|
|
|
|
PLAINTEXT_FORGED_XML = (
|
|
"<xml><MsgType><![CDATA[text]]></MsgType>"
|
|
"<Content><![CDATA[forged]]></Content>"
|
|
"<FromUserName><![CDATA[attacker]]></FromUserName></xml>"
|
|
)
|
|
|
|
def _make_channel_with_crypto(self) -> WeChatChannel:
|
|
"""Channel whose `_crypto` is set, mimicking what start() does when
|
|
encoding_aes_key + token are configured. We set _crypto directly to
|
|
avoid the network roundtrip in start()."""
|
|
config = WeComConfig(
|
|
corp_id="corp",
|
|
agent_id="1",
|
|
secret="s",
|
|
token="t",
|
|
encoding_aes_key="a" * 43,
|
|
)
|
|
channel = WeChatChannel(config, backend="wecom")
|
|
channel._crypto = MagicMock()
|
|
channel._crypto.verify_signature.return_value = (
|
|
False # default: signature won't match
|
|
)
|
|
channel._safe_process_message = AsyncMock() # type: ignore[assignment]
|
|
return channel
|
|
|
|
def _make_channel_without_crypto(self) -> WeChatChannel:
|
|
config = WeComConfig(corp_id="corp", agent_id="1", secret="s")
|
|
channel = WeChatChannel(config, backend="wecom")
|
|
# _crypto stays None (plaintext mode)
|
|
channel._safe_process_message = AsyncMock() # type: ignore[assignment]
|
|
return channel
|
|
|
|
async def test_plaintext_rejected_when_crypto_configured(self):
|
|
"""An unsigned POST on an encryption-configured channel must 403."""
|
|
channel = self._make_channel_with_crypto()
|
|
resp = await channel._handle_message(
|
|
_FakeWeChatRequest(self.PLAINTEXT_FORGED_XML)
|
|
)
|
|
assert resp.status == 403
|
|
channel._safe_process_message.assert_not_called()
|
|
|
|
async def test_missing_encrypt_rejected_even_with_crypto_present(self):
|
|
"""Even if the body has other XML, no <Encrypt> + crypto set → 403."""
|
|
channel = self._make_channel_with_crypto()
|
|
body = "<xml><MsgType><![CDATA[text]]></MsgType></xml>"
|
|
resp = await channel._handle_message(_FakeWeChatRequest(body))
|
|
assert resp.status == 403
|
|
channel._safe_process_message.assert_not_called()
|
|
|
|
async def test_invalid_signature_rejected(self):
|
|
"""Encrypted body with wrong signature → 403 (no behavior change)."""
|
|
channel = self._make_channel_with_crypto()
|
|
body = "<xml><Encrypt><![CDATA[encrypted-blob]]></Encrypt></xml>"
|
|
# crypto.verify_signature returns False by default in _make_channel_with_crypto
|
|
resp = await channel._handle_message(
|
|
_FakeWeChatRequest(
|
|
body,
|
|
query={
|
|
"msg_signature": "wrong",
|
|
"timestamp": "1",
|
|
"nonce": "n",
|
|
},
|
|
)
|
|
)
|
|
assert resp.status == 403
|
|
channel._safe_process_message.assert_not_called()
|
|
|
|
async def test_valid_signature_decrypts_and_processes(self):
|
|
"""Encrypted body with valid signature → 200, agent reached."""
|
|
channel = self._make_channel_with_crypto()
|
|
channel._crypto.verify_signature.return_value = True
|
|
channel._crypto.decrypt.return_value = (
|
|
"<xml><MsgType><![CDATA[text]]></MsgType>"
|
|
"<Content><![CDATA[legit]]></Content>"
|
|
"<FromUserName><![CDATA[user1]]></FromUserName></xml>",
|
|
"user1",
|
|
)
|
|
body = "<xml><Encrypt><![CDATA[ok]]></Encrypt></xml>"
|
|
resp = await channel._handle_message(
|
|
_FakeWeChatRequest(
|
|
body,
|
|
query={
|
|
"msg_signature": "right",
|
|
"timestamp": "1",
|
|
"nonce": "n",
|
|
},
|
|
)
|
|
)
|
|
assert resp.status == 200
|
|
channel._safe_process_message.assert_called_once()
|
|
|
|
async def test_plaintext_accepted_when_crypto_not_configured(self):
|
|
"""No-regression: plaintext mode (no crypto) keeps working."""
|
|
channel = self._make_channel_without_crypto()
|
|
resp = await channel._handle_message(
|
|
_FakeWeChatRequest(self.PLAINTEXT_FORGED_XML)
|
|
)
|
|
assert resp.status == 200
|
|
channel._safe_process_message.assert_called_once()
|
|
|
|
|
|
# ── Registration test ─────────────────────────────────────────────
|
|
|
|
|
|
class TestChannelRegistration:
|
|
def test_wechat_registered(self):
|
|
from EvoScientist.channels.channel_manager import available_channels
|
|
|
|
channels = available_channels()
|
|
assert "wechat" in channels
|