Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
@@ -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
4 changes: 3 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down
91 changes: 88 additions & 3 deletions src/dingtalk_channel_sdk/channel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
"""用法:

Expand Down Expand Up @@ -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"状态章(仅人发的消息)。"""
Expand Down
17 changes: 17 additions & 0 deletions src/dingtalk_channel_sdk/compat.py
Original file line number Diff line number Diff line change
@@ -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)
4 changes: 3 additions & 1 deletion src/dingtalk_channel_sdk/httpx.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)。"""
Expand Down Expand Up @@ -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)
5 changes: 3 additions & 2 deletions src/dingtalk_channel_sdk/media.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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"

Expand Down Expand Up @@ -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 ""
Expand Down
38 changes: 33 additions & 5 deletions src/dingtalk_channel_sdk/normalize/converters/richtext.py
Original file line number Diff line number Diff line change
@@ -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
2 changes: 1 addition & 1 deletion src/dingtalk_channel_sdk/normalize/message.py
Original file line number Diff line number Diff line change
Expand Up @@ -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":
Expand Down
8 changes: 3 additions & 5 deletions src/dingtalk_channel_sdk/reply.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
import time

from typing import Optional
from urllib.parse import quote

from .card import CardClient, CardStreamer

Expand Down Expand Up @@ -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", "")

Expand Down
Loading