[Bug Fix] Create/Update team member api 500 errror (#10479)

* Fixes issue with team_endpoints on member budget update

* refactored location of budget membership fix

* added test for _upsert_budget_membership func
This commit is contained in:
hagan 2025-06-03 15:57:33 -07:00 • committed by GitHub
parent 7c1e1790e3
commit 0f449bf038
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 444 additions and 80 deletions

View file

@ -2444,7 +2444,7 @@ class CreatePassThroughEndpoint(LiteLLMPydanticObjectBase):
class LiteLLM_TeamMembership(LiteLLMPydanticObjectBase):
user_id: str
team_id: str
budget_id: str
budget_id: Optional[str] = None
litellm_budget_table: Optional[LiteLLM_BudgetTable]

View file

@ -1,4 +1,4 @@
from typing import Any, Union
from typing import Any, Union, Optional
from litellm.proxy._types import (
GenerateKeyRequest,
@ -46,3 +46,75 @@ def _set_object_metadata_field(
_premium_user_check()
object_data.metadata = object_data.metadata or {}
object_data.metadata[field_name] = value
async def _upsert_budget_and_membership(
tx,
*,
team_id: str,
user_id: str,
max_budget: Optional[float],
existing_budget_id: Optional[str],
user_api_key_dict: UserAPIKeyAuth,
):
"""
Helper function to Create/Update or Delete the budget within the team membership
Args:
tx: The transaction object
team_id: The ID of the team
user_id: The ID of the user
max_budget: The maximum budget for the team
existing_budget_id: The ID of the existing budget, if any
user_api_key_dict: User API Key dictionary containing user information
If max_budget is None, the user's budget is removed from the team membership.
If max_budget exists, a budget is updated or created and linked to the team membership.
"""
if max_budget is None:
# disconnect the budget since max_budget is None
await tx.litellm_teammembership.update(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
data={"litellm_budget_table": {"disconnect": True}},
)
return
if existing_budget_id:
# update the existing budget
await tx.litellm_budgettable.update(
where={"budget_id": existing_budget_id},
data={"max_budget": max_budget},
)
return
# create a new budget
new_budget = await tx.litellm_budgettable.create(
data={
"max_budget": max_budget,
"created_by": user_api_key_dict.user_id or "",
"updated_by": user_api_key_dict.user_id or "",
},
include={"team_membership": True},
)
# upsert the team membership with the new/updated budget
await tx.litellm_teammembership.upsert(
where={
"user_id_team_id": {
"user_id": user_id,
"team_id": team_id,
}
},
data={
"create": {
"user_id": user_id,
"team_id": team_id,
"litellm_budget_table": {
"connect": {"budget_id": new_budget.budget_id},
},
},
"update": {
"litellm_budget_table": {
"connect": {"budget_id": new_budget.budget_id},
},
},
},
)

View file

@ -66,6 +66,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
_set_object_metadata_field,
_user_has_admin_view,
_upsert_budget_and_membership,
)
from litellm.proxy.management_endpoints.tag_management_endpoints import (
get_daily_activity,
@ -107,19 +108,22 @@ def _is_available_team(team_id: str, user_api_key_dict: UserAPIKeyAuth) -> bool:
async def get_all_team_memberships(
prisma_client: PrismaClient, team_id: List[str], user_id: Optional[str] = None
prisma_client: PrismaClient, team_ids: List[str], user_id: Optional[str] = None
) -> List[LiteLLM_TeamMembership]:
"""Get all team memberships for a given user"""
## GET ALL MEMBERSHIPS ##
if not isinstance(user_id, str):
user_id = str(user_id)
where_obj: Dict[str, Dict[str, List[str]]] = {
"team_id": {"in": team_ids}
}
if user_id is not None:
where_obj["user_id"] = {"in": [user_id]}
# if user_id is None:
# where_obj = {"team_id": {"in": team_id}}
# else:
# where_obj = {"user_id": str(user_id), "team_id": {"in": team_id}}
team_memberships = await prisma_client.db.litellm_teammembership.find_many(
where=(
{"user_id": user_id, "team_id": {"in": team_id}}
if user_id is not None
else {"team_id": {"in": team_id}}
),
where=where_obj,
include={"litellm_budget_table": True},
)
@ -165,7 +169,7 @@ async def new_team( # noqa: PLR0915
- budget_duration: Optional[str] - The duration of the budget for the team. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets)
- models: Optional[list] - A list of models associated with the team - all keys for this team_id will have at most, these models. If empty, assumes all models are allowed.
- blocked: bool - Flag indicating if the team is blocked or not - will stop all calls from keys with this team_id.
- members: Optional[List] - Control team members via `/team/member/add` and `/team/member/delete`.
- members: Optional[List] - Control team members via `/team/member/add` and `/team/member/delete`.
- tags: Optional[List[str]] - Tags for [tracking spend](https://litellm.vercel.app/docs/proxy/enterprise#tracking-spend-for-custom-tags) and/or doing [tag-based routing](https://litellm.vercel.app/docs/proxy/tag_routing).
- organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`.
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
@ -197,8 +201,8 @@ async def new_team( # noqa: PLR0915
--header 'Authorization: Bearer sk-1234' \
--header 'Content-Type: application/json' \
--data '{
"team_alias": "QA Prod Bot",
"max_budget": 0.000000001,
"team_alias": "QA Prod Bot",
"max_budget": 0.000000001,
"budget_duration": "1d"
}'
```
@ -555,7 +559,7 @@ async def update_team(
),
):
"""
Use `/team/member_add` AND `/team/member/delete` to add/remove new team members
Use `/team/member_add` AND `/team/member/delete` to add/remove new team members
You can now update team budget / rate limits via /team/update
@ -1262,28 +1266,15 @@ async def team_member_update(
break
### upsert new budget
if data.max_budget_in_team is not None:
if identified_budget_id is None:
new_budget = await prisma_client.db.litellm_budgettable.create(
data={
"max_budget": data.max_budget_in_team,
"created_by": user_api_key_dict.user_id or "",
"updated_by": user_api_key_dict.user_id or "",
}
)
await prisma_client.db.litellm_teammembership.create(
data={
"team_id": data.team_id,
"user_id": received_user_id,
"budget_id": new_budget.budget_id,
},
)
elif identified_budget_id is not None:
await prisma_client.db.litellm_budgettable.update(
where={"budget_id": identified_budget_id},
data={"max_budget": data.max_budget_in_team},
)
async with prisma_client.db.tx() as tx:
await _upsert_budget_and_membership(
tx=tx,
team_id=data.team_id,
user_id=received_user_id,
max_budget=data.max_budget_in_team,
existing_budget_id=identified_budget_id,
user_api_key_dict=user_api_key_dict,
)
### update team member role
if data.role is not None:

View file

@ -16,7 +16,7 @@ model LiteLLM_BudgetTable {
tpm_limit BigInt?
rpm_limit BigInt?
model_max_budget Json?
budget_duration String?
budget_duration String?
budget_reset_at DateTime?
created_at DateTime @default(now()) @map("created_at")
created_by String
@ -25,8 +25,8 @@ model LiteLLM_BudgetTable {
organization LiteLLM_OrganizationTable[] // multiple orgs can have the same budget
keys LiteLLM_VerificationToken[] // multiple keys can have the same budget
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
}
// Models on proxy
@ -34,7 +34,7 @@ model LiteLLM_CredentialsTable {
credential_id String @id @default(uuid())
credential_name String @unique
credential_values Json
credential_info Json?
credential_info Json?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
@ -44,9 +44,9 @@ model LiteLLM_CredentialsTable {
// Models on proxy
model LiteLLM_ProxyModelTable {
model_id String @id @default(uuid())
model_name String
model_name String
litellm_params Json
model_info Json?
model_info Json?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
@ -67,7 +67,7 @@ model LiteLLM_OrganizationTable {
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
updated_by String
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
teams LiteLLM_TeamTable[]
teams LiteLLM_TeamTable[]
users LiteLLM_UserTable[]
keys LiteLLM_VerificationToken[]
members LiteLLM_OrganizationMembership[] @relation("OrganizationToMembership")
@ -86,10 +86,10 @@ model LiteLLM_ModelTable {
}
// Assign prod keys to groups, not individuals
// Assign prod keys to groups, not individuals
model LiteLLM_TeamTable {
team_id String @id @default(uuid())
team_alias String?
team_alias String?
organization_id String?
object_permission_id String?
admins String[]
@ -102,7 +102,7 @@ model LiteLLM_TeamTable {
max_parallel_requests Int?
tpm_limit BigInt?
rpm_limit BigInt?
budget_duration String?
budget_duration String?
budget_reset_at DateTime?
blocked Boolean @default(false)
created_at DateTime @default(now()) @map("created_at")
@ -119,7 +119,7 @@ model LiteLLM_TeamTable {
// Track spend, rate limit, budget Users
model LiteLLM_UserTable {
user_id String @id
user_alias String?
user_alias String?
team_id String?
sso_user_id String? @unique
organization_id String?
@ -135,7 +135,7 @@ model LiteLLM_UserTable {
max_parallel_requests Int?
tpm_limit BigInt?
rpm_limit BigInt?
budget_duration String?
budget_duration String?
budget_reset_at DateTime?
allowed_cache_controls String[] @default([])
model_spend Json @default("{}")
@ -156,14 +156,13 @@ model LiteLLM_ObjectPermissionTable {
object_permission_id String @id @default(uuid())
mcp_servers String[] @default([])
vector_stores String[] @default([])
teams LiteLLM_TeamTable[]
verification_tokens LiteLLM_VerificationToken[]
organizations LiteLLM_OrganizationTable[]
users LiteLLM_UserTable[]
}
// Holds the MCP server configuration
// Holds the MCP server configuration
model LiteLLM_MCPServerTable {
server_id String @id @default(uuid())
alias String?
@ -171,7 +170,7 @@ model LiteLLM_MCPServerTable {
url String
transport String @default("sse")
spec_version String @default("2025-03-26")
auth_type String?
auth_type String?
created_at DateTime? @default(now()) @map("created_at")
created_by String?
updated_at DateTime? @default(now()) @updatedAt @map("updated_at")
@ -197,8 +196,8 @@ model LiteLLM_VerificationToken {
blocked Boolean?
tpm_limit BigInt?
rpm_limit BigInt?
max_budget Float?
budget_duration String?
max_budget Float?
budget_duration String?
budget_reset_at DateTime?
allowed_cache_controls String[] @default([])
allowed_routes String[] @default([])
@ -255,7 +254,7 @@ model LiteLLM_SpendLogs {
cache_hit String? @default("")
cache_key String? @default("")
request_tags Json? @default("[]")
team_id String?
team_id String?
end_user String?
requester_ip_address String?
messages Json? @default("{}")
@ -273,7 +272,7 @@ model LiteLLM_ErrorLogs {
request_id String @id @default(uuid())
startTime DateTime // Assuming start_time is a DateTime field
endTime DateTime // Assuming end_time is a DateTime field
api_base String @default("")
api_base String @default("")
model_group String @default("") // public model_name / model_group
litellm_model_name String @default("") // model passed to litellm
model_id String @default("") // ID of model in ProxyModelTable
@ -286,7 +285,7 @@ model LiteLLM_ErrorLogs {
// Beta - allow team members to request access to a model
model LiteLLM_UserNotifications {
request_id String @id
user_id String
user_id String
models String[]
justification String
status String // approved, disapproved, pending
@ -298,7 +297,7 @@ model LiteLLM_TeamMembership {
team_id String
spend Float @default(0.0)
budget_id String?
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id], onDelete: Cascade, onUpdate: Cascade)
@@id([user_id, team_id])
}
@ -316,8 +315,8 @@ model LiteLLM_OrganizationMembership {
user LiteLLM_UserTable @relation(fields: [user_id], references: [user_id])
organization LiteLLM_OrganizationTable @relation("OrganizationToMembership", fields: [organization_id], references: [organization_id])
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
@@id([user_id, organization_id])
@@unique([user_id, organization_id])
@ -350,19 +349,19 @@ model LiteLLM_AuditLog {
action String // create, update, delete
table_name String // on of LitellmTableNames.TEAM_TABLE_NAME, LitellmTableNames.USER_TABLE_NAME, LitellmTableNames.PROXY_MODEL_TABLE_NAME,
object_id String // id of the object being audited. This can be the key id, team id, user id, model id
before_value Json? // value of the row
before_value Json? // value of the row
updated_values Json? // value of the row after change
}
// Track daily user spend metrics per model and key
model LiteLLM_DailyUserSpend {
id String @id @default(uuid())
user_id String?
user_id String?
date String
api_key String
model String
model_group String?
custom_llm_provider String?
api_key String
model String
model_group String?
custom_llm_provider String?
prompt_tokens BigInt @default(0)
completion_tokens BigInt @default(0)
cache_read_input_tokens BigInt @default(0)
@ -386,10 +385,10 @@ model LiteLLM_DailyTeamSpend {
id String @id @default(uuid())
team_id String?
date String
api_key String
model String
model_group String?
custom_llm_provider String?
api_key String
model String
model_group String?
custom_llm_provider String?
prompt_tokens BigInt @default(0)
completion_tokens BigInt @default(0)
cache_read_input_tokens BigInt @default(0)
@ -411,12 +410,12 @@ model LiteLLM_DailyTeamSpend {
// Track daily team spend metrics per model and key
model LiteLLM_DailyTagSpend {
id String @id @default(uuid())
tag String?
tag String?
date String
api_key String
model String
model_group String?
custom_llm_provider String?
api_key String
model String
model_group String?
custom_llm_provider String?
prompt_tokens BigInt @default(0)
completion_tokens BigInt @default(0)
cache_read_input_tokens BigInt @default(0)
@ -454,26 +453,26 @@ model LiteLLM_ManagedFileTable {
id String @id @default(uuid())
unified_file_id String @unique // The base64 encoded unified file ID
file_object Json // Stores the OpenAIFileObject
model_mappings Json
model_mappings Json
flat_model_file_ids String[] @default([]) // Flat list of model file id's - for faster querying of model id -> unified file id
created_at DateTime @default(now())
created_by String?
created_by String?
updated_at DateTime @updatedAt
updated_by String?
@@index([unified_file_id])
}
model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use the
model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use the
id String @id @default(uuid())
unified_object_id String @unique // The base64 encoded unified file ID
model_object_id String @unique // the id returned by the backend API provider
model_object_id String @unique // the id returned by the backend API provider
file_object Json // Stores the OpenAIFileObject
file_purpose String // either 'batch' or 'fine-tune'
file_purpose String // either 'batch' or 'fine-tune'
created_at DateTime @default(now())
created_by String?
updated_at DateTime @updatedAt
updated_by String?
updated_by String?
@@index([unified_object_id])
@@index([model_object_id])

View file

@ -0,0 +1,170 @@
# tests/litellm/proxy/common_utils/test_upsert_budget_membership.py
import types
import pytest
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.management_endpoints.common_utils import (
_upsert_budget_and_membership,
)
# ---------------------------------------------------------------------------
# Fixtures: a fake Prisma transaction and a fake UserAPIKeyAuth object
# ---------------------------------------------------------------------------
@pytest.fixture
def mock_tx():
"""
Builds an object that looks just enough like the Prisma tx you use
inside _upsert_budget_and_membership.
"""
# membership “table”
membership = MagicMock()
membership.update = AsyncMock()
membership.upsert = AsyncMock()
# budget “table”
budget = MagicMock()
budget.update = AsyncMock()
# budget.create returns a fake row that has .budget_id
budget.create = AsyncMock(
return_value=types.SimpleNamespace(budget_id="new-budget-123")
)
tx = MagicMock()
tx.litellm_teammembership = membership
tx.litellm_budgettable = budget
return tx
@pytest.fixture
def fake_user():
"""Cheap stand-in for UserAPIKeyAuth."""
return types.SimpleNamespace(user_id="tester@example.com")
# TEST: max_budget is None, disconnect only
@pytest.mark.asyncio
async def test_upsert_disconnect(mock_tx, fake_user):
await _upsert_budget_and_membership(
mock_tx,
team_id="team-1",
user_id="user-1",
max_budget=None,
existing_budget_id=None,
user_api_key_dict=fake_user,
)
mock_tx.litellm_teammembership.update.assert_awaited_once_with(
where={"user_id_team_id": {"user_id": "user-1", "team_id": "team-1"}},
data={"litellm_budget_table": {"disconnect": True}},
)
mock_tx.litellm_budgettable.update.assert_not_called()
mock_tx.litellm_budgettable.create.assert_not_called()
mock_tx.litellm_teammembership.upsert.assert_not_called()
# TEST: existing budget id, update only
@pytest.mark.asyncio
async def test_upsert_update_existing(mock_tx, fake_user):
await _upsert_budget_and_membership(
mock_tx,
team_id="team-2",
user_id="user-2",
max_budget=42.0,
existing_budget_id="bud-999",
user_api_key_dict=fake_user,
)
mock_tx.litellm_budgettable.update.assert_awaited_once_with(
where={"budget_id": "bud-999"},
data={"max_budget": 42.0},
)
mock_tx.litellm_teammembership.update.assert_not_called()
mock_tx.litellm_budgettable.create.assert_not_called()
mock_tx.litellm_teammembership.upsert.assert_not_called()
# TEST: create new budget and link membership
@pytest.mark.asyncio
async def test_upsert_create_and_link(mock_tx, fake_user):
await _upsert_budget_and_membership(
mock_tx,
team_id="team-3",
user_id="user-3",
max_budget=99.9,
existing_budget_id=None,
user_api_key_dict=fake_user,
)
mock_tx.litellm_budgettable.create.assert_awaited_once_with(
data={
"max_budget": 99.9,
"created_by": fake_user.user_id,
"updated_by": fake_user.user_id,
},
include={"team_membership": True},
)
# Budget ID returned by the mocked create()
bid = mock_tx.litellm_budgettable.create.return_value.budget_id
mock_tx.litellm_teammembership.upsert.assert_awaited_once_with(
where={"user_id_team_id": {"user_id": "user-3", "team_id": "team-3"}},
data={
"create": {
"user_id": "user-3",
"team_id": "team-3",
"litellm_budget_table": {"connect": {"budget_id": bid}},
},
"update": {
"litellm_budget_table": {"connect": {"budget_id": bid}},
},
},
)
mock_tx.litellm_teammembership.update.assert_not_called()
mock_tx.litellm_budgettable.update.assert_not_called()
# TEST: create new budget and link membership, then update
@pytest.mark.asyncio
async def test_upsert_create_then_update(mock_tx, fake_user):
# FIRST CALL – create new budget and link membership
await _upsert_budget_and_membership(
mock_tx,
team_id="team-42",
user_id="user-42",
max_budget=10.0,
existing_budget_id=None,
user_api_key_dict=fake_user,
)
# capture the budget id that create() returned
created_bid = mock_tx.litellm_budgettable.create.return_value.budget_id
# sanity: we really did the create + upsert path
mock_tx.litellm_budgettable.create.assert_awaited_once()
mock_tx.litellm_teammembership.upsert.assert_awaited_once()
# SECOND CALL – pretend the same membership already exists, and
# reset call history so the next assertions are clear
mock_tx.litellm_budgettable.create.reset_mock()
mock_tx.litellm_teammembership.upsert.reset_mock()
mock_tx.litellm_budgettable.update.reset_mock()
await _upsert_budget_and_membership(
mock_tx,
team_id="team-42",
user_id="user-42",
max_budget=25.0, # new limit
existing_budget_id=created_bid, # now we say it exists
user_api_key_dict=fake_user,
)
# Now we expect ONLY an update to fire
mock_tx.litellm_budgettable.update.assert_awaited_once_with(
where={"budget_id": created_bid},
data={"max_budget": 25.0},
)
mock_tx.litellm_budgettable.create.assert_not_called()
mock_tx.litellm_teammembership.upsert.assert_not_called()

View file

@ -0,0 +1,132 @@
# tests/test_budget_endpoints.py
import os
import sys
import types
import pytest
from unittest.mock import AsyncMock, MagicMock
from fastapi.testclient import TestClient
import litellm.proxy.proxy_server as ps
from litellm.proxy.proxy_server import app
from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles, CommonProxyErrors
import litellm.proxy.management_endpoints.budget_management_endpoints as bm
sys.path.insert(
0, os.path.abspath("../../../")
) # Adds the parent directory to the system path
@pytest.fixture
def client_and_mocks(monkeypatch):
# Setup MagicMock Prisma
mock_prisma = MagicMock()
mock_table = MagicMock()
mock_table.create = AsyncMock(side_effect=lambda *, data: data)
mock_table.update = AsyncMock(side_effect=lambda *, where, data: {**where, **data})
mock_prisma.db = types.SimpleNamespace(
litellm_budgettable = mock_table,
litellm_dailyspend = mock_table,
)
# Monkeypatch Mocked Prisma client into the server module
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
# override returned auth user
fake_user = UserAPIKeyAuth(
user_id="test_user",
user_role=LitellmUserRoles.INTERNAL_USER,
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: fake_user
client = TestClient(app)
yield client, mock_prisma, mock_table
# teardown
app.dependency_overrides.clear()
monkeypatch.setattr(ps, "prisma_client", ps.prisma_client)
@pytest.mark.asyncio
async def test_new_budget_success(client_and_mocks):
client, _, mock_table = client_and_mocks
# Call /budget/new endpoint
payload = {
"budget_id": "budget_123",
"max_budget": 42.0,
"budget_duration": "30d",
}
resp = client.post("/budget/new", json=payload)
assert resp.status_code == 200, resp.text
body = resp.json()
assert body["budget_id"] == payload["budget_id"]
assert body["max_budget"] == payload["max_budget"]
assert body["budget_duration"] == payload["budget_duration"]
assert body["created_by"] == "test_user"
assert body["updated_by"] == "test_user"
mock_table.create.assert_awaited_once()
@pytest.mark.asyncio
async def test_new_budget_db_not_connected(client_and_mocks, monkeypatch):
client, mock_prisma, mock_table = client_and_mocks
# override the prisma_client that the handler imports at runtime
import litellm.proxy.proxy_server as ps
monkeypatch.setattr(ps, "prisma_client", None)
# Call /budget/new endpoint
resp = client.post("/budget/new", json={"budget_id": "no_db", "max_budget": 1.0})
assert resp.status_code == 500
detail = resp.json()["detail"]
assert detail["error"] == CommonProxyErrors.db_not_connected_error.value
@pytest.mark.asyncio
async def test_update_budget_success(client_and_mocks, monkeypatch):
client, mock_prisma, mock_table = client_and_mocks
payload = {
"budget_id": "budget_456",
"max_budget": 99.0,
"soft_budget": 50.0,
}
resp = client.post("/budget/update", json=payload)
assert resp.status_code == 200, resp.text
body = resp.json()
assert body["budget_id"] == payload["budget_id"]
assert body["max_budget"] == payload["max_budget"]
assert body["soft_budget"] == payload["soft_budget"]
assert body["updated_by"] == "test_user"
@pytest.mark.asyncio
async def test_update_budget_missing_id(client_and_mocks, monkeypatch):
client, mock_prisma, mock_table = client_and_mocks
payload = {"max_budget": 10.0}
resp = client.post("/budget/update", json=payload)
assert resp.status_code == 400, resp.text
detail = resp.json()["detail"]
assert detail["error"] == "budget_id is required"
@pytest.mark.asyncio
async def test_update_budget_db_not_connected(client_and_mocks, monkeypatch):
client, mock_prisma, mock_table = client_and_mocks
# override the prisma_client that the handler imports at runtime
import litellm.proxy.proxy_server as ps
monkeypatch.setattr(ps, "prisma_client", None)
payload = {"budget_id": "any", "max_budget": 1.0}
resp = client.post("/budget/update", json=payload)
assert resp.status_code == 500
detail = resp.json()["detail"]
assert detail["error"] == CommonProxyErrors.db_not_connected_error.value