mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vertex_ai): apply regional endpoint uplift on image generation cost path (#44679)
* fix(vertex_ai): apply regional endpoint uplift on image generation cost path completion_cost had vertex_location but never passed it to the image generation cost router, so a regional_endpoint_uplift_multiplier on a Vertex image row would be ignored. No image row carries the multiplier yet, so nothing is misbilled today. Pass the location through to the Vertex image calculator for both the token-based price and the per-image fallback. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * test(vertex_ai): cover regional image cost through the proxy logging path Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * fix(images): hand image_edit's vertex_location to its cost resolver * test(integration): audit Vertex image regional uplift billing * test(integration): require every spend row after a proxy restart * test(integration): reject extra spend rows after a proxy restart --------- Co-authored-by: Nate Armstrong <narmstrong@Nates-MacBook-Pro.local> Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
c1d639afff
commit
922ea71fe5
8 changed files with 836 additions and 1 deletions
|
|
@ -1575,6 +1575,7 @@ def completion_cost(
|
|||
optional_params=optional_params,
|
||||
call_type=call_type,
|
||||
model_info=_deployment_model_info(litellm_logging_obj, custom_pricing, router_model_id),
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
elif call_type in _VIDEO_CALL_TYPES:
|
||||
### VIDEO GENERATION COST CALCULATION ###
|
||||
|
|
|
|||
|
|
@ -876,6 +876,7 @@ def image_edit(
|
|||
**image_edit_request_params,
|
||||
"litellm_call_id": litellm_call_id,
|
||||
"model_info": model_info,
|
||||
"vertex_location": litellm_params.vertex_location,
|
||||
},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1775,6 +1775,7 @@ def calculate_image_response_cost_from_usage(
|
|||
image_response: ImageResponse,
|
||||
custom_llm_provider: str,
|
||||
model_info: ModelInfo | None = None,
|
||||
vertex_location: str | None = None,
|
||||
) -> float | None:
|
||||
"""
|
||||
Calculate image generation cost from usage metadata when available.
|
||||
|
|
@ -1858,6 +1859,7 @@ def calculate_image_response_cost_from_usage(
|
|||
usage=normalized_usage,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_info=model_info,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
return prompt_cost + completion_cost
|
||||
|
||||
|
|
@ -1919,6 +1921,7 @@ class CostCalculatorUtils:
|
|||
optional_params: dict | None = None,
|
||||
call_type: str | None = None,
|
||||
model_info: ModelInfo | None = None,
|
||||
vertex_location: str | None = None,
|
||||
) -> float:
|
||||
"""
|
||||
Route the image generation cost calculator based on the custom_llm_provider
|
||||
|
|
@ -1960,6 +1963,7 @@ class CostCalculatorUtils:
|
|||
model=model,
|
||||
image_response=completion_response,
|
||||
model_info=pricing,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
elif custom_llm_provider == litellm.LlmProviders.BEDROCK.value:
|
||||
if isinstance(completion_response, ImageResponse):
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from typing import Final
|
|||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
calculate_image_response_cost_from_usage,
|
||||
calculate_image_response_web_search_cost,
|
||||
get_vertex_regional_endpoint_uplift,
|
||||
resolve_image_model_info,
|
||||
)
|
||||
from litellm.types.utils import ImageResponse, ModelInfo
|
||||
|
|
@ -16,6 +17,7 @@ def cost_calculator(
|
|||
model: str,
|
||||
image_response: ImageResponse,
|
||||
model_info: ModelInfo | None = None,
|
||||
vertex_location: str | None = None,
|
||||
) -> float:
|
||||
"""
|
||||
Vertex AI Image Generation Cost Calculator
|
||||
|
|
@ -37,10 +39,12 @@ def cost_calculator(
|
|||
image_response=image_response,
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_info=_model_info,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
if token_based_cost is not None:
|
||||
return token_based_cost + web_search_cost
|
||||
|
||||
output_cost_per_image: Final[float] = _model_info.get("output_cost_per_image") or 0.0
|
||||
num_images: Final[int] = len(image_response.data) if image_response.data else 0
|
||||
return output_cost_per_image * num_images + web_search_cost
|
||||
uplift: Final = get_vertex_regional_endpoint_uplift(_model_info, vertex_location)
|
||||
return output_cost_per_image * num_images * uplift + web_search_cost
|
||||
|
|
|
|||
673
tests/integration/spend/test_vertex_image_regional_uplift.py
Normal file
673
tests/integration/spend/test_vertex_image_regional_uplift.py
Normal file
|
|
@ -0,0 +1,673 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import psutil
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, Scenario, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.vertex import service_account_json
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_IMAGE_BACKEND: Final = "gemini-scripted-image"
|
||||
_CHAT_BACKEND: Final = "gemini-scripted-chat"
|
||||
_PROJECT: Final = "scripted-project"
|
||||
_REGION: Final = "us-central1"
|
||||
_PROMPT: Final = "a scripted lighthouse at dusk"
|
||||
_MODEL_PATH: Final = f"/v1/projects/{_PROJECT}/locations/{_REGION}/publishers/google/models"
|
||||
_IMAGE_TARGET: Final = f"{_MODEL_PATH}/{_IMAGE_BACKEND}:generateContent"
|
||||
_CHAT_TARGET: Final = f"{_MODEL_PATH}/{_CHAT_BACKEND}:generateContent"
|
||||
_PNG_B64: Final = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8DwHwAFBQIAX8jx0gAAAABJRU5ErkJggg=="
|
||||
_PNG: Final = base64.b64decode(_PNG_B64)
|
||||
_GLOBAL_TOKEN_COST: Final = 12 * 2e-6 + 1290 * 4e-5
|
||||
_REGIONAL_TOKEN_COST: Final = (12 * 2e-6 + 1290 * 4e-5) * 1.1
|
||||
_GLOBAL_IMAGE_COST: Final = 0.05
|
||||
_REGIONAL_IMAGE_COST: Final = 0.05 * 1.1
|
||||
_REGIONAL_CHAT_COST: Final = (12 * 2e-6 + 7 * 4e-5) * 1.1
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
_CONFIG_MODEL: Final = "vertex-image-chaos"
|
||||
_BURST: Final = 30
|
||||
|
||||
|
||||
def _image_body(*, images: int, usage: bool) -> bytes:
|
||||
parts: Final = [{"inlineData": {"mimeType": "image/png", "data": _PNG_B64}} for _ in range(images)]
|
||||
usage_metadata: Final = {
|
||||
"promptTokenCount": 12,
|
||||
"candidatesTokenCount": 1290,
|
||||
"totalTokenCount": 1302,
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 12}],
|
||||
}
|
||||
return json.dumps(
|
||||
{
|
||||
"candidates": [{"content": {"role": "model", "parts": parts}, "finishReason": "STOP"}],
|
||||
**({"usageMetadata": usage_metadata} if usage else {}),
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
_CHAT_BODY: Final = json.dumps(
|
||||
{
|
||||
"candidates": [{"content": {"role": "model", "parts": [{"text": "scripted answer"}]}, "finishReason": "STOP"}],
|
||||
"usageMetadata": {"promptTokenCount": 12, "candidatesTokenCount": 7, "totalTokenCount": 19},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _scripted(*, images: int = 1, usage: bool = True, delay: float = 0) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert (request.method, request.target) == ("POST", _IMAGE_TARGET), (request.method, request.target)
|
||||
assert _PROMPT in request.body.decode(), request.body
|
||||
if delay:
|
||||
time.sleep(delay)
|
||||
return Reply(body=_image_body(images=images, usage=usage))
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _scripted_chat(request: Request) -> Reply:
|
||||
assert (request.method, request.target) == ("POST", _CHAT_TARGET), (request.method, request.target)
|
||||
assert request.headers["authorization"] == "Bearer scripted-token", request.headers
|
||||
return Reply(body=_CHAT_BODY)
|
||||
|
||||
|
||||
def _deployment(
|
||||
gateway: Gateway,
|
||||
scenario: Scenario,
|
||||
wire: Wire,
|
||||
*,
|
||||
location: Mapping[str, JsonValue],
|
||||
multiplier: JsonValue = 1.1,
|
||||
backend: str = _IMAGE_BACKEND,
|
||||
) -> str:
|
||||
return scenario.model(
|
||||
model=f"vertex_ai/{backend}",
|
||||
api_base=f"{wire.url}{_MODEL_PATH}/{backend}:generateContent",
|
||||
api_key=None,
|
||||
vertex_project=_PROJECT,
|
||||
vertex_credentials=service_account_json(_PROJECT, gateway.upstream_url.rstrip("/")),
|
||||
input_cost_per_token=2e-6,
|
||||
output_cost_per_token=4e-5,
|
||||
output_cost_per_image=0.05,
|
||||
**({} if multiplier is None else {"regional_endpoint_uplift_multiplier": multiplier}),
|
||||
**location,
|
||||
)
|
||||
|
||||
|
||||
def _chat_deployment(gateway: Gateway, scenario: Scenario, wire: Wire) -> str:
|
||||
return scenario.model(
|
||||
model=f"vertex_ai/{_CHAT_BACKEND}",
|
||||
api_base=f"{wire.url}{_MODEL_PATH}/{_CHAT_BACKEND}",
|
||||
api_key=None,
|
||||
vertex_project=_PROJECT,
|
||||
vertex_location=_REGION,
|
||||
vertex_credentials=service_account_json(_PROJECT, gateway.upstream_url.rstrip("/")),
|
||||
input_cost_per_token=2e-6,
|
||||
output_cost_per_token=4e-5,
|
||||
regional_endpoint_uplift_multiplier=1.1,
|
||||
)
|
||||
|
||||
|
||||
def _spend_rows(call_ids: tuple[str, ...], *, seconds: float = 70) -> dict[str, float]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, spend FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s)', (list(call_ids),)
|
||||
),
|
||||
lambda found: len(found) == len(call_ids),
|
||||
seconds=seconds,
|
||||
)
|
||||
assert len({str(row["request_id"]) for row in rows}) == len(rows), rows
|
||||
return {str(row["request_id"]): float(str(row["spend"])) for row in rows}
|
||||
|
||||
|
||||
def _spend(call_id: str) -> float:
|
||||
return _spend_rows((call_id,))[call_id]
|
||||
|
||||
|
||||
def _spend_by_litellm_call_id(call_id: str) -> float:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
"SELECT spend FROM \"LiteLLM_SpendLogs\" WHERE metadata->>'litellm_call_id' = %s", (call_id,)
|
||||
),
|
||||
lambda found: len(found) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return float(str(rows[0]["spend"]))
|
||||
|
||||
|
||||
def _assert_image_response(response: httpx.Response, expected: float, *, images: int = 1) -> None:
|
||||
assert response.status_code == 200, response.text
|
||||
data: Final = _JSON_OBJECT.validate_json(response.content)["data"]
|
||||
assert isinstance(data, list) and len(data) == images, response.text
|
||||
assert all(object_value(item)["b64_json"] == _PNG_B64 for item in data), response.text
|
||||
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected), response.headers
|
||||
|
||||
|
||||
def _assert_images(response: httpx.Response, call_id: str, expected: float, *, images: int = 1) -> None:
|
||||
_assert_image_response(response, expected, images=images)
|
||||
assert _spend(call_id) == pytest.approx(expected)
|
||||
|
||||
|
||||
def _assert_one_upstream_image_call(wire: Wire, *, prompt: str = _PROMPT) -> None:
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 1, [request.target for request in received]
|
||||
body: Final = _JSON_OBJECT.validate_json(received[0].body)
|
||||
contents: Final = body["contents"]
|
||||
assert isinstance(contents, list) and prompt in json.dumps(contents), body
|
||||
assert "IMAGE" in json.dumps(body["generationConfig"]), body
|
||||
|
||||
|
||||
def _generate(
|
||||
gateway: Gateway, model: str, *, key: str | None = None, extra: Mapping[str, JsonValue] | None = None
|
||||
) -> tuple[str, httpx.Response]:
|
||||
call_id: Final = uuid.uuid4().hex
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/images/generations",
|
||||
{"model": model, "prompt": _PROMPT, "n": 1, **(extra or {})},
|
||||
key=key,
|
||||
headers={"x-litellm-call-id": call_id},
|
||||
)
|
||||
return call_id, response
|
||||
|
||||
|
||||
def _edit(gateway: Gateway, model: str) -> tuple[str, httpx.Response]:
|
||||
call_id: Final = uuid.uuid4().hex
|
||||
response: Final = gateway.client.post(
|
||||
"/v1/images/edits",
|
||||
data={"model": model, "prompt": _PROMPT},
|
||||
files={"image": ("red.png", _PNG, "image/png")},
|
||||
headers={"Authorization": f"Bearer {gateway.key}", "x-litellm-call-id": call_id},
|
||||
)
|
||||
return call_id, response
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("location", "multiplier", "usage", "images", "expected"),
|
||||
[
|
||||
pytest.param({"vertex_location": "global"}, 1.1, True, 1, _GLOBAL_TOKEN_COST, id="global-usage"),
|
||||
pytest.param({"vertex_location": _REGION}, 1.1, True, 1, _REGIONAL_TOKEN_COST, id="regional-usage"),
|
||||
pytest.param({"vertex_location": _REGION}, 1.1, False, 1, _REGIONAL_IMAGE_COST, id="regional-per-image"),
|
||||
pytest.param({"vertex_location": _REGION}, 1.1, False, 2, 2 * _REGIONAL_IMAGE_COST, id="regional-two-images"),
|
||||
pytest.param({"vertex_location": "GLOBAL"}, 1.1, True, 1, _GLOBAL_TOKEN_COST, id="upper-case-global"),
|
||||
pytest.param({"vertex_location": "us"}, 1.1, True, 1, _REGIONAL_TOKEN_COST, id="multi-region"),
|
||||
pytest.param({"vertex_location": _REGION}, None, True, 1, _GLOBAL_TOKEN_COST, id="regional-without-multiplier"),
|
||||
pytest.param({"vertex_location": _REGION}, "1.1", True, 1, _REGIONAL_TOKEN_COST, id="string-multiplier"),
|
||||
pytest.param({}, 1.1, True, 1, _REGIONAL_TOKEN_COST, id="missing-location-defaults-to-us-central1"),
|
||||
pytest.param(
|
||||
{"vertex_location": ""}, 1.1, True, 1, _REGIONAL_TOKEN_COST, id="empty-location-defaults-to-us-central1"
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_image_generation_bills_the_deployment_location(
|
||||
gateway: Gateway,
|
||||
location: Mapping[str, JsonValue],
|
||||
multiplier: JsonValue,
|
||||
usage: bool,
|
||||
images: int,
|
||||
expected: float,
|
||||
) -> None:
|
||||
with wire_server(_scripted(images=images, usage=usage)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location=location, multiplier=multiplier)
|
||||
call_id, response = _generate(gateway, model, extra={"n": images})
|
||||
_assert_images(response, call_id, expected, images=images)
|
||||
_assert_one_upstream_image_call(wire)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("location", "usage", "expected"),
|
||||
[
|
||||
pytest.param("global", False, _GLOBAL_IMAGE_COST, id="global-edit"),
|
||||
pytest.param(_REGION, False, _REGIONAL_IMAGE_COST, id="regional-edit"),
|
||||
pytest.param(_REGION, True, _REGIONAL_IMAGE_COST, id="regional-edit-usage-is-per-image"),
|
||||
],
|
||||
)
|
||||
def test_image_edit_bills_the_deployment_location(
|
||||
gateway: Gateway, location: str, usage: bool, expected: float
|
||||
) -> None:
|
||||
with wire_server(_scripted(usage=usage)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": location})
|
||||
call_id, response = _edit(gateway, model)
|
||||
_assert_images(response, call_id, expected)
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 1, [request.target for request in received]
|
||||
assert _PNG_B64 in received[0].body.decode(), received[0].body[:200]
|
||||
|
||||
|
||||
def test_openai_sdk_image_generation_bills_the_regional_uplift(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": _REGION})
|
||||
call_id: Final = uuid.uuid4().hex
|
||||
client: Final = openai.OpenAI(base_url=f"{gateway.client.base_url}/v1", api_key=gateway.key, timeout=60)
|
||||
raw: Final = client.images.with_raw_response.generate(
|
||||
model=model, prompt=_PROMPT, extra_headers={"x-litellm-call-id": call_id}
|
||||
)
|
||||
assert raw.status_code == 200, raw.text
|
||||
(image,) = raw.parse().data or ()
|
||||
assert image.b64_json == _PNG_B64
|
||||
assert float(raw.headers["x-litellm-response-cost"]) == pytest.approx(_REGIONAL_TOKEN_COST), raw.headers
|
||||
assert _spend(call_id) == pytest.approx(_REGIONAL_TOKEN_COST)
|
||||
_assert_one_upstream_image_call(wire)
|
||||
|
||||
|
||||
async def test_async_openai_sdk_image_generation_keeps_global_flat(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": "global"})
|
||||
call_id: Final = uuid.uuid4().hex
|
||||
async with openai.AsyncOpenAI(
|
||||
base_url=f"{gateway.client.base_url}/v1", api_key=gateway.key, timeout=60
|
||||
) as client:
|
||||
raw: Final = await client.images.with_raw_response.generate(
|
||||
model=model, prompt=_PROMPT, extra_headers={"x-litellm-call-id": call_id}
|
||||
)
|
||||
assert raw.status_code == 200, raw.text
|
||||
(image,) = raw.parse().data or ()
|
||||
assert image.b64_json == _PNG_B64
|
||||
assert float(raw.headers["x-litellm-response-cost"]) == pytest.approx(_GLOBAL_TOKEN_COST), raw.headers
|
||||
assert await asyncio.to_thread(_spend, call_id) == pytest.approx(_GLOBAL_TOKEN_COST)
|
||||
_assert_one_upstream_image_call(wire)
|
||||
|
||||
|
||||
def test_openai_sdk_image_edit_bills_the_regional_uplift(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted(usage=False)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": _REGION})
|
||||
call_id: Final = uuid.uuid4().hex
|
||||
client: Final = openai.OpenAI(base_url=f"{gateway.client.base_url}/v1", api_key=gateway.key, timeout=60)
|
||||
raw: Final = client.images.with_raw_response.edit(
|
||||
model=model,
|
||||
image=("red.png", _PNG, "image/png"),
|
||||
prompt=_PROMPT,
|
||||
extra_headers={"x-litellm-call-id": call_id},
|
||||
)
|
||||
assert raw.status_code == 200, raw.text
|
||||
(image,) = raw.parse().data or ()
|
||||
assert image.b64_json == _PNG_B64
|
||||
assert float(raw.headers["x-litellm-response-cost"]) == pytest.approx(_REGIONAL_IMAGE_COST), raw.headers
|
||||
assert _spend(call_id) == pytest.approx(_REGIONAL_IMAGE_COST)
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 1 and _PNG_B64 in received[0].body.decode(), received
|
||||
|
||||
|
||||
def test_model_info_reads_back_the_multiplier_the_price_follows(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": _REGION})
|
||||
entries: Final = gateway.get("/model/info")["data"]
|
||||
assert isinstance(entries, list)
|
||||
(entry,) = (object_value(item) for item in entries if object_value(item)["model_name"] == model)
|
||||
info: Final = object_value(entry["model_info"])
|
||||
assert info["regional_endpoint_uplift_multiplier"] == 1.1, info
|
||||
assert object_value(entry["litellm_params"])["vertex_location"] == _REGION, entry
|
||||
|
||||
|
||||
def test_invalid_multiplier_is_refused_at_registration(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted()) as wire:
|
||||
refused: Final = gateway.request(
|
||||
"POST",
|
||||
"/model/new",
|
||||
{
|
||||
"model_name": f"integration-{uuid.uuid4().hex}",
|
||||
"litellm_params": {
|
||||
"model": f"vertex_ai/{_IMAGE_BACKEND}",
|
||||
"api_base": f"{wire.url}{_IMAGE_TARGET}",
|
||||
"vertex_project": _PROJECT,
|
||||
"vertex_location": _REGION,
|
||||
"output_cost_per_image": 0.05,
|
||||
"regional_endpoint_uplift_multiplier": "abc",
|
||||
},
|
||||
"model_info": {},
|
||||
},
|
||||
)
|
||||
assert refused.status_code == 422, refused.text
|
||||
assert "regional_endpoint_uplift_multiplier" in refused.text, refused.text
|
||||
assert wire.drain() == ()
|
||||
|
||||
|
||||
def test_identical_regional_generations_each_bill_the_uplift_once(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": _REGION})
|
||||
first_id, first = _generate(gateway, model)
|
||||
second_id, second = _generate(gateway, model)
|
||||
assert first.status_code == 200, first.text
|
||||
assert second.status_code == 200, second.text
|
||||
assert first.content == second.content, (first.text, second.text)
|
||||
rows: Final = _spend_rows((first_id, second_id))
|
||||
assert rows == {first_id: pytest.approx(_REGIONAL_TOKEN_COST), second_id: pytest.approx(_REGIONAL_TOKEN_COST)}
|
||||
assert len(wire.drain()) == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("deployment_location", "extra", "expected"),
|
||||
[
|
||||
pytest.param(_REGION, {"vertex_location": "global"}, _GLOBAL_TOKEN_COST, id="body-global-wins-over-regional"),
|
||||
pytest.param("global", {"vertex_location": 123}, _REGIONAL_TOKEN_COST, id="body-int"),
|
||||
pytest.param("global", {"vertex_location": ["us-central1"]}, _REGIONAL_TOKEN_COST, id="body-list"),
|
||||
pytest.param("global", {"vertex_location": ""}, _REGIONAL_TOKEN_COST, id="body-empty"),
|
||||
pytest.param("global", {"vertex_location": "a" * 5000}, _REGIONAL_TOKEN_COST, id="body-5kb"),
|
||||
pytest.param(
|
||||
"global",
|
||||
{"vertex_location": "global", "vertex_ai_location": _REGION},
|
||||
_GLOBAL_TOKEN_COST,
|
||||
id="body-both-aliases",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_request_body_location_overrides_the_deployment(
|
||||
gateway: Gateway, deployment_location: str, extra: Mapping[str, JsonValue], expected: float
|
||||
) -> None:
|
||||
with wire_server(_scripted()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": deployment_location})
|
||||
call_id, response = _generate(gateway, model, extra=extra)
|
||||
_assert_images(response, call_id, expected)
|
||||
_assert_one_upstream_image_call(wire)
|
||||
control_id, control = _generate(gateway, model)
|
||||
_assert_images(
|
||||
control, control_id, _GLOBAL_TOKEN_COST if deployment_location == "global" else _REGIONAL_TOKEN_COST
|
||||
)
|
||||
|
||||
|
||||
def test_stream_flag_on_image_generation_prices_the_regional_uplift(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted()) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": _REGION})
|
||||
_, response = _generate(gateway, model, extra={"stream": True})
|
||||
_assert_image_response(response, _REGIONAL_TOKEN_COST)
|
||||
_assert_one_upstream_image_call(wire)
|
||||
|
||||
|
||||
def test_openai_style_image_deployment_is_untouched(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert (request.method, request.target) == ("POST", "/images/generations"), request
|
||||
return Reply(body=json.dumps({"created": 1700000000, "data": [{"b64_json": _PNG_B64}]}).encode())
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/dall-e-3", api_base=wire.url, api_key="synthetic-image-key", output_cost_per_image=0.25
|
||||
)
|
||||
call_id, response = _generate(gateway, model)
|
||||
_assert_images(response, call_id, 0.25)
|
||||
assert len(wire.drain()) == 1
|
||||
|
||||
|
||||
def test_chat_completions_on_a_regional_chat_deployment(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted_chat) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _chat_deployment(gateway, scenario, wire)
|
||||
call_id: Final = uuid.uuid4().hex
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": _PROMPT}], "cache": {"no-cache": True}},
|
||||
headers={"x-litellm-call-id": call_id},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = _JSON_OBJECT.validate_json(response.content)
|
||||
assert "scripted answer" in response.text and response.headers["x-litellm-call-id"] == call_id, response.text
|
||||
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(_REGIONAL_CHAT_COST), (
|
||||
response.headers
|
||||
)
|
||||
assert _spend(str(body["id"])) == pytest.approx(_REGIONAL_CHAT_COST)
|
||||
assert len(wire.drain()) == 1
|
||||
|
||||
|
||||
def test_messages_on_a_regional_chat_deployment(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted_chat) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _chat_deployment(gateway, scenario, wire)
|
||||
call_id: Final = uuid.uuid4().hex
|
||||
client: Final = anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, timeout=60)
|
||||
raw: Final = client.messages.with_raw_response.create(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
extra_headers={"x-litellm-call-id": call_id},
|
||||
extra_body={"cache": {"no-cache": True}},
|
||||
)
|
||||
assert raw.status_code == 200, raw.text
|
||||
assert "scripted answer" in raw.text and raw.headers["x-litellm-call-id"] == call_id, raw.text
|
||||
assert float(raw.headers["x-litellm-response-cost"]) == pytest.approx(_REGIONAL_CHAT_COST), raw.headers
|
||||
assert _spend(str(_JSON_OBJECT.validate_json(raw.content)["id"])) == pytest.approx(_REGIONAL_CHAT_COST)
|
||||
assert len(wire.drain()) == 1
|
||||
|
||||
|
||||
def test_responses_on_a_regional_chat_deployment(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted_chat) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _chat_deployment(gateway, scenario, wire)
|
||||
call_id: Final = uuid.uuid4().hex
|
||||
client: Final = openai.OpenAI(base_url=f"{gateway.client.base_url}/v1", api_key=gateway.key, timeout=60)
|
||||
raw: Final = client.responses.with_raw_response.create(
|
||||
model=model,
|
||||
input=_PROMPT,
|
||||
extra_headers={"x-litellm-call-id": call_id},
|
||||
extra_body={"cache": {"no-cache": True}},
|
||||
)
|
||||
assert raw.status_code == 200, raw.text
|
||||
assert "scripted answer" in raw.text and raw.headers["x-litellm-call-id"] == call_id, raw.text
|
||||
assert float(raw.headers["x-litellm-response-cost"]) == pytest.approx(_REGIONAL_CHAT_COST), raw.headers
|
||||
assert _spend_by_litellm_call_id(call_id) == pytest.approx(_REGIONAL_CHAT_COST)
|
||||
assert len(wire.drain()) == 1
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Call:
|
||||
call_id: str
|
||||
route: str
|
||||
model: str
|
||||
expected: float
|
||||
|
||||
|
||||
def _burst_calls(global_model: str, regional_model: str) -> tuple[_Call, ...]:
|
||||
kinds: Final = (
|
||||
("/v1/images/generations", global_model, _GLOBAL_TOKEN_COST),
|
||||
("/v1/images/generations", regional_model, _REGIONAL_TOKEN_COST),
|
||||
("/v1/images/edits", global_model, _GLOBAL_IMAGE_COST),
|
||||
("/v1/images/edits", regional_model, _REGIONAL_IMAGE_COST),
|
||||
)
|
||||
return tuple(
|
||||
_Call(call_id=uuid.uuid4().hex, route=route, model=model, expected=expected)
|
||||
for route, model, expected in (kinds[index % len(kinds)] for index in range(_BURST))
|
||||
)
|
||||
|
||||
|
||||
def _send(client: httpx.Client, key: str, call: _Call) -> httpx.Response:
|
||||
headers: Final = {"Authorization": f"Bearer {key}", "x-litellm-call-id": call.call_id}
|
||||
if call.route == "/v1/images/edits":
|
||||
return client.post(
|
||||
call.route,
|
||||
data={"model": call.model, "prompt": _PROMPT},
|
||||
files={"image": ("red.png", _PNG, "image/png")},
|
||||
headers=headers,
|
||||
timeout=120,
|
||||
)
|
||||
return client.post(call.route, json={"model": call.model, "prompt": _PROMPT, "n": 1}, headers=headers, timeout=120)
|
||||
|
||||
|
||||
def test_mixed_burst_against_a_slow_upstream_bills_every_call_once(gateway: Gateway) -> None:
|
||||
with wire_server(_scripted(delay=0.3)) as wire, gateway.scenario() as scenario:
|
||||
global_model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": "global"})
|
||||
regional_model: Final = _deployment(gateway, scenario, wire, location={"vertex_location": _REGION})
|
||||
calls: Final = _burst_calls(global_model, regional_model)
|
||||
with ThreadPoolExecutor(max_workers=_BURST) as pool:
|
||||
responses: Final = tuple(pool.map(lambda call: _send(gateway.client, gateway.key, call), calls))
|
||||
assert [response.status_code for response in responses] == [200] * _BURST, [r.text[:120] for r in responses]
|
||||
rows: Final = _spend_rows(tuple(call.call_id for call in calls), seconds=120)
|
||||
assert rows == {call.call_id: pytest.approx(call.expected) for call in calls}
|
||||
assert len(wire.drain()) == _BURST
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
call_id: str
|
||||
status: int
|
||||
text: str
|
||||
|
||||
|
||||
async def _async_burst(base_url: str, key: str, call_ids: tuple[str, ...]) -> tuple[_Served, ...]:
|
||||
async def send(client: httpx.AsyncClient, call_id: str) -> _Served:
|
||||
response: Final = await client.post(
|
||||
"/v1/images/generations",
|
||||
json={"model": _CONFIG_MODEL, "prompt": _PROMPT, "n": 1},
|
||||
headers={"Authorization": f"Bearer {key}", "x-litellm-call-id": call_id},
|
||||
)
|
||||
return _Served(call_id=call_id, status=response.status_code, text=response.text)
|
||||
|
||||
async with httpx.AsyncClient(base_url=base_url, timeout=120, trust_env=False) as client:
|
||||
results: Final = await asyncio.gather(*(send(client, call_id) for call_id in call_ids), return_exceptions=True)
|
||||
for result in results:
|
||||
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
|
||||
return tuple(result for result in results if isinstance(result, _Served))
|
||||
|
||||
|
||||
def _chaos_config(wire: Wire, tmp_path: Path, upstream_url: str) -> Path:
|
||||
config: Final = {
|
||||
**_JSON_OBJECT.validate_python(
|
||||
yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text())
|
||||
),
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": _CONFIG_MODEL,
|
||||
"litellm_params": {
|
||||
"model": f"vertex_ai/{_IMAGE_BACKEND}",
|
||||
"api_base": f"{wire.url}{_IMAGE_TARGET}",
|
||||
"vertex_project": _PROJECT,
|
||||
"vertex_location": _REGION,
|
||||
"vertex_credentials": service_account_json(_PROJECT, upstream_url.rstrip("/")),
|
||||
"input_cost_per_token": 2e-6,
|
||||
"output_cost_per_token": 4e-5,
|
||||
"output_cost_per_image": 0.05,
|
||||
"regional_endpoint_uplift_multiplier": 1.1,
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
path: Final = tmp_path / "vertex-image-chaos.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def _artifact(tmp_path: Path, name: str, payload: Mapping[str, JsonValue]) -> None:
|
||||
(Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(tmp_path))) / name).write_text(json.dumps(payload))
|
||||
|
||||
|
||||
def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]:
|
||||
text: Final = log.read_text()
|
||||
return tuple(int(pid) for pid in _STARTED_WORKER.findall(text)), text.count("Application startup complete.")
|
||||
|
||||
|
||||
def _open_upstream_connections(pid: int, upstream: str) -> int:
|
||||
port: Final = urlsplit(upstream).port
|
||||
return sum(
|
||||
1
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
|
||||
)
|
||||
|
||||
|
||||
def _gated(release: threading.Event, held: SimpleQueue[str]) -> Callable[[Request], Reply]:
|
||||
scripted: Final = _scripted()
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
held.put(request.target)
|
||||
assert release.wait(timeout=120), "The burst was never released"
|
||||
return scripted(request)
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
@pytest.mark.timeout(420)
|
||||
async def test_worker_sigkill_mid_burst_leaves_the_sibling_billing_the_uplift(gateway: Gateway, tmp_path: Path) -> None:
|
||||
release: Final = threading.Event()
|
||||
held: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
call_ids: Final = tuple(uuid.uuid4().hex for _ in range(20))
|
||||
with wire_server(_gated(release, held)) as wire:
|
||||
config: Final = _chaos_config(wire, tmp_path, gateway.upstream_url)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
|
||||
base_url: Final = str(owned.gateway.client.base_url)
|
||||
workers, _ = eventually(
|
||||
lambda: _worker_startups(owned.log), lambda found: len(found[0]) == 2 and found[1] == 2, seconds=240
|
||||
)
|
||||
burst: Final = asyncio.create_task(_async_burst(base_url, gateway.key, call_ids))
|
||||
await asyncio.to_thread(eventually, held.qsize, lambda size: size == len(call_ids), 120)
|
||||
held_by: Final = {pid: _open_upstream_connections(pid, wire.url) for pid in workers}
|
||||
assert sum(held_by.values()) == len(call_ids), held_by
|
||||
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
|
||||
psutil.Process(victim_pid).send_signal(signal.SIGKILL)
|
||||
release.set()
|
||||
served: Final = await burst
|
||||
succeeded: Final = tuple(item.call_id for item in served if item.status == 200)
|
||||
_artifact(
|
||||
tmp_path,
|
||||
"worker-sigkill.json",
|
||||
{
|
||||
"held_by_killed_worker": held_by[victim_pid],
|
||||
"held_by_surviving_worker": held_by[survivor_pid],
|
||||
"served": len(served),
|
||||
"succeeded": len(succeeded),
|
||||
"transport_errors": len(call_ids) - len(served),
|
||||
},
|
||||
)
|
||||
assert len(succeeded) == held_by[survivor_pid], (
|
||||
held_by,
|
||||
[(item.status, item.text[:80]) for item in served],
|
||||
)
|
||||
follow_up_id, follow_up = _generate(owned.gateway, _CONFIG_MODEL)
|
||||
_assert_images(follow_up, follow_up_id, _REGIONAL_TOKEN_COST)
|
||||
rows: Final = await asyncio.to_thread(_spend_rows, succeeded, seconds=120)
|
||||
assert rows == {call_id: pytest.approx(_REGIONAL_TOKEN_COST) for call_id in succeeded}
|
||||
await asyncio.to_thread(
|
||||
eventually, lambda: _worker_startups(owned.log), lambda found: len(found[0]) == 3 and found[1] == 3, 240
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(600)
|
||||
async def test_proxy_restart_mid_burst_bills_each_landed_call_once(gateway: Gateway, tmp_path: Path) -> None:
|
||||
release: Final = threading.Event()
|
||||
held: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
call_ids: Final = tuple(uuid.uuid4().hex for _ in range(20))
|
||||
with wire_server(_gated(release, held)) as wire:
|
||||
config: Final = _chaos_config(wire, tmp_path, gateway.upstream_url)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
|
||||
base_url: Final = str(owned.gateway.client.base_url)
|
||||
eventually(
|
||||
lambda: _worker_startups(owned.log), lambda found: len(found[0]) == 2 and found[1] == 2, seconds=240
|
||||
)
|
||||
burst: Final = asyncio.create_task(_async_burst(base_url, gateway.key, call_ids))
|
||||
await asyncio.to_thread(eventually, held.qsize, lambda size: size == len(call_ids), 120)
|
||||
owned.process.terminate()
|
||||
release.set()
|
||||
served: Final = await burst
|
||||
await asyncio.to_thread(owned.process.wait, 240)
|
||||
succeeded: Final = tuple(item.call_id for item in served if item.status == 200)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as rebooted:
|
||||
follow_up_id, follow_up = _generate(rebooted.gateway, _CONFIG_MODEL)
|
||||
_assert_images(follow_up, follow_up_id, _REGIONAL_TOKEN_COST)
|
||||
landed: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, spend FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s)', (list(call_ids),)
|
||||
),
|
||||
lambda rows: set(succeeded) <= {str(row["request_id"]) for row in rows},
|
||||
seconds=120,
|
||||
)
|
||||
landed_ids: Final = tuple(str(row["request_id"]) for row in landed)
|
||||
_artifact(
|
||||
tmp_path,
|
||||
"restart-loss.json",
|
||||
{"served": len(served), "succeeded": len(succeeded), "landed": len(landed_ids)},
|
||||
)
|
||||
assert sorted(landed_ids) == sorted(succeeded), (succeeded, landed_ids)
|
||||
assert all(float(str(row["spend"])) == pytest.approx(_REGIONAL_TOKEN_COST) for row in landed), landed
|
||||
|
|
@ -1,10 +1,14 @@
|
|||
import json
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
def test_image_generation_keeps_an_internal_prefixed_kwarg_out_of_the_provider_request(
|
||||
|
|
@ -27,3 +31,48 @@ def test_image_generation_keeps_an_internal_prefixed_kwarg_out_of_the_provider_r
|
|||
sent: Final = json.loads(respx_mock.calls[0].request.content)
|
||||
assert "_litellm_undeclared_sentinel" not in sent, sent
|
||||
assert sent["prompt"] == "a red circle"
|
||||
|
||||
|
||||
def test_image_edit_prices_a_vertex_deployment_at_its_configured_location(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
api_base: Final = "http://localhost:12347/generateContent"
|
||||
respx_mock.post(api_base).mock(
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={"candidates": [{"content": {"parts": [{"inlineData": {"mimeType": "image/png", "data": "aW1n"}}]}}]},
|
||||
)
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"vertex_ai/gemini-fake-regional-edit-model",
|
||||
{
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.04,
|
||||
"regional_endpoint_uplift_multiplier": 1.1,
|
||||
},
|
||||
)
|
||||
|
||||
def cost_at(location: str) -> float:
|
||||
logging_obj: Final = Logging(
|
||||
model="gemini-fake-regional-edit-model",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type=CallTypes.image_edit.value,
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id=f"vertex-edit-{location}",
|
||||
function_id="f",
|
||||
)
|
||||
response: Final = litellm.image_edit(
|
||||
model="vertex_ai/gemini-fake-regional-edit-model",
|
||||
image=b"\x89PNG\r\n\x1a\nfakepng",
|
||||
prompt="make the circle blue",
|
||||
api_base=api_base,
|
||||
vertex_location=location,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
return logging_obj._response_cost_calculator(result=response)
|
||||
|
||||
assert cost_at("global") == pytest.approx(0.04)
|
||||
assert cost_at("us-central1") == pytest.approx(0.044)
|
||||
|
|
|
|||
|
|
@ -7279,6 +7279,64 @@ def test_response_cost_calculator_prices_proxy_vertex_calls_on_the_configured_lo
|
|||
assert cost_at("us-east5") == pytest.approx(info["regional_endpoint_uplift_multiplier"] * expected_global)
|
||||
|
||||
|
||||
def test_response_cost_calculator_prices_proxy_vertex_image_calls_on_the_configured_location(monkeypatch):
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
|
||||
from litellm.types.utils import ImageObject, ImageUsage, ImageUsageInputTokensDetails
|
||||
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"model_cost",
|
||||
{
|
||||
**get_model_cost_map(url=""),
|
||||
"vertex_ai/fake-regional-image-model": {
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"mode": "image_generation",
|
||||
"input_cost_per_token": 5e-07,
|
||||
"output_cost_per_image_token": 6e-05,
|
||||
"regional_endpoint_uplift_multiplier": 1.1,
|
||||
},
|
||||
},
|
||||
)
|
||||
monkeypatch.setenv("VERTEXAI_LOCATION", "us-east5")
|
||||
monkeypatch.setattr(litellm, "vertex_location", None)
|
||||
|
||||
def cost_at(location):
|
||||
logging_obj = LitellmLogging(
|
||||
model="fake-regional-image-model",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="image_generation",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id=f"vertex-image-loc-{location}",
|
||||
function_id="f",
|
||||
)
|
||||
logging_obj.update_environment_variables(
|
||||
model="fake-regional-image-model",
|
||||
user="",
|
||||
optional_params={"vertex_location": location},
|
||||
litellm_params={"api_base": ""},
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
response = ImageResponse(
|
||||
data=[ImageObject(b64_json="img")],
|
||||
usage=ImageUsage(
|
||||
input_tokens=100,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(image_tokens=0, text_tokens=100),
|
||||
output_tokens=1120,
|
||||
total_tokens=1220,
|
||||
),
|
||||
)
|
||||
return logging_obj._response_cost_calculator(result=response)
|
||||
|
||||
expected_global = 100 * 5e-07 + 1120 * 6e-05
|
||||
|
||||
assert cost_at("global") == pytest.approx(expected_global)
|
||||
assert cost_at("us-east5") == pytest.approx(1.1 * expected_global)
|
||||
|
||||
|
||||
def test_set_cost_breakdown_stores_vertex_location():
|
||||
"""vertex_location is recorded in the pricing basis, None for non-vertex requests."""
|
||||
from datetime import datetime
|
||||
|
|
|
|||
|
|
@ -1522,6 +1522,51 @@ def test_vertex_regional_deployment_costs_uplift_over_global(monkeypatch):
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"usage",
|
||||
[
|
||||
ImageUsage(
|
||||
input_tokens=100,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(image_tokens=0, text_tokens=100),
|
||||
output_tokens=1120,
|
||||
total_tokens=1220,
|
||||
),
|
||||
None,
|
||||
],
|
||||
ids=["token-priced", "per-image-fallback"],
|
||||
)
|
||||
def test_vertex_regional_image_generation_costs_uplift_over_global(monkeypatch, usage):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"model_cost",
|
||||
{
|
||||
**litellm.get_model_cost_map(url=""),
|
||||
"vertex_ai/fake-regional-image-model": {
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"mode": "image_generation",
|
||||
"input_cost_per_token": 5e-07,
|
||||
"output_cost_per_token": 3e-06,
|
||||
"output_cost_per_image_token": 6e-05,
|
||||
"output_cost_per_image": 0.0672,
|
||||
"regional_endpoint_uplift_multiplier": 1.1,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
def image_cost(vertex_location: str) -> float:
|
||||
return completion_cost(
|
||||
completion_response=ImageResponse(data=[ImageObject(b64_json="img")], usage=usage),
|
||||
model="vertex_ai/fake-regional-image-model",
|
||||
call_type="image_generation",
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
|
||||
global_cost: Final = image_cost("global")
|
||||
assert global_cost > 0
|
||||
assert image_cost("us-central1") == pytest.approx(global_cost * 1.10, rel=1e-9)
|
||||
|
||||
|
||||
def test_vertex_uplift_composes_with_above_128k_pricing(monkeypatch):
|
||||
"""The regional-endpoint uplift multiplies whatever rate the request priced at,
|
||||
including the above-128k dynamic rates, so a synthetic model carrying both keys
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue