test(proxy): cover skill ownership propagation

This commit is contained in:
user 2026-04-30 20:06:44 -07:00
parent a02aacf1b5
commit 2ecc79b9e9

View file

@ -1,10 +1,20 @@
from unittest.mock import AsyncMock
from unittest.mock import AsyncMock, Mock
import pytest
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
from litellm.llms.litellm_proxy.skills import handler as skills_handler
from litellm.proxy._types import LiteLLM_SkillsTable, NewSkillRequest, UserAPIKeyAuth
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
from litellm.llms.litellm_proxy.skills.transformation import (
LiteLLMSkillsTransformationHandler,
)
from litellm.proxy._types import (
LiteLLM_SkillsTable,
LitellmUserRoles,
NewSkillRequest,
UserAPIKeyAuth,
)
from litellm.proxy.common_utils import resource_ownership
from litellm.skills import main as skills_main
@pytest.fixture(autouse=True)
@ -20,6 +30,184 @@ def _skill(skill_id: str, created_by: str | None) -> LiteLLM_SkillsTable:
)
def test_should_extract_skill_auth_from_supported_metadata_fields():
auth = UserAPIKeyAuth(user_id="user-1")
assert (
skills_main._get_user_api_key_auth_from_kwargs(
{"metadata": {"user_api_key_auth": auth}}
)
is auth
)
assert (
skills_main._get_user_api_key_auth_from_kwargs(
{"metadata": {}, "litellm_metadata": {"user_api_key_auth": auth}}
)
is auth
)
assert skills_main._get_user_api_key_auth_from_kwargs({"metadata": "bad"}) is None
def test_should_extract_skill_request_metadata_with_extra_body_precedence():
body_metadata = {"purpose": "body"}
requester_metadata = {"purpose": "requester"}
assert (
skills_main._get_skill_request_metadata(
{"metadata": {"requester_metadata": requester_metadata}},
{"metadata": body_metadata},
)
== body_metadata
)
assert (
skills_main._get_skill_request_metadata(
{"metadata": {"requester_metadata": requester_metadata}},
None,
)
== requester_metadata
)
assert (
skills_main._get_skill_request_metadata(
{"metadata": {}},
{"metadata": "bad"},
)
is None
)
def test_should_forward_skill_auth_through_sdk_entrypoints(monkeypatch):
auth = UserAPIKeyAuth(user_id="user-1")
handler = Mock()
handler.create_skill_handler.return_value = "created"
handler.list_skills_handler.return_value = "listed"
handler.get_skill_handler.return_value = "got"
handler.delete_skill_handler.return_value = "deleted"
monkeypatch.setattr(skills_main, "_get_litellm_skills_handler", lambda: handler)
assert (
skills_main.create_skill(
display_title="skill",
extra_body={"metadata": {"source": "request"}},
custom_llm_provider="litellm_proxy",
metadata={"user_api_key_auth": auth},
user_id="user-1",
)
== "created"
)
assert (
skills_main.list_skills(
custom_llm_provider="litellm_proxy",
metadata={"user_api_key_auth": auth},
)
== "listed"
)
assert (
skills_main.get_skill(
"litellm_skill_1",
custom_llm_provider="litellm_proxy",
metadata={"user_api_key_auth": auth},
)
== "got"
)
assert (
skills_main.delete_skill(
"litellm_skill_1",
custom_llm_provider="litellm_proxy",
metadata={"user_api_key_auth": auth},
)
== "deleted"
)
assert handler.create_skill_handler.call_args.kwargs["metadata"] == {
"source": "request"
}
assert handler.create_skill_handler.call_args.kwargs["user_api_key_dict"] is auth
assert handler.list_skills_handler.call_args.kwargs["user_api_key_dict"] is auth
assert handler.get_skill_handler.call_args.kwargs["user_api_key_dict"] is auth
assert handler.delete_skill_handler.call_args.kwargs["user_api_key_dict"] is auth
def test_should_build_resource_owner_scopes_for_auth_context():
auth = UserAPIKeyAuth(
user_id="user-1",
team_id="team-1",
org_id="org-1",
api_key="api-key-hash",
token="token-hash",
)
assert resource_ownership.get_resource_owner_scopes(auth) == [
"user-1",
"user:user-1",
"team:team-1",
"org:org-1",
"key:api-key-hash",
]
assert resource_ownership.get_primary_resource_owner_scope(auth) == "user-1"
assert resource_ownership.user_can_access_resource_owner("team:team-1", auth)
assert resource_ownership.get_resource_owner_scopes(
UserAPIKeyAuth(token="token-hash")
) == ["key:token-hash"]
assert resource_ownership.get_resource_owner_scopes(UserAPIKeyAuth()) == [
resource_ownership.UNSCOPED_RESOURCE_OWNER_SCOPE
]
def test_should_allow_admin_and_anonymous_resource_owner_paths():
admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value)
assert resource_ownership.is_proxy_admin(admin)
assert resource_ownership.user_can_access_resource_owner(None, admin)
assert resource_ownership.user_can_access_resource_owner(None, None)
assert not resource_ownership.user_can_access_resource_owner(
None, UserAPIKeyAuth(user_id="user-1")
)
@pytest.mark.asyncio
async def test_should_forward_skill_auth_through_transformation_handler(monkeypatch):
handler = LiteLLMSkillsTransformationHandler()
auth = UserAPIKeyAuth(user_id="user-1")
create_skill = AsyncMock(return_value=_skill("litellm_skill_created", "user-1"))
list_skills = AsyncMock(return_value=[_skill("litellm_skill_listed", "user-1")])
get_skill = AsyncMock(return_value=_skill("litellm_skill_got", "user-1"))
delete_skill = AsyncMock(return_value={"id": "litellm_skill_deleted"})
monkeypatch.setattr(LiteLLMSkillsHandler, "create_skill", create_skill)
monkeypatch.setattr(LiteLLMSkillsHandler, "list_skills", list_skills)
monkeypatch.setattr(LiteLLMSkillsHandler, "get_skill", get_skill)
monkeypatch.setattr(LiteLLMSkillsHandler, "delete_skill", delete_skill)
created = await handler._async_create_skill(
display_title="skill",
metadata={"source": "request"},
user_id="user-1",
user_api_key_dict=auth,
)
listed = await handler._async_list_skills(
limit=10,
offset=2,
user_api_key_dict=auth,
)
got = await handler._async_get_skill(
"litellm_skill_got",
user_api_key_dict=auth,
)
deleted = await handler._async_delete_skill(
"litellm_skill_deleted",
user_api_key_dict=auth,
)
assert created.id == "litellm_skill_created"
assert [skill.id for skill in listed.data] == ["litellm_skill_listed"]
assert got.id == "litellm_skill_got"
assert deleted.id == "litellm_skill_deleted"
assert create_skill.await_args.kwargs["user_api_key_dict"] is auth
assert list_skills.await_args.kwargs["user_api_key_dict"] is auth
assert get_skill.await_args.kwargs["user_api_key_dict"] is auth
assert delete_skill.await_args.kwargs["user_api_key_dict"] is auth
@pytest.mark.asyncio
async def test_should_store_team_owner_for_keys_without_user_id(monkeypatch):
table = AsyncMock()