diff --git a/backend/package/yuxi/repositories/subagent_repository.py b/backend/package/yuxi/repositories/subagent_repository.py index 40858429..a8ae08de 100644 --- a/backend/package/yuxi/repositories/subagent_repository.py +++ b/backend/package/yuxi/repositories/subagent_repository.py @@ -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, - ) diff --git a/backend/package/yuxi/services/subagent_service.py b/backend/package/yuxi/services/subagent_service.py index d06121d3..18bcce49 100644 --- a/backend/package/yuxi/services/subagent_service.py +++ b/backend/package/yuxi/services/subagent_service.py @@ -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 diff --git a/backend/test/test_subagent.py b/backend/test/test_subagent.py index 40c3ef75..d2c26580 100644 --- a/backend/test/test_subagent.py +++ b/backend/test/test_subagent.py @@ -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]