| import json |
| from typing import Any, TypedDict |
|
|
| from langgraph.graph import END, StateGraph |
|
|
| from config import HF_TEXT_MODEL, HF_VISION_MODEL, MAX_AGENT_STEPS, MAX_TOOL_OUTPUT_CHARS |
| from tools.common import ( |
| AUDIO_VIDEO_EXTENSIONS, |
| IMAGE_EXTENSIONS, |
| SPREADSHEET_EXTENSIONS, |
| TEXT_EXTENSIONS, |
| extract_urls, |
| is_youtube_url, |
| normalize_answer, |
| truncate_text, |
| ) |
| from tools.executor import execute_tool |
| from tools.llm_client import classify_question_type, format_final_answer_with_llm |
| from tools.types import SolverResult |
|
|
|
|
| QUESTION_TYPE_SPECS = [ |
| { |
| "type": "direct_text", |
| "tool": "direct_answer_tool", |
| "description": "自包含纯文本、表格、反向字符串、列表筛选、简单规则或正则式问题。", |
| }, |
| { |
| "type": "python_code", |
| "tool": "python_tool", |
| "description": "带 .py 附件,需要运行或静态分析 Python 脚本得到输出。", |
| }, |
| { |
| "type": "spreadsheet", |
| "tool": "spreadsheet_tool", |
| "description": "带 .xlsx/.xls 附件,需要读取表格并计算。", |
| }, |
| { |
| "type": "wikipedia", |
| "tool": "wikipedia_tool", |
| "description": "Wikipedia、百科条目、人物作品、Featured Article、奥运表格等结构化网页问题。", |
| }, |
| { |
| "type": "sports", |
| "tool": "sports_tool", |
| "description": "体育统计、球队、赛季、球员数据问题。", |
| }, |
| { |
| "type": "web_url", |
| "tool": "web_read_tool", |
| "description": "问题中给出普通网页 URL,需要读取该 URL 内容。", |
| }, |
| { |
| "type": "web_search", |
| "tool": "web_search_tool", |
| "description": "需要开放网页检索,但没有明确可直接解析的专用工具。", |
| }, |
| { |
| "type": "attachment_text", |
| "tool": "attachment_text_tool", |
| "description": "带 .txt/.csv/.json/.md 等纯文本附件,需要读取附件内容。", |
| }, |
| { |
| "type": "audio_media", |
| "tool": "audio_tool", |
| "description": "音频转写题,例如 .mp3/.wav 附件;当前工具禁用。", |
| }, |
| { |
| "type": "video_media", |
| "tool": "video_tool", |
| "description": "视频或 YouTube 分析题;当前工具禁用。", |
| }, |
| { |
| "type": "vision_image", |
| "tool": "vision_tool", |
| "description": "图片、棋盘图或视觉识别题;当前工具禁用。", |
| }, |
| { |
| "type": "unknown", |
| "tool": "fallback", |
| "description": "无法可靠判断类型时进入兜底流程。", |
| }, |
| ] |
|
|
| QUESTION_TYPE_TO_NODE = { |
| "direct_text": "direct_text", |
| "python_code": "python_code", |
| "spreadsheet": "spreadsheet", |
| "wikipedia": "wikipedia", |
| "sports": "sports", |
| "web_url": "web_url", |
| "web_search": "web_search", |
| "attachment_text": "attachment_text", |
| "audio_media": "unsupported_media", |
| "video_media": "unsupported_media", |
| "vision_image": "unsupported_media", |
| "unknown": "fallback", |
| } |
|
|
|
|
| class GaiaWorkflowState(TypedDict, total=False): |
| question: str |
| task_id: str |
| file_name: str |
| question_type: str |
| type_confidence: str |
| type_reason: str |
| type_query: str |
| observation: dict[str, Any] |
| fallback_used: bool |
| answer: str |
| source: str |
| confidence: str |
| evidence: str |
| error: str |
| trace: list[dict[str, Any]] |
|
|
|
|
| class GaiaAgent: |
| """LangGraph 类型路由工作流 Agent。""" |
|
|
| def __init__(self): |
| print("GAIA LangGraph 类型工作流 Agent 已初始化。") |
| print(f"文本模型:{HF_TEXT_MODEL}") |
| print(f"视觉模型:{HF_VISION_MODEL or '未启用'}") |
| print(f"最大兜底工具数:{MAX_AGENT_STEPS}") |
| self.workflow = self._build_workflow() |
|
|
| def answer_task(self, question: str, task_id: str = "", file_name: str = "") -> SolverResult: |
| print(f"Agent 收到问题(前 100 个字符):{question[:100]}...") |
| initial_state: GaiaWorkflowState = { |
| "question": question, |
| "task_id": task_id, |
| "file_name": file_name, |
| "trace": [], |
| } |
| try: |
| final_state = self.workflow.invoke(initial_state) |
| except Exception as exc: |
| return SolverResult( |
| "无法确定", |
| source="langgraph.exception", |
| confidence="low", |
| evidence="", |
| error=str(exc), |
| ) |
|
|
| return SolverResult( |
| final_state.get("answer") or "无法确定", |
| source=final_state.get("source", "langgraph.final"), |
| confidence=final_state.get("confidence", "low"), |
| evidence=final_state.get("evidence", ""), |
| error=final_state.get("error", ""), |
| ) |
|
|
| def __call__(self, question: str, task_id: str = "", file_name: str = "") -> str: |
| result = self.answer_task(question, task_id=task_id, file_name=file_name) |
| print( |
| f"Agent 返回:answer={result.answer!r}, source={result.source}, " |
| f"confidence={result.confidence}, error={result.error}" |
| ) |
| return normalize_answer(result.answer or "无法确定") |
|
|
| def _build_workflow(self): |
| graph = StateGraph(GaiaWorkflowState) |
| graph.add_node("classify", self._classify_node) |
| graph.add_node("direct_text", self._direct_text_node) |
| graph.add_node("python_code", self._python_node) |
| graph.add_node("spreadsheet", self._spreadsheet_node) |
| graph.add_node("wikipedia", self._wikipedia_node) |
| graph.add_node("sports", self._sports_node) |
| graph.add_node("web_url", self._web_url_node) |
| graph.add_node("web_search", self._web_search_node) |
| graph.add_node("attachment_text", self._attachment_text_node) |
| graph.add_node("unsupported_media", self._unsupported_media_node) |
| graph.add_node("fallback", self._fallback_node) |
| graph.add_node("finalize", self._finalize_node) |
|
|
| graph.set_entry_point("classify") |
| graph.add_conditional_edges("classify", self._route_after_classification) |
| for node_name in ( |
| "direct_text", |
| "python_code", |
| "spreadsheet", |
| "wikipedia", |
| "sports", |
| "web_url", |
| "web_search", |
| "attachment_text", |
| "unsupported_media", |
| ): |
| graph.add_conditional_edges( |
| node_name, |
| self._route_after_tool, |
| {"fallback": "fallback", "finalize": "finalize"}, |
| ) |
| graph.add_edge("fallback", "finalize") |
| graph.add_edge("finalize", END) |
| return graph.compile() |
|
|
| def _classify_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| try: |
| classification = classify_question_type( |
| question=state["question"], |
| task_id=state.get("task_id", ""), |
| file_name=state.get("file_name", ""), |
| type_specs=QUESTION_TYPE_SPECS, |
| ) |
| question_type = classification.question_type |
| error = state.get("error", "") |
| trace_event = { |
| "event": "classify", |
| "question_type": question_type, |
| "confidence": classification.confidence, |
| "reason": classification.reason, |
| "query": classification.query, |
| } |
| return { |
| "question_type": question_type, |
| "type_confidence": classification.confidence, |
| "type_reason": classification.reason, |
| "type_query": classification.query, |
| "error": error, |
| "trace": self._append_trace(state, trace_event), |
| } |
| except Exception as exc: |
| trace_event = { |
| "event": "classify_error", |
| "question_type": "unknown", |
| "error": str(exc), |
| } |
| return { |
| "question_type": "unknown", |
| "type_confidence": "low", |
| "type_reason": "LLM 分类失败,进入兜底流程。", |
| "type_query": "", |
| "error": self._join_error(state.get("error", ""), f"classifier_error={exc}"), |
| "trace": self._append_trace(state, trace_event), |
| } |
|
|
| def _route_after_classification(self, state: GaiaWorkflowState) -> str: |
| return QUESTION_TYPE_TO_NODE.get(state.get("question_type", "unknown"), "fallback") |
|
|
| def _route_after_tool(self, state: GaiaWorkflowState) -> str: |
| if state.get("fallback_used"): |
| return "finalize" |
| if state.get("question_type") in {"audio_media", "video_media", "vision_image"}: |
| return "finalize" |
| observation = state.get("observation", {}) |
| if self._observation_is_useful(observation): |
| return "finalize" |
| return "fallback" |
|
|
| def _direct_text_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| return self._run_tool_node(state, "direct_answer_tool", {}) |
|
|
| def _python_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| return self._run_tool_node(state, "python_tool", {}) |
|
|
| def _spreadsheet_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| return self._run_tool_node(state, "spreadsheet_tool", {}) |
|
|
| def _wikipedia_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| query = state.get("type_query") or state["question"] |
| return self._run_tool_node(state, "wikipedia_tool", {"query": query}) |
|
|
| def _sports_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| return self._run_tool_node(state, "sports_tool", {}) |
|
|
| def _web_url_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| observation = self._read_urls_observation(state) |
| return self._state_with_observation(state, observation, "web_url") |
|
|
| def _web_search_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| observation = self._search_and_read_observation(state) |
| return self._state_with_observation(state, observation, "web_search") |
|
|
| def _attachment_text_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| return self._run_tool_node(state, "attachment_text_tool", {}) |
|
|
| def _unsupported_media_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| question_type = state.get("question_type") |
| tool_name = { |
| "audio_media": "audio_tool", |
| "video_media": "video_tool", |
| "vision_image": "vision_tool", |
| }.get(question_type, "vision_tool") |
| return self._run_tool_node(state, tool_name, {}) |
|
|
| def _fallback_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| observation = self._run_fallback_tools(state) |
| return self._state_with_observation( |
| {**state, "fallback_used": True}, |
| observation, |
| "fallback", |
| ) |
|
|
| def _finalize_node(self, state: GaiaWorkflowState) -> GaiaWorkflowState: |
| observation = state.get("observation", {}) |
| candidate_answer = str(observation.get("answer") or "").strip() |
| source = str(observation.get("source") or observation.get("tool") or "no_tool") |
| confidence = str(observation.get("confidence") or "low") |
| evidence = str(observation.get("evidence") or "") |
| final_error = state.get("error", "") |
|
|
| if candidate_answer or evidence: |
| try: |
| formatted = format_final_answer_with_llm( |
| question=state["question"], |
| question_type=state.get("question_type", "unknown"), |
| candidate_answer=candidate_answer, |
| evidence=truncate_text(evidence, MAX_TOOL_OUTPUT_CHARS), |
| source=source, |
| confidence=confidence, |
| ) |
| answer = formatted.answer or "无法确定" |
| confidence = formatted.confidence |
| final_error = self._join_error(final_error, observation.get("error", "")) |
| except Exception as exc: |
| answer = normalize_answer(candidate_answer) if candidate_answer else "无法确定" |
| final_error = self._join_error( |
| final_error, |
| observation.get("error", ""), |
| f"final_formatter_error={exc}", |
| ) |
| else: |
| answer = "无法确定" |
| final_error = self._join_error(final_error, observation.get("error", "工具没有返回可用证据。")) |
|
|
| trace = self._append_trace( |
| state, |
| { |
| "event": "finalize", |
| "answer": answer, |
| "source": source, |
| "confidence": confidence, |
| }, |
| ) |
| return { |
| "answer": normalize_answer(answer), |
| "source": f"langgraph.{state.get('question_type', 'unknown')}.{source}", |
| "confidence": confidence, |
| "evidence": self._trace_text(trace), |
| "error": final_error, |
| "trace": trace, |
| } |
|
|
| def _run_tool_node( |
| self, |
| state: GaiaWorkflowState, |
| tool_name: str, |
| args: dict[str, Any], |
| ) -> GaiaWorkflowState: |
| observation = execute_tool(tool_name, args, self._context(state)) |
| return self._state_with_observation(state, observation, tool_name) |
|
|
| def _state_with_observation( |
| self, |
| state: GaiaWorkflowState, |
| observation: dict[str, Any], |
| event_name: str, |
| ) -> GaiaWorkflowState: |
| compact_observation = self._compact_observation(observation) |
| return { |
| "observation": observation, |
| "trace": self._append_trace( |
| state, |
| { |
| "event": event_name, |
| "observation": compact_observation, |
| }, |
| ), |
| } |
|
|
| def _run_fallback_tools(self, state: GaiaWorkflowState) -> dict[str, Any]: |
| file_name = state.get("file_name", "").lower() |
| question = state["question"] |
| fallback_steps: list[tuple[str, dict[str, Any]]] = [] |
|
|
| if file_name.endswith(tuple(SPREADSHEET_EXTENSIONS)): |
| fallback_steps.append(("spreadsheet_tool", {})) |
| if file_name.endswith(".py"): |
| fallback_steps.append(("python_tool", {})) |
| if file_name.endswith(tuple(TEXT_EXTENSIONS)): |
| fallback_steps.append(("attachment_text_tool", {})) |
|
|
| fallback_steps.extend( |
| [ |
| ("direct_answer_tool", {}), |
| ("sports_tool", {}), |
| ("wikipedia_tool", {"query": state.get("type_query") or question}), |
| ] |
| ) |
|
|
| urls = extract_urls(question) |
| if urls: |
| if any(is_youtube_url(url) for url in urls): |
| fallback_steps.append(("video_tool", {})) |
| else: |
| return self._read_urls_observation(state) |
|
|
| if file_name.endswith(tuple(AUDIO_VIDEO_EXTENSIONS)): |
| fallback_steps.append(("audio_tool" if file_name.endswith(".mp3") else "video_tool", {})) |
| if file_name.endswith(tuple(IMAGE_EXTENSIONS)): |
| fallback_steps.append(("vision_tool", {})) |
|
|
| tried = [] |
| for index, (tool_name, args) in enumerate(fallback_steps, start=1): |
| if index > MAX_AGENT_STEPS: |
| break |
| observation = execute_tool(tool_name, args, self._context(state)) |
| tried.append(self._compact_observation(observation)) |
| if self._observation_is_useful(observation): |
| observation["fallback_tried"] = tried |
| return observation |
|
|
| search_observation = self._search_and_read_observation(state) |
| search_observation["fallback_tried"] = tried |
| return search_observation |
|
|
| def _read_urls_observation(self, state: GaiaWorkflowState) -> dict[str, Any]: |
| urls = [url for url in extract_urls(state["question"]) if not is_youtube_url(url)] |
| if not urls: |
| return execute_tool("video_tool", {}, self._context(state)) |
|
|
| observations = [] |
| for url in urls[:2]: |
| observations.append(execute_tool("web_read_tool", {"url": url}, self._context(state))) |
| evidence = "\n\n".join( |
| f"URL {index}: {item.get('evidence', '')}" |
| for index, item in enumerate(observations, start=1) |
| ) |
| errors = [item.get("error", "") for item in observations if item.get("error")] |
| return { |
| "tool": "web_url_workflow", |
| "ok": any(item.get("ok") for item in observations), |
| "answer": None, |
| "confidence": "medium" if any(item.get("ok") for item in observations) else "low", |
| "source": "web_url_workflow", |
| "evidence": truncate_text(evidence, MAX_TOOL_OUTPUT_CHARS), |
| "error": self._join_error(*errors), |
| } |
|
|
| def _search_and_read_observation(self, state: GaiaWorkflowState) -> dict[str, Any]: |
| query = state.get("type_query") or state["question"] |
| search_observation = execute_tool( |
| "web_search_tool", |
| {"query": query, "max_results": 5}, |
| self._context(state), |
| ) |
| evidence_parts = [str(search_observation.get("evidence") or "")] |
| errors = [str(search_observation.get("error") or "")] |
|
|
| try: |
| search_results = json.loads(str(search_observation.get("evidence") or "[]")) |
| except json.JSONDecodeError: |
| search_results = [] |
|
|
| for result in search_results[:2]: |
| url = result.get("url", "") |
| if not url or is_youtube_url(url): |
| continue |
| read_observation = execute_tool("web_read_tool", {"url": url}, self._context(state)) |
| evidence_parts.append( |
| f"--- {result.get('title', url)} ({url}) ---\n{read_observation.get('evidence', '')}" |
| ) |
| if read_observation.get("error"): |
| errors.append(str(read_observation["error"])) |
|
|
| return { |
| "tool": "web_search_workflow", |
| "ok": bool(search_results), |
| "answer": None, |
| "confidence": "medium" if search_results else "low", |
| "source": "web_search_workflow", |
| "evidence": truncate_text("\n\n".join(evidence_parts), MAX_TOOL_OUTPUT_CHARS), |
| "error": self._join_error(*errors), |
| } |
|
|
| def _context(self, state: GaiaWorkflowState) -> dict[str, str]: |
| return { |
| "question": state["question"], |
| "task_id": state.get("task_id", ""), |
| "file_name": state.get("file_name", ""), |
| } |
|
|
| def _observation_is_useful(self, observation: dict[str, Any]) -> bool: |
| if observation.get("answer"): |
| return True |
| return bool(observation.get("ok") and observation.get("evidence")) |
|
|
| def _append_trace( |
| self, |
| state: GaiaWorkflowState, |
| event: dict[str, Any], |
| ) -> list[dict[str, Any]]: |
| return list(state.get("trace", [])) + [event] |
|
|
| def _compact_observation(self, observation: dict[str, Any]) -> dict[str, Any]: |
| compact = dict(observation) |
| if compact.get("evidence"): |
| compact["evidence"] = truncate_text(str(compact["evidence"]), MAX_TOOL_OUTPUT_CHARS) |
| if compact.get("fallback_tried"): |
| compact["fallback_tried"] = [ |
| self._compact_observation(item) for item in compact["fallback_tried"] |
| ] |
| return compact |
|
|
| def _trace_text(self, trace: list[dict[str, Any]]) -> str: |
| return truncate_text(json.dumps(trace, ensure_ascii=False, indent=2), MAX_TOOL_OUTPUT_CHARS) |
|
|
| def _join_error(self, *errors: Any) -> str: |
| return "; ".join(str(error) for error in errors if str(error or "").strip()) |
|
|
|
|
| |
| BasicAgent = GaiaAgent |
|
|