ForcePilot/backend/package/yuxi/channel/message/durable_send.py

347 lines
11 KiB
Python
Raw Normal View History

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