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 名称。 BasicAgent = GaiaAgent