diff --git a/tests/e2e/quota_management/budgets/BUDGET_TEST_COVERAGE_MATRIX.md b/tests/e2e/quota_management/budgets/BUDGET_TEST_COVERAGE_MATRIX.md index a07bdf3d4e9..9c591ce277f 100644 --- a/tests/e2e/quota_management/budgets/BUDGET_TEST_COVERAGE_MATRIX.md +++ b/tests/e2e/quota_management/budgets/BUDGET_TEST_COVERAGE_MATRIX.md @@ -39,7 +39,7 @@ on the shared lifecycle (every entity it creates is deleted on teardown). | Pre-call reservation | `test_budget_reservation.py` | exercised by every enforcement test | **partial** | | Soft budget / alerts | `SlackAlerting/test_budget_alert_types.py` | `test_soft_budget_e2e::test_soft_budget_does_not_block` | **covered (new)** (block-vs-alert; the alert side-effect itself stays unit) | | Budget CRUD | `test_budget_endpoints.py` | `test_budget_crud_e2e` (roundtrip + delete) | **covered (new)** | -| Reset scheduling | `test_proxy_budget_reset.py` | `test_budget_crud_e2e::test_budget_duration_schedules_reset_on_key` | **covered (new)** (scheduling; actual zeroing is time-dependent -> unit) | +| Reset scheduling | `unit/proxy/common_utils/test_reset_budget_job.py` | `test_budget_crud_e2e::test_budget_duration_schedules_reset_on_key` | **covered (new)** (scheduling; actual zeroing is time-dependent -> unit) | | Multi-window budgets | `test_multi_budget_windows.py` | - | **gap** (window setup is fiddly; left to unit for now) | | Read budget+spend | `test_spend_management_endpoints.py` | `/key/info` asserted in CRUD + enforcement | **partial** | diff --git a/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md b/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md index e3f3bf009dc..a86f32f1837 100644 --- a/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md +++ b/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md @@ -50,7 +50,7 @@ proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`. | Endpoint | Existing | Status | Live e2e | |----------|----------|--------|----------| | `/spend/logs` (request_id / api_key) | `test_spend_management_endpoints.py` | covered | yes (primary read path; `test_spend_logs_endpoint_returns_spend` asserts 200 + spend, never 5xx) | -| `/spend/calculate` | `local_testing/test_spend_calculate_endpoint.py` | covered | yes (`test_spend_calculate_returns_nonzero_cost`) | +| `/spend/calculate` | `unit/proxy/spend_tracking/test_spend_management_endpoints.py` | covered | yes (`test_spend_calculate_returns_nonzero_cost`) | | `/spend/tags` | `test_spend_management_endpoints.py` | partial | yes (tag accuracy test) | | `/spend/logs/v2` pagination (total/total_pages/out-of-range) | `test_spend_query_optimization.py` | covered | yes (`test_spend_logs_v2_pagination_caps_pages_and_keeps_total`; filter takes the hashed token, not the raw key) | | whole spend GET surface (22 routes) | unit per-handler | partial | yes (`test_spend_routes.py` probes each for 404/5xx) | diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py deleted file mode 100644 index 16a611407b5..00000000000 --- a/tests/litellm_utils_tests/test_proxy_budget_reset.py +++ /dev/null @@ -1,349 +0,0 @@ -import asyncio -from datetime import datetime, timedelta, timezone -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest -from dotenv import load_dotenv - -load_dotenv() - -from litellm.proxy._types import LiteLLM_BudgetTableFull - - -from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob - -# Note: In our "fake" items we use dicts with fields that our fake reset functions modify. -# In a real-world scenario, these would be instances of LiteLLM_VerificationToken, LiteLLM_UserTable, etc. - - -def _attrify(d: dict): - """ - Wrap a dict so that attribute access (`.token`, `.user_id`, `.team_id`, - etc.) works alongside the existing item-access the fake_reset_* helpers - rely on. The reset job's narrow-write helpers use `getattr(item, "token", - None)` (et al), which returns None for plain dicts — that would silently - skip the row. - """ - - class _AttrDict(dict): - def __getattr__(self, k): - try: - return self[k] - except KeyError: - raise AttributeError(k) - - def __setattr__(self, k, v): - self[k] = v - - return _AttrDict(d) - - -def _wire_batcher_for_test(prisma_client, fail_commit=False): - """ - Wire prisma_client.db.batch_() to return a mock batcher whose .commit() is - awaitable and whose per-table .update()/.update_many() calls get captured. - The reset job writes every reset through prisma.db.batch_() — key/user/team - rows one by one, and the budget tier's cascade as a single transaction — so - tests must let that batch path complete. - - Only committed batches contribute to the returned list, mirroring prisma: - with fail_commit=True the transaction blows up and must persist nothing. - - Returns the list that will accumulate {table, op, where, data} dicts from - each captured write. - """ - batch_calls = [] - - def make_batcher(): - queued = [] - - class _Table: - def __init__(self, table_name): - self._table_name = table_name - - def update(self, where=None, data=None): - queued.append( - { - "table": self._table_name, - "op": "update", - "where": where, - "data": data, - } - ) - - def update_many(self, where=None, data=None): - queued.append( - { - "table": self._table_name, - "op": "update_many", - "where": where, - "data": data, - } - ) - - async def commit(): - if fail_commit: - raise RuntimeError("simulated Postgres failure committing the batch") - batch_calls.extend(queued) - - batcher = MagicMock() - batcher.litellm_verificationtoken = _Table("key") - batcher.litellm_usertable = _Table("user") - batcher.litellm_teamtable = _Table("team") - batcher.litellm_budgettable = _Table("budget") - batcher.litellm_teammembership = _Table("team_membership") - batcher.litellm_organizationtable = _Table("org") - batcher.litellm_tagtable = _Table("tag") - batcher.litellm_endusertable = _Table("enduser") - batcher.commit = commit - return batcher - - prisma_client.db.batch_ = MagicMock(side_effect=make_batcher) - return batch_calls - - -def _wire_cascade_reads_for_test(prisma_client, endusers=()): - """ - The budget tier's cascade reads the rows it is about to zero, so their - spend counters can be invalidated after the commit. Give each of those - tables an awaitable find_many so the reads resolve instead of falling into - the job's warn-and-continue path. - - End users are read by the post-commit invalidation walk rather than by - ``get_data``, so callers that care about customers pass them here. - """ - for table in ( - "litellm_teammembership", - "litellm_verificationtoken", - "litellm_organizationtable", - "litellm_tagtable", - ): - getattr(prisma_client.db, table).find_many = AsyncMock(return_value=[]) - prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=list(endusers)) - - - - - - - - -@pytest.mark.asyncio -async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance(): - """ - The happy path: every end user the tier gates is zeroed and the tier's - budget_reset_at advances, all inside one transaction. - """ - endusers = [ - _attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"}) - for i in range(1, 7) - ] - - budget1 = LiteLLM_BudgetTableFull( - **{ - "budget_id": "budget1", - "max_budget": 65.0, - "budget_duration": "2d", - "created_at": datetime.now(timezone.utc) - timedelta(days=3), - } - ) - - prisma_client = MagicMock() - - async def get_data_mock(table_name, *args, **kwargs): - if table_name == "budget": - return [budget1] - elif table_name == "enduser": - return endusers - return [] - - prisma_client.get_data = AsyncMock() - prisma_client.get_data.side_effect = get_data_mock - prisma_client.update_data = AsyncMock() - batch_calls = _wire_batcher_for_test(prisma_client) - - proxy_logging_obj = MagicMock() - proxy_logging_obj.service_logging_obj = MagicMock() - proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() - proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() - - job = ResetBudgetJob(proxy_logging_obj, prisma_client) - - await job.reset_budget_for_litellm_budget_table() - await asyncio.sleep(0.1) - - assert prisma_client.db.batch_.call_count == 1, "the cascade must be one transaction" - - enduser_writes = [c for c in batch_calls if c["table"] == "enduser"] - assert len(enduser_writes) == 1 - assert enduser_writes[0]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}} - assert enduser_writes[0]["data"] == {"spend": 0} - - budget_writes = [c for c in batch_calls if c["table"] == "budget"] - assert len(budget_writes) == 1 - assert budget_writes[0]["where"] == {"budget_id": "budget1"} - assert budget_writes[0]["data"]["budget_reset_at"] > datetime.now(timezone.utc) - - proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called() - - - - -@pytest.mark.asyncio -async def test_reset_budget_continues_other_categories_on_failure(): - """ - Test that executing the overall reset_budget() method continues to process keys, users, and teams, - even if one of the sub-categories (here, users) experiences a partial failure. - - In this simulation: - - All keys are processed successfully. - - One of the two users fails. - - All teams are processed successfully. - - We then assert that: - - update_data is called for each category with the correctly updated items. - - Each get_data call is made (indicating that one failing category did not abort the others). - """ - # Arrange dummy items for each table - key1 = {"id": "key1", "spend": 10.0, "budget_duration": 60} - key2 = {"id": "key2", "spend": 15.0, "budget_duration": 60} - user1 = { - "id": "user1", - "spend": 20.0, - "budget_duration": 120, - } # Will fail in user reset - user2 = {"id": "user2", "spend": 25.0, "budget_duration": 120} # Succeeds - team1 = {"id": "team1", "spend": 30.0, "budget_duration": 180} - team2 = {"id": "team2", "spend": 35.0, "budget_duration": 180} - enduser1 = {"user_id": "user1", "spend": 25.0, "budget_id": "budget1"} - budget1 = LiteLLM_BudgetTableFull( - **{ - "budget_id": "budget1", - "max_budget": 65.0, - "budget_duration": "2d", - "created_at": datetime.now(timezone.utc) - timedelta(days=3), - } - ) - - prisma_client = MagicMock() - - async def fake_get_data(*, table_name, query_type, **kwargs): - if table_name == "key": - return [key1, key2] - elif table_name == "user": - return [user1, user2] - elif table_name == "team": - return [team1, team2] - elif table_name == "budget": - return [budget1] - elif table_name == "enduser": - return [enduser1] - return [] - - prisma_client.get_data = AsyncMock(side_effect=fake_get_data) - prisma_client.update_data = AsyncMock() - batch_calls = _wire_batcher_for_test(prisma_client) - # ID fields required by the new write path's where clauses; _AttrDict - # lets getattr() see them alongside the item-access fake_reset_* helpers use. - for k in [key1, key2]: - k.setdefault("token", k["id"]) - for u in [user1, user2]: - u.setdefault("user_id", u["id"]) - for t in [team1, team2]: - t.setdefault("team_id", t["id"]) - key1, key2 = _attrify(key1), _attrify(key2) - user1, user2 = _attrify(user1), _attrify(user2) - team1, team2 = _attrify(team1), _attrify(team2) - enduser1 = _attrify(enduser1) - pre_reset_spend = { - **{k["token"]: k["spend"] for k in [key1, key2]}, - **{u["user_id"]: u["spend"] for u in [user2]}, - **{t["team_id"]: t["spend"] for t in [team1, team2]}, - } - _wire_cascade_reads_for_test(prisma_client, endusers=[enduser1]) - - proxy_logging_obj = MagicMock() - proxy_logging_obj.service_logging_obj = MagicMock() - proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() - proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() - - job = ResetBudgetJob(proxy_logging_obj, prisma_client) - - async def fake_reset_key(key, current_time, reset_settings=None): - key["spend"] = 0.0 - key["budget_reset_at"] = ( - current_time + timedelta(seconds=key["budget_duration"]) - ).isoformat() - return key - - async def fake_reset_user(user, current_time, reset_settings=None): - if user["id"] == "user1": - raise Exception("Simulated failure for user1") - user["spend"] = 0.0 - user["budget_reset_at"] = ( - current_time + timedelta(seconds=user["budget_duration"]) - ).isoformat() - return user - - async def fake_reset_team(team, current_time, reset_settings=None): - team["spend"] = 0.0 - team["budget_reset_at"] = ( - current_time + timedelta(seconds=team["budget_duration"]) - ).isoformat() - return team - - with ( - patch.object( - ResetBudgetJob, "_reset_budget_for_key", side_effect=fake_reset_key - ) as mock_reset_key, - patch.object( - ResetBudgetJob, "_reset_budget_for_user", side_effect=fake_reset_user - ) as mock_reset_user, - patch.object( - ResetBudgetJob, "_reset_budget_for_team", side_effect=fake_reset_team - ) as mock_reset_team, - ): - # Call the overall reset_budget method. - await job.reset_budget() - await asyncio.sleep(0.1) - - # Verify that get_data was called for each table. We can check the table names across calls. - called_tables = { - call.kwargs.get("table_name") for call in prisma_client.get_data.await_args_list - } - assert called_tables == {"key", "user", "team", "budget"} - # Customers are not part of that set: the cascade zeroes them by budget link - # and reads them only afterwards, to invalidate their cached spend. - prisma_client.db.litellm_endusertable.find_many.assert_awaited() - - # Every category writes through the batch path now, so update_data is unused. - prisma_client.update_data.assert_not_awaited() - - # The budget tier's cascade still ran despite the failing user category. - assert len([c for c in batch_calls if c["table"] == "team_membership"]) == 1 - enduser_writes = [c for c in batch_calls if c["table"] == "enduser"] - assert len(enduser_writes) == 1 - assert enduser_writes[0]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}} - assert enduser_writes[0]["data"] == {"spend": 0} - - # Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams. - # `op` separates the per-row resets from the cascade sweep, which also - # targets the key table. - key_writes = [c for c in batch_calls if c["table"] == "key" and c["op"] == "update"] - user_writes = [c for c in batch_calls if c["table"] == "user"] - team_writes = [c for c in batch_calls if c["table"] == "team"] - assert len(key_writes) == 2 - assert len(user_writes) == 1 - assert user_writes[0]["where"] == {"user_id": "user2"} - assert len(team_writes) == 2 - # Every batched write must carry only the two reset fields, never the full row. - for c in key_writes + user_writes + team_writes: - assert set(c["data"].keys()) == {"spend", "budget_reset_at"} - assert c["data"]["spend"] == { - "decrement": pre_reset_spend[next(iter(c["where"].values()))] - } - - -# --------------------------------------------------------------------------- -# Additional tests for service logger behavior (keys, users, teams, endusers) -# --------------------------------------------------------------------------- diff --git a/tests/local_testing/test_auth_utils.py b/tests/local_testing/test_auth_utils.py deleted file mode 100644 index cf6d65acd2d..00000000000 --- a/tests/local_testing/test_auth_utils.py +++ /dev/null @@ -1,190 +0,0 @@ -# What is this? -## Tests if proxy/auth/auth_utils.py works as expected - -import sys, os, asyncio, time, random, uuid -import traceback -from dotenv import load_dotenv - -load_dotenv() - -import pytest -import litellm -def test_get_end_user_id_from_request_body_always_returns_str(): - from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body - from fastapi import Request - from unittest.mock import MagicMock - - # Create a mock Request object - mock_request = MagicMock(spec=Request) - mock_request.headers = {} - - request_body = {"user": 123} - end_user_id = get_end_user_id_from_request_body( - request_body, dict(mock_request.headers) - ) - assert end_user_id == "123" - assert isinstance(end_user_id, str) - - -@pytest.mark.parametrize( - "headers, general_settings_config, request_body, expected_user_id", - [ - # Test 1: user_header_name configured and header present - ( - {"X-User-ID": "header-user-123"}, - {"user_header_name": "X-User-ID"}, - {"user": "body-user-456"}, - "header-user-123", # Header should take precedence - ), - # Test 2: user_header_name configured but header not present, fallback to body - ( - {}, - {"user_header_name": "X-User-ID"}, - {"user": "body-user-456"}, - "body-user-456", # Should fall back to body - ), - # Test 3: user_header_name not configured, should use body - ( - {"X-User-ID": "header-user-123"}, - {}, - {"user": "body-user-456"}, - "body-user-456", # Should ignore header when not configured - ), - # Test 4: user_header_name configured, header present, but no body user - ( - {"X-Custom-User": "header-only-user"}, - {"user_header_name": "X-Custom-User"}, - {"model": "gpt-4"}, - "header-only-user", # Should use header - ), - # Test 5: user_header_name configured but header is empty string - ( - {"X-User-ID": ""}, - {"user_header_name": "X-User-ID"}, - {"user": "body-user-456"}, - "body-user-456", # Should fall back to body when header is empty - ), - # Test 6: user_header_name configured with case-insensitive header - ( - {"x-user-id": "lowercase-header-user"}, - {"user_header_name": "x-user-id"}, - {"user": "body-user-456"}, - "lowercase-header-user", - ), - # Test 7: user_header_name configured but set to None - ( - {"X-User-ID": "header-user-123"}, - {"user_header_name": None}, - {"user": "body-user-456"}, - "body-user-456", # Should fall back to body when header name is None - ), - # Test 8: user_header_name is not a string - ( - {"X-User-ID": "header-user-123"}, - {"user_header_name": 123}, - {"user": "body-user-456"}, - "body-user-456", # Should fall back to body when header name is not a string - ), - # Test 9: Multiple fallback sources - litellm_metadata - ( - {}, - {"user_header_name": "X-User-ID"}, - {"litellm_metadata": {"user": "litellm-user-789"}}, - "litellm-user-789", - ), - # Test 10: Multiple fallback sources - metadata.user_id - ( - {}, - {"user_header_name": "X-User-ID"}, - {"metadata": {"user_id": "metadata-user-999"}}, - "metadata-user-999", - ), - # Test 11: Header takes precedence over all body sources - ( - {"X-User-ID": "header-priority"}, - {"user_header_name": "X-User-ID"}, - { - "user": "body-user", - "litellm_metadata": {"user": "litellm-user"}, - "metadata": {"user_id": "metadata-user"}, - }, - "header-priority", - ), - # Test 12: user_header_name is matched case-insensitively - ( - {"x-user-id": "lowercase-header-user"}, - {"user_header_name": "X-User-ID"}, - {"user": "body-user-456"}, - "lowercase-header-user", - ), - ], -) -def test_get_end_user_id_from_request_body_with_user_header_name( - headers, general_settings_config, request_body, expected_user_id -): - """Test that get_end_user_id_from_request_body respects user_header_name property""" - from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body - from fastapi import Request - from unittest.mock import MagicMock, patch - - # Create a mock Request object with headers - mock_request = MagicMock(spec=Request) - mock_request.headers = headers - - # Mock general_settings at the proxy_server module level - with patch("litellm.proxy.proxy_server.general_settings", general_settings_config): - end_user_id = get_end_user_id_from_request_body( - request_body, dict(mock_request.headers) - ) - assert end_user_id == expected_user_id - - -def test_get_end_user_id_from_request_body_no_user_found(): - """Test that function returns None when no user ID is found anywhere""" - from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body - from fastapi import Request - from unittest.mock import MagicMock, patch - - # Create a mock Request object with no relevant headers - mock_request = MagicMock(spec=Request) - mock_request.headers = {"X-Other-Header": "some-value"} - - # Mock general_settings with user_header_name that doesn't match headers - general_settings_config = {"user_header_name": "X-User-ID"} - - # Request body with no user identifiers - request_body = { - "model": "gpt-4", - "messages": [{"role": "user", "content": "hello"}], - } - - with patch("litellm.proxy.proxy_server.general_settings", general_settings_config): - end_user_id = get_end_user_id_from_request_body( - request_body, dict(mock_request.headers) - ) - assert end_user_id is None - - -def test_get_end_user_id_from_request_body_backwards_compatibility(): - """Test that function works with just request_body parameter (backwards compatibility)""" - from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body - - # Test with just request_body - should work like before - request_body = {"user": "test-user-123"} - end_user_id = get_end_user_id_from_request_body(request_body) - assert end_user_id == "test-user-123" - - # Test with litellm_metadata - request_body = {"litellm_metadata": {"user": "litellm-user-456"}} - end_user_id = get_end_user_id_from_request_body(request_body) - assert end_user_id == "litellm-user-456" - - # Test with metadata.user_id - request_body = {"metadata": {"user_id": "metadata-user-789"}} - end_user_id = get_end_user_id_from_request_body(request_body) - assert end_user_id == "metadata-user-789" - - # Test with no user - should return None - request_body = {"model": "gpt-4"} - end_user_id = get_end_user_id_from_request_body(request_body) - assert end_user_id is None diff --git a/tests/local_testing/test_llm_guard.py b/tests/local_testing/test_llm_guard.py deleted file mode 100644 index fc30f644028..00000000000 --- a/tests/local_testing/test_llm_guard.py +++ /dev/null @@ -1,215 +0,0 @@ -# What is this? -## This tests the llm guard integration - -# What is this? -## Unit test for presidio pii masking -import sys, os, asyncio, time, random -from datetime import datetime -import traceback -from dotenv import load_dotenv - -load_dotenv() - -import pytest -from fastapi import HTTPException - -import litellm -from litellm_enterprise.enterprise_callbacks.llm_guard import _ENTERPRISE_LLMGuard -from litellm import Router, mock_completion -from litellm.proxy.utils import ProxyLogging, hash_token -from litellm.proxy._types import UserAPIKeyAuth -from litellm.caching.caching import DualCache - -### UNIT TESTS FOR LLM GUARD ### - - -@pytest.mark.asyncio -async def test_llm_guard_valid_response(): - """ - A valid (is_valid=True) LLM Guard response must apply the returned - sanitized_prompt back onto the request data so the provider receives the - redacted content. - """ - litellm.llm_guard_mode = "all" - input_a_anonymizer_results = { - "sanitized_prompt": "hello world", - "is_valid": True, - "scanners": {"Regex": 0.0}, - } - llm_guard = _ENTERPRISE_LLMGuard( - mock_testing=True, mock_redacted_text=input_a_anonymizer_results - ) - - _api_key = "sk-98765" - _api_key = hash_token("sk-98765") - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) - local_cache = DualCache() - - data = { - "messages": [ - { - "role": "user", - "content": "hello world, my name is Jane Doe. My number is: 23r323r23r2wwkl", - } - ] - } - - result = await llm_guard.async_moderation_hook( - data=data, - user_api_key_dict=user_api_key_dict, - call_type="completion", - ) - - assert result is data - assert data["messages"][0]["content"] == "hello world" - - -@pytest.mark.asyncio -async def test_llm_guard_sanitizes_multimodal_and_input(): - """ - Sanitization must reach text parts of multimodal message content and the - ``input`` field (embeddings/moderation) while leaving non-text parts intact. - """ - litellm.llm_guard_mode = "all" - llm_guard = _ENTERPRISE_LLMGuard( - mock_testing=True, - mock_redacted_text={ - "sanitized_prompt": "email: [REDACTED]", - "is_valid": True, - "scanners": {"Regex": 0.0}, - }, - ) - user_api_key_dict = UserAPIKeyAuth(api_key=hash_token("sk-98765")) - - image_part = {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}} - data = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "email: person@example.com"}, - image_part, - ], - } - ] - } - result = await llm_guard.async_moderation_hook( - data=data, user_api_key_dict=user_api_key_dict, call_type="completion" - ) - assert result["messages"][0]["content"][0]["text"] == "email: [REDACTED]" - assert result["messages"][0]["content"][1] == image_part - - input_data = {"input": ["email: person@example.com", "another prompt"]} - input_result = await llm_guard.async_moderation_hook( - data=input_data, user_api_key_dict=user_api_key_dict, call_type="embeddings" - ) - assert input_result["input"] == ["email: [REDACTED]", "email: [REDACTED]"] - - -@pytest.mark.asyncio -async def test_llm_guard_error_raising(): - """ - Tests to see llm guard raises an error for a flagged response - """ - - input_b_anonymizer_results = { - "sanitized_prompt": "hello world", - "is_valid": False, - "scanners": {"Regex": 0.0}, - } - llm_guard = _ENTERPRISE_LLMGuard( - mock_testing=True, mock_redacted_text=input_b_anonymizer_results - ) - - _api_key = "sk-98765" - _api_key = hash_token("sk-98765") - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) - local_cache = DualCache() - - with pytest.raises(HTTPException) as exc_info: - await llm_guard.async_moderation_hook( - data={ - "messages": [ - { - "role": "user", - "content": "hello world, my name is Jane Doe. My number is: 23r323r23r2wwkl", - } - ] - }, - user_api_key_dict=user_api_key_dict, - call_type="completion", - ) - - assert exc_info.value.status_code == 400 - assert exc_info.value.detail == {"error": "Violated content safety policy"} - - -def test_llm_guard_key_specific_mode(): - """ - Tests to see if llm guard 'key-specific' permissions work - """ - litellm.llm_guard_mode = "key-specific" - - llm_guard = _ENTERPRISE_LLMGuard(mock_testing=True) - - _api_key = "sk-98765" - # NOT ENABLED - user_api_key_dict = UserAPIKeyAuth( - api_key=_api_key, - ) - - request_data = {} - should_proceed = llm_guard.should_proceed( - user_api_key_dict=user_api_key_dict, data=request_data - ) - - assert should_proceed == False - - # ENABLED - user_api_key_dict = UserAPIKeyAuth( - api_key=_api_key, permissions={"enable_llm_guard_check": True} - ) - - request_data = {} - - should_proceed = llm_guard.should_proceed( - user_api_key_dict=user_api_key_dict, data=request_data - ) - - assert should_proceed == True - - -def test_llm_guard_request_specific_mode(): - """ - Tests to see if llm guard 'request-specific' permissions work - """ - litellm.llm_guard_mode = "request-specific" - - llm_guard = _ENTERPRISE_LLMGuard(mock_testing=True) - - _api_key = "sk-98765" - # NOT ENABLED - user_api_key_dict = UserAPIKeyAuth( - api_key=_api_key, - ) - - request_data = {} - - should_proceed = llm_guard.should_proceed( - user_api_key_dict=user_api_key_dict, data=request_data - ) - - assert should_proceed == False - - # ENABLED - user_api_key_dict = UserAPIKeyAuth( - api_key=_api_key, permissions={"enable_llm_guard_check": True} - ) - - request_data = {"metadata": {"permissions": {"enable_llm_guard_check": True}}} - - should_proceed = llm_guard.should_proceed( - user_api_key_dict=user_api_key_dict, data=request_data - ) - - assert should_proceed == True diff --git a/tests/local_testing/test_openai_moderations_hook.py b/tests/local_testing/test_openai_moderations_hook.py deleted file mode 100644 index ce62c7ed041..00000000000 --- a/tests/local_testing/test_openai_moderations_hook.py +++ /dev/null @@ -1,179 +0,0 @@ -# What is this? -## This tests the llm guard integration - -# What is this? -## Unit test for presidio pii masking -import sys, os, asyncio, time, random -from datetime import datetime -import traceback -from dotenv import load_dotenv - -load_dotenv() - -import pytest -import litellm -from litellm.proxy.enterprise.enterprise_hooks.openai_moderation import ( - ENTERPRISE_OpenAI_Moderation, -) -from litellm import Router, mock_completion -from litellm.proxy.utils import ProxyLogging, hash_token -from litellm.proxy._types import UserAPIKeyAuth -from litellm.caching.caching import DualCache - -### UNIT TESTS FOR OpenAI Moderation ### - - -@pytest.mark.asyncio -async def test_openai_moderation_error_raising(monkeypatch): - """ - Tests to see OpenAI Moderation raises an error for a flagged response - """ - from unittest.mock import AsyncMock, MagicMock - from litellm.types.llms.openai import OpenAIModerationResponse - - litellm.openai_moderations_model_name = "omni-moderation-latest" - openai_mod = ENTERPRISE_OpenAI_Moderation() - _api_key = "sk-98765" - _api_key = hash_token("sk-98765") - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) - local_cache = DualCache() - - llm_router = litellm.Router( - model_list=[ - { - "model_name": "omni-moderation-latest", - "litellm_params": { - "model": "omni-moderation-latest", - "api_key": os.environ.get("OPENAI_API_KEY", "fake-key"), - }, - } - ] - ) - - # Mock the amoderation call to return a flagged response - mock_response = MagicMock(spec=OpenAIModerationResponse) - mock_response.results = [MagicMock(flagged=True)] - - async def mock_amoderation(*args, **kwargs): - return mock_response - - llm_router.amoderation = mock_amoderation - - import litellm.proxy.proxy_server as proxy_server - - monkeypatch.setattr(proxy_server, "llm_router", llm_router) - - with pytest.raises(Exception, match="Violated content safety policy") as exc_info: - await openai_mod.async_moderation_hook( - data={ - "messages": [ - { - "role": "user", - "content": "fuck off you're the worst", - } - ] - }, - user_api_key_dict=user_api_key_dict, - call_type="completion", - ) - e = exc_info.value - print("Got exception: ", e) - assert "Violated content safety policy" in str(e) - - -@pytest.mark.asyncio -async def test_openai_moderation_responses_api_input_field(): - """ - Tests that OpenAI Moderation works with Responses API input field via apply_guardrail. - - This test verifies that the unified guardrail interface (apply_guardrail) correctly - handles different input types: plain text strings, structured messages, and lists. - """ - from unittest.mock import patch - from litellm.types.llms.openai import ( - OpenAIModerationResponse, - OpenAIModerationResult, - ) - from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import ( - OpenAIModerationGuardrail, - ) - from litellm.types.utils import GenericGuardrailAPIInputs - - # Initialize the open-source OpenAI Moderation guardrail - openai_mod = OpenAIModerationGuardrail( - guardrail_name="openai-moderation-test", - api_key="fake-key-for-testing", - model="omni-moderation-latest", - ) - - # Mock the async_make_request to return a flagged response - mock_moderation_response = OpenAIModerationResponse( - id="modr-123", - model="omni-moderation-latest", - results=[ - OpenAIModerationResult( - flagged=True, - categories={"violence": True, "hate": False}, - category_scores={"violence": 0.95, "hate": 0.1}, - category_applied_input_types=None, - ) - ], - ) - - with patch.object( - openai_mod, "async_make_request", return_value=mock_moderation_response - ): - # Test 1: Responses API / Embeddings with texts (string input) - inputs = GenericGuardrailAPIInputs(texts=["I want to hurt people"]) - - with pytest.raises(Exception, match="Violated OpenAI moderation policy") as exc_info: - await openai_mod.apply_guardrail( - inputs=inputs, - request_data={"model": "gpt-4o", "input": "I want to hurt people"}, - input_type="request", - ) - e = exc_info.value - print("Got exception for texts input: ", e) - assert "Violated OpenAI moderation policy" in str(e) - - # Test 2: Responses API with structured_messages (list of message objects) - inputs = GenericGuardrailAPIInputs( - structured_messages=[ - {"role": "user", "content": "I want to hurt people"} - ] - ) - - with pytest.raises(Exception, match="Violated OpenAI moderation policy") as exc_info: - await openai_mod.apply_guardrail( - inputs=inputs, - request_data={ - "model": "gpt-4o", - "input": [{"role": "user", "content": "I want to hurt people"}], - }, - input_type="request", - ) - e = exc_info.value - print("Got exception for structured_messages input: ", e) - assert "Violated OpenAI moderation policy" in str(e) - - # Test 3: Chat Completions with structured_messages - inputs = GenericGuardrailAPIInputs( - structured_messages=[ - {"role": "user", "content": "I want to hurt people"} - ] - ) - - with pytest.raises(Exception, match="Violated OpenAI moderation policy") as exc_info: - await openai_mod.apply_guardrail( - inputs=inputs, - request_data={ - "model": "gpt-4o", - "messages": [{"role": "user", "content": "I want to hurt people"}], - }, - input_type="request", - ) - e = exc_info.value - print("Got exception for chat completions input: ", e) - assert "Violated OpenAI moderation policy" in str(e) - - print("✓ All Responses API moderation tests passed!") diff --git a/tests/local_testing/test_secret_detect_hook.py b/tests/local_testing/test_secret_detect_hook.py index 137d6c16101..d04804f35af 100644 --- a/tests/local_testing/test_secret_detect_hook.py +++ b/tests/local_testing/test_secret_detect_hook.py @@ -20,210 +20,18 @@ from starlette.datastructures import URL import litellm from litellm import Router, mock_completion -from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth from litellm_enterprise.enterprise_callbacks.secret_detection import ( _ENTERPRISE_SecretDetection, ) from litellm.proxy.proxy_server import chat_completion -from litellm.proxy.utils import ProxyLogging, hash_token from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE ### UNIT TESTS FOR OpenAI Moderation ### -@pytest.mark.asyncio -async def test_basic_secret_detection_chat(): - """ - Tests to see if secret detection hook will mask api keys - - - It should mask the following API_KEY = 'sk_1234567890abcdef' and OPENAI_API_KEY = 'sk_1234567890abcdef' - """ - secret_instance = _ENTERPRISE_SecretDetection() - _api_key = "sk-98765" - _api_key = hash_token("sk-98765") - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) - local_cache = DualCache() - - from litellm.proxy.proxy_server import llm_router - - test_data = { - "messages": [ - { - "role": "user", - "content": "Hey, how's it going, API_KEY = 'sk_1234567890abcdef'", - }, - { - "role": "assistant", - "content": "Hello! I'm doing well. How can I assist you today?", - }, - { - "role": "user", - "content": "this is my OPENAI_API_KEY = 'sk_1234567890abcdef'", - }, - { - "role": "user", - "content": "My hi API Key is sk-Pc4nlxVoMz41290028TbMCxx, does it seem to be in the correct format?", - }, - {"role": "user", "content": "i think it is +1 412-555-5555"}, - ], - "model": "gpt-3.5-turbo", - } - - await secret_instance.async_pre_call_hook( - cache=local_cache, - data=test_data, - user_api_key_dict=user_api_key_dict, - call_type="completion", - ) - print( - "test data after running pre_call_hook: Expect all API Keys to be masked", - test_data, - ) - - assert test_data == { - "messages": [ - {"role": "user", "content": "Hey, how's it going, API_KEY = '[REDACTED]'"}, - { - "role": "assistant", - "content": "Hello! I'm doing well. How can I assist you today?", - }, - {"role": "user", "content": "this is my OPENAI_API_KEY = '[REDACTED]'"}, - { - "role": "user", - "content": "My hi API Key is [REDACTED], does it seem to be in the correct format?", - }, - {"role": "user", "content": "i think it is +1 412-555-5555"}, - ], - "model": "gpt-3.5-turbo", - }, "Expect all API Keys to be masked" - - -@pytest.mark.asyncio -async def test_basic_secret_detection_text_completion(): - """ - Tests to see if secret detection hook will mask api keys - - - It should mask the following API_KEY = 'sk_1234567890abcdef' and OPENAI_API_KEY = 'sk_1234567890abcdef' - """ - secret_instance = _ENTERPRISE_SecretDetection() - _api_key = "sk-98765" - _api_key = hash_token("sk-98765") - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) - local_cache = DualCache() - - from litellm.proxy.proxy_server import llm_router - - test_data = { - "prompt": "Hey, how's it going, API_KEY = 'sk_1234567890abcdef', my OPENAI_API_KEY = 'sk_1234567890abcdef' and i want to know what is the weather", - "model": "gpt-3.5-turbo", - } - - await secret_instance.async_pre_call_hook( - cache=local_cache, - data=test_data, - user_api_key_dict=user_api_key_dict, - call_type="completion", - ) - - assert test_data == { - "prompt": "Hey, how's it going, API_KEY = '[REDACTED]', my OPENAI_API_KEY = '[REDACTED]' and i want to know what is the weather", - "model": "gpt-3.5-turbo", - } - print( - "test data after running pre_call_hook: Expect all API Keys to be masked", - test_data, - ) - - -@pytest.mark.asyncio -async def test_basic_secret_detection_embeddings(): - """ - Tests to see if secret detection hook will mask api keys - - - It should mask the following API_KEY = 'sk_1234567890abcdef' and OPENAI_API_KEY = 'sk_1234567890abcdef' - """ - secret_instance = _ENTERPRISE_SecretDetection() - _api_key = "sk-98765" - _api_key = hash_token("sk-98765") - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) - local_cache = DualCache() - - from litellm.proxy.proxy_server import llm_router - - test_data = { - "input": "Hey, how's it going, API_KEY = 'sk_1234567890abcdef', my OPENAI_API_KEY = 'sk_1234567890abcdef' and i want to know what is the weather", - "model": "gpt-3.5-turbo", - } - - await secret_instance.async_pre_call_hook( - cache=local_cache, - data=test_data, - user_api_key_dict=user_api_key_dict, - call_type="embedding", - ) - - assert test_data == { - "input": "Hey, how's it going, API_KEY = '[REDACTED]', my OPENAI_API_KEY = '[REDACTED]' and i want to know what is the weather", - "model": "gpt-3.5-turbo", - } - print( - "test data after running pre_call_hook: Expect all API Keys to be masked", - test_data, - ) - - -@pytest.mark.asyncio -async def test_basic_secret_detection_embeddings_list(): - """ - Tests to see if secret detection hook will mask api keys - - - It should mask the following API_KEY = 'sk_1234567890abcdef' and OPENAI_API_KEY = 'sk_1234567890abcdef' - """ - secret_instance = _ENTERPRISE_SecretDetection() - _api_key = "sk-98765" - _api_key = hash_token("sk-98765") - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) - local_cache = DualCache() - - from litellm.proxy.proxy_server import llm_router - - test_data = { - "input": [ - "hey", - "how's it going, API_KEY = 'sk_1234567890abcdef'", - "my OPENAI_API_KEY = 'sk_1234567890abcdef' and i want to know what is the weather", - ], - "model": "gpt-3.5-turbo", - } - - await secret_instance.async_pre_call_hook( - cache=local_cache, - data=test_data, - user_api_key_dict=user_api_key_dict, - call_type="embedding", - ) - - print( - "test data after running pre_call_hook: Expect all API Keys to be masked", - test_data, - ) - assert test_data == { - "input": [ - "hey", - "how's it going, API_KEY = '[REDACTED]'", - "my OPENAI_API_KEY = '[REDACTED]' and i want to know what is the weather", - ], - "model": "gpt-3.5-turbo", - } - - class testLogger(CustomLogger): def __init__(self): diff --git a/tests/local_testing/test_spend_calculate_endpoint.py b/tests/local_testing/test_spend_calculate_endpoint.py deleted file mode 100644 index 054dc398039..00000000000 --- a/tests/local_testing/test_spend_calculate_endpoint.py +++ /dev/null @@ -1,136 +0,0 @@ - -import pytest -from dotenv import load_dotenv -from fastapi import Request -from fastapi.routing import APIRoute - -import litellm -from litellm.proxy._types import SpendCalculateRequest -from litellm.proxy.spend_tracking.spend_management_endpoints import calculate_spend -from litellm.router import Router - -# this file is to test litellm/proxy - - - -@pytest.mark.asyncio -async def test_spend_calc_model_messages(): - cost_obj = await calculate_spend( - request=SpendCalculateRequest( - model="gpt-3.5-turbo", - messages=[ - {"role": "user", "content": "What is the capital of France?"}, - ], - ) - ) - - print("calculated cost", cost_obj) - cost = cost_obj["cost"] - assert cost > 0.0 - - -@pytest.mark.asyncio -async def test_spend_calc_model_on_router_messages(): - from litellm.proxy.proxy_server import llm_router as init_llm_router - - temp_llm_router = Router( - model_list=[ - { - "model_name": "special-llama-model", - "litellm_params": { - "model": "groq/openai/gpt-oss-20b", - }, - } - ] - ) - - setattr(litellm.proxy.proxy_server, "llm_router", temp_llm_router) - - cost_obj = await calculate_spend( - request=SpendCalculateRequest( - model="special-llama-model", - messages=[ - {"role": "user", "content": "What is the capital of France?"}, - ], - ) - ) - - print("calculated cost", cost_obj) - _cost = cost_obj["cost"] - - assert _cost > 0.0 - - # set router to init value - setattr(litellm.proxy.proxy_server, "llm_router", init_llm_router) - - -@pytest.mark.asyncio -async def test_spend_calc_using_response(): - cost_obj = await calculate_spend( - request=SpendCalculateRequest( - completion_response={ - "id": "chatcmpl-3bc7abcd-f70b-48ab-a16c-dfba0b286c86", - "choices": [ - { - "finish_reason": "stop", - "index": 0, - "message": { - "content": "Yooo! What's good?", - "role": "assistant", - }, - } - ], - "created": "1677652288", - "model": "groq/openai/gpt-oss-20b", - "object": "chat.completion", - "system_fingerprint": "fp_873a560973", - "usage": { - "completion_tokens": 8, - "prompt_tokens": 12, - "total_tokens": 20, - }, - } - ) - ) - - print("calculated cost", cost_obj) - cost = cost_obj["cost"] - assert cost > 0.0 - - -@pytest.mark.asyncio -async def test_spend_calc_model_alias_on_router_messages(): - from litellm.proxy.proxy_server import llm_router as init_llm_router - - temp_llm_router = Router( - model_list=[ - { - "model_name": "gpt-4o", - "litellm_params": { - "model": "gpt-4o", - }, - } - ], - model_group_alias={ - "gpt4o": "gpt-4o", - }, - ) - - setattr(litellm.proxy.proxy_server, "llm_router", temp_llm_router) - - cost_obj = await calculate_spend( - request=SpendCalculateRequest( - model="gpt4o", - messages=[ - {"role": "user", "content": "What is the capital of France?"}, - ], - ) - ) - - print("calculated cost", cost_obj) - _cost = cost_obj["cost"] - - assert _cost > 0.0 - - # set router to init value - setattr(litellm.proxy.proxy_server, "llm_router", init_llm_router) diff --git a/tests/logging_callback_tests/test_spend_logs.py b/tests/logging_callback_tests/test_spend_logs.py deleted file mode 100644 index 215bfd193ce..00000000000 --- a/tests/logging_callback_tests/test_spend_logs.py +++ /dev/null @@ -1,391 +0,0 @@ -import traceback -from litellm._uuid import uuid - -from dotenv import load_dotenv -from fastapi import Request -from fastapi.routing import APIRoute - -load_dotenv() -import io -import time - -# this file is to test litellm/proxy - -import asyncio -import datetime -import json -import logging -from typing import Optional -import pytest - -import litellm -from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload -from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload - - -@pytest.mark.parametrize( - "model_id", - ["chatcmpl-9XZmkzS1uPhRCoVdGQvBqqIbSgECt", "", None], -) -def test_spend_logs_payload(model_id: Optional[str]): - """ - Ensure only expected values are logged in spend logs payload. - """ - - input_args: dict = { - "kwargs": { - "model": "chatgpt-v-3", - "messages": [ - {"role": "system", "content": "you are a helpful assistant.\n"}, - {"role": "user", "content": "bom dia"}, - ], - "custom_llm_provider": "azure", - "optional_params": { - "stream": False, - "max_tokens": 10, - "user": "116544810872468347480", - "extra_body": {}, - }, - "litellm_params": { - "acompletion": True, - "api_key": "sk-test-mock-key-707", - "force_timeout": 600, - "logger_fn": None, - "verbose": False, - "custom_llm_provider": "azure", - "api_base": "https://openai-gpt-4-test-v-1.openai.azure.com//openai/", - "litellm_call_id": "b9929bf6-7b80-4c8c-b486-034e6ac0c8b7", - "model_alias_map": {}, - "completion_call_id": None, - "metadata": { - "tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], - "user_api_key": "sk-test-mock-api-key-123", - "user_api_key_alias": "custom-key-alias", - "user_api_end_user_max_budget": None, - "litellm_api_version": "0.0.0", - "global_max_parallel_requests": None, - "user_api_key_user_id": "116544810872468347480", - "user_api_key_org_id": "custom-org-id", - "user_api_key_team_id": "custom-team-id", - "user_api_key_team_alias": "custom-team-alias", - "user_api_key_metadata": {}, - "requester_ip_address": "127.0.0.1", - "spend_logs_metadata": {"hello": "world"}, - "headers": { - "content-type": "application/json", - "user-agent": "PostmanRuntime/7.32.3", - "accept": "*/*", - "postman-token": "92300061-eeaa-423b-a420-0b44896ecdc4", - "host": "localhost:4000", - "accept-encoding": "gzip, deflate, br", - "connection": "keep-alive", - "content-length": "163", - }, - "endpoint": "http://localhost:4000/chat/completions", - "model_group": "gpt-5-mini", - "deployment": "azure/gpt-4.1-mini", - "model_info": { - "id": "4bad40a1eb6bebd1682800f16f44b9f06c52a6703444c99c7f9f32e9de3693b4", - "db_model": False, - }, - "api_base": "https://openai-gpt-4-test-v-1.openai.azure.com/", - "caching_groups": None, - "error_information": None, - "status": "success", - "proxy_server_request": "{}", - "raw_request": "\n\nPOST Request Sent from LiteLLM:\ncurl -X POST \\\nhttps://openai-gpt-4-test-v-1.openai.azure.com//openai/ \\\n-H 'Authorization: *****' \\\n-d '{'model': 'chatgpt-v-3', 'messages': [{'role': 'system', 'content': 'you are a helpful assistant.\\n'}, {'role': 'user', 'content': 'bom dia'}], 'stream': False, 'max_tokens': 10, 'user': '116544810872468347480', 'extra_body': {}}'\n", - }, - "model_info": { - "id": "4bad40a1eb6bebd1682800f16f44b9f06c52a6703444c99c7f9f32e9de3693b4", - "db_model": False, - }, - "proxy_server_request": { - "url": "http://localhost:4000/chat/completions", - "method": "POST", - "headers": { - "content-type": "application/json", - "user-agent": "PostmanRuntime/7.32.3", - "accept": "*/*", - "postman-token": "92300061-eeaa-423b-a420-0b44896ecdc4", - "host": "localhost:4000", - "accept-encoding": "gzip, deflate, br", - "connection": "keep-alive", - "content-length": "163", - }, - "body": { - "messages": [ - { - "role": "system", - "content": "you are a helpful assistant.\n", - }, - {"role": "user", "content": "bom dia"}, - ], - "model": "gpt-5-mini", - "max_tokens": 10, - }, - }, - "preset_cache_key": None, - "no-log": False, - "stream_response": {}, - "input_cost_per_token": None, - "input_cost_per_second": None, - "output_cost_per_token": None, - "output_cost_per_second": None, - }, - "start_time": datetime.datetime(2024, 6, 7, 12, 43, 30, 307665), - "stream": False, - "user": "116544810872468347480", - "call_type": "acompletion", - "litellm_call_id": "b9929bf6-7b80-4c8c-b486-034e6ac0c8b7", - "completion_start_time": datetime.datetime(2024, 6, 7, 12, 43, 30, 954146), - "max_tokens": 10, - "extra_body": {}, - "input": [ - {"role": "system", "content": "you are a helpful assistant.\n"}, - {"role": "user", "content": "bom dia"}, - ], - "api_key": "1234", - "original_response": "", - "additional_args": { - "headers": {"Authorization": "Bearer 1234"}, - "api_base": "openai-gpt-4-test-v-1.openai.azure.com", - "acompletion": True, - "complete_input_dict": { - "model": "chatgpt-v-3", - "messages": [ - {"role": "system", "content": "you are a helpful assistant.\n"}, - {"role": "user", "content": "bom dia"}, - ], - "stream": False, - "max_tokens": 10, - "user": "116544810872468347480", - "extra_body": {}, - }, - }, - "log_event_type": "post_api_call", - "end_time": datetime.datetime(2024, 6, 7, 12, 43, 30, 954146), - "cache_hit": None, - "response_cost": 2.4999999999999998e-05, - "standard_logging_object": { - "request_tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], - "metadata": { - "user_api_key_end_user_id": "test-user", - }, - "model_map_information": { - "tpm": 1000, - "rpm": 1000, - }, - }, - }, - "response_obj": litellm.ModelResponse( - id=model_id, - choices=[ - litellm.Choices( - finish_reason="length", - index=0, - message=litellm.Message( - content="Bom dia! Como posso ajudar você", role="assistant" - ), - ) - ], - created=1717789410, - model="gpt-35-turbo", - object="chat.completion", - system_fingerprint=None, - usage=litellm.Usage( - completion_tokens=10, prompt_tokens=20, total_tokens=30 - ), - ), - "start_time": datetime.datetime(2024, 6, 7, 12, 43, 30, 308604), - "end_time": datetime.datetime(2024, 6, 7, 12, 43, 30, 954146), - } - - payload: SpendLogsPayload = get_logging_payload(**input_args) - - assert len(payload["request_id"]) > 0 - # Define the expected metadata keys - expected_metadata_keys = SpendLogsMetadata.__annotations__.keys() - - # Validate only specified metadata keys are logged - assert "metadata" in payload - assert isinstance(payload["metadata"], str) - payload["metadata"] = json.loads(payload["metadata"]) - assert set(payload["metadata"].keys()) == set(expected_metadata_keys) - - # This is crucial - used in PROD, it should pass, related issue: https://github.com/BerriAI/litellm/issues/4334 - assert ( - payload["request_tags"] == '["model-anthropic-claude-v2.1", "app-ishaan-prod"]' - ) - assert payload["metadata"]["user_api_key_org_id"] == "custom-org-id" - assert payload["metadata"]["user_api_key_team_id"] == "custom-team-id" - assert payload["metadata"]["user_api_key_team_alias"] == "custom-team-alias" - assert payload["metadata"]["user_api_key_alias"] == "custom-key-alias" - - assert payload["custom_llm_provider"] == "azure" - - -def test_spend_logs_payload_whisper(): - """ - Ensure we can write /transcription request/responses to spend logs - """ - - kwargs: dict = { - "model": "whisper-1", - "messages": [{"role": "user", "content": "audio_file"}], - "optional_params": {}, - "litellm_params": { - "api_base": "", - "metadata": { - "user_api_key": "sk-test-mock-api-key-123", - "user_api_key_alias": None, - "user_api_key_end_user_id": "test-user", - "user_api_end_user_max_budget": None, - "litellm_api_version": "1.40.19", - "global_max_parallel_requests": None, - "user_api_key_user_id": "default_user_id", - "user_api_key_org_id": None, - "user_api_key_team_id": None, - "user_api_key_team_alias": None, - "user_api_key_team_max_budget": None, - "user_api_key_team_spend": None, - "user_api_key_spend": 0.0, - "user_api_key_max_budget": None, - "user_api_key_metadata": {}, - "headers": { - "host": "localhost:4000", - "user-agent": "curl/7.88.1", - "accept": "*/*", - "content-length": "775501", - "content-type": "multipart/form-data; boundary=------------------------21d518e191326d20", - }, - "endpoint": "http://localhost:4000/v1/audio/transcriptions", - "litellm_parent_otel_span": None, - "model_group": "whisper-1", - "deployment": "whisper-1", - "model_info": { - "id": "d7761582311451c34d83d65bc8520ce5c1537ea9ef2bec13383cf77596d49eeb", - "db_model": False, - }, - "caching_groups": None, - }, - }, - "start_time": datetime.datetime(2024, 6, 26, 14, 20, 11, 313291), - "stream": False, - "user": "", - "call_type": "atranscription", - "litellm_call_id": "05921cf7-33f9-421c-aad9-33310c1e2702", - "completion_start_time": datetime.datetime(2024, 6, 26, 14, 20, 13, 653149), - "stream_options": None, - "input": "tmp-requestc8640aee-7d85-49c3-b3ef-bdc9255d8e37.wav", - "original_response": '{"text": "Four score and seven years ago, our fathers brought forth on this continent a new nation, conceived in liberty and dedicated to the proposition that all men are created equal. Now we are engaged in a great civil war, testing whether that nation, or any nation so conceived and so dedicated, can long endure."}', - "additional_args": { - "complete_input_dict": { - "model": "whisper-1", - "file": "<_io.BufferedReader name='tmp-requestc8640aee-7d85-49c3-b3ef-bdc9255d8e37.wav'>", - "language": None, - "prompt": None, - "response_format": None, - "temperature": None, - } - }, - "log_event_type": "post_api_call", - "end_time": datetime.datetime(2024, 6, 26, 14, 20, 13, 653149), - "cache_hit": None, - "response_cost": 0.00023398580000000003, - } - - response = litellm.utils.TranscriptionResponse( - text="Four score and seven years ago, our fathers brought forth on this continent a new nation, conceived in liberty and dedicated to the proposition that all men are created equal. Now we are engaged in a great civil war, testing whether that nation, or any nation so conceived and so dedicated, can long endure." - ) - - payload: SpendLogsPayload = get_logging_payload( - kwargs=kwargs, - response_obj=response, - start_time=datetime.datetime.now(), - end_time=datetime.datetime.now(), - ) - - print("payload: ", payload) - - assert payload["call_type"] == "atranscription" - assert payload["spend"] == 0.00023398580000000003 - - -def test_spend_logs_payload_with_prompts_enabled(monkeypatch): - """ - Test that messages and responses are logged in spend logs when store_prompts_in_spend_logs is enabled - """ - # Mock general_settings - from litellm.proxy.proxy_server import general_settings - - general_settings["store_prompts_in_spend_logs"] = True - - input_args: dict = { - "kwargs": { - "model": "gpt-5-mini", - "messages": [{"role": "user", "content": "Hello!"}], - "litellm_params": { - "metadata": { - "user_api_key": "fake_key", - } - }, - }, - "response_obj": litellm.ModelResponse( - id="chatcmpl-123", - choices=[ - litellm.Choices( - finish_reason="stop", - index=0, - message=litellm.Message(content="Hi there!", role="assistant"), - ) - ], - model="gpt-5-mini", - usage=litellm.Usage(completion_tokens=2, prompt_tokens=1, total_tokens=3), - ), - "start_time": datetime.datetime.now(), - "end_time": datetime.datetime.now(), - } - - # Create a standard logging payload - standard_logging_payload = { - "messages": [{"role": "user", "content": "Hello!"}], - "response": {"role": "assistant", "content": "Hi there!"}, - "metadata": { - "user_api_key_end_user_id": "test-user", - }, - "request_tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], - "model_map_information": { - "tpm": 1000, - "rpm": 1000, - }, - } - litellm_params = { - "proxy_server_request": { - "body": { - "model": "gpt-5.5", - "messages": [{"role": "user", "content": "Hello!"}], - } - } - } - input_args["kwargs"]["standard_logging_object"] = standard_logging_payload - input_args["kwargs"]["litellm_params"] = litellm_params - - payload: SpendLogsPayload = get_logging_payload(**input_args) - - print("json payload: ", json.dumps(payload, indent=4, default=str)) - - # Verify messages and response are included in payload - assert payload["response"] == json.dumps( - {"role": "assistant", "content": "Hi there!"} - ) - proxy_server_request = json.loads(payload["proxy_server_request"] or "{}") - assert proxy_server_request["model"] == "gpt-5.5" - assert proxy_server_request["messages"] == [{"role": "user", "content": "Hello!"}] - - # Clean up - reset general_settings - general_settings["store_prompts_in_spend_logs"] = False - - # Verify messages and response are not included when disabled - payload_disabled: SpendLogsPayload = get_logging_payload(**input_args) - assert payload_disabled["messages"] == "{}" - assert payload_disabled["response"] == "{}" diff --git a/tests/pass_through_unit_tests/test_claude_code_marketplace.py b/tests/pass_through_unit_tests/test_claude_code_marketplace.py deleted file mode 100644 index 2747fbbfee3..00000000000 --- a/tests/pass_through_unit_tests/test_claude_code_marketplace.py +++ /dev/null @@ -1,289 +0,0 @@ -""" -Tests for Claude Code Marketplace endpoints. - -Tests: -1. Register a plugin -2. Get marketplace.json (list enabled plugins) -""" - -import json -import time -from datetime import datetime, timezone -from unittest.mock import AsyncMock, MagicMock - -import pytest - - -import litellm -from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.proxy_server import LitellmUserRoles -from litellm.caching.caching import DualCache -from litellm.types.proxy.claude_code_endpoints import RegisterPluginRequest - -# Import the functions we're testing -from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketplace import ( - register_plugin, - get_marketplace, -) -from tests._master_key import MASTER_KEY - - -class MockPluginRecord: - """Mock plugin record that mimics Prisma model behavior.""" - - def __init__( - self, name, version, description, manifest_json, enabled=True, created_by=None - ): - self.id = f"plugin-{name}-{int(time.time())}" - self.name = name - self.version = version - self.description = description - self.manifest_json = manifest_json - self.files_json = "{}" - self.enabled = enabled - self.created_at = datetime.now(timezone.utc) - self.updated_at = datetime.now(timezone.utc) - self.created_by = created_by - - -@pytest.fixture -def mock_prisma_client(): - """Create a mock PrismaClient that doesn't require Prisma binaries.""" - # In-memory storage for plugins - plugins_store = {} - - # Create mock client - mock_client = MagicMock() - mock_client.proxy_logging_obj = MagicMock() - - # Mock the db attribute - mock_client.db = MagicMock() - - # Mock the plugin table with async methods - mock_table = MagicMock() - - async def find_unique(where): - """Mock find_unique - returns plugin if exists, None otherwise.""" - plugin_name = where.get("name") - return plugins_store.get(plugin_name) - - async def find_many(where=None): - """Mock find_many - returns list of plugins matching where clause.""" - if where is None or where == {}: - return list(plugins_store.values()) - enabled = where.get("enabled") - if enabled is not None: - return [p for p in plugins_store.values() if p.enabled == enabled] - return list(plugins_store.values()) - - async def create(data): - """Mock create - creates a new plugin.""" - plugin_name = data["name"] - manifest = data.get("manifest_json", "{}") - plugin = MockPluginRecord( - name=plugin_name, - version=data.get("version"), - description=data.get("description"), - manifest_json=manifest, - enabled=data.get("enabled", True), - created_by=data.get("created_by"), - ) - plugins_store[plugin_name] = plugin - return plugin - - async def update(where, data): - """Mock update - updates an existing plugin.""" - plugin_name = where.get("name") - if plugin_name not in plugins_store: - raise ValueError(f"Plugin {plugin_name} not found") - plugin = plugins_store[plugin_name] - # Update fields - if "version" in data: - plugin.version = data["version"] - if "description" in data: - plugin.description = data["description"] - if "manifest_json" in data: - plugin.manifest_json = data["manifest_json"] - if "enabled" in data: - plugin.enabled = data["enabled"] - if "updated_at" in data: - plugin.updated_at = data["updated_at"] - return plugin - - async def delete(where): - """Mock delete - deletes a plugin.""" - plugin_name = where.get("name") - if plugin_name in plugins_store: - del plugins_store[plugin_name] - return None - - async def connect(): - """Mock connect - no-op.""" - pass - - # Set up async mocks - mock_table.find_unique = AsyncMock(side_effect=find_unique) - mock_table.find_many = AsyncMock(side_effect=find_many) - mock_table.create = AsyncMock(side_effect=create) - mock_table.update = AsyncMock(side_effect=update) - mock_table.delete = AsyncMock(side_effect=delete) - - mock_client.db.litellm_claudecodeplugintable = mock_table - mock_client.connect = AsyncMock(side_effect=connect) - - # Store plugins_store on the mock for cleanup if needed - mock_client._plugins_store = plugins_store - - return mock_client - - -@pytest.mark.asyncio -async def test_register_plugin(mock_prisma_client): - """Test registering a plugin in the marketplace.""" - setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client) - setattr(litellm.proxy.proxy_server, "master_key", MASTER_KEY) - - await litellm.proxy.proxy_server.prisma_client.connect() - - # Create a unique plugin name for this test - plugin_name = f"test-plugin-{int(time.time())}" - - request = RegisterPluginRequest( - name=plugin_name, - source={"source": "github", "repo": "test-org/test-repo"}, - version="1.0.0", - description="Test plugin for unit tests", - ) - - user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, - api_key=MASTER_KEY, - user_id="test-user", - ) - - response = await register_plugin( - request=request, - user_api_key_dict=user_api_key_dict, - ) - - assert response.status == "success" - assert response.action == "created" - assert response.plugin.name == plugin_name - assert response.plugin.version == "1.0.0" - assert response.plugin.enabled is True - - # Verify the plugin was stored in the mock - stored_plugin = ( - await mock_prisma_client.db.litellm_claudecodeplugintable.find_unique( - where={"name": plugin_name} - ) - ) - assert stored_plugin is not None - assert stored_plugin.name == plugin_name - - # Cleanup - delete the plugin - await mock_prisma_client.db.litellm_claudecodeplugintable.delete( - where={"name": plugin_name} - ) - - -@pytest.mark.asyncio -async def test_get_marketplace(mock_prisma_client): - """Test getting marketplace.json with registered plugins.""" - setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client) - setattr(litellm.proxy.proxy_server, "master_key", MASTER_KEY) - - await litellm.proxy.proxy_server.prisma_client.connect() - - # First register a plugin - plugin_name = f"test-marketplace-plugin-{int(time.time())}" - - request = RegisterPluginRequest( - name=plugin_name, - source={"source": "github", "repo": "test-org/marketplace-test"}, - version="2.0.0", - description="Test plugin for marketplace test", - ) - - user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, - api_key=MASTER_KEY, - user_id="test-user", - ) - - await register_plugin( - request=request, - user_api_key_dict=user_api_key_dict, - ) - - # Now get the marketplace - response = await get_marketplace(request=MagicMock()) - - # Response is a JSONResponse, get the body - body = json.loads(response.body.decode()) - - assert body["name"] == "litellm" - assert "plugins" in body - - # Find our plugin in the list - our_plugin = next((p for p in body["plugins"] if p["name"] == plugin_name), None) - assert our_plugin is not None - assert our_plugin["source"] == { - "source": "github", - "repo": "test-org/marketplace-test", - } - assert our_plugin["version"] == "2.0.0" - - # Cleanup - await mock_prisma_client.db.litellm_claudecodeplugintable.delete( - where={"name": plugin_name} - ) - - -@pytest.mark.asyncio -async def test_register_plugin_git_subdir(mock_prisma_client): - """Test registering a plugin with git-subdir source type.""" - setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client) - setattr(litellm.proxy.proxy_server, "master_key", MASTER_KEY) - - await litellm.proxy.proxy_server.prisma_client.connect() - - plugin_name = f"test-subdir-plugin-{int(time.time())}" - - request = RegisterPluginRequest( - name=plugin_name, - source={ - "source": "git-subdir", - "url": "https://github.com/test-org/monorepo.git", - "path": "plugins/my-plugin", - }, - version="1.0.0", - description="Test git-subdir plugin", - ) - - user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, - api_key=MASTER_KEY, - user_id="test-user", - ) - - response = await register_plugin( - request=request, - user_api_key_dict=user_api_key_dict, - ) - - assert response.status == "success" - assert response.action == "created" - assert response.plugin.name == plugin_name - assert response.plugin.source["source"] == "git-subdir" - assert ( - response.plugin.source["url"] - == "https://github.com/test-org/monorepo.git" - ) - assert response.plugin.source["path"] == "plugins/my-plugin" - assert response.plugin.enabled is True - - # Cleanup - await mock_prisma_client.db.litellm_claudecodeplugintable.delete( - where={"name": plugin_name} - ) diff --git a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py deleted file mode 100644 index 7a3109839d8..00000000000 --- a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py +++ /dev/null @@ -1,561 +0,0 @@ -import json -from datetime import datetime -from unittest.mock import AsyncMock, Mock, patch, MagicMock -from typing import Optional - - -import fastapi -from fastapi import FastAPI -import httpx -import pytest -import litellm -from typing import AsyncGenerator -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType -from litellm.proxy.pass_through_endpoints.success_handler import ( - PassThroughEndpointLogging, -) -from litellm.proxy.pass_through_endpoints.streaming_handler import ( - PassThroughStreamingHandler, -) - -from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( - pass_through_request, -) -from fastapi import Request -from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup -from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( - HttpPassThroughEndpointHelpers, -) -from litellm.types.passthrough_endpoints.pass_through_endpoints import ( - PassthroughStandardLoggingPayload, -) - - -@pytest.fixture -def mock_request(): - # Create a mock request with headers - class QueryParams: - def __init__(self): - self._dict = {} - - def __iter__(self): - return iter(self._dict.items()) - - def items(self): - return self._dict.items() - - def keys(self): - return self._dict.keys() - - def values(self): - return self._dict.values() - - class MockRequest: - def __init__( - self, headers=None, method="POST", request_body: Optional[dict] = None - ): - self.headers = headers or {} - self.query_params = QueryParams() - self.method = method - self.request_body = request_body or {} - # Add url attribute that the actual code expects - self.url = httpx.URL("http://localhost:8000/test") - self.scope = {"type": "http", "method": method, "path": "/test"} - # Add state attribute that FastAPI requests have - self.state = type("State", (), {})() - - async def body(self) -> bytes: - return bytes(json.dumps(self.request_body), "utf-8") - - return MockRequest - - -@pytest.fixture -def mock_user_api_key_dict(): - return UserAPIKeyAuth( - api_key="test-key", - user_id="test-user", - team_id="test-team", - end_user_id="test-user", - ) - - -def test_init_kwargs_for_pass_through_endpoint_basic( - mock_request, mock_user_api_key_dict -): - """ - Basic test for init_kwargs_for_pass_through_endpoint - - - metadata should contain user_api_key, user_api_key_user_id, user_api_key_team_id, user_api_key_end_user_id from `mock_user_api_key_dict` - """ - request = mock_request() - passthrough_payload = PassthroughStandardLoggingPayload( - url="https://test.com", - request_body={}, - ) - - result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( - request=request, - user_api_key_dict=mock_user_api_key_dict, - passthrough_logging_payload=passthrough_payload, - litellm_call_id="test-call-id", - logging_obj=LiteLLMLoggingObj( - model="test-model", - messages=[], - stream=False, - call_type="test-call-type", - start_time=datetime.now(), - litellm_call_id="test-call-id", - function_id="test-function-id", - ), - ) - - assert result["call_type"] == "pass_through_endpoint" - assert result["litellm_call_id"] == "test-call-id" - assert result["passthrough_logging_payload"] == passthrough_payload - - ######################################################### - # Check metadata - ######################################################### - assert result["litellm_params"]["metadata"]["user_api_key"] == "test-key" - assert result["litellm_params"]["metadata"]["user_api_key_hash"] == "test-key" - assert result["litellm_params"]["metadata"]["user_api_key_alias"] is None - assert result["litellm_params"]["metadata"]["user_api_key_user_email"] is None - assert result["litellm_params"]["metadata"]["user_api_key_user_id"] == "test-user" - assert result["litellm_params"]["metadata"]["user_api_key_team_id"] == "test-team" - assert result["litellm_params"]["metadata"]["user_api_key_org_id"] is None - assert result["litellm_params"]["metadata"]["user_api_key_team_alias"] is None - assert ( - result["litellm_params"]["metadata"]["user_api_key_end_user_id"] == "test-user" - ) - assert result["litellm_params"]["metadata"]["user_api_key_request_route"] is None - - -def test_init_kwargs_with_litellm_metadata(mock_request, mock_user_api_key_dict): - """ - Expected behavior: litellm_metadata should be merged with default metadata - - see usage example here: https://docs.litellm.ai/docs/pass_through/anthropic_completion#send-litellm_metadata-tags - """ - request = mock_request() - parsed_body = { - "litellm_metadata": {"custom_field": "custom_value", "tags": ["tag1", "tag2"]} - } - passthrough_payload = PassthroughStandardLoggingPayload( - url="https://test.com", - request_body={}, - ) - - result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( - request=request, - user_api_key_dict=mock_user_api_key_dict, - passthrough_logging_payload=passthrough_payload, - _parsed_body=parsed_body, - litellm_call_id="test-call-id", - logging_obj=LiteLLMLoggingObj( - model="test-model", - messages=[], - stream=False, - call_type="test-call-type", - start_time=datetime.now(), - litellm_call_id="test-call-id", - function_id="test-function-id", - ), - ) - - # Check that litellm_metadata was merged with default metadata - metadata = result["litellm_params"]["metadata"] - print("metadata", metadata) - assert metadata["custom_field"] == "custom_value" - assert metadata["tags"] == ["tag1", "tag2"] - assert metadata["user_api_key"] == "test-key" - - -def test_init_kwargs_with_tags_in_header(mock_request, mock_user_api_key_dict): - """ - Tags should be added to metadata if they exist in headers - """ - request = mock_request(headers={"tags": "tag1,tag2"}) - passthrough_payload = PassthroughStandardLoggingPayload( - url="https://test.com", - request_body={}, - ) - - result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( - request=request, - user_api_key_dict=mock_user_api_key_dict, - passthrough_logging_payload=passthrough_payload, - litellm_call_id="test-call-id", - logging_obj=LiteLLMLoggingObj( - model="test-model", - messages=[], - stream=False, - call_type="test-call-type", - start_time=datetime.now(), - litellm_call_id="test-call-id", - function_id="test-function-id", - ), - ) - - # Check that tags were added to metadata - metadata = result["litellm_params"]["metadata"] - print("metadata", metadata) - assert metadata["tags"] == ["tag1", "tag2"] - - -athropic_request_body = { - "model": "claude-sonnet-4-5-20250929", - "max_tokens": 256, - "messages": [{"role": "user", "content": "Hello, world tell me 2 sentences "}], - "litellm_metadata": {"tags": ["hi", "hello"]}, -} - - -@pytest.mark.asyncio -async def test_pass_through_request_logging_failure( - mock_request, mock_user_api_key_dict -): - """ - Test that pass_through_request still returns a response even if logging raises an Exception - """ - - # Mock the logging handler to raise an error - async def mock_logging_failure(*args, **kwargs): - raise Exception("Logging failed!") - - # Create a mock response - mock_response = AsyncMock() - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - - # Add mock content - mock_response._content = b'{"mock": "response"}' - - async def mock_aread(): - return mock_response._content - - mock_response.aread = mock_aread - - # Patch both the logging handler and the httpx client - with ( - patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughEndpointLogging.pass_through_async_success_handler", - new=mock_logging_failure, - ), - patch( - "httpx.AsyncClient.send", - return_value=mock_response, - ), - patch( - "httpx.AsyncClient.request", - return_value=mock_response, - ), - ): - request = mock_request( - headers={}, method="POST", request_body=athropic_request_body - ) - response = await pass_through_request( - request=request, - target="https://exampleopenaiendpoint-production.up.railway.app/v1/messages", - custom_headers={}, - user_api_key_dict=mock_user_api_key_dict, - ) - - # Assert response was returned successfully despite logging failure - assert response.status_code == 200 - - # Verify we got the mock response content - # For FastAPI Response objects, content is accessed via the body attribute - assert response.body == b'{"mock": "response"}' - - -@pytest.mark.asyncio -async def test_pass_through_request_logging_failure_with_stream( - mock_request, mock_user_api_key_dict -): - """ - Test that pass_through_request still returns a response even if logging raises an Exception - """ - - # Mock the logging handler to raise an error - async def mock_logging_failure(*args, **kwargs): - raise Exception("Logging failed!") - - # Create a mock response - mock_response = AsyncMock() - mock_response.status_code = 200 - - # Add headers property to mock response - mock_response.headers = { - "content-type": "application/json", # Not streaming - } - - # Create mock chunks for streaming - mock_chunks = [b'{"chunk": 1}', b'{"chunk": 2}'] - mock_response.body_iterator = AsyncMock() - mock_response.body_iterator.__aiter__.return_value = mock_chunks - - # Add aread method to mock response - mock_response._content = b'{"mock": "response"}' - - async def mock_aread(): - return mock_response._content - - mock_response.aread = mock_aread - - # Patch both the logging handler and the httpx client - with ( - patch( - "litellm.proxy.pass_through_endpoints.streaming_handler.PassThroughStreamingHandler.route_streaming_logging_to_handler", - new=mock_logging_failure, - ), - patch( - "httpx.AsyncClient.send", - return_value=mock_response, - ), - patch( - "httpx.AsyncClient.request", - return_value=mock_response, - ), - ): - request = mock_request( - headers={}, method="POST", request_body=athropic_request_body - ) - response = await pass_through_request( - request=request, - target="https://exampleopenaiendpoint-production.up.railway.app/v1/messages", - custom_headers={}, - user_api_key_dict=mock_user_api_key_dict, - ) - - # Assert response was returned successfully despite logging failure - assert response.status_code == 200 - - # Check if it's a streaming response or regular response - from fastapi.responses import StreamingResponse - - if isinstance(response, StreamingResponse): - # For streaming responses in tests, we just verify it's the right type - # and status code since iterating over it is complex in test context - assert response.status_code == 200 - else: - # Non-streaming response - should have body attribute - assert hasattr(response, "body") - assert response.body == b'{"mock": "response"}' - - -def test_init_kwargs_filters_pricing_params(mock_request, mock_user_api_key_dict): - """ - Test that pricing parameters are properly filtered out from the request body - and don't get sent to the provider API. - - This ensures that custom pricing parameters like: - - cache_read_input_token_cost - - input_cost_per_token_batches - - output_cost_per_token_batches - - cache_creation_input_token_cost - etc. are removed from the request body before sending to provider. - - Regression test for: LIT-1221 - """ - request = mock_request() - - # Create a parsed body with pricing parameters that should be filtered out - parsed_body = { - "model": "gpt-5.5", - "messages": [{"role": "user", "content": "test"}], - # Standard pricing params (should be filtered) - "input_cost_per_token": 0.00002, - "output_cost_per_token": 0.00002, - "input_cost_per_second": 0.00001, - "output_cost_per_second": 0.00001, - # Cache-related pricing params (should be filtered) - "cache_read_input_token_cost": 0.00005, - "cache_creation_input_token_cost": 0.00003, - "cache_creation_input_token_cost_above_1hr": 0.00004, - # Batch pricing params (should be filtered) - "input_cost_per_token_batches": 0.00005, - "output_cost_per_token_batches": 0.00006, - # Other pricing params (should be filtered) - "input_cost_per_audio_token": 0.00001, - "output_cost_per_audio_token": 0.00001, - "input_cost_per_character": 0.000001, - "output_cost_per_character": 0.000001, - "input_cost_per_image": 0.001, - "output_cost_per_image": 0.001, - # Tiered pricing - "tiered_pricing": [{"input_cost_per_token": 0.00001}], - # This should NOT be filtered (it's a valid OpenAI parameter) - "temperature": 0.7, - "max_tokens": 100, - } - - passthrough_payload = PassthroughStandardLoggingPayload( - url="https://api.openai.com/v1/chat/completions", - request_body=parsed_body.copy(), - ) - - result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( - request=request, - user_api_key_dict=mock_user_api_key_dict, - passthrough_logging_payload=passthrough_payload, - _parsed_body=parsed_body, - litellm_call_id="test-call-id", - logging_obj=LiteLLMLoggingObj( - model="gpt-5.5", - messages=[{"role": "user", "content": "test"}], - stream=False, - call_type="completion", - start_time=datetime.now(), - litellm_call_id="test-call-id", - function_id="test-function-id", - ), - ) - - # Verify pricing parameters were filtered out from parsed_body - assert "input_cost_per_token" not in parsed_body - assert "output_cost_per_token" not in parsed_body - assert "input_cost_per_second" not in parsed_body - assert "output_cost_per_second" not in parsed_body - assert "cache_read_input_token_cost" not in parsed_body - assert "cache_creation_input_token_cost" not in parsed_body - assert "cache_creation_input_token_cost_above_1hr" not in parsed_body - assert "input_cost_per_token_batches" not in parsed_body - assert "output_cost_per_token_batches" not in parsed_body - assert "input_cost_per_audio_token" not in parsed_body - assert "output_cost_per_audio_token" not in parsed_body - assert "input_cost_per_character" not in parsed_body - assert "output_cost_per_character" not in parsed_body - assert "input_cost_per_image" not in parsed_body - assert "output_cost_per_image" not in parsed_body - assert "tiered_pricing" not in parsed_body - - # Verify valid OpenAI parameters remain in parsed_body - assert parsed_body["model"] == "gpt-5.5" - assert parsed_body["messages"] == [{"role": "user", "content": "test"}] - assert parsed_body["temperature"] == 0.7 - assert parsed_body["max_tokens"] == 100 - - # Verify pricing parameters are stored in litellm_params for internal use - litellm_params = result["litellm_params"] - assert litellm_params["input_cost_per_token"] == 0.00002 - assert litellm_params["output_cost_per_token"] == 0.00002 - # Note: Other pricing params are also stored but we test the key ones that caused the regression - - -def test_init_kwargs_client_metadata_cannot_spoof_authenticated_identity( - mock_request, mock_user_api_key_dict -): - request = mock_request() - passthrough_payload = PassthroughStandardLoggingPayload( - url="https://test.com", - request_body={}, - ) - authenticated_key = UserAPIKeyAuth( - api_key="test-key", - user_id="test-user", - team_id="test-team", - end_user_id="test-user", - key_alias="real-key", - team_alias="Real Team", - user_email="real@example.com", - org_id="real-org", - ) - - result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( - request=request, - user_api_key_dict=authenticated_key, - passthrough_logging_payload=passthrough_payload, - litellm_call_id="test-call-id", - logging_obj=LiteLLMLoggingObj( - model="test-model", - messages=[], - stream=False, - call_type="test-call-type", - start_time=datetime.now(), - litellm_call_id="test-call-id", - function_id="test-function-id", - ), - _parsed_body={ - "litellm_metadata": { - "user_api_key_org_id": "victim-org", - "user_api_key_end_user_id": "victim-end-user", - "user_api_key_user_id": "victim-user", - "user_api_key_team_id": "victim-team", - "user_api_key_team_alias": "Victim Team", - "user_api_key_alias": "victim-key", - "user_api_key_user_email": "victim@example.com", - } - }, - ) - - metadata = result["litellm_params"]["metadata"] - assert metadata["user_api_key_user_id"] == "test-user" - assert metadata["user_api_key_team_id"] == "test-team" - assert metadata["user_api_key_team_alias"] == "Real Team" - assert metadata["user_api_key_alias"] == "real-key" - assert metadata["user_api_key_user_email"] == "real@example.com" - assert metadata["user_api_key_org_id"] == "real-org" - assert metadata["user_api_key_end_user_id"] == "test-user" - - -def test_init_kwargs_no_authenticated_identity_field_is_client_settable( - mock_request, mock_user_api_key_dict -): - authenticated_key = UserAPIKeyAuth( - api_key="test-key", - user_id="test-user", - team_id="test-team", - end_user_id="test-end-user", - key_alias="real-key", - team_alias="Real Team", - user_email="real@example.com", - org_id="real-org", - organization_alias="Real Org", - project_id="real-project", - project_alias="Real Project", - spend=1.5, - max_budget=10.0, - user_spend=2.5, - user_max_budget=20.0, - team_spend=3.5, - team_max_budget=30.0, - metadata={"real": "auth-metadata"}, - ) - expected = dict( - LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( - user_api_key_dict=authenticated_key - ) - ) - assert len(expected) >= 20 - - spoofed = {key: f"SPOOFED-{key}" for key in expected} - - result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( - request=mock_request(), - user_api_key_dict=authenticated_key, - passthrough_logging_payload=PassthroughStandardLoggingPayload( - url="https://test.com", request_body={} - ), - litellm_call_id="test-call-id", - logging_obj=LiteLLMLoggingObj( - model="test-model", - messages=[], - stream=False, - call_type="test-call-type", - start_time=datetime.now(), - litellm_call_id="test-call-id", - function_id="test-function-id", - ), - _parsed_body={"litellm_metadata": dict(spoofed), "metadata": dict(spoofed)}, - ) - - metadata = result["litellm_params"]["metadata"] - survived = { - key: metadata.get(key) - for key in expected - if metadata.get(key) != expected[key] - } - assert survived == {}, f"client-supplied values survived for: {sorted(survived)}" diff --git a/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py b/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py deleted file mode 100644 index 80a7fb81fb2..00000000000 --- a/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py +++ /dev/null @@ -1,44 +0,0 @@ -import pytest - - -@pytest.mark.parametrize( - "end_user_id", - [{"litellm_metadata": {"user": "test"}}, {"metadata": {"user_id": "test"}}], -) -def test_get_user_from_metadata(end_user_id): - from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( - AnthropicPassthroughLoggingHandler, - PassthroughStandardLoggingPayload, - ) - - passthrough_logging_payload = PassthroughStandardLoggingPayload( - url="https://api.anthropic.com/v1/messages", - request_body={**end_user_id}, - response_body={ - "id": "msg_015uSaCZBvu9gUSkAmZtMfxC", - "type": "message", - "role": "assistant", - "model": "claude-sonnet-4-5-20250929", - "content": [ - { - "type": "text", - "text": "Now I'll click on the Firefox icon to launch it.", - }, - { - "type": "tool_use", - "id": "toolu_01TQsF5p7Pf4LGKyLUDDySVr", - "name": "computer", - "input": {"action": "mouse_move", "coordinate": [24, 36]}, - }, - ], - "stop_reason": "tool_use", - "stop_sequence": None, - "usage": {"input_tokens": 2202, "output_tokens": 89}, - }, - ) - - response = AnthropicPassthroughLoggingHandler._get_user_from_metadata( - passthrough_logging_payload=passthrough_logging_payload - ) - - assert response == "test" diff --git a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py deleted file mode 100644 index 7b1de9fc64b..00000000000 --- a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py +++ /dev/null @@ -1,99 +0,0 @@ -""" -Test Vertex AI Live API Passthrough Feature - -This module tests the Vertex AI Live API WebSocket passthrough functionality, -including the logging handler, cost tracking, and WebSocket message processing. -""" - -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest - -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.proxy._types import UserAPIKeyAuth - - -class TestVertexAILivePassthroughIntegration: - """Integration tests for Vertex AI Live passthrough functionality""" - - @pytest.fixture - def mock_websocket(self): - """Create a mock WebSocket for testing""" - websocket = AsyncMock() - websocket.headers = {"authorization": "Bearer test-token"} - websocket.client_state = MagicMock() - websocket.client_state.DISCONNECTED = "disconnected" - return websocket - - @pytest.fixture - def mock_user_api_key(self): - """Create a mock user API key""" - return UserAPIKeyAuth( - api_key="test-key", - user_id="test-user", - team_id="test-team", - user_role="customer", - ) - - @pytest.fixture - def mock_logging_obj(self): - """Create a mock logging object""" - mock = MagicMock(spec=LiteLLMLoggingObj) - mock.model_call_details = {} - mock.response_cost_calculator.return_value = None - return mock - - @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.websocket_passthrough_request") - @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router") - @patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.vertex_llm_base._ensure_access_token_async") - @patch("litellm.proxy.proxy_server.proxy_logging_obj") - @pytest.mark.asyncio - async def test_vertex_ai_live_websocket_passthrough_route( - self, - mock_proxy_logging_obj, - mock_ensure_access_token, - mock_router, - mock_websocket_passthrough, - mock_websocket, - mock_user_api_key, - mock_logging_obj, - ): - """Test the Vertex AI Live WebSocket passthrough route""" - from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - vertex_ai_live_websocket_passthrough, - ) - - # Mock the router methods - mock_router.get_vertex_credentials.return_value = MagicMock( - vertex_project="test-project", - vertex_location="us-central1", - vertex_credentials="test-credentials", - ) - mock_router.set_default_vertex_config.return_value = None - - # Mock the access token async call - mock_ensure_access_token.return_value = ("test-access-token", "test-project") - - # Mock the WebSocket passthrough request - it returns None, not an AsyncMock - mock_websocket_passthrough.return_value = None - - # Test the route - result = await vertex_ai_live_websocket_passthrough( - websocket=mock_websocket, user_api_key_dict=mock_user_api_key - ) - - # Verify that the WebSocket passthrough was called - mock_websocket_passthrough.assert_called_once() - - # Check the call arguments - call_args = mock_websocket_passthrough.call_args - assert call_args[1]["websocket"] == mock_websocket - assert call_args[1]["user_api_key_dict"] == mock_user_api_key - assert call_args[1]["endpoint"] == "/vertex_ai/live" - - # The result should be None since websocket_passthrough_request returns None - assert result is None - - -if __name__ == "__main__": - pytest.main([__file__]) diff --git a/tests/unit/enterprise/enterprise_callbacks/test_llm_guard.py b/tests/unit/enterprise/enterprise_callbacks/test_llm_guard.py index fd9c5496917..daab841139a 100644 --- a/tests/unit/enterprise/enterprise_callbacks/test_llm_guard.py +++ b/tests/unit/enterprise/enterprise_callbacks/test_llm_guard.py @@ -166,3 +166,74 @@ async def test_llm_guard_skips_unsupported_call_types( ) assert result is data assert data == {"messages": [{"role": "user", "content": "unchanged"}]} + + +@pytest.mark.asyncio +async def test_llm_guard_sanitizes_multimodal_and_input(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "llm_guard_mode", "all") + llm_guard: Final = _ENTERPRISE_LLMGuard( + mock_testing=True, + mock_redacted_text={ + "sanitized_prompt": "email: [REDACTED]", + "is_valid": True, + "scanners": {"Regex": 0.0}, + }, + ) + user_api_key_dict: Final = UserAPIKeyAuth(api_key=hash_token("sk-98765")) + + image_part: Final = {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}} + data: Final = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "email: person@example.com"}, + image_part, + ], + } + ] + } + result: Final = await llm_guard.async_moderation_hook( + data=data, user_api_key_dict=user_api_key_dict, call_type="completion" + ) + assert result["messages"][0]["content"][0]["text"] == "email: [REDACTED]" + assert result["messages"][0]["content"][1] == image_part + + input_data: Final = {"input": ["email: person@example.com", "another prompt"]} + input_result: Final = await llm_guard.async_moderation_hook( + data=input_data, user_api_key_dict=user_api_key_dict, call_type="embeddings" + ) + assert input_result["input"] == ["email: [REDACTED]", "email: [REDACTED]"] + + +def test_llm_guard_key_specific_mode(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "llm_guard_mode", "key-specific") + + llm_guard: Final = _ENTERPRISE_LLMGuard(mock_testing=True) + + _api_key: Final = "sk-98765" + user_api_key_dict: Final = UserAPIKeyAuth( + api_key=_api_key, + ) + + assert llm_guard.should_proceed(user_api_key_dict=user_api_key_dict, data={}) == False + + permitted_key: Final = UserAPIKeyAuth(api_key=_api_key, permissions={"enable_llm_guard_check": True}) + assert llm_guard.should_proceed(user_api_key_dict=permitted_key, data={}) == True + + +def test_llm_guard_request_specific_mode(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "llm_guard_mode", "request-specific") + + llm_guard: Final = _ENTERPRISE_LLMGuard(mock_testing=True) + + _api_key: Final = "sk-98765" + user_api_key_dict: Final = UserAPIKeyAuth( + api_key=_api_key, + ) + + assert llm_guard.should_proceed(user_api_key_dict=user_api_key_dict, data={}) == False + + permitted_key: Final = UserAPIKeyAuth(api_key=_api_key, permissions={"enable_llm_guard_check": True}) + permitted_request: Final = {"metadata": {"permissions": {"enable_llm_guard_check": True}}} + assert llm_guard.should_proceed(user_api_key_dict=permitted_key, data=permitted_request) == True diff --git a/tests/unit/enterprise/enterprise_callbacks/test_secret_detection.py b/tests/unit/enterprise/enterprise_callbacks/test_secret_detection.py index 3bdeb34c2fe..12c54210579 100644 --- a/tests/unit/enterprise/enterprise_callbacks/test_secret_detection.py +++ b/tests/unit/enterprise/enterprise_callbacks/test_secret_detection.py @@ -26,6 +26,7 @@ from litellm_enterprise.enterprise_callbacks.secret_detection import ( from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import hash_token AWS_KEY = "AKIAIOSFODNN7EXAMPLE" OPENAI_KEY = "sk-test-abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGH" @@ -989,3 +990,137 @@ async def test_legacy_nameless_instance_records_nothing(): assert data["messages"][0]["content"] == "use [REDACTED] for auth" assert "standard_logging_guardrail_information" not in data["metadata"] + + +@pytest.mark.asyncio +async def test_basic_secret_detection_chat(): + secret_instance: Final = _ENTERPRISE_SecretDetection() + user_api_key_dict: Final = UserAPIKeyAuth(api_key=hash_token("sk-98765")) + local_cache: Final = DualCache() + + test_data: Final = { + "messages": [ + { + "role": "user", + "content": "Hey, how's it going, API_KEY = 'sk_1234567890abcdef'", + }, + { + "role": "assistant", + "content": "Hello! I'm doing well. How can I assist you today?", + }, + { + "role": "user", + "content": "this is my OPENAI_API_KEY = 'sk_1234567890abcdef'", + }, + { + "role": "user", + "content": "My hi API Key is sk-Pc4nlxVoMz41290028TbMCxx, does it seem to be in the correct format?", + }, + {"role": "user", "content": "i think it is +1 412-555-5555"}, + ], + "model": "gpt-3.5-turbo", + } + + await secret_instance.async_pre_call_hook( + cache=local_cache, + data=test_data, + user_api_key_dict=user_api_key_dict, + call_type="completion", + ) + + assert test_data == { + "messages": [ + {"role": "user", "content": "Hey, how's it going, API_KEY = '[REDACTED]'"}, + { + "role": "assistant", + "content": "Hello! I'm doing well. How can I assist you today?", + }, + {"role": "user", "content": "this is my OPENAI_API_KEY = '[REDACTED]'"}, + { + "role": "user", + "content": "My hi API Key is [REDACTED], does it seem to be in the correct format?", + }, + {"role": "user", "content": "i think it is +1 412-555-5555"}, + ], + "model": "gpt-3.5-turbo", + }, "Expect all API Keys to be masked" + + +@pytest.mark.asyncio +async def test_basic_secret_detection_text_completion(): + secret_instance: Final = _ENTERPRISE_SecretDetection() + user_api_key_dict: Final = UserAPIKeyAuth(api_key=hash_token("sk-98765")) + local_cache: Final = DualCache() + + test_data: Final = { + "prompt": "Hey, how's it going, API_KEY = 'sk_1234567890abcdef', my OPENAI_API_KEY = 'sk_1234567890abcdef' and i want to know what is the weather", + "model": "gpt-3.5-turbo", + } + + await secret_instance.async_pre_call_hook( + cache=local_cache, + data=test_data, + user_api_key_dict=user_api_key_dict, + call_type="completion", + ) + + assert test_data == { + "prompt": "Hey, how's it going, API_KEY = '[REDACTED]', my OPENAI_API_KEY = '[REDACTED]' and i want to know what is the weather", + "model": "gpt-3.5-turbo", + } + + +@pytest.mark.asyncio +async def test_basic_secret_detection_embeddings(): + secret_instance: Final = _ENTERPRISE_SecretDetection() + user_api_key_dict: Final = UserAPIKeyAuth(api_key=hash_token("sk-98765")) + local_cache: Final = DualCache() + + test_data: Final = { + "input": "Hey, how's it going, API_KEY = 'sk_1234567890abcdef', my OPENAI_API_KEY = 'sk_1234567890abcdef' and i want to know what is the weather", + "model": "gpt-3.5-turbo", + } + + await secret_instance.async_pre_call_hook( + cache=local_cache, + data=test_data, + user_api_key_dict=user_api_key_dict, + call_type="embedding", + ) + + assert test_data == { + "input": "Hey, how's it going, API_KEY = '[REDACTED]', my OPENAI_API_KEY = '[REDACTED]' and i want to know what is the weather", + "model": "gpt-3.5-turbo", + } + + +@pytest.mark.asyncio +async def test_basic_secret_detection_embeddings_list(): + secret_instance: Final = _ENTERPRISE_SecretDetection() + user_api_key_dict: Final = UserAPIKeyAuth(api_key=hash_token("sk-98765")) + local_cache: Final = DualCache() + + test_data: Final = { + "input": [ + "hey", + "how's it going, API_KEY = 'sk_1234567890abcdef'", + "my OPENAI_API_KEY = 'sk_1234567890abcdef' and i want to know what is the weather", + ], + "model": "gpt-3.5-turbo", + } + + await secret_instance.async_pre_call_hook( + cache=local_cache, + data=test_data, + user_api_key_dict=user_api_key_dict, + call_type="embedding", + ) + + assert test_data == { + "input": [ + "hey", + "how's it going, API_KEY = '[REDACTED]'", + "my OPENAI_API_KEY = '[REDACTED]' and i want to know what is the weather", + ], + "model": "gpt-3.5-turbo", + } diff --git a/tests/logging_callback_tests/test_view_request_resp_logs.py b/tests/unit/integrations/gcs_bucket/test_gcs_bucket.py similarity index 51% rename from tests/logging_callback_tests/test_view_request_resp_logs.py rename to tests/unit/integrations/gcs_bucket/test_gcs_bucket.py index 249e84286d5..981ec9bcce3 100644 --- a/tests/logging_callback_tests/test_view_request_resp_logs.py +++ b/tests/unit/integrations/gcs_bucket/test_gcs_bucket.py @@ -1,28 +1,19 @@ -import io - - -import asyncio import json -import logging -import tempfile -from litellm._uuid import uuid - -from datetime import datetime, timedelta, timezone +from collections.abc import Sequence +from datetime import datetime, timezone +from typing import Final +import google.auth +import google.auth.credentials +import google.auth.transport +import httpx import pytest -import litellm -from litellm import completion -from litellm._logging import verbose_logger -from litellm.integrations.gcs_bucket.gcs_bucket import ( - GCSBucketLogger, - StandardLoggingPayload, -) -from litellm.types.utils import StandardCallbackDynamicParams +from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -# This is the response payload that GCS would return. -mock_response_data = { +mock_response_data: Final = { "id": "chatcmpl-9870a859d6df402795f75dc5fca5b2e0", "trace_id": None, "call_type": "acompletion", @@ -60,9 +51,7 @@ mock_response_data = { "model_group": "fake-openai-endpoint", "model_id": "b68d56d76b0c24ac9462ab69541e90886342508212210116e300441155f37865", "requester_ip_address": "127.0.0.1", - "messages": [ - {"role": "user", "content": [{"type": "text", "text": "very gm to u"}]} - ], + "messages": [{"role": "user", "content": [{"type": "text", "text": "very gm to u"}]}], "response": { "id": "chatcmpl-9870a859d6df402795f75dc5fca5b2e0", "created": 1677652288, @@ -111,93 +100,71 @@ mock_response_data = { } +class _StaticGoogleCredentials(google.auth.credentials.Credentials): + def refresh(self, request: google.auth.transport.Request) -> None: + self.token = "test-access-token" + + +def _google_default_credentials(scopes: Sequence[str]) -> tuple[_StaticGoogleCredentials, str]: + return _StaticGoogleCredentials(), "test-project" + + +def _gcs_logger_storing_payload_on(stored_date: str | None, monkeypatch: pytest.MonkeyPatch) -> GCSBucketLogger: + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr(google.auth, "default", _google_default_credentials) + + def storage(request: httpx.Request) -> httpx.Response: + if request.headers.get("Authorization") != "Bearer test-access-token": + return httpx.Response(401) + if not str(request.url).startswith("https://storage.googleapis.com/storage/v1/b/test-bucket/o/"): + return httpx.Response(404) + if stored_date is None or stored_date not in str(request.url): + return httpx.Response(404, text="No such object") + return httpx.Response(200, content=json.dumps(mock_response_data).encode("utf-8")) + + gcs_logger: Final = GCSBucketLogger(bucket_name="test-bucket") + gcs_logger.async_httpx_client = AsyncHTTPHandler(transport=httpx.MockTransport(storage)) + return gcs_logger + + @pytest.mark.asyncio -async def test_get_payload_current_day(): - """ - Verify that the payload is returned when it is found on the current day. - """ - gcs_logger = GCSBucketLogger() - # Use January 1, 2024 as the current day - start_time = datetime(2024, 1, 1, tzinfo=timezone.utc) - request_id = mock_response_data["id"] +async def test_get_payload_current_day(monkeypatch): + gcs_logger: Final = _gcs_logger_storing_payload_on("2024-01-01", monkeypatch) + start_time: Final = datetime(2024, 1, 1, tzinfo=timezone.utc) + request_id: Final = mock_response_data["id"] - async def fake_download(object_name: str, **kwargs) -> bytes | None: - if "2024-01-01" in object_name: - return json.dumps(mock_response_data).encode("utf-8") - return None - - gcs_logger.download_gcs_object = fake_download - - payload = await gcs_logger.get_request_response_payload( - request_id, start_time, None - ) + payload: Final = await gcs_logger.get_request_response_payload(request_id, start_time, None) assert payload is not None assert payload["id"] == request_id @pytest.mark.asyncio -async def test_get_payload_next_day(): - """ - Verify that if the payload is not found on the current day, - but is available on the next day, it is returned. - """ - gcs_logger = GCSBucketLogger() - start_time = datetime(2024, 1, 1, tzinfo=timezone.utc) - request_id = mock_response_data["id"] +async def test_get_payload_next_day(monkeypatch): + gcs_logger: Final = _gcs_logger_storing_payload_on("2024-01-02", monkeypatch) + start_time: Final = datetime(2024, 1, 1, tzinfo=timezone.utc) + request_id: Final = mock_response_data["id"] - async def fake_download(object_name: str, **kwargs) -> bytes | None: - if "2024-01-02" in object_name: - return json.dumps(mock_response_data).encode("utf-8") - return None - - gcs_logger.download_gcs_object = fake_download - - payload = await gcs_logger.get_request_response_payload( - request_id, start_time, None - ) + payload: Final = await gcs_logger.get_request_response_payload(request_id, start_time, None) assert payload is not None assert payload["id"] == request_id @pytest.mark.asyncio -async def test_get_payload_previous_day(): - """ - Verify that if the payload is not found on the current or next day, - but is available on the previous day, it is returned. - """ - gcs_logger = GCSBucketLogger() - start_time = datetime(2024, 1, 1, tzinfo=timezone.utc) - request_id = mock_response_data["id"] +async def test_get_payload_previous_day(monkeypatch): + gcs_logger: Final = _gcs_logger_storing_payload_on("2023-12-31", monkeypatch) + start_time: Final = datetime(2024, 1, 1, tzinfo=timezone.utc) + request_id: Final = mock_response_data["id"] - async def fake_download(object_name: str, **kwargs) -> bytes | None: - if "2023-12-31" in object_name: - return json.dumps(mock_response_data).encode("utf-8") - return None - - gcs_logger.download_gcs_object = fake_download - - payload = await gcs_logger.get_request_response_payload( - request_id, start_time, None - ) + payload: Final = await gcs_logger.get_request_response_payload(request_id, start_time, None) assert payload is not None assert payload["id"] == request_id @pytest.mark.asyncio -async def test_get_payload_not_found(): - """ - Verify that if none of the three days contain the payload, None is returned. - """ - gcs_logger = GCSBucketLogger() - start_time = datetime(2024, 1, 1, tzinfo=timezone.utc) - request_id = mock_response_data["id"] +async def test_get_payload_not_found(monkeypatch): + gcs_logger: Final = _gcs_logger_storing_payload_on(None, monkeypatch) + start_time: Final = datetime(2024, 1, 1, tzinfo=timezone.utc) + request_id: Final = mock_response_data["id"] - async def fake_download(object_name: str, **kwargs) -> bytes | None: - return None - - gcs_logger.download_gcs_object = fake_download - - payload = await gcs_logger.get_request_response_payload( - request_id, start_time, None - ) + payload: Final = await gcs_logger.get_request_response_payload(request_id, start_time, None) assert payload is None diff --git a/tests/unit/proxy/anthropic_endpoints/test_claude_code_marketplace.py b/tests/unit/proxy/anthropic_endpoints/test_claude_code_marketplace.py index b4780bfefeb..ccc63f12157 100644 --- a/tests/unit/proxy/anthropic_endpoints/test_claude_code_marketplace.py +++ b/tests/unit/proxy/anthropic_endpoints/test_claude_code_marketplace.py @@ -4,6 +4,7 @@ Unit tests for claude_code_marketplace.py source validation. Covers the git-subdir and archive source types added alongside the existing github and url types. """ +from typing import Final import json import pytest @@ -597,3 +598,111 @@ async def test_enable_disable_delete_plugin_reject_non_admin(): table = litellm.proxy.proxy_server.prisma_client.db.litellm_claudecodeplugintable assert (await table.find_unique(where={"name": name})).enabled is True + + +@pytest.mark.asyncio +async def test_register_plugin(): + plugin_name: Final = "test-plugin" + + request: Final = RegisterPluginRequest( + name=plugin_name, + source={"source": "github", "repo": "test-org/test-repo"}, + version="1.0.0", + description="Test plugin for unit tests", + ) + + user_api_key_dict: Final = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key=MASTER_KEY, + user_id="test-user", + ) + + response: Final = await register_plugin( + request=request, + user_api_key_dict=user_api_key_dict, + ) + + assert response.status == "success" + assert response.action == "created" + assert response.plugin.name == plugin_name + assert response.plugin.version == "1.0.0" + assert response.plugin.enabled is True + + stored_plugin: Final = await litellm.proxy.proxy_server.prisma_client.db.litellm_claudecodeplugintable.find_unique( + where={"name": plugin_name} + ) + assert stored_plugin is not None + assert stored_plugin.name == plugin_name + + +@pytest.mark.asyncio +async def test_get_marketplace(): + plugin_name: Final = "test-marketplace-plugin" + + request: Final = RegisterPluginRequest( + name=plugin_name, + source={"source": "github", "repo": "test-org/marketplace-test"}, + version="2.0.0", + description="Test plugin for marketplace test", + ) + + user_api_key_dict: Final = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key=MASTER_KEY, + user_id="test-user", + ) + + await register_plugin( + request=request, + user_api_key_dict=user_api_key_dict, + ) + + response: Final = await get_marketplace(request=MagicMock()) + + body: Final = json.loads(response.body.decode()) + + assert body["name"] == "litellm" + assert "plugins" in body + + our_plugin: Final = next((p for p in body["plugins"] if p["name"] == plugin_name), None) + assert our_plugin is not None + assert our_plugin["source"] == { + "source": "github", + "repo": "test-org/marketplace-test", + } + assert our_plugin["version"] == "2.0.0" + + +@pytest.mark.asyncio +async def test_register_plugin_git_subdir(): + plugin_name: Final = "test-subdir-plugin" + + request: Final = RegisterPluginRequest( + name=plugin_name, + source={ + "source": "git-subdir", + "url": "https://github.com/test-org/monorepo.git", + "path": "plugins/my-plugin", + }, + version="1.0.0", + description="Test git-subdir plugin", + ) + + user_api_key_dict: Final = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key=MASTER_KEY, + user_id="test-user", + ) + + response: Final = await register_plugin( + request=request, + user_api_key_dict=user_api_key_dict, + ) + + assert response.status == "success" + assert response.action == "created" + assert response.plugin.name == plugin_name + assert response.plugin.source["source"] == "git-subdir" + assert response.plugin.source["url"] == "https://github.com/test-org/monorepo.git" + assert response.plugin.source["path"] == "plugins/my-plugin" + assert response.plugin.enabled is True diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index e2cec1cc5a2..52c5fa2a41e 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -4274,3 +4274,132 @@ def test_get_model_from_request_vertex_ai_passthrough( model = get_model_from_request(request_data, route) assert model == expected_model + + +def test_get_end_user_id_from_request_body_always_returns_str(): + mock_request: Final = MagicMock(spec=Request) + mock_request.headers = {} + + request_body: Final = {"user": 123} + end_user_id: Final = get_end_user_id_from_request_body(request_body, dict(mock_request.headers)) + assert end_user_id == "123" + assert isinstance(end_user_id, str) + + +@pytest.mark.parametrize( + "headers, general_settings_config, request_body, expected_user_id", + [ + ( + {"X-User-ID": "header-user-123"}, + {"user_header_name": "X-User-ID"}, + {"user": "body-user-456"}, + "header-user-123", + ), + ( + {}, + {"user_header_name": "X-User-ID"}, + {"user": "body-user-456"}, + "body-user-456", + ), + ( + {"X-User-ID": "header-user-123"}, + {}, + {"user": "body-user-456"}, + "body-user-456", + ), + ( + {"X-Custom-User": "header-only-user"}, + {"user_header_name": "X-Custom-User"}, + {"model": "gpt-4"}, + "header-only-user", + ), + ( + {"X-User-ID": ""}, + {"user_header_name": "X-User-ID"}, + {"user": "body-user-456"}, + "body-user-456", + ), + ( + {"x-user-id": "lowercase-header-user"}, + {"user_header_name": "x-user-id"}, + {"user": "body-user-456"}, + "lowercase-header-user", + ), + ( + {"X-User-ID": "header-user-123"}, + {"user_header_name": None}, + {"user": "body-user-456"}, + "body-user-456", + ), + ( + {"X-User-ID": "header-user-123"}, + {"user_header_name": 123}, + {"user": "body-user-456"}, + "body-user-456", + ), + ( + {}, + {"user_header_name": "X-User-ID"}, + {"litellm_metadata": {"user": "litellm-user-789"}}, + "litellm-user-789", + ), + ( + {}, + {"user_header_name": "X-User-ID"}, + {"metadata": {"user_id": "metadata-user-999"}}, + "metadata-user-999", + ), + ( + {"X-User-ID": "header-priority"}, + {"user_header_name": "X-User-ID"}, + { + "user": "body-user", + "litellm_metadata": {"user": "litellm-user"}, + "metadata": {"user_id": "metadata-user"}, + }, + "header-priority", + ), + ( + {"x-user-id": "lowercase-header-user"}, + {"user_header_name": "X-User-ID"}, + {"user": "body-user-456"}, + "lowercase-header-user", + ), + ], +) +def test_get_end_user_id_from_request_body_with_user_header_name( + headers, general_settings_config, request_body, expected_user_id +): + mock_request: Final = MagicMock(spec=Request) + mock_request.headers = headers + + with patch("litellm.proxy.proxy_server.general_settings", general_settings_config): + end_user_id: Final = get_end_user_id_from_request_body(request_body, dict(mock_request.headers)) + assert end_user_id == expected_user_id + + +def test_get_end_user_id_from_request_body_no_user_found(): + mock_request: Final = MagicMock(spec=Request) + mock_request.headers = {"X-Other-Header": "some-value"} + + general_settings_config: Final = {"user_header_name": "X-User-ID"} + + request_body: Final = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hello"}], + } + + with patch("litellm.proxy.proxy_server.general_settings", general_settings_config): + end_user_id: Final = get_end_user_id_from_request_body(request_body, dict(mock_request.headers)) + assert end_user_id is None + + +def test_get_end_user_id_from_request_body_backwards_compatibility(): + cases: Final = ( + ({"user": "test-user-123"}, "test-user-123"), + ({"litellm_metadata": {"user": "litellm-user-456"}}, "litellm-user-456"), + ({"metadata": {"user_id": "metadata-user-789"}}, "metadata-user-789"), + ({"model": "gpt-4"}, None), + ) + for request_body, expected_end_user_id in cases: + assert get_end_user_id_from_request_body(request_body) == expected_end_user_id diff --git a/tests/unit/proxy/common_utils/test_reset_budget_job.py b/tests/unit/proxy/common_utils/test_reset_budget_job.py index d41bfbc7308..c17aed49014 100644 --- a/tests/unit/proxy/common_utils/test_reset_budget_job.py +++ b/tests/unit/proxy/common_utils/test_reset_budget_job.py @@ -2,7 +2,7 @@ import asyncio import json import sys import types -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Sequence from datetime import datetime, timedelta, timezone from datetime import time as dt_time from typing import Any, Dict, Final, List, Optional @@ -3953,6 +3953,166 @@ async def test_reset_budget_endusers_cascade_failure_is_all_or_nothing(): ) proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called() + +class _FrozenClock(datetime): + @classmethod + def now(cls, tz=None): + return _FROZEN_NOW if tz is None else _FROZEN_NOW.astimezone(tz) + + @classmethod + def utcnow(cls): + return _FROZEN_NOW.replace(tzinfo=None) + + +_FROZEN_NOW: Final = _FrozenClock(2024, 6, 15, 10, 30, tzinfo=timezone.utc) + + +@pytest.fixture +def frozen_reset_clock(monkeypatch: pytest.MonkeyPatch) -> datetime: + from litellm.proxy.common_utils import timezone_utils + + monkeypatch.setattr(reset_budget_job_module, "datetime", _FrozenClock) + monkeypatch.setattr(timezone_utils, "datetime", _FrozenClock) + return _FROZEN_NOW + + +async def _await_tasks_spawned_by(operation: Awaitable[None]) -> None: + already_running: Final = asyncio.all_tasks() + await operation + await asyncio.gather(*(asyncio.all_tasks() - already_running - {asyncio.current_task()})) + + +@pytest.mark.asyncio +async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance(frozen_reset_clock): + endusers: Final = [_attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"}) for i in range(1, 7)] + + budget1: Final = LiteLLM_BudgetTableFull( + **{ + "budget_id": "budget1", + "max_budget": 65.0, + "budget_duration": "2d", + "created_at": frozen_reset_clock - timedelta(days=3), + } + ) + + prisma_client: Final = MagicMock() + + async def get_data_mock( + *, table_name: str, query_type: str, reset_at: datetime, limit: int, expires: datetime | None = None + ) -> Sequence[object]: + if table_name == "budget": + return [budget1] + elif table_name == "enduser": + return endusers + return [] + + prisma_client.get_data = AsyncMock() + prisma_client.get_data.side_effect = get_data_mock + prisma_client.update_data = AsyncMock() + batch_calls: Final = _wire_batcher_for_test(prisma_client) + + proxy_logging_obj: Final = MagicMock() + proxy_logging_obj.service_logging_obj = MagicMock() + proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() + proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() + + job: Final = ResetBudgetJob(proxy_logging_obj, prisma_client) + + await _await_tasks_spawned_by(job.reset_budget_for_litellm_budget_table()) + + assert prisma_client.db.batch_.call_count == 1, "the cascade must be one transaction" + + enduser_writes: Final = [c for c in batch_calls if c["table"] == "enduser"] + assert len(enduser_writes) == 1 + assert enduser_writes[0]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}} + assert enduser_writes[0]["data"] == {"spend": 0} + + budget_writes: Final = [c for c in batch_calls if c["table"] == "budget"] + assert len(budget_writes) == 1 + assert budget_writes[0]["where"] == {"budget_id": "budget1"} + assert budget_writes[0]["data"]["budget_reset_at"] > frozen_reset_clock + + proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called() + + +@pytest.mark.asyncio +async def test_reset_budget_continues_other_categories_on_failure(frozen_reset_clock): + key1: Final = _attrify({"id": "key1", "token": "key1", "spend": 10.0, "budget_duration": "60s"}) + key2: Final = _attrify({"id": "key2", "token": "key2", "spend": 15.0, "budget_duration": "60s"}) + user1: Final = _attrify({"id": "user1", "user_id": "user1", "spend": 20.0, "budget_duration": 120}) + user2: Final = _attrify({"id": "user2", "user_id": "user2", "spend": 25.0, "budget_duration": "120s"}) + team1: Final = _attrify({"id": "team1", "team_id": "team1", "spend": 30.0, "budget_duration": "180s"}) + team2: Final = _attrify({"id": "team2", "team_id": "team2", "spend": 35.0, "budget_duration": "180s"}) + enduser1: Final = _attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}) + budget1: Final = LiteLLM_BudgetTableFull( + **{ + "budget_id": "budget1", + "max_budget": 65.0, + "budget_duration": "2d", + "created_at": frozen_reset_clock - timedelta(days=3), + } + ) + + prisma_client: Final = MagicMock() + + async def fake_get_data( + *, table_name: str, query_type: str, reset_at: datetime, limit: int, expires: datetime | None = None + ) -> Sequence[object]: + if table_name == "key": + return [key1, key2] + elif table_name == "user": + return [user1, user2] + elif table_name == "team": + return [team1, team2] + elif table_name == "budget": + return [budget1] + elif table_name == "enduser": + return [enduser1] + return [] + + prisma_client.get_data = AsyncMock(side_effect=fake_get_data) + prisma_client.update_data = AsyncMock() + batch_calls: Final = _wire_batcher_for_test(prisma_client) + pre_reset_spend: Final = { + **{k["token"]: k["spend"] for k in [key1, key2]}, + **{u["user_id"]: u["spend"] for u in [user2]}, + **{t["team_id"]: t["spend"] for t in [team1, team2]}, + } + _wire_cascade_reads_for_test(prisma_client, endusers=[enduser1]) + + proxy_logging_obj: Final = MagicMock() + proxy_logging_obj.service_logging_obj = MagicMock() + proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() + proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() + + job: Final = ResetBudgetJob(proxy_logging_obj, prisma_client) + + await _await_tasks_spawned_by(job.reset_budget()) + + called_tables: Final = {call.kwargs.get("table_name") for call in prisma_client.get_data.await_args_list} + assert called_tables == {"key", "user", "team", "budget"} + prisma_client.db.litellm_endusertable.find_many.assert_awaited() + + prisma_client.update_data.assert_not_awaited() + + assert len([c for c in batch_calls if c["table"] == "team_membership"]) == 1 + enduser_writes: Final = [c for c in batch_calls if c["table"] == "enduser"] + assert len(enduser_writes) == 1 + assert enduser_writes[0]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}} + assert enduser_writes[0]["data"] == {"spend": 0} + + key_writes: Final = [c for c in batch_calls if c["table"] == "key" and c["op"] == "update"] + user_writes: Final = [c for c in batch_calls if c["table"] == "user"] + team_writes: Final = [c for c in batch_calls if c["table"] == "team"] + assert len(key_writes) == 2 + assert len(user_writes) == 1 + assert user_writes[0]["where"] == {"user_id": "user2"} + assert len(team_writes) == 2 + for c in key_writes + user_writes + team_writes: + assert set(c["data"].keys()) == {"spend", "budget_reset_at"} + assert c["data"]["spend"] == {"decrement": pre_reset_spend[next(iter(c["where"].values()))]} + + @pytest.mark.asyncio async def test_reset_budget_teams_partial_failure(): """ diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/openai/test_moderations.py b/tests/unit/proxy/guardrails/guardrail_hooks/openai/test_moderations.py index 2c1412d0bf9..2c91569de69 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/openai/test_moderations.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/openai/test_moderations.py @@ -1026,3 +1026,66 @@ async def test_openai_moderation_records_moderation_id_as_scan_metadata(input_ty assert json.loads(headers["x-litellm-guardrail-scan-metadata"]) == [ {"guardrail": "openai-mod", "stage": stage, "provider": "openai_moderation", "scan_id": f"modr-{stage}"} ] + + +@pytest.mark.asyncio +async def test_openai_moderation_responses_api_input_field(): + from litellm.types.utils import GenericGuardrailAPIInputs + + flagged_moderation: Final = { + "id": "modr-123", + "model": "omni-moderation-latest", + "results": [ + { + "flagged": True, + "categories": {"violence": True, "hate": False}, + "category_scores": {"violence": 0.95, "hate": 0.1}, + "category_applied_input_types": None, + } + ], + } + http_client: Final = AsyncHTTPHandler() + http_client.client = httpx.AsyncClient( + transport=httpx.MockTransport(lambda _: httpx.Response(200, json=flagged_moderation)) + ) + + openai_mod: Final = OpenAIModerationGuardrail( + guardrail_name="openai-moderation-test", + api_key="fake-key-for-testing", + model="omni-moderation-latest", + ) + openai_mod.async_handler = http_client + + with pytest.raises(Exception, match="Violated OpenAI moderation policy") as texts_error: + await openai_mod.apply_guardrail( + inputs=GenericGuardrailAPIInputs(texts=["I want to hurt people"]), + request_data={"model": "gpt-4o", "input": "I want to hurt people"}, + input_type="request", + ) + assert "Violated OpenAI moderation policy" in str(texts_error.value) + + with pytest.raises(Exception, match="Violated OpenAI moderation policy") as responses_error: + await openai_mod.apply_guardrail( + inputs=GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "I want to hurt people"}] + ), + request_data={ + "model": "gpt-4o", + "input": [{"role": "user", "content": "I want to hurt people"}], + }, + input_type="request", + ) + assert "Violated OpenAI moderation policy" in str(responses_error.value) + + with pytest.raises(Exception, match="Violated OpenAI moderation policy") as chat_error: + await openai_mod.apply_guardrail( + inputs=GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "I want to hurt people"}] + ), + request_data={ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "I want to hurt people"}], + }, + input_type="request", + ) + assert "Violated OpenAI moderation policy" in str(chat_error.value) diff --git a/tests/unit/proxy/guardrails/test_guardrail_coverage.py b/tests/unit/proxy/guardrails/test_guardrail_coverage.py index 28e70ec84fa..1feae213eba 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_coverage.py +++ b/tests/unit/proxy/guardrails/test_guardrail_coverage.py @@ -12,7 +12,7 @@ inspection by the relevant guardrail hook: not just ``choices[0]``. """ -from typing import Any, Dict +from typing import Any, Dict, Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -800,6 +800,66 @@ async def test_openai_moderation_reads_model_name_at_call_time( fake_router.amoderation.assert_awaited_once_with(model=expected_model, input="hello") +@pytest.mark.asyncio +async def test_openai_moderation_error_raising(monkeypatch, respx_mock): + from enterprise.enterprise_hooks.openai_moderation import ( + ENTERPRISE_OpenAI_Moderation, + ) + from litellm.proxy.utils import hash_token + + monkeypatch.setattr(litellm, "openai_moderations_model_name", "omni-moderation-latest") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + openai_mod: Final = ENTERPRISE_OpenAI_Moderation() + user_api_key_dict: Final = UserAPIKeyAuth(api_key=hash_token("sk-98765")) + + llm_router: Final = litellm.Router( + model_list=[ + { + "model_name": "omni-moderation-latest", + "litellm_params": { + "model": "omni-moderation-latest", + "api_key": "fake-key", + }, + } + ] + ) + + moderation_route: Final = respx_mock.post("https://api.openai.com/v1/moderations").respond( + 200, + json={ + "id": "modr-123", + "model": "omni-moderation-latest", + "results": [ + { + "flagged": True, + "categories": {"harassment": True}, + "category_scores": {"harassment": 0.97}, + } + ], + }, + ) + + import litellm.proxy.proxy_server as proxy_server + + monkeypatch.setattr(proxy_server, "llm_router", llm_router) + + with pytest.raises(Exception, match="Violated content safety policy") as exc_info: + await openai_mod.async_moderation_hook( + data={ + "messages": [ + { + "role": "user", + "content": "fuck off you're the worst", + } + ] + }, + user_api_key_dict=user_api_key_dict, + call_type="completion", + ) + assert "Violated content safety policy" in str(exc_info.value) + assert moderation_route.call_count == 1 + + # ── Google Text Moderation ──────────────────────────────────────────────────── diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py index 07e6c3ac89d..93afea1b106 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -2863,3 +2863,46 @@ def test_build_complete_streaming_response(all_chunks): assert result.usage.prompt_tokens == 17 assert result.usage.completion_tokens == 249 assert result.usage.total_tokens == 266 + + +@pytest.mark.parametrize( + "end_user_id", + [{"litellm_metadata": {"user": "test"}}, {"metadata": {"user_id": "test"}}], +) +def test_get_user_from_metadata(end_user_id): + from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( + AnthropicPassthroughLoggingHandler, + PassthroughStandardLoggingPayload, + ) + + passthrough_logging_payload: Final = PassthroughStandardLoggingPayload( + url="https://api.anthropic.com/v1/messages", + request_body={**end_user_id}, + response_body={ + "id": "msg_015uSaCZBvu9gUSkAmZtMfxC", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [ + { + "type": "text", + "text": "Now I'll click on the Firefox icon to launch it.", + }, + { + "type": "tool_use", + "id": "toolu_01TQsF5p7Pf4LGKyLUDDySVr", + "name": "computer", + "input": {"action": "mouse_move", "coordinate": [24, 36]}, + }, + ], + "stop_reason": "tool_use", + "stop_sequence": None, + "usage": {"input_tokens": 2202, "output_tokens": 89}, + }, + ) + + response: Final = AnthropicPassthroughLoggingHandler._get_user_from_metadata( + passthrough_logging_payload=passthrough_logging_payload + ) + + assert response == "test" diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index d329db75fb4..d25b4470668 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -5716,6 +5716,7 @@ class TestVertexAILiveWebsocketPassthrough: "wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" ) assert passthrough_kwargs["custom_headers"]["Authorization"] == "Bearer token-abc" + assert passthrough_kwargs["endpoint"] == "/vertex_ai/live" rewriter = passthrough_kwargs["setup_model_rewriter"] assert rewriter("gemini-live") == ( "projects/proj-db/locations/global/publishers/google/models/gemini-live-2.5-flash" diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index a6ce8ee6fbc..1d761d48818 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -8,6 +8,7 @@ import zlib from collections.abc import AsyncIterator, Callable, Mapping from contextlib import ExitStack, contextmanager from dataclasses import dataclass +from datetime import datetime from io import BytesIO from types import MappingProxyType, ModuleType, SimpleNamespace from typing import Final, Optional @@ -32,6 +33,7 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._lazy_features import LazyFeature, attach_lazy_features from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import router as llm_passthrough_router from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, @@ -61,6 +63,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, EndpointType, + PassthroughStandardLoggingPayload, ) from tests._master_key import MASTER_KEY from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome @@ -9083,3 +9086,372 @@ def test_update_subpath_route_updates_registry(): del _registered_pass_through_routes[route_key] asyncio.run(_async_test()) + + +def test_init_kwargs_for_pass_through_endpoint_basic(mock_request, mock_user_api_key_dict): + request: Final = mock_request() + passthrough_payload: Final = PassthroughStandardLoggingPayload( + url="https://test.com", + request_body={}, + ) + + result: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, + user_api_key_dict=mock_user_api_key_dict, + passthrough_logging_payload=passthrough_payload, + litellm_call_id="test-call-id", + logging_obj=LiteLLMLoggingObj( + model="test-model", + messages=[], + stream=False, + call_type="test-call-type", + start_time=datetime(2026, 1, 15, 12, 0, 0), + litellm_call_id="test-call-id", + function_id="test-function-id", + ), + ) + + assert result["call_type"] == "pass_through_endpoint" + assert result["litellm_call_id"] == "test-call-id" + assert result["passthrough_logging_payload"] == passthrough_payload + + assert result["litellm_params"]["metadata"]["user_api_key"] == "test-key" + assert result["litellm_params"]["metadata"]["user_api_key_hash"] == "test-key" + assert result["litellm_params"]["metadata"]["user_api_key_alias"] is None + assert result["litellm_params"]["metadata"]["user_api_key_user_email"] is None + assert result["litellm_params"]["metadata"]["user_api_key_user_id"] == "test-user" + assert result["litellm_params"]["metadata"]["user_api_key_team_id"] == "test-team" + assert result["litellm_params"]["metadata"]["user_api_key_org_id"] is None + assert result["litellm_params"]["metadata"]["user_api_key_team_alias"] is None + assert result["litellm_params"]["metadata"]["user_api_key_end_user_id"] == "test-user" + assert result["litellm_params"]["metadata"]["user_api_key_request_route"] is None + + +def test_init_kwargs_with_litellm_metadata(mock_request, mock_user_api_key_dict): + request: Final = mock_request() + parsed_body: Final = {"litellm_metadata": {"custom_field": "custom_value", "tags": ["tag1", "tag2"]}} + passthrough_payload: Final = PassthroughStandardLoggingPayload( + url="https://test.com", + request_body={}, + ) + + result: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, + user_api_key_dict=mock_user_api_key_dict, + passthrough_logging_payload=passthrough_payload, + _parsed_body=parsed_body, + litellm_call_id="test-call-id", + logging_obj=LiteLLMLoggingObj( + model="test-model", + messages=[], + stream=False, + call_type="test-call-type", + start_time=datetime(2026, 1, 15, 12, 0, 0), + litellm_call_id="test-call-id", + function_id="test-function-id", + ), + ) + + metadata: Final = result["litellm_params"]["metadata"] + assert metadata["custom_field"] == "custom_value" + assert metadata["tags"] == ["tag1", "tag2"] + assert metadata["user_api_key"] == "test-key" + + +def test_init_kwargs_with_tags_in_header(mock_request, mock_user_api_key_dict): + request: Final = mock_request(headers={"tags": "tag1,tag2"}) + passthrough_payload: Final = PassthroughStandardLoggingPayload( + url="https://test.com", + request_body={}, + ) + + result: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, + user_api_key_dict=mock_user_api_key_dict, + passthrough_logging_payload=passthrough_payload, + litellm_call_id="test-call-id", + logging_obj=LiteLLMLoggingObj( + model="test-model", + messages=[], + stream=False, + call_type="test-call-type", + start_time=datetime(2026, 1, 15, 12, 0, 0), + litellm_call_id="test-call-id", + function_id="test-function-id", + ), + ) + + metadata: Final = result["litellm_params"]["metadata"] + assert metadata["tags"] == ["tag1", "tag2"] + + +athropic_request_body: Final = { + "model": "claude-sonnet-4-5-20250929", + "max_tokens": 256, + "messages": [{"role": "user", "content": "Hello, world tell me 2 sentences "}], + "litellm_metadata": {"tags": ["hi", "hello"]}, +} + + +@pytest.mark.asyncio +async def test_pass_through_request_logging_failure(mock_request, mock_user_api_key_dict): + + mock_response: Final = AsyncMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + + mock_response._content = b'{"mock": "response"}' + + async def mock_aread(): + return mock_response._content + + mock_response.aread = mock_aread + + with ( + patch( + "httpx.AsyncClient.send", + return_value=mock_response, + ), + patch( + "httpx.AsyncClient.request", + return_value=mock_response, + ), + ): + request: Final = mock_request(headers={}, method="POST", request_body=athropic_request_body) + response: Final = await pass_through_request( + request=request, + target="https://exampleopenaiendpoint-production.up.railway.app/v1/messages", + custom_headers={}, + user_api_key_dict=mock_user_api_key_dict, + ) + + assert response.status_code == 200 + + assert response.body == b'{"mock": "response"}' + + +@pytest.mark.asyncio +async def test_pass_through_request_logging_failure_with_stream(mock_request, mock_user_api_key_dict): + + mock_response: Final = AsyncMock() + mock_response.status_code = 200 + + mock_response.headers = { + "content-type": "application/json", + } + + mock_chunks: Final = [b'{"chunk": 1}', b'{"chunk": 2}'] + mock_response.body_iterator = AsyncMock() + mock_response.body_iterator.__aiter__.return_value = mock_chunks + + mock_response._content = b'{"mock": "response"}' + + async def mock_aread(): + return mock_response._content + + mock_response.aread = mock_aread + + with ( + patch( + "httpx.AsyncClient.send", + return_value=mock_response, + ), + patch( + "httpx.AsyncClient.request", + return_value=mock_response, + ), + ): + request: Final = mock_request(headers={}, method="POST", request_body=athropic_request_body) + response: Final = await pass_through_request( + request=request, + target="https://exampleopenaiendpoint-production.up.railway.app/v1/messages", + custom_headers={}, + user_api_key_dict=mock_user_api_key_dict, + ) + + assert response.status_code == 200 + + if isinstance(response, StreamingResponse): + assert response.status_code == 200 + else: + assert hasattr(response, "body") + assert response.body == b'{"mock": "response"}' + + +def test_init_kwargs_filters_pricing_params(mock_request, mock_user_api_key_dict): + request: Final = mock_request() + + parsed_body: Final = { + "model": "gpt-5.5", + "messages": [{"role": "user", "content": "test"}], + "input_cost_per_token": 0.00002, + "output_cost_per_token": 0.00002, + "input_cost_per_second": 0.00001, + "output_cost_per_second": 0.00001, + "cache_read_input_token_cost": 0.00005, + "cache_creation_input_token_cost": 0.00003, + "cache_creation_input_token_cost_above_1hr": 0.00004, + "input_cost_per_token_batches": 0.00005, + "output_cost_per_token_batches": 0.00006, + "input_cost_per_audio_token": 0.00001, + "output_cost_per_audio_token": 0.00001, + "input_cost_per_character": 0.000001, + "output_cost_per_character": 0.000001, + "input_cost_per_image": 0.001, + "output_cost_per_image": 0.001, + "tiered_pricing": [{"input_cost_per_token": 0.00001}], + "temperature": 0.7, + "max_tokens": 100, + } + + passthrough_payload: Final = PassthroughStandardLoggingPayload( + url="https://api.openai.com/v1/chat/completions", + request_body=parsed_body.copy(), + ) + + result: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, + user_api_key_dict=mock_user_api_key_dict, + passthrough_logging_payload=passthrough_payload, + _parsed_body=parsed_body, + litellm_call_id="test-call-id", + logging_obj=LiteLLMLoggingObj( + model="gpt-5.5", + messages=[{"role": "user", "content": "test"}], + stream=False, + call_type="completion", + start_time=datetime(2026, 1, 15, 12, 0, 0), + litellm_call_id="test-call-id", + function_id="test-function-id", + ), + ) + + assert "input_cost_per_token" not in parsed_body + assert "output_cost_per_token" not in parsed_body + assert "input_cost_per_second" not in parsed_body + assert "output_cost_per_second" not in parsed_body + assert "cache_read_input_token_cost" not in parsed_body + assert "cache_creation_input_token_cost" not in parsed_body + assert "cache_creation_input_token_cost_above_1hr" not in parsed_body + assert "input_cost_per_token_batches" not in parsed_body + assert "output_cost_per_token_batches" not in parsed_body + assert "input_cost_per_audio_token" not in parsed_body + assert "output_cost_per_audio_token" not in parsed_body + assert "input_cost_per_character" not in parsed_body + assert "output_cost_per_character" not in parsed_body + assert "input_cost_per_image" not in parsed_body + assert "output_cost_per_image" not in parsed_body + assert "tiered_pricing" not in parsed_body + + assert parsed_body["model"] == "gpt-5.5" + assert parsed_body["messages"] == [{"role": "user", "content": "test"}] + assert parsed_body["temperature"] == 0.7 + assert parsed_body["max_tokens"] == 100 + + litellm_params: Final = result["litellm_params"] + assert litellm_params["input_cost_per_token"] == 0.00002 + assert litellm_params["output_cost_per_token"] == 0.00002 + + +def test_init_kwargs_client_metadata_cannot_spoof_authenticated_identity(mock_request, mock_user_api_key_dict): + request: Final = mock_request() + passthrough_payload: Final = PassthroughStandardLoggingPayload( + url="https://test.com", + request_body={}, + ) + authenticated_key: Final = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + team_id="test-team", + end_user_id="test-user", + key_alias="real-key", + team_alias="Real Team", + user_email="real@example.com", + org_id="real-org", + ) + + result: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, + user_api_key_dict=authenticated_key, + passthrough_logging_payload=passthrough_payload, + litellm_call_id="test-call-id", + logging_obj=LiteLLMLoggingObj( + model="test-model", + messages=[], + stream=False, + call_type="test-call-type", + start_time=datetime(2026, 1, 15, 12, 0, 0), + litellm_call_id="test-call-id", + function_id="test-function-id", + ), + _parsed_body={ + "litellm_metadata": { + "user_api_key_org_id": "victim-org", + "user_api_key_end_user_id": "victim-end-user", + "user_api_key_user_id": "victim-user", + "user_api_key_team_id": "victim-team", + "user_api_key_team_alias": "Victim Team", + "user_api_key_alias": "victim-key", + "user_api_key_user_email": "victim@example.com", + } + }, + ) + + metadata: Final = result["litellm_params"]["metadata"] + assert metadata["user_api_key_user_id"] == "test-user" + assert metadata["user_api_key_team_id"] == "test-team" + assert metadata["user_api_key_team_alias"] == "Real Team" + assert metadata["user_api_key_alias"] == "real-key" + assert metadata["user_api_key_user_email"] == "real@example.com" + assert metadata["user_api_key_org_id"] == "real-org" + assert metadata["user_api_key_end_user_id"] == "test-user" + + +def test_init_kwargs_no_authenticated_identity_field_is_client_settable(mock_request, mock_user_api_key_dict): + authenticated_key: Final = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + team_id="test-team", + end_user_id="test-end-user", + key_alias="real-key", + team_alias="Real Team", + user_email="real@example.com", + org_id="real-org", + organization_alias="Real Org", + project_id="real-project", + project_alias="Real Project", + spend=1.5, + max_budget=10.0, + user_spend=2.5, + user_max_budget=20.0, + team_spend=3.5, + team_max_budget=30.0, + metadata={"real": "auth-metadata"}, + ) + expected: Final = dict( + LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=authenticated_key) + ) + assert len(expected) >= 20 + + spoofed: Final = {key: f"SPOOFED-{key}" for key in expected} + + result: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=mock_request(), + user_api_key_dict=authenticated_key, + passthrough_logging_payload=PassthroughStandardLoggingPayload(url="https://test.com", request_body={}), + litellm_call_id="test-call-id", + logging_obj=LiteLLMLoggingObj( + model="test-model", + messages=[], + stream=False, + call_type="test-call-type", + start_time=datetime(2026, 1, 15, 12, 0, 0), + litellm_call_id="test-call-id", + function_id="test-function-id", + ), + _parsed_body={"litellm_metadata": dict(spoofed), "metadata": dict(spoofed)}, + ) + + metadata: Final = result["litellm_params"]["metadata"] + survived: Final = {key: metadata.get(key) for key in expected if metadata.get(key) != expected[key]} + assert survived == {}, f"client-supplied values survived for: {sorted(survived)}" diff --git a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py index 8cc8dccd13e..2c1b5a4dbb5 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py @@ -1,3 +1,4 @@ +from typing import Final import asyncio import collections import datetime @@ -7636,6 +7637,115 @@ async def test_calculate_spend_unpriced_model_returns_400(): assert model in exc_info.value.message +@pytest.mark.asyncio +async def test_spend_calc_model_messages(): + cost_obj: Final = await spend_management_endpoints.calculate_spend( + request=SpendCalculateRequest( + model="gpt-3.5-turbo", + messages=[ + {"role": "user", "content": "What is the capital of France?"}, + ], + ) + ) + + cost: Final = cost_obj["cost"] + assert cost > 0.0 + + +@pytest.mark.asyncio +async def test_spend_calc_model_on_router_messages(monkeypatch): + temp_llm_router: Final = Router( + model_list=[ + { + "model_name": "special-llama-model", + "litellm_params": { + "model": "groq/openai/gpt-oss-20b", + }, + } + ] + ) + + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", temp_llm_router) + + cost_obj: Final = await spend_management_endpoints.calculate_spend( + request=SpendCalculateRequest( + model="special-llama-model", + messages=[ + {"role": "user", "content": "What is the capital of France?"}, + ], + ) + ) + + _cost: Final = cost_obj["cost"] + + assert _cost > 0.0 + + +@pytest.mark.asyncio +async def test_spend_calc_using_response(): + cost_obj: Final = await spend_management_endpoints.calculate_spend( + request=SpendCalculateRequest( + completion_response={ + "id": "chatcmpl-3bc7abcd-f70b-48ab-a16c-dfba0b286c86", + "choices": [ + { + "finish_reason": "stop", + "index": 0, + "message": { + "content": "Yooo! What's good?", + "role": "assistant", + }, + } + ], + "created": "1677652288", + "model": "groq/openai/gpt-oss-20b", + "object": "chat.completion", + "system_fingerprint": "fp_873a560973", + "usage": { + "completion_tokens": 8, + "prompt_tokens": 12, + "total_tokens": 20, + }, + } + ) + ) + + cost: Final = cost_obj["cost"] + assert cost > 0.0 + + +@pytest.mark.asyncio +async def test_spend_calc_model_alias_on_router_messages(monkeypatch): + temp_llm_router: Final = Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": { + "model": "gpt-4o", + }, + } + ], + model_group_alias={ + "gpt4o": "gpt-4o", + }, + ) + + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", temp_llm_router) + + cost_obj: Final = await spend_management_endpoints.calculate_spend( + request=SpendCalculateRequest( + model="gpt4o", + messages=[ + {"role": "user", "content": "What is the capital of France?"}, + ], + ) + ) + + _cost: Final = cost_obj["cost"] + + assert _cost > 0.0 + + def _admin_auth() -> UserAPIKeyAuth: return UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user") diff --git a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py index 9aa9110f7a4..4d299f05fbe 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py @@ -18,7 +18,7 @@ from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD, LITELLM_TRUNCATIO from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.llms.base_llm.ocr.transformation import OCRResponse, OCRUsageInfo from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.proxy._types import SpendLogsMetadataFields, SpendLogsPayload, UserAPIKeyAuth +from litellm.proxy._types import SpendLogsMetadata, SpendLogsMetadataFields, SpendLogsPayload, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.spend_tracking.spend_tracking_utils import ( @@ -6137,3 +6137,336 @@ class TestGetLoggingPayloadOCR: assert payload["spend"] == _OCR_RESPONSE_COST assert _additional_usage_values(payload)["pages_processed"] == 5 assert _additional_usage_values(payload)["doc_size_bytes"] == 1024 + + +@pytest.mark.parametrize( + "model_id", + ["chatcmpl-9XZmkzS1uPhRCoVdGQvBqqIbSgECt", "", None], +) +def test_spend_logs_payload(model_id: str | None): + + kwargs: Final = { + "model": "chatgpt-v-3", + "messages": [ + {"role": "system", "content": "you are a helpful assistant.\n"}, + {"role": "user", "content": "bom dia"}, + ], + "custom_llm_provider": "azure", + "optional_params": { + "stream": False, + "max_tokens": 10, + "user": "116544810872468347480", + "extra_body": {}, + }, + "litellm_params": { + "acompletion": True, + "api_key": "sk-test-mock-key-707", + "force_timeout": 600, + "logger_fn": None, + "verbose": False, + "custom_llm_provider": "azure", + "api_base": "https://openai-gpt-4-test-v-1.openai.azure.com//openai/", + "litellm_call_id": "b9929bf6-7b80-4c8c-b486-034e6ac0c8b7", + "model_alias_map": {}, + "completion_call_id": None, + "metadata": { + "tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], + "user_api_key": "sk-test-mock-api-key-123", + "user_api_key_alias": "custom-key-alias", + "user_api_end_user_max_budget": None, + "litellm_api_version": "0.0.0", + "global_max_parallel_requests": None, + "user_api_key_user_id": "116544810872468347480", + "user_api_key_org_id": "custom-org-id", + "user_api_key_team_id": "custom-team-id", + "user_api_key_team_alias": "custom-team-alias", + "user_api_key_metadata": {}, + "requester_ip_address": "127.0.0.1", + "spend_logs_metadata": {"hello": "world"}, + "headers": { + "content-type": "application/json", + "user-agent": "PostmanRuntime/7.32.3", + "accept": "*/*", + "postman-token": "92300061-eeaa-423b-a420-0b44896ecdc4", + "host": "localhost:4000", + "accept-encoding": "gzip, deflate, br", + "connection": "keep-alive", + "content-length": "163", + }, + "endpoint": "http://localhost:4000/chat/completions", + "model_group": "gpt-5-mini", + "deployment": "azure/gpt-4.1-mini", + "model_info": { + "id": "4bad40a1eb6bebd1682800f16f44b9f06c52a6703444c99c7f9f32e9de3693b4", + "db_model": False, + }, + "api_base": "https://openai-gpt-4-test-v-1.openai.azure.com/", + "caching_groups": None, + "error_information": None, + "status": "success", + "proxy_server_request": "{}", + "raw_request": "\n\nPOST Request Sent from LiteLLM:\ncurl -X POST \\\nhttps://openai-gpt-4-test-v-1.openai.azure.com//openai/ \\\n-H 'Authorization: *****' \\\n-d '{'model': 'chatgpt-v-3', 'messages': [{'role': 'system', 'content': 'you are a helpful assistant.\\n'}, {'role': 'user', 'content': 'bom dia'}], 'stream': False, 'max_tokens': 10, 'user': '116544810872468347480', 'extra_body': {}}'\n", + }, + "model_info": { + "id": "4bad40a1eb6bebd1682800f16f44b9f06c52a6703444c99c7f9f32e9de3693b4", + "db_model": False, + }, + "proxy_server_request": { + "url": "http://localhost:4000/chat/completions", + "method": "POST", + "headers": { + "content-type": "application/json", + "user-agent": "PostmanRuntime/7.32.3", + "accept": "*/*", + "postman-token": "92300061-eeaa-423b-a420-0b44896ecdc4", + "host": "localhost:4000", + "accept-encoding": "gzip, deflate, br", + "connection": "keep-alive", + "content-length": "163", + }, + "body": { + "messages": [ + { + "role": "system", + "content": "you are a helpful assistant.\n", + }, + {"role": "user", "content": "bom dia"}, + ], + "model": "gpt-5-mini", + "max_tokens": 10, + }, + }, + "preset_cache_key": None, + "no-log": False, + "stream_response": {}, + "input_cost_per_token": None, + "input_cost_per_second": None, + "output_cost_per_token": None, + "output_cost_per_second": None, + }, + "start_time": datetime.datetime(2024, 6, 7, 12, 43, 30, 307665), + "stream": False, + "user": "116544810872468347480", + "call_type": "acompletion", + "litellm_call_id": "b9929bf6-7b80-4c8c-b486-034e6ac0c8b7", + "completion_start_time": datetime.datetime(2024, 6, 7, 12, 43, 30, 954146), + "max_tokens": 10, + "extra_body": {}, + "input": [ + {"role": "system", "content": "you are a helpful assistant.\n"}, + {"role": "user", "content": "bom dia"}, + ], + "api_key": "1234", + "original_response": "", + "additional_args": { + "headers": {"Authorization": "Bearer 1234"}, + "api_base": "openai-gpt-4-test-v-1.openai.azure.com", + "acompletion": True, + "complete_input_dict": { + "model": "chatgpt-v-3", + "messages": [ + {"role": "system", "content": "you are a helpful assistant.\n"}, + {"role": "user", "content": "bom dia"}, + ], + "stream": False, + "max_tokens": 10, + "user": "116544810872468347480", + "extra_body": {}, + }, + }, + "log_event_type": "post_api_call", + "end_time": datetime.datetime(2024, 6, 7, 12, 43, 30, 954146), + "cache_hit": None, + "response_cost": 2.4999999999999998e-05, + "standard_logging_object": { + "request_tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], + "metadata": { + "user_api_key_end_user_id": "test-user", + }, + "model_map_information": { + "tpm": 1000, + "rpm": 1000, + }, + }, + } + response_obj: Final = litellm.ModelResponse( + id=model_id, + choices=[ + litellm.Choices( + finish_reason="length", + index=0, + message=litellm.Message(content="Bom dia! Como posso ajudar você", role="assistant"), + ) + ], + created=1717789410, + model="gpt-35-turbo", + object="chat.completion", + system_fingerprint=None, + usage=litellm.Usage(completion_tokens=10, prompt_tokens=20, total_tokens=30), + ) + + payload: Final[SpendLogsPayload] = get_logging_payload( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.datetime(2024, 6, 7, 12, 43, 30, 308604), + end_time=datetime.datetime(2024, 6, 7, 12, 43, 30, 954146), + ) + + assert len(payload["request_id"]) > 0 + expected_metadata_keys: Final = SpendLogsMetadata.__annotations__.keys() + + assert "metadata" in payload + assert isinstance(payload["metadata"], str) + metadata: Final[Mapping[str, object]] = json.loads(payload["metadata"]) + assert set(metadata.keys()) == set(expected_metadata_keys) + + assert payload["request_tags"] == '["model-anthropic-claude-v2.1", "app-ishaan-prod"]' + assert metadata["user_api_key_org_id"] == "custom-org-id" + assert metadata["user_api_key_team_id"] == "custom-team-id" + assert metadata["user_api_key_team_alias"] == "custom-team-alias" + assert metadata["user_api_key_alias"] == "custom-key-alias" + + assert payload["custom_llm_provider"] == "azure" + + +def test_spend_logs_payload_whisper(): + + kwargs: Final = { + "model": "whisper-1", + "messages": [{"role": "user", "content": "audio_file"}], + "optional_params": {}, + "litellm_params": { + "api_base": "", + "metadata": { + "user_api_key": "sk-test-mock-api-key-123", + "user_api_key_alias": None, + "user_api_key_end_user_id": "test-user", + "user_api_end_user_max_budget": None, + "litellm_api_version": "1.40.19", + "global_max_parallel_requests": None, + "user_api_key_user_id": "default_user_id", + "user_api_key_org_id": None, + "user_api_key_team_id": None, + "user_api_key_team_alias": None, + "user_api_key_team_max_budget": None, + "user_api_key_team_spend": None, + "user_api_key_spend": 0.0, + "user_api_key_max_budget": None, + "user_api_key_metadata": {}, + "headers": { + "host": "localhost:4000", + "user-agent": "curl/7.88.1", + "accept": "*/*", + "content-length": "775501", + "content-type": "multipart/form-data; boundary=------------------------21d518e191326d20", + }, + "endpoint": "http://localhost:4000/v1/audio/transcriptions", + "litellm_parent_otel_span": None, + "model_group": "whisper-1", + "deployment": "whisper-1", + "model_info": { + "id": "d7761582311451c34d83d65bc8520ce5c1537ea9ef2bec13383cf77596d49eeb", + "db_model": False, + }, + "caching_groups": None, + }, + }, + "start_time": datetime.datetime(2024, 6, 26, 14, 20, 11, 313291), + "stream": False, + "user": "", + "call_type": "atranscription", + "litellm_call_id": "05921cf7-33f9-421c-aad9-33310c1e2702", + "completion_start_time": datetime.datetime(2024, 6, 26, 14, 20, 13, 653149), + "stream_options": None, + "input": "tmp-requestc8640aee-7d85-49c3-b3ef-bdc9255d8e37.wav", + "original_response": '{"text": "Four score and seven years ago, our fathers brought forth on this continent a new nation, conceived in liberty and dedicated to the proposition that all men are created equal. Now we are engaged in a great civil war, testing whether that nation, or any nation so conceived and so dedicated, can long endure."}', + "additional_args": { + "complete_input_dict": { + "model": "whisper-1", + "file": "<_io.BufferedReader name='tmp-requestc8640aee-7d85-49c3-b3ef-bdc9255d8e37.wav'>", + "language": None, + "prompt": None, + "response_format": None, + "temperature": None, + } + }, + "log_event_type": "post_api_call", + "end_time": datetime.datetime(2024, 6, 26, 14, 20, 13, 653149), + "cache_hit": None, + "response_cost": 0.00023398580000000003, + } + + response: Final = litellm.utils.TranscriptionResponse( + text="Four score and seven years ago, our fathers brought forth on this continent a new nation, conceived in liberty and dedicated to the proposition that all men are created equal. Now we are engaged in a great civil war, testing whether that nation, or any nation so conceived and so dedicated, can long endure." + ) + + payload: Final[SpendLogsPayload] = get_logging_payload( + kwargs=kwargs, + response_obj=response, + start_time=datetime.datetime(2026, 1, 15, 12, 0, 0), + end_time=datetime.datetime(2026, 1, 15, 12, 0, 0), + ) + + assert payload["call_type"] == "atranscription" + assert payload["spend"] == 0.00023398580000000003 + + +def test_spend_logs_payload_with_prompts_enabled(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy.proxy_server import general_settings + + monkeypatch.setitem(general_settings, "store_prompts_in_spend_logs", True) + + kwargs: Final = { + "model": "gpt-5-mini", + "messages": [{"role": "user", "content": "Hello!"}], + "litellm_params": { + "proxy_server_request": { + "body": { + "model": "gpt-5.5", + "messages": [{"role": "user", "content": "Hello!"}], + } + } + }, + "standard_logging_object": { + "messages": [{"role": "user", "content": "Hello!"}], + "response": {"role": "assistant", "content": "Hi there!"}, + "metadata": { + "user_api_key_end_user_id": "test-user", + }, + "request_tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], + "model_map_information": { + "tpm": 1000, + "rpm": 1000, + }, + }, + } + response_obj: Final = litellm.ModelResponse( + id="chatcmpl-123", + choices=[ + litellm.Choices( + finish_reason="stop", + index=0, + message=litellm.Message(content="Hi there!", role="assistant"), + ) + ], + model="gpt-5-mini", + usage=litellm.Usage(completion_tokens=2, prompt_tokens=1, total_tokens=3), + ) + logged_at: Final = datetime.datetime(2026, 1, 15, 12, 0, 0) + + payload: Final[SpendLogsPayload] = get_logging_payload( + kwargs=kwargs, response_obj=response_obj, start_time=logged_at, end_time=logged_at + ) + + assert payload["response"] == json.dumps({"role": "assistant", "content": "Hi there!"}) + proxy_server_request: Final = json.loads(payload["proxy_server_request"] or "{}") + assert proxy_server_request["model"] == "gpt-5.5" + assert proxy_server_request["messages"] == [{"role": "user", "content": "Hello!"}] + + monkeypatch.setitem(general_settings, "store_prompts_in_spend_logs", False) + + payload_disabled: Final[SpendLogsPayload] = get_logging_payload( + kwargs=kwargs, response_obj=response_obj, start_time=logged_at, end_time=logged_at + ) + assert payload_disabled["messages"] == "{}" + assert payload_disabled["response"] == "{}"