test: move offline proxy auth, hook, spend and pass-through tests to tests/unit (#45568)

* test: move offline proxy auth, guardrail hook, spend, budget reset, GCS payload and pass-through tests into tests/unit

Relocate 53 legacy tests that pass with no network or keys into the tests/unit
files that mirror the code they exercise, drop 2 llm_guard tests already covered
by test_llm_guard_call_type_aliases, and delete the emptied legacy files.

* test: freeze the reset clock and run the real reset methods in the moved budget reset tests

Pin datetime in reset_budget_job and timezone_utils to a fixed instant and await
the service-hook tasks the job spawns instead of sleeping. Drive the per-user
failure with a malformed budget_duration row from the prisma mock so the real
_reset_budget_for_key/user/team run instead of patched replacements.

* test: use fixed timestamps in the moved pass-through and spend log tests

* test: mock the moderation HTTP call and type the moved moderation and spend log tests

Serve the OpenAI moderation response through an httpx MockTransport so the real
async_make_request runs, annotate the moved tests' locals with Final, and pass
get_logging_payload typed kwargs instead of an untyped input_args dict.

* test: drive the moved proxy tests through boundaries and annotate their locals with Final

Serve GCS downloads and the Router moderation call through httpx mocks with fake
Google credentials, let the real pass-through success handler fail on the mocked
upstream response, and fold the Vertex live route endpoint check into the existing
route test. Annotate moved locals with Final, build the budget reset rows without
rebinding, and type the remaining fakes without **kwargs
This commit is contained in:
yuneng-jiang 2026-10-09 09:49:13 -07:00 • committed by GitHub
parent 3250ccc802
commit 5e1c5c0bd1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 1649 additions and 2741 deletions

View file

@ -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** |

View file

@ -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) |

View file

@ -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)
# ---------------------------------------------------------------------------

View file

@ -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

View file

@ -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

View file

@ -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!")

View file

@ -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):

View file

@ -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)

View file

@ -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"] == "{}"

View file

@ -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}
)

View file

@ -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)}"

View file

@ -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"

View file

@ -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__])

View file

@ -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

View file

@ -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",
}

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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():
"""

View file

@ -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)

View file

@ -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 ────────────────────────────────────────────────────

View file

@ -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"

View file

@ -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"

View file

@ -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)}"

View file

@ -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")

View file

@ -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"] == "{}"