mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(vertex_ai): build a well-formed endpoint-resolution url for path-mounted custom api_base
This commit is contained in:
parent
988ae7ca80
commit
54fd69beb2
2 changed files with 55 additions and 11 deletions
|
|
@ -1,6 +1,7 @@
|
|||
import json
|
||||
from collections.abc import Coroutine, Sequence
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
|
@ -18,6 +19,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAIError, get_vertex_base_url
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import _graft_default_vertex_path
|
||||
from litellm.types.llms.openai import CreateBatchRequest
|
||||
from litellm.types.llms.vertex_ai import (
|
||||
VERTEX_CREDENTIALS_TYPES,
|
||||
|
|
@ -176,6 +178,24 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
)
|
||||
return vertex_batch_response
|
||||
|
||||
@staticmethod
|
||||
def _build_endpoint_resolution_url(api_base: str | None, model: str, vertex_location: str) -> str:
|
||||
"""
|
||||
Builds the GET url for resolving an endpoint resource (`projects/../endpoints/<id>`).
|
||||
|
||||
A custom `api_base` replaces the Google host: its `/v1`/`/v1beta1` path swallows the
|
||||
version segment (matching `_check_custom_proxy`'s grafting), any other path is kept as a
|
||||
mount prefix in front of the full default path. The `:operation` suffix convention from
|
||||
`_check_custom_proxy` does not apply to a plain resource GET.
|
||||
"""
|
||||
default_endpoint_url: Final = f"{get_vertex_base_url(vertex_location)}/v1/{model}"
|
||||
if not api_base:
|
||||
return default_endpoint_url
|
||||
api_base_path: Final = urlparse(api_base).path.rstrip("/")
|
||||
if api_base_path in ("/v1", "/v1beta1"):
|
||||
return _graft_default_vertex_path(api_base=api_base, default_url=default_endpoint_url)
|
||||
return api_base.rstrip("/") + urlparse(default_endpoint_url).path
|
||||
|
||||
def _resolve_fine_tuned_endpoint_model(
|
||||
self,
|
||||
vertex_batch_request: VertexAIBatchPredictionJob,
|
||||
|
|
@ -193,18 +213,10 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
if "/endpoints/" not in model:
|
||||
return vertex_batch_request
|
||||
|
||||
default_endpoint_url: Final = f"{get_vertex_base_url(vertex_location)}/v1/{model}"
|
||||
_, endpoint_url = self._check_custom_proxy(
|
||||
endpoint_url: Final = self._build_endpoint_resolution_url(
|
||||
api_base=api_base,
|
||||
custom_llm_provider="vertex_ai",
|
||||
gemini_api_key=None,
|
||||
endpoint=(default_endpoint_url.split(":")[-1] if len(default_endpoint_url.split(":")) > 1 else ""),
|
||||
stream=None,
|
||||
auth_header=None,
|
||||
url=default_endpoint_url,
|
||||
model=None,
|
||||
model=model,
|
||||
vertex_location=vertex_location,
|
||||
vertex_api_version="v1",
|
||||
)
|
||||
response: Final = sync_handler.get(url=endpoint_url, headers=headers)
|
||||
if response.status_code != 200:
|
||||
|
|
|
|||
|
|
@ -35,7 +35,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm.llms.vertex_ai.batches.handler import ( # noqa: E402
|
||||
VertexAIBatchPrediction,
|
||||
)
|
||||
|
|
@ -253,6 +252,39 @@ def test_create_batch_sync_resolves_fine_tuned_endpoint_to_tuned_model():
|
|||
assert sent["model"] == TUNED_MODEL_RESOURCE
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base, expected",
|
||||
[
|
||||
(
|
||||
None,
|
||||
f"https://{LOCATION}-aiplatform.googleapis.com/v1/projects/{PROJECT}"
|
||||
f"/locations/{LOCATION}/endpoints/{ENDPOINT_ID}",
|
||||
),
|
||||
(
|
||||
"https://proxy.internal",
|
||||
f"https://proxy.internal/v1/projects/{PROJECT}/locations/{LOCATION}/endpoints/{ENDPOINT_ID}",
|
||||
),
|
||||
(
|
||||
"https://proxy.internal/v1",
|
||||
f"https://proxy.internal/v1/projects/{PROJECT}/locations/{LOCATION}/endpoints/{ENDPOINT_ID}",
|
||||
),
|
||||
(
|
||||
"https://proxy.internal/vertex",
|
||||
f"https://proxy.internal/vertex/v1/projects/{PROJECT}/locations/{LOCATION}/endpoints/{ENDPOINT_ID}",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_build_endpoint_resolution_url(api_base, expected):
|
||||
"""A custom api_base must replace the Google host for the endpoint-resolution GET without
|
||||
producing a malformed url (no ':' grafting, no doubled /v1)."""
|
||||
url = VertexAIBatchPrediction._build_endpoint_resolution_url(
|
||||
api_base=api_base,
|
||||
model=f"projects/{PROJECT}/locations/{LOCATION}/endpoints/{ENDPOINT_ID}",
|
||||
vertex_location=LOCATION,
|
||||
)
|
||||
assert url == expected
|
||||
|
||||
|
||||
def test_create_batch_sync_endpoint_resolution_error_raises():
|
||||
h = _make_handler()
|
||||
client = MagicMock()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue