2026-09-09 21:55:37 +08:00
|
|
|
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(
|
2026-09-10 15:55:54 +08:00
|
|
|
self, session_id: str, user_id: int, limit: int, before: int | None = None
|
2026-09-09 21:55:37 +08:00
|
|
|
) -> list[ConversationMessage]:
|
2026-09-10 15:55:54 +08:00
|
|
|
"""按会话取一页消息(默认 `id DESC`,即最新一页)。
|
|
|
|
|
|
|
|
|
|
`before` 与 `PlatformRepository.rows` 的游标语义一致(`id < before`,取更旧的
|
|
|
|
|
一页);`None` 表示不翻页,行为与加游标前逐字相同。游标的合法性由入口的
|
|
|
|
|
`parse_cursor` 负责,这里只接受已验证的整数边界。
|
|
|
|
|
"""
|
|
|
|
|
query = select(ConversationMessage).where(
|
2026-09-09 21:55:37 +08:00
|
|
|
ConversationMessage.session_id == session_id,
|
|
|
|
|
ConversationMessage.customer_id == user_id,
|
2026-09-10 15:55:54 +08:00
|
|
|
)
|
|
|
|
|
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)))
|
2026-09-09 21:55:37 +08:00
|
|
|
|
|
|
|
|
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
|