fix: isolate startup tools and preview dialogs
This commit is contained in:
@@ -108,8 +108,12 @@ class AssistantConfig(BaseModel):
|
||||
turnConfig: dict = Field(default_factory=dict)
|
||||
startup: dict = Field(default_factory=dict)
|
||||
|
||||
# Prompt assistant reusable tools. Execution remains type-specific in the pipeline.
|
||||
# ``tools`` is the complete runtime pool (conversation + lifecycle actions).
|
||||
# ``llm_tool_ids`` limits which tools are advertised to a Prompt model. None
|
||||
# preserves compatibility for inline/test configs that historically exposed
|
||||
# every item in ``tools``.
|
||||
tools: list[RuntimeTool] = Field(default_factory=list)
|
||||
llm_tool_ids: list[str] | None = None
|
||||
knowledge_base_id: str | None = None
|
||||
knowledge_base_name: str = ""
|
||||
knowledge_base_description: str = ""
|
||||
|
||||
@@ -121,19 +121,10 @@ async def _validate_startup_actions(
|
||||
actions = body.startup.actions
|
||||
if not actions:
|
||||
return
|
||||
bound_tool_ids = set(body.tool_ids)
|
||||
for action in actions:
|
||||
if action.tool_id not in bound_tool_ids:
|
||||
raise HTTPException(400, f"启动 Action 必须先绑定工具:{action.tool_id}")
|
||||
tool = await session.get(Tool, action.tool_id)
|
||||
if not tool or tool.status != "active":
|
||||
raise HTTPException(400, f"启动 Action 引用了无效工具:{action.tool_id}")
|
||||
if action.phase == "preflight" and tool.type not in {"http", "mcp"}:
|
||||
raise HTTPException(400, "preflight Action 仅支持 HTTP 或 MCP 工具")
|
||||
if action.phase == "opening" and tool.type not in {"http", "mcp", "client"}:
|
||||
raise HTTPException(400, "opening Action 不支持该工具类型")
|
||||
if tool.type != "client":
|
||||
continue
|
||||
runtime_tool = RuntimeTool(
|
||||
id=tool.id,
|
||||
name=tool.name,
|
||||
@@ -141,10 +132,20 @@ async def _validate_startup_actions(
|
||||
type=tool.type,
|
||||
definition=tool.definition or {},
|
||||
)
|
||||
if action.phase == "preflight" and runtime_tool.type not in {"http", "mcp"}:
|
||||
raise HTTPException(400, "preflight Action 仅支持 HTTP 或 MCP 工具")
|
||||
if action.phase == "opening" and runtime_tool.type not in {
|
||||
"http",
|
||||
"mcp",
|
||||
"client",
|
||||
}:
|
||||
raise HTTPException(400, "opening Action 不支持该工具类型")
|
||||
if runtime_tool.type != "client":
|
||||
continue
|
||||
policy = policy_for_tool(runtime_tool)
|
||||
if action.required and not policy.wait_for_response:
|
||||
raise HTTPException(400, "必需的 Client 启动 Action 必须等待客户端响应")
|
||||
if tool.function_name != "show_message":
|
||||
if runtime_tool.function_name != "show_message":
|
||||
continue
|
||||
if policy.response_wait_mode != "session":
|
||||
raise HTTPException(400, "show_message 启动 Action 必须使用会话内等待")
|
||||
|
||||
@@ -84,8 +84,13 @@ class PromptBrain(BaseBrain):
|
||||
self._opening_started = False
|
||||
self._opening_finished = not bool(self._startup_actions("opening"))
|
||||
self._startup_failed = False
|
||||
llm_tool_ids = (
|
||||
set(cfg.llm_tool_ids) if cfg.llm_tool_ids is not None else None
|
||||
)
|
||||
schemas: list[FunctionSchema] = []
|
||||
for tool in cfg.tools:
|
||||
if llm_tool_ids is not None and tool.id not in llm_tool_ids:
|
||||
continue
|
||||
if tool.type == "end_call":
|
||||
schema, handler = self._make_end_call_tool(tool, runtime)
|
||||
elif tool.type in {"http", "mcp", "client"}:
|
||||
|
||||
@@ -90,20 +90,54 @@ def _secret(resource: ModelResource | None, key: str, default: str = "") -> str:
|
||||
return str((resource.secrets or {}).get(key) or default)
|
||||
|
||||
|
||||
async def _tools_for(session: AsyncSession, assistant: Assistant) -> list[RuntimeTool]:
|
||||
if assistant.type not in {"prompt", "workflow"}:
|
||||
return []
|
||||
tools = (
|
||||
async def _bound_tool_ids(session: AsyncSession, assistant_id: str) -> list[str]:
|
||||
rows = (
|
||||
await session.execute(
|
||||
select(Tool)
|
||||
.join(AssistantToolBinding, AssistantToolBinding.tool_id == Tool.id)
|
||||
.where(
|
||||
AssistantToolBinding.assistant_id == assistant.id,
|
||||
Tool.status == "active",
|
||||
)
|
||||
.order_by(AssistantToolBinding.created_at, Tool.id)
|
||||
select(AssistantToolBinding.tool_id)
|
||||
.where(AssistantToolBinding.assistant_id == assistant_id)
|
||||
.order_by(AssistantToolBinding.created_at, AssistantToolBinding.tool_id)
|
||||
)
|
||||
).scalars().all()
|
||||
return [str(tool_id) for tool_id in rows]
|
||||
|
||||
|
||||
def _startup_tool_ids(startup: dict | None) -> list[str]:
|
||||
"""Return lifecycle-only tool references in configured execution order."""
|
||||
|
||||
result: list[str] = []
|
||||
for action in (startup or {}).get("actions") or []:
|
||||
if not isinstance(action, dict):
|
||||
continue
|
||||
tool_id = str(action.get("tool_id") or action.get("toolId") or "").strip()
|
||||
if tool_id and tool_id not in result:
|
||||
result.append(tool_id)
|
||||
return result
|
||||
|
||||
|
||||
def _runtime_tool_ids(llm_tool_ids: list[str], startup: dict | None) -> list[str]:
|
||||
"""Merge both tool scopes without changing either scope's meaning."""
|
||||
|
||||
return list(dict.fromkeys([*llm_tool_ids, *_startup_tool_ids(startup)]))
|
||||
|
||||
|
||||
async def _tools_for(
|
||||
session: AsyncSession,
|
||||
assistant: Assistant,
|
||||
tool_ids: list[str],
|
||||
) -> list[RuntimeTool]:
|
||||
if assistant.type not in {"prompt", "workflow"} or not tool_ids:
|
||||
return []
|
||||
rows = (
|
||||
await session.execute(
|
||||
select(Tool)
|
||||
.where(
|
||||
Tool.id.in_(tool_ids),
|
||||
Tool.status == "active",
|
||||
)
|
||||
)
|
||||
).scalars().all()
|
||||
tool_by_id = {tool.id: tool for tool in rows}
|
||||
tools = [tool_by_id[tool_id] for tool_id in tool_ids if tool_id in tool_by_id]
|
||||
server_ids = {
|
||||
str(tool.mcp_server_id)
|
||||
for tool in tools
|
||||
@@ -217,6 +251,10 @@ async def resolve_runtime_config(
|
||||
if kb.status == "active"
|
||||
}
|
||||
|
||||
llm_tool_ids = await _bound_tool_ids(session, assistant.id)
|
||||
runtime_tool_ids = _runtime_tool_ids(llm_tool_ids, assistant.startup or {})
|
||||
runtime_tools = await _tools_for(session, assistant, runtime_tool_ids)
|
||||
|
||||
return AssistantConfig(
|
||||
name=assistant.name,
|
||||
type=assistant.type,
|
||||
@@ -228,7 +266,8 @@ async def resolve_runtime_config(
|
||||
enableInterrupt=assistant.enable_interrupt,
|
||||
turnConfig=assistant.turn_config or {},
|
||||
startup=assistant.startup or {},
|
||||
tools=await _tools_for(session, assistant),
|
||||
tools=runtime_tools,
|
||||
llm_tool_ids=llm_tool_ids,
|
||||
knowledge_base_id=assistant.knowledge_base_id,
|
||||
knowledge_base_name=knowledge_base.name if knowledge_base else "",
|
||||
knowledge_base_description=knowledge_base.description if knowledge_base else "",
|
||||
|
||||
@@ -246,6 +246,55 @@ class DifyLLMServiceTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
|
||||
class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_startup_only_tool_is_not_registered_with_llm(self):
|
||||
startup_tool = RuntimeTool(
|
||||
id="opening_message",
|
||||
name="开场确认",
|
||||
function_name="show_message",
|
||||
type="client",
|
||||
)
|
||||
conversation_tool = RuntimeTool(
|
||||
id="lookup_order",
|
||||
name="查询订单",
|
||||
function_name="lookup_order",
|
||||
type="http",
|
||||
)
|
||||
cfg = AssistantConfig(
|
||||
type="prompt",
|
||||
tools=[startup_tool, conversation_tool],
|
||||
llm_tool_ids=[conversation_tool.id],
|
||||
startup={
|
||||
"actions": [
|
||||
{
|
||||
"id": "opening_message",
|
||||
"phase": "opening",
|
||||
"tool_id": startup_tool.id,
|
||||
"required": True,
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
brain = build_brain(cfg)
|
||||
llm = FakeLLM()
|
||||
visible_schemas = []
|
||||
|
||||
await brain.setup(
|
||||
cfg,
|
||||
BrainRuntime(
|
||||
context=LLMContext(messages=[]),
|
||||
llm=llm,
|
||||
queue_frame=noop_queue_frame,
|
||||
set_system_prompt=lambda _prompt: None,
|
||||
set_tools=visible_schemas.extend,
|
||||
call_end=FakeCallEnd(),
|
||||
),
|
||||
)
|
||||
|
||||
self.assertIn(startup_tool.id, brain._tool_by_id)
|
||||
self.assertNotIn(startup_tool.function_name, llm.functions)
|
||||
self.assertIn(conversation_tool.function_name, llm.functions)
|
||||
self.assertEqual(len(visible_schemas), 1)
|
||||
|
||||
async def test_preflight_runs_multiple_server_tools_in_order(self):
|
||||
tools = [
|
||||
RuntimeTool(
|
||||
|
||||
@@ -6,14 +6,19 @@ from types import SimpleNamespace
|
||||
from fastapi import HTTPException
|
||||
from routes.assistants import _validate_startup_actions
|
||||
from schemas import AssistantUpsert
|
||||
from services.config_resolver import _runtime_tool_ids
|
||||
|
||||
|
||||
def startup_body(*, phase: str = "opening") -> AssistantUpsert:
|
||||
def startup_body(
|
||||
*,
|
||||
phase: str = "opening",
|
||||
bind_tool: bool = False,
|
||||
) -> AssistantUpsert:
|
||||
return AssistantUpsert(
|
||||
name="启动动作测试",
|
||||
type="prompt",
|
||||
runtimeMode="pipeline",
|
||||
toolIds=["tool_message"],
|
||||
toolIds=["tool_message"] if bind_tool else [],
|
||||
startup={
|
||||
"executionMode": "sequential",
|
||||
"actions": [
|
||||
@@ -45,20 +50,52 @@ class FakeSession:
|
||||
self.tool = tool
|
||||
|
||||
async def get(self, _model, tool_id):
|
||||
return self.tool if tool_id == self.tool.id else None
|
||||
return self.tool if self.tool is not None and tool_id == self.tool.id else None
|
||||
|
||||
|
||||
class StartupActionValidationTests(unittest.IsolatedAsyncioTestCase):
|
||||
def test_runtime_pool_includes_bound_and_startup_only_tools(self):
|
||||
self.assertEqual(
|
||||
_runtime_tool_ids(
|
||||
["conversation_tool", "shared_tool"],
|
||||
{
|
||||
"actions": [
|
||||
{"tool_id": "opening_tool"},
|
||||
{"toolId": "shared_tool"},
|
||||
]
|
||||
},
|
||||
),
|
||||
["conversation_tool", "shared_tool", "opening_tool"],
|
||||
)
|
||||
|
||||
def test_realtime_rejects_startup_actions(self):
|
||||
with self.assertRaisesRegex(ValueError, "Realtime"):
|
||||
AssistantUpsert(
|
||||
name="Realtime 启动动作",
|
||||
type="prompt",
|
||||
runtimeMode="realtime",
|
||||
toolIds=["tool_message"],
|
||||
startup=startup_body().startup,
|
||||
)
|
||||
|
||||
async def test_startup_tool_does_not_need_conversation_binding(self):
|
||||
tool = SimpleNamespace(
|
||||
id="tool_message",
|
||||
name="重要提示",
|
||||
function_name="show_message",
|
||||
type="client",
|
||||
status="active",
|
||||
definition={
|
||||
"config": {
|
||||
"wait_for_response": True,
|
||||
"response_wait_mode": "session",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
body = startup_body(bind_tool=False)
|
||||
await _validate_startup_actions(FakeSession(tool), body)
|
||||
self.assertEqual(body.tool_ids, [])
|
||||
|
||||
async def test_show_message_requires_session_wait(self):
|
||||
tool = SimpleNamespace(
|
||||
id="tool_message",
|
||||
|
||||
Reference in New Issue
Block a user