from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.model.conversation import ConversationFeedback, ConversationMessage from app.model.platform import AgentRun class ConversationRepository: def __init__(self, session: AsyncSession) -> None: self.session = session async def messages( self, session_id: str, user_id: int, limit: int, before: int | None = None ) -> list[ConversationMessage]: """按会话取一页消息(默认 `id DESC`,即最新一页)。 `before` 与 `PlatformRepository.rows` 的游标语义一致(`id < before`,取更旧的 一页);`None` 表示不翻页,行为与加游标前逐字相同。游标的合法性由入口的 `parse_cursor` 负责,这里只接受已验证的整数边界。 """ query = select(ConversationMessage).where( ConversationMessage.session_id == session_id, ConversationMessage.customer_id == user_id, ) if before is not None: query = query.where(ConversationMessage.id < before) return list(await self.session.scalars( query.order_by(ConversationMessage.id.desc()).limit(limit))) async def message(self, message_id: int, user_id: int) -> ConversationMessage | None: statement = select(ConversationMessage).where( ConversationMessage.id == message_id, ConversationMessage.customer_id == user_id, ).with_for_update() row: ConversationMessage | None = await self.session.scalar(statement) return row async def feedback(self, message_id: int, user_id: int) -> ConversationFeedback | None: statement = select(ConversationFeedback).where( ConversationFeedback.message_id == message_id, ConversationFeedback.customer_id == user_id, ) row: ConversationFeedback | None = await self.session.scalar(statement) return row async def run_result( self, run_id: str, user_id: int ) -> tuple[AgentRun, ConversationMessage | None] | None: run = await self.session.scalar(select(AgentRun).where( AgentRun.run_id == run_id, AgentRun.user_id == user_id)) if run is None: return None message = (await self.session.get(ConversationMessage, run.result_message_id) if run.result_message_id else None) return run, message