init
This commit is contained in:
@@ -0,0 +1,368 @@
|
||||
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]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user