Skip to content

Commit 009cc55

Browse files
committed
Merge branch 'feat-skill-tool-subagent-choice' into 'master'
feat: tool skill subagent choice See merge request 2026seiii-016/agent-base!31
2 parents 37b6cf9 + 3c53cb8 commit 009cc55

17 files changed

Lines changed: 1055 additions & 101 deletions

‎.gitignore‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,4 +5,6 @@
55
.env
66

77
__pycache__/
8-
*.py[cod]
8+
*.py[cod]
9+
Dockerfile.test
10+
pyproject.toml

‎pyproject.toml‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,17 @@ dependencies = [
1818
]
1919

2020

21+
[project.optional-dependencies]
22+
test = [
23+
"pytest>=8.0.0",
24+
"pytest-asyncio>=0.24.0",
25+
"anyio[trio]>=4.0.0",
26+
]
27+
28+
[tool.pytest.ini_options]
29+
asyncio_mode = "strict"
30+
pythonpath = ["."]
31+
2132
[[tool.uv.index]]
2233
name = "nju"
2334
url = "https://mirror.nju.edu.cn/pypi/web/simple/"

‎src/agents/base.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -27,13 +27,13 @@ def __init__(
2727
metadata: Optional[Dict[str, Any]] = None,
2828
allowed_tool_names: Optional[Iterable[str]] = None,
2929
allowed_tool_prefixes: Optional[Iterable[str]] = None,
30-
user_allowed_tools: Optional[List[str]] = None,
30+
user_disabled_tools: Optional[List[str]] = None,
3131
) -> None:
3232
self.agent_id = agent_id
3333
self.tool_registry = tool_registry
3434
self.prompt_provider = prompt_provider
3535
self.metadata = metadata or {}
36-
self.user_allowed_tools = list(user_allowed_tools) if user_allowed_tools is not None else None
36+
self.user_disabled_tools = list(user_disabled_tools) if user_disabled_tools is not None else None
3737

