mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(proxy/files): respect files_settings and model_info gcs bucket for vertex managed uploads
This commit is contained in:
parent
5d25e75f3b
commit
81b24c5059
4 changed files with 339 additions and 2 deletions
|
|
@ -54,6 +54,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
from litellm.proxy.utils import ProxyLogging, is_known_model
|
||||
from litellm.repositories.table_repositories import ManagedFileRepository
|
||||
from litellm.router import Router
|
||||
from litellm.types.router import LiteLLMParamsTypedDict
|
||||
from litellm.types.llms.openai import (
|
||||
CREATE_FILE_REQUESTS_PURPOSE,
|
||||
FileExpiresAfter,
|
||||
|
|
@ -87,9 +88,9 @@ def get_files_provider_config(
|
|||
custom_llm_provider: str,
|
||||
):
|
||||
global files_config
|
||||
if custom_llm_provider == "vertex_ai":
|
||||
return None
|
||||
if files_config is None:
|
||||
if custom_llm_provider == "vertex_ai":
|
||||
return None
|
||||
raise ValueError("files_settings is not set, set it on your config.yaml file.")
|
||||
for setting in files_config:
|
||||
if setting.get("custom_llm_provider") == custom_llm_provider:
|
||||
|
|
@ -97,6 +98,88 @@ def get_files_provider_config(
|
|||
return None
|
||||
|
||||
|
||||
def get_files_settings_bucket_for_provider(custom_llm_provider: str) -> Optional[str]:
|
||||
"""
|
||||
Return the GCS bucket configured under a top-level ``files_settings`` entry for
|
||||
a provider, or None when ``files_settings`` is unset or has no bucket for it.
|
||||
|
||||
Unlike ``get_files_provider_config`` this never raises, so it is safe to call
|
||||
on the managed-files (``target_model_names``) upload path where ``files_settings``
|
||||
is optional.
|
||||
"""
|
||||
global files_config
|
||||
if files_config is None:
|
||||
return None
|
||||
for setting in files_config:
|
||||
if setting.get("custom_llm_provider") == custom_llm_provider:
|
||||
return setting.get("gcs_bucket_name") or setting.get("bucket_name")
|
||||
return None
|
||||
|
||||
|
||||
def _deployment_is_vertex_ai(litellm_params: LiteLLMParamsTypedDict) -> bool:
|
||||
provider = litellm_params.get("custom_llm_provider")
|
||||
if provider:
|
||||
return provider == "vertex_ai"
|
||||
return str(litellm_params.get("model") or "").startswith("vertex_ai/")
|
||||
|
||||
|
||||
def _model_info_bucket(model_info: Optional[dict]) -> Optional[str]:
|
||||
if not model_info:
|
||||
return None
|
||||
for key in ("gcs_bucket_name", "bucket_name"):
|
||||
value = model_info.get(key)
|
||||
if isinstance(value, str) and value:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def resolve_vertex_files_bucket_override(
|
||||
llm_router: Optional[Router],
|
||||
target_model_names_list: List[str],
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Resolve the GCS bucket for a managed Vertex AI batch/file upload.
|
||||
|
||||
precedence (most specific wins):
|
||||
per-model litellm_params bucket > per-model model_info bucket
|
||||
> global files_settings bucket > GCS_BUCKET_NAME env var
|
||||
|
||||
The router already forwards each deployment's ``litellm_params`` to the Vertex
|
||||
upload, so a per-model ``litellm_params`` bucket is left untouched here (returns
|
||||
None so nothing is injected). This resolves the two lower-priority sources that
|
||||
would otherwise be dropped for managed uploads: per-model ``model_info`` and the
|
||||
top-level ``files_settings`` block. Env resolution stays in the transformation.
|
||||
"""
|
||||
if llm_router is None:
|
||||
return None
|
||||
|
||||
vertex_deployments = tuple(
|
||||
deployment
|
||||
for model in target_model_names_list
|
||||
for deployment in (llm_router.get_model_list(model_name=model) or [])
|
||||
if _deployment_is_vertex_ai(deployment["litellm_params"])
|
||||
)
|
||||
if not vertex_deployments:
|
||||
return None
|
||||
|
||||
if any(
|
||||
deployment["litellm_params"].get("gcs_bucket_name") or deployment["litellm_params"].get("bucket_name")
|
||||
for deployment in vertex_deployments
|
||||
):
|
||||
return None
|
||||
|
||||
model_info_bucket = next(
|
||||
(
|
||||
bucket
|
||||
for deployment in vertex_deployments
|
||||
for bucket in (_model_info_bucket(deployment.get("model_info")),)
|
||||
if bucket
|
||||
),
|
||||
None,
|
||||
)
|
||||
return model_info_bucket or get_files_settings_bucket_for_provider("vertex_ai")
|
||||
|
||||
|
||||
def get_first_json_object(file_source: Union[bytes, BinaryIO]) -> Optional[dict]:
|
||||
try:
|
||||
if isinstance(file_source, (bytes, bytearray)):
|
||||
|
|
@ -215,6 +298,13 @@ async def route_create_file(
|
|||
# Handle managed files (supports loadbalancing via llm_router.acreate_file)
|
||||
# Priority: Check for managed files BEFORE deprecated loadbalancing
|
||||
if target_model_names_list:
|
||||
if not (_create_file_request.get("gcs_bucket_name") or _create_file_request.get("bucket_name")):
|
||||
bucket_override = resolve_vertex_files_bucket_override(
|
||||
llm_router=llm_router,
|
||||
target_model_names_list=target_model_names_list,
|
||||
)
|
||||
if bucket_override:
|
||||
_create_file_request["gcs_bucket_name"] = bucket_override
|
||||
managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
|
||||
if managed_files_obj is None:
|
||||
raise ProxyException(
|
||||
|
|
|
|||
|
|
@ -402,6 +402,8 @@ class CreateFileRequest(TypedDict, total=False):
|
|||
extra_headers: Optional[Dict[str, str]]
|
||||
extra_body: Optional[Dict[str, str]]
|
||||
timeout: Optional[float]
|
||||
gcs_bucket_name: Optional[str]
|
||||
bucket_name: Optional[str]
|
||||
|
||||
|
||||
class FileContentRequest(TypedDict, total=False):
|
||||
|
|
|
|||
|
|
@ -385,6 +385,8 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
|
|||
## VERTEX AI ##
|
||||
vertex_project: Optional[str]
|
||||
vertex_location: Optional[str]
|
||||
gcs_bucket_name: Optional[str]
|
||||
bucket_name: Optional[str]
|
||||
## AWS BEDROCK / SAGEMAKER ##
|
||||
aws_access_key_id: Optional[str]
|
||||
aws_secret_access_key: Optional[str]
|
||||
|
|
|
|||
|
|
@ -2610,3 +2610,246 @@ def test_list_files_with_all_proxy_models_team_uses_openai_deployment(
|
|||
assert captured_kwargs.get("api_key") == "team-openai-key"
|
||||
assert captured_kwargs.get("custom_llm_provider") == "openai"
|
||||
proxy_logging_obj.post_call_failure_hook.assert_not_called()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def _reset_files_config():
|
||||
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
|
||||
|
||||
original = fe.files_config
|
||||
yield
|
||||
fe.files_config = original
|
||||
|
||||
|
||||
class TestVertexFilesBucketOverride:
|
||||
"""
|
||||
Regression tests for https://github.com/BerriAI/litellm/issues/33419
|
||||
|
||||
Vertex AI batch/file uploads used to ignore both the top-level files_settings
|
||||
block and per-model model_info, so the destination GCS bucket could only come
|
||||
from the GCS_BUCKET_NAME env var (shared with the gcs_bucket logging callback).
|
||||
"""
|
||||
|
||||
def test_get_files_provider_config_reads_vertex_files_settings(
|
||||
self, _reset_files_config
|
||||
):
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
get_files_provider_config,
|
||||
set_files_config,
|
||||
)
|
||||
|
||||
set_files_config(
|
||||
[{"custom_llm_provider": "vertex_ai", "gcs_bucket_name": "batch-bucket"}]
|
||||
)
|
||||
setting = get_files_provider_config(custom_llm_provider="vertex_ai")
|
||||
assert setting is not None
|
||||
assert setting["gcs_bucket_name"] == "batch-bucket"
|
||||
|
||||
def test_get_files_provider_config_vertex_without_files_settings_returns_none(
|
||||
self, _reset_files_config
|
||||
):
|
||||
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
get_files_provider_config,
|
||||
)
|
||||
|
||||
fe.files_config = None
|
||||
assert get_files_provider_config(custom_llm_provider="vertex_ai") is None
|
||||
|
||||
def test_get_files_provider_config_non_vertex_without_files_settings_raises(
|
||||
self, _reset_files_config
|
||||
):
|
||||
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
get_files_provider_config,
|
||||
)
|
||||
|
||||
fe.files_config = None
|
||||
with pytest.raises(ValueError, match="files_settings is not set"):
|
||||
get_files_provider_config(custom_llm_provider="openai")
|
||||
|
||||
def _vertex_router(self, model_info: dict = None, litellm_params_extra: dict = None):
|
||||
litellm_params = {"model": "vertex_ai/gemini-2.0-flash"}
|
||||
if litellm_params_extra:
|
||||
litellm_params.update(litellm_params_extra)
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "vertex-batch",
|
||||
"litellm_params": litellm_params,
|
||||
"model_info": {"id": "vertex-batch-id", **(model_info or {})},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
def test_resolve_uses_global_files_settings_bucket(self, _reset_files_config):
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
resolve_vertex_files_bucket_override,
|
||||
set_files_config,
|
||||
)
|
||||
|
||||
set_files_config(
|
||||
[{"custom_llm_provider": "vertex_ai", "gcs_bucket_name": "global-batch-bucket"}]
|
||||
)
|
||||
assert (
|
||||
resolve_vertex_files_bucket_override(
|
||||
llm_router=self._vertex_router(),
|
||||
target_model_names_list=["vertex-batch"],
|
||||
)
|
||||
== "global-batch-bucket"
|
||||
)
|
||||
|
||||
def test_resolve_prefers_model_info_over_files_settings(self, _reset_files_config):
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
resolve_vertex_files_bucket_override,
|
||||
set_files_config,
|
||||
)
|
||||
|
||||
set_files_config(
|
||||
[{"custom_llm_provider": "vertex_ai", "gcs_bucket_name": "global-batch-bucket"}]
|
||||
)
|
||||
assert (
|
||||
resolve_vertex_files_bucket_override(
|
||||
llm_router=self._vertex_router(
|
||||
model_info={"gcs_bucket_name": "per-model-bucket"}
|
||||
),
|
||||
target_model_names_list=["vertex-batch"],
|
||||
)
|
||||
== "per-model-bucket"
|
||||
)
|
||||
|
||||
def test_resolve_defers_to_router_when_litellm_params_pins_bucket(
|
||||
self, _reset_files_config
|
||||
):
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
resolve_vertex_files_bucket_override,
|
||||
set_files_config,
|
||||
)
|
||||
|
||||
set_files_config(
|
||||
[{"custom_llm_provider": "vertex_ai", "gcs_bucket_name": "global-batch-bucket"}]
|
||||
)
|
||||
assert (
|
||||
resolve_vertex_files_bucket_override(
|
||||
llm_router=self._vertex_router(
|
||||
litellm_params_extra={"gcs_bucket_name": "per-deployment-bucket"}
|
||||
),
|
||||
target_model_names_list=["vertex-batch"],
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
def test_resolve_ignores_non_vertex_targets(self, _reset_files_config):
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
resolve_vertex_files_bucket_override,
|
||||
set_files_config,
|
||||
)
|
||||
|
||||
set_files_config(
|
||||
[{"custom_llm_provider": "vertex_ai", "gcs_bucket_name": "global-batch-bucket"}]
|
||||
)
|
||||
openai_router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-batch",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-x"},
|
||||
}
|
||||
]
|
||||
)
|
||||
assert (
|
||||
resolve_vertex_files_bucket_override(
|
||||
llm_router=openai_router,
|
||||
target_model_names_list=["gpt-batch"],
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
def test_route_create_file_injects_files_settings_bucket_for_managed_vertex_upload(
|
||||
self, _reset_files_config
|
||||
):
|
||||
import asyncio
|
||||
|
||||
from litellm import CreateFileRequest
|
||||
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
route_create_file,
|
||||
set_files_config,
|
||||
)
|
||||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
|
||||
set_files_config(
|
||||
[{"custom_llm_provider": "vertex_ai", "gcs_bucket_name": "global-batch-bucket"}]
|
||||
)
|
||||
llm_router = self._vertex_router()
|
||||
|
||||
proxy_logging_obj = ProxyLogging(
|
||||
user_api_key_cache=DualCache(default_in_memory_ttl=1)
|
||||
)
|
||||
proxy_logging_obj._add_proxy_hooks(llm_router)
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
class DummyManagedFiles(BaseFileEndpoints):
|
||||
async def acreate_file(
|
||||
self,
|
||||
llm_router,
|
||||
create_file_request,
|
||||
target_model_names_list,
|
||||
litellm_parent_otel_span,
|
||||
user_api_key_dict,
|
||||
):
|
||||
captured["gcs_bucket_name"] = create_file_request.get("gcs_bucket_name")
|
||||
return OpenAIFileObject(
|
||||
id="file-abc123",
|
||||
object="file",
|
||||
bytes=100,
|
||||
created_at=1234567890,
|
||||
filename="mydata.jsonl",
|
||||
purpose="batch",
|
||||
status="uploaded",
|
||||
)
|
||||
|
||||
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
|
||||
raise NotImplementedError
|
||||
|
||||
async def afile_list(self, purpose, litellm_parent_otel_span):
|
||||
raise NotImplementedError
|
||||
|
||||
async def afile_delete(
|
||||
self, file_id, litellm_parent_otel_span, llm_router, **data
|
||||
):
|
||||
raise NotImplementedError
|
||||
|
||||
async def afile_content(
|
||||
self, file_id, litellm_parent_otel_span, llm_router, **data
|
||||
):
|
||||
raise NotImplementedError
|
||||
|
||||
proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="test-team",
|
||||
parent_otel_span=None,
|
||||
)
|
||||
_create_file_request = CreateFileRequest(
|
||||
file=("mydata.jsonl", b'{"model": "vertex-batch"}', "application/json"),
|
||||
purpose="batch",
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
route_create_file(
|
||||
llm_router=llm_router,
|
||||
_create_file_request=_create_file_request,
|
||||
purpose="batch",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
target_model_names_list=["vertex-batch"],
|
||||
is_router_model=False,
|
||||
router_model=None,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
)
|
||||
|
||||
assert captured["gcs_bucket_name"] == "global-batch-bucket"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue