Files
Mutual_Fund/repositories/advisor_draft.py
T

63 lines
1.9 KiB
Python
Raw Normal View History

2026-09-13 16:19:24 +08:00
"""投顾 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