OhBrian's picture
更新为 LangGraph 工作流模式
d6b82a1
Raw
History Blame Contribute Delete
20.1 kB
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