Files
Mutual_Fund/repositories/conversation_archive.py

74 lines
2.6 KiB
Python

"""conversation_archive 会话归档仓储。"""
from __future__ import annotations
from collections.abc import Iterable
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from model.conversation_archive import ConversationArchive
from repositories.base import BaseRepository
class ConversationArchiveRepo(BaseRepository):
"""提供按会话读取和幂等批量归档能力。"""
model = ConversationArchive
async def list_by_session(self, session_id: str) -> list[ConversationArchive]:
"""按消息创建顺序读取完整归档会话。"""
statement = (
select(ConversationArchive)
.where(ConversationArchive.session_id == session_id)
.order_by(ConversationArchive.create_time, ConversationArchive.id)
)
return list((await self.db.scalars(statement)).all())
async def archive_batch(self, rows: Iterable[dict]) -> int:
"""批量写入归档记录,重复的 session_id + message_id 自动跳过。"""
rows = list(rows)
if not rows:
return 0
session_ids = {row.get("session_id") for row in rows}
if len(session_ids) != 1 or None in session_ids:
raise ValueError("archive_batch 只能接收同一会话的完整消息")
if any(not row.get("message_id") for row in rows):
raise ValueError("archive_batch 的 message_id 不能为空")
session_id = rows[0]["session_id"]
message_ids = [row["message_id"] for row in rows]
existing = await self.db.scalars(
select(ConversationArchive.message_id).where(
ConversationArchive.session_id == session_id,
ConversationArchive.message_id.in_(message_ids),
)
)
existing_ids = set(existing.all())
pending = [row for row in rows if row["message_id"] not in existing_ids]
if not pending:
return 0
self.db.add_all(
[ConversationArchive(**row) for row in pending]
)
try:
await self.db.commit()
except IntegrityError:
await self.db.rollback()
# 并发归档时,唯一键冲突代表其他调用已完成该消息归档。
remaining = await self.db.scalars(
select(ConversationArchive.message_id).where(
ConversationArchive.session_id == session_id,
ConversationArchive.message_id.in_(message_ids),
)
)
if set(remaining.all()) >= set(message_ids):
return 0
raise
return len(pending)
__all__ = ["ConversationArchiveRepo"]