Files
Mutual_Fund/service/advisor_agent/draft.py
T

180 lines
5.6 KiB
Python
Raw Normal View History

2026-09-13 16:19:24 +08:00
"""投顾 Agent 草稿服务。"""
from __future__ import annotations
from decimal import Decimal
from typing import Any
import uuid
from common.common_const import (
DRAFT_STATUS_DISCARDED,
DRAFT_STATUS_DRAFT,
ERR_CODE_DRAFT_NOT_FOUND,
ERR_CODE_FORBIDDEN_CUSTOMER,
ERR_CODE_SUITABILITY_INVALID,
REPORT_DISCLAIMER,
)
from common.suitability import check_suitability
from service.advisor_agent.compliance import ensure_safe_content
from repositories.sensitive_word import SensitiveWordRepo
from model.advisor_draft import AdvisorDraft
from utils.exceptions import ApiError
def build_generated_content(content: str) -> str:
"""生成阶段强制补齐免责声明,避免重复拼接。"""
if REPORT_DISCLAIMER in content:
return content
return f"{content.rstrip()}\n\n{REPORT_DISCLAIMER}"
def _disclaimer_warning(content: str) -> tuple[bool, str | None]:
if REPORT_DISCLAIMER in content:
return True, None
return False, "草稿缺少完整免责声明,工作台发送前必须补齐并重新校验"
def _not_found() -> ApiError:
return ApiError(ERR_CODE_DRAFT_NOT_FOUND, "草稿不存在或者已废弃")
async def _resolve_sensitive_words(repo, explicit_words: list[str] | None) -> list[str]:
if explicit_words is not None:
return explicit_words
if hasattr(repo, "db"):
return await SensitiveWordRepo(repo.db).list_active_words()
return []
def _validate_structured_suitability(
structured_data: dict | None, customer_risk: str | None = None
) -> None:
if not structured_data:
return
customer_risk = customer_risk or structured_data.get("customer_risk")
if not customer_risk:
return
for item in structured_data.get("items", []):
result = check_suitability(customer_risk, item.get("risk_level", ""))
if not result.ok:
raise ApiError(ERR_CODE_SUITABILITY_INVALID, result.reason)
async def create_draft(repo, data: dict) -> AdvisorDraft:
content = data.get("content", "")
sensitive_words = await _resolve_sensitive_words(repo, data.get("sensitive_words"))
ensure_safe_content(content, sensitive_words)
disclaimer_ok, warning = _disclaimer_warning(content)
draft = AdvisorDraft(
draft_id=uuid.uuid4().hex,
customer_id=data["customer_id"],
advisor_id=data["advisor_id"],
intent=data["intent"],
title=data["title"],
content=content,
structured_data=data.get("structured_data"),
status=DRAFT_STATUS_DRAFT,
deviation=data.get("deviation"),
disclaimer_ok=disclaimer_ok,
warning=warning,
)
return await repo.add(draft)
def ensure_draft_owner(draft: Any, *, advisor_id: int) -> None:
if draft.advisor_id != advisor_id:
raise ApiError(ERR_CODE_FORBIDDEN_CUSTOMER, "无权操作该客户数据")
def summarize_draft(draft: Any) -> dict:
deviation = getattr(draft, "deviation", None)
if isinstance(deviation, Decimal):
deviation = float(deviation)
return {
"draft_id": draft.draft_id,
"customer_id": draft.customer_id,
"advisor_id": draft.advisor_id,
"intent": getattr(draft, "intent", None),
"title": getattr(draft, "title", None),
"status": draft.status,
"deviation": deviation,
"disclaimer_ok": bool(getattr(draft, "disclaimer_ok", False)),
"created_at": draft.create_time.isoformat()
if getattr(draft, "create_time", None)
else None,
"update_time": draft.update_time.isoformat()
if getattr(draft, "update_time", None)
else None,
}
def detail_draft(draft: Any) -> dict:
result = summarize_draft(draft)
result.update(
{
"content": draft.content,
"structured_data": getattr(draft, "structured_data", None),
"warning": getattr(draft, "warning", None),
}
)
return result
async def get_draft(repo, draft_id: str):
draft = await repo.get_by_draft_id(draft_id)
if draft is None or draft.status == DRAFT_STATUS_DISCARDED:
raise _not_found()
return draft
async def list_drafts(
repo,
*,
advisor_id: int | None = None,
customer_id: int | None = None,
status: str | None = None,
page: int = 1,
page_size: int = 20,
) -> dict:
page = max(1, page)
page_size = min(100, max(1, page_size))
total, items = await repo.list_drafts(
advisor_id=advisor_id,
customer_id=customer_id,
status=status,
limit=page_size,
offset=(page - 1) * page_size,
)
return {"total": total, "items": [summarize_draft(item) for item in items]}
async def save_draft(
repo,
draft_id: str,
*,
title: str | None = None,
content: str | None = None,
structured_data: dict | None = None,
customer_risk: str | None = None,
sensitive_words: list[str] | None = None,
):
draft = await get_draft(repo, draft_id)
if draft.status != DRAFT_STATUS_DRAFT:
raise _not_found()
if title is not None:
draft.title = title
if content is not None:
draft.content = content
if structured_data is not None:
draft.structured_data = structured_data
resolved_sensitive_words = await _resolve_sensitive_words(repo, sensitive_words)
ensure_safe_content(draft.content, resolved_sensitive_words)
_validate_structured_suitability(draft.structured_data, customer_risk)
draft.disclaimer_ok, draft.warning = _disclaimer_warning(draft.content)
saved = await repo.save(draft)
return detail_draft(saved)
async def discard_draft(repo, draft_id: str):
draft = await get_draft(repo, draft_id)
return await repo.discard(draft)