Merge remote-tracking branch 'origin/main' into litellm_lit6852_cli_session_spend_key

This commit is contained in:
mateo 2026-09-24 21:07:31 +00:00
commit c63ee1b131
7 changed files with 217 additions and 93 deletions

View file

@ -253,30 +253,15 @@ class VertexAIBatchTransformation:
return uris[0]
@classmethod
def _get_output_file_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str:
def _get_output_file_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str | None:
"""
Gets the output file id from the Vertex AI Batch response
Gets the output file id from the Vertex AI Batch response, None until Vertex reports outputInfo
"""
output_info: Final = response.get("outputInfo") or OutputInfo()
output_file_id: str = output_info.get("gcsOutputDirectory", "")
if output_file_id:
output_file_id = output_file_id.rstrip("/") + "/predictions.jsonl"
if output_file_id and output_file_id != "/predictions.jsonl":
return output_file_id
output_config: Final = response.get("outputConfig")
if output_config is None:
return output_file_id
gcs_destination: Final = output_config.get("gcsDestination")
if gcs_destination is None:
return output_file_id
output_uri_prefix: Final = gcs_destination.get("outputUriPrefix", "")
if output_uri_prefix.endswith("/predictions.jsonl"):
return output_uri_prefix
return output_uri_prefix.rstrip("/") + "/predictions.jsonl"
gcs_output_directory: Final = (output_info.get("gcsOutputDirectory") or "").rstrip("/")
if not gcs_output_directory:
return None
return f"{gcs_output_directory}/predictions.jsonl"
@classmethod
def _get_batch_job_status_from_vertex_ai_batch_response(

View file

@ -124,6 +124,12 @@ class _SpendIncrement(TypedDict):
increment: ReadOnly[float]
class _MemberSpendRow(TypedDict):
user_id: ReadOnly[str]
team_id: ReadOnly[str]
cost: ReadOnly[float]
class _SpendBatch(Protocol):
litellm_usertable: BatchTable
litellm_verificationtoken: BatchTable
@ -351,17 +357,22 @@ _TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS
# One statement adds every member's cost to their membership row. A missing row is created only
# while the user is still on the team's roster, so a spend flush landing after a removal never
# recreates the member.
# recreates the member. The rows travel as one JSON document, not as a numeric array: Prisma
# types a raw array parameter from the first batch a connection sees, so after an all-$0 batch
# (integers) every later fractional batch on that connection failed with "improper binary format".
_TEAM_MEMBER_SPEND_SQL: Final = """
INSERT INTO "LiteLLM_TeamMembership" (user_id, team_id, spend, total_spend)
SELECT p.user_id, p.team_id, p.cost, p.cost
FROM unnest($1::text[], $2::text[], $3::float8[]) AS p(user_id, team_id, cost)
SELECT member.user_id, member.team_id, member.cost, member.cost
FROM jsonb_to_recordset($1::jsonb) AS member(user_id text, team_id text, cost float8)
WHERE EXISTS (
SELECT 1 FROM "LiteLLM_TeamTable" t
WHERE t.team_id = p.team_id
AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))
WHERE t.team_id = member.team_id
AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', member.user_id))
)
OR EXISTS (
SELECT 1 FROM "LiteLLM_TeamMembership" m
WHERE m.user_id = member.user_id AND m.team_id = member.team_id
)
OR EXISTS (SELECT 1 FROM "LiteLLM_TeamMembership" m WHERE m.user_id = p.user_id AND m.team_id = p.team_id)
ON CONFLICT (user_id, team_id) DO UPDATE
SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend,
total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend
@ -371,15 +382,12 @@ SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend,
async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None:
# key is "team_id::<value>::user_id::<value>"; locks are taken in sorted team_id order like the team endpoints
rows: Final = sorted((key.split("::")[1], key.split("::")[3], cost) for key, cost in spend_by_member_key.items())
team_ids: Final = tuple(team_id for team_id, _user_id, _cost in rows)
for team_id in dict.fromkeys(team_ids):
for team_id in dict.fromkeys(team_id for team_id, _user_id, _cost in rows):
_ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id)
_ = await transaction.execute_raw(
_TEAM_MEMBER_SPEND_SQL,
tuple(user_id for _team_id, user_id, _cost in rows),
team_ids,
tuple(cost for _team_id, _user_id, cost in rows),
members: Final = tuple(
_MemberSpendRow(user_id=user_id, team_id=team_id, cost=cost) for team_id, user_id, cost in rows
)
_ = await transaction.execute_raw(_TEAM_MEMBER_SPEND_SQL, json.dumps(members))
def get_llm_router():

View file

@ -116,7 +116,7 @@ def test_vertex_batch_create_survives_explicit_null_output_info(gateway: Gateway
"batch",
"validating",
_encoded(INPUT_FILE_ID, model, "file-"),
_encoded(f"{OUTPUT_PREFIX}/predictions.jsonl", model, "file-"),
None,
None,
"24h",
), response.text

View file

@ -0,0 +1,106 @@
"""Team member spend keeps landing after a flush in which every cost was a whole number.
The proxy runs on a one-connection pool so every spend flush reuses the same database
connection. A batch of $0 requests (a free model here) is the whole-number batch, and the
fractional batches that follow it must still land on that connection.
The $0 batch has to be flushed on its own before the paid request is sent. The spend log
row cannot prove that, since a separate monitor writes spend logs whenever they queue up,
but the daily user spend row is written by the flush cycle right after the member spend
statement, so its arrival means the whole-number batch has already been sent.
"""
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy
from pydantic import JsonValue
SINGLE_CONNECTION_CONFIG: Final = """
model_list: []
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
database_url: os.environ/DATABASE_URL
store_model_in_db: true
proxy_batch_write_at: 1
proxy_batch_polling_interval: 1
database_connection_pool_limit: 1
router_settings:
disable_cooldowns: true
"""
def _member_row(team_id: str, user_id: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT spend, total_spend FROM "LiteLLM_TeamMembership" WHERE team_id=%s AND user_id=%s',
(team_id, user_id),
)
def _member_spend_is(rows: list[dict[str, JsonValue]], amount: float) -> bool:
return len(rows) == 1 and all(
float(str(rows[0][column])) == pytest.approx(amount) for column in ("spend", "total_spend")
)
def _daily_user_spend_rows(user_id: str) -> list[dict[str, JsonValue]]:
return read_rows('SELECT spend FROM "LiteLLM_DailyUserSpend" WHERE user_id=%s', (user_id,))
def _logged_spend(request_id: str) -> float:
rows: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (request_id,)),
lambda found: len(found) == 1,
seconds=30,
)
return float(str(rows[0]["spend"]))
def _chat(gateway: Gateway, key: str, model: str) -> str:
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "member spend control"}]},
key=key,
)
assert response.status_code == 200, response.text
return string_value(response.json()["id"])
def test_fractional_member_spend_lands_after_a_whole_number_flush_on_the_same_connection(
gateway: Gateway, tmp_path: Path
) -> None:
config: Final = tmp_path / "single_connection_proxy.yaml"
config.write_text(SINGLE_CONNECTION_CONFIG)
with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario:
free: Final = scenario.model(input_cost_per_token=0, output_cost_per_token=0, num_retries=0)
paid: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002, num_retries=0)
team: Final = scenario.team(models=[free, paid])
first: Final = scenario.user()
second: Final = scenario.user()
candidate.post(
"/team/member_add",
{"team_id": team, "member": [{"role": "user", "user_id": first}, {"role": "user", "user_id": second}]},
)
first_key: Final = scenario.key(team_id=team, user_id=first)
second_key: Final = scenario.key(team_id=team, user_id=second)
_chat(candidate, first_key, free)
flushed: Final = eventually(lambda: _daily_user_spend_rows(first), lambda rows: len(rows) == 1, seconds=30)
assert float(str(flushed[0]["spend"])) == 0
assert _member_spend_is(_member_row(team, first), 0)
paid_spend: Final = _logged_spend(_chat(candidate, second_key, paid))
assert paid_spend > 0
eventually(lambda: _member_row(team, second), lambda rows: _member_spend_is(rows, paid_spend), seconds=30)
repeat_spend: Final = _logged_spend(_chat(candidate, first_key, paid))
eventually(lambda: _member_row(team, first), lambda rows: _member_spend_is(rows, repeat_spend), seconds=30)
eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id=%s', (team,)),
lambda rows: float(str(rows[0]["spend"])) == pytest.approx(paid_spend + repeat_spend),
seconds=30,
)

View file

@ -13,6 +13,8 @@ There are no real I/O seams here; ``uuid.uuid4`` is the only nondeterministic
dependency and is patched where the displayName is asserted.
"""
from collections.abc import Mapping
from typing import Final
from unittest.mock import patch
import pytest
@ -35,8 +37,7 @@ INPUT_FILE = (
ENDPOINT_ID = "7768560373388541952"
ENDPOINT_INPUT_FILE = (
f"gs://litellm-testing-bucket/litellm-vertex-files/endpoints/{ENDPOINT_ID}/"
"e9412502-2c91-42a6-8e61-f5c294cc0fc8"
f"gs://litellm-testing-bucket/litellm-vertex-files/endpoints/{ENDPOINT_ID}/e9412502-2c91-42a6-8e61-f5c294cc0fc8"
)
@ -248,9 +249,78 @@ def test_get_input_file_id_empty_uris():
# =========================================================================== #
# _get_output_file_id_from_vertex_ai_batch_response
# _get_output_file_id_from_vertex_ai_batch_response: None until Vertex reports outputInfo
# =========================================================================== #
SHARED_OUTPUT_PREFIX: Final = "gs://bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-flash"
SUCCEEDED_OUTPUT_DIRECTORY: Final = f"{SHARED_OUTPUT_PREFIX}/prediction-model-2026-09-24T19:41:00.000000Z"
def _vertex_job(state: str) -> dict[str, object]:
return {
"name": "projects/510528649030/locations/us-central1/batchPredictionJobs/3814889423749775360",
"state": state,
"createTime": "2026-09-24T19:37:25.775603Z",
"inputConfig": {
"instancesFormat": "jsonl",
"gcsSource": {"uris": [f"{SHARED_OUTPUT_PREFIX}/0586ba52-4f8b-4988-aa8d-3573550a4b0f"]},
},
"outputConfig": {
"predictionsFormat": "jsonl",
"gcsDestination": {"outputUriPrefix": SHARED_OUTPUT_PREFIX},
},
}
@pytest.mark.parametrize(
"vertex_state,output_info_field,expected_status,expected_output_file_id",
[
("JOB_STATE_PENDING", {}, "validating", None),
("JOB_STATE_RUNNING", {"outputInfo": {}}, "in_progress", None),
("JOB_STATE_CANCELLED", {"outputInfo": None}, "cancelled", None),
(
"JOB_STATE_SUCCEEDED",
{"outputInfo": {"gcsOutputDirectory": SUCCEEDED_OUTPUT_DIRECTORY}},
"completed",
f"{SUCCEEDED_OUTPUT_DIRECTORY}/predictions.jsonl",
),
],
ids=["create_or_pending", "running", "cancelled", "succeeded"],
)
def test_transform_vertex_response_output_file_id_is_none_until_output_info(
vertex_state: str,
output_info_field: Mapping[str, object],
expected_status: str,
expected_output_file_id: str | None,
) -> None:
batch: Final = T.transform_vertex_ai_batch_response_to_openai_batch_response(
{**_vertex_job(vertex_state), **output_info_field}
)
assert batch.status == expected_status
assert batch.output_file_id == expected_output_file_id
@pytest.mark.parametrize(
"response",
[
{},
{"outputConfig": {}},
{"outputInfo": None},
{"outputInfo": {"gcsOutputDirectory": ""}},
{"outputInfo": {"gcsOutputDirectory": None}},
],
ids=[
"no_fields",
"output_config_without_destination",
"null_output_info",
"empty_output_directory",
"null_output_directory",
],
)
def test_get_output_file_id_is_none_without_output_directory(response: Mapping[str, object]) -> None:
assert T._get_output_file_id_from_vertex_ai_batch_response(response) is None
def test_get_output_file_id_from_output_info():
# outputInfo branch: rstrip trailing slash, append predictions.jsonl
@ -267,49 +337,7 @@ def test_get_output_file_id_output_info_no_trailing_slash():
)
def test_get_output_file_id_empty_output_info_falls_through_to_output_config():
# gcsOutputDirectory missing -> "" -> the "/predictions.jsonl" guard skips
# the outputInfo branch, falls through to outputConfig
resp = {
"outputInfo": {},
"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg"}},
}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_info_explicit_none_falls_through_to_output_config():
resp = {
"outputInfo": None,
"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg"}},
}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_info_explicit_none_and_no_output_config():
assert T._get_output_file_id_from_vertex_ai_batch_response({"outputInfo": None}) == ""
def test_get_output_file_id_no_output_info_and_no_output_config():
assert T._get_output_file_id_from_vertex_ai_batch_response({}) == ""
def test_get_output_file_id_output_config_missing_gcs_destination():
# outputConfig present but no gcsDestination -> returns the running "" value
assert T._get_output_file_id_from_vertex_ai_batch_response({"outputConfig": {}}) == ""
def test_get_output_file_id_output_config_already_has_suffix():
# outputUriPrefix already ends in /predictions.jsonl -> returned as-is (no double append)
resp = {"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg/predictions.jsonl"}}}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_config_strips_trailing_slash():
resp = {"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg/"}}}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_info_takes_precedence_over_output_config():
def test_get_output_file_id_output_info_ignores_output_uri_prefix():
resp = {
"outputInfo": {"gcsOutputDirectory": "gs://from-info"},
"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://from-config"}},