3838
default = _ALLOCATION.get(agent_id, {})
3939
self.allowed_tool_names = set(
@@ -57,7 +57,7 @@ async def get_tools(self) -> List[BaseTool]:
5757
ctx = ToolCallContext(
5858
agent_id=self.agent_id,
5959
metadata=self.metadata,
60-
user_allowed_tools=self.user_allowed_tools,
60+
user_disabled_tools=self.user_disabled_tools,
6161
)
6262
tools = await self.tool_registry.get_tools_for_agent(ctx)
6363
return self._filter_tools(tools)

‎src/agents/display_names.yaml‎

Lines changed: 16 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -31,15 +31,26 @@ tools:
3131
"12306_": "铁路查询"
3232
amap_: "地图服务"
3333

34-
# 可被用户禁用的工具分组(均为 MCP 外部工具)。
35-
# id 与上方 tools 中的 key 对应,用于 GET /tools/disableable 接口及 tool_ids 过滤。
34+
# 可被用户启用/禁用的工具分组(均为 MCP 外部工具)。
35+
# id 与上方 tools 中的 key 对应,用于 GET /tools/options 接口及 selections 过滤。
36+
# amap_ 为内部基础服务,不对外暴露为可选项。
3637
disableable_groups:
3738
- id: "variflight_"
3839
name: "航班查询"
3940
description: "航班信息查询"
4041
- id: "12306_"
4142
name: "铁路查询"
4243
description: "高铁/火车票信息查询"
43-
- id: "amap_"
44-
name: "地图服务"
45-
description: "POI 搜索、路线规划等地图功能"
44+
45+
# 子智能体列表,用于 GET /tools/options 接口。
46+
subagents:
47+
- id: "traffic"
48+
description: "交通规划专家,负责长途交通(机票/火车票)的查询与规划。"
49+
- id: "local_transport"
50+
description: "市内出行专家,负责目的地内各景点间的短途交通规划。"
51+
- id: "hotel"
52+
description: "酒店住宿专家,负责搜索和推荐符合需求的住宿方案。"
53+
- id: "attraction"
54+
description: "景点游玩专家,负责按天数和偏好规划景点游览安排。"
55+
- id: "food"
56+
description: "美食餐饮专家,负责为行程安排特色餐饮推荐。"

‎src/agents/skill_registry.py‎

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
import logging
33
import os
44
from dataclasses import dataclass
5-
from typing import Optional
5+
from typing import List, Optional
66

77
import yaml
88

@@ -61,9 +61,15 @@ def all_skills(self) -> list[Skill]:
6161
_registry = SkillRegistry()
6262

6363

64-
def get_skills_section() -> str:
65-
"""返回注入 supervisor system prompt 的技能元数据文本,无技能时返回空字符串。"""
64+
def get_skills_section(disabled_ids: Optional[List[str]] = None) -> str:
65+
"""返回注入 supervisor system prompt 的技能元数据文本,无技能时返回空字符串。
66+
67+
disabled_ids 为 None 或空列表时返回全部技能;提供列表时排除其中的技能。
68+
"""
6669
skills = _registry.all_skills()
70+
if disabled_ids:
71+
disabled_set = set(disabled_ids)
72+
skills = [s for s in skills if s.id not in disabled_set]
6773
if not skills:
6874
return ""
6975
lines = [

‎src/agents/subagent_tool.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@ def __init__(
2525
prompt_provider: PromptRegistry,
2626
session_context: Dict[str, Any],
2727
bus: Optional[EventBus] = None,
28-
user_allowed_tools: Optional[List[str]] = None,
28+
user_disabled_tools: Optional[List[str]] = None,
2929
):
3030
self.name = f"delegate_{agent_key}"
3131
self.description = (
@@ -56,7 +56,7 @@ def __init__(
5656
self._prompt_provider = prompt_provider
5757
self._session_ctx = session_context
5858
self._bus = bus
59-
self._user_allowed_tools = user_allowed_tools
59+
self._user_disabled_tools = user_disabled_tools
6060

6161
async def execute(self, **kwargs) -> Any:
6262
task_text = str(kwargs.get("task", ""))
@@ -72,7 +72,7 @@ async def execute(self, **kwargs) -> Any:
7272
agent_id=self._agent_key,
7373
tool_registry=self._registry,
7474
prompt_provider=self._prompt_provider,
75-
user_allowed_tools=self._user_allowed_tools,
75+
user_disabled_tools=self._user_disabled_tools,
7676
)
7777

7878
task = TaskSchema(agent=self._agent_key, task=task_text)

‎src/api/routers/tools.py‎

Lines changed: 34 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,40 @@
11
from fastapi import APIRouter, Depends
22
from src.api.deps import verify_token
3-
from src.core.display_names import get_disableable_groups
3+
from src.agents.skill_registry import _registry as skill_registry
4+
from src.core.display_names import get_disableable_groups, get_subagents
45

56
router = APIRouter()
67

78

8-
@router.get("/tools/disableable")
9-
async def list_disableable_tools(token: str = Depends(verify_token)):
10-
return get_disableable_groups()
9+
@router.get("/tools/options")
10+
async def list_options(token: str = Depends(verify_token)):
11+
"""返回所有可供用户选择的工具、技能和子智能体。
12+
13+
每项包含 type 字段(tool/skill/subagent)。
14+
tool 类额外包含 name 字段;skill 和 subagent 只有 id 和 description。
15+
"""
16+
items = []
17+
18+
for group in get_disableable_groups():
19+
items.append({
20+
"type": "tool",
21+
"id": group["id"],
22+
"name": group["name"],
23+
"description": group["description"],
24+
})
25+
26+
for skill in skill_registry.all_skills():
27+
items.append({
28+
"type": "skill",
29+
"id": skill.id,
30+
"description": skill.description,
31+
})
32+
33+
for subagent in get_subagents():
34+
items.append({
35+
"type": "subagent",
36+
"id": subagent["id"],
37+
"description": subagent["description"],
38+
})
39+
40+
return items

‎src/core/display_names.py‎

Lines changed: 17 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -12,10 +12,11 @@
1212
_tools_exact: Dict[str, str] = {}
1313
_tools_prefixes: list[tuple[str, str]] = []
1414
_disableable_groups: list = []
15+
_subagents: list = []
1516

1617

1718
def _load() -> None:
18-
global _agents, _tools_exact, _tools_prefixes, _disableable_groups
19+
global _agents, _tools_exact, _tools_prefixes, _disableable_groups, _subagents
1920

2021
path = os.path.join(
2122
os.path.dirname(__file__), "..", "agents", "display_names.yaml"
@@ -47,6 +48,15 @@ def _load() -> None:
4748
if item and item.get("id")
4849
]
4950

51+
_subagents = [
52+
{
53+
"id": str(item.get("id", "")),
54+
"description": str(item.get("description", "")),
55+
}
56+
for item in (data.get("subagents") or [])
57+
if item and item.get("id")
58+
]
59+
5060
logger.debug(
5161
"display_names loaded: agents=%d tools_exact=%d tools_prefixes=%d disableable=%d",
5262
len(_agents),
@@ -70,8 +80,13 @@ def resolve_tool(name: str) -> str:
7080

7181

7282
def get_disableable_groups() -> list:
73-
"""返回可被用户禁用的工具分组列表,供 GET /tools/disableable 使用。"""
83+
"""返回可被用户启用/禁用的工具分组列表。"""
7484
return list(_disableable_groups)
7585

7686

87+
def get_subagents() -> list:
88+
"""返回子智能体列表,供 GET /tools/options 使用。"""
89+
return list(_subagents)
90+
91+
7792
_load()

‎src/core/orchestrator.py‎

Lines changed: 30 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@
1919
)
2020
from src.agents.skill_registry import get_skills_section
2121
from src.agents.subagent_tool import SubAgentTool
22-
from src.schemas.chat import ChatCompletionRequest, FileAttachment, UserProfile
22+
from src.schemas.chat import ChatCompletionRequest, FileAttachment, SelectionItem, UserProfile
2323
from src.services.tools import create_registry
2424
from src.travel_plan import plan_store, set_bus, set_token
2525

@@ -251,6 +251,26 @@ async def _post_profile_callback(
251251
logger.exception("profile callback failed to %s", url)
252252

253253

254+
def _parse_disabled_selections(
255+
selections: Optional[List[SelectionItem]],
256+
) -> tuple[Optional[List[str]], Optional[List[str]], Optional[List[str]]]:
257+
"""从 disabled_selections 中按 type 拆分出三个独立黑名单。
258+
259+
返回 (tool_ids, skill_ids, subagent_ids),均表示需要禁用的 id 列表。
260+
某 type 在 selections 中无条目时对应值为 None(表示该类型全部开放)。
261+
selections 为 None 或空列表时三个值均为 None(全部开放)。
262+
"""
263+
if not selections:
264+
return None, None, None
265+
tool_ids = [s.id for s in selections if s.type == "tool"] or None
266+
skill_ids = [s.id for s in selections if s.type == "skill"] or None
267+
subagent_ids = [s.id for s in selections if s.type == "subagent"] or None
268+
return tool_ids, skill_ids, subagent_ids
269+
270+
271+
_parse_selections = _parse_disabled_selections # backwards compat alias
272+
273+
254274
def _load_subagent_keys() -> List[str]:
255275
path = os.path.join(
256276
os.path.dirname(__file__), "..", "agents", "tool_allocation.yaml"
@@ -308,9 +328,14 @@ async def run_agent_workflow(
308328
trace = TraceCollector(query=user_request)
309329
set_trace(trace)
310330

311-
# ---- SubAgent tools ----
331+
# ---- 解析 disabled_selections,拆分出三类黑名单 ----
332+
tool_ids, skill_ids, subagent_ids = _parse_disabled_selections(request.disabled_selections)
333+
334+
# ---- SubAgent tools(subagent_ids 为 None 时全部创建,否则排除被禁用的) ----
312335
subagent_tools: list = []
313336
for key in _SUBAGENT_KEYS:
337+
if subagent_ids is not None and key in subagent_ids:
338+
continue
314339
worker_cls = WORKER_CLASSES.get(key)
315340
if not worker_cls:
316341
continue
@@ -321,15 +346,15 @@ async def run_agent_workflow(
321346
prompt_provider=prompt_provider,
322347
session_context=session_ctx,
323348
bus=bus,
324-
user_allowed_tools=request.tool_ids,
349+
user_disabled_tools=tool_ids,
325350
)
326351
subagent_tools.append(tool)
327352
tool_registry._tool_map[tool.name] = tool # type: ignore[attr-defined]
328353

329354
# ---- Supervisor tools ----
330355
supervisor = SupervisorAgent(
331356
"supervisor", tool_registry, prompt_provider, metadata,
332-
user_allowed_tools=request.tool_ids,
357+
user_disabled_tools=tool_ids,
333358
)
334359
supervisor_tools = await supervisor.get_tools()
335360

@@ -340,7 +365,7 @@ async def run_agent_workflow(
340365
"user_request": user_request,
341366
"metadata": metadata,
342367
"user_profile_section": _format_user_profile(request.user_profile),
343-
"skills_section": get_skills_section(),
368+
"skills_section": get_skills_section(skill_ids),
344369
}
345370
supervisor_messages = supervisor.build_plan_messages(supervisor_context)
346371
if request.files:

‎src/schemas/chat.py‎

Lines changed: 19 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,12 @@
11
from pydantic import BaseModel, Field
2-
from typing import List, Dict, Any, Optional
2+
from typing import List, Dict, Any, Literal, Optional
3+
4+
5+
class SelectionItem(BaseModel):
6+
type: Literal["tool", "skill", "subagent"] = Field(
7+
..., description="类型:tool(外部工具)、skill(技能)、subagent(子智能体)。"
8+
)
9+
id: str = Field(..., description="对应类型的标识符,与 GET /tools/options 返回的 id 一致。")
310

411

512
class ChatMessage(BaseModel):
@@ -70,8 +77,17 @@ class ChatCompletionRequest(BaseModel):
7077
callback_id: Optional[int] = Field(
7178
None, description="Callback id (question_id) for the round trace."
7279
)
73-
tool_ids: Optional[List[str]] = Field(
74-
None, description="Allowlist of tool names to expose to the LLM. None means all tools."
80+
disabled_selections: Optional[List[SelectionItem]] = Field(
81+
None,
82+
description=(
83+
"本次对话中需要禁用的工具/技能/子智能体列表。"
84+
"None 或空列表表示全部开放(默认行为)。"
85+
"传入条目后,对应的工具/技能/子智能体将被禁用,其余保持可用。"
86+
"按 type 独立过滤:tool 控制工具禁用列表,skill 控制禁用的技能,"
87+
"subagent 控制禁用的子智能体。"
88+
"某 type 在列表中无条目时,该类型默认全部开放。"
89+
"列表中不存在的 id 会被静默忽略。"
90+
),
7591
)
7692
files: Optional[List[FileAttachment]] = Field(
7793
None, description="Base64-encoded file attachments injected into the last user message."

0 commit comments

Comments
 (0)