fix(keys): write soft budget and key row in one transaction

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-04 00:41:12 +00:00
parent 2102d01571
commit 1783eb3402
2 changed files with 141 additions and 65 deletions

View file

@ -2125,7 +2125,7 @@ def _validate_soft_budget_value(soft_budget: float | None) -> None:
async def _update_key_soft_budget(
prisma_client: PrismaClient,
db: "Prisma",
existing_key_row: LiteLLM_VerificationToken,
soft_budget: float | None,
changed_by: str,
@ -2137,7 +2137,7 @@ async def _update_key_soft_budget(
"updated_by": changed_by,
}
budget_where: Final[prisma_types.LiteLLM_BudgetTableWhereUniqueInput] = {"budget_id": existing_budget_id}
await BudgetRepository(prisma_client).table.update(where=budget_where, data=budget_update)
await db.litellm_budgettable.update(where=budget_where, data=budget_update)
return existing_budget_id
if soft_budget is None:
return None
@ -2146,24 +2146,20 @@ async def _update_key_soft_budget(
"created_by": changed_by,
"updated_by": changed_by,
}
created_budget: Final[prisma_models.LiteLLM_BudgetTable] = await BudgetRepository(prisma_client).table.create(
data=budget_create
)
created_budget: Final[prisma_models.LiteLLM_BudgetTable] = await db.litellm_budgettable.create(data=budget_create)
return created_budget.budget_id
async def _apply_soft_budget_update(
data: UpdateKeyRequest,
non_default_values: Mapping[str, object],
prisma_client: PrismaClient,
db: "Prisma",
existing_key_row: LiteLLM_VerificationToken,
changed_by: str,
) -> Mapping[str, object]:
if "soft_budget" not in data.model_fields_set:
return non_default_values
remaining: Final = MappingProxyType({k: v for k, v in non_default_values.items() if k != "soft_budget"})
updated_budget_id: Final = await _update_key_soft_budget(
prisma_client=prisma_client,
db=db,
existing_key_row=existing_key_row,
soft_budget=data.soft_budget,
changed_by=changed_by,
@ -2173,6 +2169,30 @@ async def _apply_soft_budget_update(
return remaining
async def _update_key_row_with_soft_budget(
prisma_client: PrismaClient,
key: str,
data: UpdateKeyRequest,
non_default_values: Mapping[str, object],
existing_key_row: LiteLLM_VerificationToken,
changed_by: str,
) -> dict[str, object]:
hashed_token: Final = _hash_token_if_needed(key)
async with prisma_client.tx() as tx:
update_values: Final = await _apply_soft_budget_update(
data=data,
non_default_values=non_default_values,
db=tx,
existing_key_row=existing_key_row,
changed_by=changed_by,
)
updated_row: Final = await tx.litellm_verificationtoken.update(
where={"token": hashed_token},
data=with_settings_updated_at(prisma_client.jsonify_object({**update_values, "token": hashed_token})),
)
return {"token": hashed_token, "data": updated_row.model_dump() if updated_row is not None else {}}
async def prepare_key_update_data(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
@ -3043,17 +3063,20 @@ async def update_key_fn(
if prisma_client is None:
raise Exception("Not connected to DB!")
update_values: Final = await _apply_soft_budget_update(
data=data,
non_default_values=non_default_values,
prisma_client=prisma_client,
existing_key_row=existing_key_row,
changed_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name
response: Final = (
await _update_key_row_with_soft_budget(
prisma_client=prisma_client,
key=key,
data=data,
non_default_values=non_default_values,
existing_key_row=existing_key_row,
changed_by=changed_by,
)
if "soft_budget" in data.model_fields_set
else await prisma_client.update_data(token=key, data={**non_default_values, "token": key})
)
_data: Final = {**update_values, "token": key}
response: Final = await prisma_client.update_data(token=key, data=_data)
# Delete - key from cache, since it's been updated!
# key updated - a new model could have been added to this key. it should not block requests after this is done
await _delete_cache_key_object(

View file

@ -17595,18 +17595,19 @@ async def test_update_key_soft_budget_updates_existing_budget_row():
)
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id="budget-123")
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_budgettable.update = AsyncMock()
mock_db = MagicMock()
mock_db.litellm_budgettable.update = AsyncMock()
mock_db.litellm_budgettable.create = AsyncMock()
result = await _update_key_soft_budget(
prisma_client=mock_prisma_client,
db=mock_db,
existing_key_row=existing_key,
soft_budget=25.0,
changed_by="user-1",
)
assert result == "budget-123"
mock_prisma_client.db.litellm_budgettable.update.assert_awaited_once_with(
mock_db.litellm_budgettable.update.assert_awaited_once_with(
where={"budget_id": "budget-123"},
data={"soft_budget": 25.0, "updated_by": "user-1"},
)
@ -17619,18 +17620,19 @@ async def test_update_key_soft_budget_clears_existing_budget_row():
)
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id="budget-123")
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_budgettable.update = AsyncMock()
mock_db = MagicMock()
mock_db.litellm_budgettable.update = AsyncMock()
mock_db.litellm_budgettable.create = AsyncMock()
result = await _update_key_soft_budget(
prisma_client=mock_prisma_client,
db=mock_db,
existing_key_row=existing_key,
soft_budget=None,
changed_by="user-1",
)
assert result == "budget-123"
mock_prisma_client.db.litellm_budgettable.update.assert_awaited_once_with(
mock_db.litellm_budgettable.update.assert_awaited_once_with(
where={"budget_id": "budget-123"},
data={"soft_budget": None, "updated_by": "user-1"},
)
@ -17645,18 +17647,19 @@ async def test_update_key_soft_budget_creates_budget_row_when_key_has_none():
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None)
created_row = MagicMock()
created_row.budget_id = "budget-new"
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_budgettable.create = AsyncMock(return_value=created_row)
mock_db = MagicMock()
mock_db.litellm_budgettable.create = AsyncMock(return_value=created_row)
mock_db.litellm_budgettable.update = AsyncMock()
result = await _update_key_soft_budget(
prisma_client=mock_prisma_client,
db=mock_db,
existing_key_row=existing_key,
soft_budget=10.5,
changed_by="user-1",
)
assert result == "budget-new"
mock_prisma_client.db.litellm_budgettable.create.assert_awaited_once_with(
mock_db.litellm_budgettable.create.assert_awaited_once_with(
data={"soft_budget": 10.5, "created_by": "user-1", "updated_by": "user-1"}
)
@ -17668,20 +17671,20 @@ async def test_update_key_soft_budget_noop_when_clearing_without_budget_row():
)
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None)
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_budgettable.create = AsyncMock()
mock_prisma_client.db.litellm_budgettable.update = AsyncMock()
mock_db = MagicMock()
mock_db.litellm_budgettable.create = AsyncMock()
mock_db.litellm_budgettable.update = AsyncMock()
result = await _update_key_soft_budget(
prisma_client=mock_prisma_client,
db=mock_db,
existing_key_row=existing_key,
soft_budget=None,
changed_by="user-1",
)
assert result is None
mock_prisma_client.db.litellm_budgettable.create.assert_not_awaited()
mock_prisma_client.db.litellm_budgettable.update.assert_not_awaited()
mock_db.litellm_budgettable.create.assert_not_awaited()
mock_db.litellm_budgettable.update.assert_not_awaited()
def test_update_key_request_accepts_soft_budget():
@ -17712,28 +17715,6 @@ def test_validate_soft_budget_value_rejects_invalid_values(invalid_value):
assert "soft_budget must be a non-negative finite number" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_apply_soft_budget_update_noop_when_field_not_set():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_apply_soft_budget_update,
)
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None)
mock_prisma_client = MagicMock()
original_values = {"max_budget": 10.0}
result = await _apply_soft_budget_update(
data=UpdateKeyRequest(key="sk-test", max_budget=10.0),
non_default_values=original_values,
prisma_client=mock_prisma_client,
existing_key_row=existing_key,
changed_by="user-1",
)
assert result is original_values
mock_prisma_client.db.litellm_budgettable.create.assert_not_called()
@pytest.mark.asyncio
async def test_apply_soft_budget_update_adds_budget_id_for_new_budget_row():
from litellm.proxy.management_endpoints.key_management_endpoints import (
@ -17743,13 +17724,14 @@ async def test_apply_soft_budget_update_adds_budget_id_for_new_budget_row():
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None)
created_row = MagicMock()
created_row.budget_id = "budget-created-456"
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_budgettable.create = AsyncMock(return_value=created_row)
mock_db = MagicMock()
mock_db.litellm_budgettable.create = AsyncMock(return_value=created_row)
mock_db.litellm_budgettable.update = AsyncMock()
result = await _apply_soft_budget_update(
data=UpdateKeyRequest(key="sk-test", soft_budget=25.0),
non_default_values={"soft_budget": 25.0},
prisma_client=mock_prisma_client,
db=mock_db,
existing_key_row=existing_key,
changed_by="user-1",
)
@ -17764,24 +17746,95 @@ async def test_apply_soft_budget_update_keeps_existing_budget_id_out_of_token_up
)
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id="budget-123")
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_budgettable.update = AsyncMock()
mock_db = MagicMock()
mock_db.litellm_budgettable.update = AsyncMock()
mock_db.litellm_budgettable.create = AsyncMock()
result = await _apply_soft_budget_update(
data=UpdateKeyRequest(key="sk-test", soft_budget=40.0),
non_default_values={"soft_budget": 40.0, "max_budget": 100.0},
prisma_client=mock_prisma_client,
db=mock_db,
existing_key_row=existing_key,
changed_by="user-1",
)
assert dict(result) == {"max_budget": 100.0}
mock_prisma_client.db.litellm_budgettable.update.assert_awaited_once_with(
mock_db.litellm_budgettable.update.assert_awaited_once_with(
where={"budget_id": "budget-123"},
data={"soft_budget": 40.0, "updated_by": "user-1"},
)
@pytest.mark.asyncio
async def test_update_key_row_with_soft_budget_updates_budget_and_key_in_transaction():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_update_key_row_with_soft_budget,
)
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None)
created_row = MagicMock(budget_id="budget-new")
updated_row = MagicMock()
updated_row.model_dump.return_value = {"token": "hashed", "budget_id": "budget-new"}
tx = MagicMock()
tx.litellm_budgettable.create = AsyncMock(return_value=created_row)
tx.litellm_verificationtoken.update = AsyncMock(return_value=updated_row)
tx_context = MagicMock()
tx_context.__aenter__ = AsyncMock(return_value=tx)
tx_context.__aexit__ = AsyncMock(return_value=None)
prisma_client = MagicMock()
prisma_client.tx.return_value = tx_context
prisma_client.jsonify_object = lambda data: dict(data)
result = await _update_key_row_with_soft_budget(
prisma_client=prisma_client,
key="sk-test",
data=UpdateKeyRequest(key="sk-test", soft_budget=25.0),
non_default_values={"soft_budget": 25.0},
existing_key_row=existing_key,
changed_by="user-1",
)
assert set(result) == {"token", "data"}
assert result["data"] == {"token": "hashed", "budget_id": "budget-new"}
tx.litellm_verificationtoken.update.assert_awaited_once()
update_call = tx.litellm_verificationtoken.update.await_args
assert update_call.kwargs["where"] == {"token": result["token"]}
assert update_call.kwargs["data"]["budget_id"] == "budget-new"
assert "soft_budget" not in update_call.kwargs["data"]
@pytest.mark.asyncio
async def test_update_key_row_with_soft_budget_propagates_transaction_error():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_update_key_row_with_soft_budget,
)
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None)
created_row = MagicMock(budget_id="budget-new")
tx = MagicMock()
tx.litellm_budgettable.create = AsyncMock(return_value=created_row)
tx.litellm_verificationtoken.update = AsyncMock(side_effect=RuntimeError("update failed"))
tx_context = MagicMock()
tx_context.__aenter__ = AsyncMock(return_value=tx)
tx_context.__aexit__ = AsyncMock(return_value=None)
prisma_client = MagicMock()
prisma_client.tx.return_value = tx_context
prisma_client.jsonify_object = lambda data: dict(data)
with pytest.raises(RuntimeError, match="update failed"):
await _update_key_row_with_soft_budget(
prisma_client=prisma_client,
key="sk-test",
data=UpdateKeyRequest(key="sk-test", soft_budget=25.0),
non_default_values={"soft_budget": 25.0},
existing_key_row=existing_key,
changed_by="user-1",
)
tx_context.__aexit__.assert_awaited_once()
assert tx_context.__aexit__.await_args.args[0] is RuntimeError
def test_generate_key_request_blank_team_id_is_personal():
"""The UI Team-field clear submits team_id=""; it must count as no team (LIT-3925)."""
from litellm.proxy._types import RegenerateKeyRequest