95 lines
2.7 KiB
Python
95 lines
2.7 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import logging
|
||
|
|
import time
|
||
|
|
from collections import OrderedDict
|
||
|
|
from dataclasses import dataclass, field
|
||
|
|
|
||
|
|
from yuxi.channels.models import DeliveryResult
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass
|
||
|
|
class SentMessage:
|
||
|
|
message_id: str
|
||
|
|
chat_id: str
|
||
|
|
content: str = ""
|
||
|
|
sent_at: float = field(default_factory=time.time)
|
||
|
|
updated_at: float = 0.0
|
||
|
|
update_count: int = 0
|
||
|
|
status: str = "sent"
|
||
|
|
error: str = ""
|
||
|
|
|
||
|
|
|
||
|
|
class MessageCache:
|
||
|
|
def __init__(self, max_messages: int = 500, ttl_s: float = 3600.0):
|
||
|
|
self._max = max_messages
|
||
|
|
self._ttl = ttl_s
|
||
|
|
self._messages: OrderedDict[str, SentMessage] = OrderedDict()
|
||
|
|
|
||
|
|
def record_sent(self, result: DeliveryResult, chat_id: str, content: str = "") -> SentMessage | None:
|
||
|
|
if not result.message_id:
|
||
|
|
return None
|
||
|
|
|
||
|
|
msg = SentMessage(
|
||
|
|
message_id=result.message_id,
|
||
|
|
chat_id=chat_id,
|
||
|
|
content=content[:500],
|
||
|
|
status="sent" if result.success else "failed",
|
||
|
|
error=result.error or "",
|
||
|
|
)
|
||
|
|
self._messages[result.message_id] = msg
|
||
|
|
self._messages.move_to_end(result.message_id)
|
||
|
|
|
||
|
|
while len(self._messages) > self._max:
|
||
|
|
self._messages.popitem(last=False)
|
||
|
|
|
||
|
|
return msg
|
||
|
|
|
||
|
|
def record_update(self, message_id: str, content: str = "") -> SentMessage | None:
|
||
|
|
msg = self._messages.get(message_id)
|
||
|
|
if msg is None:
|
||
|
|
return None
|
||
|
|
msg.updated_at = time.time()
|
||
|
|
msg.update_count += 1
|
||
|
|
if content:
|
||
|
|
msg.content = content[:500]
|
||
|
|
return msg
|
||
|
|
|
||
|
|
def get(self, message_id: str) -> SentMessage | None:
|
||
|
|
return self._messages.get(message_id)
|
||
|
|
|
||
|
|
def get_by_chat(self, chat_id: str) -> list[SentMessage]:
|
||
|
|
return [m for m in self._messages.values() if m.chat_id == chat_id]
|
||
|
|
|
||
|
|
def cleanup_expired(self) -> int:
|
||
|
|
now = time.time()
|
||
|
|
expired = [mid for mid, msg in self._messages.items() if now - msg.sent_at > self._ttl]
|
||
|
|
for mid in expired:
|
||
|
|
del self._messages[mid]
|
||
|
|
return len(expired)
|
||
|
|
|
||
|
|
@property
|
||
|
|
def snapshot(self) -> dict:
|
||
|
|
return {
|
||
|
|
"total": len(self._messages),
|
||
|
|
"max": self._max,
|
||
|
|
"ttl_s": self._ttl,
|
||
|
|
"recent": [
|
||
|
|
{
|
||
|
|
"message_id": m.message_id,
|
||
|
|
"chat_id": m.chat_id,
|
||
|
|
"sent_at": m.sent_at,
|
||
|
|
"updates": m.update_count,
|
||
|
|
"status": m.status,
|
||
|
|
}
|
||
|
|
for m in list(self._messages.values())[-20:]
|
||
|
|
],
|
||
|
|
}
|
||
|
|
|
||
|
|
def clear(self) -> int:
|
||
|
|
count = len(self._messages)
|
||
|
|
self._messages.clear()
|
||
|
|
return count
|