125 lines
6.3 KiB
Python
125 lines
6.3 KiB
Python
|
|
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,
|
||
|
|
}
|