View file

@ -26,12 +26,12 @@ def test_output_file_id_uses_predictions_jsonl_with_output_info():
)
def test_output_file_id_falls_back_to_output_uri_prefix_with_predictions_jsonl():
def test_output_file_id_is_none_until_output_info():
response = {
"outputInfo": {},
"outputConfig": {
"gcsDestination": {
"outputUriPrefix": "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-456"
"outputUriPrefix": "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro"
}
},
}
@ -42,10 +42,7 @@ def test_output_file_id_falls_back_to_output_uri_prefix_with_predictions_jsonl()
)
)
assert (
output_file_id
== "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-456/predictions.jsonl"
)
assert output_file_id is None
def test_vertex_ai_cancel_batch():

View file

@ -994,11 +994,11 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster
assert lock_statement is _TEAM_ADVISORY_LOCK_SQL
assert locked_team_id == team_id
assert "pg_advisory_xact_lock(hashtext($1))" in lock_statement
statement, user_ids, team_ids, costs = spend_call.args
statement, members = spend_call.args
assert statement is _TEAM_MEMBER_SPEND_SQL
assert (list(user_ids), list(team_ids), list(costs)) == ([user_id], [team_id], [response_cost])
assert json.loads(members) == [{"user_id": user_id, "team_id": team_id, "cost": response_cost}]
assert 'INSERT INTO "LiteLLM_TeamMembership"' in statement
assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))" in statement
assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', member.user_id))" in statement
assert "ON CONFLICT (user_id, team_id) DO UPDATE" in statement
assert 'spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend' in statement
assert 'total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend' in statement
@ -1007,7 +1007,7 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster
@pytest.mark.asyncio
async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_user():
"""
The member spend statement touches rows in the order of its input arrays, so the batch
The member spend statement touches rows in the order of its input rows, so the batch
is handed over sorted by (team_id, user_id), with each cost kept next to its member, and
each distinct team is locked once, in `sorted(team_ids)` order, the order /team/delete
locks in, so a concurrent flush and delete cannot deadlock. `eng` and `eng2` pin that:
@ -1034,13 +1034,13 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u
)
*lock_calls, spend_call = mock_transaction.execute_raw.await_args_list
_statement, user_ids, team_ids, costs = spend_call.args
_statement, members = spend_call.args
assert [lock_call.args for lock_call in lock_calls] == [
(_TEAM_ADVISORY_LOCK_SQL, "eng"),
(_TEAM_ADVISORY_LOCK_SQL, "eng-b"),
(_TEAM_ADVISORY_LOCK_SQL, "eng2"),
]
assert list(zip(team_ids, user_ids, costs)) == [
assert [(row["team_id"], row["user_id"], row["cost"]) for row in json.loads(members)] == [
("eng", "user_x", 0.3),
("eng", "user_y", 0.2),
("eng-b", "user_x", 0.4),