mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #34847 from BerriAI/litellm_fix_vertex_batch_read_per_model_bucket
fix(vertex_ai): honor per-model gcs_bucket_name on managed-file read path
This commit is contained in:
commit
6fe1e73699
2 changed files with 177 additions and 52 deletions
|
|
@ -1,7 +1,9 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from urllib.parse import unquote
|
||||
from typing import Any, Coroutine, Optional, Tuple, Union
|
||||
from typing import Any, Coroutine, Mapping, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -10,6 +12,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import (
|
|||
GCSBucketBase,
|
||||
GCSLoggingConfig,
|
||||
)
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
VERTEX_AI_MANAGED_GCS_PREFIX,
|
||||
should_allow_legacy_cloud_file_ids,
|
||||
|
|
@ -39,6 +42,35 @@ class VertexAIFilesHandler(GCSBucketBase):
|
|||
llm_provider=LlmProviders.VERTEX_AI,
|
||||
)
|
||||
|
||||
def _resolve_read_gcs_config(
|
||||
self,
|
||||
litellm_params: Mapping[str, object] | None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""
|
||||
Resolve the GCS bucket and service-account credentials for the read/content path.
|
||||
|
||||
Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` /
|
||||
``bucket_name`` and ``vertex_credentials``), mirroring the write path in
|
||||
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global
|
||||
``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch
|
||||
run entirely at the model-group level, so output written to a per-model bucket is
|
||||
readable without setting the global env vars.
|
||||
"""
|
||||
params: Mapping[str, object] = litellm_params or {}
|
||||
bucket_candidate = params.get("gcs_bucket_name") or params.get("bucket_name")
|
||||
configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME")
|
||||
|
||||
credentials = params.get("vertex_credentials") or vertex_credentials
|
||||
if isinstance(credentials, dict):
|
||||
path_service_account: str | None = json.dumps(credentials)
|
||||
elif isinstance(credentials, str):
|
||||
path_service_account = credentials
|
||||
else:
|
||||
path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT")
|
||||
|
||||
return configured_bucket_name, path_service_account
|
||||
|
||||
def _extract_bucket_and_object_from_file_id(
|
||||
self,
|
||||
file_id: str,
|
||||
|
|
@ -91,7 +123,17 @@ class VertexAIFilesHandler(GCSBucketBase):
|
|||
if not file_id:
|
||||
raise ValueError("file_id is required in file_content_request")
|
||||
|
||||
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(kwargs={})
|
||||
configured_bucket_name, path_service_account = self._resolve_read_gcs_config(
|
||||
litellm_params=litellm_params,
|
||||
vertex_credentials=vertex_credentials,
|
||||
)
|
||||
dynamic_params = StandardCallbackDynamicParams(
|
||||
gcs_bucket_name=configured_bucket_name,
|
||||
gcs_path_service_account=path_service_account,
|
||||
)
|
||||
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(
|
||||
kwargs={"standard_callback_dynamic_params": dynamic_params}
|
||||
)
|
||||
bucket_name, object_path = self._extract_bucket_and_object_from_file_id(
|
||||
file_id=file_id,
|
||||
configured_bucket_name=gcs_logging_config["bucket_name"],
|
||||
|
|
|
|||
|
|
@ -31,10 +31,7 @@ class TestVertexAIFilesHandler:
|
|||
def test_extract_bucket_and_object_from_file_id_standard_path(self):
|
||||
"""Test extraction of bucket and object from URL-encoded file_id with standard path"""
|
||||
# Sample file_id with nested folder structure
|
||||
file_id = (
|
||||
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
|
||||
"%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
|
||||
)
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
|
||||
|
||||
bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id(
|
||||
file_id=file_id,
|
||||
|
|
@ -105,21 +102,14 @@ class TestVertexAIFilesHandler:
|
|||
async def test_afile_content_success(self):
|
||||
"""Test successful async file content retrieval"""
|
||||
# Setup test data
|
||||
file_id = (
|
||||
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
|
||||
"%2Fuploads%2Fabc-test-file.txt"
|
||||
)
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt"
|
||||
expected_content = b"test file content"
|
||||
|
||||
file_content_request = FileContentRequest(
|
||||
file_id=file_id, extra_headers=None, extra_body=None
|
||||
)
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
# Mock the download_gcs_object method
|
||||
with (
|
||||
patch.object(
|
||||
self.handler, "download_gcs_object", new_callable=AsyncMock
|
||||
) as mock_download,
|
||||
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
|
||||
patch.object(
|
||||
self.handler,
|
||||
"get_gcs_logging_config",
|
||||
|
|
@ -148,15 +138,9 @@ class TestVertexAIFilesHandler:
|
|||
# Verify the download was called with correct parameters
|
||||
mock_download.assert_called_once()
|
||||
call_args = mock_download.call_args
|
||||
assert (
|
||||
call_args.kwargs["object_name"]
|
||||
== "litellm-vertex-files/uploads/abc-test-file.txt"
|
||||
)
|
||||
assert call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-test-file.txt"
|
||||
assert "standard_callback_dynamic_params" in call_args.kwargs
|
||||
assert (
|
||||
call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"]
|
||||
== "test-bucket"
|
||||
)
|
||||
assert call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] == "test-bucket"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_missing_file_id(self):
|
||||
|
|
@ -164,9 +148,7 @@ class TestVertexAIFilesHandler:
|
|||
file_content_request = FileContentRequest(extra_headers=None, extra_body=None)
|
||||
|
||||
# Should raise ValueError for missing file_id
|
||||
with pytest.raises(
|
||||
ValueError, match="file_id is required in file_content_request"
|
||||
):
|
||||
with pytest.raises(ValueError, match="file_id is required in file_content_request"):
|
||||
await self.handler.afile_content(
|
||||
file_content_request=file_content_request,
|
||||
vertex_credentials=None,
|
||||
|
|
@ -179,20 +161,13 @@ class TestVertexAIFilesHandler:
|
|||
@pytest.mark.asyncio
|
||||
async def test_afile_content_download_failure(self):
|
||||
"""Test async file content retrieval when download fails"""
|
||||
file_id = (
|
||||
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
|
||||
"%2Fuploads%2Fabc-test-file.txt"
|
||||
)
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt"
|
||||
|
||||
file_content_request = FileContentRequest(
|
||||
file_id=file_id, extra_headers=None, extra_body=None
|
||||
)
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
# Mock download to return None (failure)
|
||||
with (
|
||||
patch.object(
|
||||
self.handler, "download_gcs_object", new_callable=AsyncMock
|
||||
) as mock_download,
|
||||
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
|
||||
patch.object(
|
||||
self.handler,
|
||||
"get_gcs_logging_config",
|
||||
|
|
@ -216,14 +191,130 @@ class TestVertexAIFilesHandler:
|
|||
max_retries=3,
|
||||
)
|
||||
|
||||
def test_resolve_read_gcs_config_prefers_per_model_bucket(self, monkeypatch):
|
||||
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
|
||||
monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json")
|
||||
|
||||
bucket, service_account = self.handler._resolve_read_gcs_config(
|
||||
litellm_params={
|
||||
"gcs_bucket_name": "my-model-bucket",
|
||||
"vertex_credentials": "/model/sa.json",
|
||||
},
|
||||
vertex_credentials=None,
|
||||
)
|
||||
|
||||
assert bucket == "my-model-bucket"
|
||||
assert service_account == "/model/sa.json"
|
||||
|
||||
def test_resolve_read_gcs_config_falls_back_to_env(self, monkeypatch):
|
||||
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
|
||||
monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json")
|
||||
|
||||
bucket, service_account = self.handler._resolve_read_gcs_config(litellm_params={}, vertex_credentials=None)
|
||||
|
||||
assert bucket == "env-default-bucket"
|
||||
assert service_account == "/env/sa.json"
|
||||
|
||||
def test_resolve_read_gcs_config_serializes_dict_credentials(self, monkeypatch):
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
|
||||
_, service_account = self.handler._resolve_read_gcs_config(
|
||||
litellm_params={"gcs_bucket_name": "my-model-bucket"},
|
||||
vertex_credentials={"type": "service_account", "project_id": "p"},
|
||||
)
|
||||
|
||||
assert service_account == '{"type": "service_account", "project_id": "p"}'
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_honors_per_model_bucket_over_env(self, monkeypatch):
|
||||
"""
|
||||
Regression for #32640: a batch output written to a per-model gcs_bucket_name must be
|
||||
readable even when the global GCS_BUCKET_NAME points at a different bucket. Before the
|
||||
fix the read path resolved the bucket from env only and raised
|
||||
"file_id bucket does not match the configured storage bucket".
|
||||
"""
|
||||
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
|
||||
file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl"
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
with (
|
||||
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
|
||||
patch.object(
|
||||
self.handler,
|
||||
"get_or_create_vertex_instance",
|
||||
new_callable=AsyncMock,
|
||||
return_value=object(),
|
||||
),
|
||||
):
|
||||
mock_download.return_value = b"batch output"
|
||||
|
||||
result = await self.handler.afile_content(
|
||||
file_content_request=file_content_request,
|
||||
vertex_credentials="/model/sa.json",
|
||||
vertex_project="test-project",
|
||||
vertex_location="us-central1",
|
||||
timeout=60.0,
|
||||
max_retries=0,
|
||||
litellm_params={
|
||||
"gcs_bucket_name": "my-model-bucket",
|
||||
"vertex_credentials": "/model/sa.json",
|
||||
},
|
||||
)
|
||||
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
assert result.response.content == b"batch output"
|
||||
|
||||
dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"]
|
||||
assert dynamic_params["gcs_bucket_name"] == "my-model-bucket"
|
||||
assert dynamic_params["gcs_path_service_account"] == "/model/sa.json"
|
||||
assert mock_download.call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-batch-output.jsonl"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_reads_without_global_env_bucket(self, monkeypatch):
|
||||
"""
|
||||
Regression for #32640: with no global GCS_BUCKET_NAME set, a model-group-level
|
||||
deployment (per-model gcs_bucket_name) must still be readable. Before the fix the read
|
||||
path raised "GCS_BUCKET_NAME is not set in the environment".
|
||||
"""
|
||||
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
|
||||
file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl"
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
with (
|
||||
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
|
||||
patch.object(
|
||||
self.handler,
|
||||
"get_or_create_vertex_instance",
|
||||
new_callable=AsyncMock,
|
||||
return_value=object(),
|
||||
),
|
||||
):
|
||||
mock_download.return_value = b"batch output"
|
||||
|
||||
result = await self.handler.afile_content(
|
||||
file_content_request=file_content_request,
|
||||
vertex_credentials="/model/sa.json",
|
||||
vertex_project="test-project",
|
||||
vertex_location="us-central1",
|
||||
timeout=60.0,
|
||||
max_retries=0,
|
||||
litellm_params={"gcs_bucket_name": "my-model-bucket"},
|
||||
)
|
||||
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"]
|
||||
assert dynamic_params["gcs_bucket_name"] == "my-model-bucket"
|
||||
|
||||
def test_file_content_sync_success(self):
|
||||
"""Test successful sync file content retrieval"""
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
|
||||
expected_content = b"test file content"
|
||||
|
||||
file_content_request = FileContentRequest(
|
||||
file_id=file_id, extra_headers=None, extra_body=None
|
||||
)
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
# Create expected response
|
||||
mock_response = httpx.Response(
|
||||
|
|
@ -261,25 +352,17 @@ class TestVertexAIFilesHandler:
|
|||
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
|
||||
expected_content = b"test file content"
|
||||
|
||||
file_content_request = FileContentRequest(
|
||||
file_id=file_id, extra_headers=None, extra_body=None
|
||||
)
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
# Mock the afile_content method
|
||||
with patch.object(
|
||||
self.handler, "afile_content", new_callable=AsyncMock
|
||||
) as mock_afile_content:
|
||||
with patch.object(self.handler, "afile_content", new_callable=AsyncMock) as mock_afile_content:
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=expected_content,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request(
|
||||
method="GET", url="gs://test-bucket/test-file.txt"
|
||||
),
|
||||
)
|
||||
mock_afile_content.return_value = HttpxBinaryResponseContent(
|
||||
response=mock_response
|
||||
request=httpx.Request(method="GET", url="gs://test-bucket/test-file.txt"),
|
||||
)
|
||||
mock_afile_content.return_value = HttpxBinaryResponseContent(response=mock_response)
|
||||
|
||||
# Call the method with _is_async=True
|
||||
result = self.handler.file_content(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue