mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
test: fake the provider at the HTTP boundary in the batch model_group test
The test-quality gate flagged the first version for patching an SDK internal (litellm.batches.main.vertex_ai_batches_instance) and for writing litellm.callbacks directly. Drive an openai-compatible deployment through respx instead, so the retrieve call and the usage accounting that reads the completed batch's output file both run for real, and install the collector with monkeypatch so nothing leaks into the next test. Signed-off-by: mynkyu <mynkyu@bucketplace.net>
This commit is contained in:
parent
2286bf3eca
commit
e630f21d16
1 changed files with 89 additions and 57 deletions
|
|
@ -1,35 +1,81 @@
|
|||
"""
|
||||
model_group attribution on router batch retrieval.
|
||||
|
||||
Batch token usage is accounted on the *retrieve* call, not on create: the
|
||||
provider only knows the token counts once the job finishes, so
|
||||
`LiteLLMBatch.usage` arrives on `aretrieve_batch` and that is the record the
|
||||
spend log tokens land on.
|
||||
Batch token usage is accounted on the *retrieve* call, not on create: a provider
|
||||
only reports token counts once the job finishes, so the usage is read off the
|
||||
completed batch's output file during retrieve logging and that is the spend log
|
||||
row the tokens land on.
|
||||
|
||||
`aretrieve_batch` is addressed by batch_id, so the request carries no model,
|
||||
and the router fans the lookup out across its deployments. These tests lock
|
||||
that the winning deployment's model group is stamped on the emitted
|
||||
StandardLoggingPayload, so `/global/activity/model` - which groups the spend
|
||||
logs by `model_group` - can attribute those tokens instead of bucketing every
|
||||
batch under "".
|
||||
A batch is retrieved by id, so the request carries no model and the router fans
|
||||
the lookup out across its deployments. These tests lock that the answering
|
||||
deployment's model group is stamped on the emitted StandardLoggingPayload, so
|
||||
`/global/activity/model` - which groups the spend logs by `model_group` - can
|
||||
attribute those tokens instead of bucketing every batch under "".
|
||||
|
||||
The provider is faked at the HTTP boundary, so the whole retrieve + usage
|
||||
accounting path runs for real.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import MagicMock, patch
|
||||
import json
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
import litellm.batches.main as bm
|
||||
from litellm import Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import LiteLLMBatch, Usage
|
||||
|
||||
MODEL_GROUP = "vertex-gemini-2.5-flash-lite-dev"
|
||||
DEPLOYMENT_MODEL = "vertex_ai/gemini-2.5-flash-lite"
|
||||
MODEL_GROUP = "gemini-batch-group"
|
||||
DEPLOYMENT_MODEL = "openai/gpt-4o-mini"
|
||||
API_BASE = "http://localhost:4001/v1"
|
||||
BATCH_ID = "batch-1"
|
||||
ROWS = 2
|
||||
TOKENS_PER_ROW = 600
|
||||
|
||||
COMPLETED_BATCH = {
|
||||
"id": BATCH_ID,
|
||||
"object": "batch",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"errors": None,
|
||||
"input_file_id": "file-in-1",
|
||||
"completion_window": "24h",
|
||||
"status": "completed",
|
||||
"output_file_id": "file-out-1",
|
||||
"error_file_id": None,
|
||||
"created_at": 0,
|
||||
"completed_at": 1,
|
||||
"request_counts": {"total": ROWS, "completed": ROWS, "failed": 0},
|
||||
"metadata": None,
|
||||
}
|
||||
|
||||
OUTPUT_JSONL = "\n".join(
|
||||
json.dumps(
|
||||
{
|
||||
"id": f"req-{row}",
|
||||
"custom_id": f"row-{row}",
|
||||
"response": {
|
||||
"status_code": 200,
|
||||
"body": {
|
||||
"id": f"chatcmpl-{row}",
|
||||
"object": "chat.completion",
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}
|
||||
],
|
||||
"usage": {"prompt_tokens": 500, "completion_tokens": 100, "total_tokens": TOKENS_PER_ROW},
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
for row in range(ROWS)
|
||||
)
|
||||
|
||||
|
||||
class _PayloadCollector(CustomLogger):
|
||||
"""Captures the StandardLoggingPayload the spend log is built from."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.payloads = []
|
||||
|
|
@ -37,6 +83,14 @@ class _PayloadCollector(CustomLogger):
|
|||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
self.payloads.append(kwargs.get("standard_logging_object"))
|
||||
|
||||
async def retrieve_batch_payload(self) -> dict:
|
||||
for _ in range(100): # the success handler runs as a background task
|
||||
for payload in self.payloads:
|
||||
if payload and payload.get("call_type") == "aretrieve_batch":
|
||||
return payload
|
||||
await asyncio.sleep(0.05)
|
||||
raise AssertionError(f"no aretrieve_batch payload was emitted: {self.payloads}")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def router():
|
||||
|
|
@ -46,9 +100,8 @@ def router():
|
|||
"model_name": MODEL_GROUP,
|
||||
"litellm_params": {
|
||||
"model": DEPLOYMENT_MODEL,
|
||||
"vertex_project": "fake-project",
|
||||
"vertex_location": "us-central1",
|
||||
"vertex_credentials": "fake-creds",
|
||||
"api_base": API_BASE,
|
||||
"api_key": "sk-fake",
|
||||
},
|
||||
}
|
||||
]
|
||||
|
|
@ -56,63 +109,42 @@ def router():
|
|||
|
||||
|
||||
@pytest.fixture
|
||||
def collector():
|
||||
def collector(monkeypatch):
|
||||
logger = _PayloadCollector()
|
||||
previous = litellm.callbacks
|
||||
litellm.callbacks = [logger]
|
||||
try:
|
||||
yield logger
|
||||
finally:
|
||||
litellm.callbacks = previous
|
||||
monkeypatch.setattr(litellm, "callbacks", [logger])
|
||||
return logger
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def vertex_retrieve():
|
||||
"""Mock the vertex provider seam - the only real network boundary."""
|
||||
batch = LiteLLMBatch(
|
||||
id="batch-1",
|
||||
completion_window="24h",
|
||||
created_at=0,
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="file-1",
|
||||
object="batch",
|
||||
status="completed",
|
||||
usage=Usage(prompt_tokens=1000, completion_tokens=200, total_tokens=1200),
|
||||
)
|
||||
seam = MagicMock(name="vertex_ai_batches_instance")
|
||||
seam.retrieve_batch.return_value = batch
|
||||
with patch.object(bm, "vertex_ai_batches_instance", seam):
|
||||
yield seam
|
||||
|
||||
|
||||
async def _collected_payload(collector) -> dict:
|
||||
for _ in range(50): # the success handler runs as a background task
|
||||
payloads = [p for p in collector.payloads if p is not None]
|
||||
if payloads:
|
||||
return payloads[-1]
|
||||
await asyncio.sleep(0.05)
|
||||
raise AssertionError(f"no StandardLoggingPayload was emitted: {collector.payloads}")
|
||||
def provider():
|
||||
"""Fake the provider at the HTTP boundary: the completed batch plus the
|
||||
output file the usage accounting reads."""
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
respx_mock.get(f"{API_BASE}/batches/{BATCH_ID}").mock(return_value=httpx.Response(200, json=COMPLETED_BATCH))
|
||||
respx_mock.get(f"{API_BASE}/files/file-out-1/content").mock(return_value=httpx.Response(200, text=OUTPUT_JSONL))
|
||||
yield respx_mock
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aretrieve_batch_without_model_stamps_model_group(router, collector, vertex_retrieve):
|
||||
async def test_aretrieve_batch_without_model_stamps_model_group(router, collector, provider):
|
||||
"""
|
||||
The proxy retrieves a managed batch by id only - no `model` in the request.
|
||||
The router fans out over its deployments, so the model group is only known
|
||||
from the deployment that answered.
|
||||
"""
|
||||
response = await router.aretrieve_batch(batch_id="batch-1")
|
||||
response = await router.aretrieve_batch(batch_id=BATCH_ID)
|
||||
|
||||
assert response.usage.total_tokens == 1200
|
||||
payload = await _collected_payload(collector)
|
||||
assert response.id == BATCH_ID
|
||||
payload = await collector.retrieve_batch_payload()
|
||||
assert payload["total_tokens"] == ROWS * TOKENS_PER_ROW
|
||||
assert payload["model"] == DEPLOYMENT_MODEL
|
||||
assert payload["model_group"] == MODEL_GROUP
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aretrieve_batch_with_model_stamps_requested_model_group(router, collector, vertex_retrieve):
|
||||
async def test_aretrieve_batch_with_model_stamps_requested_model_group(router, collector, provider):
|
||||
"""An explicitly requested model group is what gets logged."""
|
||||
await router.aretrieve_batch(model=MODEL_GROUP, batch_id="batch-1")
|
||||
await router.aretrieve_batch(model=MODEL_GROUP, batch_id=BATCH_ID)
|
||||
|
||||
payload = await _collected_payload(collector)
|
||||
payload = await collector.retrieve_batch_payload()
|
||||
assert payload["model_group"] == MODEL_GROUP
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue