From 38fce78425c1527205f848b3079fe3597a7f6128 Mon Sep 17 00:00:00 2001 From: rookiestar28 <151893693+rookiestar28@users.noreply.github.com> Date: Thu, 12 Feb 2026 10:58:32 +0800 Subject: [PATCH] feat(connector-security): enforce strict ack-by-default callback lifecycle and public-safe token serialization in R75 transport contract --- connector/transport_contract.py | 710 +++++++++++++++++++++++++++ tests/test_r75_transport_contract.py | 663 +++++++++++++++++++++++++ 2 files changed, 1373 insertions(+) create mode 100644 connector/transport_contract.py create mode 100644 tests/test_r75_transport_contract.py diff --git a/connector/transport_contract.py b/connector/transport_contract.py new file mode 100644 index 0000000..d9904a6 --- /dev/null +++ b/connector/transport_contract.py @@ -0,0 +1,710 @@ +""" +Shared Transport Contract (R75). + +Platform-agnostic primitives for connector session lifecycle, +event streaming, callback delivery, and token management. + +These contracts define deterministic behavior for retries, dedupe, +timeout budgets, event ordering, and fail-closed auth — independent +of any specific chat platform (Telegram, Discord, LINE, WhatsApp, Kakao, WeChat). +""" + +from __future__ import annotations + +import enum +import hashlib +import logging +import time +import uuid +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional, Sequence + +logger = logging.getLogger("connector.transport_contract") + +# --------------------------------------------------------------------------- +# WP2 — Session Contract +# --------------------------------------------------------------------------- + + +class SessionState(str, enum.Enum): + """Valid session states with deterministic transitions.""" + + PENDING = "pending" + ACTIVE = "active" + EXPIRED = "expired" + REVOKED = "revoked" + + @classmethod + def valid_transitions(cls) -> Dict["SessionState", List["SessionState"]]: + return { + cls.PENDING: [cls.ACTIVE, cls.EXPIRED, cls.REVOKED], + cls.ACTIVE: [cls.EXPIRED, cls.REVOKED], + cls.EXPIRED: [], # terminal + cls.REVOKED: [], # terminal + } + + +@dataclass +class SessionInfo: + """Session metadata for a connector connection.""" + + session_id: str = field(default_factory=lambda: uuid.uuid4().hex[:16]) + platform: str = "" + state: str = SessionState.PENDING.value + created_at: float = field(default_factory=time.time) + activated_at: Optional[float] = None + expired_at: Optional[float] = None + revoked_at: Optional[float] = None + ttl_sec: int = 86400 # 24h default + metadata: Dict[str, Any] = field(default_factory=dict) + + +class SessionContract: + """ + Manages session lifecycle with explicit state transitions. + + States: pending -> active -> expired / revoked + Terminal states: expired, revoked (no transitions out). + """ + + TERMINAL_STATES = frozenset( + {SessionState.EXPIRED.value, SessionState.REVOKED.value} + ) + + def __init__(self, default_ttl_sec: int = 86400): + self._sessions: Dict[str, SessionInfo] = {} + self._default_ttl = default_ttl_sec + + def create( + self, + platform: str, + *, + ttl_sec: Optional[int] = None, + metadata: Optional[Dict[str, Any]] = None, + ) -> SessionInfo: + """Create a new session in PENDING state.""" + session = SessionInfo( + platform=platform, + ttl_sec=ttl_sec or self._default_ttl, + metadata=metadata or {}, + ) + self._sessions[session.session_id] = session + return session + + def activate(self, session_id: str) -> SessionInfo: + """Transition: PENDING -> ACTIVE.""" + session = self._get_or_raise(session_id) + self._transition(session, SessionState.ACTIVE) + session.activated_at = time.time() + return session + + def expire(self, session_id: str) -> SessionInfo: + """Transition: PENDING|ACTIVE -> EXPIRED.""" + session = self._get_or_raise(session_id) + self._transition(session, SessionState.EXPIRED) + session.expired_at = time.time() + return session + + def revoke(self, session_id: str) -> SessionInfo: + """Transition: PENDING|ACTIVE -> REVOKED.""" + session = self._get_or_raise(session_id) + self._transition(session, SessionState.REVOKED) + session.revoked_at = time.time() + return session + + def get(self, session_id: str) -> Optional[SessionInfo]: + """Get session by ID, auto-expiring if TTL exceeded.""" + session = self._sessions.get(session_id) + if session and session.state not in self.TERMINAL_STATES: + if time.time() - session.created_at > session.ttl_sec: + self._transition(session, SessionState.EXPIRED) + session.expired_at = time.time() + return session + + def is_active(self, session_id: str) -> bool: + s = self.get(session_id) + return s is not None and s.state == SessionState.ACTIVE.value + + def _get_or_raise(self, session_id: str) -> SessionInfo: + session = self._sessions.get(session_id) + if session is None: + raise SessionError(f"Session not found: {session_id}") + return session + + def _transition(self, session: SessionInfo, target: SessionState) -> None: + current = SessionState(session.state) + allowed = SessionState.valid_transitions()[current] + if target not in allowed: + raise SessionError( + f"Invalid transition: {current.value} -> {target.value} " + f"(allowed: {[s.value for s in allowed]})" + ) + session.state = target.value + + +class SessionError(Exception): + """Raised for invalid session operations.""" + + +# --------------------------------------------------------------------------- +# WP2 — Event Stream Contract (SSE reconnect/resume) +# --------------------------------------------------------------------------- + + +@dataclass +class StreamEvent: + """Normalized event in an event stream.""" + + event_id: str = field(default_factory=lambda: uuid.uuid4().hex[:12]) + event_type: str = "message" + data: Dict[str, Any] = field(default_factory=dict) + timestamp: float = field(default_factory=time.time) + sequence: int = 0 + + +class ReconnectPolicy: + """Bounded reconnect/backoff policy for event streams.""" + + DEFAULT_INITIAL_DELAY_MS = 1000 + DEFAULT_MAX_DELAY_MS = 30000 + DEFAULT_JITTER_MS = 500 + DEFAULT_MAX_RETRIES = 10 + + def __init__( + self, + initial_delay_ms: int = DEFAULT_INITIAL_DELAY_MS, + max_delay_ms: int = DEFAULT_MAX_DELAY_MS, + jitter_ms: int = DEFAULT_JITTER_MS, + max_retries: int = DEFAULT_MAX_RETRIES, + ): + self.initial_delay_ms = initial_delay_ms + self.max_delay_ms = max_delay_ms + self.jitter_ms = jitter_ms + self.max_retries = max_retries + + def compute_delay_ms(self, attempt: int) -> int: + """Exponential backoff with jitter, capped at max_delay.""" + import random + + if attempt >= self.max_retries: + return -1 # signal: stop retrying + base = min(self.initial_delay_ms * (2**attempt), self.max_delay_ms) + jitter = random.randint(0, self.jitter_ms) + return base + jitter + + def should_retry(self, attempt: int) -> bool: + return attempt < self.max_retries + + +class EventStreamContract: + """ + Manages bounded event buffering with replay/resume support. + + - Bounded retention (max events in buffer). + - Resume from Last-Event-ID equivalent. + - Sequence-ordered delivery. + """ + + DEFAULT_MAX_BUFFER = 500 + + def __init__( + self, + max_buffer: int = DEFAULT_MAX_BUFFER, + reconnect_policy: Optional[ReconnectPolicy] = None, + ): + self._buffer: List[StreamEvent] = [] + self._max_buffer = max_buffer + self._sequence_counter = 0 + self.reconnect_policy = reconnect_policy or ReconnectPolicy() + + def emit(self, event_type: str, data: Dict[str, Any]) -> StreamEvent: + """Emit a new event into the stream buffer.""" + self._sequence_counter += 1 + event = StreamEvent( + event_type=event_type, + data=data, + sequence=self._sequence_counter, + ) + self._buffer.append(event) + # Evict oldest if over capacity + if len(self._buffer) > self._max_buffer: + self._buffer = self._buffer[-self._max_buffer :] + return event + + def replay_from(self, last_event_id: str) -> List[StreamEvent]: + """ + Return events after the given last_event_id. + If ID not found in buffer, return all buffered events. + """ + idx = None + for i, evt in enumerate(self._buffer): + if evt.event_id == last_event_id: + idx = i + break + if idx is not None: + return list(self._buffer[idx + 1 :]) + # ID not in buffer (gap): return everything + return list(self._buffer) + + def get_all(self) -> List[StreamEvent]: + return list(self._buffer) + + @property + def latest_sequence(self) -> int: + return self._sequence_counter + + +# --------------------------------------------------------------------------- +# WP3 — Callback Contract (ack, deferred, idempotency) +# --------------------------------------------------------------------------- + + +class CallbackState(str, enum.Enum): + """Callback delivery states.""" + + PENDING = "pending" + ACKNOWLEDGED = "acknowledged" + DELIVERED = "delivered" + EXPIRED = "expired" + FAILED = "failed" + + +@dataclass +class CallbackRecord: + """Tracks a callback delivery attempt.""" + + callback_id: str = field(default_factory=lambda: uuid.uuid4().hex[:16]) + idempotency_key: str = "" + state: str = CallbackState.PENDING.value + created_at: float = field(default_factory=time.time) + acknowledged_at: Optional[float] = None + delivered_at: Optional[float] = None + expired_at: Optional[float] = None + ttl_sec: int = 300 # 5min default expiry + attempts: int = 0 + max_attempts: int = 3 + # Default strict contract: must ack before deliver unless opt-out is explicit. + require_ack: bool = True + allow_direct_delivery: bool = False + payload_hash: str = "" + metadata: Dict[str, Any] = field(default_factory=dict) + + +class CallbackContract: + """ + Manages callback delivery lifecycle with idempotency. + + Two modes controlled by ``require_ack`` and ``allow_direct_delivery``: + + - **Strict mode (default)**: ``require_ack=True`` and + ``allow_direct_delivery=False``. ``deliver()`` rejects pending + callbacks until ``acknowledge()`` succeeds. + - **Compatibility mode (explicit opt-out)**: set + ``allow_direct_delivery=True`` (or ``require_ack=False``) to allow + direct delivery from pending state. + + Common behaviour across both modes: + - Idempotency key based dedupe. + - Single-use delivery (delivered -> cannot deliver again). + - Bounded max attempts. + """ + + DEFAULT_ACK_WINDOW_SEC = 3 + DEFAULT_CALLBACK_TTL_SEC = 300 + DEFAULT_MAX_ATTEMPTS = 3 + + def __init__( + self, + ack_window_sec: int = DEFAULT_ACK_WINDOW_SEC, + callback_ttl_sec: int = DEFAULT_CALLBACK_TTL_SEC, + max_attempts: int = DEFAULT_MAX_ATTEMPTS, + ): + self._records: Dict[str, CallbackRecord] = {} + self._idempotency_index: Dict[str, str] = {} # key -> callback_id + self._ack_window_sec = ack_window_sec + self._callback_ttl_sec = callback_ttl_sec + self._max_attempts = max_attempts + + def create( + self, + *, + idempotency_key: str = "", + payload: Optional[Dict[str, Any]] = None, + ttl_sec: Optional[int] = None, + require_ack: Optional[bool] = None, + allow_direct_delivery: bool = False, + metadata: Optional[Dict[str, Any]] = None, + ) -> CallbackRecord: + """ + Create a callback record. If idempotency_key matches an existing + non-terminal record, returns the existing one (dedupe). + + Args: + require_ack: Optional explicit ack requirement override. + None means "derive from allow_direct_delivery". + allow_direct_delivery: Explicit compatibility opt-out. + When True, direct deliver from pending is permitted. + """ + effective_require_ack, effective_allow_direct = self._resolve_ack_policy( + require_ack=require_ack, + allow_direct_delivery=allow_direct_delivery, + ) + + # Dedupe check + if idempotency_key: + existing_id = self._idempotency_index.get(idempotency_key) + if existing_id and existing_id in self._records: + existing = self._records[existing_id] + if existing.state not in ( + CallbackState.EXPIRED.value, + CallbackState.FAILED.value, + ): + return existing + + record = CallbackRecord( + idempotency_key=idempotency_key, + ttl_sec=ttl_sec or self._callback_ttl_sec, + max_attempts=self._max_attempts, + require_ack=effective_require_ack, + allow_direct_delivery=effective_allow_direct, + payload_hash=self._hash_payload(payload) if payload else "", + metadata=metadata or {}, + ) + self._records[record.callback_id] = record + if idempotency_key: + self._idempotency_index[idempotency_key] = record.callback_id + return record + + def acknowledge(self, callback_id: str) -> CallbackRecord: + """Mark callback as acknowledged within ack window.""" + record = self._get_or_raise(callback_id) + if record.state == CallbackState.PENDING.value: + elapsed = time.time() - record.created_at + if elapsed > self._ack_window_sec: + record.state = CallbackState.EXPIRED.value + record.expired_at = time.time() + raise CallbackError( + f"Ack window expired ({elapsed:.1f}s > {self._ack_window_sec}s)" + ) + self._expire_if_needed(record) + if record.state != CallbackState.PENDING.value: + raise CallbackError( + f"Cannot ack callback in state '{record.state}' (must be pending)" + ) + record.state = CallbackState.ACKNOWLEDGED.value + record.acknowledged_at = time.time() + return record + + def deliver(self, callback_id: str) -> CallbackRecord: + """ + Mark callback as delivered (final, single-use). + + In strict mode (``require_ack=True``), rejects delivery of + unacknowledged callbacks. + """ + record = self._get_or_raise(callback_id) + self._expire_if_needed(record) + + # Strict mode: enforce ack-before-deliver + if record.require_ack and record.state == CallbackState.PENDING.value: + raise CallbackError( + f"Cannot deliver: require_ack=True but callback is still " + f"pending (must acknowledge first)" + ) + + if record.state not in ( + CallbackState.PENDING.value, + CallbackState.ACKNOWLEDGED.value, + ): + raise CallbackError(f"Cannot deliver callback in state '{record.state}'") + record.state = CallbackState.DELIVERED.value + record.delivered_at = time.time() + record.attempts += 1 + return record + + def record_attempt(self, callback_id: str) -> CallbackRecord: + """Record a delivery attempt; fail if max attempts exceeded.""" + record = self._get_or_raise(callback_id) + record.attempts += 1 + if record.attempts >= record.max_attempts: + record.state = CallbackState.FAILED.value + return record + + def get(self, callback_id: str) -> Optional[CallbackRecord]: + record = self._records.get(callback_id) + if record: + self._expire_if_needed(record) + return record + + def get_by_idempotency_key(self, key: str) -> Optional[CallbackRecord]: + cb_id = self._idempotency_index.get(key) + if cb_id: + return self.get(cb_id) + return None + + def _get_or_raise(self, callback_id: str) -> CallbackRecord: + record = self._records.get(callback_id) + if record is None: + raise CallbackError(f"Callback not found: {callback_id}") + return record + + def _expire_if_needed(self, record: CallbackRecord) -> None: + if record.state in ( + CallbackState.DELIVERED.value, + CallbackState.EXPIRED.value, + CallbackState.FAILED.value, + ): + return + elapsed = time.time() - record.created_at + if ( + record.require_ack + and record.state == CallbackState.PENDING.value + and elapsed > self._ack_window_sec + ): + record.state = CallbackState.EXPIRED.value + record.expired_at = time.time() + return + if elapsed > record.ttl_sec: + record.state = CallbackState.EXPIRED.value + record.expired_at = time.time() + + @staticmethod + def _hash_payload(payload: Dict[str, Any]) -> str: + import json + + raw = json.dumps(payload, sort_keys=True, default=str) + return hashlib.sha256(raw.encode()).hexdigest()[:16] + + @staticmethod + def _resolve_ack_policy( + *, + require_ack: Optional[bool], + allow_direct_delivery: bool, + ) -> tuple[bool, bool]: + if require_ack is None: + # Default strict, explicit opt-out via allow_direct_delivery=True. + return (not allow_direct_delivery, allow_direct_delivery) + if require_ack and allow_direct_delivery: + raise CallbackError( + "Invalid callback policy: require_ack=True conflicts " + "with allow_direct_delivery=True" + ) + if not require_ack: + # Explicit legacy/permissive request always enables direct delivery. + return (False, True) + return (True, False) + + +class CallbackError(Exception): + """Raised for invalid callback operations.""" + + +# --------------------------------------------------------------------------- +# WP4 — Token Contract (source precedence, fail-closed, redaction) +# --------------------------------------------------------------------------- + + +class TokenValidity(str, enum.Enum): + """Token validation result.""" + + VALID = "valid" + MISSING = "missing" + INVALID = "invalid" + EXPIRED = "expired" + + +@dataclass +class TokenSource: + """A token source with explicit precedence.""" + + name: str + env_var: str + precedence: int # lower = higher priority + required: bool = True + redaction_pattern: str = "***" # how to mask in logs + + +@dataclass +class PublicTokenResult: + """Public-safe token resolution view (no raw token field).""" + + validity: str = TokenValidity.MISSING.value + source_name: str = "" + masked_value: str = "" + + def to_dict(self) -> Dict[str, Any]: + return { + "validity": self.validity, + "source_name": self.source_name, + "masked_value": self.masked_value, + } + + +@dataclass +class TokenResult: + """Result of token resolution.""" + + validity: str = TokenValidity.MISSING.value + source_name: str = "" + masked_value: str = "" + _raw_value: str = field(default="", repr=False) + + @property + def raw_value(self) -> str: + """Internal-use access to the resolved token value.""" + return self._raw_value + + def to_public(self) -> PublicTokenResult: + return PublicTokenResult( + validity=self.validity, + source_name=self.source_name, + masked_value=self.masked_value, + ) + + def to_public_dict(self) -> Dict[str, Any]: + """Serialize for logs/API/audit — hard-excludes raw token.""" + return self.to_public().to_dict() + + def to_dict(self) -> Dict[str, Any]: + """Backward-compatible alias of to_public_dict().""" + return self.to_public_dict() + + +class TokenContract: + """ + Manages token source precedence and fail-closed behavior. + + - Explicit source precedence (env var priority order). + - Fail-closed: missing/invalid required token -> reject. + - Redaction: tokens are always masked in logs/errors/audit. + """ + + def __init__(self, sources: Sequence[TokenSource]): + # Sort by precedence (lower = higher priority) + self._sources = sorted(sources, key=lambda s: s.precedence) + + def resolve(self, env: Optional[Dict[str, str]] = None) -> TokenResult: + """ + Resolve token from configured sources in precedence order. + Returns the first non-empty token found. + """ + import os + + lookup = env if env is not None else os.environ + + for source in self._sources: + value = lookup.get(source.env_var, "").strip() + if value: + return TokenResult( + validity=TokenValidity.VALID.value, + source_name=source.name, + masked_value=self._mask(value, source.redaction_pattern), + _raw_value=value, + ) + + # No token found + return TokenResult( + validity=TokenValidity.MISSING.value, + source_name="", + masked_value="", + _raw_value="", + ) + + def validate_or_reject(self, env: Optional[Dict[str, str]] = None) -> TokenResult: + """ + Resolve and validate token. Raises if required token is missing. + Fail-closed behavior: no token = reject. + """ + result = self.resolve(env) + if result.validity == TokenValidity.MISSING.value: + required_sources = [s for s in self._sources if s.required] + if required_sources: + raise TokenError( + f"Required token missing. Checked sources: " + f"{[s.env_var for s in required_sources]}. " + f"Fail-closed: request rejected." + ) + return result + + def get_precedence_table(self) -> List[Dict[str, Any]]: + """Return the precedence table for documentation/diagnostics.""" + return [ + { + "precedence": s.precedence, + "name": s.name, + "env_var": s.env_var, + "required": s.required, + } + for s in self._sources + ] + + @staticmethod + def _mask(value: str, pattern: str = "***") -> str: + """Mask a token value, showing first 4 and last 2 chars if long enough.""" + if len(value) <= 8: + return pattern + return f"{value[:4]}{pattern}{value[-2:]}" + + +class TokenError(Exception): + """Raised when required token is missing or invalid (fail-closed).""" + + +# --------------------------------------------------------------------------- +# WP3 — Deterministic Retry Policy +# --------------------------------------------------------------------------- + + +@dataclass +class RetryPolicy: + """Deterministic retry policy for callback/webhook delivery.""" + + max_retries: int = 3 + initial_delay_sec: float = 1.0 + max_delay_sec: float = 30.0 + backoff_factor: float = 2.0 + retry_on_status: frozenset = field( + default_factory=lambda: frozenset({408, 429, 500, 502, 503, 504}) + ) + + def compute_delay(self, attempt: int) -> float: + """Compute delay for the given attempt number (0-indexed).""" + if attempt >= self.max_retries: + return -1.0 + delay = min( + self.initial_delay_sec * (self.backoff_factor**attempt), + self.max_delay_sec, + ) + return delay + + def should_retry(self, attempt: int, status_code: Optional[int] = None) -> bool: + if attempt >= self.max_retries: + return False + if status_code is not None and status_code not in self.retry_on_status: + return False + return True + + +# --------------------------------------------------------------------------- +# Error Envelope (normalized across all transports) +# --------------------------------------------------------------------------- + + +@dataclass +class TransportError: + """Normalized error envelope for connector transport.""" + + code: str # machine-readable code, e.g. "session_expired", "token_missing" + message: str + retryable: bool = False + details: Dict[str, Any] = field(default_factory=dict) + timestamp: float = field(default_factory=time.time) + + def to_dict(self) -> Dict[str, Any]: + return { + "code": self.code, + "message": self.message, + "retryable": self.retryable, + "details": self.details, + "timestamp": self.timestamp, + } diff --git a/tests/test_r75_transport_contract.py b/tests/test_r75_transport_contract.py new file mode 100644 index 0000000..c7f8b99 --- /dev/null +++ b/tests/test_r75_transport_contract.py @@ -0,0 +1,663 @@ +""" +R75 Shared Transport Contract — Contract Tests. + +Tests cover: +- WP2: Session lifecycle + event stream (replay, reconnect policy) +- WP3: Callback contract (ack, delivery, idempotency, expiry, retry) +- WP4: Token precedence + fail-closed +- Regression: existing connector flows unaffected +""" + +import time +import unittest +from unittest.mock import patch + +from connector.transport_contract import ( + CallbackContract, + CallbackError, + CallbackRecord, + CallbackState, + EventStreamContract, + ReconnectPolicy, + RetryPolicy, + SessionContract, + SessionError, + SessionInfo, + SessionState, + StreamEvent, + TokenContract, + TokenError, + TokenResult, + TokenSource, + TokenValidity, + TransportError, +) + +# ========================================================================= +# WP2 — Session Contract Tests +# ========================================================================= + + +class TestSessionContract(unittest.TestCase): + """Session lifecycle state transitions.""" + + def test_create_session_pending(self): + sc = SessionContract() + session = sc.create("telegram") + self.assertEqual(session.state, SessionState.PENDING.value) + self.assertEqual(session.platform, "telegram") + self.assertIsNotNone(session.session_id) + + def test_activate_from_pending(self): + sc = SessionContract() + session = sc.create("discord") + activated = sc.activate(session.session_id) + self.assertEqual(activated.state, SessionState.ACTIVE.value) + self.assertIsNotNone(activated.activated_at) + + def test_expire_from_pending(self): + sc = SessionContract() + session = sc.create("line") + expired = sc.expire(session.session_id) + self.assertEqual(expired.state, SessionState.EXPIRED.value) + + def test_expire_from_active(self): + sc = SessionContract() + session = sc.create("kakao") + sc.activate(session.session_id) + expired = sc.expire(session.session_id) + self.assertEqual(expired.state, SessionState.EXPIRED.value) + + def test_revoke_from_active(self): + sc = SessionContract() + session = sc.create("whatsapp") + sc.activate(session.session_id) + revoked = sc.revoke(session.session_id) + self.assertEqual(revoked.state, SessionState.REVOKED.value) + + def test_invalid_transition_expired_to_active(self): + sc = SessionContract() + session = sc.create("test") + sc.expire(session.session_id) + with self.assertRaises(SessionError) as ctx: + sc.activate(session.session_id) + self.assertIn("Invalid transition", str(ctx.exception)) + + def test_invalid_transition_revoked(self): + sc = SessionContract() + session = sc.create("test") + sc.revoke(session.session_id) + with self.assertRaises(SessionError): + sc.expire(session.session_id) + + def test_double_activate_rejected(self): + """Cannot transition active -> active.""" + sc = SessionContract() + session = sc.create("test") + sc.activate(session.session_id) + with self.assertRaises(SessionError): + sc.activate(session.session_id) + + def test_session_not_found(self): + sc = SessionContract() + with self.assertRaises(SessionError): + sc.activate("nonexistent") + + def test_auto_expire_on_ttl(self): + """Session auto-expires when TTL is exceeded.""" + sc = SessionContract() + session = sc.create("test", ttl_sec=10) + sc.activate(session.session_id) + # Force created_at far enough in the past to exceed TTL + session.created_at = time.time() - 20 + result = sc.get(session.session_id) + self.assertEqual(result.state, SessionState.EXPIRED.value) + + def test_is_active(self): + sc = SessionContract() + session = sc.create("test") + self.assertFalse(sc.is_active(session.session_id)) + sc.activate(session.session_id) + self.assertTrue(sc.is_active(session.session_id)) + + def test_is_active_after_revoke(self): + sc = SessionContract() + session = sc.create("test") + sc.activate(session.session_id) + sc.revoke(session.session_id) + self.assertFalse(sc.is_active(session.session_id)) + + def test_metadata_preserved(self): + sc = SessionContract() + session = sc.create("test", metadata={"user_id": "u123"}) + self.assertEqual(session.metadata["user_id"], "u123") + + +# ========================================================================= +# WP2 — Event Stream Contract Tests +# ========================================================================= + + +class TestEventStreamContract(unittest.TestCase): + """Event stream buffering, replay, and reconnect policy.""" + + def test_emit_event(self): + es = EventStreamContract() + evt = es.emit("message", {"text": "hello"}) + self.assertEqual(evt.event_type, "message") + self.assertEqual(evt.sequence, 1) + self.assertEqual(evt.data["text"], "hello") + + def test_sequence_monotonic(self): + es = EventStreamContract() + e1 = es.emit("a", {}) + e2 = es.emit("b", {}) + e3 = es.emit("c", {}) + self.assertEqual(e1.sequence, 1) + self.assertEqual(e2.sequence, 2) + self.assertEqual(e3.sequence, 3) + + def test_replay_from_event_id(self): + es = EventStreamContract() + e1 = es.emit("a", {"n": 1}) + e2 = es.emit("b", {"n": 2}) + e3 = es.emit("c", {"n": 3}) + replayed = es.replay_from(e1.event_id) + self.assertEqual(len(replayed), 2) + self.assertEqual(replayed[0].event_id, e2.event_id) + self.assertEqual(replayed[1].event_id, e3.event_id) + + def test_replay_from_last_returns_empty(self): + es = EventStreamContract() + e1 = es.emit("a", {}) + replayed = es.replay_from(e1.event_id) + self.assertEqual(len(replayed), 0) + + def test_replay_from_unknown_id_returns_all(self): + es = EventStreamContract() + es.emit("a", {}) + es.emit("b", {}) + replayed = es.replay_from("nonexistent") + self.assertEqual(len(replayed), 2) + + def test_buffer_bounded(self): + es = EventStreamContract(max_buffer=5) + for i in range(10): + es.emit("e", {"i": i}) + self.assertEqual(len(es.get_all()), 5) + # Oldest events evicted + self.assertEqual(es.get_all()[0].data["i"], 5) + + def test_latest_sequence(self): + es = EventStreamContract() + self.assertEqual(es.latest_sequence, 0) + es.emit("a", {}) + es.emit("b", {}) + self.assertEqual(es.latest_sequence, 2) + + +class TestReconnectPolicy(unittest.TestCase): + """Reconnect/backoff policy for event streams.""" + + def test_should_retry_within_limit(self): + policy = ReconnectPolicy(max_retries=3) + self.assertTrue(policy.should_retry(0)) + self.assertTrue(policy.should_retry(2)) + self.assertFalse(policy.should_retry(3)) + + def test_exponential_backoff(self): + policy = ReconnectPolicy( + initial_delay_ms=100, max_delay_ms=10000, jitter_ms=0, max_retries=5 + ) + d0 = policy.compute_delay_ms(0) + d1 = policy.compute_delay_ms(1) + d2 = policy.compute_delay_ms(2) + self.assertEqual(d0, 100) # 100 * 2^0 + self.assertEqual(d1, 200) # 100 * 2^1 + self.assertEqual(d2, 400) # 100 * 2^2 + + def test_max_delay_cap(self): + policy = ReconnectPolicy( + initial_delay_ms=1000, max_delay_ms=5000, jitter_ms=0, max_retries=10 + ) + d = policy.compute_delay_ms(8) # 1000 * 2^8 = 256000, capped at 5000 + self.assertEqual(d, 5000) + + def test_exceeded_retries_returns_negative(self): + policy = ReconnectPolicy(max_retries=2) + self.assertEqual(policy.compute_delay_ms(2), -1) + + def test_jitter_bounded(self): + policy = ReconnectPolicy( + initial_delay_ms=100, max_delay_ms=10000, jitter_ms=50, max_retries=5 + ) + delays = {policy.compute_delay_ms(0) for _ in range(50)} + for d in delays: + self.assertGreaterEqual(d, 100) + self.assertLessEqual(d, 150) + + +# ========================================================================= +# WP3 — Callback Contract Tests +# ========================================================================= + + +class TestCallbackContract(unittest.TestCase): + """Callback delivery lifecycle and idempotency.""" + + def test_create_callback(self): + cc = CallbackContract() + record = cc.create() + self.assertEqual(record.state, CallbackState.PENDING.value) + self.assertIsNotNone(record.callback_id) + + def test_acknowledge_within_window(self): + cc = CallbackContract(ack_window_sec=10) + record = cc.create() + acked = cc.acknowledge(record.callback_id) + self.assertEqual(acked.state, CallbackState.ACKNOWLEDGED.value) + self.assertIsNotNone(acked.acknowledged_at) + + def test_acknowledge_expired_window(self): + cc = CallbackContract(ack_window_sec=0) + record = cc.create() + record.created_at = time.time() - 1 # force past window + with self.assertRaises(CallbackError) as ctx: + cc.acknowledge(record.callback_id) + self.assertIn("Ack window expired", str(ctx.exception)) + + def test_deliver_from_acknowledged(self): + cc = CallbackContract(ack_window_sec=10) + record = cc.create() + cc.acknowledge(record.callback_id) + delivered = cc.deliver(record.callback_id) + self.assertEqual(delivered.state, CallbackState.DELIVERED.value) + + def test_deliver_from_pending_default_mode(self): + """Default mode is strict: pending must ack before deliver.""" + cc = CallbackContract() + record = cc.create() + self.assertTrue(record.require_ack) + self.assertFalse(record.allow_direct_delivery) + with self.assertRaises(CallbackError) as ctx: + cc.deliver(record.callback_id) + self.assertIn("must acknowledge first", str(ctx.exception)) + + def test_allow_direct_delivery_opt_out(self): + """Compatibility opt-out allows direct pending delivery.""" + cc = CallbackContract() + record = cc.create(allow_direct_delivery=True) + self.assertFalse(record.require_ack) + self.assertTrue(record.allow_direct_delivery) + delivered = cc.deliver(record.callback_id) + self.assertEqual(delivered.state, CallbackState.DELIVERED.value) + + def test_cannot_deliver_expired(self): + cc = CallbackContract(callback_ttl_sec=0) + record = cc.create() + record.created_at = time.time() - 1 + with self.assertRaises(CallbackError): + cc.deliver(record.callback_id) + + def test_single_use_delivery(self): + """Once delivered, cannot deliver again.""" + cc = CallbackContract() + record = cc.create(allow_direct_delivery=True) + cc.deliver(record.callback_id) + with self.assertRaises(CallbackError): + cc.deliver(record.callback_id) + + def test_idempotency_dedupe(self): + cc = CallbackContract() + r1 = cc.create(idempotency_key="req-abc") + r2 = cc.create(idempotency_key="req-abc") + self.assertEqual(r1.callback_id, r2.callback_id) + + def test_idempotency_new_after_terminal(self): + """After expiry, same idempotency key creates new record.""" + cc = CallbackContract(callback_ttl_sec=0) + r1 = cc.create(idempotency_key="req-xyz") + r1.created_at = time.time() - 1 # force expire + cc.get(r1.callback_id) # triggers expiry check + r2 = cc.create(idempotency_key="req-xyz") + self.assertNotEqual(r1.callback_id, r2.callback_id) + + def test_max_attempts_fail(self): + cc = CallbackContract(max_attempts=2) + record = cc.create() + cc.record_attempt(record.callback_id) + result = cc.record_attempt(record.callback_id) + self.assertEqual(result.state, CallbackState.FAILED.value) + + def test_get_by_idempotency_key(self): + cc = CallbackContract() + r1 = cc.create(idempotency_key="lookup-test") + found = cc.get_by_idempotency_key("lookup-test") + self.assertIsNotNone(found) + self.assertEqual(found.callback_id, r1.callback_id) + + def test_get_by_idempotency_key_not_found(self): + cc = CallbackContract() + self.assertIsNone(cc.get_by_idempotency_key("nope")) + + def test_payload_hash(self): + cc = CallbackContract() + r1 = cc.create(payload={"action": "generate", "params": {"seed": 42}}) + self.assertTrue(len(r1.payload_hash) > 0) + + def test_callback_not_found(self): + cc = CallbackContract() + with self.assertRaises(CallbackError): + cc.acknowledge("nonexistent") + + # -- Strict mode (require_ack=True) tests -------------------------- + + def test_strict_mode_rejects_direct_deliver(self): + """Explicit strict mode: deliver() rejects pending (must ack first).""" + cc = CallbackContract(ack_window_sec=10) + record = cc.create(require_ack=True) + self.assertTrue(record.require_ack) + with self.assertRaises(CallbackError) as ctx: + cc.deliver(record.callback_id) + self.assertIn("require_ack=True", str(ctx.exception)) + self.assertIn("must acknowledge first", str(ctx.exception)) + + def test_strict_mode_ack_then_deliver(self): + """require_ack=True: ack -> deliver succeeds.""" + cc = CallbackContract(ack_window_sec=10) + record = cc.create(require_ack=True) + cc.acknowledge(record.callback_id) + delivered = cc.deliver(record.callback_id) + self.assertEqual(delivered.state, CallbackState.DELIVERED.value) + + def test_strict_mode_expired_ack_window_on_deliver(self): + """Strict pending callback auto-expires when ack window is missed.""" + cc = CallbackContract(ack_window_sec=5) + record = cc.create(require_ack=True) + record.created_at = time.time() - 10 # far past ack window + with self.assertRaises(CallbackError): + cc.deliver(record.callback_id) + # Verify state is now expired + refreshed = cc.get(record.callback_id) + self.assertEqual(refreshed.state, CallbackState.EXPIRED.value) + + def test_conflicting_policy_rejected(self): + """Cannot request strict + direct mode simultaneously.""" + cc = CallbackContract() + with self.assertRaises(CallbackError) as ctx: + cc.create(require_ack=True, allow_direct_delivery=True) + self.assertIn("conflicts", str(ctx.exception)) + + +# ========================================================================= +# WP4 — Token Contract Tests +# ========================================================================= + + +class TestTokenContract(unittest.TestCase): + """Token source precedence and fail-closed behavior.""" + + def _sources(self): + return [ + TokenSource(name="primary", env_var="TOKEN_PRIMARY", precedence=1), + TokenSource(name="fallback", env_var="TOKEN_FALLBACK", precedence=2), + ] + + def test_resolve_primary(self): + tc = TokenContract(self._sources()) + env = {"TOKEN_PRIMARY": "pk-abc123xyz", "TOKEN_FALLBACK": "fk-999"} + result = tc.resolve(env) + self.assertEqual(result.validity, TokenValidity.VALID.value) + self.assertEqual(result.source_name, "primary") + self.assertEqual(result.raw_value, "pk-abc123xyz") + + def test_resolve_fallback(self): + tc = TokenContract(self._sources()) + env = {"TOKEN_FALLBACK": "fk-longtoken99"} + result = tc.resolve(env) + self.assertEqual(result.source_name, "fallback") + + def test_resolve_missing(self): + tc = TokenContract(self._sources()) + result = tc.resolve({}) + self.assertEqual(result.validity, TokenValidity.MISSING.value) + + def test_fail_closed_rejects(self): + tc = TokenContract(self._sources()) + with self.assertRaises(TokenError) as ctx: + tc.validate_or_reject({}) + self.assertIn("Fail-closed", str(ctx.exception)) + + def test_fail_closed_with_optional_sources(self): + sources = [ + TokenSource( + name="optional", env_var="OPT_TOKEN", precedence=1, required=False + ), + ] + tc = TokenContract(sources) + result = tc.validate_or_reject({}) + self.assertEqual(result.validity, TokenValidity.MISSING.value) + # No exception because not required + + def test_masking_short_token(self): + tc = TokenContract(self._sources()) + env = {"TOKEN_PRIMARY": "short"} + result = tc.resolve(env) + self.assertEqual(result.masked_value, "***") + self.assertNotIn("short", result.masked_value) + + def test_masking_long_token(self): + tc = TokenContract(self._sources()) + env = {"TOKEN_PRIMARY": "pk-abcdefghijklmnop"} + result = tc.resolve(env) + self.assertTrue(result.masked_value.startswith("pk-a")) + self.assertTrue(result.masked_value.endswith("op")) + self.assertIn("***", result.masked_value) + # Must not contain full token + self.assertNotEqual(result.masked_value, "pk-abcdefghijklmnop") + + def test_precedence_table(self): + tc = TokenContract(self._sources()) + table = tc.get_precedence_table() + self.assertEqual(len(table), 2) + self.assertEqual(table[0]["name"], "primary") + self.assertEqual(table[1]["name"], "fallback") + + def test_whitespace_token_ignored(self): + tc = TokenContract(self._sources()) + env = {"TOKEN_PRIMARY": " ", "TOKEN_FALLBACK": "real-token-value"} + result = tc.resolve(env) + self.assertEqual(result.source_name, "fallback") + + def test_to_dict_excludes_raw_value(self): + """TokenResult serialized views hard-exclude raw_value.""" + tc = TokenContract(self._sources()) + env = {"TOKEN_PRIMARY": "pk-secret-token-12345"} + result = tc.resolve(env) + # Verify raw_value is populated internally + self.assertEqual(result.raw_value, "pk-secret-token-12345") + # Verify to_dict() excludes it + d = result.to_dict() + self.assertNotIn("raw_value", d) + self.assertNotIn("pk-secret-token-12345", str(d)) + # Verify other fields present + self.assertEqual(d["validity"], TokenValidity.VALID.value) + self.assertIn("***", d["masked_value"]) + + def test_to_public_dict_matches_public_contract(self): + tc = TokenContract(self._sources()) + env = {"TOKEN_PRIMARY": "pk-secret-token-12345"} + result = tc.resolve(env) + public = result.to_public() + d = result.to_public_dict() + self.assertEqual(d, public.to_dict()) + self.assertFalse(hasattr(public, "raw_value")) + + +# ========================================================================= +# WP3 — Retry Policy Tests +# ========================================================================= + + +class TestRetryPolicy(unittest.TestCase): + """Deterministic retry policy for callback delivery.""" + + def test_should_retry_within_limit(self): + rp = RetryPolicy(max_retries=3) + self.assertTrue(rp.should_retry(0)) + self.assertTrue(rp.should_retry(2)) + self.assertFalse(rp.should_retry(3)) + + def test_should_not_retry_on_4xx(self): + rp = RetryPolicy() + self.assertFalse(rp.should_retry(0, status_code=400)) + self.assertFalse(rp.should_retry(0, status_code=404)) + + def test_should_retry_on_5xx(self): + rp = RetryPolicy() + self.assertTrue(rp.should_retry(0, status_code=500)) + self.assertTrue(rp.should_retry(0, status_code=503)) + + def test_should_retry_on_429(self): + rp = RetryPolicy() + self.assertTrue(rp.should_retry(0, status_code=429)) + + def test_compute_delay_backoff(self): + rp = RetryPolicy(initial_delay_sec=1.0, backoff_factor=2.0, max_retries=5) + self.assertAlmostEqual(rp.compute_delay(0), 1.0) + self.assertAlmostEqual(rp.compute_delay(1), 2.0) + self.assertAlmostEqual(rp.compute_delay(2), 4.0) + + def test_compute_delay_capped(self): + rp = RetryPolicy( + initial_delay_sec=1.0, max_delay_sec=5.0, backoff_factor=2.0, max_retries=10 + ) + self.assertAlmostEqual(rp.compute_delay(8), 5.0) + + def test_compute_delay_exceeded(self): + rp = RetryPolicy(max_retries=2) + self.assertEqual(rp.compute_delay(2), -1.0) + + +# ========================================================================= +# Transport Error Envelope Tests +# ========================================================================= + + +class TestTransportError(unittest.TestCase): + """Normalized error envelope.""" + + def test_to_dict(self): + err = TransportError( + code="session_expired", + message="Session has expired", + retryable=False, + details={"session_id": "abc123"}, + ) + d = err.to_dict() + self.assertEqual(d["code"], "session_expired") + self.assertFalse(d["retryable"]) + self.assertIn("session_id", d["details"]) + + def test_retryable_error(self): + err = TransportError( + code="rate_limited", + message="Too many requests", + retryable=True, + ) + self.assertTrue(err.retryable) + + +# ========================================================================= +# Regression — Existing Connector Unaffected +# ========================================================================= + + +class TestExistingConnectorRegression(unittest.TestCase): + """Verify existing connector contracts remain importable and unchanged.""" + + def test_command_request_unchanged(self): + from connector.contract import CommandRequest + + req = CommandRequest( + platform="telegram", + sender_id="123", + channel_id="456", + username="test", + message_id="789", + text="/help", + timestamp=time.time(), + ) + self.assertEqual(req.platform, "telegram") + + def test_command_response_unchanged(self): + from connector.contract import CommandResponse + + resp = CommandResponse(text="OK", files=[], buttons=[]) + self.assertEqual(resp.text, "OK") + + def test_platform_abc_unchanged(self): + from connector.contract import Platform + + p = Platform() + # Methods exist + self.assertTrue(hasattr(p, "start")) + self.assertTrue(hasattr(p, "stop")) + self.assertTrue(hasattr(p, "send_message")) + self.assertTrue(hasattr(p, "send_image")) + + def test_transport_contract_does_not_modify_existing(self): + """Import transport_contract alongside existing contract — no conflict.""" + from connector.contract import CommandRequest + from connector.transport_contract import SessionContract + + # Both importable, no namespace collision + sc = SessionContract() + req = CommandRequest( + platform="test", + sender_id="s", + channel_id="c", + username="u", + message_id="m", + text="t", + timestamp=0, + ) + self.assertIsNotNone(sc) + self.assertIsNotNone(req) + + +# ========================================================================= +# Kakao Disabled Path Regression +# ========================================================================= + + +class TestKakaoDisabledRegression(unittest.TestCase): + """When Kakao is not configured, no side effects on other platforms.""" + + def test_no_kakao_env_no_error(self): + """Missing Kakao env vars do not break connector config loading.""" + # Ensure no Kakao-specific env vars exist + import os + + from connector.config import load_config + + for key in list(os.environ.keys()): + if "KAKAO" in key.upper(): + os.environ.pop(key) + + config = load_config() + # Config loads fine without Kakao + self.assertIsNotNone(config) + + def test_session_contract_platform_agnostic(self): + """Session contract works for any platform name including kakao.""" + sc = SessionContract() + for platform in ["telegram", "discord", "line", "whatsapp", "kakao", "wechat"]: + session = sc.create(platform) + self.assertEqual(session.platform, platform) + sc.activate(session.session_id) + self.assertTrue(sc.is_active(session.session_id)) + + +if __name__ == "__main__": + unittest.main()