diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 935b96a0e39..781ef28f5f6 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -107,6 +107,13 @@ async def create_batch( # noqa: PLR0915 route_type="acreate_batch", ) + # Remap proxy-injected headers to extra_headers; + # batch TypedDicts use extra_headers, not headers (fixes #27641) + if "headers" in data: + _proxy_headers = data.pop("headers") + if _proxy_headers: + data["extra_headers"] = {**(data.get("extra_headers") or {}), **_proxy_headers} + ## check if model is a loadbalanced model router_model: Optional[str] = None is_router_model = False @@ -394,6 +401,13 @@ async def retrieve_batch( # noqa: PLR0915 route_type="aretrieve_batch", ) + # Remap proxy-injected headers to extra_headers; + # batch TypedDicts use extra_headers, not headers (fixes #27641) + if "headers" in data: + _proxy_headers = data.pop("headers") + if _proxy_headers: + data["extra_headers"] = {**(data.get("extra_headers") or {}), **_proxy_headers} + # FIX: First, try to read from ManagedObjectTable for consistent state managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") from litellm.proxy.proxy_server import prisma_client @@ -653,6 +667,13 @@ async def list_batches( route_type="alist_batches", ) + # Remap proxy-injected headers to extra_headers; + # batch TypedDicts use extra_headers, not headers (fixes #27641) + if "headers" in data: + _proxy_headers = data.pop("headers") + if _proxy_headers: + data["extra_headers"] = {**(data.get("extra_headers") or {}), **_proxy_headers} + # Try to use managed objects table for listing batches (returns encoded IDs) managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is not None and hasattr( @@ -852,6 +873,13 @@ async def cancel_batch( proxy_config=proxy_config, ) + # Remap proxy-injected headers to extra_headers; + # batch TypedDicts use extra_headers, not headers (fixes #27641) + if "headers" in data: + _proxy_headers = data.pop("headers") + if _proxy_headers: + data["extra_headers"] = {**(data.get("extra_headers") or {}), **_proxy_headers} + # SCENARIO 1: Batch ID is encoded with model info if model_from_id is not None: credentials = get_credentials_for_model(