mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
refactor(rust): give Messages configs typed litellm params (#45611)
* refactor(rust): give Messages configs typed litellm params * feat(rust): port _normalize_system_role_messages for Messages hosts mid_conversation_system mirrors Python's module function for function: a leading run of system turns is hoisted into the top-level system, a later turn stays in place when the model map flags supports_mid_conversation_system and otherwise becomes a user turn where it was, never between a tool_use and its tool_result, and billing blocks are stripped from the hoisted system. The capability joins MessagesModelCapabilities and the bridge projection. Azure AI uses it in place of fold_system_role_messages, which hoisted every turn. * fix(rust): keep the Bedrock runtime endpoint behind a blank api_base * refactor(rust): import LitellmParams from its crate instead of a re-export * refactor(rust): template the Messages request decode errors Both decode failures go through ErrorDetail::invalid with the subject and the serde error as source instead of a preformatted sentence. * refactor(rust_bridge): read the module globals the param specs name The Messages host no longer keeps its own list of litellm globals; a spec whose spellings the call leaves out names the global to read, so the next family that has one needs no host change. * refactor(rust): give each route lookup Python's implicit None ProviderConfigManager's per-route methods list only the providers the route serves and fall through to None; the Rust lookups now do the same, so a new LlmProviders variant touches only the route that serves it.
This commit is contained in:
parent
3d2abf9a44
commit
ee295aea68
39 changed files with 901 additions and 240 deletions
4
litellm-rust/Cargo.lock
generated
4
litellm-rust/Cargo.lock
generated
|
|
@ -4031,6 +4031,7 @@ dependencies = [
|
|||
"litellm-inference-testing",
|
||||
"litellm-llms",
|
||||
"litellm-llms-types",
|
||||
"litellm-router-types",
|
||||
"litellm-secrets",
|
||||
"litellm-tracing",
|
||||
"reqwest 0.12.28",
|
||||
|
|
@ -4151,6 +4152,7 @@ dependencies = [
|
|||
"litellm-http",
|
||||
"litellm-llms-types",
|
||||
"litellm-python-compat",
|
||||
"litellm-router-types",
|
||||
"litellm-secrets",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
|
|
@ -4222,6 +4224,7 @@ dependencies = [
|
|||
"litellm-inference-transcription",
|
||||
"litellm-llms",
|
||||
"litellm-llms-types",
|
||||
"litellm-router-types",
|
||||
"litellm-secrets",
|
||||
"litellm-secrets-aws",
|
||||
"litellm-secrets-types",
|
||||
|
|
@ -4270,6 +4273,7 @@ dependencies = [
|
|||
"litellm-config",
|
||||
"litellm-inference",
|
||||
"litellm-inference-messages",
|
||||
"litellm-router-types",
|
||||
"rstest",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -79,6 +79,7 @@ fn project(
|
|||
api_key: deployment.api_key.clone(),
|
||||
api_base: deployment.api_base.clone(),
|
||||
custom_llm_provider: deployment.custom_llm_provider.clone(),
|
||||
litellm_params: deployment.litellm_params.clone(),
|
||||
extra_headers: None,
|
||||
provider_specific_header: anthropic_api_headers(headers),
|
||||
timeout: deployment.timeout,
|
||||
|
|
|
|||
|
|
@ -33,13 +33,7 @@ pub(super) fn chat_completions_provider(provider: LlmProviders) -> Option<ChatPr
|
|||
LlmProviders::Anthropic => Some(ChatProvider::Anthropic),
|
||||
LlmProviders::Bedrock => Some(ChatProvider::Bedrock),
|
||||
LlmProviders::OpenaiLike => Some(ChatProvider::OpenaiLike),
|
||||
LlmProviders::AwsTextract
|
||||
| LlmProviders::AzureAi
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ litellm-http.workspace = true
|
|||
litellm-inference.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-llms-types.workspace = true
|
||||
litellm-router-types.workspace = true
|
||||
litellm-secrets.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
reqwest.workspace = true
|
||||
|
|
|
|||
|
|
@ -44,13 +44,7 @@ pub(crate) fn messages_provider(provider: LlmProviders) -> Option<MessagesProvid
|
|||
LlmProviders::Anthropic => Some(MessagesProvider::Anthropic),
|
||||
LlmProviders::AzureAi => Some(MessagesProvider::AzureAi),
|
||||
LlmProviders::Bedrock => Some(MessagesProvider::Bedrock),
|
||||
LlmProviders::AwsTextract
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -70,7 +70,9 @@ impl MessagesRoute {
|
|||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
body: serde_json::to_value(&body).map_err(serialize_failure)?,
|
||||
body: provider
|
||||
.config()
|
||||
.wire_body(serde_json::to_value(&body).map_err(serialize_failure)?),
|
||||
},
|
||||
request_context,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -15,7 +15,8 @@ use std::sync::Arc;
|
|||
|
||||
pub use litellm_inference::RouteError as Error;
|
||||
pub use types::{
|
||||
MessagesCall, MessagesCallResponse, MessagesSettings, MessagesShaping, messages_body,
|
||||
MessagesCall, MessagesCallResponse, MessagesSettings, MessagesShaping, litellm_params,
|
||||
messages_body,
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ fn prepare_provider_request(
|
|||
body,
|
||||
api_key,
|
||||
api_base,
|
||||
litellm_params,
|
||||
extra_headers,
|
||||
provider_specific_header,
|
||||
timeout,
|
||||
|
|
@ -98,6 +99,7 @@ fn prepare_provider_request(
|
|||
forwarded,
|
||||
api_key.as_deref(),
|
||||
&transformed.model,
|
||||
&litellm_params,
|
||||
&env_lookup,
|
||||
)?;
|
||||
let environment = ValidatedEnvironment {
|
||||
|
|
@ -108,11 +110,13 @@ fn prepare_provider_request(
|
|||
auth: validated.auth,
|
||||
};
|
||||
|
||||
let url = if transformed.params.stream == Some(true) {
|
||||
config.complete_stream_url(api_base.as_deref(), &transformed.model, &env_lookup)?
|
||||
} else {
|
||||
config.get_complete_url(api_base.as_deref(), &transformed.model, &env_lookup)?
|
||||
};
|
||||
let url = config.get_complete_url(
|
||||
api_base.as_deref(),
|
||||
&transformed.model,
|
||||
&litellm_params,
|
||||
transformed.params.stream == Some(true),
|
||||
&env_lookup,
|
||||
)?;
|
||||
|
||||
Ok(ProviderMessagesRequest {
|
||||
provider,
|
||||
|
|
@ -227,6 +231,7 @@ mod tests {
|
|||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
litellm_params: Default::default(),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
|
|
@ -253,6 +258,7 @@ mod tests {
|
|||
api_key: Some("sk-test".into()),
|
||||
api_base: Some("https://anthropic.test".into()),
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
litellm_params: Default::default(),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
|
|
@ -360,6 +366,7 @@ mod tests {
|
|||
api_key: Some("sk-test".into()),
|
||||
api_base: Some("https://resource.services.ai.azure.com".into()),
|
||||
custom_llm_provider: custom_llm_provider.map(Into::into),
|
||||
litellm_params: Default::default(),
|
||||
extra_headers: Some(Map::from_iter([("x-priority".into(), json!("extra"))])),
|
||||
provider_specific_header: Some(configured),
|
||||
timeout: None,
|
||||
|
|
|
|||
|
|
@ -2,11 +2,12 @@ use std::time::Duration;
|
|||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
|
||||
use litellm_llms::{ErrorDetail, base_llm::messages::context::MessagesModelCapabilities};
|
||||
use litellm_llms_types::{
|
||||
formats::messages::{MessagesRequest, MessagesResponse},
|
||||
headers::ProviderSpecificHeaders,
|
||||
};
|
||||
use litellm_router_types::LitellmParams;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
|
|
@ -17,6 +18,7 @@ pub struct MessagesCall {
|
|||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub litellm_params: LitellmParams,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
|
|
@ -27,8 +29,16 @@ pub fn messages_body(body: Map<String, Value>) -> Result<MessagesRequest, Error>
|
|||
serde_json::from_value(Value::Object(body)).map_err(invalid_request)
|
||||
}
|
||||
|
||||
/// The caller's litellm params, projected by a host from the keys [`LitellmParams::fields`]
|
||||
/// names. A key present with a value of the wrong type is a request error, as it is for
|
||||
/// Python's `GenericLiteLLMParams(**kwargs)`.
|
||||
pub fn litellm_params(fields: Map<String, Value>) -> Result<LitellmParams, Error> {
|
||||
serde_json::from_value(Value::Object(fields))
|
||||
.map_err(|err| Error::InvalidRequest(ErrorDetail::invalid("litellm params", err)))
|
||||
}
|
||||
|
||||
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}").into())
|
||||
Error::InvalidRequest(ErrorDetail::invalid("Anthropic messages request", err))
|
||||
}
|
||||
|
||||
pub type MessagesCallResponse =
|
||||
|
|
@ -126,6 +136,7 @@ mod tests {
|
|||
supports_output_config: true,
|
||||
supports_sampling_params: false,
|
||||
supports_speed: true,
|
||||
supports_mid_conversation_system: false,
|
||||
effort_tiers: SupportedEffortTiers {
|
||||
minimal: false,
|
||||
low: true,
|
||||
|
|
|
|||
|
|
@ -206,7 +206,7 @@ async fn messages_cache_identity_follows_resolved_configuration_and_request_call
|
|||
let cache = ScopedCache::new(cache.clone(), CacheScope::Shared);
|
||||
support::messages_route(secrets.clone()).with_cache(cache).execute(MessagesCall {
|
||||
body: serde_json::from_value(json!({"model":"anthropic/cache-test-model","messages":[{"role":"user","content":"hello"}],"max_tokens":32})).unwrap(),
|
||||
api_key:None,api_base:None,custom_llm_provider:None,extra_headers:None,provider_specific_header:None,timeout:None,shaping:Default::default(),
|
||||
api_key:None,api_base:None,custom_llm_provider:None,litellm_params:Default::default(),extra_headers:None,provider_specific_header:None,timeout:None,shaping:Default::default(),
|
||||
}, &hooks, None).await.unwrap();
|
||||
}
|
||||
assert_eq!(hooks.calls.load(Ordering::SeqCst), 4);
|
||||
|
|
|
|||
|
|
@ -81,6 +81,7 @@ fn call() -> MessagesCall {
|
|||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
litellm_params: Default::default(),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
|
|
|
|||
|
|
@ -58,6 +58,7 @@ async fn credentials_become_exactly_one_auth_header(
|
|||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
api_key: api_key.map(Into::into),
|
||||
api_base: Some(upstream.uri()),
|
||||
extra_headers: headers(extra_headers.iter().copied()),
|
||||
|
|
@ -85,6 +86,7 @@ async fn a_call_without_credentials_fails_before_sending(
|
|||
|
||||
let error = run(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
|
|
@ -124,6 +126,7 @@ async fn each_provider_posts_to_its_messages_endpoint(
|
|||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: provider.map(Into::into),
|
||||
litellm_params: Default::default(),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(format!("{}{base_suffix}", upstream.uri())),
|
||||
..with_model(call, model)
|
||||
|
|
@ -154,6 +157,7 @@ async fn unsupported_providers_are_rejected_before_sending(
|
|||
) {
|
||||
let error = run(MessagesCall {
|
||||
custom_llm_provider: provider.map(Into::into),
|
||||
litellm_params: Default::default(),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(UNREACHABLE_BASE.into()),
|
||||
..with_model(call, model)
|
||||
|
|
@ -206,6 +210,7 @@ async fn cache_scope_removal_is_selected_by_the_provider(
|
|||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
api_key: Some("sk-azure".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
body: body(json!({
|
||||
|
|
@ -349,6 +354,7 @@ async fn caller_protocol_headers_win_over_the_defaults(call: MessagesCall, #[cas
|
|||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
extra_headers: headers([
|
||||
|
|
@ -400,6 +406,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
|
|||
},
|
||||
body: call.body.clone(),
|
||||
custom_llm_provider: call.custom_llm_provider.clone(),
|
||||
litellm_params: Default::default(),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: call.timeout,
|
||||
|
|
@ -577,6 +584,7 @@ async fn metadata_is_reduced_to_the_user_id(call: MessagesCall, #[case] provider
|
|||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
|
|
@ -638,6 +646,7 @@ async fn system_message_folding_is_selected_by_the_provider(
|
|||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
api_key: Some("sk-azure".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
|
|
@ -702,6 +711,7 @@ async fn provider_validation_runs_before_caller_parameter_removal(
|
|||
let result = run(with_fields(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
api_key: Some("sk-test".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
|
|
|
|||
|
|
@ -104,6 +104,7 @@ async fn the_provider_message_is_returned(call: MessagesCall, #[case] provider:
|
|||
|
||||
let message = run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ async fn the_credential_and_base_come_from_the_secret_source(
|
|||
secrets.clone(),
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
..call
|
||||
},
|
||||
)
|
||||
|
|
@ -189,6 +190,7 @@ async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall)
|
|||
Arc::new(RecordingSecrets::new([("AZURE_API_KEY", "sk-azure")])),
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
litellm_params: Default::default(),
|
||||
..call
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -346,6 +346,7 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte(
|
|||
.execute(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
litellm_params: Default::default(),
|
||||
..streaming(call, upstream.uri())
|
||||
},
|
||||
&(),
|
||||
|
|
@ -470,6 +471,7 @@ async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) {
|
|||
let host = RecordingStreamHost::new(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
litellm_params: Default::default(),
|
||||
..streaming(call, upstream.uri())
|
||||
},
|
||||
usize::MAX,
|
||||
|
|
|
|||
|
|
@ -21,14 +21,7 @@ pub(crate) fn prepare_request(
|
|||
Some("MISTRAL_AZURE_API_BASE"),
|
||||
),
|
||||
LlmProviders::AzureAi => (None, Some("AZURE_AI_API_BASE")),
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::AwsTextract
|
||||
| LlmProviders::Bedrock
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => (None, None),
|
||||
_ => (None, None),
|
||||
};
|
||||
let secret = |name: &str| secrets.truthy(name);
|
||||
let dynamic_api_key = credentials.dynamic_api_key.or_else(|| {
|
||||
|
|
|
|||
|
|
@ -225,10 +225,7 @@ pub(crate) fn resolve_provider_config(
|
|||
OcrConfigKind::VertexDeepSeek
|
||||
}
|
||||
LlmProviders::VertexAi => OcrConfigKind::VertexAi,
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::Bedrock
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike => {
|
||||
_ => {
|
||||
return Err(Error::InvalidProvider(
|
||||
provider.custom_llm_provider.to_string(),
|
||||
));
|
||||
|
|
|
|||
|
|
@ -17,15 +17,7 @@ use litellm_inference::provider::resolve_llm_provider;
|
|||
fn provider_config(provider: LlmProviders) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
|
||||
match provider {
|
||||
LlmProviders::Bedrock => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG),
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::AwsTextract
|
||||
| LlmProviders::AzureAi
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ litellm-framer.workspace = true
|
|||
litellm-http.workspace = true
|
||||
litellm-secrets.workspace = true
|
||||
litellm-python-compat.workspace = true
|
||||
litellm-router-types.workspace = true
|
||||
base64.workspace = true
|
||||
bytes.workspace = true
|
||||
data-url = "0.3.2"
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ use litellm_http::request::{
|
|||
use litellm_llms_types::{
|
||||
formats::messages::{
|
||||
ContentBlock, ContentBlockType, EffortLevel, Message, MessageContent, MessagesTool,
|
||||
SystemPrompt,
|
||||
},
|
||||
providers::anthropic::{AnthropicBeta, BetaSet},
|
||||
recognized::Recognized,
|
||||
|
|
@ -385,6 +386,33 @@ pub fn is_encrypted_reasoning_block(block: &ContentBlock) -> bool {
|
|||
field.is_some_and(|value| value.starts_with(ENCRYPTED_REASONING_SIGNATURE_PREFIX))
|
||||
}
|
||||
|
||||
const BILLING_HEADER_PREFIX: &str = "x-anthropic-billing-header:";
|
||||
|
||||
fn is_billing_header_block(block: &ContentBlock) -> bool {
|
||||
block.block_type == Some(ContentBlockType::Text)
|
||||
&& block
|
||||
.text
|
||||
.as_deref()
|
||||
.is_some_and(|text| text.starts_with(BILLING_HEADER_PREFIX))
|
||||
}
|
||||
|
||||
/// Python's `AnthropicMessagesConfig._filter_billing_headers_from_system`: the Claude Code
|
||||
/// attribution blocks the first-party API reads, dropped for hosts that reject them. `None`
|
||||
/// when nothing else was in the system prompt.
|
||||
pub fn filter_billing_headers_from_system(system: SystemPrompt) -> Option<SystemPrompt> {
|
||||
match system {
|
||||
SystemPrompt::Text(text) if text.starts_with(BILLING_HEADER_PREFIX) => None,
|
||||
SystemPrompt::Text(text) => Some(SystemPrompt::Text(text)),
|
||||
SystemPrompt::Blocks(blocks) => {
|
||||
let kept: Vec<ContentBlock> = blocks
|
||||
.into_iter()
|
||||
.filter(|block| !is_billing_header_block(block))
|
||||
.collect();
|
||||
(!kept.is_empty()).then_some(SystemPrompt::Blocks(kept))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn strip_encrypted_reasoning_blocks(messages: Vec<Message>) -> Vec<Message> {
|
||||
retain_blocks(messages, |block| !is_encrypted_reasoning_block(block))
|
||||
}
|
||||
|
|
@ -1858,6 +1886,7 @@ mod tests {
|
|||
supports_output_config: false,
|
||||
supports_sampling_params: true,
|
||||
supports_speed: false,
|
||||
supports_mid_conversation_system: false,
|
||||
effort_tiers: tiers(false, false, false, false, false, false),
|
||||
}
|
||||
);
|
||||
|
|
|
|||
|
|
@ -3,9 +3,10 @@ use litellm_llms_types::{
|
|||
formats::messages::{
|
||||
ContextEdit, ContextManagement, Message, MessagesOptionalParams, MessagesRequest, Speed,
|
||||
},
|
||||
providers::anthropic::{AnthropicBeta, BetaSet},
|
||||
providers::anthropic::{AnthropicBeta, BetaProvider, BetaSet},
|
||||
recognized::Recognized,
|
||||
};
|
||||
use litellm_router_types::LitellmParams;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::{handler::shape_anthropic_messages_request, thinking::translate_thinking};
|
||||
|
|
@ -47,6 +48,8 @@ impl BaseMessagesConfig for AnthropicMessagesConfig {
|
|||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_litellm_params: &LitellmParams,
|
||||
_stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(complete_anthropic_url(api_base, env_lookup))
|
||||
|
|
@ -76,6 +79,7 @@ impl BaseMessagesConfig for AnthropicMessagesConfig {
|
|||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
_model: &str,
|
||||
_litellm_params: &LitellmParams,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
let headers = match optionally_handle_anthropic_oauth(headers, api_key) {
|
||||
|
|
@ -150,6 +154,15 @@ pub(crate) fn update_headers_with_anthropic_beta(
|
|||
merge_beta_headers(headers, feature_betas(request))
|
||||
}
|
||||
|
||||
/// The betas a request needs on a host other than the first-party API: the Anthropic set
|
||||
/// filtered and renamed through the host's `BetaProvider` policy.
|
||||
pub fn provider_feature_betas(request: &MessagesRequest, provider: BetaProvider) -> BetaSet {
|
||||
feature_betas(request)
|
||||
.iter()
|
||||
.filter_map(|beta| beta.on(provider))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn feature_betas(request: &MessagesRequest) -> BetaSet {
|
||||
let params = &request.params;
|
||||
let tools = params.tools.as_deref();
|
||||
|
|
@ -815,7 +828,13 @@ mod tests {
|
|||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(
|
||||
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(api_base, "claude", &env(vars)),
|
||||
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(
|
||||
api_base,
|
||||
"claude",
|
||||
&LitellmParams::default(),
|
||||
false,
|
||||
&env(vars),
|
||||
),
|
||||
Ok(expected.to_string())
|
||||
);
|
||||
}
|
||||
|
|
@ -833,6 +852,7 @@ mod tests {
|
|||
headers(forwarded),
|
||||
api_key,
|
||||
"claude",
|
||||
&LitellmParams::default(),
|
||||
&env(vars),
|
||||
)
|
||||
}
|
||||
|
|
@ -1108,8 +1128,20 @@ mod tests {
|
|||
requested.borrow_mut().push(name.to_string());
|
||||
None
|
||||
};
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.validate_environment(Vec::new(), None, "claude", &record);
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.validate_environment(
|
||||
Vec::new(),
|
||||
None,
|
||||
"claude",
|
||||
&LitellmParams::default(),
|
||||
&record,
|
||||
);
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(
|
||||
None,
|
||||
"claude",
|
||||
&LitellmParams::default(),
|
||||
false,
|
||||
&record,
|
||||
);
|
||||
let requested = requested.into_inner();
|
||||
assert!(!requested.is_empty());
|
||||
let undeclared: Vec<&String> = requested
|
||||
|
|
|
|||
|
|
@ -1,9 +1,7 @@
|
|||
use litellm_auth::{CredentialPlacement, SecretValue};
|
||||
use litellm_http::request::{has_bearer_auth, has_header};
|
||||
use litellm_llms_types::formats::messages::{
|
||||
CacheControl, ContentBlock, Message, MessageContent, MessagesOptionalParams, MessagesRequest,
|
||||
SystemPrompt,
|
||||
};
|
||||
use litellm_llms_types::formats::messages::MessagesRequest;
|
||||
use litellm_router_types::LitellmParams;
|
||||
|
||||
use crate::{
|
||||
Error,
|
||||
|
|
@ -20,7 +18,7 @@ use crate::{
|
|||
auth::{AuthScheme, Headers, ValidatedEnvironment},
|
||||
messages::{
|
||||
context::MessagesTransformContext,
|
||||
normalization::fold_system_role_messages,
|
||||
normalization::{normalize_system_role_messages, strip_cache_control_scope},
|
||||
transformation::{BaseMessagesConfig, MESSAGES_PATH_SUFFIX},
|
||||
},
|
||||
},
|
||||
|
|
@ -47,6 +45,8 @@ impl BaseMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_litellm_params: &LitellmParams,
|
||||
_stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
complete_azure_anthropic_url(api_base, env_lookup)
|
||||
|
|
@ -57,20 +57,14 @@ impl BaseMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
request: MessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<MessagesRequest, Error> {
|
||||
let request = fold_system_role_messages(request);
|
||||
transform_messages_request(
|
||||
MessagesRequest {
|
||||
messages: request
|
||||
.messages
|
||||
.into_iter()
|
||||
.map(strip_scope_from_message)
|
||||
.collect(),
|
||||
params: MessagesOptionalParams {
|
||||
system: request.params.system.map(strip_scope_from_system),
|
||||
..request.params
|
||||
},
|
||||
..request
|
||||
},
|
||||
strip_cache_control_scope(normalize_system_role_messages(
|
||||
request,
|
||||
context
|
||||
.thinking
|
||||
.capabilities
|
||||
.supports_mid_conversation_system,
|
||||
)),
|
||||
context,
|
||||
)
|
||||
}
|
||||
|
|
@ -86,6 +80,7 @@ impl BaseMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
_model: &str,
|
||||
_litellm_params: &LitellmParams,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
if has_header(&headers, API_KEY_PLACEMENT.header_name()) || has_bearer_auth(&headers) {
|
||||
|
|
@ -129,37 +124,6 @@ pub fn complete_azure_anthropic_url(
|
|||
Ok(format!("{with_anthropic}{MESSAGES_PATH_SUFFIX}"))
|
||||
}
|
||||
|
||||
fn strip_scope_from_block(block: ContentBlock) -> ContentBlock {
|
||||
ContentBlock {
|
||||
cache_control: block.cache_control.map(|cache_control| CacheControl {
|
||||
scope: None,
|
||||
..cache_control
|
||||
}),
|
||||
..block
|
||||
}
|
||||
}
|
||||
|
||||
fn strip_scope_from_system(system: SystemPrompt) -> SystemPrompt {
|
||||
match system {
|
||||
SystemPrompt::Blocks(blocks) => {
|
||||
SystemPrompt::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
|
||||
}
|
||||
text => text,
|
||||
}
|
||||
}
|
||||
|
||||
fn strip_scope_from_message(message: Message) -> Message {
|
||||
Message {
|
||||
content: match message.content {
|
||||
MessageContent::Blocks(blocks) => {
|
||||
MessageContent::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
|
||||
}
|
||||
text => text,
|
||||
},
|
||||
..message
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms_types::formats::messages::MessagesResponse;
|
||||
|
|
@ -169,7 +133,9 @@ mod tests {
|
|||
use litellm_auth::CredentialPlacement;
|
||||
|
||||
use super::*;
|
||||
use crate::base_llm::messages::context::MessagesModelCapabilities;
|
||||
use crate::base_llm::messages::{
|
||||
context::MessagesModelCapabilities, mid_conversation_system::CONVERTED_SYSTEM_NOTE,
|
||||
};
|
||||
|
||||
fn request_from(value: serde_json::Value) -> MessagesRequest {
|
||||
serde_json::from_value(value).expect("valid request")
|
||||
|
|
@ -268,6 +234,7 @@ mod tests {
|
|||
.collect(),
|
||||
api_key,
|
||||
"claude",
|
||||
&LitellmParams::default(),
|
||||
&|_| None,
|
||||
)
|
||||
.unwrap()
|
||||
|
|
@ -427,62 +394,101 @@ mod tests {
|
|||
assert_eq!(transformed, body);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_request_folds_system_role_message_into_top_level_system() {
|
||||
let request = request_from(json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 256,
|
||||
"system": [{"type": "text", "text": "base system"}],
|
||||
"messages": [
|
||||
{"role": "user", "content": "fix the bug"},
|
||||
{"role": "system", "content": "Available agent types: claude"}
|
||||
]
|
||||
}));
|
||||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
transformed["messages"],
|
||||
json!([{"role": "user", "content": "fix the bug"}])
|
||||
);
|
||||
assert_eq!(
|
||||
transformed["system"],
|
||||
json!([
|
||||
{"type": "text", "text": "base system"},
|
||||
{"type": "text", "text": "Available agent types: claude"}
|
||||
])
|
||||
);
|
||||
fn context_with(supports_mid_conversation_system: bool) -> MessagesTransformContext {
|
||||
MessagesTransformContext::new(
|
||||
MessagesModelCapabilities {
|
||||
supports_mid_conversation_system,
|
||||
..MessagesModelCapabilities::default()
|
||||
},
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_request_folds_system_role_when_no_top_level_system() {
|
||||
let request = request_from(json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 256,
|
||||
#[rstest]
|
||||
#[case::leading_turn_joins_the_top_level_system(
|
||||
json!({
|
||||
"system": [{"type": "text", "text": "base system"}],
|
||||
"messages": [
|
||||
{"role": "user", "content": [{"type": "text", "text": "hi"}]},
|
||||
{"role": "system", "content": [{"type": "text", "text": "sys block"}]}
|
||||
{"role": "system", "content": "Available agent types: claude"},
|
||||
{"role": "user", "content": "fix the bug"}
|
||||
]
|
||||
}));
|
||||
}),
|
||||
false,
|
||||
json!([{"type": "text", "text": "base system"}, {"type": "text", "text": "Available agent types: claude"}]),
|
||||
json!([{"role": "user", "content": "fix the bug"}]),
|
||||
)]
|
||||
#[case::leading_turn_becomes_the_system_when_there_is_none(
|
||||
json!({"messages": [
|
||||
{"role": "system", "content": [{"type": "text", "text": "sys block"}]},
|
||||
{"role": "user", "content": [{"type": "text", "text": "hi"}]}
|
||||
]}),
|
||||
false,
|
||||
json!([{"type": "text", "text": "sys block"}]),
|
||||
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]),
|
||||
)]
|
||||
#[case::billing_block_in_a_leading_turn_is_dropped(
|
||||
json!({"messages": [
|
||||
{"role": "system", "content": [{"type": "text", "text": "x-anthropic-billing-header: cc_version=1"}, {"type": "text", "text": "keep"}]},
|
||||
{"role": "user", "content": "hi"}
|
||||
]}),
|
||||
false,
|
||||
json!([{"type": "text", "text": "keep"}]),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
)]
|
||||
#[case::later_turn_is_converted_in_place_for_a_model_without_the_capability(
|
||||
json!({
|
||||
"system": "be terse",
|
||||
"messages": [
|
||||
{"role": "user", "content": "fix the bug"},
|
||||
{"role": "system", "content": "reminder"}
|
||||
]
|
||||
}),
|
||||
false,
|
||||
json!("be terse"),
|
||||
json!([
|
||||
{"role": "user", "content": "fix the bug"},
|
||||
{"role": "user", "content": [
|
||||
{"type": "text", "text": CONVERTED_SYSTEM_NOTE},
|
||||
{"type": "text", "text": "reminder"}
|
||||
]}
|
||||
]),
|
||||
)]
|
||||
#[case::later_turn_stays_for_a_model_with_the_capability(
|
||||
json!({
|
||||
"system": "be terse",
|
||||
"messages": [
|
||||
{"role": "user", "content": "fix the bug"},
|
||||
{"role": "system", "content": "reminder"}
|
||||
]
|
||||
}),
|
||||
true,
|
||||
json!("be terse"),
|
||||
json!([
|
||||
{"role": "user", "content": "fix the bug"},
|
||||
{"role": "system", "content": "reminder"}
|
||||
]),
|
||||
)]
|
||||
fn transform_request_normalizes_system_role_messages_like_python(
|
||||
#[case] body: serde_json::Value,
|
||||
#[case] supports_mid_conversation_system: bool,
|
||||
#[case] system: serde_json::Value,
|
||||
#[case] messages: serde_json::Value,
|
||||
) {
|
||||
let mut body = body;
|
||||
body["model"] = json!("claude-sonnet-4-5");
|
||||
body["max_tokens"] = json!(256);
|
||||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.transform_anthropic_messages_request(
|
||||
request_from(body),
|
||||
&context_with(supports_mid_conversation_system),
|
||||
)
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
transformed["messages"],
|
||||
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
|
||||
);
|
||||
assert_eq!(
|
||||
transformed["system"],
|
||||
json!([{"type": "text", "text": "sys block"}])
|
||||
);
|
||||
assert_eq!(transformed["system"], system);
|
||||
assert_eq!(transformed["messages"], messages);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -607,9 +613,16 @@ mod tests {
|
|||
Vec::new(),
|
||||
None,
|
||||
"claude",
|
||||
&LitellmParams::default(),
|
||||
&record,
|
||||
);
|
||||
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(
|
||||
None,
|
||||
"claude",
|
||||
&LitellmParams::default(),
|
||||
false,
|
||||
&record,
|
||||
);
|
||||
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
|
||||
let requested = requested.into_inner();
|
||||
assert!(!requested.is_empty());
|
||||
let undeclared: Vec<&String> = requested
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@
|
|||
|
||||
use litellm_auth::{AuthServices, CredentialPlacement, SecretValue, TokenProviderHandle};
|
||||
use litellm_auth_aws::{AwsCredentialSource, SigV4Signer};
|
||||
use litellm_auth_gcp::VertexConfig;
|
||||
use litellm_http::request::with_header;
|
||||
|
||||
pub type Headers = Vec<(String, String)>;
|
||||
|
|
@ -24,6 +25,9 @@ pub enum AuthScheme {
|
|||
/// A bearer acquired when the request is sent, from a token source such as a cloud SDK
|
||||
/// or a caller-supplied callable.
|
||||
Token { provider: TokenProviderHandle },
|
||||
/// A Google access token for the Vertex AI project the config names, acquired when the
|
||||
/// request is sent through the shared GCP token cache.
|
||||
GcpAccessToken { config: Box<VertexConfig> },
|
||||
/// AWS SigV4 over the bytes that go on the wire, so the handler signs after the body is
|
||||
/// serialized.
|
||||
AwsSigV4 {
|
||||
|
|
@ -73,6 +77,13 @@ pub async fn resolve_auth(
|
|||
signer: None,
|
||||
})
|
||||
}
|
||||
AuthScheme::GcpAccessToken { config } => {
|
||||
let token = services.gcp.access_token(&config, env_lookup).await?;
|
||||
Ok(Authenticated {
|
||||
headers: with_credential(headers, CredentialPlacement::Bearer, &token),
|
||||
signer: None,
|
||||
})
|
||||
}
|
||||
AuthScheme::AwsSigV4 {
|
||||
region,
|
||||
service,
|
||||
|
|
|
|||
|
|
@ -40,6 +40,8 @@ pub struct MessagesModelCapabilities {
|
|||
#[serde(default)]
|
||||
pub supports_speed: bool,
|
||||
#[serde(default)]
|
||||
pub supports_mid_conversation_system: bool,
|
||||
#[serde(default)]
|
||||
pub effort_tiers: SupportedEffortTiers,
|
||||
}
|
||||
|
||||
|
|
@ -57,6 +59,7 @@ impl Default for MessagesModelCapabilities {
|
|||
supports_output_config: false,
|
||||
supports_sampling_params: true,
|
||||
supports_speed: false,
|
||||
supports_mid_conversation_system: false,
|
||||
effort_tiers: SupportedEffortTiers::default(),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,201 @@
|
|||
use litellm_llms_types::formats::messages::{
|
||||
ContentBlock, ContentBlockType, Message, MessageContent, SystemPrompt,
|
||||
};
|
||||
use serde_json::Map;
|
||||
|
||||
/// Python's `CONVERTED_SYSTEM_NOTE`.
|
||||
pub const CONVERTED_SYSTEM_NOTE: &str = "Operator note (not from the user): the following was originally a mid-conversation system-role reminder.";
|
||||
|
||||
const SYSTEM_ROLE: &str = "system";
|
||||
const USER_ROLE: &str = "user";
|
||||
|
||||
pub fn as_system_content_blocks(value: Option<SystemPrompt>) -> Vec<ContentBlock> {
|
||||
match value {
|
||||
None => Vec::new(),
|
||||
Some(SystemPrompt::Text(text)) => vec![ContentBlock::text(text)],
|
||||
Some(SystemPrompt::Blocks(blocks)) => blocks,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn message_content_blocks(content: MessageContent) -> Vec<ContentBlock> {
|
||||
match content {
|
||||
MessageContent::Text(text) => vec![ContentBlock::text(text)],
|
||||
MessageContent::Blocks(blocks) => blocks,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_system_role_message(message: &Message) -> bool {
|
||||
message.role == SYSTEM_ROLE
|
||||
}
|
||||
|
||||
pub fn system_role_message_as_user(message: Message) -> Message {
|
||||
Message {
|
||||
role: USER_ROLE.into(),
|
||||
content: MessageContent::Blocks(
|
||||
[ContentBlock::text(CONVERTED_SYSTEM_NOTE)]
|
||||
.into_iter()
|
||||
.chain(message_content_blocks(message.content))
|
||||
.collect(),
|
||||
),
|
||||
extra: Map::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn opens_with_tool_results(message: &Message) -> bool {
|
||||
message.role == USER_ROLE
|
||||
&& matches!(
|
||||
&message.content,
|
||||
MessageContent::Blocks(blocks)
|
||||
if blocks.first().and_then(|block| block.block_type.as_ref()) == Some(&ContentBlockType::ToolResult)
|
||||
)
|
||||
}
|
||||
|
||||
fn system_run_placed_after_tool_results(
|
||||
system_run: Vec<Message>,
|
||||
follower_run: Vec<Message>,
|
||||
) -> Vec<Message> {
|
||||
let mut follower_run = follower_run.into_iter();
|
||||
match follower_run.next() {
|
||||
Some(first) if opens_with_tool_results(&first) => [first]
|
||||
.into_iter()
|
||||
.chain(system_run)
|
||||
.chain(follower_run)
|
||||
.collect(),
|
||||
first => system_run
|
||||
.into_iter()
|
||||
.chain(first)
|
||||
.chain(follower_run)
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
fn runs(messages: Vec<Message>) -> Vec<Vec<Message>> {
|
||||
messages
|
||||
.into_iter()
|
||||
.fold(Vec::new(), |mut runs: Vec<Vec<Message>>, message| {
|
||||
match runs.last_mut() {
|
||||
Some(run)
|
||||
if is_system_role_message(&run[0]) == is_system_role_message(&message) =>
|
||||
{
|
||||
run.push(message);
|
||||
}
|
||||
_ => runs.push(vec![message]),
|
||||
}
|
||||
runs
|
||||
})
|
||||
}
|
||||
|
||||
fn system_turns_after_tool_results(messages: Vec<Message>) -> Vec<Message> {
|
||||
let mut runs = runs(messages).into_iter();
|
||||
let mut placed = Vec::new();
|
||||
while let Some(run) = runs.next() {
|
||||
if !is_system_role_message(&run[0]) {
|
||||
placed.extend(run);
|
||||
continue;
|
||||
}
|
||||
let follower_run = runs.next().unwrap_or_default();
|
||||
placed.extend(system_run_placed_after_tool_results(run, follower_run));
|
||||
}
|
||||
placed
|
||||
}
|
||||
|
||||
pub fn convert_mid_conversation_system_turns(messages: Vec<Message>) -> Vec<Message> {
|
||||
system_turns_after_tool_results(messages)
|
||||
.into_iter()
|
||||
.map(|message| {
|
||||
if is_system_role_message(&message) {
|
||||
system_role_message_as_user(message)
|
||||
} else {
|
||||
message
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn messages(value: Value) -> Vec<Message> {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
fn roles(messages: &[Message]) -> Vec<&str> {
|
||||
messages
|
||||
.iter()
|
||||
.map(|message| message.role.as_str())
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_system_turns(
|
||||
json!([{"role": "user", "content": "a"}, {"role": "assistant", "content": "b"}]),
|
||||
vec!["user", "assistant"],
|
||||
)]
|
||||
#[case::a_turn_between_messages_becomes_a_user_turn_in_place(
|
||||
json!([
|
||||
{"role": "user", "content": "a"},
|
||||
{"role": "assistant", "content": "b"},
|
||||
{"role": "system", "content": "reminder"},
|
||||
{"role": "user", "content": "c"}
|
||||
]),
|
||||
vec!["user", "assistant", "user", "user"],
|
||||
)]
|
||||
#[case::a_run_between_a_tool_use_and_its_result_moves_after_the_result(
|
||||
json!([
|
||||
{"role": "user", "content": "a"},
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "t", "name": "n", "input": {}}]},
|
||||
{"role": "system", "content": "reminder"},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t", "content": "ok"}, {"type": "text", "text": "next"}]},
|
||||
{"role": "user", "content": "d"}
|
||||
]),
|
||||
vec!["user", "assistant", "user", "user", "user"],
|
||||
)]
|
||||
fn turns_are_converted_in_place_and_never_split_a_tool_call_from_its_result(
|
||||
#[case] input: Value,
|
||||
#[case] expected_roles: Vec<&str>,
|
||||
) {
|
||||
let converted = convert_mid_conversation_system_turns(messages(input));
|
||||
|
||||
assert_eq!(roles(&converted), expected_roles);
|
||||
assert!(
|
||||
converted
|
||||
.iter()
|
||||
.all(|message| !is_system_role_message(message))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn a_converted_turn_carries_the_note_then_the_original_content_only() {
|
||||
let converted = convert_mid_conversation_system_turns(messages(json!([
|
||||
{"role": "user", "content": "a"},
|
||||
{"role": "system", "content": "reminder", "name": "ignored"}
|
||||
])));
|
||||
|
||||
assert_eq!(
|
||||
serde_json::to_value(&converted[1]).unwrap(),
|
||||
json!({"role": "user", "content": [
|
||||
{"type": "text", "text": CONVERTED_SYSTEM_NOTE},
|
||||
{"type": "text", "text": "reminder"}
|
||||
]})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn the_tool_result_turn_keeps_its_place_before_the_moved_run() {
|
||||
let converted = convert_mid_conversation_system_turns(messages(json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "t", "name": "n", "input": {}}]},
|
||||
{"role": "system", "content": "reminder"},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t", "content": "ok"}]}
|
||||
])));
|
||||
|
||||
assert!(opens_with_tool_results(&converted[1]));
|
||||
assert_eq!(
|
||||
serde_json::to_value(&converted[2]).unwrap()["content"][1],
|
||||
json!({"type": "text", "text": "reminder"})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
pub mod context;
|
||||
pub mod mid_conversation_system;
|
||||
pub mod normalization;
|
||||
pub mod streaming;
|
||||
pub mod transformation;
|
||||
|
|
|
|||
|
|
@ -1,47 +1,103 @@
|
|||
use litellm_llms_types::formats::messages::{
|
||||
ContentBlock, Message, MessageContent, MessagesOptionalParams, MessagesRequest, SystemPrompt,
|
||||
CacheControl, ContentBlock, Message, MessageContent, MessagesOptionalParams, MessagesRequest,
|
||||
SystemPrompt,
|
||||
};
|
||||
|
||||
const SYSTEM_ROLE: &str = "system";
|
||||
use crate::{
|
||||
anthropic::common_utils::filter_billing_headers_from_system,
|
||||
base_llm::messages::mid_conversation_system::{
|
||||
as_system_content_blocks, convert_mid_conversation_system_turns, is_system_role_message,
|
||||
message_content_blocks,
|
||||
},
|
||||
};
|
||||
|
||||
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
|
||||
match content {
|
||||
MessageContent::Text(text) => vec![ContentBlock::text(text)],
|
||||
MessageContent::Blocks(blocks) => blocks,
|
||||
}
|
||||
}
|
||||
|
||||
fn system_into_blocks(system: Option<SystemPrompt>) -> Vec<ContentBlock> {
|
||||
match system {
|
||||
None => Vec::new(),
|
||||
Some(SystemPrompt::Text(text)) => vec![ContentBlock::text(text)],
|
||||
Some(SystemPrompt::Blocks(blocks)) => blocks,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn fold_system_role_messages(request: MessagesRequest) -> MessagesRequest {
|
||||
if !request.messages.iter().any(|msg| msg.role == SYSTEM_ROLE) {
|
||||
return request;
|
||||
}
|
||||
|
||||
let (system_messages, chat_messages): (Vec<Message>, Vec<Message>) = request
|
||||
/// Python's `_normalize_system_role_messages`: a leading run of `role: "system"` entries is
|
||||
/// hoisted into the top-level `system`, a later one stays in place when the model accepts it
|
||||
/// and otherwise becomes a user turn where it was, and billing blocks are stripped from the
|
||||
/// top-level `system` either way.
|
||||
pub fn normalize_system_role_messages(
|
||||
request: MessagesRequest,
|
||||
supports_mid_conversation_system: bool,
|
||||
) -> MessagesRequest {
|
||||
let leading_count = request
|
||||
.messages
|
||||
.into_iter()
|
||||
.partition(|msg| msg.role == SYSTEM_ROLE);
|
||||
|
||||
let folded_system: Vec<ContentBlock> = system_into_blocks(request.params.system)
|
||||
.into_iter()
|
||||
.chain(
|
||||
system_messages
|
||||
.iter()
|
||||
.take_while(|message| is_system_role_message(message))
|
||||
.count();
|
||||
let mut messages = request.messages.into_iter();
|
||||
let hoisted: Vec<Message> = messages.by_ref().take(leading_count).collect();
|
||||
let remaining: Vec<Message> = messages.collect();
|
||||
let remaining = if supports_mid_conversation_system {
|
||||
remaining
|
||||
} else {
|
||||
convert_mid_conversation_system_turns(remaining)
|
||||
};
|
||||
let system = if hoisted.is_empty() {
|
||||
request.params.system
|
||||
} else {
|
||||
Some(SystemPrompt::Blocks(
|
||||
as_system_content_blocks(request.params.system)
|
||||
.into_iter()
|
||||
.flat_map(|msg| content_into_blocks(msg.content)),
|
||||
)
|
||||
.collect();
|
||||
|
||||
.chain(
|
||||
hoisted
|
||||
.into_iter()
|
||||
.flat_map(|message| message_content_blocks(message.content)),
|
||||
)
|
||||
.collect(),
|
||||
))
|
||||
};
|
||||
MessagesRequest {
|
||||
messages: chat_messages,
|
||||
messages: remaining,
|
||||
params: MessagesOptionalParams {
|
||||
system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)),
|
||||
system: system.and_then(filter_billing_headers_from_system),
|
||||
..request.params
|
||||
},
|
||||
..request
|
||||
}
|
||||
}
|
||||
|
||||
fn strip_scope_from_block(block: ContentBlock) -> ContentBlock {
|
||||
ContentBlock {
|
||||
cache_control: block.cache_control.map(|cache_control| CacheControl {
|
||||
scope: None,
|
||||
..cache_control
|
||||
}),
|
||||
..block
|
||||
}
|
||||
}
|
||||
|
||||
fn strip_scope_from_system(system: SystemPrompt) -> SystemPrompt {
|
||||
match system {
|
||||
SystemPrompt::Blocks(blocks) => {
|
||||
SystemPrompt::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
|
||||
}
|
||||
text => text,
|
||||
}
|
||||
}
|
||||
|
||||
fn strip_scope_from_message(message: Message) -> Message {
|
||||
Message {
|
||||
content: match message.content {
|
||||
MessageContent::Blocks(blocks) => {
|
||||
MessageContent::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
|
||||
}
|
||||
text => text,
|
||||
},
|
||||
..message
|
||||
}
|
||||
}
|
||||
|
||||
/// Python's `_remove_scope_from_cache_control`: hosts other than the first-party API reject
|
||||
/// `cache_control.scope`, so it is dropped from the system prompt and every message block.
|
||||
pub fn strip_cache_control_scope(request: MessagesRequest) -> MessagesRequest {
|
||||
MessagesRequest {
|
||||
messages: request
|
||||
.messages
|
||||
.into_iter()
|
||||
.map(strip_scope_from_message)
|
||||
.collect(),
|
||||
params: MessagesOptionalParams {
|
||||
system: request.params.system.map(strip_scope_from_system),
|
||||
..request.params
|
||||
},
|
||||
..request
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
use litellm_llms_types::formats::messages::{MessagesRequest, MessagesResponse};
|
||||
use litellm_router_types::LitellmParams;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::context::MessagesTransformContext;
|
||||
|
||||
|
|
@ -20,18 +22,11 @@ pub trait BaseMessagesConfig: Sync {
|
|||
&self,
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
litellm_params: &LitellmParams,
|
||||
stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error>;
|
||||
|
||||
fn complete_stream_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
self.get_complete_url(api_base, model, env_lookup)
|
||||
}
|
||||
|
||||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: MessagesRequest,
|
||||
|
|
@ -58,6 +53,7 @@ pub trait BaseMessagesConfig: Sync {
|
|||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
model: &str,
|
||||
litellm_params: &LitellmParams,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error>;
|
||||
|
||||
|
|
@ -75,6 +71,12 @@ pub trait BaseMessagesConfig: Sync {
|
|||
fn request_headers(&self, headers: Headers, _request: &MessagesRequest) -> Headers {
|
||||
headers
|
||||
}
|
||||
|
||||
/// The JSON that goes on the wire, for a host whose body differs from the typed request:
|
||||
/// Python's configs `pop("model")` when the model is addressed by the URL.
|
||||
fn wire_body(&self, body: Value) -> Value {
|
||||
body
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -94,6 +96,8 @@ mod tests {
|
|||
&self,
|
||||
_api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_litellm_params: &LitellmParams,
|
||||
_stream: bool,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(String::new())
|
||||
|
|
@ -104,6 +108,7 @@ mod tests {
|
|||
headers: Headers,
|
||||
_api_key: Option<&str>,
|
||||
_model: &str,
|
||||
_litellm_params: &LitellmParams,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
Ok(ValidatedEnvironment {
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ use litellm_llms_types::formats::messages::{
|
|||
MessagesRequest,
|
||||
streaming::{MessagesStreamEvent, MessagesStreamUsage},
|
||||
};
|
||||
use litellm_router_types::LitellmParams;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
|
|
@ -74,16 +75,26 @@ fn bearer_token(
|
|||
fn invoke_url(
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
aws: &AwsParams,
|
||||
stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
path: &str,
|
||||
) -> String {
|
||||
let (model_id, model_region) =
|
||||
bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model));
|
||||
let region = resolve_bedrock_region(model_region.as_deref(), &AwsParams::default(), env_lookup);
|
||||
let endpoint = api_base
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string)
|
||||
let region = resolve_bedrock_region(model_region.as_deref(), aws, env_lookup);
|
||||
let path = if stream {
|
||||
INVOKE_STREAM_PATH
|
||||
} else {
|
||||
INVOKE_PATH
|
||||
};
|
||||
let configured = |value: Option<&str>| {
|
||||
value
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string)
|
||||
};
|
||||
let endpoint = configured(api_base)
|
||||
.or_else(|| configured(aws.aws_bedrock_runtime_endpoint.as_deref()))
|
||||
.or_else(|| env_lookup(AWS_BEDROCK_RUNTIME_ENDPOINT))
|
||||
.unwrap_or_else(|| BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", ®ion));
|
||||
format!("{}/model/{model_id}/{path}", endpoint.trim_end_matches('/'))
|
||||
|
|
@ -102,18 +113,17 @@ impl BaseMessagesConfig for AmazonAnthropicClaudeMessagesConfig {
|
|||
&self,
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
litellm_params: &LitellmParams,
|
||||
stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(invoke_url(api_base, model, env_lookup, INVOKE_PATH))
|
||||
}
|
||||
|
||||
fn complete_stream_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(invoke_url(api_base, model, env_lookup, INVOKE_STREAM_PATH))
|
||||
Ok(invoke_url(
|
||||
api_base,
|
||||
model,
|
||||
&litellm_params.aws,
|
||||
stream,
|
||||
env_lookup,
|
||||
))
|
||||
}
|
||||
|
||||
fn transform_anthropic_messages_request(
|
||||
|
|
@ -137,6 +147,7 @@ impl BaseMessagesConfig for AmazonAnthropicClaudeMessagesConfig {
|
|||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
model: &str,
|
||||
litellm_params: &LitellmParams,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
if let Some(token) = bearer_token(api_key, env_lookup) {
|
||||
|
|
@ -150,13 +161,13 @@ impl BaseMessagesConfig for AmazonAnthropicClaudeMessagesConfig {
|
|||
}
|
||||
let (_, model_region) =
|
||||
bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model));
|
||||
let params = AwsParams::default();
|
||||
let params = &litellm_params.aws;
|
||||
Ok(ValidatedEnvironment {
|
||||
headers,
|
||||
auth: AuthScheme::AwsSigV4 {
|
||||
region: resolve_bedrock_region(model_region.as_deref(), ¶ms, env_lookup),
|
||||
region: resolve_bedrock_region(model_region.as_deref(), params, env_lookup),
|
||||
service: BEDROCK_SERVICE,
|
||||
credentials: Box::new(AwsCredentialSource::from_params(¶ms, env_lookup)),
|
||||
credentials: Box::new(AwsCredentialSource::from_params(params, env_lookup)),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
|
@ -503,19 +514,129 @@ mod tests {
|
|||
assert_eq!(sse, expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_uses_the_streaming_url_only_for_streams() {
|
||||
let env = |_: &str| -> Option<String> { None };
|
||||
let config = AmazonAnthropicClaudeMessagesConfig;
|
||||
fn in_region(region: &str) -> LitellmParams {
|
||||
LitellmParams {
|
||||
aws: AwsParams {
|
||||
aws_region_name: Some(region.into()),
|
||||
..AwsParams::default()
|
||||
},
|
||||
..LitellmParams::default()
|
||||
}
|
||||
}
|
||||
|
||||
fn with_runtime_endpoint(endpoint: &str) -> LitellmParams {
|
||||
LitellmParams {
|
||||
aws: AwsParams {
|
||||
aws_bedrock_runtime_endpoint: Some(endpoint.into()),
|
||||
..in_region("us-west-2").aws
|
||||
},
|
||||
..LitellmParams::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::invoke_in_the_params_region(
|
||||
None,
|
||||
in_region("us-west-2"),
|
||||
false,
|
||||
None,
|
||||
"https://bedrock-runtime.us-west-2.amazonaws.com/model/anthropic.claude-3/invoke"
|
||||
)]
|
||||
#[case::stream_in_the_params_region(
|
||||
None,
|
||||
in_region("us-west-2"),
|
||||
true,
|
||||
None,
|
||||
"https://bedrock-runtime.us-west-2.amazonaws.com/model/anthropic.claude-3/invoke-with-response-stream"
|
||||
)]
|
||||
#[case::api_base_outranks_the_params_endpoint(
|
||||
Some("https://base.test/"),
|
||||
with_runtime_endpoint("https://params.test"),
|
||||
false,
|
||||
None,
|
||||
"https://base.test/model/anthropic.claude-3/invoke"
|
||||
)]
|
||||
#[case::a_blank_api_base_does_not_hide_the_params_endpoint(
|
||||
Some(" "),
|
||||
with_runtime_endpoint("https://params.test"),
|
||||
false,
|
||||
Some("https://env.test"),
|
||||
"https://params.test/model/anthropic.claude-3/invoke"
|
||||
)]
|
||||
#[case::params_endpoint_outranks_the_environment(
|
||||
None,
|
||||
with_runtime_endpoint("https://params.test"),
|
||||
false,
|
||||
Some("https://env.test"),
|
||||
"https://params.test/model/anthropic.claude-3/invoke"
|
||||
)]
|
||||
#[case::environment_outranks_the_region_template(
|
||||
None,
|
||||
in_region("us-west-2"),
|
||||
false,
|
||||
Some("https://env.test"),
|
||||
"https://env.test/model/anthropic.claude-3/invoke"
|
||||
)]
|
||||
fn url_follows_python_endpoint_precedence_and_the_stream_path(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] litellm_params: LitellmParams,
|
||||
#[case] stream: bool,
|
||||
#[case] env_endpoint: Option<&str>,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
let env = |name: &str| {
|
||||
(name == AWS_BEDROCK_RUNTIME_ENDPOINT)
|
||||
.then_some(env_endpoint)
|
||||
.flatten()
|
||||
.map(str::to_string)
|
||||
};
|
||||
assert_eq!(
|
||||
config
|
||||
.get_complete_url(None, "anthropic.claude-3", &env)
|
||||
AmazonAnthropicClaudeMessagesConfig
|
||||
.get_complete_url(
|
||||
api_base,
|
||||
"anthropic.claude-3",
|
||||
&litellm_params,
|
||||
stream,
|
||||
&env
|
||||
)
|
||||
.unwrap(),
|
||||
config
|
||||
.complete_stream_url(None, "anthropic.claude-3", &env)
|
||||
.unwrap()
|
||||
.replace(INVOKE_STREAM_PATH, INVOKE_PATH)
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sigv4_scope_and_credentials_come_from_the_litellm_params() {
|
||||
let litellm_params = LitellmParams {
|
||||
aws: AwsParams {
|
||||
aws_access_key_id: Some("AKIAPARAMS".into()),
|
||||
aws_secret_access_key: Some("params-secret".into()),
|
||||
..in_region("eu-central-1").aws
|
||||
},
|
||||
..LitellmParams::default()
|
||||
};
|
||||
let validated = AmazonAnthropicClaudeMessagesConfig
|
||||
.validate_environment(
|
||||
Vec::new(),
|
||||
None,
|
||||
"anthropic.claude-3",
|
||||
&litellm_params,
|
||||
&|_| None,
|
||||
)
|
||||
.unwrap();
|
||||
let AuthScheme::AwsSigV4 {
|
||||
region,
|
||||
credentials,
|
||||
..
|
||||
} = validated.auth
|
||||
else {
|
||||
panic!("expected SigV4, got {:?}", validated.auth);
|
||||
};
|
||||
let AwsCredentialSource::HostSupplied(credentials) = *credentials else {
|
||||
panic!("expected the params' static keys, got {credentials:?}");
|
||||
};
|
||||
assert_eq!(
|
||||
(region.as_str(), credentials.access_key_id()),
|
||||
("eu-central-1", "AKIAPARAMS")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -538,6 +659,7 @@ mod tests {
|
|||
vec![("authorization".into(), "Bearer forwarded".into())],
|
||||
api_key,
|
||||
"anthropic.claude-3",
|
||||
&LitellmParams::default(),
|
||||
&env,
|
||||
)
|
||||
.unwrap();
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_llms::base_llm::messages::normalization::fold_system_role_messages;
|
||||
use litellm_llms::base_llm::messages::normalization::normalize_system_role_messages;
|
||||
use litellm_llms_types::formats::messages::MessagesRequest;
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
|
@ -7,42 +7,47 @@ use serde_json::{Value, json};
|
|||
#[case::text_system(json!("existing"), json!([{"type": "text", "text": "existing"}]))]
|
||||
#[case::block_system(json!([{ "type": "future", "payload": 7 }]), json!([{ "type": "future", "payload": 7 }]))]
|
||||
#[case::no_system(Value::Null, json!([]))]
|
||||
fn folding_preserves_block_fields_order_and_unrelated_request_fields(
|
||||
fn hoisting_preserves_block_fields_order_and_unrelated_request_fields(
|
||||
#[case] system: Value,
|
||||
#[case] initial_blocks: Value,
|
||||
) {
|
||||
let cache_control = json!({"type": "ephemeral", "scope": "global", "future": true});
|
||||
let folded_block = json!({"type": "text", "text": "second", "cache_control": cache_control});
|
||||
let leading_block = json!({"type": "text", "text": "second", "cache_control": cache_control});
|
||||
let user = json!({"role": "user", "content": "hello", "future_message": 42});
|
||||
let later_turn = json!({"role": "system", "content": "mid-conversation reminder"});
|
||||
let request: MessagesRequest = serde_json::from_value(json!({
|
||||
"model": "test-model",
|
||||
"max_tokens": 64,
|
||||
"system": system,
|
||||
"messages": [
|
||||
{"role": "system", "content": "first"},
|
||||
{"role": "system", "content": [leading_block]},
|
||||
user,
|
||||
{"role": "system", "content": [folded_block]}
|
||||
later_turn
|
||||
],
|
||||
"future_request": {"nested": true}
|
||||
}))
|
||||
.unwrap();
|
||||
let folded = fold_system_role_messages(request);
|
||||
let normalized = normalize_system_role_messages(request, true);
|
||||
let expected_blocks: Vec<Value> = initial_blocks
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.cloned()
|
||||
.chain([json!({"type": "text", "text": "first"}), folded_block])
|
||||
.chain([json!({"type": "text", "text": "first"}), leading_block])
|
||||
.collect();
|
||||
assert_eq!(
|
||||
serde_json::to_value(&folded).unwrap(),
|
||||
serde_json::to_value(&normalized).unwrap(),
|
||||
json!({
|
||||
"model": "test-model",
|
||||
"max_tokens": 64,
|
||||
"system": expected_blocks,
|
||||
"messages": [user],
|
||||
"messages": [user, later_turn],
|
||||
"future_request": {"nested": true}
|
||||
})
|
||||
);
|
||||
assert_eq!(fold_system_role_messages(folded.clone()), folded);
|
||||
assert_eq!(
|
||||
normalize_system_role_messages(normalized.clone(), true),
|
||||
normalized
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ litellm-llms.workspace = true
|
|||
litellm-secrets = { workspace = true, features = ["aws", "azure", "google", "hashicorp", "cyberark"] }
|
||||
litellm-secrets-types.workspace = true
|
||||
litellm-llms-types.workspace = true
|
||||
litellm-router-types.workspace = true
|
||||
litellm-host-python.workspace = true
|
||||
litellm-token-counter = { path = "../token-counter", default-features = false }
|
||||
pyo3.workspace = true
|
||||
|
|
|
|||
|
|
@ -60,7 +60,7 @@ impl PythonSettings {
|
|||
}
|
||||
}
|
||||
|
||||
fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult<bool> {
|
||||
pub(crate) fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult<bool> {
|
||||
if !error.is_instance_of::<PyModuleNotFoundError>(py) {
|
||||
return Ok(false);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,11 +5,12 @@ use bytes::Bytes;
|
|||
use litellm_host_python::{InvokeError, PythonBinding, from_py, present, to_py};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_inference_messages::{
|
||||
Error, MessagesCall, MessagesSettings, MessagesShaping, messages_body,
|
||||
Error, MessagesCall, MessagesSettings, MessagesShaping, litellm_params, messages_body,
|
||||
route::{Messages, MessagesStreamHead},
|
||||
};
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
|
||||
use litellm_llms_types::headers::ProviderSpecificHeaders;
|
||||
use litellm_router_types::LitellmParams;
|
||||
use pyo3::{
|
||||
exceptions::{PyException, PyValueError},
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
|
|
@ -21,6 +22,7 @@ use serde_json::{Map, Value};
|
|||
use crate::{
|
||||
errors::{RustUpstreamError, route_error_to_pyerr},
|
||||
marshal::{optional_timeout, project_optional_fields, public_response, python_timeout_seconds},
|
||||
python_settings::missing_module,
|
||||
};
|
||||
|
||||
const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host";
|
||||
|
|
@ -63,6 +65,33 @@ fn merge_headers(
|
|||
(!merged.is_empty()).then_some(merged)
|
||||
}
|
||||
|
||||
/// Reads `litellm.<name>`, the module global a param spec names, the way
|
||||
/// `VertexBase.get_vertex_ai_project` reads `litellm.vertex_project`.
|
||||
fn module_global<'py>(py: Python<'py>, name: &str) -> PyResult<Option<Bound<'py, PyAny>>> {
|
||||
let module = match py.import("litellm") {
|
||||
Ok(module) => module,
|
||||
Err(error) if missing_module(py, &error, "litellm")? => return Ok(None),
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
let value = module.getattr(name)?;
|
||||
Ok((!value.is_none()).then_some(value))
|
||||
}
|
||||
|
||||
/// The caller's litellm params, read from the kwargs by the names the typed params declare,
|
||||
/// so a key the configs do not read is never converted; a spec whose spellings the call all
|
||||
/// leaves out falls back to the module global it names, which stays the host's concern.
|
||||
fn project_litellm_params<'py>(
|
||||
argument: impl Fn(&str) -> PyResult<Option<Bound<'py, PyAny>>>,
|
||||
global: impl Fn(&str) -> PyResult<Option<Bound<'py, PyAny>>>,
|
||||
) -> PyResult<Result<LitellmParams, Error>> {
|
||||
let fields = project_optional_fields(LitellmParams::fields(), &argument)?;
|
||||
let globals = LitellmParams::specs()
|
||||
.filter(|spec| !spec.wire.iter().any(|name| fields.contains_key(*name)))
|
||||
.filter_map(|spec| spec.module_global);
|
||||
let folded = project_optional_fields(globals, &global)?;
|
||||
Ok(litellm_params(fields.into_iter().chain(folded).collect()))
|
||||
}
|
||||
|
||||
fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
||||
match error {
|
||||
Error::Transport(TransportError::Http { status, body }) => {
|
||||
|
|
@ -132,15 +161,19 @@ impl MessagesPythonHost {
|
|||
let api_base = string("api_base")?;
|
||||
let extra_headers = self.merged_headers(py, arguments)?;
|
||||
let provider_specific_header = self.provider_specific_header(py, arguments)?;
|
||||
Ok(messages_body(body).map(|body| MessagesCall {
|
||||
body,
|
||||
api_key,
|
||||
api_base,
|
||||
extra_headers,
|
||||
provider_specific_header,
|
||||
custom_llm_provider,
|
||||
timeout: optional_timeout(timeout),
|
||||
shaping,
|
||||
let litellm_params = project_litellm_params(argument, |name| module_global(py, name))?;
|
||||
Ok(messages_body(body).and_then(|body| {
|
||||
Ok(MessagesCall {
|
||||
body,
|
||||
api_key,
|
||||
api_base,
|
||||
extra_headers,
|
||||
provider_specific_header,
|
||||
custom_llm_provider,
|
||||
litellm_params: litellm_params?,
|
||||
timeout: optional_timeout(timeout),
|
||||
shaping,
|
||||
})
|
||||
}))
|
||||
}
|
||||
|
||||
|
|
@ -351,6 +384,106 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::model_and_credentials_are_projected(
|
||||
json!({"model": "m", "api_key": "k", "api_base": "b", "api_version": "v", "timeout": 5, "messages": []}),
|
||||
Ok(json!({"model": "m", "api_key": "k", "api_base": "b", "api_version": "v"})),
|
||||
)]
|
||||
#[case::aws_keys_are_projected(
|
||||
json!({"model": "m", "aws_region_name": "eu-central-1", "aws_access_key_id": "AKIA", "messages": [{"role": "user"}]}),
|
||||
Ok(json!({"model": "m", "aws_region_name": "eu-central-1", "aws_access_key_id": "AKIA"})),
|
||||
)]
|
||||
#[case::vertex_keys_are_projected_in_both_spellings(
|
||||
json!({"model": "m", "vertex_project": "p", "vertex_ai_location": "us-east5", "vertex_credentials": {"type": "service_account"}}),
|
||||
Ok(json!({"model": "m", "vertex_project": "p", "vertex_ai_location": "us-east5", "vertex_credentials": {"type": "service_account"}})),
|
||||
)]
|
||||
#[case::an_explicit_none_is_absent(json!({"model": "m", "aws_region_name": null}), Ok(json!({"model": "m"})))]
|
||||
#[case::a_wrong_type_is_a_request_error(json!({"model": "m", "aws_region_name": 7}), Err(()))]
|
||||
#[case::a_missing_model_is_a_request_error(json!({"aws_region_name": "eu-central-1"}), Err(()))]
|
||||
fn litellm_params_are_projected_by_their_declared_names(
|
||||
#[case] kwargs: Value,
|
||||
#[case] expected: Result<Value, ()>,
|
||||
) {
|
||||
assert_projection(kwargs, json!({}), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::global_fills_a_missing_vertex_project(
|
||||
json!({"model": "m"}),
|
||||
json!({"vertex_project": "from-global", "vertex_location": "us-east5"}),
|
||||
Ok(json!({"model": "m", "vertex_project": "from-global", "vertex_location": "us-east5"})),
|
||||
)]
|
||||
#[case::the_call_wins_over_the_global(
|
||||
json!({"model": "m", "vertex_project": "from-call"}),
|
||||
json!({"vertex_project": "from-global"}),
|
||||
Ok(json!({"model": "m", "vertex_project": "from-call"})),
|
||||
)]
|
||||
#[case::an_explicit_none_in_the_call_still_falls_back(
|
||||
json!({"model": "m", "vertex_project": null}),
|
||||
json!({"vertex_project": "from-global"}),
|
||||
Ok(json!({"model": "m", "vertex_project": "from-global"})),
|
||||
)]
|
||||
#[case::a_global_of_the_wrong_type_is_a_request_error(
|
||||
json!({"model": "m"}),
|
||||
json!({"vertex_location": 5}),
|
||||
Err(()),
|
||||
)]
|
||||
fn module_globals_fill_the_litellm_params_the_call_leaves_out(
|
||||
#[case] kwargs: Value,
|
||||
#[case] globals: Value,
|
||||
#[case] expected: Result<Value, ()>,
|
||||
) {
|
||||
assert_projection(kwargs, globals, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::a_global_without_a_spec_is_never_read(
|
||||
json!({"model": "m"}),
|
||||
json!({"aws_region_name": "eu-central-1", "api_key": "k"}),
|
||||
Ok(json!({"model": "m"})),
|
||||
)]
|
||||
#[case::only_the_spec_named_global_is_read(
|
||||
json!({"model": "m"}),
|
||||
json!({"vertex_ai_project": "legacy-global", "vertex_project": "from-global"}),
|
||||
Ok(json!({"model": "m", "vertex_project": "from-global"})),
|
||||
)]
|
||||
fn only_the_globals_the_specs_name_are_consulted(
|
||||
#[case] kwargs: Value,
|
||||
#[case] globals: Value,
|
||||
#[case] expected: Result<Value, ()>,
|
||||
) {
|
||||
assert_projection(kwargs, globals, expected);
|
||||
}
|
||||
|
||||
fn assert_projection(kwargs: Value, globals: Value, expected: Result<Value, ()>) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let dict = |value: &Value| {
|
||||
to_py(py, value)
|
||||
.unwrap()
|
||||
.into_bound(py)
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap()
|
||||
};
|
||||
let kwargs = dict(&kwargs);
|
||||
let globals = dict(&globals);
|
||||
let bound = PyDict::new(py);
|
||||
let projected = project_litellm_params(
|
||||
|name| present(&kwargs, &bound, name),
|
||||
|name| present(&globals, &bound, name),
|
||||
)
|
||||
.unwrap();
|
||||
match (projected, expected) {
|
||||
(Ok(projected), Ok(fields)) => assert_eq!(
|
||||
projected,
|
||||
litellm_params(serde_json::from_value(fields).unwrap()).unwrap()
|
||||
),
|
||||
(Err(error), Err(())) => assert!(matches!(error, Error::InvalidRequest(_))),
|
||||
(projected, expected) => panic!("got {projected:?}, expected {expected:?}"),
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)]
|
||||
#[case::missing_field(Error::MissingField("max_tokens"), true)]
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ repository.workspace = true
|
|||
litellm-config.workspace = true
|
||||
litellm-inference.workspace = true
|
||||
litellm-inference-messages.workspace = true
|
||||
litellm-router-types.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_inference_messages::MessagesShaping;
|
||||
use litellm_router_types::LitellmParams;
|
||||
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub struct Deployment {
|
||||
|
|
@ -8,6 +9,7 @@ pub struct Deployment {
|
|||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub litellm_params: LitellmParams,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ impl Router {
|
|||
.map(|value| value.expose().to_string()),
|
||||
api_base: model.litellm_params.api_base.clone(),
|
||||
custom_llm_provider: model.litellm_params.custom_llm_provider.clone(),
|
||||
litellm_params: model.litellm_params.clone(),
|
||||
..Deployment::default()
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -29,6 +29,32 @@ fn configuration_preserves_deployment_parameters(#[case] parameters: &str) {
|
|||
assert_eq!(deployment.custom_llm_provider, params.custom_llm_provider);
|
||||
assert_eq!(deployment.timeout, Deployment::default().timeout);
|
||||
assert_eq!(deployment.shaping, Deployment::default().shaping);
|
||||
assert_eq!(deployment.litellm_params, *params);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn configuration_types_the_provider_params_of_a_deployment() {
|
||||
let config = Config::from_yaml(
|
||||
"model_list:
|
||||
- model_name: public-model
|
||||
litellm_params:
|
||||
model: bedrock/anthropic.claude-3
|
||||
aws_region_name: eu-central-1
|
||||
aws_bedrock_runtime_endpoint: https://runtime.example
|
||||
azure_ad_token: ignored-by-every-group",
|
||||
)
|
||||
.unwrap();
|
||||
let router = Router::from_model_list(&config.model_list);
|
||||
let params = &router.get("public-model").unwrap().litellm_params;
|
||||
|
||||
assert_eq!(
|
||||
(
|
||||
params.aws.aws_region_name.as_deref(),
|
||||
params.aws.aws_bedrock_runtime_endpoint.as_deref(),
|
||||
params.aws.aws_access_key_id.as_deref(),
|
||||
),
|
||||
(Some("eu-central-1"), Some("https://runtime.example"), None)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -72,6 +98,7 @@ fn programmatic_deployments_preserve_overrides_and_last_entry_wins() {
|
|||
api_key: Some("test-key".into()),
|
||||
api_base: Some("https://provider.example/v1".into()),
|
||||
custom_llm_provider: Some("test-provider".into()),
|
||||
litellm_params: Default::default(),
|
||||
timeout: Some(Duration::from_secs(7)),
|
||||
shaping: MessagesShaping {
|
||||
settings: MessagesSettings {
|
||||
|
|
|
|||
|
|
@ -31,5 +31,6 @@ def anthropic_model_capabilities(model: str, custom_llm_provider: str | None) ->
|
|||
"supports_output_config": supports("supports_output_config"),
|
||||
"supports_sampling_params": AnthropicModelInfo._supports_sampling_params(resolved_model), # pyright: ignore[reportPrivateUsage] # same gate the handler applies
|
||||
"supports_speed": AnthropicConfig._model_supports_speed_param(resolved_model, provider), # pyright: ignore[reportPrivateUsage] # same gate the handler applies
|
||||
"supports_mid_conversation_system": supports("supports_mid_conversation_system"),
|
||||
"effort_tiers": {level: tier(level) for level in ("minimal", "low", "medium", "high", "xhigh", "max")},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -32,12 +32,14 @@ def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeyp
|
|||
supports_output_config=True,
|
||||
supports_xhigh_reasoning_effort=True,
|
||||
supports_sampling_params=False,
|
||||
supports_mid_conversation_system=True,
|
||||
)
|
||||
|
||||
capabilities: Final = anthropic_model_capabilities("anthropic/claude-test-adaptive", None)
|
||||
|
||||
assert capabilities["supports_adaptive_thinking"]
|
||||
assert capabilities["supports_output_config"]
|
||||
assert capabilities["supports_mid_conversation_system"]
|
||||
assert not capabilities["supports_legacy_thinking"]
|
||||
assert not capabilities["supports_sampling_params"]
|
||||
assert capabilities["effort_tiers"] == {
|
||||
|
|
@ -56,6 +58,7 @@ def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> Non
|
|||
assert capabilities["supports_sampling_params"]
|
||||
assert not capabilities["supports_reasoning"]
|
||||
assert not capabilities["supports_adaptive_thinking"]
|
||||
assert not capabilities["supports_mid_conversation_system"]
|
||||
assert capabilities["effort_tiers"] == dict.fromkeys(("minimal", "low", "medium", "high", "xhigh", "max"), False)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue