mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(anthropic): validate federated credentials off the event loop in async create_file and create_batch
This commit is contained in:
parent
230b7f6bea
commit
142092ae87
2 changed files with 193 additions and 26 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue