import html
import json
import re
from typing import Any, Dict, List
from langchain_core.messages import HumanMessage, SystemMessage, AIMessage
from .state import AgentState
from config import Config
from services.core.prompt_manager import get_prompt_manager
from services.core.template_matcher import get_template_matcher
from services.core.sql_prompt_manager import get_sql_prompt_manager
from tools.sr_api_tool import SrApiQueryTool
def _short(value, max_len: int = 500) -> str:
text = str(value)
return text if len(text) <= max_len else text[:max_len] + "..."
def _trace(state: AgentState, label: str, value: Any | None = None, *, clip: bool = True) -> None:
if not state.context.get("debug_node_trace"):
return
if value is None:
print(label)
else:
print(label, _short(value) if clip else value)
FOLLOW_UP_HINTS = (
"那",
"那么",
"然后",
"改成",
"改为",
"换成",
"只看",
"那如果",
"how about",
"what about",
"same",
"also",
)
TOPN_RE = re.compile(r"\btop\s*(\d+)\b", re.IGNORECASE)
TOPN_CN_RE = re.compile(r"前\s*(?:\d+|[一二三四五六七八九十百千万]+)")
NUMBER_RE = re.compile(r"\b\d{8,14}\b")
DATE_RE = re.compile(r"\b\d{4}-\d{2}-\d{2}\b|\b\d{1,2}/\d{1,2}(?:/\d{2,4})?\b")
AGGREGATE_ENGLISH_RE = re.compile(
r"\b(?:count|summary|summarize|aggregate|sum)\b|\bgroup\s+by\b",
re.IGNORECASE,
)
HISTORY_ENGLISH_RE = re.compile(r"\b(?:history|historical|changelog)\b", re.IGNORECASE)
def _has_topn_hint(text: str) -> bool:
lowered = (text or "").lower()
if TOPN_RE.search(lowered):
return True
if TOPN_CN_RE.search(text or ""):
return True
return any(token in (text or "") for token in ["排名", "最高", "最大", "最小"])
def _last_human_message(state: AgentState) -> HumanMessage | None:
for message in reversed(state.messages):
if isinstance(message, HumanMessage):
return message
return None
def _looks_like_follow_up(text: str) -> bool:
lowered = (text or "").strip().lower()
return any(hint in lowered for hint in FOLLOW_UP_HINTS)
def _detect_query_mode(text: str) -> str:
lowered = (text or "").lower()
if not lowered:
return "detail"
if HISTORY_ENGLISH_RE.search(lowered) or any(token in lowered for token in ["历史", "变更记录", "历史变更", "change record", "change log", "changes"]):
return "history"
if AGGREGATE_ENGLISH_RE.search(lowered) or any(token in lowered for token in ["聚合", "统计", "汇总", "计数", "数量", "多少", "几个", "分组", "求和", "总计", "合计"]):
return "aggregate"
if _has_topn_hint(text):
return "topn"
return "detail"
def _extract_query_entities(text: str, prompt_data: Dict[str, Any] | None = None) -> Dict[str, Any]:
lowered = (text or "").lower()
entity_numbers = NUMBER_RE.findall(text or "")
entity_dates = DATE_RE.findall(text or "")
top_n = None
match = TOPN_RE.search(lowered)
if match:
try:
top_n = int(match.group(1))
except Exception:
top_n = None
elif _has_topn_hint(text):
top_n = 10
sort_direction = "desc"
if any(token in lowered for token in ["从小到大", "升序", "ascending", "asc"]):
sort_direction = "asc"
elif any(token in lowered for token in ["从大到小", "降序", "descending", "desc"]):
sort_direction = "desc"
countries: List[str] = []
regions: List[str] = []
if prompt_data:
additional_fields = (((prompt_data.get("field_mapping_reference") or {}).get("additional_fields") or {}))
countries = list((((additional_fields.get("ship_to_country") or {}).get("values")) or []))
regions = list((((additional_fields.get("region") or {}).get("values")) or []))
words = re.findall(r"\b[A-Z]{2,10}\b", text or "")
matched_countries = [word for word in words if word in countries]
matched_regions = [word for word in words if word in regions]
return {
"numbers": entity_numbers,
"dates": entity_dates,
"top_n": top_n,
"sort_direction": sort_direction,
"country_codes": matched_countries,
"regions": matched_regions,
"mentions_eta_info": "eta信息" in lowered or "eta info" in lowered,
"mentions_history": any(token in lowered for token in ["history", "historical", "changelog", "历史", "变更记录", "历史变更"]),
}
def _get_default_table_name() -> str | None:
cfg = Config.get_section("ragflow")
table_name = str(cfg.get("default_table_name") or "").strip()
return table_name or None
def _looks_like_json(text: str) -> bool:
stripped = (text or "").strip()
return stripped.startswith("{") or stripped.startswith("[")
def _try_json_loads(value: Any) -> Any:
if isinstance(value, (dict, list)):
return value
if isinstance(value, str) and _looks_like_json(value):
try:
return json.loads(value)
except Exception:
return value
return value
def _parse_sr_api_result(raw_result: Any) -> Any:
parsed = _try_json_loads(raw_result)
if isinstance(parsed, dict) and "text" in parsed:
text_payload = _try_json_loads(parsed.get("text"))
parsed = {**parsed, "text": text_payload}
return parsed
def _extract_result_rows(value: Any) -> list[Any] | None:
if isinstance(value, list):
return value
if isinstance(value, dict):
for key in ("data", "rows", "records", "items", "list", "result", "values"):
rows = value.get(key)
if isinstance(rows, list):
return rows
nested = value.get("text")
if isinstance(nested, (dict, list)):
return _extract_result_rows(nested)
return None
def _is_empty_sr_api_result(raw_result: Any) -> bool:
parsed = _parse_sr_api_result(raw_result)
rows = _extract_result_rows(parsed)
if rows is not None:
return len(rows) == 0
if isinstance(parsed, dict):
total = parsed.get("total")
if isinstance(total, int):
return total == 0
text_payload = parsed.get("text")
if isinstance(text_payload, dict):
total = text_payload.get("total")
if isinstance(total, int):
return total == 0
return False
def _format_empty_result_response(query: str, llm_response: str) -> str:
"""将空结果的 LLM 回复格式化为纯文本格式,与 endpoints._build_rich_answer_html 保持一致"""
safe_query = html.escape((query or "").strip())
safe_response = html.escape((llm_response or "").strip())
return (
f"Question: {safe_query}\n"
f"{safe_response}\n"
f"Rows: 0"
)
def _default_normalizer_prompt() -> str:
return (
"You are a translation and normalization assistant. "
"Convert the user's input to a clear, grammatically correct English sentence suitable for SQL intent. "
"Preserve business identifiers, codes, country abbreviations, order numbers, and field aliases exactly when possible. "
"Return only the final English sentence without extra explanations."
)
def process_input(state: AgentState) -> AgentState:
"""处理用户输入"""
_trace(state, "[process_input][in] messages=", _short(state.messages))
last_message = _last_human_message(state)
if last_message:
state.original_input = str(last_message.content)
state.context["is_follow_up"] = _looks_like_follow_up(state.original_input)
state.set_current_step("processed")
_trace(state, "[process_input][out] current_step=", state.current_step)
return state
def normalize_input(state: AgentState, model) -> AgentState:
"""将用户输入规范化为保留业务标识的标准英文语句"""
last_message = _last_human_message(state)
if not last_message:
return state
_trace(state, "[normalize_input][in] user_input=", _short(last_message.content))
prompt_manager = get_prompt_manager()
normalizer_prompt = (
prompt_manager.get("user", "english_normalizer")
or prompt_manager.get("system", "english_normalizer")
or _default_normalizer_prompt()
)
system_prompt = SystemMessage(content=normalizer_prompt)
try:
response = model.invoke([system_prompt, HumanMessage(content=last_message.content)])
normalized = response.content if hasattr(response, "content") else str(response)
except Exception as exc:
normalized = str(last_message.content)
state.add_error(f"normalize_input_failed:{exc}")
state.original_input = str(last_message.content)
state.normalized_input = normalized.strip() or str(last_message.content)
state.sync_context()
state.set_current_step("normalized")
_trace(state, "[normalize_input][out] normalized=", _short(state.normalized_input))
return state
def classify_query_mode(state: AgentState) -> AgentState:
"""识别查询模式:detail / aggregate / topn / history。"""
text = "\n".join(filter(None, [state.original_input, state.normalized_input]))
state.query_mode = _detect_query_mode(text)
if not state.intent:
state.intent = "sql_query" if state.original_input else "general"
state.query_entities = _extract_query_entities(text)
state.sync_context()
state.set_current_step("query_mode_classified")
_trace(state, "[classify_query_mode][out] query_mode=", state.query_mode)
_trace(state, "[classify_query_mode][out] query_entities=", _short(state.query_entities))
return state
def match_table(state: AgentState) -> AgentState:
"""根据规范化输入检索候选表,并在追问场景下回退到上一轮表或配置默认表。"""
query = state.normalized_input or state.original_input
if not query:
return state
_trace(state, "[match_table][in] query=", _short(query))
matcher = get_template_matcher()
match_result = matcher.match(query)
table_name = (match_result or {}).get("table_name")
candidate_tables = list((match_result or {}).get("candidates") or [])
if not table_name and state.context.get("is_follow_up"):
last_context = state.context.get("last_context") or {}
fallback_table = last_context.get("table_name") or ((last_context.get("table_match") or {}).get("table_name"))
if fallback_table:
table_name = fallback_table
candidate_tables = candidate_tables or [{"table_name": fallback_table, "source": "last_context"}]
match_result = {
"table_name": fallback_table,
"candidates": candidate_tables,
"raw": {"source": "last_context"},
}
state.context["table_match_fallback"] = "last_context"
if not table_name:
default_table = _get_default_table_name()
if default_table:
table_name = default_table
candidate_tables = candidate_tables or [{"table_name": default_table, "source": "config_default"}]
match_result = {
"table_name": default_table,
"candidates": candidate_tables,
"raw": {"source": "config_default"},
}
state.context["table_match_fallback"] = "config_default"
state.context["default_table_name"] = default_table
state.table_match = dict(match_result or {})
state.candidate_tables = candidate_tables
state.table_name = table_name
state.sync_context()
state.set_current_step("table_matched")
_trace(state, "[match_table][out] table_name=", state.table_name)
return state
def load_sql_prompt(state: AgentState) -> AgentState:
"""加载目标表对应的 SQL prompt JSON。"""
if not state.table_name:
_trace(state, "[load_sql_prompt][skip] missing table_name")
return state
prompt_manager = get_sql_prompt_manager()
prompt_data, source = prompt_manager.get_prompt_with_source(state.table_name)
if not prompt_data:
state.add_error(f"sql_prompt_not_found:{state.table_name}")
_trace(state, "[load_sql_prompt][skip] prompt not found for table=", state.table_name)
return state
state.sql_prompt = prompt_data
state.sql_prompt_source = source
state.sync_context()
state.set_current_step("sql_prompt_loaded")
_trace(state, f"[load_sql_prompt][out] table_name={state.table_name} source={source}")
return state
def build_sql_plan(state: AgentState) -> AgentState:
"""构建结构化 SQL 计划,为最终 SQL 生成提供显式上下文。"""
prompt_data = state.sql_prompt or {}
business_rules = (prompt_data.get("business_logic_rules") or {})
data_model = (prompt_data.get("data_model_specification") or {})
meta = (prompt_data.get("meta") or {})
default_fields = ((data_model.get("mandatory_display_fields") or {}).get("default_fields")) or ""
text = "\n".join(filter(None, [state.original_input, state.normalized_input]))
extracted = _extract_query_entities(text, prompt_data)
if state.query_entities:
extracted = {**state.query_entities, **{k: v for k, v in extracted.items() if v not in (None, [], {}, "")}}
state.query_entities = extracted
state.sql_plan = {
"intent": state.intent or "sql_query",
"query_mode": state.query_mode or "detail",
"selected_table": state.table_name,
"candidate_tables": [item.get("table_name", item) for item in state.candidate_tables],
"data_source": meta.get("data_source"),
"domain": meta.get("domain"),
"default_select_fields": default_fields,
"default_filters": list(business_rules.get("default_filters") or []),
"aggregate_rules": dict(business_rules.get("aggregate_rules") or {}),
"top_n_rules": dict(business_rules.get("top_n_rules") or {}),
"query_entities": extracted,
"previous_context": {
key: (state.context.get("last_context") or {}).get(key)
for key in ("table_name", "query_mode", "final_sql", "sql_plan")
if (state.context.get("last_context") or {}).get(key) is not None
},
}
state.sync_context()
state.set_current_step("sql_plan_built")
_trace(state, "[build_sql_plan][out] sql_plan=", _short(state.sql_plan))
return state
def generate_sql(state: AgentState, model) -> AgentState:
"""根据表 prompt + 结构化计划生成 SQL。"""
if not state.table_name or not state.normalized_input:
_trace(state, "[generate_sql][skip] missing table_name or normalized_input")
return state
prompt_data = state.sql_prompt
if not prompt_data:
_trace(state, "[generate_sql][skip] missing sql_prompt")
return state
_trace(state, "[generate_sql][in] table_name=", state.table_name)
_trace(state, "[generate_sql][in] query_mode=", state.query_mode)
prompt_text = json.dumps(prompt_data, ensure_ascii=False, indent=2)
plan_text = json.dumps(state.sql_plan or {}, ensure_ascii=False, indent=2)
prompt_manager = get_prompt_manager()
system_template = prompt_manager.get("system", "sql_mysql_select_only")
system_content = system_template.format(table_prompt_json=prompt_text)
user_content = (
f"Original user question: {state.original_input}\n"
f"Normalized user question: {state.normalized_input}\n"
f"Detected query mode: {state.query_mode or 'detail'}\n"
f"SQL planning context JSON:\n{plan_text}\n"
"Generate the best SQL for the selected table and query mode. "
"If the query mode is topn and the plan contains top_n, LIMIT is allowed and required. "
"If update_date is used as a filter, do not add data_flag. "
"Return only the final SQL."
)
response = model.invoke([SystemMessage(content=system_content), HumanMessage(content=user_content)])
sql_text = response.content if hasattr(response, "content") else str(response)
state.final_sql = sql_text.strip()
state.sync_context()
state.set_current_step("sql_generated")
_trace(state, "[generate_sql][out] sql=", state.final_sql, clip=False)
return state
def execute_sql(state: AgentState) -> AgentState:
"""在需要时执行生成后的 SQL。skip_sr_api=True 时跳过执行。"""
if not state.final_sql:
_trace(state, "[execute_sql][skip] missing final_sql")
state.set_current_step("sql_execution_skipped")
return state
if state.skip_sr_api:
_trace(state, "[execute_sql][skip] skip_sr_api=true")
state.set_current_step("sql_execution_skipped")
return state
try:
tool = SrApiQueryTool()
state.sr_api_result = tool.run(json.dumps({"sql": state.final_sql}, ensure_ascii=False))
_trace(state, "[execute_sql][out] sr_api_result=", _short(state.sr_api_result))
except Exception as exc:
state.add_error(f"sql_execution_failed:{exc}")
_trace(state, "[execute_sql][error]", exc)
state.sync_context()
state.set_current_step("sql_executed")
return state
def check_empty_result(state: AgentState) -> AgentState:
"""检查 SQL 执行结果是否为空,设置 is_empty_result 标记。"""
sr_api_result = state.sr_api_result
if not sr_api_result:
state.context["is_empty_result"] = None
state.context["result_checked"] = False
_trace(state, "[check_empty_result][skip] no sr_api_result")
state.set_current_step("result_checked")
return state
is_empty = _is_empty_sr_api_result(sr_api_result)
state.context["is_empty_result"] = is_empty
state.context["result_checked"] = True
if is_empty:
_trace(state, "[check_empty_result][out] is_empty=True")
else:
result_rows = _extract_result_rows(sr_api_result)
row_count = len(result_rows) if result_rows else 0
state.context["result_row_count"] = row_count
_trace(state, f"[check_empty_result][out] is_empty=False, row_count={row_count}")
state.sync_context()
state.set_current_step("result_checked")
return state
def generate_response(state: AgentState, model) -> AgentState:
"""使用 SQL 执行结果、SQL 本身或模型回退生成最终回复。"""
_trace(state, "[generate_response][in] context_keys=", list((state.context or {}).keys()))
# 优先使用 context 中的 is_empty_result(由 check_empty_result 节点设置)
is_empty_result = state.context.get("is_empty_result")
sr_api_result = state.sr_api_result
# 如果有执行结果且标记为空
if sr_api_result and is_empty_result is True:
sql_plan_text = json.dumps(state.sql_plan or {}, ensure_ascii=False, indent=2)
fallback_system = SystemMessage(
content=(
"You are a friendly business query assistant. "
"The query executed successfully but returned no data. "
"Answer the user in a concise and helpful way. "
"IMPORTANT RULES:\n"
"1. DO NOT show any SQL statements, technical field names, or database terminology to the user\n"
"2. Use business language that non-technical users can understand\n"
"3. Clearly state that no matching data was found\n"
"4. Provide specific suggestions about which conditions might be too restrictive\n"
"5. Use the query context to suggest alternatives, but express them in plain language\n"
"6. For example, say 'try removing the country filter' instead of 'remove ship_to_country condition'\n"
"7. For example, say 'try searching all records instead of just the latest' instead of 'remove data_flag filter'"
)
)
fallback_user = HumanMessage(
content=(
f"Original user question: {state.original_input}\n"
f"Query mode: {state.query_mode or 'detail'}\n"
f"SQL plan context (for your reference only, DO NOT show to user):\n{sql_plan_text}\n"
"Please answer the user in plain business language without any SQL or technical terms."
)
)
try:
response = model.invoke([fallback_system, fallback_user])
llm_content = response.content if hasattr(response, "content") else str(response)
state.messages.append(response)
formatted_html = _format_empty_result_response(
state.original_input,
llm_content
)
state.context["formatted_answer"] = formatted_html
state.context["response_source"] = "model_empty_result_fallback"
_trace(state, "[generate_response][out] source=model_empty_result_fallback")
except Exception as exc:
state.add_error(f"empty_result_fallback_failed:{exc}")
fixed_content = "未查询到符合条件的数据,请尝试调整筛选条件后再查询。"
state.messages.append(AIMessage(content=fixed_content))
formatted_html = _format_empty_result_response(
state.original_input,
fixed_content
)
state.context["formatted_answer"] = formatted_html
state.context["response_source"] = "empty_result_fixed_fallback"
_trace(state, "[generate_response][out] source=empty_result_fixed_fallback")
state.sync_context()
state.set_current_step("response_generated")
return state
# 有执行结果且不为空
if sr_api_result:
state.context["is_empty_result"] = False
state.context["response_source"] = "sr_api_result"
state.messages.append(AIMessage(content=str(sr_api_result)))
_trace(state, "[generate_response][out] source=sr_api_result")
state.sync_context()
state.set_current_step("response_generated")
return state
# 没有执行结果,返回 SQL(skip_sr_api=True 的情况)
final_sql = state.final_sql
if final_sql:
state.context["response_source"] = "final_sql"
state.messages.append(AIMessage(content=final_sql))
_trace(state, "[generate_response][out] source=final_sql")
state.sync_context()
state.set_current_step("response_generated")
return state
# 兜底:使用模型生成回复
if state.messages:
response = model.invoke(state.messages)
state.messages.append(response)
state.context["response_source"] = "model_invoke"
_trace(state, "[generate_response][out] source=model_invoke")
state.sync_context()
state.set_current_step("response_generated")
return state