main.py 4.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198
  1. from __future__ import annotations
  2. """TravelMind CLI入口:读取用户旅行需求(支持多行输入),构建多角色规划图并执行完整工作流。"""
  3. import asyncio
  4. import logging
  5. import sys
  6. from time import perf_counter
  7. from rich.console import Console
  8. from app.graph.planner_builder import (
  9. build_planner_graph,
  10. )
  11. from app.logging_config import (
  12. configure_logging,
  13. new_trace_id,
  14. )
  15. console = Console()
  16. logger = logging.getLogger(__name__)
  17. def get_user_query() -> str:
  18. """从命令行参数或终端读取旅行需求(支持多行输入,空行结束)。"""
  19. command_line_query = " ".join(
  20. sys.argv[1:]
  21. ).strip()
  22. if command_line_query:
  23. return command_line_query
  24. console.print(
  25. "[bold]请输入旅行需求"
  26. "(支持多行,输入空行结束):[/bold]"
  27. )
  28. lines: list[str] = []
  29. line_count = 0
  30. while True:
  31. prefix = "> " if line_count == 0 else "… "
  32. line = input(prefix)
  33. stripped = line.strip()
  34. if not stripped:
  35. if lines:
  36. break
  37. # 第一行就为空,继续等待。
  38. continue
  39. lines.append(stripped)
  40. line_count += 1
  41. return " ".join(lines)
  42. async def run_workflow(
  43. user_query: str,
  44. ) -> dict:
  45. """执行完整TravelMind工作流。"""
  46. graph = await build_planner_graph()
  47. return await graph.ainvoke(
  48. {
  49. "user_query": user_query,
  50. "errors": [],
  51. "missing_fields": [],
  52. "planning_attempts": 0,
  53. "review_attempts": 0,
  54. "revision_feedback": [],
  55. "workflow_status": "running",
  56. },
  57. config={
  58. # 图中存在重新规划循环。
  59. "recursion_limit": 80,
  60. },
  61. )
  62. async def main() -> None:
  63. """CLI主入口:配置日志、生成trace_id、读取用户需求并执行完整工作流。"""
  64. configure_logging()
  65. trace_id = new_trace_id()
  66. user_query = get_user_query()
  67. if not user_query:
  68. console.print(
  69. "[bold red]旅行需求不能为空。[/bold red]"
  70. )
  71. raise SystemExit(1)
  72. console.rule(
  73. "[bold blue]TravelMind 多角色旅行规划"
  74. )
  75. console.print(
  76. f"[dim]trace_id:{trace_id}[/dim]"
  77. )
  78. started_at = perf_counter()
  79. logger.info(
  80. "TravelMind工作流开始,用户输入长度=%s",
  81. len(user_query),
  82. )
  83. try:
  84. result = await run_workflow(
  85. user_query
  86. )
  87. except Exception as exc:
  88. elapsed = perf_counter() - started_at
  89. logger.exception(
  90. "TravelMind工作流异常,"
  91. "elapsed_seconds=%.2f",
  92. elapsed,
  93. )
  94. console.print(
  95. "[bold red]工作流执行失败:[/bold red]"
  96. f"{type(exc).__name__}: {exc}"
  97. )
  98. raise SystemExit(1) from exc
  99. elapsed = perf_counter() - started_at
  100. workflow_status = result.get(
  101. "workflow_status"
  102. )
  103. logger.info(
  104. "TravelMind工作流结束,"
  105. "status=%s,"
  106. "planning_attempts=%s,"
  107. "review_attempts=%s,"
  108. "elapsed_seconds=%.2f",
  109. workflow_status,
  110. result.get("planning_attempts"),
  111. result.get("review_attempts"),
  112. elapsed,
  113. )
  114. final_answer = result.get(
  115. "final_answer"
  116. )
  117. console.print()
  118. if final_answer:
  119. console.print(final_answer)
  120. else:
  121. console.print(
  122. "[bold red]"
  123. "系统没有生成最终回答。"
  124. "[/bold red]"
  125. )
  126. console.rule("[bold]执行状态")
  127. console.print(
  128. {
  129. "trace_id": trace_id,
  130. "工作流状态": workflow_status,
  131. "Planner执行次数": result.get(
  132. "planning_attempts"
  133. ),
  134. "Reviewer执行次数": result.get(
  135. "review_attempts"
  136. ),
  137. "Supervisor决策": result.get(
  138. "supervisor_decision"
  139. ),
  140. "总耗时(秒)": round(
  141. elapsed,
  142. 2,
  143. ),
  144. "系统错误": result.get(
  145. "errors",
  146. [],
  147. ),
  148. }
  149. )
  150. if __name__ == "__main__":
  151. asyncio.run(main())