feat(connector-security): enforce strict ack-by-default callback lifecycle and public-safe token serialization in R75 transport contract

This commit is contained in:
rookiestar28
2026-02-12 10:58:32 +08:00
parent 721dd42c1e
commit 38fce78425
2 changed files with 1373 additions and 0 deletions
+710
View File
@@ -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,
}
+663
View File
@@ -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()