test_requirement_graph.py 1.6 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970
  1. from __future__ import annotations
  2. """测试需求解析阶段子图:验证自然语言解析为TravelRequest的完整流程与澄清分支。"""
  3. import asyncio
  4. from rich.console import Console
  5. from app.graph.builder import (
  6. build_requirement_graph,
  7. )
  8. console = Console()
  9. async def run_case(
  10. graph,
  11. user_query: str,
  12. ) -> None:
  13. console.rule("[bold]用户输入")
  14. console.print(user_query)
  15. result = await graph.ainvoke(
  16. {
  17. "user_query": user_query,
  18. "errors": [],
  19. "missing_fields": [],
  20. }
  21. )
  22. console.rule("[bold]Graph输出")
  23. console.print(result["final_answer"])
  24. request = result.get("travel_request")
  25. if request is not None:
  26. console.rule("[bold]结构化需求")
  27. console.print_json(
  28. request.model_dump_json(
  29. indent=2,
  30. )
  31. )
  32. async def main() -> None:
  33. graph = build_requirement_graph()
  34. await run_case(
  35. graph,
  36. (
  37. "我和妻子计划2026年8月10日"
  38. "从上海去成都,5天4晚,"
  39. "总预算5000元。喜欢熊猫、"
  40. "自然景观和川菜,不想安排太赶。"
  41. "酒店希望靠近地铁,最好住在市中心,住宿每晚不超过"
  42. "300元。机票价格相差不大的情况下,优先考虑落地双流机场"
  43. "只优先考虑直达。"
  44. ),
  45. )
  46. # 可选测试:信息不完整时是否追问。
  47. await run_case(
  48. graph,
  49. "我想从上海去成都旅游,喜欢美食。",
  50. )
  51. if __name__ == "__main__":
  52. asyncio.run(main())