这是一个批量整理提交,包含以下主要改动: 1. 删除多处冗余的空行和未使用的导入 2. 修复文件末尾缺少换行符的问题 3. 调整部分模块的导入顺序与代码排版 4. 修复部分配置默认值与策略逻辑 5. 新增多个功能模块与辅助工具 6. 完善异常处理与日志记录 7. 修复速率限制、消息缓存、权限校验等逻辑bug 8. 废弃部分旧有API与配置项并添加警告提示
143 lines
4.0 KiB
Python
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()
|