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()