fix(SkillsHandler): correct mapping for skills availability to provider

This commit is contained in:
Fred Chasin 2026-03-21 14:55:17 -07:00
parent 4d543360d1
commit 863be177b6
5 changed files with 263 additions and 143 deletions

View file

@ -58,7 +58,6 @@ class LiteLLMSkillsHandler:
async def create_skill(
data: NewSkillRequest,
user_id: Optional[str] = None,
source: str = "custom",
) -> LiteLLM_SkillsTable:
"""
Create a new skill in the LiteLLM database.
@ -66,7 +65,6 @@ class LiteLLMSkillsHandler:
Args:
data: NewSkillRequest with skill details
user_id: Optional user ID for tracking
source: Source of the skill ("custom", "gateway", etc.)
Returns:
LiteLLM_SkillsTable record
@ -80,7 +78,6 @@ class LiteLLMSkillsHandler:
"display_title": data.display_title,
"description": data.description,
"instructions": data.instructions,
"source": source,
"created_by": user_id,
"updated_by": user_id,
}
@ -113,7 +110,6 @@ class LiteLLMSkillsHandler:
async def list_skills(
limit: int = 20,
offset: int = 0,
source: Optional[str] = None,
) -> List[LiteLLM_SkillsTable]:
"""
List skills from the LiteLLM database.
@ -121,7 +117,6 @@ class LiteLLMSkillsHandler:
Args:
limit: Maximum number of skills to return
offset: Number of skills to skip
source: Optional filter by source ("custom", "gateway", etc.)
Returns:
List of LiteLLM_SkillsTable records
@ -129,16 +124,10 @@ class LiteLLMSkillsHandler:
prisma_client = await LiteLLMSkillsHandler._get_prisma_client()
verbose_logger.debug(
f"LiteLLMSkillsHandler: Listing skills with limit={limit}, offset={offset}, source={source}"
f"LiteLLMSkillsHandler: Listing skills with limit={limit}, offset={offset}"
)
# Build where clause
where_clause: Dict[str, Any] = {}
if source is not None:
where_clause["source"] = source
skills = await prisma_client.db.litellm_skillstable.find_many(
where=where_clause if where_clause else None,
take=limit,
skip=offset,
order={"created_at": "desc"},
@ -227,3 +216,76 @@ class LiteLLMSkillsHandler:
f"LiteLLMSkillsHandler: Error fetching skill {skill_id}: {e}"
)
return None
@staticmethod
async def save_provider_skill_id(
skill_id: str,
provider: str,
provider_skill_id: str,
) -> None:
"""
Save a provider-assigned skill ID for a LiteLLM skill.
When a skill is used with a provider that has a native skills API
(e.g. Anthropic), the provider returns its own skill ID. We store
that mapping so subsequent calls can reuse it without re-creating.
Stored in metadata._provider_skill_ids.{provider} = provider_skill_id
Args:
skill_id: The LiteLLM skill ID
provider: Provider name (e.g. "anthropic")
provider_skill_id: The ID assigned by the provider
"""
import json
prisma_client = await LiteLLMSkillsHandler._get_prisma_client()
# Fetch current metadata
skill = await prisma_client.db.litellm_skillstable.find_unique(
where={"skill_id": skill_id}
)
if skill is None:
verbose_logger.warning(
f"LiteLLMSkillsHandler: Cannot save provider ID - skill {skill_id} not found"
)
return
metadata = skill.metadata if isinstance(skill.metadata, dict) else {}
if isinstance(metadata, str):
metadata = json.loads(metadata)
provider_ids = metadata.get("_provider_skill_ids", {})
provider_ids[provider] = provider_skill_id
metadata["_provider_skill_ids"] = provider_ids
await prisma_client.db.litellm_skillstable.update(
where={"skill_id": skill_id},
data={"metadata": json.dumps(metadata)},
)
verbose_logger.debug(
f"LiteLLMSkillsHandler: Saved {provider} skill ID "
f"'{provider_skill_id}' for skill {skill_id}"
)
@staticmethod
def get_provider_skill_id(
skill: LiteLLM_SkillsTable,
provider: str,
) -> Optional[str]:
"""
Get the provider-assigned skill ID from a skill's metadata.
Args:
skill: The LiteLLM skill record
provider: Provider name (e.g. "anthropic")
Returns:
The provider's skill ID, or None if not yet registered
"""
if not skill.metadata or not isinstance(skill.metadata, dict):
return None
provider_ids = skill.metadata.get("_provider_skill_ids", {})
return provider_ids.get(provider)

