from datetime import UTC, datetime from sqlalchemy import select, update from sqlalchemy.engine import CursorResult from sqlalchemy.ext.asyncio import AsyncSession from app.core.errors import InvalidStateError, SessionNotFoundError from app.model.session import ConversationSession class SessionRepository: def __init__(self, session: AsyncSession) -> None: self.session = session async def owned( self, session_id: str, user_id: int, *, lock: bool = False ) -> ConversationSession: query = select(ConversationSession).where(ConversationSession.session_id == session_id, ConversationSession.user_id == user_id) if lock: query = query.with_for_update() row = await self.session.scalar(query) if row is None: raise SessionNotFoundError("会话不存在") return row async def require_active(self, session_id: str, user_id: int, agent_type: str) -> None: row = await self.owned(session_id, user_id, lock=True) if row.status != "active": raise InvalidStateError("会话已结束") if row.agent_type != agent_type: raise InvalidStateError("Agent 与会话不一致") row.message_count += 1 row.last_active_at = datetime.now(UTC).replace(tzinfo=None) async def advance_clarification(self, session_id: str, expected: int) -> bool: if expected >= 10: return False result = await self.session.execute(update(ConversationSession).where( ConversationSession.session_id == session_id, ConversationSession.clarification_round == expected, ConversationSession.status == "active", ).values(clarification_round=expected + 1)) return isinstance(result, CursorResult) and result.rowcount == 1