diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..c888e25 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,28 @@ +name: CI + +on: + push: + branches: ["**"] + pull_request: + branches: ["**"] + +jobs: + test: + name: Python Test + runs-on: ubuntu-latest + strategy: + matrix: + python-version: ["3.8", "3.9", "3.10", "3.11", "3.12"] + steps: + - uses: actions/checkout@v4 + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + cache: "pip" + - name: Install dependencies + run: | + python -m pip install --upgrade pip + pip install -e ".[dev]" + - name: Test + run: pytest -v diff --git a/pyproject.toml b/pyproject.toml index b5e44df..c3da2ac 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -10,7 +10,9 @@ readme = "README.md" authors = [{name = "DingTalk Team"}] license = {text = "MIT"} requires-python = ">=3.8" -dependencies = [] +dependencies = [ + "websockets>=10.0", +] [project.optional-dependencies] redis = ["redis>=4.5.0"] diff --git a/src/dingtalk_channel_sdk/channel.py b/src/dingtalk_channel_sdk/channel.py index 4fb8b74..52abaa1 100644 --- a/src/dingtalk_channel_sdk/channel.py +++ b/src/dingtalk_channel_sdk/channel.py @@ -4,11 +4,16 @@ import asyncio import json +import os +import tempfile +import threading import time import urllib.request from datetime import timedelta from typing import Any, Awaitable, Callable, List, Optional +from .compat import to_thread + from .safety.batching import BatchConfig, BatchedMessage, MessageBatcher from .bot_identity import BotIdentity, BotIdentityProvider from .card import CardClient @@ -38,6 +43,16 @@ RejectHandler = Callable[[RejectEvent], Awaitable[None]] +class _SSRFSafeRedirectHandler(urllib.request.HTTPRedirectHandler): + def __init__(self, allowlist: Optional[List[str]] = None): + super().__init__() + self.allowlist = allowlist + + def redirect_request(self, req, fp, code, msg, headers, newurl): + assert_public_url(newurl, allowlist=self.allowlist) + return super().redirect_request(req, fp, code, msg, headers, newurl) + + class DingTalkChannel: """用法: @@ -215,10 +230,80 @@ async def download_file(self, url: str, timeout: float = 60.0) -> bytes: def _fetch() -> bytes: req = urllib.request.Request(url, headers={"User-Agent": USER_AGENT}) - with urllib.request.urlopen(req, timeout=timeout) as resp: - return resp.read() + opener = urllib.request.build_opener(_SSRFSafeRedirectHandler(allowlist=self.cfg.ssrf_allowlist)) + with opener.open(req, timeout=timeout) as resp: + data = resp.read() + cl = resp.headers.get("Content-Length") + if cl is not None: + try: + expected = int(cl) + if len(data) != expected: + raise OSError(f"download truncated: expected {expected} bytes, got {len(data)}") + except ValueError: + pass + return data + + return await to_thread(_fetch) + + async def download_file_to_file(self, url: str, dest_path: str, timeout: float = 60.0) -> int: + """流式下载文件到本地路径,不整块载入内存。 + + SSRF 防护同 download_file;父目录必须已存在;先写同目录临时文件再 + 原子重命名,失败不落半截文件。返回写入的字节数。 + """ + assert_public_url(url, allowlist=self.cfg.ssrf_allowlist) - return await asyncio.to_thread(_fetch) + def _fetch(cancel_event: threading.Event) -> int: + dest = os.path.abspath(dest_path) + parent = os.path.dirname(dest) + if not os.path.isdir(parent): + raise FileNotFoundError(f"parent directory does not exist: {parent}") + n = 0 + fd, tmp = tempfile.mkstemp(prefix="." + os.path.basename(dest) + ".tmp-", dir=parent) + tmp_open = True + try: + with os.fdopen(fd, "wb") as out: + tmp_open = False + req = urllib.request.Request(url, headers={"User-Agent": USER_AGENT}) + opener = urllib.request.build_opener(_SSRFSafeRedirectHandler(allowlist=self.cfg.ssrf_allowlist)) + with opener.open(req, timeout=timeout) as resp: + cl = resp.headers.get("Content-Length") + expected: Optional[int] = None + if cl is not None: + try: + expected = int(cl) + except ValueError: + expected = None + while True: + if cancel_event.is_set(): + raise RuntimeError("download cancelled") + chunk = resp.read(64 * 1024) + if not chunk: + break + out.write(chunk) + n += len(chunk) + if expected is not None and n != expected: + raise OSError(f"download truncated: expected {expected} bytes, got {n}") + if cancel_event.is_set(): + raise RuntimeError("download cancelled") + os.replace(tmp, dest) + tmp = None + return n + finally: + if tmp_open: + os.close(fd) + if tmp is not None: + try: + os.remove(tmp) + except OSError: + pass + + cancel_ev = threading.Event() + try: + return await to_thread(_fetch, cancel_ev) + except asyncio.CancelledError: + cancel_ev.set() + raise async def mark_thinking(self, conversation_id: str, msg_id: str) -> None: """在用户消息上打"🤔Thinking"状态章(仅人发的消息)。""" diff --git a/src/dingtalk_channel_sdk/compat.py b/src/dingtalk_channel_sdk/compat.py new file mode 100644 index 0000000..99481af --- /dev/null +++ b/src/dingtalk_channel_sdk/compat.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +import asyncio +import functools +import sys +from typing import Any, Callable, TypeVar + +T = TypeVar("T") + + +async def to_thread(func: Callable[..., T], *args: Any, **kwargs: Any) -> T: + """Python 3.8+ compatible asyncio.to_thread helper.""" + if sys.version_info >= (3, 9): + return await asyncio.to_thread(func, *args, **kwargs) + loop = asyncio.get_running_loop() + pfunc = functools.partial(func, *args, **kwargs) + return await loop.run_in_executor(None, pfunc) diff --git a/src/dingtalk_channel_sdk/httpx.py b/src/dingtalk_channel_sdk/httpx.py index 20f48a9..225cadb 100644 --- a/src/dingtalk_channel_sdk/httpx.py +++ b/src/dingtalk_channel_sdk/httpx.py @@ -7,6 +7,8 @@ import urllib.request from typing import Any, Dict, Optional +from .compat import to_thread + class ApiError(Exception): """钉钉 API 错误;is_qps_limit 判定 403 + code 含 QpsLimit(SPEC §6)。""" @@ -49,4 +51,4 @@ def _request_sync(method: str, url: str, headers: Dict[str, str], body: Optional async def http_json(method: str, url: str, headers: Optional[Dict[str, str]] = None, body: Optional[dict] = None) -> Dict[str, Any]: - return await asyncio.to_thread(_request_sync, method, url, headers or {}, body) + return await to_thread(_request_sync, method, url, headers or {}, body) diff --git a/src/dingtalk_channel_sdk/media.py b/src/dingtalk_channel_sdk/media.py index 1eee630..5f72ce8 100644 --- a/src/dingtalk_channel_sdk/media.py +++ b/src/dingtalk_channel_sdk/media.py @@ -8,6 +8,7 @@ import urllib.parse import urllib.request +from .compat import to_thread from .config import Config DEFAULT_OAPI_BASE = "https://oapi.dingtalk.com" @@ -36,7 +37,7 @@ async def upload_media(self, media_type: str, filename: str, data: bytes, conten media_type: image | file | video | voice """ - token = await asyncio.to_thread(self._get_token) + token = await to_thread(self._get_token) if not content_type: content_type = "image/jpeg" if media_type == "image" else "application/octet-stream" @@ -64,7 +65,7 @@ def _do_upload() -> dict: except urllib.error.HTTPError as e: raise RuntimeError(f"media/upload: http {e.code} {e.read().decode('utf-8', 'replace')}") from e - out = await asyncio.to_thread(_do_upload) + out = await to_thread(_do_upload) if out.get("errcode") not in (0, None): raise RuntimeError(f"media/upload: errcode={out.get('errcode')} {out.get('errmsg', '')}") media_id = out.get("media_id") or "" diff --git a/src/dingtalk_channel_sdk/normalize/converters/richtext.py b/src/dingtalk_channel_sdk/normalize/converters/richtext.py index 67f428f..d5bebd5 100644 --- a/src/dingtalk_channel_sdk/normalize/converters/richtext.py +++ b/src/dingtalk_channel_sdk/normalize/converters/richtext.py @@ -1,21 +1,49 @@ -"""富文本(richText)转换器:拼接正文与 @提及。""" +"""富文本(richText)转换器:拼接正文、@提及与内嵌媒体资源。""" from __future__ import annotations from typing import Any, Dict, List, Tuple -def convert_rich_text(content: Dict[str, Any]) -> Tuple[str, List[Dict[str, Any]]]: - """从 richText 数组提取拼接文本与 @提及(userId / 手机号)。""" +def convert_rich_text(content: Dict[str, Any]) -> Tuple[str, List[Dict[str, Any]], List[Dict[str, Any]]]: + """从 richText 数组提取拼接文本、@提及(userId / 手机号)与内嵌媒体资源。 + + 富文本附件区能力:picture/file 段提取为资源 + (兼容 downloadCode / pictureDownloadCode / picture)。脏数据防御:段值非字符串或下载码为空时跳过 + 该段,不影响其余段落;同一下载码在单条消息内去重。 + """ mentions: List[Dict[str, Any]] = [] + resources: List[Dict[str, Any]] = [] + seen: set = set() parts: List[str] = [] for item in (content or {}).get("richText", []): + if not isinstance(item, dict): + continue item_type = item.get("type", "") if item_type == "text": - parts.append(item.get("text", "")) + text = item.get("text", "") + if isinstance(text, str): + parts.append(text) elif item_type == "at": for uid in item.get("atUserIds") or []: mentions.append({"userId": uid}) for mob in item.get("atMobiles") or []: mentions.append({"userId": mob, "name": mob}) - return "".join(parts), mentions + elif item_type == "picture": + code = item.get("downloadCode") + if not isinstance(code, str) or not code: + code = item.get("pictureDownloadCode") + if not isinstance(code, str) or not code: + code = item.get("picture") + if isinstance(code, str) and code and code not in seen: + seen.add(code) + resources.append({"type": "image", "downloadCode": code}) + elif item_type == "file": + code = item.get("downloadCode") + name = item.get("fileName") + if isinstance(code, str) and code and code not in seen: + seen.add(code) + resources.append( + {"type": "file", "downloadCode": code, "fileName": name if isinstance(name, str) else ""} + ) + return "".join(parts), mentions, resources diff --git a/src/dingtalk_channel_sdk/normalize/message.py b/src/dingtalk_channel_sdk/normalize/message.py index 231ba2c..60db17d 100644 --- a/src/dingtalk_channel_sdk/normalize/message.py +++ b/src/dingtalk_channel_sdk/normalize/message.py @@ -67,7 +67,7 @@ def parse_content( if msg_type == "text": text = convert_text(content) elif msg_type == "richText": - text, mentions = convert_rich_text(content) + text, mentions, resources = convert_rich_text(content) elif msg_type == "picture": text, resources = convert_picture(content) elif msg_type == "file": diff --git a/src/dingtalk_channel_sdk/reply.py b/src/dingtalk_channel_sdk/reply.py index c514dd7..f675301 100644 --- a/src/dingtalk_channel_sdk/reply.py +++ b/src/dingtalk_channel_sdk/reply.py @@ -6,7 +6,6 @@ import time from typing import Optional -from urllib.parse import quote from .card import CardClient, CardStreamer @@ -166,11 +165,10 @@ async def download_url(self, download_code: str, msg_id: str) -> str: """换取消息附件下载地址(E9)。""" token = await self.tokens.get() out = await http_json( - "GET", - f"{self.cfg.api_base}/v1.0/robot/messageFiles/download" - f"?downloadCode={quote(download_code)}&messageId={quote(msg_id)}" - f"&robotCode={quote(self.cfg.client_id)}", + "POST", + f"{self.cfg.api_base}/v1.0/robot/messageFiles/download", {"x-acs-dingtalk-access-token": token}, + {"downloadCode": download_code, "robotCode": self.cfg.client_id}, ) return out.get("downloadUrl", "") diff --git a/tests/test_download_to_file.py b/tests/test_download_to_file.py new file mode 100644 index 0000000..0021694 --- /dev/null +++ b/tests/test_download_to_file.py @@ -0,0 +1,159 @@ +"""download_file_to_file 流式落盘单测。""" + +from __future__ import annotations + +import asyncio +import json +import threading +import time +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +import pytest + +from dingtalk_channel_sdk import DingTalkChannel + +MEDIA = b"dingtalk-media-bytes" * 512 + + +def _fake_server(): + class Handler(BaseHTTPRequestHandler): + def do_GET(self): + if self.path == "/v1.0/oauth2/accessToken": + self._json(200, {"accessToken": "tok-1", "expireIn": 7200}) + elif self.path == "/media.bin": + body = MEDIA + self.send_response(200) + self.send_header("Content-Type", "application/octet-stream") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + elif self.path == "/truncated.bin": + self.send_response(200) + self.send_header("Content-Type", "application/octet-stream") + self.send_header("Content-Length", "1000") + self.end_headers() + self.wfile.write(b"only 10 bytes") + elif self.path == "/slow.bin": + self.send_response(200) + self.send_header("Content-Type", "application/octet-stream") + self.send_header("Content-Length", str(len(MEDIA))) + self.end_headers() + time.sleep(0.3) + self.wfile.write(MEDIA) + else: + self._json(404, {}) + + def _json(self, code, obj): + body = json.dumps(obj).encode() + self.send_response(code) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, *args): + pass + + return ThreadingHTTPServer(("127.0.0.1", 0), Handler) + + +def _channel(base, allowlist=None): + return DingTalkChannel( + client_id="ding-test", + client_secret="s", + api_base=base, + ssrf_allowlist=allowlist if allowlist is not None else ["127.0.0.1"], + ) + + +async def test_download_file_to_file(tmp_path): + server = _fake_server() + threading.Thread(target=server.serve_forever, daemon=True).start() + try: + base = f"http://127.0.0.1:{server.server_address[1]}" + ch = _channel(base) + + dest = tmp_path / "media.bin" + n = await ch.download_file_to_file(f"{base}/media.bin", str(dest)) + assert n == len(MEDIA) + assert dest.read_bytes() == MEDIA + + # 原有内存下载语义保持不变 + assert await ch.download_file(f"{base}/media.bin") == MEDIA + finally: + server.shutdown() + + +async def test_download_file_to_file_missing_parent(tmp_path): + ch = DingTalkChannel(client_id="a", client_secret="b", ssrf_allowlist=["127.0.0.1"]) + dest = tmp_path / "no-such-dir" / "media.bin" + with pytest.raises(FileNotFoundError): + await ch.download_file_to_file("http://127.0.0.1:1/x", str(dest)) + assert not (tmp_path / "no-such-dir").exists() + + +async def test_download_file_to_file_ssrf_blocked(tmp_path): + ch = DingTalkChannel(client_id="a", client_secret="b") + with pytest.raises(Exception): + await ch.download_file_to_file("http://127.0.0.1:1/x", str(tmp_path / "x.bin")) + + +async def test_download_file_truncated_fails_cleanly(tmp_path): + server = _fake_server() + threading.Thread(target=server.serve_forever, daemon=True).start() + try: + base = f"http://127.0.0.1:{server.server_address[1]}" + ch = _channel(base) + dest = tmp_path / "truncated.bin" + with pytest.raises(OSError, match="download truncated"): + await ch.download_file_to_file(f"{base}/truncated.bin", str(dest)) + assert not dest.exists(), "截断文件不应落盘" + finally: + server.shutdown() + + +async def test_download_file_ssrf_redirect_bypass(tmp_path): + target = _fake_server() + threading.Thread(target=target.serve_forever, daemon=True).start() + + class RedirectHandler(BaseHTTPRequestHandler): + def do_GET(self): + target_port = target.server_address[1] + self.send_response(302) + self.send_header("Location", f"http://127.0.0.1:{target_port}/media.bin") + self.end_headers() + + def log_message(self, *args): + pass + + redirect_server = ThreadingHTTPServer(("127.0.0.1", 0), RedirectHandler) + threading.Thread(target=redirect_server.serve_forever, daemon=True).start() + + try: + redirect_base = f"http://127.0.0.1:{redirect_server.server_address[1]}" + ch = _channel(redirect_base, allowlist=[f"127.0.0.1:{redirect_server.server_address[1]}"]) + dest = tmp_path / "redirect.bin" + with pytest.raises(Exception): + await ch.download_file_to_file(f"{redirect_base}/redirect", str(dest)) + assert not dest.exists() + finally: + redirect_server.shutdown() + target.shutdown() + + +async def test_download_file_cancellation_does_not_overwrite(tmp_path): + server = _fake_server() + threading.Thread(target=server.serve_forever, daemon=True).start() + try: + base = f"http://127.0.0.1:{server.server_address[1]}" + ch = _channel(base) + dest = tmp_path / "slow.bin" + task = asyncio.create_task(ch.download_file_to_file(f"{base}/slow.bin", str(dest))) + await asyncio.sleep(0.05) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + await asyncio.sleep(0.4) + assert not dest.exists(), "已取消的任务不应覆盖目标文件" + finally: + server.shutdown() diff --git a/tests/test_normalize_converters.py b/tests/test_normalize_converters.py index 084a10a..62da11d 100644 --- a/tests/test_normalize_converters.py +++ b/tests/test_normalize_converters.py @@ -136,3 +136,45 @@ def test_convert_reply_quote_matrix(): def test_convert_reply_body_only(): assert parse_content("reply", {"text": "只有正文"}, [])[0] == "只有正文" assert parse_content("reply", {}, [])[0] == "[引用消息]" + + +# ── richText 附件资源 ── + + +def test_convert_rich_text_resources(): + content = { + "richText": [ + {"type": "text", "text": "图1 "}, + {"type": "picture", "downloadCode": "dc-1"}, + {"type": "picture", "pictureDownloadCode": "dc-1"}, + {"type": "picture", "picture": "dc-1"}, + {"type": "picture", "pictureDownloadCode": "dc-2"}, + {"type": "picture", "picture": "dc-3"}, + {"type": "file", "downloadCode": "dc-4", "fileName": "report.pdf"}, + {"type": "text", "text": " 图2"}, + ] + } + text, resources, _ = parse_content("richText", content, []) + assert text == "图1 图2" + assert resources == [ + {"type": "image", "downloadCode": "dc-1"}, + {"type": "image", "downloadCode": "dc-2"}, + {"type": "image", "downloadCode": "dc-3"}, + {"type": "file", "downloadCode": "dc-4", "fileName": "report.pdf"}, + ] + + +def test_convert_rich_text_resources_dirty_data(): + content = { + "richText": [ + {"type": "picture", "picture": 123}, + {"type": "picture", "picture": ""}, + {"type": "picture"}, + {"type": "file", "downloadCode": 42}, + {"type": "text", "text": "ok"}, + "junk-segment", + ] + } + text, resources, _ = parse_content("richText", content, []) + assert text == "ok" + assert resources == []