from dataclasses import dataclass, field from typing import Any, Dict, List, Optional from langchain_core.messages import BaseMessage @dataclass class AgentState: """代理工作流的状态定义""" messages: List[BaseMessage] = field(default_factory=list) current_step: str = "start" context: Dict[str, Any] = field(default_factory=dict) intent: Optional[str] = None original_input: str = "" normalized_input: str = "" query_mode: str = "" query_entities: Dict[str, Any] = field(default_factory=dict) candidate_tables: List[Dict[str, Any]] = field(default_factory=list) table_match: Dict[str, Any] = field(default_factory=dict) table_name: Optional[str] = None sql_prompt: Dict[str, Any] = field(default_factory=dict) sql_prompt_source: str = "" sql_plan: Dict[str, Any] = field(default_factory=dict) final_sql: str = "" sr_api_result: Any = None skip_sr_api: bool = False validation_errors: List[str] = field(default_factory=list) errors: List[str] = field(default_factory=list) def __post_init__(self) -> None: self.context = dict(self.context or {}) self.messages = list(self.messages or []) self.intent = self.context.get("intent", self.intent) self.original_input = str(self.context.get("original_input") or self.original_input or "") self.normalized_input = str(self.context.get("normalized_input") or self.normalized_input or "") self.query_mode = str(self.context.get("query_mode") or self.query_mode or "") self.query_entities = dict(self.context.get("query_entities") or self.query_entities or {}) self.candidate_tables = list(self.context.get("candidate_tables") or self.candidate_tables or []) self.table_match = dict(self.context.get("table_match") or self.table_match or {}) self.table_name = self.context.get("table_name") or self.table_name or self.table_match.get("table_name") self.sql_prompt = dict(self.context.get("sql_prompt") or self.sql_prompt or {}) self.sql_prompt_source = str(self.context.get("sql_prompt_source") or self.sql_prompt_source or "") self.sql_plan = dict(self.context.get("sql_plan") or self.sql_plan or {}) self.final_sql = str(self.context.get("final_sql") or self.final_sql or "") self.sr_api_result = self.context.get("sr_api_result", self.sr_api_result) self.skip_sr_api = bool(self.context.get("skip_sr_api", self.skip_sr_api)) self.validation_errors = list(self.context.get("validation_errors") or self.validation_errors or []) self.errors = list(self.context.get("errors") or self.errors or []) self.sync_context() def sync_context(self) -> Dict[str, Any]: """将显式状态字段回写到兼容 context。""" self.context["current_step"] = self.current_step self.context["skip_sr_api"] = self.skip_sr_api optional_values = { "intent": self.intent, "original_input": self.original_input, "normalized_input": self.normalized_input, "query_mode": self.query_mode, "query_entities": self.query_entities, "candidate_tables": self.candidate_tables, "table_match": self.table_match, "table_name": self.table_name, "sql_prompt": self.sql_prompt, "sql_prompt_source": self.sql_prompt_source, "sql_plan": self.sql_plan, "final_sql": self.final_sql, "sr_api_result": self.sr_api_result, "validation_errors": self.validation_errors, "errors": self.errors, } for key, value in optional_values.items(): empty = value in (None, "", [], {}) if empty: self.context.pop(key, None) else: self.context[key] = value return self.context def set_current_step(self, step: str) -> None: self.current_step = step self.sync_context() def add_error(self, message: str) -> None: if message and message not in self.errors: self.errors.append(message) self.sync_context() def apply_graph_result(self, result: Any) -> "AgentState": """兼容 LangGraph 返回 dict 或 AgentState 两种形式。""" if isinstance(result, AgentState): return result if isinstance(result, dict): self.messages = result.get("messages", self.messages) self.current_step = result.get("current_step", self.current_step) self.context.update(result.get("context", {})) self.intent = self.context.get("intent") self.original_input = str(self.context.get("original_input") or self.original_input) self.normalized_input = str(self.context.get("normalized_input") or self.normalized_input) self.query_mode = str(self.context.get("query_mode") or self.query_mode) self.query_entities = dict(self.context.get("query_entities") or self.query_entities) self.candidate_tables = list(self.context.get("candidate_tables") or self.candidate_tables) self.table_match = dict(self.context.get("table_match") or self.table_match) self.table_name = self.context.get("table_name") or self.table_name or self.table_match.get("table_name") self.sql_prompt = dict(self.context.get("sql_prompt") or self.sql_prompt) self.sql_prompt_source = str(self.context.get("sql_prompt_source") or self.sql_prompt_source) self.sql_plan = dict(self.context.get("sql_plan") or self.sql_plan) self.final_sql = str(self.context.get("final_sql") or self.final_sql) self.sr_api_result = self.context.get("sr_api_result", self.sr_api_result) self.skip_sr_api = bool(self.context.get("skip_sr_api", self.skip_sr_api)) self.validation_errors = list(self.context.get("validation_errors") or self.validation_errors) self.errors = list(self.context.get("errors") or self.errors) self.sync_context() return self def to_result(self) -> Dict[str, Any]: """输出与现有 API 兼容的结果结构。""" self.sync_context() return { "messages": self.messages, "current_step": self.current_step, "context": self.context, }