mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(SkillsHandler): Fix errors and skills handling between models
This commit is contained in:
parent
b3603f5cc6
commit
ac7db4d2b2
6 changed files with 201 additions and 47 deletions
|
|
@ -71,6 +71,17 @@ class LiteLLMSkillsHandler:
|
|||
"""
|
||||
prisma_client = await LiteLLMSkillsHandler._get_prisma_client()
|
||||
|
||||
# Enforce unique display_title
|
||||
if data.display_title:
|
||||
existing = await prisma_client.db.litellm_skillstable.find_first(
|
||||
where={"display_title": data.display_title}
|
||||
)
|
||||
if existing is not None:
|
||||
raise ValueError(
|
||||
f"A skill with display_title '{data.display_title}' already exists "
|
||||
f"(id: {existing.skill_id}). Skill names must be unique."
|
||||
)
|
||||
|
||||
skill_id = f"litellm_skill_{uuid.uuid4()}"
|
||||
|
||||
skill_data: Dict[str, Any] = {
|
||||
|
|
|
|||
|
|
@ -165,12 +165,27 @@ def get_provider_from_model(model: str) -> str:
|
|||
"""
|
||||
Determine the provider from a model string.
|
||||
|
||||
Uses LiteLLM's get_llm_provider to resolve the provider.
|
||||
First checks the proxy router's model list to resolve aliases
|
||||
(e.g., "claude-sonnet" -> "anthropic/claude-sonnet-4-20250514"),
|
||||
then uses get_llm_provider on the resolved model.
|
||||
"""
|
||||
resolved_model = model
|
||||
|
||||
# Try to resolve through the router's model list
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if llm_router is not None:
|
||||
deployments = llm_router.get_model_list(model_name=model)
|
||||
if deployments:
|
||||
resolved_model = deployments[0]["litellm_params"]["model"]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model=model)
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model=resolved_model)
|
||||
return custom_llm_provider or "openai"
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
|
|
|
|||
|
|
@ -126,10 +126,13 @@ async def _handle_litellm_create_skill(
|
|||
)
|
||||
|
||||
# Create skill in DB
|
||||
skill_record = await LiteLLMSkillsHandler.create_skill(
|
||||
data=skill_request,
|
||||
user_id=user_api_key_dict.user_id,
|
||||
)
|
||||
try:
|
||||
skill_record = await LiteLLMSkillsHandler.create_skill(
|
||||
data=skill_request,
|
||||
user_id=user_api_key_dict.user_id,
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=409, detail=str(e))
|
||||
|
||||
verbose_proxy_logger.debug(f"Created LiteLLM skill: {skill_record.skill_id}")
|
||||
|
||||
|
|
|
|||
|
|
@ -97,42 +97,63 @@ class SkillsInjectionHook(CustomLogger):
|
|||
f"SkillsInjectionHook: Processing {len(skills)} skills"
|
||||
)
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.llms.litellm_proxy.skills.skill_applicator import (
|
||||
SkillApplicator,
|
||||
get_provider_from_model,
|
||||
)
|
||||
|
||||
model = data.get("model", "")
|
||||
provider = get_provider_from_model(model)
|
||||
applicator = SkillApplicator()
|
||||
|
||||
litellm_skills: List[LiteLLM_SkillsTable] = []
|
||||
anthropic_skills: List[Dict[str, Any]] = []
|
||||
|
||||
# Separate skills by prefix
|
||||
# Classify and validate skills
|
||||
for skill in skills:
|
||||
if not isinstance(skill, dict):
|
||||
continue
|
||||
|
||||
skill_id = skill.get("skill_id", "")
|
||||
if skill_id.startswith("litellm_"):
|
||||
# Fetch from LiteLLM DB
|
||||
|
||||
if skill_id.startswith("litellm_skill_"):
|
||||
# LiteLLM gateway-managed skill — fetch from DB
|
||||
db_skill = await self._fetch_skill_from_db(skill_id)
|
||||
if db_skill:
|
||||
litellm_skills.append(db_skill)
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
f"SkillsInjectionHook: Skill '{skill_id}' not found in LiteLLM DB"
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Skill not found: {skill_id}",
|
||||
)
|
||||
elif skill_id.startswith("skill_"):
|
||||
# Native Anthropic skill — only allowed with native-skills providers
|
||||
if not applicator.supports_native_skills(provider):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Anthropic skill '{skill_id}' cannot be used with "
|
||||
f"model '{model}' (provider '{provider}' does not support "
|
||||
f"native skills). Use a litellm_skill_* ID instead.",
|
||||
)
|
||||
else:
|
||||
# Native Anthropic skill - pass through
|
||||
anthropic_skills.append(skill)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid skill_id '{skill_id}'. Must start with "
|
||||
f"'litellm_skill_' (gateway skill) or 'skill_' (Anthropic native).",
|
||||
)
|
||||
|
||||
if len(litellm_skills) > 0:
|
||||
# Determine provider to pick the right strategy
|
||||
from litellm.llms.litellm_proxy.skills.skill_applicator import (
|
||||
SkillApplicator,
|
||||
get_provider_from_model,
|
||||
)
|
||||
|
||||
model = data.get("model", "")
|
||||
provider = get_provider_from_model(model)
|
||||
applicator = SkillApplicator()
|
||||
# When the request comes through /v1/messages (anthropic_messages),
|
||||
# we must inject into the top-level 'system' param because
|
||||
# anthropic_messages() has separate 'messages' and 'system' params.
|
||||
use_anthropic_format = call_type == "anthropic_messages"
|
||||
|
||||
if applicator.supports_native_skills(provider):
|
||||
# Native skills path: convert to tools + system prompt
|
||||
use_anthropic_format = call_type == "anthropic_messages"
|
||||
data = self._process_for_messages_api(
|
||||
data=data,
|
||||
litellm_skills=litellm_skills,
|
||||
|
|
@ -140,11 +161,16 @@ class SkillsInjectionHook(CustomLogger):
|
|||
)
|
||||
else:
|
||||
# Non-native path: inject into system prompt only
|
||||
data = await applicator.apply_skills(
|
||||
data=data,
|
||||
skills=litellm_skills,
|
||||
provider=provider,
|
||||
)
|
||||
skill_contents = []
|
||||
for skill in litellm_skills:
|
||||
content = applicator._format_skill_content(skill)
|
||||
if content:
|
||||
skill_contents.append(content)
|
||||
|
||||
if skill_contents:
|
||||
data = self.prompt_handler.inject_skill_content_to_messages(
|
||||
data, skill_contents, use_anthropic_format=use_anthropic_format
|
||||
)
|
||||
# Remove container (not supported by underlying providers)
|
||||
data.pop("container", None)
|
||||
|
||||
|
|
|
|||
|
|
@ -53,6 +53,7 @@ class TestCreateSkill:
|
|||
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_skillstable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_skillstable.create = AsyncMock(
|
||||
return_value=_make_prisma_skill()
|
||||
)
|
||||
|
|
@ -87,6 +88,7 @@ class TestCreateSkill:
|
|||
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_skillstable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_skillstable.create = AsyncMock(
|
||||
return_value=_make_prisma_skill()
|
||||
)
|
||||
|
|
@ -123,6 +125,36 @@ class TestCreateSkill:
|
|||
with pytest.raises(ValueError, match="Prisma client"):
|
||||
await LiteLLMSkillsHandler.create_skill(data=request)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_skill_duplicate_title_raises(self):
|
||||
"""Test that creating a skill with a duplicate display_title raises ValueError."""
|
||||
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
|
||||
|
||||
existing_skill = MagicMock()
|
||||
existing_skill.skill_id = "litellm_skill_existing"
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_skillstable.find_first = AsyncMock(
|
||||
return_value=existing_skill
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
LiteLLMSkillsHandler,
|
||||
"_get_prisma_client",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_prisma,
|
||||
):
|
||||
request = NewSkillRequest(
|
||||
display_title="Duplicate Name",
|
||||
instructions="test",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="already exists"):
|
||||
await LiteLLMSkillsHandler.create_skill(data=request)
|
||||
|
||||
# Should never reach create
|
||||
mock_prisma.db.litellm_skillstable.create.assert_not_called()
|
||||
|
||||
|
||||
class TestListSkills:
|
||||
"""Tests for LiteLLMSkillsHandler.list_skills."""
|
||||
|
|
|
|||
|
|
@ -166,8 +166,10 @@ class TestPreCallHookOptIn:
|
|||
assert "container" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_skill_logged_but_continues(self):
|
||||
"""Test that missing skills don't break the request."""
|
||||
async def test_missing_skill_raises_404(self):
|
||||
"""Test that referencing a nonexistent skill raises 404."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
hook = SkillsInjectionHook()
|
||||
|
||||
with patch.object(
|
||||
|
|
@ -181,6 +183,34 @@ class TestPreCallHookOptIn:
|
|||
},
|
||||
}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert "litellm_skill_nonexistent" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_native_skill_on_native_provider(self):
|
||||
"""Test that Anthropic skill_ IDs pass through on native providers."""
|
||||
hook = SkillsInjectionHook()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.skill_applicator.get_provider_from_model",
|
||||
return_value="anthropic",
|
||||
):
|
||||
data = {
|
||||
"model": "claude-sonnet-4-20250514",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"container": {
|
||||
"skills": [{"skill_id": "skill_01abc123", "type": "anthropic"}]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
|
|
@ -188,31 +218,68 @@ class TestPreCallHookOptIn:
|
|||
call_type="completion",
|
||||
)
|
||||
|
||||
# Should still succeed, just without skill injection
|
||||
# Anthropic native skill should pass through
|
||||
assert isinstance(result, dict)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_litellm_skill_treated_as_anthropic_native(self):
|
||||
"""Test that skills without litellm_ prefix are treated as native Anthropic."""
|
||||
async def test_anthropic_native_skill_on_non_native_provider_fails(self):
|
||||
"""Test that Anthropic skill_ IDs fail on non-native providers."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
hook = SkillsInjectionHook()
|
||||
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"container": {
|
||||
"skills": [{"skill_id": "anthropic_native_skill_123", "type": "anthropic"}]
|
||||
},
|
||||
}
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.skill_applicator.get_provider_from_model",
|
||||
return_value="openai",
|
||||
):
|
||||
data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"container": {
|
||||
"skills": [{"skill_id": "skill_01abc123", "type": "anthropic"}]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
# No litellm skills found, so no processing should happen
|
||||
assert isinstance(result, dict)
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "does not support native skills" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_skill_id_prefix_fails(self):
|
||||
"""Test that skill IDs with unrecognized prefixes fail."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
hook = SkillsInjectionHook()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.skill_applicator.get_provider_from_model",
|
||||
return_value="openai",
|
||||
):
|
||||
data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"container": {
|
||||
"skills": [{"skill_id": "random_garbage_id", "type": "anthropic"}]
|
||||
},
|
||||
}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Invalid skill_id" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
class TestSystemPromptInjection:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue