mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(vertex_ai): price passthrough calls on the URL's serving location
This commit is contained in:
parent
7b6f537855
commit
c549cddada
5 changed files with 119 additions and 17 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue