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:
mubashir1osmani 2026-09-09 14:33:29 -04:00
parent 1183b2abc6
commit 697da95bf4
4 changed files with 85 additions and 10 deletions

View file

@ -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",

View file

@ -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 []

View file

@ -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",
[

View file

@ -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):