"""advisor_report 建议报告仓储:本地镜像查询 + 列表(工作台本地,不直连 Agent 表)。""" from __future__ import annotations from sqlalchemy import func, select from model.advisor_report import AdvisorReport from repositories.base import BaseRepository class AdvisorReportRepo(BaseRepository): model = AdvisorReport async def get_by_report_id(self, report_id: str) -> AdvisorReport | None: return await self.db.scalar( select(AdvisorReport).where(AdvisorReport.report_id == report_id) ) async def get_by_draft_id(self, draft_id: str) -> AdvisorReport | None: """按草稿 draft_id 取本地镜像(草稿快照 / 已发送报告)。""" return await self.db.scalar( select(AdvisorReport).where(AdvisorReport.draft_id == draft_id) ) async def list_by_advisor( self, *, advisor_id: int, customer_id: int | None = None, send_status: str | None = None, limit: int = 100, offset: int = 0, ) -> list[AdvisorReport]: conds = [AdvisorReport.advisor_id == advisor_id] if customer_id is not None: conds.append(AdvisorReport.customer_id == customer_id) if send_status: conds.append(AdvisorReport.send_status == send_status) stmt = ( select(AdvisorReport) .where(*conds) .order_by(AdvisorReport.id.desc()) .limit(limit) .offset(offset) ) return list((await self.db.scalars(stmt)).all()) async def count_by_advisor( self, *, advisor_id: int, customer_id: int | None = None, send_status: str | None = None, ) -> int: conds = [AdvisorReport.advisor_id == advisor_id] if customer_id is not None: conds.append(AdvisorReport.customer_id == customer_id) if send_status: conds.append(AdvisorReport.send_status == send_status) stmt = select(func.count()).select_from(AdvisorReport).where(*conds) return (await self.db.scalar(stmt)) or 0 async def list_by_customer( self, customer_id: int, *, advisor_id: int | None = None, limit: int = 100, offset: int = 0, ) -> list[AdvisorReport]: """某客户的历史建议报告(360 视图用)。""" conditions = [AdvisorReport.customer_id == customer_id] if advisor_id is not None: conditions.append(AdvisorReport.advisor_id == advisor_id) stmt = ( select(AdvisorReport) .where(*conditions) .order_by(AdvisorReport.id.desc()) .limit(limit) .offset(offset) ) return list((await self.db.scalars(stmt)).all())