182 lines
5.7 KiB
Python
182 lines
5.7 KiB
Python
"""投顾 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
|
|
from utils.pagination import normalize_pagination, pagination_result
|
|
|
|
|
|
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 = 10,
|
|
) -> dict:
|
|
page, page_size, offset = normalize_pagination(page, page_size)
|
|
total, items = await repo.list_drafts(
|
|
advisor_id=advisor_id,
|
|
customer_id=customer_id,
|
|
status=status,
|
|
limit=page_size,
|
|
offset=offset,
|
|
)
|
|
return pagination_result(
|
|
[summarize_draft(item) for item in items], total, page=page, page_size=page_size
|
|
)
|
|
|
|
|
|
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)
|