ForcePilot/backend/package/yuxi/channels/adapters/line/webhook.py
Kris 1f78c44b03 refactor: 整理并清理项目中的冗余代码与格式问题
这是一个批量整理提交,包含以下主要改动:
1.  删除多处冗余的空行和未使用的导入
2.  修复文件末尾缺少换行符的问题
3.  调整部分模块的导入顺序与代码排版
4.  修复部分配置默认值与策略逻辑
5.  新增多个功能模块与辅助工具
6.  完善异常处理与日志记录
7.  修复速率限制、消息缓存、权限校验等逻辑bug
8.  废弃部分旧有API与配置项并添加警告提示
2026-05-12 14:51:53 +08:00

143 lines
4.0 KiB
Python

from __future__ import annotations
import asyncio
import base64
import hashlib
import hmac
import json
import time
from collections import OrderedDict
from dataclasses import dataclass, field
LINE_SIGNATURE_HEADER = "x-line-signature"
LINE_WEBHOOK_MAX_BODY_BYTES = 64 * 1024
_REPLAY_WINDOW_SECONDS = 600
_REPLAY_CACHE_MAX_ENTRIES = 4096
class RetryableError(Exception):
pass
class ReplayDetectedError(RetryableError):
pass
class WebhookBodyTooLargeError(Exception):
pass
class WebhookParseError(Exception):
pass
class ConcurrentRequestRejectedError(Exception):
pass
def validate_line_signature(raw_body: bytes, signature: str, channel_secret: str) -> bool:
if not signature or not channel_secret:
return False
if len(raw_body) > LINE_WEBHOOK_MAX_BODY_BYTES:
return False
try:
computed = hmac.new(
key=channel_secret.encode("utf-8"),
msg=raw_body if isinstance(raw_body, bytes) else raw_body.encode("utf-8"),
digestmod=hashlib.sha256,
).digest()
computed_b64 = base64.b64encode(computed).decode("utf-8")
return hmac.compare_digest(computed_b64, signature)
except Exception:
return False
def parse_webhook_body(raw_body: bytes) -> list[dict] | None:
try:
body_str = raw_body.decode("utf-8") if isinstance(raw_body, bytes) else raw_body
data = json.loads(body_str)
return data.get("events", [])
except (json.JSONDecodeError, UnicodeDecodeError):
return None
@dataclass
class WebhookReplayGuard:
_window_seconds: float = _REPLAY_WINDOW_SECONDS
_max_entries: int = _REPLAY_CACHE_MAX_ENTRIES
_seen_hashes: OrderedDict = field(default_factory=OrderedDict)
def check_and_claim(self, signature_hash: str) -> None:
now = time.time()
cutoff = now - self._window_seconds
expired = [k for k, ts in self._seen_hashes.items() if ts < cutoff]
for k in expired:
self._seen_hashes.pop(k, None)
if signature_hash in self._seen_hashes:
raise ReplayDetectedError(f"Replay detected for signature {signature_hash[:16]}...")
self._seen_hashes[signature_hash] = now
while len(self._seen_hashes) > self._max_entries:
self._seen_hashes.popitem(last=False)
def clear(self) -> None:
self._seen_hashes.clear()
@dataclass
class MultiAccountSignatureRouter:
_secrets: OrderedDict = field(default_factory=OrderedDict)
def register_account(self, account_id: str, channel_secret: str) -> None:
self._secrets[account_id] = channel_secret
def unregister_account(self, account_id: str) -> None:
self._secrets.pop(account_id, None)
def list_accounts(self) -> list[str]:
return list(self._secrets.keys())
def match_signature(self, raw_body: bytes, signature: str) -> str | None:
for account_id, secret in self._secrets.items():
if validate_line_signature(raw_body, signature, secret):
return account_id
return None
def clear(self) -> None:
self._secrets.clear()
@dataclass
class WebhookConcurrencyGuard:
_max_concurrent: int = 1
_semaphores: dict[str, asyncio.Semaphore] = field(default_factory=dict)
async def acquire(self, path: str) -> bool:
if path not in self._semaphores:
self._semaphores[path] = asyncio.Semaphore(self._max_concurrent)
return await self._semaphores[path].acquire()
def release(self, path: str) -> None:
if path in self._semaphores:
try:
self._semaphores[path].release()
except ValueError:
pass
def in_flight_count(self, path: str | None = None) -> int:
if path:
sem = self._semaphores.get(path)
if sem is None:
return 0
return self._max_concurrent - sem._value
return sum(self._max_concurrent - sem._value for sem in self._semaphores.values())
def clear(self) -> None:
self._semaphores.clear()