diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py index 3bc09139f8f..4d2a1e18eb5 100644 --- a/litellm/llms/vertex_ai/files/handler.py +++ b/litellm/llms/vertex_ai/files/handler.py @@ -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"], diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py index 453a0c14bf9..5e854bbad70 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py @@ -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(