Files
more_dots/tests/test_sql_workflow_refactor.py
T
2026-03-24 18:07:22 +08:00

369 lines
15 KiB
Python

from langchain_core.messages import HumanMessage
from agent.agents.conversation import ConversationAgent
from agent.core import nodes
from agent.core.state import AgentState
from services.core.prompt_manager import PromptManager
from services.core.sql_prompt_manager import SqlPromptManager
from workflows.workflow_manager import WorkflowManager, WorkflowType
class FakeResponse:
def __init__(self, content: str, tool_calls=None):
self.content = content
self.tool_calls = tool_calls or []
class FakeModel:
def invoke(self, messages):
if len(messages) == 2 and getattr(messages[0], "content", "").startswith("You are a translation and normalization assistant"):
return FakeResponse(messages[1].content)
if len(messages) == 2 and "table prompt JSON" in getattr(messages[0], "content", ""):
user_content = getattr(messages[1], "content", "")
lowered = user_content.lower()
if "query mode: topn" in lowered or "top 20" in lowered:
return FakeResponse("SELECT ship_to_country as country, count(soid) as qty FROM dwd_ai.apbo_eta_ful WHERE data_flag = 'Newest' GROUP BY ship_to_country ORDER BY qty DESC LIMIT 20")
if "query mode: aggregate" in lowered:
return FakeResponse("SELECT region, count(soid) as qty FROM dwd_ai.apbo_eta_ful WHERE data_flag = 'Newest' GROUP BY region ORDER BY qty DESC")
return FakeResponse("SELECT service_order_id FROM dwd_ai.apbo_eta_ful WHERE data_flag = 'Newest'")
return FakeResponse("fallback")
class FakeTemplateMatcher:
def match(self, normalized_text: str):
return {
"table_name": "apbo_eta_ful",
"candidates": [{"table_name": "apbo_eta_ful", "metadata": {"table": "apbo_eta_ful"}}],
"raw": {"query": normalized_text},
}
class FakeSqlPromptManager:
def get_prompt(self, table_name: str):
manager = SqlPromptManager()
return manager.get_prompt(table_name)
class EmptyTemplateMatcher:
def match(self, normalized_text: str):
return {
"table_name": None,
"candidates": [],
"raw": {"query": normalized_text},
}
class FirstTurnOnlyTemplateMatcher:
def match(self, normalized_text: str):
lowered = (normalized_text or "").lower()
if "4020438779" in lowered or "query so" in lowered:
return FakeTemplateMatcher().match(normalized_text)
return EmptyTemplateMatcher().match(normalized_text)
def test_prompt_manager_resolves_project_prompt_file():
manager = PromptManager()
prompt_text = manager.get("system", "sql_mysql_select_only")
assert "expert SQL generator" in prompt_text
assert "Only output a single MySQL SELECT statement." in prompt_text
def test_sql_prompt_manager_resolves_project_sql_prompt_dir():
manager = SqlPromptManager()
prompt = manager.get_prompt("apbo_eta_ful")
assert prompt is not None
assert prompt["meta"]["data_source"] == "dwd_ai.apbo_eta_ful"
def test_apbo_eta_ful_prompt_uses_identifier_specific_where_filters():
prompt = SqlPromptManager().get_prompt("apbo_eta_ful")
rule_text = prompt["business_logic_rules"]["soid_or_service_order_id"]
examples = prompt["examples"]
assert "禁止使用WHERE soid" in rule_text
assert "必须使用(service_order_id = 'xxx' or soid = 'xxx')" in rule_text
assert "service_order_id" in examples["history_records_all_fields"]["sql"]
assert " or soid " in examples["history_records_all_fields"]["sql"]
assert " or soid " in examples["specific_fields_query"]["sql"]
assert " or soid " in examples["newest_status_all_fields"]["sql"]
def test_query_mode_and_plan_build_for_topn():
prompt = SqlPromptManager().get_prompt("apbo_eta_ful")
state = AgentState(
messages=[HumanMessage(content="CC为LT ,by country 查询 top 20")],
context={"original_input": "CC为LT ,by country 查询 top 20", "normalized_input": "CC=LT by country top 20", "table_name": "apbo_eta_ful"},
)
state.sql_prompt = prompt
state = nodes.classify_query_mode(state)
state = nodes.build_sql_plan(state)
assert state.query_mode == "topn"
assert state.sql_plan["query_entities"]["top_n"] == 20
assert state.sql_plan["selected_table"] == "apbo_eta_ful"
assert "top_n_rules" in state.sql_plan
def test_query_mode_country_does_not_trigger_aggregate():
state = AgentState(
messages=[HumanMessage(content="查询 AU country 的 DC premier stock backlog")],
context={
"original_input": "查询 AU country 的 DC premier stock backlog",
"normalized_input": "List of backlogs with values for AU Country DC premier stock",
},
)
state = nodes.classify_query_mode(state)
assert state.query_mode == "detail"
def test_query_mode_explicit_aggregate_keyword_still_matches():
state = AgentState(
messages=[HumanMessage(content="按 country 统计 backlog 数量")],
context={
"original_input": "按 country 统计 backlog 数量",
"normalized_input": "Count backlog quantity by country",
},
)
state = nodes.classify_query_mode(state)
assert state.query_mode == "aggregate"
def test_query_mode_current_keyword_is_not_misclassified_as_topn():
state = AgentState(
messages=[HumanMessage(content="汇总统计当前的bo数量")],
context={
"original_input": "汇总统计当前的bo数量",
"normalized_input": "Aggregate summary of the current backlog order count",
},
)
state = nodes.classify_query_mode(state)
assert state.query_mode == "aggregate"
assert state.query_entities.get("top_n") is None
def test_match_table_falls_back_to_config_default(monkeypatch):
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: EmptyTemplateMatcher())
monkeypatch.setattr(
"agent.core.nodes.Config.get_section",
lambda section: {"default_table_name": "apbo_eta_ful"} if section == "ragflow" else {},
)
state = AgentState(
messages=[HumanMessage(content="random question")],
context={"original_input": "random question", "normalized_input": "random question"},
)
state = nodes.match_table(state)
assert state.table_name == "apbo_eta_ful"
assert state.context["table_match_fallback"] == "config_default"
assert state.context["default_table_name"] == "apbo_eta_ful"
def test_generate_response_falls_back_to_model_when_sr_api_result_empty():
class EmptyResultModel:
def invoke(self, messages):
return FakeResponse("未查询到符合条件的数据,请尝试调整筛选条件。")
state = AgentState(
messages=[HumanMessage(content="查询 AU 的 backlog")],
context={
"original_input": "查询 AU 的 backlog",
"normalized_input": "Query backlog for AU",
"final_sql": "SELECT * FROM dwd_ai.apbo_eta_ful WHERE ship_to_country='AU'",
"sr_api_result": '{"status_code": 200, "text": "{\\"code\\":\\"0\\",\\"data\\":[],\\"msg\\":\\"操作成功\\",\\"total\\":0}"}',
},
)
state.final_sql = state.context["final_sql"]
state.sr_api_result = state.context["sr_api_result"]
state = nodes.generate_response(state, EmptyResultModel())
assert state.messages[-1].content.startswith("未查询到符合条件的数据")
assert state.context["response_source"] == "model_empty_result_fallback"
assert state.context["is_empty_result"] is True
def test_conversation_agent_run_keeps_final_sql_in_context(monkeypatch):
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FakeTemplateMatcher())
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
agent = ConversationAgent()
result = agent.run("查询 SO 4020438779 的 eta 信息", skip_sr_api=True)
assert result["context"]["table_name"] == "apbo_eta_ful"
assert result["context"]["query_mode"] == "detail"
assert result["context"]["final_sql"].startswith("SELECT")
assert result["messages"][-1].content == result["context"]["final_sql"]
def test_conversation_agent_requires_explicit_memory_for_follow_up(monkeypatch):
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FirstTurnOnlyTemplateMatcher())
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
agent = ConversationAgent()
first = agent.run("查询 SO 4020438779 的 eta 信息", skip_sr_api=True)
second = agent.run("改成 by country top 20", skip_sr_api=True)
third = agent.run(
"改成 by country top 20",
skip_sr_api=True,
conversation_history=first["conversation_history"],
last_context=first["context"].get("last_context"),
)
assert first["context"]["table_name"] == "apbo_eta_ful"
assert second["context"]["table_name"] == "apbo_eta_ful"
assert second["context"]["table_match_fallback"] == "config_default"
assert third["context"]["table_name"] == "apbo_eta_ful"
assert third["context"]["table_match_fallback"] == "last_context"
assert third["context"]["query_mode"] == "topn"
assert "LIMIT 20" in third["context"]["final_sql"]
def test_workflow_manager_follow_up_does_not_reuse_session_last_context(monkeypatch):
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FirstTurnOnlyTemplateMatcher())
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
manager = WorkflowManager(enable_multi_turn=False)
first = manager.execute_workflow(
WorkflowType.CONVERSATION,
"查询 SO 4020438779 的 eta 信息",
session_id="cid-1",
skip_sr_api=True,
)
second = manager.execute_workflow(
WorkflowType.CONVERSATION,
"改成 by country top 20",
session_id="cid-1",
skip_sr_api=True,
)
first_context = first["result"]["context"]
second_context = second["result"]["context"]
session_info = manager.get_session_info("cid-1")
assert first_context["table_name"] == "apbo_eta_ful"
assert second_context["table_name"] == "apbo_eta_ful"
assert second_context["table_match_fallback"] == "config_default"
assert second_context["query_mode"] == "topn"
assert "LIMIT 20" in second_context["final_sql"]
assert "last_context" not in session_info
assert "conversation_history" not in session_info
def test_workflow_manager_follow_up_reuses_session_last_context_when_enabled(monkeypatch):
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FirstTurnOnlyTemplateMatcher())
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
manager = WorkflowManager(enable_multi_turn=True)
first = manager.execute_workflow(
WorkflowType.CONVERSATION,
"查询 SO 4020438779 的 eta 信息",
session_id="cid-enabled",
skip_sr_api=True,
)
second = manager.execute_workflow(
WorkflowType.CONVERSATION,
"改成 by country top 20",
session_id="cid-enabled",
skip_sr_api=True,
)
first_context = first["result"]["context"]
second_context = second["result"]["context"]
session_info = manager.get_session_info("cid-enabled")
assert first_context["table_name"] == "apbo_eta_ful"
assert second_context["table_name"] == "apbo_eta_ful"
assert second_context["table_match_fallback"] == "last_context"
assert second_context["query_mode"] == "topn"
assert "LIMIT 20" in second_context["final_sql"]
assert session_info["last_context"]["table_name"] == "apbo_eta_ful"
assert len(session_info["conversation_history"]) >= 4
def test_workflow_manager_isolates_conversation_memory_by_session(monkeypatch):
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FirstTurnOnlyTemplateMatcher())
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
manager = WorkflowManager(enable_multi_turn=False)
manager.execute_workflow(
WorkflowType.CONVERSATION,
"查询 SO 4020438779 的 eta 信息",
session_id="session-a",
skip_sr_api=True,
)
second = manager.execute_workflow(
WorkflowType.CONVERSATION,
"改成 by country top 20",
session_id="session-b",
skip_sr_api=True,
)
assert second["result"]["context"]["table_name"] == "apbo_eta_ful"
assert second["result"]["context"]["table_match_fallback"] == "config_default"
def test_conversation_agent_run_is_silent_without_debug_prints(monkeypatch, capsys):
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FakeTemplateMatcher())
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
agent = ConversationAgent()
agent.run("查询 SO 4020438779 的 eta 信息", skip_sr_api=True)
captured = capsys.readouterr()
assert captured.out == ""
def test_conversation_agent_run_prints_node_trace_when_enabled(monkeypatch, capsys):
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FakeTemplateMatcher())
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
agent = ConversationAgent()
agent.run("查询 SO 4020438779 的 eta 信息", skip_sr_api=True, debug_node_trace=True)
captured = capsys.readouterr()
assert "[process_input][in]" in captured.out
assert "[generate_response][out] source=final_sql" in captured.out
def test_generate_sql_trace_outputs_full_sql_when_enabled(capsys):
class LongSqlModel:
def invoke(self, messages):
return FakeResponse("SELECT " + ", ".join(f"col_{idx}" for idx in range(150)) + " FROM dwd_ai.apbo_eta_ful")
state = AgentState(
messages=[HumanMessage(content="查询 long sql")],
context={
"original_input": "查询 long sql",
"normalized_input": "Query long sql",
"table_name": "apbo_eta_ful",
"sql_plan": {"selected_table": "apbo_eta_ful"},
"debug_node_trace": True,
},
)
state.sql_prompt = SqlPromptManager().get_prompt("apbo_eta_ful")
state = nodes.generate_sql(state, LongSqlModel())
captured = capsys.readouterr()
assert "[generate_sql][out] sql=" in captured.out
assert "col_149" in captured.out
assert "..." not in captured.out.split("[generate_sql][out] sql=", 1)[1]