261 lines
9.5 KiB
Python
261 lines
9.5 KiB
Python
#!/usr/bin/env python
|
||
# -*- coding: utf-8 -*-
|
||
"""
|
||
钉钉 Stream 模式事件监听客户端
|
||
|
||
取代 HTTP 回调模式,通过 WebSocket 长连接接收钉钉通讯录变更事件。
|
||
无需公网回调地址,无需配置 callback_token/callback_aes_key/corp_id。
|
||
|
||
使用方式:
|
||
1. 在钉钉开放平台应用配置中,开启 Stream 模式推送
|
||
2. 后端启动后自动建立长连接,接收事件推送
|
||
|
||
Stream 模式事件数据格式与 HTTP 推送不同:
|
||
HTTP: { "EventType": "...", "UserId": [...], "DeptId": [...] }
|
||
Stream: event.headers.event_type 为事件类型,
|
||
event.data 为事件体 { "timeStamp": "...", "deptId": [...] / "userId": [...] }
|
||
"""
|
||
import asyncio
|
||
import logging
|
||
from collections import OrderedDict
|
||
from datetime import datetime
|
||
from typing import Optional
|
||
|
||
import dingtalk_stream
|
||
from dingtalk_stream import AckMessage
|
||
from dingtalk_stream import EventMessage
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
_DEDUP_MAX = 256
|
||
|
||
|
||
class _EventDedup:
|
||
"""基于 OrderedDict 的有界去重缓存,保留最近 N 个 event_id"""
|
||
|
||
def __init__(self, maxlen: int = _DEDUP_MAX):
|
||
self._seen: OrderedDict[str, None] = OrderedDict()
|
||
self._maxlen = maxlen
|
||
|
||
def is_duplicate(self, event_id: str) -> bool:
|
||
if event_id in self._seen:
|
||
self._seen.move_to_end(event_id)
|
||
return True
|
||
self._seen[event_id] = None
|
||
if len(self._seen) > self._maxlen:
|
||
self._seen.popitem(last=False)
|
||
return False
|
||
|
||
|
||
class DingtalkStreamEventHandler(dingtalk_stream.EventHandler):
|
||
"""钉钉 Stream 事件处理器 - 处理通讯录变更事件"""
|
||
|
||
_dedup = _EventDedup()
|
||
|
||
async def process(self, event: EventMessage) -> tuple:
|
||
try:
|
||
event_type = event.headers.event_type if event.headers else ''
|
||
event_data = event.data if isinstance(event.data, dict) else {}
|
||
|
||
event_id = getattr(event.headers, 'event_id', None) or ''
|
||
if event_id and self._dedup.is_duplicate(event_id):
|
||
logger.info("跳过重复事件: event_id=%s, type=%s", event_id, event_type)
|
||
return AckMessage.STATUS_OK, 'OK'
|
||
|
||
DingtalkStreamManager._record_event(event_type)
|
||
|
||
if not event_type:
|
||
logger.warning(f"收到未知事件: {event_data}")
|
||
return AckMessage.STATUS_OK, 'OK'
|
||
|
||
event_body = event_data
|
||
|
||
logger.info(f"收到钉钉 Stream 事件: {event_type}")
|
||
|
||
from core.dingtalk_sync.callback_handler import DingtalkCallbackHandler
|
||
from core.dingtalk_sync.callback_handler import DEPT_EVENTS, USER_EVENTS
|
||
|
||
from app.config_manager import config_manager
|
||
config = await config_manager.get_group("sync_dingtalk")
|
||
enable_dept = config.get("enable_dept_event") == "true"
|
||
enable_user = config.get("enable_user_event") == "true"
|
||
|
||
if event_type in DEPT_EVENTS and enable_dept:
|
||
normalized = self._normalize_dept_event(event_body)
|
||
await DingtalkCallbackHandler.handle_event(event_type, normalized)
|
||
|
||
elif event_type in USER_EVENTS and enable_user:
|
||
normalized = self._normalize_user_event(event_body)
|
||
await DingtalkCallbackHandler.handle_event(event_type, normalized)
|
||
|
||
else:
|
||
logger.info(f"忽略未启用的事件: {event_type}")
|
||
|
||
except Exception as e:
|
||
logger.error(f"处理 Stream 事件失败: {e}", exc_info=True)
|
||
|
||
return AckMessage.STATUS_OK, 'OK'
|
||
|
||
@staticmethod
|
||
def _normalize_dept_event(event_body: dict) -> dict:
|
||
"""将 Stream 模式的部门事件字段转为 HTTP 回调格式"""
|
||
dept_ids = event_body.get('deptId') or event_body.get('DeptId') or []
|
||
return {"DeptId": dept_ids}
|
||
|
||
@staticmethod
|
||
def _normalize_user_event(event_body: dict) -> dict:
|
||
"""将 Stream 模式的用户事件字段转为 HTTP 回调格式"""
|
||
user_ids = event_body.get('userId') or event_body.get('UserId') or []
|
||
return {"UserId": user_ids}
|
||
|
||
|
||
class DingtalkStreamManager:
|
||
"""钉钉 Stream 模式连接管理器"""
|
||
|
||
_client: Optional[dingtalk_stream.DingTalkStreamClient] = None
|
||
_task: Optional[asyncio.Task] = None
|
||
_connected: bool = False
|
||
_started_at: Optional[str] = None
|
||
|
||
# 事件统计
|
||
_event_stats = {
|
||
"total_events": 0,
|
||
"last_event_type": None,
|
||
"last_event_time": None,
|
||
}
|
||
|
||
@classmethod
|
||
def _record_event(cls, event_type: str):
|
||
cls._event_stats["total_events"] += 1
|
||
cls._event_stats["last_event_type"] = event_type
|
||
cls._event_stats["last_event_time"] = datetime.now().isoformat()
|
||
|
||
@classmethod
|
||
async def start(cls):
|
||
"""启动 Stream 连接,使用同步配置中的 app_key/app_secret"""
|
||
from app.config_manager import config_manager
|
||
|
||
config = await config_manager.get_group("sync_dingtalk")
|
||
app_key = config.get("app_key")
|
||
app_secret = config.get("app_secret")
|
||
|
||
if not app_key or not app_secret:
|
||
logger.warning("钉钉同步凭证未配置(app_key/app_secret),Stream 模式未启动")
|
||
return
|
||
|
||
credential = dingtalk_stream.Credential(app_key, app_secret)
|
||
client = dingtalk_stream.DingTalkStreamClient(credential)
|
||
|
||
handler = DingtalkStreamEventHandler()
|
||
client.register_all_event_handler(handler)
|
||
|
||
cls._client = client
|
||
cls._connected = True
|
||
cls._started_at = datetime.now().isoformat()
|
||
cls._task = asyncio.create_task(cls._run(client))
|
||
logger.info("钉钉 Stream 模式已启动(app_key=%s)", app_key[:6] + "***")
|
||
|
||
@classmethod
|
||
async def _run(cls, client: dingtalk_stream.DingTalkStreamClient):
|
||
"""
|
||
运行 Stream 客户端(带超时和重连)。
|
||
|
||
SDK 的 open_connection() 使用同步 requests(短暂阻塞可接受),
|
||
但 websockets.connect() 可能因为网络/防火墙问题长时间卡住,
|
||
因此对整个 start() 加超时保护,超时后自动重试。
|
||
"""
|
||
import json as _json
|
||
import websockets as _ws
|
||
from urllib.parse import quote_plus
|
||
|
||
client.pre_start()
|
||
while cls._connected:
|
||
try:
|
||
connection = await asyncio.get_event_loop().run_in_executor(
|
||
None, client.open_connection
|
||
)
|
||
if not connection:
|
||
logger.error("钉钉 Stream: open_connection 返回空")
|
||
await asyncio.sleep(10)
|
||
continue
|
||
|
||
logger.info("钉钉 Stream endpoint: %s", connection.get('endpoint', ''))
|
||
uri = f'{connection["endpoint"]}?ticket={quote_plus(connection["ticket"])}'
|
||
|
||
async with asyncio.timeout(30):
|
||
websocket = await _ws.connect(uri)
|
||
|
||
logger.info("钉钉 Stream WebSocket 连接成功")
|
||
client.websocket = websocket
|
||
|
||
asyncio.create_task(client.keepalive(websocket))
|
||
async for raw_message in websocket:
|
||
json_message = _json.loads(raw_message)
|
||
asyncio.create_task(client.background_task(json_message))
|
||
|
||
except asyncio.CancelledError:
|
||
logger.info("钉钉 Stream 任务已取消")
|
||
break
|
||
except TimeoutError:
|
||
logger.warning("钉钉 Stream WebSocket 连接超时(30s),将重试")
|
||
await asyncio.sleep(5)
|
||
continue
|
||
except (_ws.exceptions.ConnectionClosedError, ConnectionError, OSError) as e:
|
||
logger.warning(f"钉钉 Stream 连接断开: {e},10s 后重连")
|
||
await asyncio.sleep(10)
|
||
continue
|
||
except Exception as e:
|
||
logger.error(f"钉钉 Stream 异常: {e}", exc_info=True)
|
||
await asyncio.sleep(5)
|
||
continue
|
||
|
||
cls._connected = False
|
||
|
||
@classmethod
|
||
async def stop(cls):
|
||
"""停止 Stream 连接"""
|
||
cls._connected = False
|
||
if cls._task:
|
||
cls._task.cancel()
|
||
cls._task = None
|
||
if cls._client:
|
||
cls._client = None
|
||
cls._started_at = None
|
||
logger.info("钉钉 Stream 模式已停止")
|
||
|
||
@classmethod
|
||
def is_running(cls) -> bool:
|
||
"""
|
||
检查 Stream 是否正在运行。
|
||
dingtalk_stream SDK 的 start() 可能在连接建立后返回(task done),
|
||
但连接实际仍然活跃,因此用 _connected 标志而非仅依赖 task 状态。
|
||
"""
|
||
if cls._connected and cls._task is not None:
|
||
if cls._task.done():
|
||
ex = cls._task.exception() if not cls._task.cancelled() else None
|
||
if ex or cls._task.cancelled():
|
||
cls._connected = False
|
||
return False
|
||
return True
|
||
return True
|
||
return False
|
||
|
||
@classmethod
|
||
def status(cls) -> dict:
|
||
"""获取 Stream 模式状态"""
|
||
return {
|
||
"stream_mode": True,
|
||
"running": cls.is_running(),
|
||
"started_at": cls._started_at,
|
||
}
|
||
|
||
@classmethod
|
||
def get_event_stats(cls) -> dict:
|
||
"""获取 Stream 事件统计"""
|
||
return {
|
||
**cls.status(),
|
||
"total_events": cls._event_stats.get("total_events", 0),
|
||
"last_event_type": cls._event_stats.get("last_event_type"),
|
||
"last_event_time": cls._event_stats.get("last_event_time"),
|
||
}
|