"""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"]