diff --git a/tests/test_litellm/llms/litellm_proxy/test_skills_ownership.py b/tests/test_litellm/llms/litellm_proxy/test_skills_ownership.py index f6946e519d1..45ed172a845 100644 --- a/tests/test_litellm/llms/litellm_proxy/test_skills_ownership.py +++ b/tests/test_litellm/llms/litellm_proxy/test_skills_ownership.py @@ -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()