fix(vertex_ai): price passthrough calls on the URL's serving location

This commit is contained in:
mateo-berri 2026-08-19 16:44:52 -07:00
parent 7b6f537855
commit c549cddada
5 changed files with 119 additions and 17 deletions

View file

@ -999,9 +999,6 @@ def generic_cost_per_token(
prompt_cost *= uplift
completion_cost *= uplift
## VERTEX REGIONAL-ENDPOINT UPLIFT
# Applied as a flat multiplier across all token costs for the request
# when the Vertex AI endpoint serving it is non-global.
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
if vertex_uplift != 1.0:
prompt_cost *= vertex_uplift
@ -1109,8 +1106,6 @@ def get_token_type_cost_breakdown(
cache_read_cost *= uplift
cache_creation_cost *= uplift
# Same flat uplift for Vertex AI non-global endpoints, keeping per-type
# costs reconciled with the totals for regional Vertex deployments.
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
if vertex_uplift != 1.0:
reasoning_cost *= vertex_uplift

View file

@ -164,8 +164,6 @@ def cost_per_character(
usage=usage,
)
# Applied once here; the cost_per_token fallbacks above are called without
# vertex_location so the uplift can never compound.
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
return prompt_cost * vertex_uplift, completion_cost * vertex_uplift

View file

@ -10,6 +10,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import VERTEX_BATCH_PREDICTION_JOBS_ROUTE
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.vertex_ai.common_utils import get_vertex_location_from_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator as VertexModelResponseIterator,
)
@ -60,6 +61,7 @@ class VertexPassthroughLoggingHandler:
request_body: dict | None = None,
**kwargs,
) -> PassThroughEndpointLoggingTypedDict:
vertex_location: Final = get_vertex_location_from_url(url_route)
if "predictLongRunning" in url_route:
model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
@ -82,6 +84,7 @@ class VertexPassthroughLoggingHandler:
model=model,
custom_llm_provider="vertex_ai",
call_type="create_video",
vertex_location=vertex_location,
)
# Set response_cost in _hidden_params to prevent recalculation
@ -123,6 +126,7 @@ class VertexPassthroughLoggingHandler:
end_time=end_time,
logging_obj=logging_obj,
custom_llm_provider=VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route),
vertex_location=vertex_location,
)
return {
@ -190,6 +194,7 @@ class VertexPassthroughLoggingHandler:
end_time=end_time,
logging_obj=logging_obj,
custom_llm_provider="vertex_ai",
vertex_location=vertex_location,
)
return {
@ -206,6 +211,7 @@ class VertexPassthroughLoggingHandler:
model="vertex_ai/search_api",
custom_llm_provider="vertex_ai",
call_type="vector_store_search",
vertex_location=vertex_location,
)
standard_pass_through_response_object: Final[StandardPassThroughResponseObject] = {
@ -302,6 +308,7 @@ class VertexPassthroughLoggingHandler:
completion_response=litellm_prediction_response,
model=model,
custom_llm_provider="vertex_ai",
vertex_location=get_vertex_location_from_url(url_route),
)
kwargs["response_cost"] = response_cost
@ -381,6 +388,7 @@ class VertexPassthroughLoggingHandler:
completion_response=litellm_embedding_response,
model=model,
custom_llm_provider=custom_llm_provider,
vertex_location=get_vertex_location_from_url(url_route),
)
kwargs["response_cost"] = response_cost
@ -438,6 +446,7 @@ class VertexPassthroughLoggingHandler:
end_time=end_time,
logging_obj=litellm_logging_obj,
custom_llm_provider=VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route),
vertex_location=get_vertex_location_from_url(url_route),
)
return {
@ -591,6 +600,7 @@ class VertexPassthroughLoggingHandler:
end_time: datetime,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str,
vertex_location: str | None,
) -> dict:
"""
Create the standard logging object for Vertex passthrough generateContent (streaming and non-streaming)
@ -601,6 +611,7 @@ class VertexPassthroughLoggingHandler:
completion_response=litellm_model_response,
model=model,
custom_llm_provider="vertex_ai",
vertex_location=vertex_location,
)
kwargs["response_cost"] = response_cost

View file

@ -1109,6 +1109,98 @@ class TestVertexAIPassThroughHandler:
assert result["kwargs"].get("model") == "gemini-embedding-2-preview"
mock_completion_cost.assert_called_once()
@pytest.mark.parametrize("streaming", [False, True])
def test_vertex_passthrough_handler_prices_regional_endpoint_with_uplift(self, monkeypatch, streaming):
"""
Passthrough cost is computed inside the handler and stored as response_cost before the
logging cost resolver runs, so the handler itself must read the serving location out of
the passthrough URL; otherwise regional Vertex passthrough traffic bills at the global
rate (#34393).
"""
import datetime
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import (
VertexPassthroughLoggingHandler,
)
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(
litellm,
"model_cost",
{
**litellm.get_model_cost_map(url=""),
"vertex_ai/gemini-fake-regional": {
"litellm_provider": "vertex_ai",
"mode": "chat",
"input_cost_per_token": 1e-06,
"output_cost_per_token": 2e-06,
"regional_endpoint_uplift_multiplier": 1.1,
},
},
)
response_body: Final = {
"candidates": [
{
"content": {"parts": [{"text": "hello"}], "role": "model"},
"finishReason": "STOP",
}
],
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 20,
"totalTokenCount": 30,
},
}
def cost_for(location: str) -> float:
url_route: Final = (
f"https://{location}-aiplatform.googleapis.com/v1/projects/p/locations/{location}"
"/publishers/google/models/gemini-fake-regional:"
f"{'streamGenerateContent' if streaming else 'generateContent'}"
)
mock_logging_obj: Final = Mock()
mock_logging_obj.litellm_call_id = "call-id"
mock_logging_obj.model_call_details = {}
mock_logging_obj.optional_params = {}
start_time: Final = datetime.datetime.now()
end_time: Final = datetime.datetime.now()
if streaming:
result = VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks(
litellm_logging_obj=mock_logging_obj,
passthrough_success_handler_obj=Mock(),
url_route=url_route,
request_body={},
endpoint_type="vertex_ai",
start_time=start_time,
all_chunks=[json.dumps(response_body)],
model=None,
end_time=end_time,
)
else:
mock_httpx_response: Final = Mock()
mock_httpx_response.json.return_value = response_body
mock_httpx_response.headers = {}
mock_httpx_response.status_code = 200
result = VertexPassthroughLoggingHandler.vertex_passthrough_handler(
httpx_response=mock_httpx_response,
logging_obj=mock_logging_obj,
url_route=url_route,
result="test-result",
start_time=start_time,
end_time=end_time,
cache_hit=False,
)
return result["kwargs"]["response_cost"]
global_cost: Final = cost_for("global")
regional_cost: Final = cost_for("us-east5")
assert global_cost == pytest.approx(10 * 1e-06 + 20 * 2e-06, rel=1e-9)
assert regional_cost == pytest.approx(global_cost * 1.10, rel=1e-9), (
"regional Vertex passthrough traffic must bill at 1.1x the global rate"
)
class TestVertexAIDiscoveryPassThroughHandler:
"""

View file

@ -1781,16 +1781,22 @@ def test_vertex_uplift_composes_with_above_128k_pricing(monkeypatch):
including the above-128k dynamic rates, so a synthetic model carrying both keys
prices regional above-128k usage at 1.1x the above-128k rate."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
litellm.model_cost["vertex_ai/fake-regional-128k-model"] = {
"litellm_provider": "vertex_ai",
"mode": "chat",
"input_cost_per_token": 1e-06,
"output_cost_per_token": 2e-06,
"input_cost_per_token_above_128k_tokens": 2e-06,
"output_cost_per_token_above_128k_tokens": 4e-06,
"regional_endpoint_uplift_multiplier": 1.1,
}
monkeypatch.setattr(
litellm,
"model_cost",
{
**litellm.get_model_cost_map(url=""),
"vertex_ai/fake-regional-128k-model": {
"litellm_provider": "vertex_ai",
"mode": "chat",
"input_cost_per_token": 1e-06,
"output_cost_per_token": 2e-06,
"input_cost_per_token_above_128k_tokens": 2e-06,
"output_cost_per_token_above_128k_tokens": 4e-06,
"regional_endpoint_uplift_multiplier": 1.1,
},
},
)
usage = Usage(prompt_tokens=200_000, completion_tokens=10, total_tokens=200_010)
global_prompt, global_completion = cost_per_token(