refactor(subagents): 重构服务层,统一工具解析与过滤逻辑

- 合并 normalize/filter 函数为 filter_specs_by_names,移除冗余类型检查
- exists_name 改用 SELECT COUNT(*) 仅查计数
- update 改用字典迭代批量赋值,删除未使用的 upsert
- 新增 list_all_specs 到 repository 层
- 修复 get_subagents_from_names 空列表返回 tuple 的 bug
- 更新相关测试
This commit is contained in:
Wenjie Zhang 2026-03-20 01:31:00 +08:00
parent 446eb652bb
commit 89aed450ab
3 changed files with 129 additions and 70 deletions

View File

@ -18,14 +18,23 @@ class SubAgentRepository:
result = await self.db.execute(select(SubAgent).order_by(SubAgent.updated_at.desc()))
return list(result.scalars().all())
async def list_all_specs(self) -> list[dict[str, Any]]:
"""获取所有 SubAgent 运行规格,按 updated_at 降序"""
items = await self.list_all()
return [item.to_subagent_spec() for item in items]
async def get_by_name(self, name: str) -> SubAgent | None:
"""根据名称获取 SubAgent"""
result = await self.db.execute(select(SubAgent).where(SubAgent.name == name))
return result.scalar_one_or_none()
async def exists_name(self, name: str) -> bool:
"""检查名称是否存在"""
return (await self.get_by_name(name)) is not None
"""检查名称是否存在(仅查询计数,不获取完整数据)"""
from sqlalchemy import select, func
result = await self.db.execute(
select(func.count()).select_from(SubAgent).where(SubAgent.name == name)
)
return result.scalar() > 0
async def create(
self,
@ -67,14 +76,17 @@ class SubAgentRepository:
model_provided: bool = False,
updated_by: str | None,
) -> SubAgent:
if description is not None:
item.description = description
if system_prompt is not None:
item.system_prompt = system_prompt
if tools is not None:
item.tools = tools
# 批量更新非空字段
updates = {
"description": description,
"system_prompt": system_prompt,
"tools": tools,
}
if model_provided:
item.model = model
updates["model"] = model
for field, value in updates.items():
if value is not None:
setattr(item, field, value)
item.updated_by = updated_by
item.updated_at = utc_now_naive()
await self.db.commit()
@ -85,28 +97,3 @@ class SubAgentRepository:
"""删除 SubAgent"""
await self.db.delete(item)
await self.db.commit()
async def upsert(self, data: dict[str, Any], created_by: str | None) -> SubAgent:
"""Upsert 操作,如果存在则更新,否则创建"""
name = data["name"]
existing = await self.get_by_name(name)
if existing:
return await self.update(
existing,
description=data.get("description", existing.description),
system_prompt=data.get("system_prompt", existing.system_prompt),
tools=data.get("tools", existing.tools),
model=data.get("model", existing.model),
model_provided="model" in data,
updated_by=created_by,
)
else:
return await self.create(
name=name,
description=data["description"],
system_prompt=data["system_prompt"],
tools=data.get("tools"),
model=data.get("model"),
is_builtin=data.get("is_builtin", False),
created_by=created_by,
)

View File

