feat: add system tools and state updates

This commit is contained in:
Xin Wang
2026-08-04 15:27:39 +08:00
parent 5bf5987fe4
commit 731a372df9
33 changed files with 1787 additions and 124 deletions

View File

@@ -154,6 +154,46 @@ class BrainRegistryTests(unittest.TestCase):
)
self.assertIn("user_name", assistant.dynamic_variable_definitions)
def test_system_tools_are_prompt_pipeline_only(self):
assistant = AssistantUpsert(
name="prompt",
type="prompt",
systemTools=["end_conversation", "skip_turn", "end_conversation"],
)
self.assertEqual(assistant.system_tools, ["end_conversation", "skip_turn"])
workflow = AssistantUpsert(
name="workflow",
type="workflow",
systemTools=["update_state"],
graph={},
)
self.assertEqual(workflow.system_tools, [])
with self.assertRaises(ValueError):
AssistantUpsert(
name="realtime prompt",
type="prompt",
runtimeMode="realtime",
systemTools=["skip_turn"],
)
def test_system_tools_reject_unknown_kind(self):
with self.assertRaises(ValueError):
AssistantUpsert(
name="prompt",
type="prompt",
systemTools=["update_state", "magic"],
)
def test_prompt_update_state_requires_a_declared_variable(self):
with self.assertRaisesRegex(ValueError, "必须声明至少一个动态变量"):
AssistantUpsert(
name="prompt",
type="prompt",
systemTools=["update_state"],
)
def test_workflow_keeps_dynamic_variables_and_tool_bindings(self):
assistant = AssistantUpsert(
name="workflow",
@@ -757,7 +797,7 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(prompts, ["面板状态true"])
self.assertEqual(queued_frames, [])
async def test_end_call_tool_is_owned_by_prompt_brain(self):
async def test_end_conversation_system_tool_is_owned_by_prompt_brain(self):
brain = build_brain(
AssistantConfig(
type="prompt",
@@ -766,9 +806,10 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
id="end-call",
name="结束通话",
function_name="end_call",
type="end_call",
type="system",
definition={
"config": {
"kind": "end_conversation",
"message_type": "none",
"capture_reason": True,
}
@@ -792,8 +833,13 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
id="end-call",
name="结束通话",
function_name="end_call",
type="end_call",
definition={"config": {"capture_reason": True}},
type="system",
definition={
"config": {
"kind": "end_conversation",
"capture_reason": True,
}
},
)
],
),
@@ -819,13 +865,18 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
self.assertTrue(call_end.response_started)
self.assertEqual(params.result["action"], "ending_call")
async def test_end_call_waits_for_prompt_generated_closing_speech(self):
async def test_system_end_conversation_waits_for_generated_closing_speech(self):
tool = RuntimeTool(
id="end-call",
name="结束通话",
function_name="end_call",
type="end_call",
definition={"config": {"message_type": "none"}},
type="system",
definition={
"config": {
"kind": "end_conversation",
"message_type": "none",
}
},
)
cfg = AssistantConfig(type="prompt", tools=[tool])
brain = build_brain(cfg)
@@ -855,6 +906,279 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
self.assertTrue(call_end.armed)
self.assertTrue(call_end.waited_for_text)
async def test_system_tools_register_with_fixed_names(self):
cfg = AssistantConfig(
type="prompt",
dynamic_variable_definitions={
"user_name": {
"type": "string",
"required": False,
"default": None,
}
},
system_tools=[
"end_conversation",
"update_state",
"skip_turn",
"request_human_handoff",
],
)
brain = build_brain(cfg)
llm = FakeLLM()
visible_tools = []
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
llm=llm,
queue_frame=lambda _frame: None,
set_system_prompt=lambda _prompt: None,
set_tools=lambda tools: visible_tools.extend(tools or []),
call_end=FakeCallEnd(),
),
)
self.assertEqual(
[tool.name for tool in visible_tools],
[
"end_conversation",
"update_state",
"skip_turn",
"request_human_handoff",
],
)
update_schema = next(
tool for tool in visible_tools if tool.name == "update_state"
)
self.assertEqual(
update_schema.properties,
{
"user_name": {
"type": "string",
"description": "更新动态变量 user_name。",
}
},
)
for name in (
"end_conversation",
"update_state",
"skip_turn",
"request_human_handoff",
):
self.assertIn(name, llm.functions)
async def test_end_conversation_ends_call_after_generated_speech(self):
cfg = AssistantConfig(
type="prompt",
system_tools=["end_conversation"],
)
brain = build_brain(cfg)
llm = FakeLLM()
call_end = FakeCallEnd()
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
llm=llm,
queue_frame=lambda _frame: None,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=call_end,
),
)
await brain.on_assistant_text_start("turn-1")
params = FakeFunctionParams({"reason": "用户已完成咨询"})
await llm.functions["end_conversation"](params)
self.assertEqual(call_end.reason, "用户已完成咨询")
self.assertTrue(call_end.ending)
self.assertFalse(call_end.finished)
self.assertEqual(params.result["action"], "ending_call")
self.assertFalse(params.properties.run_llm)
await brain.on_assistant_text_end(
"turn-1",
"好的,再见。",
False,
)
self.assertFalse(call_end.finished)
self.assertTrue(call_end.armed)
async def test_end_conversation_finishes_when_no_speech(self):
cfg = AssistantConfig(
type="prompt",
system_tools=["end_conversation"],
)
brain = build_brain(cfg)
llm = FakeLLM()
call_end = FakeCallEnd()
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
llm=llm,
queue_frame=lambda _frame: None,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=call_end,
),
)
await brain.on_assistant_text_start("turn-1")
await llm.functions["end_conversation"](FakeFunctionParams({}))
await brain.on_assistant_text_end("turn-1", "", False)
self.assertTrue(call_end.finished)
async def test_update_state_updates_declared_variables(self):
cfg = AssistantConfig(
type="prompt",
prompt="当前用户名: {{user_name}}",
dynamic_variable_definitions={
"user_name": {"type": "string", "required": False, "default": None}
},
system_tools=["update_state"],
)
brain = build_brain(cfg)
llm = FakeLLM()
prompts = []
queued_frames = []
async def collect(frame):
queued_frames.append(frame)
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
llm=llm,
queue_frame=collect,
set_system_prompt=prompts.append,
set_tools=lambda _tools: None,
call_end=FakeCallEnd(),
),
)
params = FakeFunctionParams({"user_name": "王小明"})
await llm.functions["update_state"](params)
self.assertEqual(params.result["status"], "success")
self.assertEqual(params.result["changed"], ["user_name"])
self.assertIsNone(params.properties)
self.assertEqual(
[frame.message for frame in queued_frames],
[
{
"type": "session-variables",
"reason": "update_state",
"variables": {"user_name": "王小明"},
"changed": ["user_name"],
}
],
)
self.assertTrue(prompts[-1].endswith("当前用户名: 王小明"))
async def test_update_state_rejects_undeclared_variable(self):
cfg = AssistantConfig(
type="prompt",
system_tools=["update_state"],
)
brain = build_brain(cfg)
llm = FakeLLM()
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
llm=llm,
queue_frame=lambda _frame: None,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=FakeCallEnd(),
),
)
params = FakeFunctionParams({"unknown_key": "x"})
await llm.functions["update_state"](params)
self.assertEqual(params.result["status"], "error")
self.assertIn("未声明", params.result["message"])
async def test_skip_turn_suppresses_response(self):
cfg = AssistantConfig(
type="prompt",
system_tools=["skip_turn"],
)
brain = build_brain(cfg)
llm = FakeLLM()
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
llm=llm,
queue_frame=lambda _frame: None,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=FakeCallEnd(),
),
)
params = FakeFunctionParams({"reason": "无需回复"})
await llm.functions["skip_turn"](params)
self.assertEqual(params.result["action"], "skip_turn")
self.assertEqual(params.result["reason"], "无需回复")
self.assertFalse(params.properties.run_llm)
async def test_request_human_handoff_keeps_call_available(self):
cfg = AssistantConfig(
type="prompt",
system_tools=["request_human_handoff"],
)
brain = build_brain(cfg)
llm = FakeLLM()
call_end = FakeCallEnd()
queued_frames = []
async def collect(frame):
queued_frames.append(frame)
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
llm=llm,
queue_frame=collect,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=call_end,
),
)
await brain.on_assistant_text_start("turn-1")
params = FakeFunctionParams({"reason": "用户要求人工客服"})
await llm.functions["request_human_handoff"](params)
self.assertEqual(
[frame.message for frame in queued_frames],
[
{
"type": "handoff-requested",
"source": "prompt-system-tool",
"reason": "用户要求人工客服",
"message": "用户请求转接人工服务。",
}
],
)
self.assertEqual(call_end.reason, "")
self.assertFalse(call_end.ending)
self.assertEqual(params.result["action"], "human_handoff_requested")
self.assertIsNone(params.properties)
await brain.on_assistant_text_end("turn-1", "", False)
self.assertFalse(call_end.finished)
async def test_http_tool_renders_secrets_and_updates_prompt_variable(self):
requests = []
@@ -962,6 +1286,195 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
async def test_agent_system_tools_are_scoped_and_update_declared_state(self):
cfg = prepare_dynamic_config(
AssistantConfig(
type="workflow",
graph={
"specVersion": 3,
"settings": {},
"nodes": [
{"id": "start", "type": "start", "data": {}},
{
"id": "agent",
"type": "agent",
"data": {
"systemTools": [
"update_state",
"skip_turn",
"request_human_handoff",
],
"stateVariableNames": ["customer_name"],
},
},
],
"edges": [],
},
dynamic_variable_definitions={
"customer_name": {
"type": "string",
"required": False,
"default": None,
},
"internal_note": {
"type": "string",
"required": False,
"default": None,
},
},
),
{},
assistant_id="asst_workflow_system_tools",
)
brain = WorkflowBrain(cfg)
queued = []
async def queue_frame(frame):
queued.append(frame)
brain._runtime = BrainRuntime(
context=LLMContext(messages=[]),
llm=FakeLLM(),
queue_frame=queue_frame,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=FakeCallEnd(),
)
brain._agent_stage = SimpleNamespace(
node_config=lambda node_id, *, functions, leading_messages: {
"name": node_id,
"functions": functions,
"leading_messages": leading_messages,
}
)
brain._emit_variables = AsyncMock()
brain._refresh_agent_prompt = AsyncMock()
node_config = brain._agent_config("agent")
functions = {tool.name: tool for tool in node_config["functions"]}
self.assertEqual(
list(functions),
["update_state", "skip_turn", "request_human_handoff"],
)
self.assertEqual(
list(functions["update_state"].properties),
["customer_name"],
)
result = await functions["update_state"].handler(
{"customer_name": "李白"},
None,
)
self.assertEqual(result["changed"], ["customer_name"])
self.assertEqual(brain._store.public_values(), {"customer_name": "李白"})
brain._emit_variables.assert_awaited_once_with(
reason="update_state",
node_id="agent",
changed=["customer_name"],
)
brain._refresh_agent_prompt.assert_awaited_once_with("agent")
unauthorized = await functions["update_state"].handler(
{"internal_note": "不可写"},
None,
)
self.assertEqual(unauthorized["status"], "error")
self.assertIn("未获当前节点授权", unauthorized["message"])
handoff = await functions["request_human_handoff"].handler(
{"reason": "用户要求人工"},
None,
)
self.assertEqual(handoff["status"], "requested")
self.assertFalse(brain._runtime.call_end.ending)
event = next(
frame.message
for frame in queued
if isinstance(frame, OutputTransportMessageUrgentFrame)
)
self.assertEqual(event["type"], "handoff-requested")
self.assertEqual(event["source"], "workflow-system-tool")
async def test_update_state_node_applies_assignments_atomically(self):
cfg = prepare_dynamic_config(
AssistantConfig(
type="workflow",
graph={
"specVersion": 3,
"settings": {},
"nodes": [
{"id": "start", "type": "start", "data": {}},
{
"id": "set_state",
"type": "update_state",
"data": {
"assignments": {
"confirmed": True,
"display_name": "{{source_name}}",
}
},
},
],
"edges": [],
},
dynamic_variable_definitions={
"source_name": {"type": "string", "default": "李白"},
"display_name": {
"type": "string",
"required": False,
"default": None,
},
"confirmed": {"type": "boolean", "default": False},
},
),
{},
assistant_id="asst_workflow_update_state_node",
)
brain = WorkflowBrain(cfg)
brain._emit_node_active = AsyncMock()
brain._emit_variables = AsyncMock()
await brain._enter_update_state("set_state")
self.assertEqual(brain._store.public_values()["display_name"], "李白")
self.assertTrue(brain._store.public_values()["confirmed"])
brain._emit_variables.assert_awaited_once_with(
reason="update_state",
node_id="set_state",
changed=["confirmed", "display_name"],
)
async def test_workflow_end_conversation_waits_for_generated_speech(self):
brain = WorkflowBrain(
{
"specVersion": 3,
"settings": {},
"nodes": [{"id": "start", "type": "start", "data": {}}],
"edges": [],
}
)
call_end = FakeCallEnd()
brain._runtime = BrainRuntime(
context=LLMContext(messages=[]),
llm=FakeLLM(),
queue_frame=noop_queue_frame,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=call_end,
)
tool = brain._workflow_end_conversation_tool()
await brain.on_assistant_text_start("turn-1")
result = await tool.handler({"reason": "用户告别"}, None)
self.assertEqual(result["action"], "ending_call")
self.assertTrue(call_end.ending)
self.assertEqual(call_end.reason, "用户告别")
self.assertTrue(getattr(tool.handler, "_suppress_followup_llm"))
await brain.on_assistant_text_end("turn-1", "再见。", False)
self.assertTrue(call_end.armed)
self.assertTrue(call_end.waited_for_text)
async def test_flow_manager_dispatches_native_vision_without_auxiliary_handler(self):
manager = object.__new__(ConfiguredFlowManager)
fallback_transition = AsyncMock()

