chore(proxy): scope skills and container resources

This commit is contained in:
user
2026-04-30 18:23:58 -07:00
parent 05e6402bdb
commit ad9aa43e86
10 changed files with 797 additions and 55 deletions
+50 -10
View File
@@ -9,7 +9,13 @@ import uuid
from typing import Any, Dict, List, Optional
from litellm._logging import verbose_logger
from litellm.proxy._types import LiteLLM_SkillsTable, NewSkillRequest
from litellm.proxy._types import LiteLLM_SkillsTable, NewSkillRequest, UserAPIKeyAuth
from litellm.proxy.common_utils.resource_ownership import (
get_primary_resource_owner_scope,
get_resource_owner_scopes,
is_proxy_admin,
user_can_access_resource_owner,
)
def _prisma_skill_to_litellm(prisma_skill) -> LiteLLM_SkillsTable:
@@ -58,6 +64,7 @@ class LiteLLMSkillsHandler:
async def create_skill(
data: NewSkillRequest,
user_id: Optional[str] = None,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
) -> LiteLLM_SkillsTable:
"""
Create a new skill in the LiteLLM database.
@@ -72,6 +79,7 @@ class LiteLLMSkillsHandler:
prisma_client = await LiteLLMSkillsHandler._get_prisma_client()
skill_id = f"litellm_skill_{uuid.uuid4()}"
owner = get_primary_resource_owner_scope(user_api_key_dict) or user_id
skill_data: Dict[str, Any] = {
"skill_id": skill_id,
@@ -79,8 +87,8 @@ class LiteLLMSkillsHandler:
"description": data.description,
"instructions": data.instructions,
"source": "custom",
"created_by": user_id,
"updated_by": user_id,
"created_by": owner,
"updated_by": owner,
}
# Handle metadata
@@ -111,6 +119,7 @@ class LiteLLMSkillsHandler:
async def list_skills(
limit: int = 20,
offset: int = 0,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
) -> List[LiteLLM_SkillsTable]:
"""
List skills from the LiteLLM database.
@@ -128,16 +137,28 @@ class LiteLLMSkillsHandler:
f"LiteLLMSkillsHandler: Listing skills with limit={limit}, offset={offset}"
)
find_many_kwargs: Dict[str, Any] = {
"take": limit,
"skip": offset,
"order": {"created_at": "desc"},
}
if user_api_key_dict is not None and not is_proxy_admin(user_api_key_dict):
owner_scopes = get_resource_owner_scopes(user_api_key_dict)
if not owner_scopes:
return []
find_many_kwargs["where"] = {"created_by": {"in": owner_scopes}}
skills = await prisma_client.db.litellm_skillstable.find_many(
take=limit,
skip=offset,
order={"created_at": "desc"},
**find_many_kwargs
)
return [_prisma_skill_to_litellm(s) for s in skills]
@staticmethod
async def get_skill(skill_id: str) -> LiteLLM_SkillsTable:
async def get_skill(
skill_id: str,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
) -> LiteLLM_SkillsTable:
"""
Get a skill by ID from the LiteLLM database.
@@ -161,10 +182,18 @@ class LiteLLMSkillsHandler:
if skill is None:
raise ValueError(f"Skill not found: {skill_id}")
if not user_can_access_resource_owner(
getattr(skill, "created_by", None), user_api_key_dict
):
raise ValueError(f"Skill not found: {skill_id}")
return _prisma_skill_to_litellm(skill)
@staticmethod
async def delete_skill(skill_id: str) -> Dict[str, str]:
async def delete_skill(
skill_id: str,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
) -> Dict[str, str]:
"""
Delete a skill by ID from the LiteLLM database.
@@ -189,13 +218,21 @@ class LiteLLMSkillsHandler:
if skill is None:
raise ValueError(f"Skill not found: {skill_id}")
if not user_can_access_resource_owner(
getattr(skill, "created_by", None), user_api_key_dict
):
raise ValueError(f"Skill not found: {skill_id}")
# Delete the skill
await prisma_client.db.litellm_skillstable.delete(where={"skill_id": skill_id})
return {"id": skill_id, "type": "skill_deleted"}
@staticmethod
async def fetch_skill_from_db(skill_id: str) -> Optional[LiteLLM_SkillsTable]:
async def fetch_skill_from_db(
skill_id: str,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
) -> Optional[LiteLLM_SkillsTable]:
"""
Fetch a skill from the database (used by skills injection hook).
@@ -209,7 +246,10 @@ class LiteLLMSkillsHandler:
LiteLLM_SkillsTable or None if not found
"""
try:
return await LiteLLMSkillsHandler.get_skill(skill_id)
return await LiteLLMSkillsHandler.get_skill(
skill_id,
user_api_key_dict=user_api_key_dict,
)
except ValueError:
return None
except Exception as e:
@@ -18,6 +18,7 @@ from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
class LiteLLMSkillsTransformationHandler:
@@ -44,6 +45,7 @@ class LiteLLMSkillsTransformationHandler:
file_type: Optional[str] = None,
metadata: Optional[Dict[str, Any]] = None,
user_id: Optional[str] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
_is_async: bool = False,
logging_obj: Optional["LiteLLMLoggingObj"] = None,
litellm_call_id: Optional[str] = None,
@@ -99,6 +101,7 @@ class LiteLLMSkillsTransformationHandler:
file_type=file_type,
metadata=metadata,
user_id=user_id,
user_api_key_dict=user_api_key_dict,
)
import asyncio
@@ -113,6 +116,7 @@ class LiteLLMSkillsTransformationHandler:
file_type=file_type,
metadata=metadata,
user_id=user_id,
user_api_key_dict=user_api_key_dict,
)
)
@@ -126,6 +130,7 @@ class LiteLLMSkillsTransformationHandler:
file_type: Optional[str] = None,
metadata: Optional[Dict[str, Any]] = None,
user_id: Optional[str] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
) -> Skill:
"""Async implementation of create_skill."""
# Lazy import to avoid SDK dependency on proxy
@@ -145,6 +150,7 @@ class LiteLLMSkillsTransformationHandler:
db_skill = await LiteLLMSkillsHandler.create_skill(
data=skill_request,
user_id=user_id,
user_api_key_dict=user_api_key_dict,
)
return self._db_skill_to_response(db_skill)
@@ -156,6 +162,7 @@ class LiteLLMSkillsTransformationHandler:
_is_async: bool = False,
logging_obj: Optional["LiteLLMLoggingObj"] = None,
litellm_call_id: Optional[str] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
**kwargs,
) -> Union[ListSkillsResponse, Coroutine[Any, Any, ListSkillsResponse]]:
"""
@@ -182,18 +189,27 @@ class LiteLLMSkillsTransformationHandler:
)
if _is_async:
return self._async_list_skills(limit=limit, offset=offset)
return self._async_list_skills(
limit=limit,
offset=offset,
user_api_key_dict=user_api_key_dict,
)
import asyncio
return asyncio.get_event_loop().run_until_complete(
self._async_list_skills(limit=limit, offset=offset)
self._async_list_skills(
limit=limit,
offset=offset,
user_api_key_dict=user_api_key_dict,
)
)
async def _async_list_skills(
self,
limit: int = 20,
offset: int = 0,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
) -> ListSkillsResponse:
"""Async implementation of list_skills."""
# Lazy import to avoid SDK dependency on proxy
@@ -202,6 +218,7 @@ class LiteLLMSkillsTransformationHandler:
db_skills = await LiteLLMSkillsHandler.list_skills(
limit=limit,
offset=offset,
user_api_key_dict=user_api_key_dict,
)
skills = [self._db_skill_to_response(s) for s in db_skills]
@@ -217,6 +234,7 @@ class LiteLLMSkillsTransformationHandler:
_is_async: bool = False,
logging_obj: Optional["LiteLLMLoggingObj"] = None,
litellm_call_id: Optional[str] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
**kwargs,
) -> Union[Skill, Coroutine[Any, Any, Skill]]:
"""
@@ -242,20 +260,33 @@ class LiteLLMSkillsTransformationHandler:
)
if _is_async:
return self._async_get_skill(skill_id=skill_id)
return self._async_get_skill(
skill_id=skill_id,
user_api_key_dict=user_api_key_dict,
)
import asyncio
return asyncio.get_event_loop().run_until_complete(
self._async_get_skill(skill_id=skill_id)
self._async_get_skill(
skill_id=skill_id,
user_api_key_dict=user_api_key_dict,
)
)
async def _async_get_skill(self, skill_id: str) -> Skill:
async def _async_get_skill(
self,
skill_id: str,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
) -> Skill:
"""Async implementation of get_skill."""
# Lazy import to avoid SDK dependency on proxy
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
db_skill = await LiteLLMSkillsHandler.get_skill(skill_id=skill_id)
db_skill = await LiteLLMSkillsHandler.get_skill(
skill_id=skill_id,
user_api_key_dict=user_api_key_dict,
)
return self._db_skill_to_response(db_skill)
def delete_skill_handler(
@@ -264,6 +295,7 @@ class LiteLLMSkillsTransformationHandler:
_is_async: bool = False,
logging_obj: Optional["LiteLLMLoggingObj"] = None,
litellm_call_id: Optional[str] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
**kwargs,
) -> Union[DeleteSkillResponse, Coroutine[Any, Any, DeleteSkillResponse]]:
"""
@@ -289,20 +321,33 @@ class LiteLLMSkillsTransformationHandler:
)
if _is_async:
return self._async_delete_skill(skill_id=skill_id)
return self._async_delete_skill(
skill_id=skill_id,
user_api_key_dict=user_api_key_dict,
)
import asyncio
return asyncio.get_event_loop().run_until_complete(
self._async_delete_skill(skill_id=skill_id)
self._async_delete_skill(
skill_id=skill_id,
user_api_key_dict=user_api_key_dict,
)
)
async def _async_delete_skill(self, skill_id: str) -> DeleteSkillResponse:
async def _async_delete_skill(
self,
skill_id: str,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
) -> DeleteSkillResponse:
"""Async implementation of delete_skill."""
# Lazy import to avoid SDK dependency on proxy
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
result = await LiteLLMSkillsHandler.delete_skill(skill_id=skill_id)
result = await LiteLLMSkillsHandler.delete_skill(
skill_id=skill_id,
user_api_key_dict=user_api_key_dict,
)
return DeleteSkillResponse(
id=result["id"],
type=result.get("type", "skill_deleted"),
@@ -0,0 +1,75 @@
from typing import List, Optional
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
def is_proxy_admin(user_api_key_dict: Optional[UserAPIKeyAuth]) -> bool:
if user_api_key_dict is None:
return False
return (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
def get_resource_owner_scopes(
user_api_key_dict: Optional[UserAPIKeyAuth],
) -> List[str]:
"""
Return ownership scopes that may access a user-created proxy resource.
Raw user_id is included for rows created before scope prefixes existed.
Prefixes avoid collisions when falling back to team/org/key ownership for
keys that do not have a user_id.
"""
if user_api_key_dict is None:
return []
scopes: List[str] = []
def _add(scope: Optional[str]) -> None:
if scope and scope not in scopes:
scopes.append(scope)
if user_api_key_dict.user_id:
_add(user_api_key_dict.user_id)
_add(f"user:{user_api_key_dict.user_id}")
if user_api_key_dict.team_id:
_add(f"team:{user_api_key_dict.team_id}")
if user_api_key_dict.org_id:
_add(f"org:{user_api_key_dict.org_id}")
if user_api_key_dict.api_key:
_add(f"key:{user_api_key_dict.api_key}")
return scopes
def get_primary_resource_owner_scope(
user_api_key_dict: Optional[UserAPIKeyAuth],
) -> Optional[str]:
if user_api_key_dict is None:
return None
if user_api_key_dict.user_id:
return user_api_key_dict.user_id
if user_api_key_dict.team_id:
return f"team:{user_api_key_dict.team_id}"
if user_api_key_dict.org_id:
return f"org:{user_api_key_dict.org_id}"
if user_api_key_dict.api_key:
return f"key:{user_api_key_dict.api_key}"
return None
def user_can_access_resource_owner(
owner: Optional[str],
user_api_key_dict: Optional[UserAPIKeyAuth],
) -> bool:
if user_api_key_dict is None:
return True
if is_proxy_admin(user_api_key_dict):
return True
if owner is None:
return False
return owner in get_resource_owner_scopes(user_api_key_dict)
+27 -2
View File
@@ -14,6 +14,11 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_headers,
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.container_endpoints.ownership import (
assert_user_can_access_container,
filter_container_list_response,
record_container_owner,
)
router = APIRouter()
@@ -98,7 +103,7 @@ async def create_container(
# Process request using ProxyBaseLLMRequestProcessing
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@@ -123,6 +128,11 @@ async def create_container(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
return await record_container_owner(
response=response,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
@router.get(
@@ -191,7 +201,7 @@ async def list_containers(
# Process request using ProxyBaseLLMRequestProcessing
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@@ -216,6 +226,11 @@ async def list_containers(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
return await filter_container_list_response(
response=response,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
@router.get(
@@ -280,6 +295,11 @@ async def retrieve_container(
)
# Add custom_llm_provider to data
_, custom_llm_provider = await assert_user_can_access_container(
container_id=container_id,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
data["custom_llm_provider"] = custom_llm_provider
# Process request using ProxyBaseLLMRequestProcessing
@@ -374,6 +394,11 @@ async def delete_container(
)
# Add custom_llm_provider to data
_, custom_llm_provider = await assert_user_can_access_container(
container_id=container_id,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
data["custom_llm_provider"] = custom_llm_provider
# Process request using ProxyBaseLLMRequestProcessing
@@ -19,7 +19,9 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_headers,
get_custom_llm_provider_from_request_query,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.proxy.container_endpoints.ownership import (
assert_user_can_access_container,
)
def _load_endpoints_config() -> Dict:
@@ -176,14 +178,11 @@ async def _process_binary_request(
# Build litellm_params - credentials are resolved by provider config from env
litellm_params = GenericLiteLLMParams()
# Decode container ID and extract provider info
decoded = ResponsesAPIRequestUtils._decode_container_id(container_id)
original_container_id = decoded.get("response_id", container_id)
# If container ID has encoded provider info and user didn't explicitly set provider, use it
decoded_provider = decoded.get("custom_llm_provider")
if decoded_provider and custom_llm_provider == "openai":
custom_llm_provider = decoded_provider
original_container_id, custom_llm_provider = await assert_user_can_access_container(
container_id=container_id,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
# Get the provider config
container_provider_config = _get_container_provider_config(custom_llm_provider)
@@ -284,16 +283,13 @@ async def _process_multipart_upload_request(
or "openai"
)
# Decode container ID and extract provider info
decoded = ResponsesAPIRequestUtils._decode_container_id(container_id)
original_container_id = decoded.get("response_id", container_id)
_, custom_llm_provider = await assert_user_can_access_container(
container_id=container_id,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
# If container ID has encoded provider info and user didn't explicitly set provider, use it
decoded_provider = decoded.get("custom_llm_provider")
if decoded_provider and custom_llm_provider == "openai":
custom_llm_provider = decoded_provider
data["container_id"] = original_container_id # Use decoded original ID
data["container_id"] = container_id
data["custom_llm_provider"] = custom_llm_provider
processor = ProxyBaseLLMRequestProcessing(data=data)
@@ -361,18 +357,11 @@ async def _process_request(
# Decode container_id if present in path_params
if "container_id" in path_params:
decoded = ResponsesAPIRequestUtils._decode_container_id(
path_params["container_id"]
_, custom_llm_provider = await assert_user_can_access_container(
container_id=path_params["container_id"],
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
original_container_id = decoded.get("response_id", path_params["container_id"])
# If container ID has encoded provider info and user didn't explicitly set provider, use it
decoded_provider = decoded.get("custom_llm_provider")
if decoded_provider and custom_llm_provider == "openai":
custom_llm_provider = decoded_provider
# Update path_params with decoded original ID
data["container_id"] = original_container_id
data["custom_llm_provider"] = custom_llm_provider
@@ -0,0 +1,256 @@
from typing import Any, Dict, List, Optional, Set, Tuple
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.resource_ownership import (
get_primary_resource_owner_scope,
get_resource_owner_scopes,
is_proxy_admin,
user_can_access_resource_owner,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
CONTAINER_OBJECT_PURPOSE = "container"
def _container_model_object_id(
original_container_id: str,
custom_llm_provider: str,
) -> str:
return f"{CONTAINER_OBJECT_PURPOSE}:{custom_llm_provider}:{original_container_id}"
def decode_container_id_for_ownership(
container_id: str,
custom_llm_provider: str,
) -> Tuple[str, str]:
decoded = ResponsesAPIRequestUtils._decode_container_id(container_id)
original_container_id = decoded.get("response_id", container_id)
decoded_provider = decoded.get("custom_llm_provider")
if decoded_provider and custom_llm_provider == "openai":
custom_llm_provider = decoded_provider
return original_container_id, custom_llm_provider
def _get_response_id(response: Any) -> Optional[str]:
if response is None:
return None
if isinstance(response, dict):
value = response.get("id")
else:
value = getattr(response, "id", None)
return value if isinstance(value, str) else None
def _dump_response(response: Any) -> Dict[str, Any]:
if isinstance(response, dict):
return response
if hasattr(response, "model_dump"):
return response.model_dump()
if hasattr(response, "dict"):
return response.dict()
return {"id": _get_response_id(response)}
async def _get_prisma_client():
from litellm.proxy.proxy_server import prisma_client
return prisma_client
async def record_container_owner(
response: Any,
user_api_key_dict: UserAPIKeyAuth,
custom_llm_provider: str,
) -> Any:
container_id = _get_response_id(response)
owner = get_primary_resource_owner_scope(user_api_key_dict)
prisma_client = await _get_prisma_client()
if is_proxy_admin(user_api_key_dict) and (
container_id is None or owner is None or prisma_client is None
):
return response
if container_id is None or owner is None or prisma_client is None:
raise HTTPException(status_code=500, detail="Unable to track container")
original_container_id, resolved_provider = decode_container_id_for_ownership(
container_id,
custom_llm_provider,
)
model_object_id = _container_model_object_id(
original_container_id,
resolved_provider,
)
file_object = _dump_response(response)
file_object["custom_llm_provider"] = resolved_provider
file_object["provider_container_id"] = original_container_id
try:
existing = await prisma_client.db.litellm_managedobjecttable.find_unique(
where={"model_object_id": model_object_id}
)
if existing is not None:
if getattr(existing, "file_purpose", None) != CONTAINER_OBJECT_PURPOSE:
raise HTTPException(status_code=500, detail="Unable to track container")
if not user_can_access_resource_owner(
getattr(existing, "created_by", None), user_api_key_dict
):
raise HTTPException(status_code=403, detail="Forbidden")
await prisma_client.db.litellm_managedobjecttable.update(
where={"model_object_id": model_object_id},
data={
"unified_object_id": container_id,
"file_object": file_object,
"updated_by": owner,
},
)
else:
await prisma_client.db.litellm_managedobjecttable.create(
data={
"unified_object_id": container_id,
"model_object_id": model_object_id,
"file_object": file_object,
"file_purpose": CONTAINER_OBJECT_PURPOSE,
"created_by": owner,
"updated_by": owner,
}
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.warning(
"Failed to record container ownership for container_id=%s: %s",
model_object_id,
e,
)
raise HTTPException(status_code=500, detail="Unable to track container")
return response
async def _get_container_owner(
original_container_id: str,
custom_llm_provider: str,
) -> Optional[str]:
prisma_client = await _get_prisma_client()
if prisma_client is None:
return None
row = await prisma_client.db.litellm_managedobjecttable.find_first(
where={
"model_object_id": _container_model_object_id(
original_container_id,
custom_llm_provider,
),
"file_purpose": CONTAINER_OBJECT_PURPOSE,
}
)
return getattr(row, "created_by", None) if row is not None else None
async def assert_user_can_access_container(
container_id: str,
user_api_key_dict: UserAPIKeyAuth,
custom_llm_provider: str,
) -> Tuple[str, str]:
original_container_id, resolved_provider = decode_container_id_for_ownership(
container_id,
custom_llm_provider,
)
if is_proxy_admin(user_api_key_dict):
return original_container_id, resolved_provider
owner = await _get_container_owner(original_container_id, resolved_provider)
if not user_can_access_resource_owner(owner, user_api_key_dict):
raise HTTPException(status_code=403, detail="Forbidden")
return original_container_id, resolved_provider
def _get_container_list_data(response: Any) -> Optional[List[Any]]:
if response is None:
return None
if isinstance(response, dict):
data = response.get("data")
else:
data = getattr(response, "data", None)
return data if isinstance(data, list) else None
def _set_container_list_data(response: Any, data: List[Any]) -> Any:
if isinstance(response, dict):
response["data"] = data
if data:
response["first_id"] = _get_response_id(data[0])
response["last_id"] = _get_response_id(data[-1])
else:
response["first_id"] = None
response["last_id"] = None
return response
response.data = data
response.first_id = _get_response_id(data[0]) if data else None
response.last_id = _get_response_id(data[-1]) if data else None
return response
async def _get_allowed_container_ids(
user_api_key_dict: UserAPIKeyAuth,
custom_llm_provider: str,
) -> Set[str]:
prisma_client = await _get_prisma_client()
if prisma_client is None:
return set()
owner_scopes = get_resource_owner_scopes(user_api_key_dict)
if not owner_scopes:
return set()
rows = await prisma_client.db.litellm_managedobjecttable.find_many(
where={
"file_purpose": CONTAINER_OBJECT_PURPOSE,
"created_by": {"in": owner_scopes},
}
)
return {
row.model_object_id
for row in rows
if getattr(row, "model_object_id", None) is not None
}
async def filter_container_list_response(
response: Any,
user_api_key_dict: UserAPIKeyAuth,
custom_llm_provider: str,
) -> Any:
if is_proxy_admin(user_api_key_dict):
return response
data = _get_container_list_data(response)
if data is None:
return response
allowed_container_ids = await _get_allowed_container_ids(
user_api_key_dict,
custom_llm_provider,
)
filtered: List[Any] = []
for item in data:
container_id = _get_response_id(item)
if container_id is None:
continue
original_container_id, resolved_provider = decode_container_id_for_ownership(
container_id,
custom_llm_provider,
)
if (
_container_model_object_id(original_container_id, resolved_provider)
in allowed_container_ids
):
filtered.append(item)
return _set_container_list_data(response, filtered)
+11 -3
View File
@@ -110,7 +110,10 @@ class SkillsInjectionHook(CustomLogger):
skill_id = skill.get("skill_id", "")
if skill_id.startswith("litellm_"):
# Fetch from LiteLLM DB
db_skill = await self._fetch_skill_from_db(skill_id)
db_skill = await self._fetch_skill_from_db(
skill_id,
user_api_key_dict=user_api_key_dict,
)
if db_skill:
litellm_skills.append(db_skill)
else:
@@ -276,7 +279,9 @@ class SkillsInjectionHook(CustomLogger):
return data
async def _fetch_skill_from_db(
self, skill_id: str
self,
skill_id: str,
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[LiteLLM_SkillsTable]:
"""
Fetch a skill from the LiteLLM database.
@@ -290,7 +295,10 @@ class SkillsInjectionHook(CustomLogger):
try:
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
return await LiteLLMSkillsHandler.fetch_skill_from_db(skill_id)
return await LiteLLMSkillsHandler.fetch_skill_from_db(
skill_id,
user_api_key_dict=user_api_key_dict,
)
except Exception as e:
verbose_proxy_logger.warning(
f"SkillsInjectionHook: Error fetching skill {skill_id}: {e}"
+28 -1
View File
@@ -34,6 +34,29 @@ DEFAULT_ANTHROPIC_API_BASE = "https://api.anthropic.com/v1"
_litellm_skills_handler = None
def _get_user_api_key_auth_from_kwargs(kwargs: Dict[str, Any]) -> Optional[Any]:
for metadata_key in ("metadata", "litellm_metadata"):
metadata = kwargs.get(metadata_key)
if isinstance(metadata, dict) and metadata.get("user_api_key_auth") is not None:
return metadata["user_api_key_auth"]
return None
def _get_skill_request_metadata(
kwargs: Dict[str, Any],
extra_body: Optional[Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
if extra_body and isinstance(extra_body.get("metadata"), dict):
return extra_body["metadata"]
metadata = kwargs.get("metadata")
if isinstance(metadata, dict) and isinstance(
metadata.get("requester_metadata"), dict
):
return metadata["requester_metadata"]
return None
def _get_litellm_skills_handler():
"""Lazy initialization of LiteLLM skills handler to avoid import overhead."""
global _litellm_skills_handler
@@ -165,8 +188,9 @@ def create_skill(
return _get_litellm_skills_handler().create_skill_handler(
display_title=display_title,
files=files,
metadata=extra_body.get("metadata") if extra_body else None,
metadata=_get_skill_request_metadata(kwargs, extra_body),
user_id=kwargs.get("user_id"),
user_api_key_dict=_get_user_api_key_auth_from_kwargs(kwargs),
_is_async=_is_async,
logging_obj=litellm_logging_obj,
litellm_call_id=litellm_call_id,
@@ -348,6 +372,7 @@ def list_skills(
return _get_litellm_skills_handler().list_skills_handler(
limit=limit or 20,
offset=0,
user_api_key_dict=_get_user_api_key_auth_from_kwargs(kwargs),
_is_async=_is_async,
logging_obj=litellm_logging_obj,
litellm_call_id=litellm_call_id,
@@ -523,6 +548,7 @@ def get_skill(
if custom_llm_provider == LlmProviders.LITELLM_PROXY.value:
return _get_litellm_skills_handler().get_skill_handler(
skill_id=skill_id,
user_api_key_dict=_get_user_api_key_auth_from_kwargs(kwargs),
_is_async=_is_async,
logging_obj=litellm_logging_obj,
litellm_call_id=litellm_call_id,
@@ -690,6 +716,7 @@ def delete_skill(
if custom_llm_provider == LlmProviders.LITELLM_PROXY.value:
return _get_litellm_skills_handler().delete_skill_handler(
skill_id=skill_id,
user_api_key_dict=_get_user_api_key_auth_from_kwargs(kwargs),
_is_async=_is_async,
logging_obj=litellm_logging_obj,
litellm_call_id=litellm_call_id,
@@ -0,0 +1,159 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.container_endpoints import ownership
from litellm.types.containers.main import ContainerListResponse, ContainerObject
def _container(container_id: str) -> ContainerObject:
return ContainerObject(
id=container_id,
object="container",
created_at=1,
status="active",
)
@pytest.mark.asyncio
async def test_should_record_container_owner_with_original_provider_id(monkeypatch):
table = AsyncMock()
table.find_unique.return_value = None
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(user_id="user-1")
response = _container("cntr_provider")
await ownership.record_container_owner(
response=response,
user_api_key_dict=auth,
custom_llm_provider="openai",
)
table.create.assert_awaited_once()
data = table.create.await_args.kwargs["data"]
assert data["model_object_id"] == "container:openai:cntr_provider"
assert data["file_purpose"] == ownership.CONTAINER_OBJECT_PURPOSE
assert data["created_by"] == "user-1"
@pytest.mark.asyncio
async def test_should_record_team_owner_for_keys_without_user_id(monkeypatch):
table = AsyncMock()
table.find_unique.return_value = None
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(team_id="team-1")
await ownership.record_container_owner(
response=_container("cntr_provider"),
user_api_key_dict=auth,
custom_llm_provider="openai",
)
data = table.create.await_args.kwargs["data"]
assert data["created_by"] == "team:team-1"
assert data["updated_by"] == "team:team-1"
@pytest.mark.asyncio
async def test_should_deny_container_access_for_different_owner(monkeypatch):
table = AsyncMock()
table.find_first.return_value = SimpleNamespace(created_by="user-2")
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(user_id="user-1")
with pytest.raises(HTTPException) as exc:
await ownership.assert_user_can_access_container(
container_id="cntr_provider",
user_api_key_dict=auth,
custom_llm_provider="openai",
)
assert exc.value.status_code == 403
@pytest.mark.asyncio
async def test_should_not_reassign_existing_container_to_different_owner(monkeypatch):
table = AsyncMock()
table.find_unique.return_value = SimpleNamespace(
file_purpose=ownership.CONTAINER_OBJECT_PURPOSE,
created_by="user-2",
)
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(user_id="user-1")
with pytest.raises(HTTPException) as exc:
await ownership.record_container_owner(
response=_container("cntr_existing"),
user_api_key_dict=auth,
custom_llm_provider="openai",
)
assert exc.value.status_code == 403
table.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_should_filter_container_list_to_owned_records(monkeypatch):
table = AsyncMock()
table.find_many.return_value = [
SimpleNamespace(model_object_id="container:openai:cntr_owned"),
]
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(user_id="user-1")
response = ContainerListResponse(
object="list",
data=[_container("cntr_owned"), _container("cntr_other")],
has_more=False,
)
filtered = await ownership.filter_container_list_response(
response=response,
user_api_key_dict=auth,
custom_llm_provider="openai",
)
assert [item.id for item in filtered.data] == ["cntr_owned"]
assert filtered.first_id == "cntr_owned"
assert filtered.last_id == "cntr_owned"
where = table.find_many.await_args.kwargs["where"]
assert where["file_purpose"] == ownership.CONTAINER_OBJECT_PURPOSE
assert where["created_by"]["in"] == ["user-1", "user:user-1"]
@@ -0,0 +1,118 @@
from unittest.mock import AsyncMock
import pytest
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
from litellm.proxy._types import LiteLLM_SkillsTable, NewSkillRequest, UserAPIKeyAuth
def _skill(skill_id: str, created_by: str | None) -> LiteLLM_SkillsTable:
return LiteLLM_SkillsTable(
skill_id=skill_id,
display_title="skill",
created_by=created_by,
)
@pytest.mark.asyncio
async def test_should_store_team_owner_for_keys_without_user_id(monkeypatch):
table = AsyncMock()
table.create.side_effect = lambda data: _skill(data["skill_id"], data["created_by"])
prisma_client = type(
"Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()}
)()
monkeypatch.setattr(
LiteLLMSkillsHandler,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(team_id="team-1")
skill = await LiteLLMSkillsHandler.create_skill(
data=NewSkillRequest(display_title="skill"),
user_api_key_dict=auth,
)
assert skill.created_by == "team:team-1"
assert table.create.await_args.kwargs["data"]["updated_by"] == "team:team-1"
@pytest.mark.asyncio
async def test_should_filter_list_skills_to_authenticated_owner_scopes(monkeypatch):
table = AsyncMock()
table.find_many.return_value = [_skill("litellm_skill_owner", "user-1")]
prisma_client = type(
"Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()}
)()
monkeypatch.setattr(
LiteLLMSkillsHandler,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(user_id="user-1", team_id="team-1")
skills = await LiteLLMSkillsHandler.list_skills(user_api_key_dict=auth)
assert [skill.skill_id for skill in skills] == ["litellm_skill_owner"]
table.find_many.assert_awaited_once()
where = table.find_many.await_args.kwargs["where"]
assert where["created_by"]["in"] == [
"user-1",
"user:user-1",
"team:team-1",
]
@pytest.mark.asyncio
async def test_should_hide_skill_from_different_owner(monkeypatch):
table = AsyncMock()
table.find_unique.return_value = _skill("litellm_skill_other", "user-2")
prisma_client = type(
"Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()}
)()
monkeypatch.setattr(
LiteLLMSkillsHandler,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(user_id="user-1")
with pytest.raises(ValueError, match="Skill not found"):
await LiteLLMSkillsHandler.get_skill(
"litellm_skill_other",
user_api_key_dict=auth,
)
@pytest.mark.asyncio
async def test_should_scope_skill_injection_fetch_to_authenticated_user(monkeypatch):
from litellm.proxy.hooks.litellm_skills.main import SkillsInjectionHook
fetch = AsyncMock(return_value=None)
monkeypatch.setattr(LiteLLMSkillsHandler, "fetch_skill_from_db", fetch)
auth = UserAPIKeyAuth(user_id="user-1")
hook = SkillsInjectionHook()
data = {
"container": {
"skills": [
{"skill_id": "litellm_skill_other"},
]
}
}
response = await hook.async_pre_call_hook(
user_api_key_dict=auth,
cache=AsyncMock(),
data=data,
call_type="completion",
)
assert response == data
fetch.assert_awaited_once_with(
"litellm_skill_other",
user_api_key_dict=auth,
)