View file

@ -2,12 +2,16 @@
Skill Applicator for Gateway Skills.
Handles provider-specific strategies for applying skills to LLM requests.
Routes skills to appropriate injection method based on model provider.
Uses get_llm_provider() to resolve models to providers, then checks the
centralized beta headers config to determine if the model's provider
supports native skills (skills-2025-10-02 beta). If not, falls back to
system prompt injection.
"""
from typing import Any, Dict, List, Optional
from typing import List, Optional
from litellm._logging import verbose_logger
from litellm.constants import ANTHROPIC_SKILLS_API_BETA_VERSION
from litellm.proxy._types import LiteLLM_SkillsTable
@ -15,43 +19,12 @@ class SkillApplicator:
"""
Applies gateway skills to LLM requests using provider-specific strategies.
Strategies:
- system_prompt: Inject skill content into system message (OpenAI, Azure, Bedrock, etc.)
- tool_conversion: Convert skills to tools + system prompt (existing behavior)
Provider resolution is delegated to litellm.get_llm_provider().
Native skills support is determined by the centralized beta headers
config (anthropic_beta_headers_config.json) if the provider maps
skills-2025-10-02 to a non-null value, native skills are supported.
"""
# Provider to strategy mapping
PROVIDER_STRATEGIES: Dict[str, str] = {
# System prompt injection
"openai": "system_prompt",
"azure": "system_prompt",
"azure_ai": "system_prompt",
"bedrock": "system_prompt",
"vertex_ai": "system_prompt",
"vertex_ai_beta": "system_prompt",
"gemini": "system_prompt",
"ollama": "system_prompt",
"ollama_chat": "system_prompt",
"groq": "system_prompt",
"together_ai": "system_prompt",
"deepseek": "system_prompt",
"fireworks_ai": "system_prompt",
"mistral": "system_prompt",
"cohere": "system_prompt",
"cohere_chat": "system_prompt",
"ai21": "system_prompt",
"replicate": "system_prompt",
"sagemaker": "system_prompt",
"perplexity": "system_prompt",
"anyscale": "system_prompt",
"openrouter": "system_prompt",
"huggingface": "system_prompt",
"text-completion-openai": "system_prompt",
"text-completion-codestral": "system_prompt",
# Tool conversion (existing Anthropic behavior)
"anthropic": "tool_conversion",
}
def __init__(self):
from litellm.llms.litellm_proxy.skills.prompt_injection import (
SkillPromptInjectionHandler,
@ -59,17 +32,17 @@ class SkillApplicator:
self.prompt_handler = SkillPromptInjectionHandler()
def get_strategy(self, provider: str) -> str:
def supports_native_skills(self, provider: str) -> bool:
"""
Get the skill application strategy for a provider.
Args:
provider: The LLM provider name
Returns:
Strategy name ("system_prompt" or "tool_conversion")
Check if a provider supports native skills by consulting the
centralized beta headers config.
"""
return self.PROVIDER_STRATEGIES.get(provider, "system_prompt")
from litellm.anthropic_beta_headers_manager import is_beta_header_supported
return is_beta_header_supported(
beta_header=ANTHROPIC_SKILLS_API_BETA_VERSION,
provider=provider,
)
async def apply_skills(
self,
@ -78,12 +51,12 @@ class SkillApplicator:
provider: str,
) -> dict:
"""
Apply skills to a request based on provider strategy.
Apply skills to a request based on provider.
Args:
data: The request data dict
skills: List of skills to apply
provider: The LLM provider name
provider: The LLM provider name (from get_llm_provider)
Returns:
Modified request data with skills applied
@ -91,22 +64,18 @@ class SkillApplicator:
if not skills:
return data
strategy = self.get_strategy(provider)
if self.supports_native_skills(provider):
verbose_logger.debug(
f"SkillApplicator: Applying {len(skills)} skills via native API "
f"for provider={provider}"
)
return self._apply_tool_conversion_strategy(data, skills)
verbose_logger.debug(
f"SkillApplicator: Applying {len(skills)} skills with strategy={strategy} "
f"SkillApplicator: Applying {len(skills)} skills via system prompt "
f"for provider={provider}"
)
if strategy == "system_prompt":
return self._apply_system_prompt_strategy(data, skills)
elif strategy == "tool_conversion":
return self._apply_tool_conversion_strategy(data, skills)
else:
verbose_logger.warning(
f"SkillApplicator: Unknown strategy {strategy}, using system_prompt"
)
return self._apply_system_prompt_strategy(data, skills)
return self._apply_system_prompt_strategy(data, skills)
def _apply_system_prompt_strategy(
self,
@ -116,7 +85,7 @@ class SkillApplicator:
"""
Apply skills by injecting content into system prompt.
This strategy appends skill content to the system message:
Format:
---
## Skill: {display_title}
**Description:** {description}
@ -124,13 +93,6 @@ class SkillApplicator:
### Instructions
{SKILL.md body content}
---
Args:
data: The request data dict
skills: List of skills to apply
Returns:
Modified request data with skill content in system message
"""
skill_contents: List[str] = []
@ -142,7 +104,6 @@ class SkillApplicator:
if not skill_contents:
return data
# Inject into system message
return self.prompt_handler.inject_skill_content_to_messages(
data, skill_contents, use_anthropic_format=False
)
@ -153,26 +114,14 @@ class SkillApplicator:
skills: List[LiteLLM_SkillsTable],
) -> dict:
"""
Apply skills by converting to tools (existing Anthropic behavior).
This uses the existing SkillPromptInjectionHandler logic for
tool conversion and skill content injection.
Args:
data: The request data dict
skills: List of skills to apply
Returns:
Modified request data with tools and skill content
Apply skills by converting to Anthropic-style tools + system prompt.
"""
tools = data.get("tools", [])
skill_contents: List[str] = []
for skill in skills:
# Convert skill to Anthropic-style tool
tools.append(self.prompt_handler.convert_skill_to_anthropic_tool(skill))
# Extract skill content
content = self.prompt_handler.extract_skill_content(skill)
if content:
skill_contents.append(content)
@ -181,7 +130,6 @@ class SkillApplicator:
data["tools"] = tools
if skill_contents:
# For Anthropic, use top-level 'system' param
data = self.prompt_handler.inject_skill_content_to_messages(
data, skill_contents, use_anthropic_format=True
)
@ -191,33 +139,15 @@ class SkillApplicator:
def _format_skill_content(self, skill: LiteLLM_SkillsTable) -> Optional[str]:
"""
Format skill content for system prompt injection.
Produces:
---
## Skill: {display_title}
**Description:** {description}
### Instructions
{extracted content or instructions}
---
Args:
skill: The skill to format
Returns:
Formatted skill content string, or None if no content
"""
# Try to extract content from file first
content = self.prompt_handler.extract_skill_content(skill)
# Fall back to instructions
if not content:
content = skill.instructions
if not content:
return None
# Build formatted skill section
title = skill.display_title or skill.skill_id
parts = [f"## Skill: {title}"]
@ -235,13 +165,7 @@ def get_provider_from_model(model: str) -> str:
"""
Determine the provider from a model string.
Uses LiteLLM's get_llm_provider function to resolve the provider.
Args:
model: The model identifier
Returns:
The provider name (e.g., "openai", "anthropic", "bedrock")
Uses LiteLLM's get_llm_provider to resolve the provider.
"""
try:
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
@ -252,5 +176,4 @@ def get_provider_from_model(model: str) -> str:
verbose_logger.warning(
f"SkillApplicator: Failed to determine provider for model {model}: {e}"
)
# Default to OpenAI-compatible
return "openai"

View file

@ -124,7 +124,6 @@ async def _handle_litellm_create_skill(
skill_record = await LiteLLMSkillsHandler.create_skill(
data=skill_request,
user_id=user_api_key_dict.user_id,
source="litellm",
)
verbose_proxy_logger.debug(
@ -134,7 +133,7 @@ async def _handle_litellm_create_skill(
return Skill(
id=skill_record.skill_id,
display_title=skill_record.display_title,
source="litellm",
source=skill_record.source,
latest_version=skill_record.latest_version,
created_at=skill_record.created_at.isoformat() if skill_record.created_at else "",
updated_at=skill_record.updated_at.isoformat() if skill_record.updated_at else "",

View file

@ -73,7 +73,7 @@ class TestCreateSkill:
)
result = await LiteLLMSkillsHandler.create_skill(
data=request, user_id="user1", source="custom"
data=request, user_id="user1"
)
assert isinstance(result, LiteLLM_SkillsTable)
@ -169,8 +169,8 @@ class TestListSkills:
assert results == []
@pytest.mark.asyncio
async def test_list_skills_with_source_filter(self):
"""Test listing skills passes source filter to Prisma."""
async def test_list_skills_with_pagination(self):
"""Test listing skills passes limit and offset to Prisma."""
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
mock_prisma = MagicMock()
@ -182,10 +182,11 @@ class TestListSkills:
new_callable=AsyncMock,
return_value=mock_prisma,
):
await LiteLLMSkillsHandler.list_skills(source="custom")
await LiteLLMSkillsHandler.list_skills(limit=5, offset=10)
call_args = mock_prisma.db.litellm_skillstable.find_many.call_args
assert call_args[1]["where"] == {"source": "custom"}
assert call_args[1]["take"] == 5
assert call_args[1]["skip"] == 10
class TestGetSkill:
@ -307,3 +308,145 @@ class TestFetchSkillFromDb:
):
result = await LiteLLMSkillsHandler.fetch_skill_from_db("any_id")
assert result is None
class TestProviderSkillIds:
"""Tests for provider skill ID save/retrieve."""
@pytest.mark.asyncio
async def test_save_provider_skill_id(self):
"""Test saving an Anthropic skill ID for a LiteLLM skill."""
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
mock_skill = MagicMock()
mock_skill.metadata = {}
mock_prisma = MagicMock()
mock_prisma.db.litellm_skillstable.find_unique = AsyncMock(
return_value=mock_skill
)
mock_prisma.db.litellm_skillstable.update = AsyncMock()
with patch.object(
LiteLLMSkillsHandler,
"_get_prisma_client",
new_callable=AsyncMock,
return_value=mock_prisma,
):
await LiteLLMSkillsHandler.save_provider_skill_id(
skill_id="litellm_skill_abc",
provider="anthropic",
provider_skill_id="sk_ant_123",
)
call_args = mock_prisma.db.litellm_skillstable.update.call_args
assert call_args[1]["where"] == {"skill_id": "litellm_skill_abc"}
import json
saved_metadata = json.loads(call_args[1]["data"]["metadata"])
assert saved_metadata["_provider_skill_ids"]["anthropic"] == "sk_ant_123"
@pytest.mark.asyncio
async def test_save_provider_skill_id_preserves_existing_metadata(self):
"""Test that saving a provider ID doesn't clobber existing metadata."""
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
mock_skill = MagicMock()
mock_skill.metadata = {"custom_key": "custom_value"}
mock_prisma = MagicMock()
mock_prisma.db.litellm_skillstable.find_unique = AsyncMock(
return_value=mock_skill
)
mock_prisma.db.litellm_skillstable.update = AsyncMock()
with patch.object(
LiteLLMSkillsHandler,
"_get_prisma_client",
new_callable=AsyncMock,
return_value=mock_prisma,
):
await LiteLLMSkillsHandler.save_provider_skill_id(
skill_id="litellm_skill_abc",
provider="anthropic",
provider_skill_id="sk_ant_456",
)
call_args = mock_prisma.db.litellm_skillstable.update.call_args
import json
saved_metadata = json.loads(call_args[1]["data"]["metadata"])
assert saved_metadata["custom_key"] == "custom_value"
assert saved_metadata["_provider_skill_ids"]["anthropic"] == "sk_ant_456"
@pytest.mark.asyncio
async def test_save_provider_skill_id_skill_not_found(self):
"""Test that saving provider ID for missing skill doesn't raise."""
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
mock_prisma = MagicMock()
mock_prisma.db.litellm_skillstable.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_skillstable.update = AsyncMock()
with patch.object(
LiteLLMSkillsHandler,
"_get_prisma_client",
new_callable=AsyncMock,
return_value=mock_prisma,
):
# Should not raise
await LiteLLMSkillsHandler.save_provider_skill_id(
skill_id="nonexistent",
provider="anthropic",
provider_skill_id="sk_ant_789",
)
mock_prisma.db.litellm_skillstable.update.assert_not_called()
def test_get_provider_skill_id_found(self):
"""Test retrieving a saved Anthropic skill ID."""
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
skill = LiteLLM_SkillsTable(
skill_id="litellm_skill_abc",
metadata={"_provider_skill_ids": {"anthropic": "sk_ant_123"}},
)
result = LiteLLMSkillsHandler.get_provider_skill_id(skill, "anthropic")
assert result == "sk_ant_123"
def test_get_provider_skill_id_not_found(self):
"""Test retrieving provider ID when none saved."""
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
skill = LiteLLM_SkillsTable(
skill_id="litellm_skill_abc",
metadata={},
)
result = LiteLLMSkillsHandler.get_provider_skill_id(skill, "anthropic")
assert result is None
def test_get_provider_skill_id_no_metadata(self):
"""Test retrieving provider ID when metadata is None."""
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
skill = LiteLLM_SkillsTable(
skill_id="litellm_skill_abc",
metadata=None,
)
result = LiteLLMSkillsHandler.get_provider_skill_id(skill, "anthropic")
assert result is None
def test_get_provider_skill_id_different_provider(self):
"""Test that provider IDs are namespaced per provider."""
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
skill = LiteLLM_SkillsTable(
skill_id="litellm_skill_abc",
metadata={"_provider_skill_ids": {"anthropic": "sk_ant_123"}},
)
assert LiteLLMSkillsHandler.get_provider_skill_id(skill, "anthropic") == "sk_ant_123"
assert LiteLLMSkillsHandler.get_provider_skill_id(skill, "openai") is None

View file

@ -451,25 +451,18 @@ class TestSkillApplicator:
assert len(result["messages"]) == 1
assert result["messages"][0]["role"] == "user"
def test_get_strategy_known_providers(self):
"""Test strategy mapping for known providers."""
def test_supports_native_skills(self):
"""Test native skills support detection."""
from litellm.llms.litellm_proxy.skills.skill_applicator import SkillApplicator
applicator = SkillApplicator()
assert applicator.get_strategy("openai") == "system_prompt"
assert applicator.get_strategy("azure") == "system_prompt"
assert applicator.get_strategy("bedrock") == "system_prompt"
assert applicator.get_strategy("gemini") == "system_prompt"
assert applicator.get_strategy("anthropic") == "tool_conversion"
def test_get_strategy_unknown_defaults_to_system_prompt(self):
"""Test unknown provider defaults to system_prompt."""
from litellm.llms.litellm_proxy.skills.skill_applicator import SkillApplicator
applicator = SkillApplicator()
assert applicator.get_strategy("some_new_provider") == "system_prompt"
assert applicator.supports_native_skills("anthropic") is True
assert applicator.supports_native_skills("openai") is False
assert applicator.supports_native_skills("azure") is False
assert applicator.supports_native_skills("bedrock") is False
assert applicator.supports_native_skills("gemini") is False
assert applicator.supports_native_skills("some_new_provider") is False
class TestSkillContentExtraction: