mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(vertex_ai): stop grafting batch storage and job urls onto the deployment api_base
A deployment api_base points at online inference, commonly a full .../endpoints/<id>:rawPredict resource url. The batch file upload used it as the GCS storage host and batch create/retrieve/list/cancel grafted job paths onto it, both producing urls Google answers with an HTML 404. Uploads now always target storage.googleapis.com, and batch operations ignore a resource-shaped api_base (path containing /projects/) while still honoring host-level or /v1 gateway mounts.
This commit is contained in:
parent
1183b2abc6
commit
697da95bf4
4 changed files with 85 additions and 10 deletions
|
|
@ -72,6 +72,19 @@ class _VertexEndpointPayloadView(TypedDict):
|
|||
payload: ReadOnly[_VertexEndpointResponse]
|
||||
|
||||
|
||||
def _gateway_api_base_or_none(api_base: str | None) -> str | None:
|
||||
"""
|
||||
A deployment `api_base` whose path names a concrete Vertex resource (contains `/projects/`,
|
||||
e.g. the `.../endpoints/<id>:rawPredict` url configured for online inference) is not a Vertex
|
||||
API gateway; grafting `batchPredictionJobs` or resource-GET paths onto it can only produce
|
||||
urls Google answers with an HTML 404 (LIT-7386). Batch operations ignore it and use the real
|
||||
Vertex host; only a host-level or `/v1`-style gateway mount passes through.
|
||||
"""
|
||||
if api_base and "/projects/" in urlparse(api_base).path:
|
||||
return None
|
||||
return api_base
|
||||
|
||||
|
||||
def _vertex_batch_payload(response: _VertexBatchJsonSource) -> VertexBatchPredictionResponse:
|
||||
return response.json()
|
||||
|
||||
|
|
@ -126,11 +139,12 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
vertex_location=vertex_location or "us-central1",
|
||||
)
|
||||
)
|
||||
gateway_api_base: Final = _gateway_api_base_or_none(api_base)
|
||||
vertex_batch_request: Final = self._resolve_fine_tuned_endpoint_model(
|
||||
vertex_batch_request=transformed_batch_request,
|
||||
headers=headers,
|
||||
sync_handler=sync_handler,
|
||||
api_base=api_base,
|
||||
api_base=gateway_api_base,
|
||||
vertex_location=vertex_location or "us-central1",
|
||||
)
|
||||
|
||||
|
|
@ -145,7 +159,7 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
endpoint = ""
|
||||
|
||||
_, api_base = self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
api_base=gateway_api_base,
|
||||
custom_llm_provider="vertex_ai",
|
||||
gemini_api_key=None,
|
||||
endpoint=endpoint,
|
||||
|
|
@ -325,7 +339,7 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
endpoint = ""
|
||||
|
||||
_, api_base = self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
api_base=_gateway_api_base_or_none(api_base),
|
||||
custom_llm_provider="vertex_ai",
|
||||
gemini_api_key=None,
|
||||
endpoint=endpoint,
|
||||
|
|
@ -481,7 +495,7 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
endpoint = ""
|
||||
|
||||
_, api_base = self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
api_base=_gateway_api_base_or_none(api_base),
|
||||
custom_llm_provider="vertex_ai",
|
||||
gemini_api_key=None,
|
||||
endpoint=endpoint,
|
||||
|
|
@ -579,7 +593,7 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
cancel_api_base_default: Final = f"{retrieve_api_base_default}:cancel"
|
||||
|
||||
_, api_base = self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
api_base=_gateway_api_base_or_none(api_base),
|
||||
custom_llm_provider="vertex_ai",
|
||||
gemini_api_key=None,
|
||||
endpoint="cancel",
|
||||
|
|
|
|||
|
|
@ -809,11 +809,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
object_name = f"{object_prefix}/{object_name}"
|
||||
encoded_object_name: Final = encode_gcs_object_name_for_url(object_name)
|
||||
endpoint: Final = f"upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={encoded_object_name}"
|
||||
api_base = api_base or "https://storage.googleapis.com"
|
||||
if not api_base:
|
||||
raise ValueError("api_base is required")
|
||||
|
||||
return f"{api_base}/{endpoint}"
|
||||
return f"https://storage.googleapis.com/{endpoint}"
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]:
|
||||
return []
|
||||
|
|
|
|||
|
|
@ -257,6 +257,45 @@ def test_create_batch_sync_resolves_fine_tuned_endpoint_to_tuned_model():
|
|||
assert sent["model"] == TUNED_MODEL_RESOURCE
|
||||
|
||||
|
||||
def test_create_batch_sync_ignores_resource_shaped_api_base():
|
||||
"""A deployment api_base like `.../endpoints/<id>:rawPredict` targets online inference, not
|
||||
the Vertex API root; grafting batch urls onto it yields guaranteed 404s, so batch operations
|
||||
must fall back to the default Vertex host (LIT-7386)."""
|
||||
h = _make_handler()
|
||||
client = MagicMock()
|
||||
client.post.return_value = _http_response()
|
||||
raw_predict_api_base = (
|
||||
f"https://{LOCATION}-aiplatform.googleapis.com/v1/projects/{PROJECT}"
|
||||
f"/locations/{LOCATION}/endpoints/{ENDPOINT_ID}:rawPredict"
|
||||
)
|
||||
|
||||
with (
|
||||
patch(f"{HMOD}._get_httpx_client", return_value=client),
|
||||
patch(f"{HMOD}.safe_get", return_value=_endpoint_get_response()) as safe_get,
|
||||
):
|
||||
out = h.create_batch(
|
||||
_is_async=False,
|
||||
create_batch_data=ENDPOINT_CREATE_DATA,
|
||||
api_base=raw_predict_api_base,
|
||||
vertex_credentials=None,
|
||||
vertex_project=PROJECT,
|
||||
vertex_location=LOCATION,
|
||||
timeout=600.0,
|
||||
max_retries=None,
|
||||
)
|
||||
|
||||
assert isinstance(out, LiteLLMBatch)
|
||||
resolution_url = safe_get.call_args.args[1]
|
||||
assert ":rawPredict" not in resolution_url
|
||||
assert resolution_url == (
|
||||
f"https://{LOCATION}-aiplatform.googleapis.com/v1/projects/{PROJECT}"
|
||||
f"/locations/{LOCATION}/endpoints/{ENDPOINT_ID}"
|
||||
)
|
||||
assert h._check_custom_proxy.call_args.kwargs["api_base"] is None
|
||||
sent = json.loads(client.post.call_args.kwargs["data"])
|
||||
assert sent["model"] == TUNED_MODEL_RESOURCE
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base, expected",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -158,6 +158,32 @@ class TestCreateFileUrl:
|
|||
assert ".." not in object_name
|
||||
assert "?" not in object_name
|
||||
|
||||
def test_should_upload_to_gcs_host_even_when_deployment_sets_api_base(self, config):
|
||||
"""The deployment api_base points at the inference endpoint (often a full
|
||||
`.../endpoints/<id>:rawPredict` URL); grafting the GCS upload onto it produces a
|
||||
guaranteed 404 from Google, so the storage host must stay storage.googleapis.com
|
||||
(LIT-7386)."""
|
||||
url = config.get_complete_file_url(
|
||||
api_base=(
|
||||
"https://us-central1-aiplatform.googleapis.com/v1/projects/my-project"
|
||||
"/locations/us-central1/endpoints/6335039103326748672:rawPredict"
|
||||
),
|
||||
api_key=None,
|
||||
model="",
|
||||
optional_params={},
|
||||
litellm_params={
|
||||
"gcs_bucket_name": "my-bucket",
|
||||
"model": "vertex_ai/gemini/6335039103326748672",
|
||||
},
|
||||
data={
|
||||
"file": ("batch.jsonl", b'{"body": {"model": "gemini-2.5-flash"}}', "application/jsonl"),
|
||||
"purpose": "batch",
|
||||
},
|
||||
)
|
||||
assert url.startswith("https://storage.googleapis.com/upload/storage/v1/b/my-bucket/o?")
|
||||
assert "aiplatform" not in url
|
||||
assert "rawPredict" not in url
|
||||
|
||||
|
||||
class TestBatchObjectNaming:
|
||||
def test_should_store_publisher_model_under_publishers_path(self, config):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue