mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(rust): clarify Bedrock bridge configuration
Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
f13425bb4f
commit
30e810abc5
6 changed files with 15 additions and 7 deletions
|
|
@ -1,6 +1,7 @@
|
|||
use litellm_core::CoreError;
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::messages::transformation::MessagesAuthStrategy;
|
||||
use litellm_core::providers::bedrock::AWS_REGION_NAME;
|
||||
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use serde_json::Value;
|
||||
|
||||
|
|
@ -29,10 +30,10 @@ pub(super) fn prepare_messages_call(
|
|||
|
||||
let config = messages_provider_config(provider)
|
||||
.ok_or_else(|| CoreError::InvalidProvider(provider.to_string()))?;
|
||||
let region_override = request.aws_region_name;
|
||||
let env_lookup = |key: &str| {
|
||||
if key == "AWS_REGION_NAME" {
|
||||
return request
|
||||
.aws_region_name
|
||||
if key == AWS_REGION_NAME {
|
||||
return region_override
|
||||
.map(str::to_string)
|
||||
.or_else(|| std::env::var(key).ok());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ pub struct AnthropicMessage {
|
|||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct AnthropicMessagesRequest {
|
||||
#[serde(skip_serializing_if = "String::is_empty")]
|
||||
#[serde(default, skip_serializing_if = "String::is_empty")]
|
||||
pub model: String,
|
||||
pub messages: Vec<AnthropicMessage>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
|
|||
|
|
@ -6,4 +6,5 @@
|
|||
pub mod audio_transcription;
|
||||
pub mod aws_base;
|
||||
mod constants;
|
||||
pub use constants::AWS_REGION_NAME;
|
||||
pub mod messages;
|
||||
|
|
|
|||
|
|
@ -2283,7 +2283,6 @@ class BaseLLMHTTPHandler:
|
|||
from litellm.rust_bridge import messages as rust_messages_bridge
|
||||
|
||||
upstream_body = {key: value for key, value in request_body.items() if key != "stream"}
|
||||
upstream_body.setdefault("model", model)
|
||||
try:
|
||||
rust_response = await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -111,7 +111,7 @@ def messages(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
**({"aws_region_name": aws_region_name} if aws_region_name is not None else {}),
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -137,5 +137,5 @@ async def amessages(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
**({"aws_region_name": aws_region_name} if aws_region_name is not None else {}),
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ class RecordingMessages:
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
aws_region_name: str | None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -56,6 +57,7 @@ class RecordingMessages:
|
|||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"aws_region_name": aws_region_name,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
|
@ -74,6 +76,7 @@ class RecordingAsyncMessages:
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
aws_region_name: str | None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -84,6 +87,7 @@ class RecordingAsyncMessages:
|
|||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"aws_region_name": aws_region_name,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
|
@ -192,6 +196,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
|
|||
"custom_llm_provider": "azure_ai",
|
||||
"extra_headers": {"anthropic-beta": "token-efficient-tools-2025-02-19"},
|
||||
"timeout_seconds": 42.0,
|
||||
"aws_region_name": None,
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -213,6 +218,7 @@ async def test_amessages_wrapper_forwards_args():
|
|||
assert response == FAKE_MESSAGES_RESPONSE
|
||||
assert bridge.calls[0]["model"] == "claude-sonnet-4-5"
|
||||
assert bridge.calls[0]["timeout_seconds"] == 12.5
|
||||
assert bridge.calls[0]["aws_region_name"] is None
|
||||
|
||||
|
||||
def _gate(**overrides):
|
||||
|
|
@ -248,6 +254,7 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
assert call["api_base"] == "https://resource.services.ai.azure.com/anthropic"
|
||||
assert call["extra_headers"] == {"x-api-key": "sk-azure", "anthropic-version": "2023-06-01"}
|
||||
assert call["timeout_seconds"] == 30.0
|
||||
assert call["aws_region_name"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue