fix(vertex): address lyria review feedback

This commit is contained in:
Emerson Gomes 2026-06-19 17:44:54 -05:00
parent 3ead9d1688
commit 514e9a1ee6
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
5 changed files with 75 additions and 19 deletions

View file

@ -45512,6 +45512,7 @@
"supports_tool_choice": true
},
"vertex_ai/lyria-002": {
"audio_seconds_per_prediction": 30,
"litellm_provider": "vertex_ai",
"max_audio_length_hours": 0.009111111111111111,
"max_audio_per_prompt": 4,

View file

@ -44,9 +44,7 @@ else:
PassThroughEndpointLogging = Any
LiteLLMBatch = Any
# Define EndpointType locally to avoid import issues
EndpointType = Any
_LYRIA_SECONDS_PER_AUDIO_PREDICTION = 30
class VertexPassthroughLoggingHandler:
@ -271,11 +269,11 @@ class VertexPassthroughLoggingHandler:
_json_response: Final[dict[str, object]] = httpx_response.json()
litellm_prediction_response: ModelResponse | EmbeddingResponse | ImageResponse = ModelResponse()
if VertexPassthroughLoggingHandler._is_lyria_predict_response(
if VertexPassthroughLoggingHandler._is_audio_predict_response(
model=model,
json_response=_json_response,
):
return VertexPassthroughLoggingHandler._handle_lyria_predict_response(
return VertexPassthroughLoggingHandler._handle_audio_predict_response(
json_response=_json_response,
logging_obj=logging_obj,
model=model,
@ -335,23 +333,19 @@ class VertexPassthroughLoggingHandler:
}
@staticmethod
def _handle_lyria_predict_response(
def _handle_audio_predict_response(
json_response: dict,
logging_obj: LiteLLMLoggingObj,
model: str,
kwargs: dict,
) -> PassThroughEndpointLoggingTypedDict:
prediction_count: Final = (
VertexPassthroughLoggingHandler._get_lyria_audio_prediction_count(
json_response=json_response
)
prediction_count: Final = VertexPassthroughLoggingHandler._get_audio_prediction_count(
json_response=json_response
)
model_info: Final = litellm.model_cost.get(f"vertex_ai/{model}", {})
response_cost: Final = (
model_info.get("output_cost_per_second", 0.0)
* _LYRIA_SECONDS_PER_AUDIO_PREDICTION
* prediction_count
)
VertexPassthroughLoggingHandler._get_audio_prediction_unit_cost(model=model)
or 0.0
) * prediction_count
logging_obj.model = model
logging_obj.model_call_details["model"] = model
@ -372,17 +366,31 @@ class VertexPassthroughLoggingHandler:
}
@staticmethod
def _is_lyria_predict_response(model: str, json_response: dict) -> bool:
def _is_audio_predict_response(model: str, json_response: dict) -> bool:
return (
model == "lyria-002"
and VertexPassthroughLoggingHandler._get_lyria_audio_prediction_count(
VertexPassthroughLoggingHandler._get_audio_prediction_count(
json_response=json_response
)
> 0
and VertexPassthroughLoggingHandler._get_audio_prediction_unit_cost(
model=model
)
is not None
)
@staticmethod
def _get_lyria_audio_prediction_count(json_response: dict) -> int:
def _get_audio_prediction_unit_cost(model: str) -> float | None:
model_info: Final = litellm.model_cost.get(f"vertex_ai/{model}", {})
output_cost_per_second: Final = model_info.get("output_cost_per_second")
audio_seconds_per_prediction: Final = model_info.get("audio_seconds_per_prediction")
if not isinstance(output_cost_per_second, (int, float)) or not isinstance(
audio_seconds_per_prediction, (int, float)
):
return None
return float(output_cost_per_second * audio_seconds_per_prediction)
@staticmethod
def _get_audio_prediction_count(json_response: dict) -> int:
predictions: Final = json_response.get("predictions")
if not isinstance(predictions, list):
return 0

View file

@ -45512,6 +45512,7 @@
"supports_tool_choice": true
},
"vertex_ai/lyria-002": {
"audio_seconds_per_prediction": 30,
"litellm_provider": "vertex_ai",
"max_audio_length_hours": 0.009111111111111111,
"max_audio_per_prompt": 4,

View file

@ -16,7 +16,10 @@ def test_lyria_predict_response_preserves_audio_response_and_logs_cost(
monkeypatch.setitem(
litellm.model_cost,
"vertex_ai/lyria-002",
{"output_cost_per_second": 0.002},
{
"audio_seconds_per_prediction": 30,
"output_cost_per_second": 0.002,
},
)
logging_obj = MagicMock()
logging_obj.model_call_details = {}
@ -66,3 +69,44 @@ def test_lyria_predict_response_preserves_audio_response_and_logs_cost(
assert result["kwargs"]["response_cost"] == pytest.approx(0.12)
assert logging_obj.model == "lyria-002"
assert logging_obj.model_call_details["response_cost"] == pytest.approx(0.12)
def test_audio_predict_response_uses_model_map_metadata(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setitem(
litellm.model_cost,
"vertex_ai/music-audio-preview",
{
"audio_seconds_per_prediction": 12,
"output_cost_per_second": 0.5,
},
)
logging_obj = MagicMock()
logging_obj.model_call_details = {}
response = httpx.Response(
status_code=200,
json={
"predictions": [
{
"audioContent": "clip",
"mimeType": "audio/wav",
}
]
},
)
result = VertexPassthroughLoggingHandler.vertex_passthrough_handler(
httpx_response=response,
logging_obj=logging_obj,
url_route="/v1/projects/test/locations/us-central1/publishers/google/models/music-audio-preview:predict",
result=response.text,
start_time=datetime.now(),
end_time=datetime.now(),
cache_hit=False,
request_body={"instances": [{"prompt": "ambient piano"}]},
)
assert result["kwargs"]["model"] == "music-audio-preview"
assert result["kwargs"]["response_cost"] == pytest.approx(6.0)
assert logging_obj.model_call_details["response_cost"] == pytest.approx(6.0)

View file

@ -949,6 +949,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"max_tokens": {"type": "number"},
"metadata": {"type": "object"},
"provider_specific_entry": {"type": "object"},
"audio_seconds_per_prediction": {"type": "number"},
"mode": {
"type": "string",
"enum": [
@ -2852,6 +2853,7 @@ def test_vertex_ai_lyria_models_in_cost_map():
assert lyria_2["litellm_provider"] == "vertex_ai"
assert clip["litellm_provider"] == "vertex_ai"
assert pro["litellm_provider"] == "vertex_ai"
assert lyria_2["audio_seconds_per_prediction"] == 30
assert lyria_2["output_cost_per_second"] == 0.002
assert lyria_2["supported_modalities"] == ["text"]
assert lyria_2["supported_output_modalities"] == ["audio"]