fix(anthropic): validate federated credentials off the event loop in async create_file and create_batch

This commit is contained in:
mateo-berri 2026-08-29 14:40:22 -07:00
parent 230b7f6bea
commit 142092ae87
2 changed files with 193 additions and 26 deletions

View file

@ -238,7 +238,7 @@ class _AsyncFilesEnvironmentValidator(Protocol):
async def _avalidate_files_environment(
provider_config: BaseFilesConfig,
provider_config: BaseFilesConfig | BaseBatchesConfig,
*,
headers: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
model: str,
@ -3550,6 +3550,19 @@ class BaseLLMHTTPHandler:
"""
Creates a file using Gemini's two-step upload process
"""
if _is_async:
return self._avalidate_and_create_file(
create_file_data=create_file_data,
litellm_params=litellm_params,
provider_config=provider_config,
headers=headers,
api_base=api_base,
api_key=api_key,
logging_obj=logging_obj,
client=client,
timeout=timeout,
)
# get config from model, custom llm provider
headers = provider_config.validate_environment(
api_key=api_key,
@ -3579,18 +3592,6 @@ class BaseLLMHTTPHandler:
optional_params={},
)
if _is_async:
return self.async_create_file(
transformed_request=transformed_request,
litellm_params=litellm_params,
provider_config=provider_config,
headers=headers,
api_base=api_base,
logging_obj=logging_obj,
client=client,
timeout=timeout,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client()
else:
@ -3718,6 +3719,54 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params_with_url,
)
async def _avalidate_and_create_file(
self,
*,
create_file_data: CreateFileRequest,
litellm_params: dict, # mutable-ok: mirrors the create_file contract this dispatches for
provider_config: BaseFilesConfig,
headers: dict, # mutable-ok: mirrors the create_file contract this dispatches for
api_base: str | None,
api_key: str | None,
logging_obj: LiteLLMLoggingObj,
client: HTTPHandler | AsyncHTTPHandler | None,
timeout: float | httpx.Timeout | None,
) -> OpenAIFileObject:
validated_headers: Final = await _avalidate_files_environment(
provider_config,
headers=headers,
model="",
messages=[],
optional_params={},
litellm_params=litellm_params,
api_key=api_key,
)
complete_api_base: Final = provider_config.get_complete_file_url(
api_base=api_base,
api_key=api_key,
model="",
optional_params={},
litellm_params=litellm_params,
data=create_file_data,
)
if complete_api_base is None:
raise ValueError("api_base is required for create_file")
return await self.async_create_file(
transformed_request=provider_config.transform_create_file_request(
model="",
create_file_data=create_file_data,
litellm_params=litellm_params,
optional_params={},
),
litellm_params=litellm_params,
provider_config=provider_config,
headers=validated_headers,
api_base=complete_api_base,
logging_obj=logging_obj,
client=client,
timeout=timeout,
)
async def async_create_file(
self,
transformed_request: Union[bytes, str, dict, "TwoStepFileUploadConfig"],
@ -3968,6 +4017,20 @@ class BaseLLMHTTPHandler:
if model is None:
raise ValueError("model is required for create_batch")
if _is_async:
return self._avalidate_and_create_batch(
create_batch_data=create_batch_data,
litellm_params=litellm_params,
provider_config=provider_config,
headers=headers,
api_base=api_base,
api_key=api_key,
logging_obj=logging_obj,
client=client,
timeout=timeout,
model=model,
)
headers = provider_config.validate_environment(
api_key=api_key,
headers=headers,
@ -3996,19 +4059,6 @@ class BaseLLMHTTPHandler:
optional_params={},
)
if _is_async:
return self.async_create_batch(
transformed_request=transformed_request,
litellm_params=litellm_params,
provider_config=provider_config,
headers=headers,
api_base=api_base,
logging_obj=logging_obj,
client=client,
timeout=timeout,
create_batch_data=create_batch_data,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client()
else:
@ -4145,6 +4195,56 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params,
)
async def _avalidate_and_create_batch(
self,
*,
create_batch_data: "CreateBatchRequest",
litellm_params: dict, # mutable-ok: mirrors the create_batch contract this dispatches for
provider_config: "BaseBatchesConfig",
headers: dict, # mutable-ok: mirrors the create_batch contract this dispatches for
api_base: str | None,
api_key: str | None,
logging_obj: "LiteLLMLoggingObj",
client: Union["HTTPHandler", "AsyncHTTPHandler"] | None,
timeout: float | httpx.Timeout | None,
model: str,
) -> "LiteLLMBatch":
validated_headers: Final = await _avalidate_files_environment(
provider_config,
headers=headers,
model=model,
messages=[],
optional_params={},
litellm_params=litellm_params,
api_key=api_key,
)
complete_api_base: Final = provider_config.get_complete_batch_url(
api_base=api_base,
api_key=api_key,
model=model,
optional_params={},
litellm_params=litellm_params,
data=create_batch_data,
)
if complete_api_base is None:
raise ValueError("api_base is required for create_batch")
return await self.async_create_batch(
transformed_request=provider_config.transform_create_batch_request(
model=model,
create_batch_data=create_batch_data,
litellm_params=litellm_params,
optional_params={},
),
litellm_params=litellm_params,
provider_config=provider_config,
headers=validated_headers,
api_base=complete_api_base,
logging_obj=logging_obj,
client=client,
timeout=timeout,
create_batch_data=create_batch_data,
)
async def async_create_batch(
self,
transformed_request: bytes | str | dict,

View file

@ -18,7 +18,9 @@ from litellm.llms.base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
BaseAudioTranscriptionConfig,
)
from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.files.transformation import BaseFilesConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
BaseLLMHTTPHandler,
@ -3105,3 +3107,68 @@ async def test_a_provider_that_keeps_rejecting_is_not_retried_forever_on_the_asy
)
assert len(recorder.bodies) == 2
def _async_client_returning(response: Mock) -> AsyncMock:
client = AsyncMock(spec=AsyncHTTPHandler)
client.post.return_value = response
return client
@pytest.mark.asyncio
async def test_create_file_async_awaits_the_provider_credential_hook_instead_of_blocking():
provider_config = Mock(spec=BaseFilesConfig)
provider_config.validate_environment.side_effect = AssertionError("sync validate_environment ran on the event loop")
provider_config.avalidate_environment = AsyncMock(return_value={"x-api-key": "federated"})
provider_config.get_complete_file_url.return_value = "https://files.example/v1/files"
provider_config.transform_create_file_request.return_value = {"file": ("batch.jsonl", b"{}", "application/jsonl")}
file_object = object()
provider_config.transform_create_file_response.return_value = file_object
client = _async_client_returning(Mock(spec=httpx.Response))
result = await BaseLLMHTTPHandler().create_file(
create_file_data={"file": b"{}", "purpose": "batch"},
litellm_params={},
provider_config=provider_config,
headers={},
api_base=None,
api_key=None,
logging_obj=Mock(),
_is_async=True,
client=client,
)
assert result is file_object
provider_config.validate_environment.assert_not_called()
provider_config.avalidate_environment.assert_awaited_once()
assert client.post.call_args.kwargs["headers"] == {"x-api-key": "federated"}
assert client.post.call_args.kwargs["url"] == "https://files.example/v1/files"
@pytest.mark.asyncio
async def test_create_batch_async_validates_credentials_off_the_event_loop():
provider_config = Mock(spec=BaseBatchesConfig)
provider_config.validate_environment.side_effect = lambda **_: {"x-validated-on": str(threading.get_ident())}
provider_config.get_complete_batch_url.return_value = "https://batches.example/v1/messages/batches"
provider_config.transform_create_batch_request.return_value = {"requests": []}
batch = object()
provider_config.transform_create_batch_response.return_value = batch
client = _async_client_returning(Mock(spec=httpx.Response))
result = await BaseLLMHTTPHandler().create_batch(
create_batch_data={"input_file_id": "file_1", "endpoint": "/v1/chat/completions", "completion_window": "24h"},
litellm_params={},
provider_config=provider_config,
headers={},
api_base=None,
api_key=None,
logging_obj=Mock(),
_is_async=True,
client=client,
model="claude-sonnet-4-5",
)
assert result is batch
validated_on = client.post.call_args.kwargs["headers"]["x-validated-on"]
assert validated_on != str(threading.get_ident())
assert client.post.call_args.kwargs["url"] == "https://batches.example/v1/messages/batches"