import asyncio import logging import random import time import uuid from collections.abc import Callable, Coroutine from dataclasses import dataclass, field from enum import StrEnum from typing import Any logger = logging.getLogger(__name__) class DurableStrategy(StrEnum): REQUIRED = "required" BEST_EFFORT = "best_effort" DISABLED = "disabled" class MessageSendState(StrEnum): IDLE = "idle" RENDERING = "rendering" PREVIEWING = "previewing" SENDING = "sending" SENT = "sent" SUPPRESSED = "suppressed" PARTIAL_FAILED = "partial_failed" FAILED = "failed" UNKNOWN_AFTER_SEND = "unknown_after_send" FINALIZING = "finalizing" EDITING = "editing" EDITED = "edited" DELETING = "deleting" DELETED = "deleted" CANCELLED = "cancelled" class MessageReceiptPartKind(StrEnum): TEXT = "text" MEDIA = "media" VOICE = "voice" CARD = "card" PREVIEW = "preview" UNKNOWN = "unknown" @dataclass class MessageReceiptPart: platform_message_id: str kind: MessageReceiptPartKind = MessageReceiptPartKind.UNKNOWN index: int = 0 thread_id: str | None = None reply_to_id: str | None = None @dataclass class DurableMessageReceipt: primary_platform_message_id: str = "" platform_message_ids: list[str] = field(default_factory=list) parts: list[MessageReceiptPart] = field(default_factory=list) thread_id: str | None = None reply_to_id: str | None = None edit_token: str | None = None delete_token: str | None = None sent_at: float = 0.0 metadata: dict[str, Any] = field(default_factory=dict) def __post_init__(self): if self.sent_at == 0.0: self.sent_at = time.time() if not self.platform_message_ids and self.primary_platform_message_id: self.platform_message_ids = [self.primary_platform_message_id] @dataclass class MessageSendContext: target_id: str content: str id: str = "" channel: str = "" account_id: str | None = None reply_to_id: str | None = None thread_id: str | None = None strategy: DurableStrategy = DurableStrategy.BEST_EFFORT receipt: DurableMessageReceipt | None = None previous_receipt: DurableMessageReceipt | None = None state: MessageSendState = MessageSendState.IDLE error: str | None = None attempt: int = 1 retry_count: int = 0 max_retries: int = 3 min_delay_ms: int = 300 max_delay_ms: int = 30_000 jitter: float = 0.0 metadata: dict[str, Any] = field(default_factory=dict) parts: list["MessageSendContext"] = field(default_factory=list) _on_commit: Callable[..., Coroutine[Any, Any, None]] | None = field(default=None, repr=False) _on_fail: Callable[..., Coroutine[Any, Any, None]] | None = field(default=None, repr=False) def __post_init__(self): if not self.id: self.id = f"{self.channel or 'msg'}:{self.target_id}:{uuid.uuid4().hex[:8]}" @property def is_terminal(self) -> bool: return self.state in ( MessageSendState.SENT, MessageSendState.SUPPRESSED, MessageSendState.PARTIAL_FAILED, MessageSendState.FAILED, MessageSendState.CANCELLED, ) async def render(self) -> str: self.state = MessageSendState.RENDERING return self.content def _backoff_delay(self, attempt: int) -> int: base = self.min_delay_ms * (2 ** (attempt - 1)) delay = min(base, self.max_delay_ms) if self.jitter > 0: offset = (random.random() * 2 - 1) * self.jitter delay = int(delay * (1 + offset)) return max(0, delay) async def send(self, send_fn: Callable[..., Coroutine[Any, Any, str | None]]) -> DurableMessageReceipt | None: self.state = MessageSendState.SENDING last_error = None total_attempts = self.max_retries + 1 for attempt_idx in range(total_attempts): self.attempt = attempt_idx + 1 try: message_id = await send_fn(self.content) if message_id: self.state = MessageSendState.SENT self.receipt = DurableMessageReceipt(primary_platform_message_id=message_id) return self.receipt self.state = MessageSendState.SENT return None except Exception as e: last_error = str(e) self.retry_count = attempt_idx + 1 logger.warning( "Message send attempt %d/%d failed: %s", attempt_idx + 1, total_attempts, e, ) if attempt_idx < total_attempts - 1: delay = self._backoff_delay(attempt_idx + 1) if delay > 0: await asyncio.sleep(delay / 1000) self.state = MessageSendState.FAILED self.error = last_error return None async def send_batch( self, contents: list[str], send_fn: Callable[..., Coroutine[Any, Any, str | None]], ) -> list[DurableMessageReceipt | None]: if not contents: return [] results: list[DurableMessageReceipt | None] = [] self.parts.clear() failed_count = 0 for i, content in enumerate(contents): part = MessageSendContext( target_id=self.target_id, content=content, id=f"{self.id}#{i}", channel=self.channel, account_id=self.account_id, reply_to_id=self.reply_to_id, thread_id=self.thread_id, strategy=self.strategy, max_retries=self.max_retries, min_delay_ms=self.min_delay_ms, max_delay_ms=self.max_delay_ms, jitter=self.jitter, ) receipt = await part.send(send_fn) results.append(receipt) self.parts.append(part) if receipt is None and part.state == MessageSendState.FAILED: failed_count += 1 total = len(contents) if failed_count == total: self.state = MessageSendState.FAILED self.error = f"All {total} parts failed" elif failed_count > 0: self.state = MessageSendState.PARTIAL_FAILED self.error = f"{failed_count}/{total} parts failed" else: self.state = MessageSendState.SENT return results def mark_suppressed(self, reason: str = "") -> None: self.state = MessageSendState.SUPPRESSED self.error = reason async def edit( self, edit_fn: Callable[..., Coroutine[Any, Any, str | None]], new_content: str ) -> DurableMessageReceipt | None: if self.receipt is None: logger.warning("Cannot edit message without receipt") return None self.state = MessageSendState.EDITING try: new_id = await edit_fn(self.receipt.primary_platform_message_id, new_content) if new_id: self.receipt.primary_platform_message_id = new_id if new_id not in self.receipt.platform_message_ids: self.receipt.platform_message_ids.append(new_id) self.state = MessageSendState.EDITED return self.receipt except Exception as e: self.state = MessageSendState.FAILED self.error = str(e) return None async def delete(self, delete_fn: Callable[..., Coroutine[Any, Any, None]]) -> bool: if self.receipt is None: return False self.state = MessageSendState.DELETING try: await delete_fn(self.receipt.primary_platform_message_id) self.state = MessageSendState.DELETED return True except Exception as e: self.state = MessageSendState.FAILED self.error = str(e) return False def mark_cancelled(self) -> None: self.state = MessageSendState.CANCELLED def mark_unknown_after_send(self) -> None: self.state = MessageSendState.UNKNOWN_AFTER_SEND async def commit(self) -> None: if self._on_commit: await self._on_commit(self.receipt) async def fail(self, error: Exception | None = None) -> None: if self._on_fail: if error is None: error = Exception(self.error or "send failed") await self._on_fail(error) class OutboundBridge: def __init__( self, send_text_fn: Callable[..., Coroutine[Any, Any, str | None]] | None = None, send_media_fn: Callable[..., Coroutine[Any, Any, str | None]] | None = None, send_payload_fn: Callable[..., Coroutine[Any, Any, str | None]] | None = None, ): self._send_text = send_text_fn self._send_media = send_media_fn self._send_payload = send_payload_fn async def text( self, target_id: str, content: str, *, reply_to_id: str | None = None, thread_id: str | None = None, ) -> str | None: if not self._send_text: raise RuntimeError("OutboundBridge: send_text not configured") return await self._send_text(target_id, content, reply_to_id=reply_to_id, thread_id=thread_id) async def media( self, target_id: str, media_url: str, text: str = "", *, reply_to_id: str | None = None, thread_id: str | None = None, audio_as_voice: bool = False, ) -> str | None: if not self._send_media: raise RuntimeError("OutboundBridge: send_media not configured") return await self._send_media( target_id, media_url, text, reply_to_id=reply_to_id, thread_id=thread_id, audio_as_voice=audio_as_voice ) async def payload( self, target_id: str, payload: Any, *, reply_to_id: str | None = None, thread_id: str | None = None, ) -> str | None: if not self._send_payload: raise RuntimeError("OutboundBridge: send_payload not configured") return await self._send_payload(target_id, payload, reply_to_id=reply_to_id, thread_id=thread_id) class DurableSendContextManager: def __init__( self, ctx: "MessageSendContext", *, on_commit: Callable[..., Coroutine[Any, Any, None]] | None = None, on_fail: Callable[..., Coroutine[Any, Any, None]] | None = None, ): self.ctx = ctx self.ctx._on_commit = on_commit self.ctx._on_fail = on_fail async def __aenter__(self) -> "MessageSendContext": return self.ctx async def __aexit__(self, exc_type, exc_val, exc_tb) -> bool: if exc_type is not None: await self.ctx.fail(exc_val) return False async def send_durable_message_batch( ctx: "MessageSendContext", send_fn: Callable[..., Coroutine[Any, Any, str | None]], *, on_commit: Callable[..., Coroutine[Any, Any, None]] | None = None, on_fail: Callable[..., Coroutine[Any, Any, None]] | None = None, ) -> "DurableMessageReceipt | None": if on_commit: ctx._on_commit = on_commit if on_fail: ctx._on_fail = on_fail await ctx.render() result = await ctx.send(send_fn) if result is not None and ctx.state in (MessageSendState.SENT, MessageSendState.SUPPRESSED): await ctx.commit() else: await ctx.fail() return result