nodes.py 6.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287
  1. from __future__ import annotations
  2. """需求解析节点及基础控制节点:解析/标准化TravelRequest、澄清、错误处理与路由判断。"""
  3. from datetime import date, timedelta
  4. from typing import Literal
  5. from app.agents.requirement_agent import (
  6. RequirementAgent,
  7. )
  8. from app.graph.state import TravelState
  9. from app.schemas.travel_request import (
  10. TravelRequest,
  11. )
  12. FIELD_LABELS = {
  13. "origin_city": "出发城市",
  14. "destination_city": "目的城市",
  15. "departure_date": "出发日期",
  16. "return_date": "返程日期",
  17. }
  18. def normalize_and_validate_request(
  19. request: TravelRequest,
  20. ) -> tuple[
  21. TravelRequest,
  22. list[str],
  23. list[str],
  24. ]:
  25. """补全可确定推导的字段,并检查关键信息。
  26. 返回:
  27. 1. 标准化后的需求;
  28. 2. 缺失字段;
  29. 3. 错误信息。
  30. """
  31. updates: dict[str, object] = {}
  32. errors: list[str] = []
  33. departure_date = request.departure_date
  34. return_date = request.return_date
  35. # 根据出发日期和天数推算返程日期。
  36. if (
  37. departure_date is not None
  38. and return_date is None
  39. ):
  40. if request.nights is not None:
  41. updates["return_date"] = (
  42. departure_date
  43. + timedelta(days=request.nights)
  44. )
  45. elif request.trip_days is not None:
  46. updates["return_date"] = (
  47. departure_date
  48. + timedelta(
  49. days=request.trip_days - 1
  50. )
  51. )
  52. normalized = request.model_copy(
  53. update=updates
  54. )
  55. # 日期完整时,补全天数和晚数。
  56. if (
  57. normalized.departure_date is not None
  58. and normalized.return_date is not None
  59. ):
  60. nights = (
  61. normalized.return_date
  62. - normalized.departure_date
  63. ).days
  64. if normalized.nights is None:
  65. normalized = normalized.model_copy(
  66. update={"nights": nights}
  67. )
  68. if normalized.trip_days is None:
  69. normalized = normalized.model_copy(
  70. update={"trip_days": nights + 1}
  71. )
  72. missing_fields: list[str] = []
  73. for field_name in (
  74. "origin_city",
  75. "destination_city",
  76. "departure_date",
  77. "return_date",
  78. ):
  79. if getattr(normalized, field_name) is None:
  80. missing_fields.append(field_name)
  81. if (
  82. normalized.origin_city
  83. and normalized.destination_city
  84. and normalized.origin_city
  85. == normalized.destination_city
  86. ):
  87. errors.append(
  88. "出发城市和目的城市不能相同。"
  89. )
  90. if (
  91. normalized.departure_date is not None
  92. and normalized.departure_date
  93. < date.today()
  94. ):
  95. errors.append(
  96. "出发日期早于当前日期。"
  97. )
  98. if (
  99. normalized.departure_date is not None
  100. and normalized.return_date is not None
  101. and normalized.return_date
  102. <= normalized.departure_date
  103. ):
  104. errors.append(
  105. "返程日期必须晚于出发日期。"
  106. )
  107. return normalized, missing_fields, errors
  108. def make_parse_request_node(
  109. agent: RequirementAgent,
  110. ):
  111. """创建带依赖的需求解析节点。"""
  112. async def parse_request_node(
  113. state: TravelState,
  114. ) -> dict:
  115. """需求解析节点:调用RequirementAgent解析用户输入为结构化TravelRequest,失败时写入errors。"""
  116. user_query = state.get(
  117. "user_query",
  118. "",
  119. ).strip()
  120. if not user_query:
  121. return {
  122. "errors": [
  123. "用户旅行需求不能为空。"
  124. ],
  125. "missing_fields": [],
  126. }
  127. try:
  128. request = await agent.parse(user_query)
  129. (
  130. normalized_request,
  131. missing_fields,
  132. errors,
  133. ) = normalize_and_validate_request(
  134. request
  135. )
  136. return {
  137. "travel_request": (
  138. normalized_request
  139. ),
  140. "missing_fields": missing_fields,
  141. "errors": errors,
  142. }
  143. except Exception as exc:
  144. return {
  145. "missing_fields": [],
  146. "errors": [
  147. "需求解析失败:"
  148. f"{type(exc).__name__}: {exc}"
  149. ],
  150. }
  151. return parse_request_node
  152. def route_after_parse(
  153. state: TravelState,
  154. ) -> Literal[
  155. "clarify",
  156. "ready",
  157. "error",
  158. ]:
  159. """根据需求解析结果决定下一节点。"""
  160. if state.get("errors"):
  161. return "error"
  162. if state.get("missing_fields"):
  163. return "clarify"
  164. return "ready"
  165. def clarification_node(
  166. state: TravelState,
  167. ) -> dict:
  168. """生成需要用户补充的问题。"""
  169. missing_fields = state.get(
  170. "missing_fields",
  171. [],
  172. )
  173. labels = [
  174. FIELD_LABELS.get(field, field)
  175. for field in missing_fields
  176. ]
  177. question = (
  178. "为了继续规划,请补充:"
  179. + "、".join(labels)
  180. + "。"
  181. )
  182. return {
  183. "needs_clarification": True,
  184. "clarification_question": question,
  185. "final_answer": question,
  186. }
  187. def error_node(
  188. state: TravelState,
  189. ) -> dict:
  190. """向用户返回确定性校验错误。"""
  191. errors = state.get("errors", [])
  192. message = (
  193. "旅行需求存在以下问题:"
  194. + ";".join(errors)
  195. + " 请修改后重新提交。"
  196. )
  197. return {
  198. "needs_clarification": True,
  199. "final_answer": message,
  200. }
  201. def ready_node(
  202. state: TravelState,
  203. ) -> dict:
  204. """本阶段的完成节点。
  205. 后续这里会连接航班、酒店和景点查询节点。
  206. """
  207. request = state["travel_request"]
  208. travelers = (
  209. f"{request.adults}位成人"
  210. f"、{request.children}位儿童"
  211. )
  212. budget_text = (
  213. f"{request.total_budget:.0f}"
  214. f"{request.currency}"
  215. if request.total_budget is not None
  216. else "未设置总预算"
  217. )
  218. message = (
  219. "需求解析完成:"
  220. f"{request.origin_city}"
  221. f" → {request.destination_city},"
  222. f"{request.departure_date}"
  223. f" 至 {request.return_date},"
  224. f"{request.trip_days}天"
  225. f"{request.nights}晚,"
  226. f"{travelers},"
  227. f"{budget_text}。"
  228. )
  229. return {
  230. "needs_clarification": False,
  231. "final_answer": message,
  232. }