View File

@@ -360,6 +360,13 @@ class WorkflowGraphTests(unittest.TestCase):
def test_agent_effective_config_inherits_then_switches_to_override(self):
graph = valid_graph()
agent = next(node for node in graph["nodes"] if node["id"] == "agent")
agent["data"].update(
{
"systemTools": ["update_state", "skip_turn"],
"stateVariableNames": ["customer", "order_status"],
}
)
graph["settings"].update(
{
"defaultLlmResourceId": "llm_global",
@@ -389,6 +396,11 @@ class WorkflowGraphTests(unittest.TestCase):
)
self.assertTrue(engine.uses_vision())
self.assertEqual(inherited.tool_ids, ("tool_global",))
self.assertEqual(inherited.system_tools, ("update_state", "skip_turn"))
self.assertEqual(
inherited.state_variable_names,
("customer", "order_status"),
)
self.assertEqual(inherited.knowledge_mode, "on_demand")
self.assertFalse(inherited.enable_interrupt)
self.assertEqual(
@@ -417,6 +429,7 @@ class WorkflowGraphTests(unittest.TestCase):
self.assertIsNone(custom.vision_model_resource_id)
self.assertFalse(engine.uses_vision())
self.assertEqual(custom.tool_ids, ("tool_agent",))
self.assertEqual(custom.system_tools, ("update_state", "skip_turn"))
self.assertEqual(custom.knowledge_mode, "disabled")
self.assertTrue(custom.enable_interrupt)
self.assertEqual(
@@ -424,6 +437,84 @@ class WorkflowGraphTests(unittest.TestCase):
"smart_turn",
)
def test_update_state_node_and_agent_state_scope_are_validated(self):
graph = valid_graph()
graph["nodes"].insert(
1,
{
"id": "set_state",
"type": "update_state",
"data": {
"name": "记录状态",
"assignments": {"confirmed": True},
},
},
)
graph["edges"][0]["target"] = "set_state"
graph["edges"].insert(
1,
{
"id": "after_state",
"source": "set_state",
"target": "agent",
"data": {"mode": "always", "priority": 0},
},
)
agent = next(node for node in graph["nodes"] if node["id"] == "agent")
agent["data"].update(
{
"systemTools": ["update_state"],
"stateVariableNames": ["customer"],
}
)
body = AssistantUpsert(
name="状态工作流",
type="workflow",
dynamicVariableDefinitions={
"confirmed": {
"type": "boolean",
"required": False,
"default": False,
},
"customer": {
"type": "string",
"required": False,
"default": None,
},
},
graph=graph,
)
_validate_workflow(body)
normalized_state = next(
node for node in body.graph["nodes"] if node["type"] == "update_state"
)
self.assertEqual(normalized_state["data"]["assignments"], {"confirmed": True})
normalized_state["data"]["assignments"] = {"missing": "x"}
with self.assertRaisesRegex(HTTPException, "未声明变量:missing"):
_validate_workflow(body)
def test_agent_update_state_tool_requires_an_authorized_variable(self):
graph = valid_graph()
agent = next(node for node in graph["nodes"] if node["id"] == "agent")
agent["data"]["systemTools"] = ["update_state"]
body = AssistantUpsert(
name="缺少授权",
type="workflow",
dynamicVariableDefinitions={
"customer": {
"type": "string",
"required": False,
"default": None,
}
},
graph=graph,
)
with self.assertRaisesRegex(HTTPException, "必须授权至少一个变量"):
_validate_workflow(body)
def test_vision_resource_creates_isolated_runtime_config(self):
base = AssistantConfig(type="workflow", model="text-only")
resource = RuntimeModelResource(