"""投顾 Agent 草稿仓储。""" from __future__ import annotations from sqlalchemy import select from agent.advisor_agent.drafts import ( DRAFT_STATUS_DISCARDED, DRAFT_STATUS_DRAFT, ) from model.advisor_draft import AdvisorDraft from repositories.base import BaseRepository class AdvisorDraftRepo(BaseRepository): model = AdvisorDraft async def get_by_draft_id(self, draft_id: str) -> AdvisorDraft | None: return await self.db.scalar( select(AdvisorDraft).where(AdvisorDraft.draft_id == draft_id) ) async def save(self, draft: AdvisorDraft) -> AdvisorDraft: self.db.add(draft) await self.db.commit() await self.db.refresh(draft) return draft async def list_drafts( self, *, advisor_id: int | None = None, customer_id: int | None = None, status: str | None = None, limit: int = 20, offset: int = 0, ) -> tuple[int, list[AdvisorDraft]]: filters = [] if advisor_id is not None: filters.append(AdvisorDraft.advisor_id == advisor_id) if customer_id is not None: filters.append(AdvisorDraft.customer_id == customer_id) if status is not None: filters.append(AdvisorDraft.status == status) total = await self.count(where=filters) items = await self.list( where=filters, order_by=AdvisorDraft.update_time.desc(), limit=limit, offset=offset, ) return total, items async def discard(self, draft: AdvisorDraft) -> AdvisorDraft: if draft.status == DRAFT_STATUS_DISCARDED: return draft if draft.status != DRAFT_STATUS_DRAFT: raise ValueError("草稿状态无效") draft.status = DRAFT_STATUS_DISCARDED await self.db.commit() await self.db.refresh(draft) return draft