mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
3250ccc802
commit
5e1c5c0bd1
26 changed files with 1649 additions and 2741 deletions
|
|
@ -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** |
|
||||
|
||||
|
|
|
|||
|
|
@ -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) |
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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!")
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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"] == "{}"
|
||||
|
|
@ -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}
|
||||
)
|
||||
|
|
@ -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)}"
|
||||
|
|
@ -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"
|
||||
|
|
@ -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__])
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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"] == "{}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue