mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(SkillsHandler): Fix skillshandler gateway and add additional testes
This commit is contained in:
parent
0116ece0ef
commit
b3603f5cc6
4 changed files with 661 additions and 11 deletions
|
|
@ -9,7 +9,8 @@ Supports two modes controlled by litellm_settings.skills_mode:
|
|||
from typing import Literal, Optional
|
||||
|
||||
import orjson
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, UploadFile
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from starlette.datastructures import UploadFile
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import NewSkillRequest, UserAPIKeyAuth
|
||||
|
|
@ -68,9 +69,10 @@ async def _handle_litellm_create_skill(
|
|||
display_title_override = form_data.get("display_title")
|
||||
|
||||
# Get files from form data
|
||||
files_data = form_data.get("files[]", [])
|
||||
# get_form_data strips [] suffix, so "files[]" becomes "files"
|
||||
files_data = form_data.get("files", [])
|
||||
if not files_data:
|
||||
files_data = form_data.get("files", [])
|
||||
files_data = form_data.get("files[]", [])
|
||||
|
||||
if not files_data:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -119,17 +119,35 @@ class SkillsInjectionHook(CustomLogger):
|
|||
# Native Anthropic skill - pass through
|
||||
anthropic_skills.append(skill)
|
||||
|
||||
# Check if using messages API spec (anthropic_messages call type)
|
||||
# Messages API always uses Anthropic-style tool format
|
||||
use_anthropic_format = call_type == "anthropic_messages"
|
||||
|
||||
if len(litellm_skills) > 0:
|
||||
data = self._process_for_messages_api(
|
||||
data=data,
|
||||
litellm_skills=litellm_skills,
|
||||
use_anthropic_format=use_anthropic_format,
|
||||
# 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()
|
||||
|
||||
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,
|
||||
use_anthropic_format=use_anthropic_format,
|
||||
)
|
||||
else:
|
||||
# Non-native path: inject into system prompt only
|
||||
data = await applicator.apply_skills(
|
||||
data=data,
|
||||
skills=litellm_skills,
|
||||
provider=provider,
|
||||
)
|
||||
# Remove container (not supported by underlying providers)
|
||||
data.pop("container", None)
|
||||
|
||||
return data
|
||||
|
||||
def _process_for_messages_api(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,346 @@
|
|||
"""
|
||||
Tests for the create skill endpoint form data parsing.
|
||||
|
||||
Simulates actual curl/multipart uploads to verify the endpoint
|
||||
correctly handles file uploads in litellm mode.
|
||||
"""
|
||||
|
||||
import io
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import UploadFile
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.proxy._types import LiteLLM_SkillsTable, UserAPIKeyAuth
|
||||
|
||||
|
||||
def _make_upload_file(filename: str, content: bytes) -> UploadFile:
|
||||
"""Create a FastAPI UploadFile matching what multipart form parsing produces."""
|
||||
return UploadFile(
|
||||
filename=filename,
|
||||
file=io.BytesIO(content),
|
||||
headers=Headers({"content-type": "application/octet-stream"}),
|
||||
)
|
||||
|
||||
|
||||
def _make_db_skill(**overrides) -> LiteLLM_SkillsTable:
|
||||
"""Create a mock DB skill record."""
|
||||
from datetime import datetime
|
||||
|
||||
defaults = {
|
||||
"skill_id": "litellm_skill_test123",
|
||||
"display_title": "Test Skill",
|
||||
"description": None,
|
||||
"instructions": "Test instructions",
|
||||
"source": "custom",
|
||||
"file_content": b"fake-zip",
|
||||
"file_name": "skill.zip",
|
||||
"file_type": "application/zip",
|
||||
"created_at": datetime(2026, 3, 21),
|
||||
"updated_at": datetime(2026, 3, 21),
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return LiteLLM_SkillsTable(**defaults)
|
||||
|
||||
|
||||
SKILL_MD_CONTENT = b"""---
|
||||
name: test-skill
|
||||
description: A test skill
|
||||
---
|
||||
|
||||
Test instructions here.
|
||||
"""
|
||||
|
||||
|
||||
class TestCreateSkillFormParsing:
|
||||
"""Tests that simulate actual curl multipart uploads."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_files_bracket_key_single_upload(self):
|
||||
"""
|
||||
Simulate: curl -F "files[]=@SKILL.md;filename=skill/SKILL.md"
|
||||
|
||||
get_form_data strips [] so key becomes "files" with value in a list.
|
||||
"""
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
upload = _make_upload_file("skill/SKILL.md", SKILL_MD_CONTENT)
|
||||
|
||||
mock_request = MagicMock()
|
||||
# get_form_data strips [] and appends to list
|
||||
mock_request.form = AsyncMock(return_value={"files[]": upload})
|
||||
|
||||
mock_db_skill = _make_db_skill(display_title="test-skill")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"files": [upload], "display_title": "Test Skill"},
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.handler.LiteLLMSkillsHandler.create_skill",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_db_skill,
|
||||
):
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
result = await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
assert result.id == "litellm_skill_test123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_files_key_without_brackets(self):
|
||||
"""
|
||||
Test that files under plain "files" key also works.
|
||||
"""
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
upload = _make_upload_file("skill/SKILL.md", SKILL_MD_CONTENT)
|
||||
mock_request = MagicMock()
|
||||
|
||||
mock_db_skill = _make_db_skill()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"files": [upload], "display_title": "Test"},
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.handler.LiteLLMSkillsHandler.create_skill",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_db_skill,
|
||||
):
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
result = await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
assert result.id == "litellm_skill_test123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_upload_file_not_in_list(self):
|
||||
"""
|
||||
Test that a single UploadFile (not wrapped in list) is handled.
|
||||
"""
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
upload = _make_upload_file("skill/SKILL.md", SKILL_MD_CONTENT)
|
||||
mock_request = MagicMock()
|
||||
|
||||
mock_db_skill = _make_db_skill()
|
||||
|
||||
# Single UploadFile, not in a list
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"files": upload, "display_title": "Test"},
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.handler.LiteLLMSkillsHandler.create_skill",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_db_skill,
|
||||
):
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
result = await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
assert result.id == "litellm_skill_test123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_files_returns_400(self):
|
||||
"""Test that missing files returns 400."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"display_title": "Test"},
|
||||
):
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "No files provided" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_files_list_returns_400(self):
|
||||
"""Test that empty files list returns 400."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"files": [], "display_title": "Test"},
|
||||
):
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tuple_format_files(self):
|
||||
"""Test that (filename, content) tuple format works."""
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_db_skill = _make_db_skill()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={
|
||||
"files": [("skill/SKILL.md", SKILL_MD_CONTENT)],
|
||||
"display_title": "Test",
|
||||
},
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.handler.LiteLLMSkillsHandler.create_skill",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_db_skill,
|
||||
):
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
result = await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
assert result.id == "litellm_skill_test123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_files(self):
|
||||
"""Test uploading SKILL.md plus additional files."""
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
skill_md = _make_upload_file("skill/SKILL.md", SKILL_MD_CONTENT)
|
||||
helper_py = _make_upload_file("skill/helper.py", b"def helper(): return 42")
|
||||
mock_request = MagicMock()
|
||||
mock_db_skill = _make_db_skill()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={
|
||||
"files": [skill_md, helper_py],
|
||||
"display_title": "Multi File Skill",
|
||||
},
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.handler.LiteLLMSkillsHandler.create_skill",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_db_skill,
|
||||
) as mock_create:
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
result = await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
assert result.id == "litellm_skill_test123"
|
||||
# Verify the create was called with file content
|
||||
call_args = mock_create.call_args
|
||||
data = call_args[1]["data"]
|
||||
assert data.file_content is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_frontmatter_returns_400(self):
|
||||
"""Test that SKILL.md with invalid frontmatter returns 400."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
bad_skill_md = b"""---
|
||||
description: Missing required name field
|
||||
---
|
||||
|
||||
Body content.
|
||||
"""
|
||||
upload = _make_upload_file("skill/SKILL.md", bad_skill_md)
|
||||
mock_request = MagicMock()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"files": [upload], "display_title": "Test"},
|
||||
):
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_display_title_override(self):
|
||||
"""Test that display_title from form overrides frontmatter name."""
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
upload = _make_upload_file("skill/SKILL.md", SKILL_MD_CONTENT)
|
||||
mock_request = MagicMock()
|
||||
mock_db_skill = _make_db_skill(display_title="Custom Title")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={
|
||||
"files": [upload],
|
||||
"display_title": "Custom Title",
|
||||
},
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.handler.LiteLLMSkillsHandler.create_skill",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_db_skill,
|
||||
) as mock_create:
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
call_args = mock_create.call_args
|
||||
data = call_args[1]["data"]
|
||||
assert data.display_title == "Custom Title"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_display_title_falls_back_to_frontmatter_name(self):
|
||||
"""Test that without display_title override, frontmatter name is used."""
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
_handle_litellm_create_skill,
|
||||
)
|
||||
|
||||
upload = _make_upload_file("skill/SKILL.md", SKILL_MD_CONTENT)
|
||||
mock_request = MagicMock()
|
||||
mock_db_skill = _make_db_skill(display_title="test-skill")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.anthropic_endpoints.skills_endpoints.get_form_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"files": [upload]},
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.handler.LiteLLMSkillsHandler.create_skill",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_db_skill,
|
||||
) as mock_create:
|
||||
user = UserAPIKeyAuth(api_key="test-key")
|
||||
await _handle_litellm_create_skill(mock_request, user)
|
||||
|
||||
call_args = mock_create.call_args
|
||||
data = call_args[1]["data"]
|
||||
# Should fall back to frontmatter name "test-skill"
|
||||
assert data.display_title == "test-skill"
|
||||
|
|
@ -591,3 +591,287 @@ class TestMessagesAPIProcessing:
|
|||
"_litellm_code_execution_enabled"
|
||||
)
|
||||
assert "_skill_files" in result.get("litellm_metadata", {})
|
||||
|
||||
|
||||
class TestPreCallProviderRouting:
|
||||
"""
|
||||
Tests that the pre-call hook routes to the correct strategy based on provider.
|
||||
|
||||
Non-native providers (OpenAI, Azure, Bedrock, etc.):
|
||||
- Skill content injected into system prompt
|
||||
- No skill tools added
|
||||
- container removed
|
||||
|
||||
Native providers (Anthropic, azure_ai, databricks):
|
||||
- Skill converted to Anthropic-style tool
|
||||
- Skill content injected into system prompt
|
||||
- Code execution tool added (if skill has files)
|
||||
- container removed
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_gets_system_prompt_only(self):
|
||||
"""OpenAI should get system prompt injection, no tools from skill."""
|
||||
hook = SkillsInjectionHook()
|
||||
skill = _make_skill()
|
||||
|
||||
with patch.object(
|
||||
hook, "_fetch_skill_from_db", new_callable=AsyncMock, return_value=skill
|
||||
):
|
||||
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": "litellm_skill_test1"}]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
# System prompt should be injected
|
||||
system_msgs = [
|
||||
m for m in result["messages"] if m.get("role") == "system"
|
||||
]
|
||||
assert len(system_msgs) == 1
|
||||
assert "Test Skill" in system_msgs[0]["content"]
|
||||
|
||||
# No tools should be added for OpenAI
|
||||
assert "tools" not in result
|
||||
|
||||
# container should be removed
|
||||
assert "container" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_gets_system_prompt_only(self):
|
||||
"""Azure should get system prompt injection, no tools from skill."""
|
||||
hook = SkillsInjectionHook()
|
||||
skill = _make_skill()
|
||||
|
||||
with patch.object(
|
||||
hook, "_fetch_skill_from_db", new_callable=AsyncMock, return_value=skill
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.skill_applicator.get_provider_from_model",
|
||||
return_value="azure",
|
||||
):
|
||||
data = {
|
||||
"model": "azure/gpt-4o",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"container": {
|
||||
"skills": [{"skill_id": "litellm_skill_test1"}]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
system_msgs = [
|
||||
m for m in result["messages"] if m.get("role") == "system"
|
||||
]
|
||||
assert len(system_msgs) == 1
|
||||
assert "tools" not in result
|
||||
assert "container" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_gets_system_prompt_only(self):
|
||||
"""Bedrock should get system prompt injection, no tools from skill."""
|
||||
hook = SkillsInjectionHook()
|
||||
skill = _make_skill()
|
||||
|
||||
with patch.object(
|
||||
hook, "_fetch_skill_from_db", new_callable=AsyncMock, return_value=skill
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.skill_applicator.get_provider_from_model",
|
||||
return_value="bedrock",
|
||||
):
|
||||
data = {
|
||||
"model": "bedrock/anthropic.claude-v2",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"container": {
|
||||
"skills": [{"skill_id": "litellm_skill_test1"}]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
system_msgs = [
|
||||
m for m in result["messages"] if m.get("role") == "system"
|
||||
]
|
||||
assert len(system_msgs) == 1
|
||||
assert "tools" not in result
|
||||
assert "container" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_gets_tools_and_system_prompt(self):
|
||||
"""Anthropic should get tool conversion + system prompt injection."""
|
||||
hook = SkillsInjectionHook()
|
||||
skill = _make_skill()
|
||||
|
||||
with patch.object(
|
||||
hook, "_fetch_skill_from_db", new_callable=AsyncMock, return_value=skill
|
||||
):
|
||||
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": "litellm_skill_test1"}]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
# Tools should be added (skill tool)
|
||||
assert "tools" in result
|
||||
assert len(result["tools"]) >= 1
|
||||
|
||||
# container should be removed
|
||||
assert "container" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_preserves_existing_system_message(self):
|
||||
"""OpenAI skill injection should append to existing system message."""
|
||||
hook = SkillsInjectionHook()
|
||||
skill = _make_skill()
|
||||
|
||||
with patch.object(
|
||||
hook, "_fetch_skill_from_db", new_callable=AsyncMock, return_value=skill
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.litellm_proxy.skills.skill_applicator.get_provider_from_model",
|
||||
return_value="openai",
|
||||
):
|
||||
data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
"container": {
|
||||
"skills": [{"skill_id": "litellm_skill_test1"}]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
system_msg = result["messages"][0]
|
||||
assert system_msg["role"] == "system"
|
||||
# Original content preserved
|
||||
assert system_msg["content"].startswith("You are a helpful assistant.")
|
||||
# Skill content appended
|
||||
assert "Test Skill" in system_msg["content"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_no_tools_even_with_code_files(self):
|
||||
"""OpenAI should NOT get tools even if skill has Python files."""
|
||||
hook = SkillsInjectionHook()
|
||||
skill = _make_skill_with_code()
|
||||
|
||||
with patch.object(
|
||||
hook, "_fetch_skill_from_db", new_callable=AsyncMock, return_value=skill
|
||||
):
|
||||
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": "litellm_skill_code1"}]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
# System prompt should have skill content
|
||||
system_msgs = [
|
||||
m for m in result["messages"] if m.get("role") == "system"
|
||||
]
|
||||
assert len(system_msgs) == 1
|
||||
|
||||
# No tools — OpenAI just gets the prompt
|
||||
assert "tools" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_skills_openai(self):
|
||||
"""Multiple skills should all be injected into system prompt for OpenAI."""
|
||||
hook = SkillsInjectionHook()
|
||||
skill1 = _make_skill(
|
||||
skill_id="litellm_skill_a",
|
||||
display_title="Skill Alpha",
|
||||
instructions="Alpha instructions.",
|
||||
)
|
||||
skill2 = _make_skill(
|
||||
skill_id="litellm_skill_b",
|
||||
display_title="Skill Beta",
|
||||
instructions="Beta instructions.",
|
||||
)
|
||||
|
||||
async def _fetch(skill_id):
|
||||
return {"litellm_skill_a": skill1, "litellm_skill_b": skill2}.get(skill_id)
|
||||
|
||||
with patch.object(hook, "_fetch_skill_from_db", side_effect=_fetch):
|
||||
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": "litellm_skill_a"},
|
||||
{"skill_id": "litellm_skill_b"},
|
||||
]
|
||||
},
|
||||
}
|
||||
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
system_content = result["messages"][0]["content"]
|
||||
assert "Skill Alpha" in system_content
|
||||
assert "Skill Beta" in system_content
|
||||
assert "tools" not in result
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue