Files
more_dots/agent/nodes.py
T

108 lines
4.3 KiB
Python
Raw Normal View History

2026-02-26 18:06:17 +08:00
import json
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage, AIMessage
2026-02-26 13:43:44 +08:00
from .state import AgentState
2026-02-26 18:06:17 +08:00
from services.prompt_manager import get_prompt_manager
from services.template_matcher import get_template_matcher
2026-03-02 15:35:02 +08:00
from services.sql_prompt_manager import get_sql_prompt_manager
2026-02-26 18:06:17 +08:00
from tools.sr_api_tool import SrApiQueryTool
2026-02-26 13:43:44 +08:00
2026-03-02 15:35:02 +08:00
def _short(value, max_len: int = 500) -> str:
text = str(value)
return text if len(text) <= max_len else text[:max_len] + "..."
2026-02-26 13:43:44 +08:00
def process_input(state: AgentState) -> AgentState:
"""处理用户输入"""
2026-03-02 15:35:02 +08:00
print("[process_input][in] messages=", _short(state.messages))
2026-02-26 13:43:44 +08:00
state.current_step = "processed"
2026-03-02 15:35:02 +08:00
print("[process_input][out] current_step=", state.current_step)
2026-02-26 13:43:44 +08:00
return state
def generate_response(state: AgentState, model) -> AgentState:
"""使用 LLM 生成回复"""
2026-03-02 15:35:02 +08:00
print("[generate_response][in] context_keys=", list((state.context or {}).keys()))
2026-02-26 18:06:17 +08:00
sr_api_result = state.context.get("sr_api_result")
if sr_api_result:
state.messages.append(AIMessage(content=str(sr_api_result)))
2026-03-02 15:35:02 +08:00
print("[generate_response][out] source=sr_api_result")
2026-02-26 18:06:17 +08:00
return state
final_sql = state.context.get("final_sql")
if final_sql:
state.messages.append(AIMessage(content=final_sql))
2026-03-02 15:35:02 +08:00
print("[generate_response][out] source=final_sql")
2026-02-26 18:06:17 +08:00
return state
2026-02-26 13:43:44 +08:00
if state.messages:
response = model.invoke(state.messages)
state.messages.append(response)
2026-03-02 15:35:02 +08:00
print("[generate_response][out] source=model_invoke")
2026-02-26 13:43:44 +08:00
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
2026-03-02 15:35:02 +08:00
print("[normalize_input][in] user_input=", _short(last_message.content))
2026-02-26 18:06:17 +08:00
prompt_manager = get_prompt_manager()
2026-03-02 15:35:02 +08:00
normalizer_prompt = (
prompt_manager.get("system", "english_normalizer")
or prompt_manager.get("user", "english_normalizer")
2026-02-26 13:43:44 +08:00
)
2026-03-02 15:35:02 +08:00
system_prompt = SystemMessage(content=normalizer_prompt)
2026-02-26 13:43:44 +08:00
response = model.invoke([system_prompt, HumanMessage(content=last_message.content)])
normalized = response.content if hasattr(response, "content") else str(response)
2026-03-02 15:35:02 +08:00
print("[normalize_input][out] normalized=", _short(normalized))
2026-02-26 13:43:44 +08:00
state.context["original_input"] = last_message.content
state.context["normalized_input"] = normalized
2026-02-26 18:06:17 +08:00
matcher = get_template_matcher()
2026-02-26 13:43:44 +08:00
state.context["table_match"] = matcher.match(normalized)
2026-03-02 15:35:02 +08:00
print("[normalize_input][out] table_match=", _short(state.context.get("table_match")))
2026-02-26 13:43:44 +08:00
return state
2026-02-26 18:06:17 +08:00
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:
2026-03-02 15:35:02 +08:00
print("[generate_sql][skip] missing table_name or normalized")
2026-02-26 18:06:17 +08:00
return state
2026-03-02 15:35:02 +08:00
print("[generate_sql][in] table_name=", table_name)
print("[generate_sql][in] normalized=", _short(normalized))
prompt_manager = get_sql_prompt_manager()
2026-02-26 18:06:17 +08:00
prompt_data = prompt_manager.get_prompt(table_name)
if not prompt_data:
2026-03-02 15:35:02 +08:00
print("[generate_sql][skip] prompt not found for table=", table_name)
2026-02-26 18:06:17 +08:00
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)
2026-03-02 15:35:02 +08:00
print("[generate_sql][out] sql=", _short(sql_text))
2026-02-26 18:06:17 +08:00
state.context["final_sql"] = sql_text
2026-03-02 15:35:02 +08:00
if not state.context.get("skip_sr_api"):
tool = SrApiQueryTool()
state.context["sr_api_result"] = tool.run(json.dumps({"sql": sql_text}, ensure_ascii=False))
print("[generate_sql][out] sr_api_result=", _short(state.context.get("sr_api_result")))
2026-02-26 18:06:17 +08:00
return state