mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-05 08:07:05 +00:00
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:
parent
2102d01571
commit
1783eb3402
2 changed files with 141 additions and 65 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue