229 lines
9.5 KiB
Python
229 lines
9.5 KiB
Python
"""Conversation-level investment-goal entity extraction and gap handling."""
|
||
|
||
import re
|
||
from datetime import UTC, datetime
|
||
from decimal import Decimal
|
||
from typing import Any, Protocol
|
||
|
||
from sqlalchemy import select
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from app.infrastructure.db import SessionFactory
|
||
from app.model.audit import InteractionAudit
|
||
from app.model.conversation import ConversationMessage
|
||
from app.model.goal_conversation import AdvisorGoalConversationExtraction
|
||
from app.model.session import ConversationSession
|
||
|
||
EXTRACTION_VERSION = "investment-goal-v1"
|
||
MAX_CLARIFICATION_ROUNDS = 10
|
||
REQUIRED_FIELDS = (
|
||
"annualized_return_lower_pct",
|
||
"annualized_return_upper_pct",
|
||
"max_drawdown_pct",
|
||
"liquidity_requirement",
|
||
"investment_horizon_months",
|
||
"benchmark_name",
|
||
)
|
||
|
||
FIELD_LABELS = {
|
||
"annualized_return_lower_pct": "年化收益目标下限",
|
||
"annualized_return_upper_pct": "年化收益目标上限",
|
||
"max_drawdown_pct": "最大回撤容忍度",
|
||
"liquidity_requirement": "资金流动性要求",
|
||
"investment_horizon_months": "投资期限",
|
||
"benchmark_name": "业绩比较基准",
|
||
}
|
||
|
||
_NUMBER = r"\d+(?:\.\d+)?"
|
||
|
||
|
||
class SessionFactoryLike(Protocol):
|
||
def __call__(self) -> AsyncSession: ...
|
||
|
||
|
||
def _number(value: str) -> float:
|
||
return float(Decimal(value))
|
||
|
||
|
||
def extract_goal_entities(text: str) -> dict[str, Any]:
|
||
"""Extract explicit goal facts from one user message.
|
||
|
||
This parser intentionally accepts only explicit financial statements. It does not
|
||
infer a risk level, expected return, or benchmark from vague language.
|
||
"""
|
||
result: dict[str, Any] = {}
|
||
return_match = re.search(
|
||
rf"(?:年化)?(?:收益目标|收益率|收益)\s*(?:为|在|约为|目标为)?\s*({_NUMBER})\s*%?\s*"
|
||
rf"(?:至|到|[-~~])\s*({_NUMBER})\s*%?",
|
||
text,
|
||
re.IGNORECASE,
|
||
)
|
||
if return_match:
|
||
lower, upper = map(_number, return_match.groups())
|
||
result["annualized_return_lower_pct"] = min(lower, upper)
|
||
result["annualized_return_upper_pct"] = max(lower, upper)
|
||
else:
|
||
single_return = re.search(
|
||
rf"(?:年化)?(?:收益目标|收益率|收益)\s*"
|
||
rf"(?:是|为|约为|目标为)?\s*({_NUMBER})\s*%",
|
||
text,
|
||
re.IGNORECASE,
|
||
)
|
||
if single_return:
|
||
value = _number(single_return.group(1))
|
||
result["annualized_return_lower_pct"] = value
|
||
result["annualized_return_upper_pct"] = value
|
||
|
||
drawdown = re.search(rf"(?:最大)?回撤\s*(?:容忍|控制|不超过|为|约为)?\s*({_NUMBER})\s*%?", text)
|
||
if drawdown:
|
||
result["max_drawdown_pct"] = _number(drawdown.group(1))
|
||
|
||
if re.search(r"随时|每天|每日|当天|日内", text):
|
||
result["liquidity_requirement"] = "daily"
|
||
else:
|
||
days = re.search(
|
||
r"(?:流动性|用钱|变现|取出).{0,12}?"
|
||
r"([一两两三四五六七八九十0-9]+)\s*天",
|
||
text,
|
||
)
|
||
if days is None:
|
||
days = re.search(
|
||
r"([一两三四五六七八九十0-9]+)\s*天(?:内|可用|能用)", text
|
||
)
|
||
if days:
|
||
count = _chinese_number(days.group(1))
|
||
if count <= 7:
|
||
result["liquidity_requirement"] = "within_7_days"
|
||
elif count <= 30:
|
||
result["liquidity_requirement"] = "within_30_days"
|
||
else:
|
||
result["liquidity_requirement"] = "over_30_days"
|
||
|
||
months = re.search(rf"({_NUMBER})\s*(?:年|岁)", text)
|
||
if months:
|
||
result["investment_horizon_months"] = int(_number(months.group(1)) * 12)
|
||
else:
|
||
months = re.search(rf"({_NUMBER})\s*个?月", text)
|
||
if months:
|
||
result["investment_horizon_months"] = int(_number(months.group(1)))
|
||
|
||
benchmark = re.search(
|
||
r"(?:业绩比较基准|比较基准|基准)\s*(?:是|为|选)?\s*"
|
||
r"([\u4e00-\u9fffA-Za-z0-9_-]{2,32})",
|
||
text,
|
||
)
|
||
if benchmark:
|
||
result["benchmark_name"] = benchmark.group(1).strip(",。;、")
|
||
else:
|
||
for candidate in ("沪深300指数", "沪深300", "中证500", "中证全指", "上证指数"):
|
||
if candidate in text:
|
||
result["benchmark_name"] = candidate
|
||
break
|
||
return result
|
||
|
||
|
||
def _chinese_number(value: str) -> int:
|
||
if value.isdigit():
|
||
return int(value)
|
||
digits = {"一": 1, "二": 2, "两": 2, "三": 3, "四": 4, "五": 5,
|
||
"六": 6, "七": 7, "八": 8, "九": 9, "十": 10}
|
||
if value == "十":
|
||
return 10
|
||
if value.startswith("十"):
|
||
return 10 + digits.get(value[1:], 0)
|
||
if value.endswith("十"):
|
||
return digits.get(value[0], 0) * 10
|
||
if "十" in value:
|
||
left, right = value.split("十", 1)
|
||
return digits.get(left, 0) * 10 + digits.get(right, 0)
|
||
return digits.get(value, 0)
|
||
|
||
|
||
class GoalConversationService:
|
||
"""Merge explicit facts across a session and persist an internal gap snapshot."""
|
||
|
||
def __init__(
|
||
self, session_factory: SessionFactoryLike = SessionFactory,
|
||
) -> None:
|
||
self.session_factory = session_factory
|
||
|
||
async def process(
|
||
self, *, session_id: str, customer_id: int, trace_id: str, message: str
|
||
) -> dict[str, Any]:
|
||
async with self.session_factory() as session, session.begin():
|
||
session_row = await session.scalar(select(ConversationSession).where(
|
||
ConversationSession.session_id == session_id,
|
||
ConversationSession.user_id == customer_id,
|
||
).with_for_update())
|
||
if session_row is None:
|
||
return self._result({}, list(REQUIRED_FIELDS), "partial")
|
||
messages = list(await session.scalars(select(ConversationMessage).where(
|
||
ConversationMessage.session_id == session_id,
|
||
ConversationMessage.customer_id == customer_id,
|
||
ConversationMessage.role == "user",
|
||
).order_by(ConversationMessage.created_at, ConversationMessage.id)))
|
||
merged: dict[str, Any] = {}
|
||
for row in messages:
|
||
merged.update(extract_goal_entities(row.content))
|
||
merged.update(extract_goal_entities(message))
|
||
missing = [field for field in REQUIRED_FIELDS if field not in merged]
|
||
status = "complete" if not missing else (
|
||
"clarification_limit"
|
||
if session_row.clarification_round >= MAX_CLARIFICATION_ROUNDS
|
||
else "partial"
|
||
)
|
||
if status == "partial":
|
||
session_row.clarification_round += 1
|
||
now = datetime.now(UTC).replace(tzinfo=None)
|
||
message_row = await session.scalar(select(ConversationMessage).where(
|
||
ConversationMessage.session_id == session_id,
|
||
ConversationMessage.customer_id == customer_id,
|
||
ConversationMessage.trace_id == trace_id,
|
||
ConversationMessage.role == "user",
|
||
).order_by(ConversationMessage.id.desc()).limit(1))
|
||
if message_row is not None:
|
||
exists = await session.scalar(select(AdvisorGoalConversationExtraction).where(
|
||
AdvisorGoalConversationExtraction.message_id == message_row.id))
|
||
if exists is None:
|
||
session.add(AdvisorGoalConversationExtraction(
|
||
message_id=message_row.id, session_id=session_id,
|
||
customer_id=customer_id, extraction_version=EXTRACTION_VERSION,
|
||
extracted_fields=merged, missing_fields=missing, status=status,
|
||
created_at=now, updated_at=now,
|
||
))
|
||
session.add(InteractionAudit(
|
||
actor_type="agent", actor_id=customer_id, session_id=session_id,
|
||
portal="agent", action_type="advisor.goal_gap_evaluated",
|
||
detail={"trace_id": trace_id, "status": status,
|
||
"missing_count": len(missing)}, created_at=now,
|
||
))
|
||
return self._result(merged, missing, status)
|
||
|
||
@staticmethod
|
||
def _result(fields: dict[str, Any], missing: list[str], status: str) -> dict[str, Any]:
|
||
return {"fields": fields, "missing": missing, "status": status}
|
||
|
||
@staticmethod
|
||
def customer_prompt(result: dict[str, Any]) -> str:
|
||
missing = result.get("missing", [])
|
||
if result.get("status") == "clarification_limit":
|
||
return "投资目标信息仍不完整,暂时无法继续采集,请转人工顾问协助确认。"
|
||
if missing:
|
||
labels = [FIELD_LABELS[field] for field in missing if field in FIELD_LABELS]
|
||
return "为了形成投资目标,还需要确认:" + "、".join(labels) + "。请一次补充这些信息。"
|
||
fields = result["fields"]
|
||
return (
|
||
f"投资目标信息已收集完整:年化收益目标 {fields['annualized_return_lower_pct']}%-"
|
||
f"{fields['annualized_return_upper_pct']}%,最大回撤 {fields['max_drawdown_pct']}%,"
|
||
f"流动性 {fields['liquidity_requirement']},"
|
||
f"期限 {fields['investment_horizon_months']} 个月,"
|
||
f"业绩比较基准 {fields['benchmark_name']}。请确认后再生成目标书。"
|
||
)
|
||
|
||
@staticmethod
|
||
def should_collect(text: str) -> bool:
|
||
"""Distinguish goal collection from a request to view an existing goal."""
|
||
return bool(extract_goal_entities(text)) or bool(re.search(
|
||
r"制定|设定|创建|填写|采集|规划|目标书", text
|
||
))
|