refactor: 会话级 skill resolver 与 skills 加载链路重构
This commit is contained in:
parent
c1afed8d3d
commit
326f87d67f
@ -35,6 +35,8 @@ Skills 管理模块用于集中维护可供 Agent 只读引用的技能包。
|
||||
## Agent 运行时行为
|
||||
|
||||
1. `context.skills` 用于配置技能 slug 列表。
|
||||
2. 运行时仅暴露选中 skills 到 `/skills/<slug>/...`。
|
||||
3. `/skills` 路径只读,不允许写入、编辑、上传。
|
||||
4. 变更在下一次对话请求生效。
|
||||
2. 运行时按会话构建 `SkillResolver` 快照(同一会话首次构建,后续复用)。
|
||||
3. 运行时仅暴露快照中的可见 skills 到 `/skills/<slug>/...`。
|
||||
4. `/skills` 路径只读,不允许写入、编辑、上传。
|
||||
5. 同会话内若 `context.skills` 变化会触发快照重建。
|
||||
6. 后台修改 skills 内容后,已有会话不会自动刷新,需新会话或调整 `context.skills` 才生效。
|
||||
|
||||
@ -75,7 +75,7 @@ async def get_default_agent(current_user: User = Depends(get_required_user)):
|
||||
default_agent_id = conf.default_agent_id
|
||||
# 如果没有设置默认智能体,尝试获取第一个可用的智能体
|
||||
if not default_agent_id:
|
||||
agents = await agent_manager.get_agents_info()
|
||||
agents = await agent_manager.get_agents_info(include_configurable_items=False)
|
||||
if agents:
|
||||
default_agent_id = agents[0].get("id", "")
|
||||
|
||||
@ -94,7 +94,7 @@ async def set_default_agent(request_data: dict = Body(...), current_user=Depends
|
||||
raise HTTPException(status_code=422, detail="缺少必需的 agent_id 字段")
|
||||
|
||||
# 验证智能体是否存在
|
||||
agents = await agent_manager.get_agents_info()
|
||||
agents = await agent_manager.get_agents_info(include_configurable_items=False)
|
||||
agent_ids = [agent.get("id", "") for agent in agents]
|
||||
|
||||
if agent_id not in agent_ids:
|
||||
@ -137,7 +137,7 @@ async def call(query: str = Body(...), meta: dict = Body(None), current_user: Us
|
||||
@chat.get("/agent")
|
||||
async def get_agent(current_user: User = Depends(get_required_user)):
|
||||
"""获取所有可用智能体的基本信息(需要登录)"""
|
||||
agents_info = await agent_manager.get_agents_info()
|
||||
agents_info = await agent_manager.get_agents_info(include_configurable_items=False)
|
||||
|
||||
# Return agents with basic information (without configurable_items for performance)
|
||||
agents = [
|
||||
|
||||
@ -76,10 +76,11 @@ async def list_skills_route(
|
||||
@skills.get("/dependency-options")
|
||||
async def get_skill_dependency_options_route(
|
||||
_current_user: User = Depends(get_superadmin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取 skill 依赖项可选列表(仅超级管理员)。"""
|
||||
try:
|
||||
return {"success": True, "data": get_skill_dependency_options()}
|
||||
return {"success": True, "data": await get_skill_dependency_options(db)}
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get skill dependency options: {e}")
|
||||
raise HTTPException(status_code=500, detail="获取 skill 依赖选项失败")
|
||||
|
||||
@ -4,7 +4,7 @@ from fastapi import FastAPI
|
||||
|
||||
from src.services.task_service import tasker
|
||||
from src.services.mcp_service import init_mcp_servers
|
||||
from src.services.skill_service import init_skills_cache
|
||||
from src.services.run_queue_service import close_queue_clients, get_redis_client
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.knowledge import knowledge_base
|
||||
from src.utils import logger
|
||||
@ -28,12 +28,6 @@ async def lifespan(app: FastAPI):
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize MCP servers during startup: {e}")
|
||||
|
||||
# 初始化 Skills 缓存
|
||||
try:
|
||||
await init_skills_cache()
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize skills cache during startup: {e}")
|
||||
|
||||
# 初始化知识库管理器
|
||||
try:
|
||||
await knowledge_base.initialize()
|
||||
|
||||
@ -39,9 +39,9 @@ class AgentManager(metaclass=SingletonMeta):
|
||||
for agent_id in self._classes.keys():
|
||||
self.get_agent(agent_id, reload=True)
|
||||
|
||||
async def get_agents_info(self):
|
||||
async def get_agents_info(self, include_configurable_items: bool = True):
|
||||
agents = self.get_agents()
|
||||
return await asyncio.gather(*[a.get_info() for a in agents])
|
||||
return await asyncio.gather(*[a.get_info(include_configurable_items=include_configurable_items) for a in agents])
|
||||
|
||||
def auto_discover_agents(self):
|
||||
"""自动发现并注册 src/agents/ 下的所有智能体。
|
||||
|
||||
@ -6,7 +6,8 @@ from typing import Any
|
||||
from deepagents.backends import CompositeBackend, FilesystemBackend, StateBackend
|
||||
from deepagents.backends.protocol import EditResult, FileDownloadResponse, FileUploadResponse, WriteResult
|
||||
|
||||
from src.services.skill_service import get_expanded_visible_skill_slugs, get_skills_root_dir, is_valid_skill_slug
|
||||
from src.services.skill_resolver import normalize_selected_skills
|
||||
from src.services.skill_service import get_skills_root_dir, is_valid_skill_slug
|
||||
|
||||
|
||||
class SelectedSkillsReadonlyBackend(FilesystemBackend):
|
||||
@ -121,11 +122,23 @@ class SelectedSkillsReadonlyBackend(FilesystemBackend):
|
||||
|
||||
def create_agent_composite_backend(runtime) -> CompositeBackend:
|
||||
"""为 agent 构建 backend:默认 StateBackend + /skills 路由只读 backend。"""
|
||||
selected_skills = getattr(runtime.context, "skills", None)
|
||||
visible_skills = get_expanded_visible_skill_slugs(selected_skills or [])
|
||||
visible_skills = _get_visible_skills_from_runtime(runtime)
|
||||
return CompositeBackend(
|
||||
default=StateBackend(runtime),
|
||||
routes={
|
||||
"/skills/": SelectedSkillsReadonlyBackend(selected_slugs=visible_skills),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _get_visible_skills_from_runtime(runtime) -> list[str]:
|
||||
state = getattr(runtime, "state", None)
|
||||
if isinstance(state, dict):
|
||||
snapshot = state.get("skill_session_snapshot")
|
||||
if isinstance(snapshot, dict):
|
||||
visible = snapshot.get("visible_skills")
|
||||
if isinstance(visible, list):
|
||||
return [slug for slug in visible if isinstance(slug, str) and is_valid_skill_slug(slug)]
|
||||
|
||||
selected = getattr(runtime.context, "skills", None) or []
|
||||
return normalize_selected_skills(selected)
|
||||
|
||||
@ -12,6 +12,7 @@ from langgraph.graph.state import CompiledStateGraph
|
||||
|
||||
from src import config as sys_config
|
||||
from src.agents.common.context import BaseContext
|
||||
from src.services.skill_resolver import get_skill_options_from_db
|
||||
from src.utils import logger
|
||||
|
||||
|
||||
@ -43,9 +44,15 @@ class BaseAgent:
|
||||
"""Get the agent's class name."""
|
||||
return self.__class__.__name__
|
||||
|
||||
async def get_info(self):
|
||||
async def get_info(self, include_configurable_items: bool = True):
|
||||
# Load metadata from file
|
||||
metadata = self.load_metadata()
|
||||
configurable_items = {}
|
||||
if include_configurable_items:
|
||||
configurable_items = self.context_schema.get_configurable_items()
|
||||
if "skills" in configurable_items:
|
||||
configurable_items["skills"] = dict(configurable_items["skills"])
|
||||
configurable_items["skills"]["options"] = await get_skill_options_from_db()
|
||||
|
||||
# Merge metadata with class attributes, metadata takes precedence
|
||||
return {
|
||||
@ -53,7 +60,7 @@ class BaseAgent:
|
||||
"name": metadata.get("name", getattr(self, "name", "Unknown")),
|
||||
"description": metadata.get("description", getattr(self, "description", "Unknown")),
|
||||
"examples": metadata.get("examples", []),
|
||||
"configurable_items": self.context_schema.get_configurable_items(),
|
||||
"configurable_items": configurable_items,
|
||||
"has_checkpointer": await self.check_checkpointer(),
|
||||
"capabilities": getattr(self, "capabilities", []), # 智能体能力列表
|
||||
}
|
||||
|
||||
@ -10,7 +10,6 @@ import yaml
|
||||
|
||||
from src import config as sys_config
|
||||
from src.services.mcp_service import get_mcp_server_names
|
||||
from src.services.skill_service import get_skill_options
|
||||
from src.utils import logger
|
||||
|
||||
from .tools import gen_tool_info, get_buildin_tools
|
||||
@ -96,7 +95,7 @@ class BaseContext:
|
||||
default_factory=list,
|
||||
metadata={
|
||||
"name": "Skills",
|
||||
"options": lambda: get_skill_options(),
|
||||
"options": [],
|
||||
"description": "可选技能列表(由超级管理员维护)。运行时仅挂载并只读暴露选中的 skills。",
|
||||
"type": "list",
|
||||
},
|
||||
|
||||
@ -14,11 +14,15 @@ from langgraph.types import Command
|
||||
from src.agents.common import load_chat_model
|
||||
from src.agents.common.tools import get_buildin_tools, get_kb_based_tools
|
||||
from src.services.mcp_service import get_enabled_mcp_tools
|
||||
from src.services.skill_service import (
|
||||
get_dependency_bundle_for_activated_skills,
|
||||
get_skill_prompt_metadata_by_slugs,
|
||||
is_valid_skill_slug,
|
||||
from src.services.skill_resolver import (
|
||||
SkillSessionSnapshot,
|
||||
build_dependency_bundle,
|
||||
collect_prompt_metadata,
|
||||
is_snapshot_match_selected_skills,
|
||||
normalize_selected_skills,
|
||||
resolve_session_snapshot,
|
||||
)
|
||||
from src.services.skill_service import is_valid_skill_slug
|
||||
from src.utils.datetime_utils import shanghai_now
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
@ -40,6 +44,7 @@ def _activated_skills_reducer(left: list[str] | None, right: list[str] | None) -
|
||||
|
||||
class RuntimeConfigState(AgentState):
|
||||
activated_skills: NotRequired[Annotated[list[str], _activated_skills_reducer]]
|
||||
skill_session_snapshot: NotRequired[SkillSessionSnapshot]
|
||||
|
||||
|
||||
class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
@ -122,6 +127,7 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]
|
||||
) -> ModelResponse:
|
||||
runtime_context = request.runtime.context
|
||||
snapshot, request = await self._ensure_skill_snapshot(request)
|
||||
overrides: dict[str, Any] = {}
|
||||
|
||||
# 1. 模型覆盖(可选)
|
||||
@ -131,10 +137,11 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
|
||||
# 2. 工具覆盖(可选)
|
||||
if self.enable_tools_override:
|
||||
activated_skills = request.state.get("activated_skills", []) if request.state else []
|
||||
state = request.state if isinstance(request.state, dict) else {}
|
||||
activated_skills = state.get("activated_skills", [])
|
||||
if not isinstance(activated_skills, list):
|
||||
activated_skills = []
|
||||
deps_bundle = get_dependency_bundle_for_activated_skills(activated_skills)
|
||||
deps_bundle = build_dependency_bundle(snapshot, activated_skills)
|
||||
enabled_tools = await self.get_tools_from_context(
|
||||
runtime_context,
|
||||
extra_tool_names=deps_bundle["tools"],
|
||||
@ -161,7 +168,10 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
configured_skills = getattr(runtime_context, self.skills_context_name, None) or []
|
||||
if self.enable_skills_prompt_override and configured_skills:
|
||||
if self._supports_skill_prompt(request):
|
||||
skills_meta = get_skill_prompt_metadata_by_slugs(configured_skills)
|
||||
skills_for_prompt = configured_skills
|
||||
if snapshot and isinstance(snapshot.get("visible_skills"), list):
|
||||
skills_for_prompt = snapshot.get("visible_skills") or []
|
||||
skills_meta = collect_prompt_metadata(snapshot, skills_for_prompt)
|
||||
skills_section = self._build_skills_section(skills_meta)
|
||||
merged_system_prompt = f"{merged_system_prompt}\n\n{skills_section}"
|
||||
else:
|
||||
@ -255,6 +265,9 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
slug = self._extract_skill_slug_from_skill_md_path(file_path)
|
||||
if not slug:
|
||||
return result
|
||||
if not self._is_visible_skill_slug(request, slug):
|
||||
logger.warning(f"RuntimeConfigMiddleware: deny skill activation for invisible slug: {slug}")
|
||||
return result
|
||||
logger.debug(f"RuntimeConfigMiddleware: activated skill by read_file: {slug}")
|
||||
return self._merge_activated_skill_update(result, slug)
|
||||
|
||||
@ -272,9 +285,36 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
slug = self._extract_skill_slug_from_skill_md_path(file_path)
|
||||
if not slug:
|
||||
return result
|
||||
if not self._is_visible_skill_slug(request, slug):
|
||||
logger.warning(f"RuntimeConfigMiddleware: deny skill activation for invisible slug: {slug}")
|
||||
return result
|
||||
logger.debug(f"RuntimeConfigMiddleware: activated skill by read_file: {slug}")
|
||||
return self._merge_activated_skill_update(result, slug)
|
||||
|
||||
async def _ensure_skill_snapshot(self, request: ModelRequest) -> tuple[SkillSessionSnapshot | None, ModelRequest]:
|
||||
runtime_context = request.runtime.context
|
||||
configured_skills = getattr(runtime_context, self.skills_context_name, None) or []
|
||||
normalized_skills = normalize_selected_skills(configured_skills)
|
||||
state = request.state if isinstance(request.state, dict) else {}
|
||||
|
||||
snapshot = self._get_skill_snapshot_from_state(state)
|
||||
if is_snapshot_match_selected_skills(snapshot, normalized_skills):
|
||||
return snapshot, request
|
||||
|
||||
try:
|
||||
snapshot = await resolve_session_snapshot(normalized_skills)
|
||||
except Exception as e:
|
||||
logger.warning(f"RuntimeConfigMiddleware: failed to resolve skill snapshot, fallback empty: {e}")
|
||||
snapshot = None
|
||||
|
||||
if isinstance(request.state, dict):
|
||||
if snapshot:
|
||||
request.state["skill_session_snapshot"] = snapshot
|
||||
else:
|
||||
request.state.pop("skill_session_snapshot", None)
|
||||
|
||||
return snapshot, request
|
||||
|
||||
def _extract_skill_slug_from_skill_md_path(self, file_path: Any) -> str | None:
|
||||
if not isinstance(file_path, str):
|
||||
return None
|
||||
@ -304,6 +344,29 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
|
||||
return result
|
||||
|
||||
def _is_visible_skill_slug(self, request: ToolCallRequest, slug: str) -> bool:
|
||||
snapshot = self._get_skill_snapshot_from_state(request.state)
|
||||
if snapshot:
|
||||
visible_skills = snapshot.get("visible_skills")
|
||||
if isinstance(visible_skills, list):
|
||||
return slug in visible_skills
|
||||
|
||||
configured_skills = getattr(request.runtime.context, self.skills_context_name, None) or []
|
||||
normalized = normalize_selected_skills(configured_skills)
|
||||
return slug in normalized
|
||||
|
||||
def _get_skill_snapshot_from_state(self, state: Any) -> SkillSessionSnapshot | None:
|
||||
if not isinstance(state, dict):
|
||||
return None
|
||||
snapshot = state.get("skill_session_snapshot")
|
||||
if not isinstance(snapshot, dict):
|
||||
return None
|
||||
visible_skills = snapshot.get("visible_skills")
|
||||
selected_skills = snapshot.get("selected_skills")
|
||||
if not isinstance(visible_skills, list) or not isinstance(selected_skills, list):
|
||||
return None
|
||||
return snapshot
|
||||
|
||||
def _supports_skill_prompt(self, request: ModelRequest) -> bool:
|
||||
"""仅当请求工具中包含 read_file 时,才注入 skills 指引。"""
|
||||
for tool in request.tools or []:
|
||||
|
||||
226
src/services/skill_resolver.py
Normal file
226
src/services/skill_resolver.py
Normal file
@ -0,0 +1,226 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TypedDict
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.repositories.skill_repository import SkillRepository
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import Skill
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
|
||||
class SkillPromptMetadata(TypedDict):
|
||||
name: str
|
||||
description: str
|
||||
path: str
|
||||
|
||||
|
||||
class SkillDependencyNode(TypedDict):
|
||||
tools: list[str]
|
||||
mcps: list[str]
|
||||
skills: list[str]
|
||||
|
||||
|
||||
class SkillSessionSnapshot(TypedDict):
|
||||
selected_skills: list[str]
|
||||
visible_skills: list[str]
|
||||
prompt_metadata: dict[str, SkillPromptMetadata]
|
||||
dependency_map: dict[str, SkillDependencyNode]
|
||||
|
||||
|
||||
def normalize_selected_skills(selected_skills: list[str] | None) -> list[str]:
|
||||
return _normalize_string_list(selected_skills)
|
||||
|
||||
|
||||
def is_snapshot_match_selected_skills(
|
||||
snapshot: SkillSessionSnapshot | None,
|
||||
selected_skills: list[str] | None,
|
||||
) -> bool:
|
||||
if not snapshot:
|
||||
return False
|
||||
current = snapshot.get("selected_skills")
|
||||
if not isinstance(current, list):
|
||||
return False
|
||||
return current == normalize_selected_skills(selected_skills)
|
||||
|
||||
|
||||
async def resolve_session_snapshot(
|
||||
selected_skills: list[str] | None,
|
||||
*,
|
||||
db: AsyncSession | None = None,
|
||||
) -> SkillSessionSnapshot:
|
||||
normalized_selected = normalize_selected_skills(selected_skills)
|
||||
skills = await _list_skills_from_db(db)
|
||||
prompt_metadata, dependency_map = _build_maps(skills)
|
||||
visible_skills = expand_skill_closure(normalized_selected, dependency_map)
|
||||
return {
|
||||
"selected_skills": normalized_selected,
|
||||
"visible_skills": visible_skills,
|
||||
"prompt_metadata": prompt_metadata,
|
||||
"dependency_map": dependency_map,
|
||||
}
|
||||
|
||||
|
||||
def collect_prompt_metadata(
|
||||
snapshot: SkillSessionSnapshot | None,
|
||||
slugs: list[str] | None,
|
||||
) -> list[SkillPromptMetadata]:
|
||||
if not snapshot or not slugs:
|
||||
return []
|
||||
|
||||
prompt_metadata = snapshot.get("prompt_metadata") or {}
|
||||
result: list[SkillPromptMetadata] = []
|
||||
seen: set[str] = set()
|
||||
for slug in slugs:
|
||||
if not isinstance(slug, str):
|
||||
continue
|
||||
normalized = slug.strip()
|
||||
if not normalized or normalized in seen:
|
||||
continue
|
||||
seen.add(normalized)
|
||||
|
||||
item = prompt_metadata.get(normalized)
|
||||
if not item:
|
||||
logger.debug(f"Skill slug not found in session snapshot, skip prompt metadata: {normalized}")
|
||||
continue
|
||||
result.append(dict(item))
|
||||
return result
|
||||
|
||||
|
||||
def build_dependency_bundle(
|
||||
snapshot: SkillSessionSnapshot | None,
|
||||
activated_slugs: list[str] | None,
|
||||
) -> dict[str, list[str]]:
|
||||
if not snapshot:
|
||||
return {"tools": [], "mcps": [], "skills": []}
|
||||
|
||||
dependency_map = snapshot.get("dependency_map") or {}
|
||||
closure = expand_skill_closure(activated_slugs or [], dependency_map)
|
||||
tools: list[str] = []
|
||||
mcps: list[str] = []
|
||||
seen_tools: set[str] = set()
|
||||
seen_mcps: set[str] = set()
|
||||
|
||||
for slug in closure:
|
||||
dep = dependency_map.get(slug, {})
|
||||
for tool_name in dep.get("tools", []):
|
||||
if tool_name in seen_tools:
|
||||
continue
|
||||
seen_tools.add(tool_name)
|
||||
tools.append(tool_name)
|
||||
for mcp_name in dep.get("mcps", []):
|
||||
if mcp_name in seen_mcps:
|
||||
continue
|
||||
seen_mcps.add(mcp_name)
|
||||
mcps.append(mcp_name)
|
||||
|
||||
return {"tools": tools, "mcps": mcps, "skills": closure}
|
||||
|
||||
|
||||
def expand_skill_closure(
|
||||
slugs: list[str] | None,
|
||||
dependency_map: dict[str, SkillDependencyNode],
|
||||
) -> list[str]:
|
||||
ordered_roots = _normalize_string_list(slugs)
|
||||
if not ordered_roots:
|
||||
return []
|
||||
|
||||
result: list[str] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
def dfs(slug: str, stack: set[str]) -> None:
|
||||
if slug in stack:
|
||||
logger.warning(f"Cycle detected in skill dependencies, skip: {' -> '.join([*stack, slug])}")
|
||||
return
|
||||
if slug in seen:
|
||||
return
|
||||
|
||||
node = dependency_map.get(slug)
|
||||
if not node:
|
||||
logger.warning(f"Skill dependency target not found in DB snapshot, skip: {slug}")
|
||||
return
|
||||
|
||||
seen.add(slug)
|
||||
result.append(slug)
|
||||
next_stack = set(stack)
|
||||
next_stack.add(slug)
|
||||
for dep in node.get("skills", []):
|
||||
dfs(dep, next_stack)
|
||||
|
||||
for root in ordered_roots:
|
||||
dfs(root, set())
|
||||
return result
|
||||
|
||||
|
||||
async def get_skill_options_from_db(
|
||||
*,
|
||||
db: AsyncSession | None = None,
|
||||
) -> list[dict[str, str]]:
|
||||
items = await _list_skills_from_db(db)
|
||||
return [
|
||||
{
|
||||
"id": item.slug,
|
||||
"name": item.name,
|
||||
"description": item.description,
|
||||
}
|
||||
for item in items
|
||||
]
|
||||
|
||||
|
||||
async def get_skill_slug_set_from_db(
|
||||
*,
|
||||
db: AsyncSession | None = None,
|
||||
) -> set[str]:
|
||||
items = await _list_skills_from_db(db)
|
||||
return {item.slug for item in items}
|
||||
|
||||
|
||||
def _normalize_string_list(values: list[str] | None) -> list[str]:
|
||||
if not values:
|
||||
return []
|
||||
normalized: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for value in values:
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
item = value.strip()
|
||||
if not item or item in seen:
|
||||
continue
|
||||
seen.add(item)
|
||||
normalized.append(item)
|
||||
return normalized
|
||||
|
||||
|
||||
def _build_maps(skills: list[Skill]) -> tuple[dict[str, SkillPromptMetadata], dict[str, SkillDependencyNode]]:
|
||||
prompt_metadata: dict[str, SkillPromptMetadata] = {}
|
||||
dependency_map: dict[str, SkillDependencyNode] = {}
|
||||
for item in skills:
|
||||
prompt_metadata[item.slug] = {
|
||||
"name": item.name,
|
||||
"description": item.description,
|
||||
"path": f"/skills/{item.slug}/SKILL.md",
|
||||
}
|
||||
dependency_map[item.slug] = {
|
||||
"tools": _normalize_string_list(item.tool_dependencies or []),
|
||||
"mcps": _normalize_string_list(item.mcp_dependencies or []),
|
||||
"skills": _normalize_string_list(item.skill_dependencies or []),
|
||||
}
|
||||
return prompt_metadata, dependency_map
|
||||
|
||||
|
||||
async def _list_skills_from_db(db: AsyncSession | None) -> list[Skill]:
|
||||
if db is not None:
|
||||
repo = SkillRepository(db)
|
||||
return await repo.list_all()
|
||||
|
||||
try:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
repo = SkillRepository(session)
|
||||
return await repo.list_all()
|
||||
except RuntimeError:
|
||||
# 在非 FastAPI 生命周期场景(如 worker/脚本)按需初始化
|
||||
pg_manager.initialize()
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
repo = SkillRepository(session)
|
||||
return await repo.list_all()
|
||||
@ -14,9 +14,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from src import config as sys_config
|
||||
from src.repositories.skill_repository import SkillRepository
|
||||
from src.services.mcp_service import get_mcp_server_names
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import Skill
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
SKILL_SLUG_PATTERN = re.compile(r"^[a-z0-9]+(-[a-z0-9]+)*$")
|
||||
SKILL_NAME_PATTERN = SKILL_SLUG_PATTERN
|
||||
@ -52,11 +50,6 @@ TEXT_FILE_EXTENSIONS = {
|
||||
".tsx",
|
||||
}
|
||||
|
||||
_skill_options_cache: list[dict[str, str]] = []
|
||||
_skill_prompt_metadata_cache: dict[str, dict[str, str]] = {}
|
||||
_skill_dependency_cache: dict[str, dict[str, list[str]]] = {}
|
||||
|
||||
|
||||
def _normalize_string_list(values: list[str] | None) -> list[str]:
|
||||
if not values:
|
||||
return []
|
||||
@ -98,146 +91,19 @@ def get_skills_root_dir() -> Path:
|
||||
return root
|
||||
|
||||
|
||||
def get_skill_options() -> list[dict[str, str]]:
|
||||
"""返回技能选项缓存(用于 BaseContext configurable options)。"""
|
||||
return list(_skill_options_cache)
|
||||
|
||||
|
||||
def get_skill_dependency_options() -> dict[str, list[str]]:
|
||||
async def get_skill_dependency_options(db: AsyncSession) -> dict[str, list[str]]:
|
||||
repo = SkillRepository(db)
|
||||
items = await repo.list_all()
|
||||
return {
|
||||
"tools": _get_buildin_tool_names(),
|
||||
"mcps": get_mcp_server_names(),
|
||||
"skills": [item["id"] for item in _skill_options_cache],
|
||||
"skills": [item.slug for item in items],
|
||||
}
|
||||
|
||||
|
||||
def get_skill_prompt_metadata_by_slugs(slugs: list[str]) -> list[dict[str, str]]:
|
||||
"""按 slug 顺序返回 skills prompt 元数据(仅缓存,无 IO)。"""
|
||||
if not slugs:
|
||||
return []
|
||||
|
||||
result: list[dict[str, str]] = []
|
||||
seen: set[str] = set()
|
||||
for slug in slugs:
|
||||
if slug in seen:
|
||||
continue
|
||||
seen.add(slug)
|
||||
|
||||
item = _skill_prompt_metadata_cache.get(slug)
|
||||
if not item:
|
||||
logger.debug(f"Skill slug not found in cache, skip prompt metadata: {slug}")
|
||||
continue
|
||||
|
||||
result.append(dict(item))
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def expand_skill_closure(slugs: list[str]) -> list[str]:
|
||||
"""递归展开 skill 依赖(仅缓存,无 IO),去重保序并去环。"""
|
||||
ordered_roots = _normalize_string_list(slugs)
|
||||
if not ordered_roots:
|
||||
return []
|
||||
|
||||
result: list[str] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
def dfs(slug: str, stack: set[str]) -> None:
|
||||
if slug in stack:
|
||||
logger.warning(f"Cycle detected in skill dependencies, skip: {' -> '.join([*stack, slug])}")
|
||||
return
|
||||
if slug in seen:
|
||||
return
|
||||
|
||||
node = _skill_dependency_cache.get(slug)
|
||||
if not node:
|
||||
logger.warning(f"Skill dependency target not found in cache, skip: {slug}")
|
||||
return
|
||||
|
||||
seen.add(slug)
|
||||
result.append(slug)
|
||||
next_stack = set(stack)
|
||||
next_stack.add(slug)
|
||||
for dep in node.get("skills", []):
|
||||
dfs(dep, next_stack)
|
||||
|
||||
for root in ordered_roots:
|
||||
dfs(root, set())
|
||||
return result
|
||||
|
||||
|
||||
def get_dependency_bundle_for_activated_skills(activated_slugs: list[str]) -> dict[str, list[str]]:
|
||||
closure = expand_skill_closure(activated_slugs)
|
||||
tools: list[str] = []
|
||||
mcps: list[str] = []
|
||||
seen_tools: set[str] = set()
|
||||
seen_mcps: set[str] = set()
|
||||
|
||||
for slug in closure:
|
||||
dep = _skill_dependency_cache.get(slug, {})
|
||||
for tool_name in dep.get("tools", []):
|
||||
if tool_name in seen_tools:
|
||||
continue
|
||||
seen_tools.add(tool_name)
|
||||
tools.append(tool_name)
|
||||
for server_name in dep.get("mcps", []):
|
||||
if server_name in seen_mcps:
|
||||
continue
|
||||
seen_mcps.add(server_name)
|
||||
mcps.append(server_name)
|
||||
return {"tools": tools, "mcps": mcps, "skills": closure}
|
||||
|
||||
|
||||
def get_expanded_visible_skill_slugs(selected_slugs: list[str]) -> list[str]:
|
||||
"""展开运行时可见 skills(根 skills + 递归依赖)。"""
|
||||
return expand_skill_closure(selected_slugs)
|
||||
|
||||
|
||||
def _set_skill_options_cache(items: list[Skill]) -> None:
|
||||
global _skill_options_cache, _skill_prompt_metadata_cache, _skill_dependency_cache
|
||||
_skill_options_cache = [
|
||||
{
|
||||
"id": item.slug,
|
||||
"name": item.name,
|
||||
"description": item.description,
|
||||
}
|
||||
for item in items
|
||||
]
|
||||
_skill_prompt_metadata_cache = {
|
||||
item.slug: {
|
||||
"name": item.name,
|
||||
"description": item.description,
|
||||
"path": f"/skills/{item.slug}/SKILL.md",
|
||||
}
|
||||
for item in items
|
||||
}
|
||||
_skill_dependency_cache = {
|
||||
item.slug: {
|
||||
"tools": _normalize_string_list(item.tool_dependencies or []),
|
||||
"mcps": _normalize_string_list(item.mcp_dependencies or []),
|
||||
"skills": _normalize_string_list(item.skill_dependencies or []),
|
||||
}
|
||||
for item in items
|
||||
}
|
||||
|
||||
|
||||
async def init_skills_cache() -> None:
|
||||
"""启动时加载技能缓存,避免每次构建 configurable_items 触发 DB IO。"""
|
||||
try:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
repo = SkillRepository(session)
|
||||
items = await repo.list_all()
|
||||
_set_skill_options_cache(items)
|
||||
logger.info(f"Loaded skills cache with {len(items)} items")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to initialize skills cache: {e}")
|
||||
|
||||
|
||||
async def list_skills(db: AsyncSession) -> list[Skill]:
|
||||
repo = SkillRepository(db)
|
||||
items = await repo.list_all()
|
||||
_set_skill_options_cache(items)
|
||||
return items
|
||||
return await repo.list_all()
|
||||
|
||||
|
||||
def _validate_dependencies(
|
||||
@ -246,6 +112,7 @@ def _validate_dependencies(
|
||||
tool_dependencies: list[str],
|
||||
mcp_dependencies: list[str],
|
||||
skill_dependencies: list[str],
|
||||
available_skill_slugs: set[str],
|
||||
) -> tuple[list[str], list[str], list[str]]:
|
||||
tools = _normalize_string_list(tool_dependencies)
|
||||
mcps = _normalize_string_list(mcp_dependencies)
|
||||
@ -261,8 +128,7 @@ def _validate_dependencies(
|
||||
if invalid_mcps:
|
||||
raise ValueError(f"存在无效 MCP 依赖: {', '.join(invalid_mcps)}")
|
||||
|
||||
available_skills = {item["id"] for item in _skill_options_cache}
|
||||
invalid_skills = [name for name in skills if name not in available_skills]
|
||||
invalid_skills = [name for name in skills if name not in available_skill_slugs]
|
||||
if invalid_skills:
|
||||
raise ValueError(f"存在无效 skill 依赖: {', '.join(invalid_skills)}")
|
||||
|
||||
@ -283,24 +149,23 @@ async def update_skill_dependencies(
|
||||
) -> Skill:
|
||||
item = await get_skill_or_raise(db, slug)
|
||||
repo = SkillRepository(db)
|
||||
# 写操作前先同步一次缓存,确保依赖校验基于最新技能集合。
|
||||
_set_skill_options_cache(await repo.list_all())
|
||||
skill_items = await repo.list_all()
|
||||
available_skill_slugs = {skill.slug for skill in skill_items}
|
||||
tools, mcps, skills = _validate_dependencies(
|
||||
slug=slug,
|
||||
tool_dependencies=tool_dependencies,
|
||||
mcp_dependencies=mcp_dependencies,
|
||||
skill_dependencies=skill_dependencies,
|
||||
available_skill_slugs=available_skill_slugs,
|
||||
)
|
||||
|
||||
updated = await repo.update_dependencies(
|
||||
return await repo.update_dependencies(
|
||||
item,
|
||||
tool_dependencies=tools,
|
||||
mcp_dependencies=mcps,
|
||||
skill_dependencies=skills,
|
||||
updated_by=updated_by,
|
||||
)
|
||||
_set_skill_options_cache(await repo.list_all())
|
||||
return updated
|
||||
|
||||
|
||||
def _validate_skill_name(name: str) -> str:
|
||||
@ -499,8 +364,6 @@ async def import_skill_zip(
|
||||
shutil.rmtree(final_dir, ignore_errors=True)
|
||||
raise
|
||||
|
||||
items = await repo.list_all()
|
||||
_set_skill_options_cache(items)
|
||||
return item
|
||||
|
||||
|
||||
@ -572,7 +435,6 @@ async def create_skill_node(
|
||||
if parsed_name is not None and parsed_desc is not None:
|
||||
repo = SkillRepository(db)
|
||||
await repo.update_metadata(item, name=parsed_name, description=parsed_desc, updated_by=updated_by)
|
||||
_set_skill_options_cache(await repo.list_all())
|
||||
|
||||
|
||||
async def update_skill_file(
|
||||
@ -603,7 +465,6 @@ async def update_skill_file(
|
||||
if parsed_name is not None and parsed_desc is not None:
|
||||
repo = SkillRepository(db)
|
||||
await repo.update_metadata(item, name=parsed_name, description=parsed_desc, updated_by=updated_by)
|
||||
_set_skill_options_cache(await repo.list_all())
|
||||
|
||||
|
||||
async def delete_skill_node(db: AsyncSession, *, slug: str, relative_path: str) -> None:
|
||||
@ -664,5 +525,3 @@ async def delete_skill(db: AsyncSession, *, slug: str) -> None:
|
||||
|
||||
if trash_dir and trash_dir.exists():
|
||||
shutil.rmtree(trash_dir, ignore_errors=True)
|
||||
|
||||
_set_skill_options_cache(await repo.list_all())
|
||||
|
||||
@ -10,7 +10,6 @@ from langgraph.types import Command
|
||||
|
||||
import src.agents.common.middlewares.runtime_config_middleware as runtime_middleware
|
||||
from src.agents.common.middlewares.runtime_config_middleware import RuntimeConfigMiddleware
|
||||
from src.services import skill_service
|
||||
|
||||
|
||||
@dataclass
|
||||
@ -34,18 +33,40 @@ class _FakeRequest:
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _FakeToolCallRequest:
|
||||
tool_call: dict[str, Any]
|
||||
runtime: Any
|
||||
state: dict[str, Any]
|
||||
|
||||
|
||||
async def _echo_handler(request):
|
||||
return request
|
||||
|
||||
|
||||
def _build_request(*, skills: list[str], tools: list[str], system_prompt: str = "你是助手") -> _FakeRequest:
|
||||
context = SimpleNamespace(system_prompt=system_prompt, skills=skills)
|
||||
def _build_request(*, skills: list[str], tools: list[str], system_prompt: str = "你是助手", state=None) -> _FakeRequest:
|
||||
context = SimpleNamespace(system_prompt=system_prompt, skills=skills, tools=[], knowledges=[], mcps=[])
|
||||
runtime = SimpleNamespace(context=context)
|
||||
return _FakeRequest(
|
||||
runtime=runtime,
|
||||
tools=[_FakeTool(name=name) for name in tools],
|
||||
system_message=SystemMessage(content=[{"type": "text", "text": "base"}]),
|
||||
state={},
|
||||
state=state or {},
|
||||
)
|
||||
|
||||
|
||||
def _build_tool_request(*, skills: list[str], visible_skills: list[str], file_path: str) -> _FakeToolCallRequest:
|
||||
return _FakeToolCallRequest(
|
||||
tool_call={"name": "read_file", "args": {"file_path": file_path}},
|
||||
runtime=SimpleNamespace(context=SimpleNamespace(skills=skills)),
|
||||
state={
|
||||
"skill_session_snapshot": {
|
||||
"selected_skills": skills,
|
||||
"visible_skills": visible_skills,
|
||||
"prompt_metadata": {},
|
||||
"dependency_map": {},
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@ -62,24 +83,31 @@ def _build_middleware() -> RuntimeConfigMiddleware:
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _FakeToolCallRequest:
|
||||
tool_call: dict[str, Any]
|
||||
def _build_snapshot(selected: list[str], metadata: dict[str, dict[str, str]] | None = None) -> dict[str, Any]:
|
||||
return {
|
||||
"selected_skills": selected,
|
||||
"visible_skills": selected,
|
||||
"prompt_metadata": metadata or {},
|
||||
"dependency_map": {},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_injects_skills_section_when_skills_configured_and_read_file_available(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(
|
||||
runtime_middleware,
|
||||
"get_skill_prompt_metadata_by_slugs",
|
||||
lambda _slugs: [
|
||||
async def fake_resolve(selected):
|
||||
assert selected == ["research-report"]
|
||||
return _build_snapshot(
|
||||
["research-report"],
|
||||
{
|
||||
"name": "research-report",
|
||||
"description": "Write structured research reports",
|
||||
"path": "/skills/research-report/SKILL.md",
|
||||
}
|
||||
],
|
||||
)
|
||||
"research-report": {
|
||||
"name": "research-report",
|
||||
"description": "Write structured research reports",
|
||||
"path": "/skills/research-report/SKILL.md",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
monkeypatch.setattr(runtime_middleware, "resolve_session_snapshot", fake_resolve)
|
||||
middleware = _build_middleware()
|
||||
request = _build_request(skills=["research-report"], tools=["read_file"])
|
||||
|
||||
@ -92,14 +120,11 @@ async def test_injects_skills_section_when_skills_configured_and_read_file_avail
|
||||
assert "Read `/skills/research-report/SKILL.md` for full instructions" in prompt
|
||||
assert "Recognize when a skill applies" in prompt
|
||||
assert "当前时间:" in prompt
|
||||
assert "skill_session_snapshot" in result.state
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_skills_section_when_context_skills_empty(monkeypatch: pytest.MonkeyPatch):
|
||||
def _should_not_call(_slugs: list[str]):
|
||||
raise AssertionError("should not query skills metadata when context.skills is empty")
|
||||
|
||||
monkeypatch.setattr(runtime_middleware, "get_skill_prompt_metadata_by_slugs", _should_not_call)
|
||||
async def test_skips_skills_section_when_context_skills_empty():
|
||||
middleware = _build_middleware()
|
||||
request = _build_request(skills=[], tools=["read_file"])
|
||||
|
||||
@ -117,11 +142,20 @@ async def test_skips_skills_section_without_read_file_and_logs_warning(monkeypat
|
||||
warning=lambda message: warnings.append(message),
|
||||
)
|
||||
monkeypatch.setattr(runtime_middleware, "logger", fake_logger)
|
||||
monkeypatch.setattr(
|
||||
runtime_middleware,
|
||||
"get_skill_prompt_metadata_by_slugs",
|
||||
lambda _slugs: (_ for _ in ()).throw(AssertionError("should not query metadata without read_file")),
|
||||
)
|
||||
|
||||
async def fake_resolve(_selected):
|
||||
return _build_snapshot(
|
||||
["research-report"],
|
||||
{
|
||||
"research-report": {
|
||||
"name": "research-report",
|
||||
"description": "Write structured research reports",
|
||||
"path": "/skills/research-report/SKILL.md",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
monkeypatch.setattr(runtime_middleware, "resolve_session_snapshot", fake_resolve)
|
||||
middleware = _build_middleware()
|
||||
request = _build_request(skills=["research-report"], tools=["write_file"])
|
||||
|
||||
@ -134,22 +168,16 @@ async def test_skips_skills_section_without_read_file_and_logs_warning(monkeypat
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_injects_skills_in_input_order_with_dedup_and_invalid_slug_skipped(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(
|
||||
skill_service,
|
||||
"_skill_prompt_metadata_cache",
|
||||
{
|
||||
"beta": {
|
||||
"name": "beta",
|
||||
"description": "beta skill",
|
||||
"path": "/skills/beta/SKILL.md",
|
||||
async def fake_resolve(_selected):
|
||||
return _build_snapshot(
|
||||
["beta", "missing", "alpha", "beta"],
|
||||
{
|
||||
"beta": {"name": "beta", "description": "beta skill", "path": "/skills/beta/SKILL.md"},
|
||||
"alpha": {"name": "alpha", "description": "alpha skill", "path": "/skills/alpha/SKILL.md"},
|
||||
},
|
||||
"alpha": {
|
||||
"name": "alpha",
|
||||
"description": "alpha skill",
|
||||
"path": "/skills/alpha/SKILL.md",
|
||||
},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
monkeypatch.setattr(runtime_middleware, "resolve_session_snapshot", fake_resolve)
|
||||
middleware = _build_middleware()
|
||||
request = _build_request(skills=["beta", "missing", "alpha", "beta"], tools=["read_file"])
|
||||
|
||||
@ -168,11 +196,10 @@ async def test_injects_skills_in_input_order_with_dedup_and_invalid_slug_skipped
|
||||
@pytest.mark.asyncio
|
||||
async def test_awrap_tool_call_activates_skill_when_read_skill_md():
|
||||
middleware = _build_middleware()
|
||||
request = _FakeToolCallRequest(
|
||||
tool_call={
|
||||
"name": "read_file",
|
||||
"args": {"file_path": "/skills/research-report/SKILL.md"},
|
||||
}
|
||||
request = _build_tool_request(
|
||||
skills=["research-report"],
|
||||
visible_skills=["research-report"],
|
||||
file_path="/skills/research-report/SKILL.md",
|
||||
)
|
||||
|
||||
async def _handler(_request):
|
||||
@ -187,11 +214,10 @@ async def test_awrap_tool_call_activates_skill_when_read_skill_md():
|
||||
@pytest.mark.asyncio
|
||||
async def test_awrap_tool_call_skips_invalid_skill_slug_path():
|
||||
middleware = _build_middleware()
|
||||
request = _FakeToolCallRequest(
|
||||
tool_call={
|
||||
"name": "read_file",
|
||||
"args": {"file_path": "/skills/../SKILL.md"},
|
||||
}
|
||||
request = _build_tool_request(
|
||||
skills=["research-report"],
|
||||
visible_skills=["research-report"],
|
||||
file_path="/skills/../SKILL.md",
|
||||
)
|
||||
|
||||
async def _handler(_request):
|
||||
@ -204,11 +230,10 @@ async def test_awrap_tool_call_skips_invalid_skill_slug_path():
|
||||
@pytest.mark.asyncio
|
||||
async def test_awrap_tool_call_merges_with_existing_command_update():
|
||||
middleware = _build_middleware()
|
||||
request = _FakeToolCallRequest(
|
||||
tool_call={
|
||||
"name": "read_file",
|
||||
"args": {"file_path": "/skills/research-report/SKILL.md"},
|
||||
}
|
||||
request = _build_tool_request(
|
||||
skills=["research-report"],
|
||||
visible_skills=["research-report"],
|
||||
file_path="/skills/research-report/SKILL.md",
|
||||
)
|
||||
|
||||
async def _handler(_request):
|
||||
@ -219,6 +244,22 @@ async def test_awrap_tool_call_merges_with_existing_command_update():
|
||||
assert result.update["activated_skills"] == ["a", "research-report"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_awrap_tool_call_denies_invisible_skill():
|
||||
middleware = _build_middleware()
|
||||
request = _build_tool_request(
|
||||
skills=["research-report"],
|
||||
visible_skills=["alpha"],
|
||||
file_path="/skills/research-report/SKILL.md",
|
||||
)
|
||||
|
||||
async def _handler(_request):
|
||||
return ToolMessage(content="ok", tool_call_id="tc-1")
|
||||
|
||||
result = await middleware.awrap_tool_call(request, _handler)
|
||||
assert isinstance(result, ToolMessage)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_call_injects_dependency_tools_and_mcps_after_activation(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(
|
||||
@ -229,8 +270,8 @@ async def test_model_call_injects_dependency_tools_and_mcps_after_activation(mon
|
||||
monkeypatch.setattr(runtime_middleware, "get_kb_based_tools", lambda db_names=None: [])
|
||||
monkeypatch.setattr(
|
||||
runtime_middleware,
|
||||
"get_dependency_bundle_for_activated_skills",
|
||||
lambda activated: {"tools": ["dep-tool"], "mcps": ["mcp-a"], "skills": activated},
|
||||
"build_dependency_bundle",
|
||||
lambda _snapshot, activated: {"tools": ["dep-tool"], "mcps": ["mcp-a"], "skills": activated},
|
||||
)
|
||||
|
||||
async def fake_get_enabled_mcp_tools(server_name: str):
|
||||
@ -248,17 +289,10 @@ async def test_model_call_injects_dependency_tools_and_mcps_after_activation(mon
|
||||
enable_skills_prompt_override=False,
|
||||
)
|
||||
|
||||
context = SimpleNamespace(system_prompt="x", skills=[], tools=[], knowledges=[], mcps=[])
|
||||
request = _FakeRequest(
|
||||
runtime=SimpleNamespace(context=context),
|
||||
tools=[
|
||||
_FakeTool(name="calculator"),
|
||||
_FakeTool(name="dep-tool"),
|
||||
_FakeTool(name="mcp_tool"),
|
||||
_FakeTool(name="read_file"),
|
||||
],
|
||||
system_message=SystemMessage(content=[{"type": "text", "text": "base"}]),
|
||||
state={"activated_skills": ["alpha"]},
|
||||
request = _build_request(
|
||||
skills=[],
|
||||
tools=["calculator", "dep-tool", "mcp_tool", "read_file"],
|
||||
state={"activated_skills": ["alpha"], "skill_session_snapshot": _build_snapshot([])},
|
||||
)
|
||||
|
||||
result = await middleware.awrap_model_call(request, _echo_handler)
|
||||
@ -266,3 +300,53 @@ async def test_model_call_injects_dependency_tools_and_mcps_after_activation(mon
|
||||
assert "dep-tool" in tool_names
|
||||
assert "mcp_tool" in tool_names
|
||||
assert "calculator" not in tool_names
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_call_reuses_snapshot_until_skills_changed(monkeypatch: pytest.MonkeyPatch):
|
||||
called = {"count": 0}
|
||||
|
||||
async def fake_resolve(selected):
|
||||
called["count"] += 1
|
||||
return _build_snapshot(selected)
|
||||
|
||||
monkeypatch.setattr(runtime_middleware, "resolve_session_snapshot", fake_resolve)
|
||||
|
||||
middleware = _build_middleware()
|
||||
req1 = _build_request(skills=["alpha"], tools=["read_file"], state={})
|
||||
res1 = await middleware.awrap_model_call(req1, _echo_handler)
|
||||
assert called["count"] == 1
|
||||
|
||||
req2 = _build_request(skills=["alpha"], tools=["read_file"], state=res1.state)
|
||||
await middleware.awrap_model_call(req2, _echo_handler)
|
||||
assert called["count"] == 1
|
||||
|
||||
req3 = _build_request(skills=["beta"], tools=["read_file"], state=res1.state)
|
||||
await middleware.awrap_model_call(req3, _echo_handler)
|
||||
assert called["count"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_injects_dependency_skills_into_prompt(monkeypatch: pytest.MonkeyPatch):
|
||||
async def fake_resolve(_selected):
|
||||
return {
|
||||
"selected_skills": ["alpha"],
|
||||
"visible_skills": ["alpha", "beta"],
|
||||
"prompt_metadata": {
|
||||
"alpha": {"name": "alpha", "description": "alpha desc", "path": "/skills/alpha/SKILL.md"},
|
||||
"beta": {"name": "beta", "description": "beta desc", "path": "/skills/beta/SKILL.md"},
|
||||
},
|
||||
"dependency_map": {
|
||||
"alpha": {"tools": [], "mcps": [], "skills": ["beta"]},
|
||||
"beta": {"tools": [], "mcps": [], "skills": []},
|
||||
},
|
||||
}
|
||||
|
||||
monkeypatch.setattr(runtime_middleware, "resolve_session_snapshot", fake_resolve)
|
||||
middleware = _build_middleware()
|
||||
request = _build_request(skills=["alpha"], tools=["read_file"])
|
||||
|
||||
result = await middleware.awrap_model_call(request, _echo_handler)
|
||||
prompt = _extract_appended_prompt(result)
|
||||
assert "- **alpha**: alpha desc" in prompt
|
||||
assert "- **beta**: beta desc" in prompt
|
||||
|
||||
85
test/test_skill_resolver.py
Normal file
85
test/test_skill_resolver.py
Normal file
@ -0,0 +1,85 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services import skill_resolver as resolver
|
||||
from src.storage.postgres.models_business import Skill
|
||||
|
||||
|
||||
def test_expand_skill_closure_and_dependency_bundle():
|
||||
dependency_map = {
|
||||
"alpha": {"tools": ["t1"], "mcps": ["m1"], "skills": ["beta"]},
|
||||
"beta": {"tools": ["t2"], "mcps": ["m2"], "skills": ["gamma"]},
|
||||
"gamma": {"tools": ["t3"], "mcps": [], "skills": []},
|
||||
}
|
||||
snapshot: resolver.SkillSessionSnapshot = {
|
||||
"selected_skills": ["alpha"],
|
||||
"visible_skills": ["alpha", "beta", "gamma"],
|
||||
"prompt_metadata": {},
|
||||
"dependency_map": dependency_map,
|
||||
}
|
||||
|
||||
closure = resolver.expand_skill_closure(["alpha"], dependency_map)
|
||||
assert closure == ["alpha", "beta", "gamma"]
|
||||
|
||||
bundle = resolver.build_dependency_bundle(snapshot, ["alpha"])
|
||||
assert bundle["skills"] == ["alpha", "beta", "gamma"]
|
||||
assert bundle["tools"] == ["t1", "t2", "t3"]
|
||||
assert bundle["mcps"] == ["m1", "m2"]
|
||||
|
||||
|
||||
def test_expand_skill_closure_cycle():
|
||||
dependency_map = {
|
||||
"alpha": {"tools": [], "mcps": [], "skills": ["beta"]},
|
||||
"beta": {"tools": [], "mcps": [], "skills": ["alpha"]},
|
||||
}
|
||||
assert resolver.expand_skill_closure(["alpha"], dependency_map) == ["alpha", "beta"]
|
||||
|
||||
|
||||
def test_collect_prompt_metadata_order_and_dedup():
|
||||
snapshot: resolver.SkillSessionSnapshot = {
|
||||
"selected_skills": ["beta", "alpha"],
|
||||
"visible_skills": ["beta", "alpha"],
|
||||
"prompt_metadata": {
|
||||
"beta": {"name": "beta", "description": "beta skill", "path": "/skills/beta/SKILL.md"},
|
||||
"alpha": {"name": "alpha", "description": "alpha skill", "path": "/skills/alpha/SKILL.md"},
|
||||
},
|
||||
"dependency_map": {},
|
||||
}
|
||||
result = resolver.collect_prompt_metadata(snapshot, ["beta", "missing", "alpha", "beta"])
|
||||
assert [item["name"] for item in result] == ["beta", "alpha"]
|
||||
assert [item["path"] for item in result] == ["/skills/beta/SKILL.md", "/skills/alpha/SKILL.md"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_session_snapshot_and_selected_change(monkeypatch: pytest.MonkeyPatch):
|
||||
async def fake_list_skills(_db=None):
|
||||
return [
|
||||
Skill(
|
||||
slug="alpha",
|
||||
name="alpha",
|
||||
description="a",
|
||||
tool_dependencies=[],
|
||||
mcp_dependencies=[],
|
||||
skill_dependencies=["beta"],
|
||||
dir_path="skills/alpha",
|
||||
),
|
||||
Skill(
|
||||
slug="beta",
|
||||
name="beta",
|
||||
description="b",
|
||||
tool_dependencies=[],
|
||||
mcp_dependencies=[],
|
||||
skill_dependencies=[],
|
||||
dir_path="skills/beta",
|
||||
),
|
||||
]
|
||||
|
||||
monkeypatch.setattr(resolver, "_list_skills_from_db", fake_list_skills)
|
||||
|
||||
snapshot = await resolver.resolve_session_snapshot([" alpha ", "alpha"])
|
||||
assert snapshot["selected_skills"] == ["alpha"]
|
||||
assert snapshot["visible_skills"] == ["alpha", "beta"]
|
||||
|
||||
assert resolver.is_snapshot_match_selected_skills(snapshot, ["alpha"]) is True
|
||||
assert resolver.is_snapshot_match_selected_skills(snapshot, ["beta"]) is False
|
||||
@ -100,14 +100,14 @@ def test_update_skill_file_passes_operator(monkeypatch):
|
||||
|
||||
|
||||
def test_dependency_options_route(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"server.routers.skill_router.get_skill_dependency_options",
|
||||
lambda: {
|
||||
async def fake_get_skill_dependency_options(_db):
|
||||
return {
|
||||
"tools": ["calculator"],
|
||||
"mcps": ["mcp-a"],
|
||||
"skills": ["demo"],
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
monkeypatch.setattr("server.routers.skill_router.get_skill_dependency_options", fake_get_skill_dependency_options)
|
||||
|
||||
app = _build_app(allow_superadmin=True)
|
||||
client = TestClient(app)
|
||||
|
||||
@ -43,72 +43,29 @@ def test_validate_skill_slug():
|
||||
svc.validate_skill_slug("../bad")
|
||||
|
||||
|
||||
def test_get_skill_prompt_metadata_by_slugs_dedup_and_skip_missing(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(
|
||||
svc,
|
||||
"_skill_prompt_metadata_cache",
|
||||
{
|
||||
"alpha": {"name": "alpha", "description": "a", "path": "/skills/alpha/SKILL.md"},
|
||||
"beta": {"name": "beta", "description": "b", "path": "/skills/beta/SKILL.md"},
|
||||
},
|
||||
)
|
||||
|
||||
result = svc.get_skill_prompt_metadata_by_slugs(["beta", "missing", "alpha", "beta"])
|
||||
assert [item["name"] for item in result] == ["beta", "alpha"]
|
||||
assert [item["path"] for item in result] == ["/skills/beta/SKILL.md", "/skills/alpha/SKILL.md"]
|
||||
|
||||
|
||||
def test_get_skill_dependency_options(monkeypatch: pytest.MonkeyPatch):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_skill_dependency_options(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(svc, "_get_buildin_tool_names", lambda: ["calculator", "search"])
|
||||
monkeypatch.setattr(svc, "get_mcp_server_names", lambda: ["mcp-a", "mcp-b"])
|
||||
monkeypatch.setattr(
|
||||
svc,
|
||||
"_skill_options_cache",
|
||||
[
|
||||
{"id": "alpha", "name": "alpha", "description": "a"},
|
||||
{"id": "beta", "name": "beta", "description": "b"},
|
||||
],
|
||||
)
|
||||
|
||||
result = svc.get_skill_dependency_options()
|
||||
class FakeRepo:
|
||||
def __init__(self, _db):
|
||||
pass
|
||||
|
||||
async def list_all(self):
|
||||
return [
|
||||
Skill(slug="alpha", name="alpha", description="a", dir_path="skills/alpha"),
|
||||
Skill(slug="beta", name="beta", description="b", dir_path="skills/beta"),
|
||||
]
|
||||
|
||||
monkeypatch.setattr(svc, "SkillRepository", FakeRepo)
|
||||
|
||||
result = await svc.get_skill_dependency_options(None)
|
||||
assert result["tools"] == ["calculator", "search"]
|
||||
assert result["mcps"] == ["mcp-a", "mcp-b"]
|
||||
assert result["skills"] == ["alpha", "beta"]
|
||||
|
||||
|
||||
def test_expand_skill_closure_and_dependency_bundle(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(
|
||||
svc,
|
||||
"_skill_dependency_cache",
|
||||
{
|
||||
"alpha": {"tools": ["t1"], "mcps": ["m1"], "skills": ["beta"]},
|
||||
"beta": {"tools": ["t2"], "mcps": ["m2"], "skills": ["gamma"]},
|
||||
"gamma": {"tools": ["t3"], "mcps": [], "skills": []},
|
||||
},
|
||||
)
|
||||
|
||||
closure = svc.expand_skill_closure(["alpha"])
|
||||
assert closure == ["alpha", "beta", "gamma"]
|
||||
|
||||
bundle = svc.get_dependency_bundle_for_activated_skills(["alpha"])
|
||||
assert bundle["skills"] == ["alpha", "beta", "gamma"]
|
||||
assert bundle["tools"] == ["t1", "t2", "t3"]
|
||||
assert bundle["mcps"] == ["m1", "m2"]
|
||||
|
||||
|
||||
def test_expand_skill_closure_cycle(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(
|
||||
svc,
|
||||
"_skill_dependency_cache",
|
||||
{
|
||||
"alpha": {"tools": [], "mcps": [], "skills": ["beta"]},
|
||||
"beta": {"tools": [], "mcps": [], "skills": ["alpha"]},
|
||||
},
|
||||
)
|
||||
# 不应抛异常,并且去重保序
|
||||
assert svc.expand_skill_closure(["alpha"]) == ["alpha", "beta"]
|
||||
|
||||
|
||||
def test_resolve_relative_path_blocks_traversal(tmp_path: Path):
|
||||
skill_dir = tmp_path / "skill"
|
||||
skill_dir.mkdir(parents=True, exist_ok=True)
|
||||
@ -158,9 +115,6 @@ async def test_import_skill_zip_conflict_rewrite_name(tmp_path: Path, monkeypatc
|
||||
self.__class__.created_item = item
|
||||
return item
|
||||
|
||||
async def list_all(self) -> list[Skill]:
|
||||
return [self.__class__.created_item] if self.__class__.created_item else []
|
||||
|
||||
monkeypatch.setattr(svc, "SkillRepository", FakeRepo)
|
||||
|
||||
zip_bytes = _build_zip(
|
||||
@ -230,9 +184,6 @@ async def test_update_skill_md_syncs_metadata(tmp_path: Path, monkeypatch: pytes
|
||||
updates["updated_by"] = updated_by
|
||||
return item
|
||||
|
||||
async def list_all(self) -> list[Skill]:
|
||||
return [item]
|
||||
|
||||
monkeypatch.setattr(svc, "get_skill_or_raise", fake_get_skill_or_raise)
|
||||
monkeypatch.setattr(svc, "SkillRepository", FakeRepo)
|
||||
|
||||
@ -271,14 +222,6 @@ async def test_update_skill_dependencies(monkeypatch: pytest.MonkeyPatch):
|
||||
)
|
||||
monkeypatch.setattr(svc, "_get_buildin_tool_names", lambda: ["calculator"])
|
||||
monkeypatch.setattr(svc, "get_mcp_server_names", lambda: ["mcp-a"])
|
||||
monkeypatch.setattr(
|
||||
svc,
|
||||
"_skill_options_cache",
|
||||
[
|
||||
{"id": "alpha", "name": "alpha", "description": "a"},
|
||||
{"id": "beta", "name": "beta", "description": "b"},
|
||||
],
|
||||
)
|
||||
|
||||
async def fake_get_skill_or_raise(_db, slug: str):
|
||||
assert slug == "alpha"
|
||||
|
||||
@ -49,11 +49,17 @@ def test_selected_skills_backend_readonly_and_visible_only_selected(tmp_path, mo
|
||||
def test_composite_backend_mounts_skills_under_prefix(tmp_path, monkeypatch):
|
||||
_prepare_skills_dir(tmp_path)
|
||||
monkeypatch.setattr(skills_backend, "get_skills_root_dir", lambda: tmp_path)
|
||||
monkeypatch.setattr(skills_backend, "get_expanded_visible_skill_slugs", lambda slugs: ["alpha", "beta"])
|
||||
|
||||
runtime = SimpleNamespace(
|
||||
context=SimpleNamespace(skills=["alpha"]),
|
||||
state={},
|
||||
state={
|
||||
"skill_session_snapshot": {
|
||||
"selected_skills": ["alpha"],
|
||||
"visible_skills": ["alpha", "beta"],
|
||||
"prompt_metadata": {},
|
||||
"dependency_map": {},
|
||||
}
|
||||
},
|
||||
)
|
||||
composite = skills_backend.create_agent_composite_backend(runtime)
|
||||
|
||||
@ -67,3 +73,17 @@ def test_composite_backend_mounts_skills_under_prefix(tmp_path, monkeypatch):
|
||||
|
||||
denied = composite.write("/skills/alpha/new.md", "x")
|
||||
assert denied.error and "read-only" in denied.error
|
||||
|
||||
|
||||
def test_composite_backend_fallbacks_to_context_skills_when_snapshot_missing(tmp_path, monkeypatch):
|
||||
_prepare_skills_dir(tmp_path)
|
||||
monkeypatch.setattr(skills_backend, "get_skills_root_dir", lambda: tmp_path)
|
||||
|
||||
runtime = SimpleNamespace(
|
||||
context=SimpleNamespace(skills=["alpha"]),
|
||||
state={},
|
||||
)
|
||||
composite = skills_backend.create_agent_composite_backend(runtime)
|
||||
skills_root = composite.ls_info("/skills/")
|
||||
skill_paths = sorted(entry.get("path") for entry in skills_root)
|
||||
assert skill_paths == ["/skills/alpha/"]
|
||||
|
||||
@ -372,6 +372,7 @@ import ModelSelectorComponent from '@/components/ModelSelectorComponent.vue'
|
||||
import { useAgentStore } from '@/stores/agent'
|
||||
import { useUserStore } from '@/stores/user'
|
||||
import { useDatabaseStore } from '@/stores/database'
|
||||
import { skillApi } from '@/apis/skill_api'
|
||||
import { storeToRefs } from 'pinia'
|
||||
|
||||
// Props
|
||||
@ -399,6 +400,7 @@ watch(
|
||||
async (val) => {
|
||||
if (val) {
|
||||
databaseStore.loadDatabases().catch(() => {})
|
||||
loadLiveSkillOptions().catch(() => {})
|
||||
if (selectedAgentId.value) {
|
||||
try {
|
||||
await agentStore.fetchAgentDetail(selectedAgentId.value, true)
|
||||
@ -430,6 +432,7 @@ const tempSelectedValues = ref([])
|
||||
const selectionSearchText = ref('')
|
||||
const systemPromptEditMode = ref(false)
|
||||
const activeTab = ref('basic')
|
||||
const liveSkillOptions = ref([])
|
||||
|
||||
const isEmptyConfig = computed(() => {
|
||||
return !selectedAgentId.value || Object.keys(configurableItems.value).length === 0
|
||||
@ -470,6 +473,24 @@ const segmentedOptions = computed(() => {
|
||||
return options
|
||||
})
|
||||
|
||||
const loadLiveSkillOptions = async () => {
|
||||
if (!userStore.isAdmin) {
|
||||
liveSkillOptions.value = []
|
||||
return
|
||||
}
|
||||
try {
|
||||
const result = await skillApi.listSkills()
|
||||
const rows = result?.data || []
|
||||
liveSkillOptions.value = rows.map((item) => ({
|
||||
id: item.slug,
|
||||
name: item.slug,
|
||||
description: item.description || ''
|
||||
}))
|
||||
} catch (error) {
|
||||
console.warn('加载实时 Skills 列表失败:', error)
|
||||
}
|
||||
}
|
||||
|
||||
// 通用选项获取与处理
|
||||
const getConfigOptions = (value) => {
|
||||
if (value?.template_metadata?.kind === 'tools') {
|
||||
@ -478,6 +499,9 @@ const getConfigOptions = (value) => {
|
||||
if (value?.template_metadata?.kind === 'knowledges') {
|
||||
return databaseStore.databases || []
|
||||
}
|
||||
if (value?.template_metadata?.kind === 'skills') {
|
||||
return liveSkillOptions.value.length > 0 ? liveSkillOptions.value : value?.options || []
|
||||
}
|
||||
return value?.options || []
|
||||
}
|
||||
|
||||
@ -636,12 +660,8 @@ const openSelectionModal = async (key) => {
|
||||
console.error('加载知识库列表失败:', error)
|
||||
}
|
||||
}
|
||||
if (configurableItems.value[key]?.template_metadata?.kind === 'skills' && selectedAgentId.value) {
|
||||
try {
|
||||
await agentStore.fetchAgentDetail(selectedAgentId.value, true)
|
||||
} catch (error) {
|
||||
console.error('刷新 Skills 列表失败:', error)
|
||||
}
|
||||
if (configurableItems.value[key]?.template_metadata?.kind === 'skills') {
|
||||
await loadLiveSkillOptions()
|
||||
}
|
||||
const currentValues = agentConfig.value[key] || []
|
||||
tempSelectedValues.value = [...currentValues]
|
||||
|
||||
Loading…
Reference in New Issue
Block a user