mirror of
https://github.com/rookiestar28/ComfyUI-OpenClaw.git
synced 2026-08-14 00:48:07 +00:00
feat(connector-security): enforce strict ack-by-default callback lifecycle and public-safe token serialization in R75 transport contract
This commit is contained in:
@@ -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,
|
||||
}
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user