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