74 lines
2.6 KiB
Python
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"]
|