83 lines
2.9 KiB
Python
83 lines
2.9 KiB
Python
import json
|
|
|
|
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage, AIMessage
|
|
from .state import AgentState
|
|
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:
|
|
"""处理用户输入"""
|
|
state.current_step = "processed"
|
|
return state
|
|
|
|
|
|
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)
|
|
return state
|
|
|
|
|
|
def normalize_input(state: AgentState, model) -> AgentState:
|
|
"""将用户输入规范化为标准英文语句"""
|
|
if not state.messages:
|
|
return state
|
|
|
|
last_message = state.messages[-1]
|
|
if not isinstance(last_message, HumanMessage):
|
|
return state
|
|
|
|
prompt_manager = get_prompt_manager()
|
|
system_prompt = SystemMessage(
|
|
content=prompt_manager.get("system", "english_normalizer")
|
|
)
|
|
|
|
response = model.invoke([system_prompt, HumanMessage(content=last_message.content)])
|
|
normalized = response.content if hasattr(response, "content") else str(response)
|
|
|
|
state.context["original_input"] = last_message.content
|
|
state.context["normalized_input"] = normalized
|
|
|
|
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
|