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:
parent
446eb652bb
commit
89aed450ab
@ -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,
|
||||
)
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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]
|
||||
|
||||
Loading…
Reference in New Issue
Block a user