| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287 |
- from __future__ import annotations
- """需求解析节点及基础控制节点:解析/标准化TravelRequest、澄清、错误处理与路由判断。"""
- from datetime import date, timedelta
- from typing import Literal
- from app.agents.requirement_agent import (
- RequirementAgent,
- )
- from app.graph.state import TravelState
- from app.schemas.travel_request import (
- TravelRequest,
- )
- FIELD_LABELS = {
- "origin_city": "出发城市",
- "destination_city": "目的城市",
- "departure_date": "出发日期",
- "return_date": "返程日期",
- }
- def normalize_and_validate_request(
- request: TravelRequest,
- ) -> tuple[
- TravelRequest,
- list[str],
- list[str],
- ]:
- """补全可确定推导的字段,并检查关键信息。
- 返回:
- 1. 标准化后的需求;
- 2. 缺失字段;
- 3. 错误信息。
- """
- updates: dict[str, object] = {}
- errors: list[str] = []
- departure_date = request.departure_date
- return_date = request.return_date
- # 根据出发日期和天数推算返程日期。
- if (
- departure_date is not None
- and return_date is None
- ):
- if request.nights is not None:
- updates["return_date"] = (
- departure_date
- + timedelta(days=request.nights)
- )
- elif request.trip_days is not None:
- updates["return_date"] = (
- departure_date
- + timedelta(
- days=request.trip_days - 1
- )
- )
- normalized = request.model_copy(
- update=updates
- )
- # 日期完整时,补全天数和晚数。
- if (
- normalized.departure_date is not None
- and normalized.return_date is not None
- ):
- nights = (
- normalized.return_date
- - normalized.departure_date
- ).days
- if normalized.nights is None:
- normalized = normalized.model_copy(
- update={"nights": nights}
- )
- if normalized.trip_days is None:
- normalized = normalized.model_copy(
- update={"trip_days": nights + 1}
- )
- missing_fields: list[str] = []
- for field_name in (
- "origin_city",
- "destination_city",
- "departure_date",
- "return_date",
- ):
- if getattr(normalized, field_name) is None:
- missing_fields.append(field_name)
- if (
- normalized.origin_city
- and normalized.destination_city
- and normalized.origin_city
- == normalized.destination_city
- ):
- errors.append(
- "出发城市和目的城市不能相同。"
- )
- if (
- normalized.departure_date is not None
- and normalized.departure_date
- < date.today()
- ):
- errors.append(
- "出发日期早于当前日期。"
- )
- if (
- normalized.departure_date is not None
- and normalized.return_date is not None
- and normalized.return_date
- <= normalized.departure_date
- ):
- errors.append(
- "返程日期必须晚于出发日期。"
- )
- return normalized, missing_fields, errors
- def make_parse_request_node(
- agent: RequirementAgent,
- ):
- """创建带依赖的需求解析节点。"""
- async def parse_request_node(
- state: TravelState,
- ) -> dict:
- """需求解析节点:调用RequirementAgent解析用户输入为结构化TravelRequest,失败时写入errors。"""
- user_query = state.get(
- "user_query",
- "",
- ).strip()
- if not user_query:
- return {
- "errors": [
- "用户旅行需求不能为空。"
- ],
- "missing_fields": [],
- }
- try:
- request = await agent.parse(user_query)
- (
- normalized_request,
- missing_fields,
- errors,
- ) = normalize_and_validate_request(
- request
- )
- return {
- "travel_request": (
- normalized_request
- ),
- "missing_fields": missing_fields,
- "errors": errors,
- }
- except Exception as exc:
- return {
- "missing_fields": [],
- "errors": [
- "需求解析失败:"
- f"{type(exc).__name__}: {exc}"
- ],
- }
- return parse_request_node
- def route_after_parse(
- state: TravelState,
- ) -> Literal[
- "clarify",
- "ready",
- "error",
- ]:
- """根据需求解析结果决定下一节点。"""
- if state.get("errors"):
- return "error"
- if state.get("missing_fields"):
- return "clarify"
- return "ready"
- def clarification_node(
- state: TravelState,
- ) -> dict:
- """生成需要用户补充的问题。"""
- missing_fields = state.get(
- "missing_fields",
- [],
- )
- labels = [
- FIELD_LABELS.get(field, field)
- for field in missing_fields
- ]
- question = (
- "为了继续规划,请补充:"
- + "、".join(labels)
- + "。"
- )
- return {
- "needs_clarification": True,
- "clarification_question": question,
- "final_answer": question,
- }
- def error_node(
- state: TravelState,
- ) -> dict:
- """向用户返回确定性校验错误。"""
- errors = state.get("errors", [])
- message = (
- "旅行需求存在以下问题:"
- + ";".join(errors)
- + " 请修改后重新提交。"
- )
- return {
- "needs_clarification": True,
- "final_answer": message,
- }
- def ready_node(
- state: TravelState,
- ) -> dict:
- """本阶段的完成节点。
- 后续这里会连接航班、酒店和景点查询节点。
- """
- request = state["travel_request"]
- travelers = (
- f"{request.adults}位成人"
- f"、{request.children}位儿童"
- )
- budget_text = (
- f"{request.total_budget:.0f}"
- f"{request.currency}"
- if request.total_budget is not None
- else "未设置总预算"
- )
- message = (
- "需求解析完成:"
- f"{request.origin_city}"
- f" → {request.destination_city},"
- f"{request.departure_date}"
- f" 至 {request.return_date},"
- f"{request.trip_days}天"
- f"{request.nights}晚,"
- f"{travelers},"
- f"{budget_text}。"
- )
- return {
- "needs_clarification": False,
- "final_answer": message,
- }
|