@ -7,8 +7,8 @@ from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.repositories.subagent_repository import SubAgentRepository
from yuxi.services.mcp_service import get_tools_from_all_servers
from yuxi.storage.postgres.manager import pg_manager
from yuxi.utils import logger
# SubAgent specs cache for get_subagent_specs
_subagent_specs_cache: list[dict[str, Any]] | None = None
@ -91,46 +91,47 @@ async def get_subagent_specs(db: AsyncSession | None = None) -> list[dict[str, A
return deepcopy(_subagent_specs_cache)
async with _get_session(db) as session:
repo = SubAgentRepository(session)
subagents = await repo.list_all()
_subagent_specs_cache = [sa.to_subagent_spec() for sa in subagents]
_subagent_specs_cache = await repo.list_all_specs()
return deepcopy(_subagent_specs_cache)
def invalidate_subagent_specs_cache() -> None:
def clear_specs_cache() -> None:
"""清除 subagent specs 缓存"""
global _subagent_specs_cache
_subagent_specs_cache = None
async def get_subagents_from_names(selected_names: Any, *, db: AsyncSession | None = None) -> list[dict[str, Any]]:
"""根据名称获取 subagent specs含工具解析"""
specs = await get_subagent_specs(db)
def resolve_subagent_tools(specs: list[dict[str, Any]], available_tools: list[Any]) -> list[dict[str, Any]]:
"""将 subagent specs 中的工具名称解析为实际工具实例"""
available_by_name = {tool.name: tool for tool in available_tools if hasattr(tool, "name")}
if not selected_names:
return []
selected_set = set(selected_names)
available = {spec["name"] for spec in specs if isinstance(spec.get("name"), str)}
matched = [spec for spec in specs if spec.get("name") in selected_set]
missing = [n for n in selected_names if n not in available]
if missing:
logger.warning(f"Configured subagents not found, skip: {missing}")
# 处理工具
# 仅从子智能体配置中的工具名称进行解析;不做 Tavily/MCP 特殊注入。
from yuxi.agents.common.toolkits import get_all_tool_instances
all_tools = get_all_tool_instances()
all_tool_names = {tool.name: tool for tool in all_tools}
resolved_specs = []
for spec in specs:
for spec in matched:
resolved_spec = dict(spec)
tool_names = spec.get("tools", [])
resolved_spec["tools"] = [
available_by_name[name] for name in tool_names if isinstance(name, str) and name in available_by_name
all_tool_names[name] for name in tool_names if name in all_tool_names
]
resolved_specs.append(resolved_spec)
return resolved_specs
async def _get_available_tools() -> list[Any]:
"""获取所有可用的工具实例"""
from yuxi.agents.common.toolkits.buildin.tools import _create_tavily_search
tools = []
# 添加 tavily_search 工具
tavily = _create_tavily_search()
if tavily:
tools.append(tavily)
# 添加 MCP 工具
mcp_tools = await get_tools_from_all_servers()
tools.extend(mcp_tools)
return tools
async def get_all_subagents(db: AsyncSession | None = None) -> list[dict[str, Any]]:
"""获取所有 SubAgent含禁用的"""
async with _get_session(db) as session:
@ -164,7 +165,7 @@ async def create_subagent(
is_builtin=False,
created_by=created_by,
)
invalidate_subagent_specs_cache()
clear_specs_cache()
return item.to_dict()
@ -191,7 +192,7 @@ async def update_subagent(
model_provided="model" in data,
updated_by=updated_by,
)
invalidate_subagent_specs_cache()
clear_specs_cache()
return item.to_dict()
@ -205,5 +206,5 @@ async def delete_subagent(name: str, db: AsyncSession | None = None) -> bool:
if item.is_builtin:
raise ValueError("内置 SubAgent 不可删除")
await repo.delete(item)
invalidate_subagent_specs_cache()
clear_specs_cache()
return True

View File

@ -415,16 +415,12 @@ class TestSubAgentService:
"tools": ["tool_a"],
}
class MockSubAgent:
def to_subagent_spec(self):
return mock_spec
class MockRepo:
def __init__(self, session):
pass
async def list_all(self):
return [MockSubAgent()]
async def list_all_specs(self):
return [mock_spec]
@asynccontextmanager
async def mock_session_context(*args, **kwargs):
@ -435,7 +431,6 @@ class TestSubAgentService:
monkeypatch.setattr(service_module, "SubAgentRepository", MockRepo)
monkeypatch.setattr(service_module, "pg_manager", MockPgManager())
monkeypatch.setattr(service_module, "_get_available_tools", AsyncMock(return_value=[]))
result = await service_module.get_subagent_specs()
@ -460,7 +455,7 @@ class TestSubAgentService:
second = await service_module.get_subagent_specs()
assert second[0]["tools"] == ["tool_a"]
service_module.invalidate_subagent_specs_cache()
service_module.clear_specs_cache()
def test_resolve_subagent_tools_does_not_mutate_input(self):
from yuxi.services import subagent_service as service_module
@ -551,3 +546,79 @@ class TestSubAgentModel:
spec = agent.to_subagent_spec()
assert "model" not in spec
class TestDeepAgentSubagentSelection:
def test_filter_specs_by_names_empty_selection_returns_empty(self):
from yuxi.services.subagent_service import filter_specs_by_names
specs = [
{"name": "research-agent", "description": "r"},
{"name": "critique-agent", "description": "c"},
]
filtered, missing = filter_specs_by_names(specs, [])
assert filtered == []
assert missing == []
def test_filter_specs_by_names_none_selection_returns_all(self):
from yuxi.services.subagent_service import filter_specs_by_names
specs = [
{"name": "research-agent", "description": "r"},
{"name": "critique-agent", "description": "c"},
]
filtered, missing = filter_specs_by_names(specs, None)
assert filtered == specs
assert missing == []
def test_filter_specs_by_names_returns_subset_and_missing(self):
from yuxi.services.subagent_service import filter_specs_by_names
specs = [
{"name": "research-agent", "description": "r"},
{"name": "critique-agent", "description": "c"},
{"name": 123, "description": "invalid"},
]
filtered, missing = filter_specs_by_names(
specs,
["research-agent", "missing-agent", "research-agent", ""],
)
assert [item["name"] for item in filtered] == ["research-agent"]
assert missing == ["missing-agent"]
@pytest.mark.asyncio
async def test_get_subagents_from_names_filters_and_resolves_tools(self, monkeypatch):
from yuxi.services import subagent_service as service_module
async def fake_get_specs(_db=None):
return [
{
"name": "research-agent",
"description": "r",
"system_prompt": "s",
"tools": ["tool_a"],
},
{
"name": "critique-agent",
"description": "c",
"system_prompt": "s",
"tools": [],
},
]
mock_tool = MagicMock()
mock_tool.name = "tool_a"
monkeypatch.setattr(service_module, "get_subagent_specs", fake_get_specs)
monkeypatch.setattr("yuxi.agents.common.toolkits.get_all_tool_instances", lambda: [mock_tool])
resolved_specs = await service_module.get_subagents_from_names(["research-agent", "missing-agent"])
assert [item["name"] for item in resolved_specs] == ["research-agent"]
assert resolved_specs[0]["tools"] == [mock_tool]