init
This commit is contained in:
@@ -19,11 +19,13 @@ class ConversationAgent(BaseAgent):
|
||||
|
||||
workflow.add_node("analyze_intent", self._analyze_intent)
|
||||
workflow.add_node("normalize_input", self._normalize_input)
|
||||
workflow.add_node("generate_sql", self._generate_sql)
|
||||
workflow.add_node("generate_response", self._generate_response)
|
||||
workflow.add_node("update_context", self._update_context)
|
||||
|
||||
workflow.add_edge("analyze_intent", "normalize_input")
|
||||
workflow.add_edge("normalize_input", "generate_response")
|
||||
workflow.add_edge("normalize_input", "generate_sql")
|
||||
workflow.add_edge("generate_sql", "generate_response")
|
||||
workflow.add_edge("generate_response", "update_context")
|
||||
workflow.add_edge("update_context", END)
|
||||
|
||||
@@ -61,6 +63,11 @@ class ConversationAgent(BaseAgent):
|
||||
state.current_step = "response_generated"
|
||||
return state
|
||||
|
||||
def _generate_sql(self, state: AgentState) -> AgentState:
|
||||
"""生成 SQL"""
|
||||
from . import nodes
|
||||
return nodes.generate_sql(state, self.model)
|
||||
|
||||
def _update_context(self, state: AgentState) -> AgentState:
|
||||
"""更新对话上下文与历史"""
|
||||
for message in state.messages:
|
||||
|
||||
+7
-1
@@ -20,10 +20,12 @@ class BaseAgent:
|
||||
|
||||
workflow.add_node("process_input", nodes.process_input)
|
||||
workflow.add_node("normalize_input", self._normalize_input)
|
||||
workflow.add_node("generate_sql", self._generate_sql)
|
||||
workflow.add_node("generate_response", self._generate_response)
|
||||
|
||||
workflow.add_edge("process_input", "normalize_input")
|
||||
workflow.add_edge("normalize_input", "generate_response")
|
||||
workflow.add_edge("normalize_input", "generate_sql")
|
||||
workflow.add_edge("generate_sql", "generate_response")
|
||||
workflow.add_edge("generate_response", END)
|
||||
|
||||
workflow.set_entry_point("process_input")
|
||||
@@ -38,6 +40,10 @@ class BaseAgent:
|
||||
"""规范化用户输入"""
|
||||
return nodes.normalize_input(state, self.model)
|
||||
|
||||
def _generate_sql(self, state: AgentState) -> AgentState:
|
||||
"""生成 SQL"""
|
||||
return nodes.generate_sql(state, self.model)
|
||||
|
||||
def run(self, user_input: str, **kwargs) -> Dict[str, Any]:
|
||||
"""运行代理并处理用户输入"""
|
||||
initial_state = AgentState(
|
||||
|
||||
+44
-5
@@ -1,7 +1,11 @@
|
||||
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage
|
||||
import json
|
||||
|
||||
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage, AIMessage
|
||||
from .state import AgentState
|
||||
from services.prompt_manager import PromptManager
|
||||
from services.template_matcher import TemplateMatcher
|
||||
from services.prompt_manager import get_prompt_manager
|
||||
from services.template_matcher import get_template_matcher
|
||||
from services.sql_prompt_manager import SqlPromptManager
|
||||
from tools.sr_api_tool import SrApiQueryTool
|
||||
|
||||
|
||||
def process_input(state: AgentState) -> AgentState:
|
||||
@@ -12,6 +16,14 @@ def process_input(state: AgentState) -> AgentState:
|
||||
|
||||
def generate_response(state: AgentState, model) -> AgentState:
|
||||
"""使用 LLM 生成回复"""
|
||||
sr_api_result = state.context.get("sr_api_result")
|
||||
if sr_api_result:
|
||||
state.messages.append(AIMessage(content=str(sr_api_result)))
|
||||
return state
|
||||
final_sql = state.context.get("final_sql")
|
||||
if final_sql:
|
||||
state.messages.append(AIMessage(content=final_sql))
|
||||
return state
|
||||
if state.messages:
|
||||
response = model.invoke(state.messages)
|
||||
state.messages.append(response)
|
||||
@@ -27,7 +39,7 @@ def normalize_input(state: AgentState, model) -> AgentState:
|
||||
if not isinstance(last_message, HumanMessage):
|
||||
return state
|
||||
|
||||
prompt_manager = PromptManager()
|
||||
prompt_manager = get_prompt_manager()
|
||||
system_prompt = SystemMessage(
|
||||
content=prompt_manager.get("system", "english_normalizer")
|
||||
)
|
||||
@@ -38,6 +50,33 @@ def normalize_input(state: AgentState, model) -> AgentState:
|
||||
state.context["original_input"] = last_message.content
|
||||
state.context["normalized_input"] = normalized
|
||||
|
||||
matcher = TemplateMatcher()
|
||||
matcher = get_template_matcher()
|
||||
state.context["table_match"] = matcher.match(normalized)
|
||||
return state
|
||||
|
||||
|
||||
def generate_sql(state: AgentState, model) -> AgentState:
|
||||
"""根据表名与提示词生成 SQL"""
|
||||
table_match = state.context.get("table_match") or {}
|
||||
table_name = table_match.get("table_name")
|
||||
normalized = state.context.get("normalized_input")
|
||||
|
||||
if not table_name or not normalized:
|
||||
return state
|
||||
|
||||
prompt_manager = SqlPromptManager()
|
||||
prompt_data = prompt_manager.get_prompt(table_name)
|
||||
if not prompt_data:
|
||||
return state
|
||||
|
||||
prompt_text = json.dumps(prompt_data, ensure_ascii=False, indent=2)
|
||||
system_template = get_prompt_manager().get("system", "sql_mysql_select_only")
|
||||
system_content = system_template.format(table_prompt_json=prompt_text)
|
||||
user_content = f"User question (normalized English): {normalized}"
|
||||
response = model.invoke([SystemMessage(content=system_content), HumanMessage(content=user_content)])
|
||||
sql_text = response.content if hasattr(response, "content") else str(response)
|
||||
|
||||
state.context["final_sql"] = sql_text
|
||||
tool = SrApiQueryTool()
|
||||
state.context["sr_api_result"] = tool.run(json.dumps({"sql": sql_text}, ensure_ascii=False))
|
||||
return state
|
||||
|
||||
+8
-1
@@ -28,10 +28,12 @@ class ToolAgent(BaseAgent):
|
||||
workflow = StateGraph(AgentState)
|
||||
|
||||
workflow.add_node("normalize_input", self._normalize_input)
|
||||
workflow.add_node("generate_sql", self._generate_sql)
|
||||
workflow.add_node("agent", self._agent_node)
|
||||
workflow.add_node("tools", self.tool_node)
|
||||
|
||||
workflow.add_edge("normalize_input", "agent")
|
||||
workflow.add_edge("normalize_input", "generate_sql")
|
||||
workflow.add_edge("generate_sql", "agent")
|
||||
workflow.add_edge("tools", "agent")
|
||||
|
||||
workflow.add_conditional_edges(
|
||||
@@ -65,6 +67,11 @@ class ToolAgent(BaseAgent):
|
||||
|
||||
return state
|
||||
|
||||
def _generate_sql(self, state: AgentState) -> AgentState:
|
||||
"""生成 SQL"""
|
||||
from . import nodes
|
||||
return nodes.generate_sql(state, self.model)
|
||||
|
||||
def _should_use_tools(self, state: AgentState) -> str:
|
||||
"""判断是否需要使用工具"""
|
||||
last_message = state.messages[-1]
|
||||
|
||||
Reference in New Issue
Block a user