mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
merge: main into litellm_lit8439_invoke_native_extensions
This commit is contained in:
commit
b5ec6be015
114 changed files with 11232 additions and 516 deletions
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -108,6 +108,7 @@ jobs:
|
|||
tests/test_litellm/completion_extras
|
||||
tests/test_litellm/containers
|
||||
tests/test_litellm/endpoints
|
||||
tests/test_litellm/files
|
||||
tests/test_litellm/images
|
||||
tests/test_litellm/interactions
|
||||
tests/test_litellm/messages
|
||||
|
|
|
|||
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -3414,6 +3414,7 @@ dependencies = [
|
|||
name = "litellm-types"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
]
|
||||
|
|
|
|||
274
litellm-rust/crates/core-utils/src/dot_notation_indexing.rs
Normal file
274
litellm-rust/crates/core-utils/src/dot_notation_indexing.rs
Normal file
|
|
@ -0,0 +1,274 @@
|
|||
use serde_json::Value;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
enum Segment {
|
||||
Field(String),
|
||||
Every,
|
||||
Index(usize),
|
||||
}
|
||||
|
||||
fn parse_segments(path: &str) -> Option<Vec<Segment>> {
|
||||
let mut segments = Vec::new();
|
||||
let mut rest = path;
|
||||
while !rest.is_empty() {
|
||||
if let Some(after_open) = rest.strip_prefix('[') {
|
||||
let (inside, after) = after_open.split_once(']')?;
|
||||
segments.push(match inside {
|
||||
"*" => Segment::Every,
|
||||
index => Segment::Index(index.trim().parse().ok()?),
|
||||
});
|
||||
rest = after.strip_prefix('.').unwrap_or(after);
|
||||
continue;
|
||||
}
|
||||
let end = rest.find(['.', '[']).unwrap_or(rest.len());
|
||||
let (field, after) = rest.split_at(end);
|
||||
if !field.is_empty() {
|
||||
segments.push(Segment::Field(field.to_string()));
|
||||
}
|
||||
rest = after.strip_prefix('.').unwrap_or(after);
|
||||
}
|
||||
Some(segments)
|
||||
}
|
||||
|
||||
fn without_path(value: Value, segments: &[Segment]) -> Value {
|
||||
let Some((segment, tail)) = segments.split_first() else {
|
||||
return value;
|
||||
};
|
||||
match (segment, value) {
|
||||
(Segment::Field(name), Value::Object(object)) => Value::Object(
|
||||
object
|
||||
.into_iter()
|
||||
.filter_map(|(key, item)| {
|
||||
if key != *name {
|
||||
return Some((key, item));
|
||||
}
|
||||
(!tail.is_empty()).then(|| (key, without_path(item, tail)))
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
(Segment::Every, Value::Array(items)) => Value::Array(
|
||||
items
|
||||
.into_iter()
|
||||
.map(|item| without_path(item, tail))
|
||||
.collect(),
|
||||
),
|
||||
(Segment::Index(index), Value::Array(items)) => Value::Array(
|
||||
items
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(position, item)| {
|
||||
if position == *index {
|
||||
without_path(item, tail)
|
||||
} else {
|
||||
item
|
||||
}
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
(_, value) => value,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn delete_nested_value(value: Value, path: &str) -> Value {
|
||||
match parse_segments(path) {
|
||||
Some(segments) => without_path(value, &segments),
|
||||
None => value,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[fixture]
|
||||
fn body() -> Value {
|
||||
json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
})
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::top_level_field("top", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}}
|
||||
}))]
|
||||
#[case::whole_object("meta", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::nested_field("meta.inner.drop", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::trailing_dot("meta.inner.drop.", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::leading_and_doubled_dots(".meta..inner.drop", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_in_every_element("tools[*].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::whole_array_field_in_every_element("tools[*].arr", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"]},
|
||||
{"name": "t1", "examples": ["b"]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_in_indexed_element("tools[1].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::padded_index("tools[ 1 ].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_right_after_bracket("tools[0]examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::index_then_wildcard("tools[0].arr[*].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::wildcard_then_index_only_where_it_exists("tools[*].arr[1].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::nested_wildcards("tools[*].arr[*].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
fn deletes_the_addressed_field(body: Value, #[case] path: &str, #[case] expected: Value) {
|
||||
assert_eq!(delete_nested_value(body, path), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty_path("")]
|
||||
#[case::missing_field("missing")]
|
||||
#[case::missing_parent("missing.field")]
|
||||
#[case::field_through_a_scalar("top.value")]
|
||||
#[case::field_on_an_array("tools.name")]
|
||||
#[case::index_on_an_object("meta[0].user")]
|
||||
#[case::wildcard_on_an_object("meta[*].user")]
|
||||
#[case::wildcard_over_scalars("tools[*].examples[*].name")]
|
||||
#[case::index_out_of_range("tools[5].name")]
|
||||
#[case::every_element_itself("tools[*]")]
|
||||
#[case::indexed_element_itself("tools[0]")]
|
||||
#[case::nested_element_itself("tools[*].arr[0]")]
|
||||
#[case::negative_index("tools[-1].name")]
|
||||
#[case::non_numeric_index("tools[x].name")]
|
||||
#[case::empty_index("tools[].name")]
|
||||
#[case::unclosed_bracket("top[0")]
|
||||
fn leaves_the_value_untouched(body: Value, #[case] path: &str) {
|
||||
assert_eq!(delete_nested_value(body.clone(), path), body);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::wildcards_indices_and_nesting(
|
||||
json!({"tools": [
|
||||
{"name": "t0", "configs": [{"id": "c0", "remove_me": 1, "keep": 1}, {"id": "c1", "remove_me": 2, "keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
|
||||
{"name": "t1", "configs": [{"id": "c0", "remove_me": 3, "keep": 3}, {"id": "c1", "remove_me": 4, "keep": 4}], "metadata": {"drop_this": 2, "preserve": 2}},
|
||||
{"name": "t2", "configs": [{"id": "c0", "remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
|
||||
]}),
|
||||
&["tools[*].configs[1].remove_me", "tools[1].metadata.drop_this", "tools[*].configs[*].id"],
|
||||
json!({"tools": [
|
||||
{"name": "t0", "configs": [{"remove_me": 1, "keep": 1}, {"keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
|
||||
{"name": "t1", "configs": [{"remove_me": 3, "keep": 3}, {"keep": 4}], "metadata": {"preserve": 2}},
|
||||
{"name": "t2", "configs": [{"remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
|
||||
]}),
|
||||
)]
|
||||
#[case::simple_and_wildcard_nesting(
|
||||
json!({
|
||||
"tools": [{"name": "t1", "simple_nested": {"remove": 1, "keep": 2}, "complex": [{"nested": {"remove": 3, "keep": 4}}]}],
|
||||
"top_level_remove": "should_go",
|
||||
"top_level_keep": "should_stay"
|
||||
}),
|
||||
&["tools[*].simple_nested.remove", "tools[*].complex[*].nested.remove"],
|
||||
json!({
|
||||
"tools": [{"name": "t1", "simple_nested": {"keep": 2}, "complex": [{"nested": {"keep": 4}}]}],
|
||||
"top_level_remove": "should_go",
|
||||
"top_level_keep": "should_stay"
|
||||
}),
|
||||
)]
|
||||
#[case::triple_nested_wildcards(
|
||||
json!({"tools": [{"name": "t1", "arr1": [
|
||||
{"arr2": [{"field": 1, "keep": 1}, {"field": 2, "keep": 2}]},
|
||||
{"arr2": [{"field": 3, "keep": 3}]}
|
||||
]}]}),
|
||||
&["tools[*].arr1[*].arr2[*].field"],
|
||||
json!({"tools": [{"name": "t1", "arr1": [
|
||||
{"arr2": [{"keep": 1}, {"keep": 2}]},
|
||||
{"arr2": [{"keep": 3}]}
|
||||
]}]}),
|
||||
)]
|
||||
fn applies_paths_in_sequence(
|
||||
#[case] value: Value,
|
||||
#[case] paths: &[&str],
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let deleted = paths
|
||||
.iter()
|
||||
.fold(value, |value, path| delete_nested_value(value, path));
|
||||
assert_eq!(deleted, expected);
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,93 @@
|
|||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub fn get_provider_specific_headers(
|
||||
provider_specific_header: Option<&ProviderSpecificHeaders>,
|
||||
custom_llm_provider: &str,
|
||||
) -> Map<String, Value> {
|
||||
let entries: &[ProviderSpecificHeader] = match provider_specific_header {
|
||||
None => &[],
|
||||
Some(ProviderSpecificHeaders::One(entry)) => std::slice::from_ref(entry),
|
||||
Some(ProviderSpecificHeaders::Many(entries)) => entries,
|
||||
};
|
||||
entries
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
entry
|
||||
.custom_llm_provider
|
||||
.split(',')
|
||||
.any(|scoped| scoped.trim() == custom_llm_provider)
|
||||
})
|
||||
.flat_map(|entry| entry.extra_headers.clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::single_entry_for_the_provider(
|
||||
json!({"custom_llm_provider": "anthropic", "extra_headers": {"Authorization": "Bearer t", "Custom-Header": "v"}}),
|
||||
json!({"Authorization": "Bearer t", "Custom-Header": "v"}),
|
||||
)]
|
||||
#[case::single_entry_for_another_provider(
|
||||
json!({"custom_llm_provider": "openai", "extra_headers": {"Authorization": "Bearer t"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::provider_in_a_comma_separated_scope(
|
||||
json!({"custom_llm_provider": "bedrock,anthropic,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}),
|
||||
json!({"anthropic-beta": "context-1m-2025-08-07"}),
|
||||
)]
|
||||
#[case::provider_missing_from_a_comma_separated_scope(
|
||||
json!({"custom_llm_provider": "bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::scope_with_spaces(
|
||||
json!({"custom_llm_provider": "bedrock, anthropic , vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({"anthropic-beta": "test"}),
|
||||
)]
|
||||
#[case::scope_names_must_match_exactly(
|
||||
json!({"custom_llm_provider": "anthropic_text", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::entries_scope_independently(
|
||||
json!([
|
||||
{"custom_llm_provider": "anthropic,bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}},
|
||||
{"custom_llm_provider": "bedrock", "extra_headers": {"x-bedrock-only": "no"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"authorization": "Bearer sk-ant-oat01-fake-token"}}
|
||||
]),
|
||||
json!({"anthropic-beta": "context-1m-2025-08-07", "authorization": "Bearer sk-ant-oat01-fake-token"}),
|
||||
)]
|
||||
#[case::later_entries_win(
|
||||
json!([
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "first"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "second"}}
|
||||
]),
|
||||
json!({"x-scoped": "second"}),
|
||||
)]
|
||||
#[case::empty_list(json!([]), json!({}))]
|
||||
#[case::entry_without_scope(json!({"extra_headers": {"x-scoped": "yes"}}), json!({}))]
|
||||
#[case::entry_without_headers(json!({"custom_llm_provider": "anthropic"}), json!({}))]
|
||||
fn provider_specific_headers_match_the_scoped_provider(
|
||||
#[case] configured: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let configured: ProviderSpecificHeaders = serde_json::from_value(configured).unwrap();
|
||||
assert_eq!(
|
||||
Value::Object(get_provider_specific_headers(
|
||||
Some(&configured),
|
||||
"anthropic"
|
||||
)),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn no_configured_headers_match_nothing() {
|
||||
assert_eq!(get_provider_specific_headers(None, "anthropic"), Map::new());
|
||||
}
|
||||
}
|
||||
|
|
@ -1,7 +1,9 @@
|
|||
pub mod call_arguments;
|
||||
pub mod core_helpers;
|
||||
pub mod dot_notation_indexing;
|
||||
pub mod exception_mapping_utils;
|
||||
pub mod get_llm_provider_logic;
|
||||
pub mod get_provider_specific_headers;
|
||||
pub mod params;
|
||||
pub mod prompt_templates;
|
||||
pub mod secret_redaction;
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
use litellm_http::request::string_headers as shared_string_headers;
|
||||
pub(super) use litellm_http::request::{has_bearer_auth, has_header, truncate_error_body};
|
||||
pub(super) use litellm_http::request::truncate_error_body;
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
|
|
@ -18,8 +20,34 @@ pub enum Error {
|
|||
Transport(#[from] litellm_http::transport::Error),
|
||||
#[error(transparent)]
|
||||
Headers(#[from] litellm_http::request::HeaderError),
|
||||
#[error(transparent)]
|
||||
Secret(#[from] SecretError),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, thiserror::Error)]
|
||||
#[error(transparent)]
|
||||
pub struct SecretError(Arc<litellm_secrets::Error>);
|
||||
|
||||
impl SecretError {
|
||||
pub fn source_error(&self) -> &litellm_secrets::Error {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl From<litellm_secrets::Error> for Error {
|
||||
fn from(error: litellm_secrets::Error) -> Self {
|
||||
Self::Secret(SecretError(Arc::new(error)))
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for SecretError {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
Arc::ptr_eq(&self.0, &other.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for SecretError {}
|
||||
|
||||
impl From<LlmError> for Error {
|
||||
fn from(error: LlmError) -> Self {
|
||||
match error {
|
||||
|
|
|
|||
|
|
@ -12,6 +12,9 @@ mod common_utils;
|
|||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_secrets::source::EnvironmentSecrets;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
|
||||
use serde_json::Value;
|
||||
|
|
@ -31,9 +34,12 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
|
|||
api_base: request.api_base.map(Into::into),
|
||||
custom_llm_provider: request.custom_llm_provider.map(Into::into),
|
||||
extra_headers: request.extra_headers,
|
||||
provider_specific_header: request.provider_specific_header,
|
||||
timeout: request.timeout,
|
||||
shaping: request.shaping,
|
||||
};
|
||||
match litellm_host::run::run(messages_machine(), &LocalMessagesHost::new(call)).await? {
|
||||
let secrets = Arc::new(EnvironmentSecrets::python_compatible());
|
||||
match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? {
|
||||
MessagesOutput::Message(message) => Ok(*message),
|
||||
MessagesOutput::Streamed => Err(Error::Unsupported(
|
||||
"streamed responses need a streaming host",
|
||||
|
|
|
|||
|
|
@ -1,51 +1,102 @@
|
|||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_llms::base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, MessagesAuthStrategy,
|
||||
use litellm_core_utils::{
|
||||
dot_notation_indexing::delete_nested_value,
|
||||
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
|
||||
get_provider_specific_headers::get_provider_specific_headers,
|
||||
settings::Lookup,
|
||||
};
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request,
|
||||
base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, MessagesTransformContext,
|
||||
},
|
||||
};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers},
|
||||
common_utils::{messages_provider_config, string_headers},
|
||||
};
|
||||
use crate::messages::types::{MessagesRequest, ProviderMessagesRequest};
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: MessagesRequest<'_>,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
|
||||
pub(super) struct ResolvedProvider<'a> {
|
||||
pub(super) model: &'a str,
|
||||
pub(super) provider: &'a str,
|
||||
pub(super) config: &'static dyn BaseAnthropicMessagesConfig,
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<ResolvedProvider<'a>, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
request
|
||||
.custom_llm_provider
|
||||
.map(|provider| CustomLlmProvider {
|
||||
model: request.model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for messages request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let model = provider_info.model.to_string();
|
||||
let provider = provider_info.custom_llm_provider;
|
||||
|
||||
let config = messages_provider_config(provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
Ok(ResolvedProvider {
|
||||
model,
|
||||
provider,
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
let headers =
|
||||
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: MessagesRequest<'_>,
|
||||
resolved: ResolvedProvider<'_>,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let ResolvedProvider {
|
||||
model,
|
||||
provider,
|
||||
config,
|
||||
} = resolved;
|
||||
let model = model.to_string();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
|
||||
let typed_request: AnthropicMessagesRequest =
|
||||
serde_json::from_value(request.body).map_err(|err| {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
})?;
|
||||
let transformed = config.transform_anthropic_messages_request(AnthropicMessagesRequest {
|
||||
model: model.clone(),
|
||||
..typed_request
|
||||
})?;
|
||||
serde_json::from_value(request.body).map_err(invalid_request)?;
|
||||
let sanitized = shape_anthropic_messages_request(
|
||||
AnthropicMessagesRequest {
|
||||
model: model.clone(),
|
||||
..typed_request
|
||||
},
|
||||
request.shaping.reasoning_auto_summary,
|
||||
)?;
|
||||
let trimmed =
|
||||
without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?;
|
||||
let transformed = config.transform_anthropic_messages_request(
|
||||
trimmed,
|
||||
&MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params),
|
||||
)?;
|
||||
|
||||
let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider);
|
||||
let forwarded = string_headers(Some(
|
||||
request
|
||||
.extra_headers
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(scoped)
|
||||
.collect(),
|
||||
))?;
|
||||
let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?;
|
||||
let headers = config.request_headers(
|
||||
with_default_headers(authenticated, config.default_headers()),
|
||||
&transformed,
|
||||
);
|
||||
|
||||
let body = serde_json::to_value(transformed).map_err(|err| {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
|
|
@ -65,33 +116,371 @@ pub(super) fn prepare_provider_request(
|
|||
})
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
config: &dyn BaseAnthropicMessagesConfig,
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
let mut headers = string_headers(extra_headers)?;
|
||||
|
||||
let auth_strategy = config.auth_strategy();
|
||||
let already_authorized = has_header(&headers, auth_strategy.header_name())
|
||||
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
|
||||
if !already_authorized {
|
||||
let api_key = config.resolve_api_key(api_key, env_lookup)?;
|
||||
let auth_header = match auth_strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
headers.push(auth_header);
|
||||
}
|
||||
|
||||
for (name, value) in config.default_headers() {
|
||||
if !has_header(&headers, name) {
|
||||
headers.push((name.to_string(), value.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(headers)
|
||||
fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
}
|
||||
|
||||
fn without_additional_drop_params(
|
||||
request: AnthropicMessagesRequest,
|
||||
paths: &[String],
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
if paths.is_empty() {
|
||||
return Ok(request);
|
||||
}
|
||||
let Value::Object(fields) = serde_json::to_value(request).map_err(invalid_request)? else {
|
||||
return Err(Error::InvalidRequest(
|
||||
"Anthropic messages request did not serialize to an object".to_string(),
|
||||
));
|
||||
};
|
||||
let (required, optional): (Map<String, Value>, Map<String, Value>) = fields
|
||||
.into_iter()
|
||||
.partition(|(key, _)| matches!(key.as_str(), "model" | "messages"));
|
||||
let trimmed = paths.iter().fold(Value::Object(optional), |body, path| {
|
||||
delete_nested_value(body, path)
|
||||
});
|
||||
let merged: Map<String, Value> = required
|
||||
.into_iter()
|
||||
.chain(trimmed.as_object().cloned().unwrap_or_default())
|
||||
.collect();
|
||||
serde_json::from_value(Value::Object(merged)).map_err(invalid_request)
|
||||
}
|
||||
|
||||
fn with_default_headers(
|
||||
headers: Vec<(String, String)>,
|
||||
defaults: &[(&str, &str)],
|
||||
) -> Vec<(String, String)> {
|
||||
let missing: Vec<(String, String)> = defaults
|
||||
.iter()
|
||||
.filter(|(name, _)| {
|
||||
!headers
|
||||
.iter()
|
||||
.any(|(header, _)| header.eq_ignore_ascii_case(name))
|
||||
})
|
||||
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
|
||||
.collect();
|
||||
headers.into_iter().chain(missing).collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::messages::types::MessagesShaping;
|
||||
|
||||
#[fixture]
|
||||
fn shaping() -> MessagesShaping {
|
||||
MessagesShaping::default()
|
||||
}
|
||||
|
||||
fn prepare(request: MessagesRequest<'_>) -> Result<ProviderMessagesRequest, Error> {
|
||||
prepare_with_secrets(request, &|_: &str| None)
|
||||
}
|
||||
|
||||
fn prepare_with_secrets(
|
||||
request: MessagesRequest<'_>,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let resolved = resolve_provider(request.model, request.custom_llm_provider)?;
|
||||
prepare_provider_request(request, resolved, secrets)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
)]
|
||||
#[case::auth_token(
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "token")],
|
||||
&[("authorization", "Bearer token")],
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
)]
|
||||
#[case::api_base(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_API_BASE", "https://gateway.test")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://gateway.test/v1/messages"
|
||||
)]
|
||||
#[case::sdk_base_url(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_BASE_URL", "https://sdk.test")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://sdk.test/v1/messages"
|
||||
)]
|
||||
fn credentials_and_base_come_from_the_resolved_secrets(
|
||||
shaping: MessagesShaping,
|
||||
#[case] secrets: &[(&str, &str)],
|
||||
#[case] expected_auth: &[(&str, &str)],
|
||||
#[case] expected_url: &str,
|
||||
) {
|
||||
let lookup = |name: &str| {
|
||||
secrets
|
||||
.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
};
|
||||
let prepared = prepare_with_secrets(
|
||||
MessagesRequest {
|
||||
model: "claude-test",
|
||||
body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
shaping,
|
||||
},
|
||||
&lookup,
|
||||
)
|
||||
.unwrap();
|
||||
let auth: Vec<(&str, &str)> = prepared
|
||||
.upstream_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization"))
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
.collect();
|
||||
assert_eq!(
|
||||
(auth.as_slice(), prepared.url.as_str()),
|
||||
(expected_auth, expected_url)
|
||||
);
|
||||
}
|
||||
|
||||
fn prepared_body(body: Value, shaping: MessagesShaping) -> Result<Value, Error> {
|
||||
prepare(MessagesRequest {
|
||||
model: "anthropic/claude-test",
|
||||
body,
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some("https://anthropic.test"),
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
shaping,
|
||||
})
|
||||
.map(|prepared| prepared.body)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::nothing_forwarded(
|
||||
&[],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
)]
|
||||
#[case::forwarded_header_wins_in_any_case(
|
||||
&[("X-Version", "custom"), ("x-api-key", "k")],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
&[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")],
|
||||
)]
|
||||
#[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])]
|
||||
fn default_headers_fill_only_missing_names(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] defaults: &[(&str, &str)],
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let owned = |headers: &[(&str, &str)]| -> Vec<(String, String)> {
|
||||
headers
|
||||
.iter()
|
||||
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
|
||||
.collect()
|
||||
};
|
||||
assert_eq!(
|
||||
with_default_headers(owned(forwarded), defaults),
|
||||
owned(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::top_level_and_nested_paths(
|
||||
json!({
|
||||
"max_tokens": 1024,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048},
|
||||
"context_management": {"edits": [{"type": "clear_thinking_20251015"}]},
|
||||
"metadata": {"user_id": "u1"},
|
||||
"tools": [{"name": "lookup", "input_schema": {"type": "object"}, "input_examples": [{"q": "x"}]}]
|
||||
}),
|
||||
&["thinking", "context_management", "tools[*].input_examples"],
|
||||
json!({
|
||||
"max_tokens": 1024,
|
||||
"metadata": {"user_id": "u1"},
|
||||
"tools": [{"name": "lookup", "input_schema": {"type": "object"}}]
|
||||
}),
|
||||
)]
|
||||
#[case::no_paths(
|
||||
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
|
||||
&[],
|
||||
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
|
||||
)]
|
||||
#[case::model_and_messages_are_never_dropped(
|
||||
json!({"max_tokens": 16}),
|
||||
&["model", "messages", "messages[0].content"],
|
||||
json!({"max_tokens": 16}),
|
||||
)]
|
||||
fn prepared_body_drops_configured_paths(
|
||||
shaping: MessagesShaping,
|
||||
#[case] fields: Value,
|
||||
#[case] additional_drop_params: &[&str],
|
||||
#[case] expected_fields: Value,
|
||||
) {
|
||||
let with_messages = |fields: Value| -> Value {
|
||||
let Value::Object(fields) = fields else {
|
||||
unreachable!()
|
||||
};
|
||||
Value::Object(
|
||||
[
|
||||
("model".to_string(), json!("claude-test")),
|
||||
(
|
||||
"messages".to_string(),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
),
|
||||
]
|
||||
.into_iter()
|
||||
.chain(fields)
|
||||
.collect(),
|
||||
)
|
||||
};
|
||||
let shaping = MessagesShaping {
|
||||
additional_drop_params: additional_drop_params
|
||||
.iter()
|
||||
.map(ToString::to_string)
|
||||
.collect(),
|
||||
..shaping
|
||||
};
|
||||
assert_eq!(
|
||||
prepared_body(with_messages(fields), shaping),
|
||||
Ok(with_messages(expected_fields))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::model_prefix_picks_the_provider(
|
||||
"azure_ai/claude-test",
|
||||
None,
|
||||
&[("x-priority", "extra"), ("x-scoped", "azure_ai")]
|
||||
)]
|
||||
#[case::explicit_provider(
|
||||
"claude-test",
|
||||
Some("anthropic"),
|
||||
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
|
||||
)]
|
||||
#[case::provider_prefix_on_an_anthropic_model(
|
||||
"anthropic/claude-test",
|
||||
None,
|
||||
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
|
||||
)]
|
||||
fn provider_specific_headers_follow_the_resolved_provider(
|
||||
shaping: MessagesShaping,
|
||||
#[case] model: &str,
|
||||
#[case] custom_llm_provider: Option<&str>,
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let configured: ProviderSpecificHeaders = serde_json::from_value(json!([
|
||||
{"custom_llm_provider": "azure_ai", "extra_headers": {"x-scoped": "azure_ai"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}}
|
||||
]))
|
||||
.unwrap();
|
||||
let prepared = prepare(MessagesRequest {
|
||||
model,
|
||||
body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some("https://resource.services.ai.azure.com"),
|
||||
custom_llm_provider,
|
||||
extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()),
|
||||
provider_specific_header: Some(configured),
|
||||
timeout: None,
|
||||
shaping,
|
||||
})
|
||||
.unwrap();
|
||||
let caller_headers: Vec<(&str, &str)> = prepared
|
||||
.upstream_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped"))
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
.collect();
|
||||
assert_eq!(caller_headers, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn prepared_body_carries_the_provider_stripped_model(shaping: MessagesShaping) {
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "anthropic/claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Ok(json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn dropped_thinking_display_is_not_restored_by_auto_summary(shaping: MessagesShaping) {
|
||||
let shaping = MessagesShaping {
|
||||
reasoning_auto_summary: true,
|
||||
additional_drop_params: vec!["thinking.display".to_string()],
|
||||
..shaping
|
||||
};
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Ok(json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048}
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn dropping_an_invalid_metadata_user_id_does_not_skip_its_validation(shaping: MessagesShaping) {
|
||||
let shaping = MessagesShaping {
|
||||
additional_drop_params: vec!["metadata.user_id".to_string()],
|
||||
..shaping
|
||||
};
|
||||
assert!(matches!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16,
|
||||
"metadata": {"user_id": 123}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Err(Error::InvalidRequest(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn prepared_body_rejects_invalid_metadata_before_the_call(shaping: MessagesShaping) {
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16,
|
||||
"metadata": {"user_id": 123}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Err(Error::InvalidRequest(
|
||||
"metadata.user_id must be a string, got 123".to_string()
|
||||
))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,7 @@
|
|||
use std::{sync::Mutex, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, Mutex},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_auth::SecretValue;
|
||||
|
|
@ -9,15 +12,19 @@ use litellm_host::{
|
|||
machine::{HostChannel, MachineFault, RouteMachine},
|
||||
route::Route,
|
||||
};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
|
||||
utils::ProviderSpecificHeaders,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::messages_provider_config,
|
||||
handler::{decode_response, network, provider_error, send},
|
||||
prepare::prepare_provider_request,
|
||||
types::MessagesRequest,
|
||||
prepare::{prepare_provider_request, resolve_provider},
|
||||
types::{MessagesRequest, MessagesShaping},
|
||||
};
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
|
||||
|
|
@ -38,7 +45,9 @@ pub struct MessagesCall {
|
|||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
impl MessagesCall {
|
||||
|
|
@ -120,22 +129,33 @@ impl Host<Messages> for LocalMessagesHost {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn messages_machine() -> MessagesMachine {
|
||||
RouteMachine::new(|host| Box::pin(execute(host)))
|
||||
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
|
||||
}
|
||||
|
||||
async fn execute(host: MessagesHost) -> Result<MessagesOutput, Error> {
|
||||
async fn execute(
|
||||
host: MessagesHost,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
|
||||
let stream = call.streams();
|
||||
let request = prepare_provider_request(MessagesRequest {
|
||||
model: &call.model,
|
||||
body: Value::Object(call.body.clone()),
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers.clone(),
|
||||
timeout: call.timeout,
|
||||
})?;
|
||||
let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
|
||||
let secrets = secrets.resolve(resolved.config.secret_names()).await?;
|
||||
let request = prepare_provider_request(
|
||||
MessagesRequest {
|
||||
model: &call.model,
|
||||
body: Value::Object(call.body.clone()),
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers.clone(),
|
||||
provider_specific_header: call.provider_specific_header.clone(),
|
||||
timeout: call.timeout,
|
||||
shaping: call.shaping.clone(),
|
||||
},
|
||||
resolved,
|
||||
secrets.as_ref(),
|
||||
)?;
|
||||
if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER {
|
||||
return Err(Error::Unsupported("streaming messages for this provider"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
use std::time::Duration;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_http::request::{has_bearer_auth, has_header};
|
||||
use litellm_secrets::{SecretValue, source::SecretSource};
|
||||
use serde_json::{Map, Value, json};
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
|
|
@ -8,12 +11,132 @@ use tokio::{
|
|||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{
|
||||
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
|
||||
},
|
||||
common_utils::{messages_provider_config, string_headers, truncate_error_body},
|
||||
messages,
|
||||
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
|
||||
};
|
||||
use crate::messages::types::MessagesRequest;
|
||||
use crate::messages::types::{MessagesRequest, MessagesShaping};
|
||||
|
||||
struct RecordingSecrets {
|
||||
values: Vec<(&'static str, String)>,
|
||||
fails: bool,
|
||||
requested: std::sync::Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl RecordingSecrets {
|
||||
fn new(values: Vec<(&'static str, String)>, fails: bool) -> Self {
|
||||
Self {
|
||||
values,
|
||||
fails,
|
||||
requested: std::sync::Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl SecretSource for RecordingSecrets {
|
||||
fn get_secret_str<'a>(
|
||||
&'a self,
|
||||
name: &'a str,
|
||||
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
|
||||
Box::pin(async move {
|
||||
self.requested.lock().unwrap().push(name.to_string());
|
||||
if self.fails {
|
||||
return Err(litellm_secrets::Error::ManagedSecretMissing);
|
||||
}
|
||||
Ok(self
|
||||
.values
|
||||
.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| SecretValue::new(value.clone())))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn secrets_call() -> MessagesCall {
|
||||
let Value::Object(body) = json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}) else {
|
||||
unreachable!("literal object")
|
||||
};
|
||||
MessagesCall {
|
||||
model: "claude-sonnet-4-5".into(),
|
||||
body,
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn route_reads_the_provider_credential_and_base_from_the_secret_source() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
let addr = listener.local_addr().expect("addr");
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
||||
let request = read_http_request(&mut socket).await;
|
||||
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
let secrets = Arc::new(RecordingSecrets::new(
|
||||
vec![
|
||||
("ANTHROPIC_API_KEY", "sk-from-manager".to_string()),
|
||||
("ANTHROPIC_BASE_URL", format!("http://{addr}")),
|
||||
],
|
||||
false,
|
||||
));
|
||||
|
||||
let output = litellm_host::run::run(
|
||||
messages_machine(secrets.clone()),
|
||||
&LocalMessagesHost::new(secrets_call()),
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
let request = server.await.expect("server task completes");
|
||||
assert!(
|
||||
request
|
||||
.to_ascii_lowercase()
|
||||
.contains("x-api-key: sk-from-manager"),
|
||||
"{request}"
|
||||
);
|
||||
let requested = secrets.requested.lock().unwrap().clone();
|
||||
assert_eq!(
|
||||
requested,
|
||||
messages_provider_config("anthropic")
|
||||
.unwrap()
|
||||
.secret_names()
|
||||
.iter()
|
||||
.map(ToString::to_string)
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn route_surfaces_a_secret_manager_failure_before_the_call() {
|
||||
let Err(error) = litellm_host::run::run(
|
||||
messages_machine(Arc::new(RecordingSecrets::new(Vec::new(), true))),
|
||||
&LocalMessagesHost::new(secrets_call()),
|
||||
)
|
||||
.await
|
||||
else {
|
||||
panic!("a secret manager failure fails the call");
|
||||
};
|
||||
assert!(
|
||||
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
|
||||
"{error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
async fn read_http_request(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
|
|
@ -159,7 +282,9 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -215,7 +340,9 @@ async fn messages_round_trip_builds_native_anthropic_request() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -268,7 +395,9 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -322,7 +451,9 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("entra id request succeeds without api key");
|
||||
|
|
@ -346,7 +477,9 @@ async fn messages_requires_auth_when_no_key_and_no_header() {
|
|||
api_base: Some("http://127.0.0.1:1"),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_millis(50)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("missing auth errors");
|
||||
|
|
@ -384,7 +517,9 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("falls back to api key");
|
||||
|
|
@ -425,7 +560,9 @@ async fn messages_maps_provider_error_status_to_http_error() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("provider error propagates");
|
||||
|
|
@ -445,7 +582,9 @@ async fn messages_rejects_unsupported_provider() {
|
|||
api_base: Some("http://127.0.0.1:1"),
|
||||
custom_llm_provider: Some("openai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_millis(50)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("unsupported provider errors");
|
||||
|
|
|
|||
|
|
@ -1,8 +1,25 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig;
|
||||
use litellm_llms::{
|
||||
anthropic::common_utils::AnthropicModelCapabilities,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
};
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct MessagesShaping {
|
||||
#[serde(default)]
|
||||
pub capabilities: AnthropicModelCapabilities,
|
||||
#[serde(default)]
|
||||
pub drop_params: bool,
|
||||
#[serde(default)]
|
||||
pub reasoning_auto_summary: bool,
|
||||
#[serde(default)]
|
||||
pub additional_drop_params: Vec<String>,
|
||||
}
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
|
|
@ -10,7 +27,9 @@ pub struct MessagesRequest<'a> {
|
|||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
pub struct ProviderMessagesRequest {
|
||||
|
|
@ -22,3 +41,86 @@ pub struct ProviderMessagesRequest {
|
|||
pub upstream_headers: Vec<(String, String)>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::nothing_projected(json!({}), MessagesShaping::default())]
|
||||
#[case::only_drop_params(
|
||||
json!({"drop_params": true}),
|
||||
MessagesShaping { drop_params: true, ..MessagesShaping::default() },
|
||||
)]
|
||||
#[case::only_reasoning_auto_summary(
|
||||
json!({"reasoning_auto_summary": true}),
|
||||
MessagesShaping { reasoning_auto_summary: true, ..MessagesShaping::default() },
|
||||
)]
|
||||
#[case::only_additional_drop_params(
|
||||
json!({"additional_drop_params": ["tools[*].input_examples"]}),
|
||||
MessagesShaping {
|
||||
additional_drop_params: vec!["tools[*].input_examples".to_string()],
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
)]
|
||||
#[case::partial_capabilities(
|
||||
json!({"capabilities": {"supports_reasoning": true}}),
|
||||
MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
)]
|
||||
#[case::everything_the_python_host_projects(
|
||||
json!({
|
||||
"capabilities": {
|
||||
"supports_reasoning": true,
|
||||
"supports_adaptive_thinking": true,
|
||||
"thinking_always_on": false,
|
||||
"supports_legacy_thinking": false,
|
||||
"supports_output_config": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_speed": true,
|
||||
"effort_tiers": {"minimal": false, "low": true, "medium": true, "high": true, "xhigh": true, "max": false}
|
||||
},
|
||||
"drop_params": true,
|
||||
"reasoning_auto_summary": true,
|
||||
"additional_drop_params": ["metadata.user_id", "thinking"]
|
||||
}),
|
||||
MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
thinking_always_on: false,
|
||||
supports_legacy_thinking: false,
|
||||
supports_output_config: true,
|
||||
supports_sampling_params: false,
|
||||
supports_speed: true,
|
||||
effort_tiers: SupportedEffortTiers {
|
||||
minimal: false,
|
||||
low: true,
|
||||
medium: true,
|
||||
high: true,
|
||||
xhigh: true,
|
||||
max: false,
|
||||
},
|
||||
},
|
||||
drop_params: true,
|
||||
reasoning_auto_summary: true,
|
||||
additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()],
|
||||
},
|
||||
)]
|
||||
fn shaping_deserializes_with_defaults_for_absent_fields(
|
||||
#[case] projected: Value,
|
||||
#[case] expected: MessagesShaping,
|
||||
) {
|
||||
let shaping: MessagesShaping = serde_json::from_value(projected).unwrap();
|
||||
assert_eq!(shaping, expected);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
1568
litellm-rust/crates/llms/src/anthropic/common_utils.rs
Normal file
1568
litellm-rust/crates/llms/src/anthropic/common_utils.rs
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -0,0 +1,270 @@
|
|||
use litellm_types::llms::anthropic_messages::anthropic_request::{
|
||||
AnthropicMessage, AnthropicMessagesRequest,
|
||||
};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::{
|
||||
anthropic::common_utils::{
|
||||
flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks,
|
||||
strip_provider_specific_fields,
|
||||
},
|
||||
base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
pub fn shape_anthropic_messages_request(
|
||||
request: AnthropicMessagesRequest,
|
||||
reasoning_auto_summary: bool,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
Ok(AnthropicMessagesRequest {
|
||||
messages: sanitize_anthropic_messages(request.messages),
|
||||
metadata: request
|
||||
.metadata
|
||||
.as_ref()
|
||||
.map(validate_anthropic_api_metadata)
|
||||
.transpose()?,
|
||||
thinking: with_reasoning_auto_summary(request.thinking, reasoning_auto_summary),
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
fn sanitize_anthropic_messages(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
|
||||
strip_provider_specific_fields(flatten_unencrypted_web_search_results(
|
||||
sanitize_tool_use_ids(strip_empty_content_blocks(messages)),
|
||||
))
|
||||
}
|
||||
|
||||
fn validate_anthropic_api_metadata(metadata: &Value) -> Result<Value, Error> {
|
||||
let Value::Object(fields) = metadata else {
|
||||
return Err(Error::InvalidRequest(format!(
|
||||
"metadata must be an object, got {metadata}"
|
||||
)));
|
||||
};
|
||||
match fields.get("user_id") {
|
||||
None | Some(Value::Null) => Ok(json!({})),
|
||||
Some(Value::String(user_id)) => Ok(json!({"user_id": user_id})),
|
||||
Some(other) => Err(Error::InvalidRequest(format!(
|
||||
"metadata.user_id must be a string, got {other}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn with_reasoning_auto_summary(thinking: Option<Value>, enabled: bool) -> Option<Value> {
|
||||
let Some(Value::Object(thinking)) = thinking else {
|
||||
return thinking;
|
||||
};
|
||||
if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") {
|
||||
return Some(Value::Object(thinking));
|
||||
}
|
||||
Some(Value::Object(
|
||||
thinking
|
||||
.into_iter()
|
||||
.filter(|(key, _)| key != "display")
|
||||
.chain([("display".to_string(), json!("summarized"))])
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn messages(value: Value) -> Vec<AnthropicMessage> {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
fn request(body: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(body).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty_text_next_to_a_tool_use(
|
||||
json!([{"role": "assistant", "content": [
|
||||
{"type": "text", "text": " "},
|
||||
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
|
||||
]}]),
|
||||
json!([{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
|
||||
]}]),
|
||||
)]
|
||||
#[case::cross_provider_tool_ids(
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
|
||||
]),
|
||||
)]
|
||||
#[case::replayed_unencrypted_web_search_results(
|
||||
json!([
|
||||
{"role": "user", "content": "latest litellm version?"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "latest litellm version"}},
|
||||
{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_1", "content": [{
|
||||
"type": "web_search_result",
|
||||
"url": "https://github.com/BerriAI/litellm/releases",
|
||||
"title": "Releases",
|
||||
"page_age": null,
|
||||
"encrypted_content": "",
|
||||
"snippet": "Latest release v1.95.0"
|
||||
}]}
|
||||
]},
|
||||
{"role": "user", "content": "which version?"}
|
||||
]),
|
||||
json!([
|
||||
{"role": "user", "content": "latest litellm version?"},
|
||||
{"role": "assistant", "content": [{
|
||||
"type": "text",
|
||||
"text": "Web search results for 'latest litellm version':\n\nTitle: Releases\nURL: https://github.com/BerriAI/litellm/releases\nSnippet: Latest release v1.95.0"
|
||||
}]},
|
||||
{"role": "user", "content": "which version?"}
|
||||
]),
|
||||
)]
|
||||
#[case::replayed_provider_specific_fields(
|
||||
json!([
|
||||
{"role": "assistant", "content": [{
|
||||
"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"},
|
||||
"provider_specific_fields": {"signature": "sig_abc"}
|
||||
}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
|
||||
]),
|
||||
)]
|
||||
#[case::ids_are_normalized_before_web_search_results_flatten(
|
||||
json!([
|
||||
{"role": "user", "content": "run it"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "", "signature": "sig"},
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}, "provider_specific_fields": {"x": 1}},
|
||||
{"type": "server_tool_use", "id": "srv.1", "name": "web_search", "input": {"query": "q"}, "provider_specific_fields": {"x": 2}},
|
||||
{"type": "web_search_tool_result", "tool_use_id": "srv.1", "provider_specific_fields": {"x": 3}, "content": [
|
||||
{"type": "web_search_result", "url": "u", "title": "", "encrypted_content": "", "provider_specific_fields": {"x": 4}}
|
||||
]}
|
||||
]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": " "}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "user", "content": "run it"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}},
|
||||
{"type": "server_tool_use", "id": "srv_1", "name": "web_search", "input": {"query": "q"}},
|
||||
{"type": "text", "text": "Web search results:\n\nURL: u"}
|
||||
]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
|
||||
]),
|
||||
)]
|
||||
fn sanitize_anthropic_messages_cleans_replayed_history(
|
||||
#[case] history: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
serde_json::to_value(sanitize_anthropic_messages(messages(history))).unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::keeps_only_user_id(json!({"user_id": "u-1", "trace_id": "internal"}), Ok(json!({"user_id": "u-1"})))]
|
||||
#[case::null_user_id(json!({"user_id": null, "trace_id": "internal"}), Ok(json!({})))]
|
||||
#[case::no_user_id(json!({"trace_id": "internal"}), Ok(json!({})))]
|
||||
#[case::empty(json!({}), Ok(json!({})))]
|
||||
#[case::numeric_user_id(
|
||||
json!({"user_id": 123}),
|
||||
Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string())),
|
||||
)]
|
||||
#[case::boolean_user_id(
|
||||
json!({"user_id": true}),
|
||||
Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string())),
|
||||
)]
|
||||
#[case::not_an_object(
|
||||
json!(["u-1"]),
|
||||
Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string())),
|
||||
)]
|
||||
fn validate_anthropic_api_metadata_passes_only_a_string_user_id(
|
||||
#[case] metadata: Value,
|
||||
#[case] expected: Result<Value, Error>,
|
||||
) {
|
||||
assert_eq!(validate_anthropic_api_metadata(&metadata), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::adaptive(
|
||||
Some(json!({"type": "adaptive", "budget_tokens": 5000})),
|
||||
true,
|
||||
Some(json!({"type": "adaptive", "budget_tokens": 5000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::enabled(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::no_type(Some(json!({})), true, Some(json!({"display": "summarized"})))]
|
||||
#[case::display_omitted_is_overridden(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "omitted"})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::display_summarized_is_kept(
|
||||
Some(json!({"type": "enabled", "display": "summarized"})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "display": "summarized"})),
|
||||
)]
|
||||
#[case::disabled_thinking(Some(json!({"type": "disabled"})), true, Some(json!({"type": "disabled"})))]
|
||||
#[case::flag_off(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
false,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
)]
|
||||
#[case::flag_off_keeps_callers_display(
|
||||
Some(json!({"type": "enabled", "display": "omitted"})),
|
||||
false,
|
||||
Some(json!({"type": "enabled", "display": "omitted"})),
|
||||
)]
|
||||
#[case::no_thinking(None, true, None)]
|
||||
#[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))]
|
||||
fn reasoning_auto_summary_marks_active_thinking_as_summarized(
|
||||
#[case] thinking: Option<Value>,
|
||||
#[case] enabled: bool,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shaping_cleans_messages_metadata_and_thinking() {
|
||||
let sanitized = shape_anthropic_messages_request(
|
||||
request(json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "assistant", "content": [
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}
|
||||
]}],
|
||||
"metadata": {"user_id": "u", "trace_id": "t"},
|
||||
"thinking": {"type": "enabled", "budget_tokens": 1024},
|
||||
"safeguards": [{"type": "dangerous_tool_use"}]
|
||||
})),
|
||||
true,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(sanitized).unwrap(),
|
||||
json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}
|
||||
]}],
|
||||
"metadata": {"user_id": "u"},
|
||||
"thinking": {"type": "enabled", "budget_tokens": 1024, "display": "summarized"},
|
||||
"safeguards": [{"type": "dangerous_tool_use"}]
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,643 @@
|
|||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
anthropic::{
|
||||
ANTHROPIC_OAUTH_TOKEN_PREFIX,
|
||||
common_utils::{
|
||||
ANTHROPIC_OAUTH_BETA_HEADER, beta, has_advisor_tool, is_anthropic_oauth_key,
|
||||
is_tool_search_used, join_beta_values, requires_native_compaction_beta,
|
||||
split_beta_values,
|
||||
},
|
||||
},
|
||||
base_llm::anthropic_messages::transformation::Headers,
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
const BETA_HEADER: &str = "anthropic-beta";
|
||||
const AUTHORIZATION: &str = "authorization";
|
||||
const API_KEY_HEADER: &str = "x-api-key";
|
||||
const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access";
|
||||
|
||||
fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
|
||||
headers
|
||||
.iter()
|
||||
.find(|(header, _)| header.eq_ignore_ascii_case(name))
|
||||
.map(|(_, value)| value.as_str())
|
||||
}
|
||||
|
||||
fn without(headers: Headers, names: &[&str]) -> Headers {
|
||||
headers
|
||||
.into_iter()
|
||||
.filter(|(header, _)| !names.iter().any(|name| header.eq_ignore_ascii_case(name)))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn existing_betas(headers: &[(String, String)]) -> impl Iterator<Item = String> + '_ {
|
||||
headers
|
||||
.iter()
|
||||
.filter(|(header, _)| header.eq_ignore_ascii_case(BETA_HEADER))
|
||||
.flat_map(|(_, value)| split_beta_values(Some(value)))
|
||||
}
|
||||
|
||||
fn with_oauth_bearer(headers: Headers, bearer: String) -> Headers {
|
||||
let beta =
|
||||
join_beta_values(existing_betas(&headers).chain([ANTHROPIC_OAUTH_BETA_HEADER.to_string()]));
|
||||
without(headers, &[API_KEY_HEADER, AUTHORIZATION, BETA_HEADER])
|
||||
.into_iter()
|
||||
.chain([
|
||||
(AUTHORIZATION.to_string(), bearer),
|
||||
(BETA_HEADER.to_string(), beta),
|
||||
(DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string()),
|
||||
])
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
value.map(str::trim).filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub fn authenticate(
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, litellm_auth::Error> {
|
||||
if let Some(forwarded) = header_value(&headers, AUTHORIZATION)
|
||||
&& forwarded
|
||||
.strip_prefix("Bearer ")
|
||||
.is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX))
|
||||
{
|
||||
let bearer = forwarded.to_string();
|
||||
return Ok(with_oauth_bearer(headers, bearer));
|
||||
}
|
||||
if let Some(key) = api_key.filter(|key| key.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) {
|
||||
return Ok(with_oauth_bearer(headers, format!("Bearer {key}")));
|
||||
}
|
||||
if header_value(&headers, API_KEY_HEADER).is_some()
|
||||
|| header_value(&headers, AUTHORIZATION).is_some()
|
||||
{
|
||||
return Ok(headers);
|
||||
}
|
||||
let resolved_key = non_empty(api_key)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty()));
|
||||
let auth = match resolved_key {
|
||||
Some(key) if is_anthropic_oauth_key(&key) => {
|
||||
(AUTHORIZATION.to_string(), format!("Bearer {key}"))
|
||||
}
|
||||
Some(key) => (API_KEY_HEADER.to_string(), key),
|
||||
None => match env_lookup(ANTHROPIC_AUTH_TOKEN_ENV).filter(|value| !value.trim().is_empty())
|
||||
{
|
||||
Some(token) => (AUTHORIZATION.to_string(), format!("Bearer {token}")),
|
||||
None => {
|
||||
return Err(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
});
|
||||
}
|
||||
},
|
||||
};
|
||||
Ok(headers.into_iter().chain([auth]).collect())
|
||||
}
|
||||
|
||||
fn context_management_betas(
|
||||
context_management: Option<&Value>,
|
||||
) -> impl Iterator<Item = &'static str> {
|
||||
let edits = context_management
|
||||
.and_then(|value| value.get("edits"))
|
||||
.and_then(Value::as_array)
|
||||
.map(Vec::as_slice)
|
||||
.unwrap_or(&[]);
|
||||
let (compact, other) = edits.iter().fold((false, false), |(compact, other), edit| {
|
||||
match edit.get("type").and_then(Value::as_str) {
|
||||
Some("compact_20260112") => (true, other),
|
||||
_ => (compact, true),
|
||||
}
|
||||
});
|
||||
compact
|
||||
.then_some(beta::COMPACT_2026_01_12)
|
||||
.into_iter()
|
||||
.chain(other.then_some(beta::CONTEXT_MANAGEMENT_2025_06_27))
|
||||
}
|
||||
|
||||
fn uses_structured_output(request: &AnthropicMessagesRequest) -> bool {
|
||||
request.output_format.is_some()
|
||||
|| request
|
||||
.output_config
|
||||
.as_ref()
|
||||
.and_then(|config| config.get("format"))
|
||||
.is_some_and(|format| !format.is_null())
|
||||
}
|
||||
|
||||
fn messages_carry_output_config(request: &AnthropicMessagesRequest) -> bool {
|
||||
request
|
||||
.messages
|
||||
.iter()
|
||||
.any(|message| message.extra.contains_key("output_config"))
|
||||
}
|
||||
|
||||
pub fn feature_betas(request: &AnthropicMessagesRequest) -> Vec<&'static str> {
|
||||
let tools = request.tools.as_deref();
|
||||
[
|
||||
requires_native_compaction_beta(request.compaction.as_ref(), &request.messages)
|
||||
.then_some(beta::COMPACT_2026_09_04),
|
||||
uses_structured_output(request).then_some(beta::STRUCTURED_OUTPUT),
|
||||
(request.speed.as_deref() == Some("fast")).then_some(beta::FAST_MODE_2026_02_01),
|
||||
messages_carry_output_config(request).then_some(beta::PER_TURN_CONTROL_2026_07_01),
|
||||
has_advisor_tool(tools).then_some(beta::ADVISOR_TOOL_2026_03_01),
|
||||
is_tool_search_used(tools).then_some(beta::ADVANCED_TOOL_USE_2025_11_20),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(context_management_betas(
|
||||
request.context_management.as_ref(),
|
||||
))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn with_feature_betas(headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
let existing = existing_betas(&headers).collect::<Vec<_>>();
|
||||
let features = feature_betas(request);
|
||||
if existing.is_empty() && features.is_empty() {
|
||||
return headers;
|
||||
}
|
||||
let merged = join_beta_values(
|
||||
existing
|
||||
.into_iter()
|
||||
.chain(features.into_iter().map(str::to_string)),
|
||||
);
|
||||
without(headers, &[BETA_HEADER])
|
||||
.into_iter()
|
||||
.chain([(BETA_HEADER.to_string(), merged)])
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
const OAUTH_TOKEN: &str = "sk-ant-oat01-token";
|
||||
const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token";
|
||||
const REGULAR_KEY: &str = "sk-ant-api03-regular";
|
||||
const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true");
|
||||
|
||||
type Env = &'static [(&'static str, &'static str)];
|
||||
|
||||
fn request(fields: Value) -> AnthropicMessagesRequest {
|
||||
let mut body =
|
||||
json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]});
|
||||
body.as_object_mut()
|
||||
.unwrap()
|
||||
.extend(fields.as_object().unwrap().clone());
|
||||
serde_json::from_value(body).unwrap()
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn betas(values: &[&str]) -> String {
|
||||
values.join(",")
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn no_env() -> Env {
|
||||
&[]
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn full_env() -> Env {
|
||||
&[
|
||||
("ANTHROPIC_API_KEY", "sk-env"),
|
||||
("ANTHROPIC_AUTH_TOKEN", "env-token"),
|
||||
]
|
||||
}
|
||||
|
||||
fn authenticate_with(
|
||||
forwarded: &[(&str, &str)],
|
||||
api_key: Option<&str>,
|
||||
env: Env,
|
||||
) -> Result<Headers, litellm_auth::Error> {
|
||||
let lookup = |name: &str| {
|
||||
env.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
};
|
||||
authenticate(headers(forwarded), api_key, &lookup)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_drops_forwarded_and_deployment_keys(
|
||||
&[("X-Api-Key", REGULAR_KEY), ("Authorization", OAUTH_BEARER)],
|
||||
Some(REGULAR_KEY),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_in_uppercase_authorization_header(
|
||||
&[("AUTHORIZATION", OAUTH_BEARER)],
|
||||
None,
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_keeps_unrelated_headers_in_place(
|
||||
&[("anthropic-version", "2023-06-01"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
OAUTH_BEARER,
|
||||
&[("anthropic-version", "2023-06-01")],
|
||||
)]
|
||||
#[case::forwarded_bearer_wins_over_an_oauth_api_key(
|
||||
&[("authorization", OAUTH_BEARER)],
|
||||
Some("sk-ant-oat01-deployment"),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_authenticates_as_a_bearer(&[], Some(OAUTH_TOKEN), OAUTH_BEARER, &[])]
|
||||
#[case::api_key_removes_a_forwarded_x_api_key(
|
||||
&[("x-api-key", OAUTH_TOKEN)],
|
||||
Some(OAUTH_TOKEN),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_replaces_a_forwarded_non_oauth_bearer(
|
||||
&[("Authorization", "Bearer some-proxy-token")],
|
||||
Some(OAUTH_TOKEN),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
fn oauth_token_is_the_whole_credential(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] expected_bearer: &str,
|
||||
#[case] kept: &[(&str, &str)],
|
||||
full_env: Env,
|
||||
) {
|
||||
let expected = kept
|
||||
.iter()
|
||||
.copied()
|
||||
.chain([
|
||||
("authorization", expected_bearer),
|
||||
("anthropic-beta", ANTHROPIC_OAUTH_BETA_HEADER),
|
||||
BROWSER_ACCESS,
|
||||
])
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, full_env).unwrap(),
|
||||
headers(&expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_merges_a_differently_cased_beta_header(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::forwarded_bearer_dedupes_an_existing_oauth_beta(
|
||||
&[("anthropic-beta", "web-search-2025-03-05, oauth-2025-04-20"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::api_key_merges_the_existing_beta_header(
|
||||
&[("anthropic-beta", " web-search-2025-03-05 ,")],
|
||||
Some(OAUTH_TOKEN),
|
||||
)]
|
||||
#[case::forwarded_bearer_unions_every_beta_header_casing(
|
||||
&[("anthropic-beta", "oauth-2025-04-20"), ("ANTHROPIC-BETA", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
fn oauth_beta_merges_into_existing_betas(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
no_env: Env,
|
||||
) {
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, no_env).unwrap(),
|
||||
headers(&[
|
||||
("authorization", OAUTH_BEARER),
|
||||
(
|
||||
"anthropic-beta",
|
||||
&betas(&[ANTHROPIC_OAUTH_BETA_HEADER, "web-search-2025-03-05"])
|
||||
),
|
||||
BROWSER_ACCESS,
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::x_api_key_over_the_deployment_key(&[("x-api-key", "caller-key")], Some("sk-other"))]
|
||||
#[case::uppercase_x_api_key(&[("X-API-KEY", "caller-key")], None)]
|
||||
#[case::non_oauth_bearer(&[("Authorization", "Bearer some-proxy-token")], None)]
|
||||
#[case::non_oauth_bearer_over_a_regular_api_key(
|
||||
&[("authorization", "Bearer sk-ant-api03-forwarded")],
|
||||
Some(REGULAR_KEY),
|
||||
)]
|
||||
#[case::oauth_token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)]
|
||||
#[case::oauth_token_behind_a_lowercase_bearer_scheme(
|
||||
&[("authorization", "bearer sk-ant-oat01-token")],
|
||||
None,
|
||||
)]
|
||||
fn forwarded_auth_header_is_kept_untouched(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
full_env: Env,
|
||||
) {
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, full_env).unwrap(),
|
||||
headers(forwarded)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key_param(Some("sk-param"), &[], ("x-api-key", "sk-param"))]
|
||||
#[case::api_key_param_over_env_key_and_auth_token(
|
||||
Some("sk-param"),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("x-api-key", "sk-param"),
|
||||
)]
|
||||
#[case::env_key_without_a_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
|
||||
#[case::env_key_when_the_param_is_empty(Some(""), &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
|
||||
#[case::env_key_when_the_param_is_whitespace(
|
||||
Some(" "),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env")],
|
||||
("x-api-key", "sk-env"),
|
||||
)]
|
||||
#[case::env_key_over_auth_token(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("x-api-key", "sk-env"),
|
||||
)]
|
||||
#[case::auth_token_as_a_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("authorization", "Bearer env-token"),
|
||||
)]
|
||||
#[case::auth_token_when_the_env_key_is_whitespace(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("authorization", "Bearer env-token"),
|
||||
)]
|
||||
#[case::oauth_env_key_as_a_plain_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-ant-oat01-env")],
|
||||
("authorization", "Bearer sk-ant-oat01-env"),
|
||||
)]
|
||||
fn credential_is_resolved_after_the_existing_headers(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] env: Env,
|
||||
#[case] expected: (&str, &str),
|
||||
) {
|
||||
let forwarded = [("anthropic-beta", "web-search-2025-03-05")];
|
||||
assert_eq!(
|
||||
authenticate_with(&forwarded, api_key, env).unwrap(),
|
||||
headers(&[forwarded[0], expected])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_credentials(&[], None, &[])]
|
||||
#[case::empty_api_key(&[], Some(""), &[])]
|
||||
#[case::whitespace_only_env_values(
|
||||
&[],
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")],
|
||||
)]
|
||||
#[case::unrelated_forwarded_headers(&[("anthropic-beta", "web-search-2025-03-05")], None, &[])]
|
||||
fn missing_credentials_are_an_auth_error(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] env: Env,
|
||||
) {
|
||||
assert!(matches!(
|
||||
authenticate_with(forwarded, api_key, env),
|
||||
Err(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
})
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_features(json!({}), &[])]
|
||||
#[case::output_format(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
|
||||
#[case::null_output_format(json!({"output_format": null}), &[])]
|
||||
#[case::output_config_format(
|
||||
json!({"output_config": {"format": {"type": "json_schema"}, "effort": "xhigh"}}),
|
||||
&[beta::STRUCTURED_OUTPUT]
|
||||
)]
|
||||
#[case::null_output_config_format(json!({"output_config": {"format": null}}), &[])]
|
||||
#[case::top_level_output_config_without_format(json!({"output_config": {"effort": "high"}}), &[])]
|
||||
#[case::fast_speed(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
|
||||
#[case::standard_speed(json!({"speed": "standard"}), &[])]
|
||||
#[case::compaction_param(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::empty_compaction_param(json!({"compaction": {}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::signed_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": "sig"}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[beta::COMPACT_2026_09_04]
|
||||
)]
|
||||
#[case::unsigned_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": ""}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[]
|
||||
)]
|
||||
#[case::advisor_tool(
|
||||
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}]}),
|
||||
&[beta::ADVISOR_TOOL_2026_03_01]
|
||||
)]
|
||||
#[case::no_tools(json!({"tools": []}), &[])]
|
||||
#[case::regex_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_regex_20251119"}]}),
|
||||
&[beta::ADVANCED_TOOL_USE_2025_11_20]
|
||||
)]
|
||||
#[case::bm25_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_bm25_20251119"}]}),
|
||||
&[beta::ADVANCED_TOOL_USE_2025_11_20]
|
||||
)]
|
||||
#[case::unrelated_server_tool(json!({"tools": [{"type": "web_search_20250305", "name": "web_search"}]}), &[])]
|
||||
#[case::only_compact_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
|
||||
&[beta::COMPACT_2026_01_12]
|
||||
)]
|
||||
#[case::only_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}]}}),
|
||||
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::compact_and_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_tool_uses_20250919"}]}}),
|
||||
&[beta::COMPACT_2026_01_12, beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::edit_without_a_type(json!({"context_management": {"edits": [{}]}}), &[beta::CONTEXT_MANAGEMENT_2025_06_27])]
|
||||
#[case::empty_edits(json!({"context_management": {"edits": []}}), &[])]
|
||||
#[case::context_management_without_edits(json!({"context_management": {}}), &[])]
|
||||
#[case::per_message_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
#[case::per_message_null_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": null}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
fn feature_betas_follow_the_request(#[case] fields: Value, #[case] expected: &[&str]) {
|
||||
assert_eq!(feature_betas(&request(fields)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_betas(&[("x-api-key", "k"), ("anthropic-version", "2023-06-01")], json!({}))]
|
||||
#[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], json!({}))]
|
||||
fn headers_without_any_beta_value_are_untouched(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
with_feature_betas(headers(input), &request(fields)),
|
||||
headers(input)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::feature_beta_is_appended(
|
||||
&[("x-api-key", "k")],
|
||||
json!({"speed": "fast"}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
)]
|
||||
#[case::existing_betas_are_normalized_without_features(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05, interleaved-thinking-2025-05-14 ,web-search-2025-03-05"), ("x-api-key", "k")],
|
||||
json!({}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "interleaved-thinking-2025-05-14,web-search-2025-03-05")],
|
||||
)]
|
||||
#[case::existing_advisor_beta_is_kept_without_an_advisor_tool(
|
||||
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
|
||||
json!({"tools": []}),
|
||||
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
|
||||
)]
|
||||
#[case::feature_already_sent_is_not_duplicated(
|
||||
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
json!({"speed": "fast"}),
|
||||
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
)]
|
||||
fn feature_betas_merge_into_the_headers(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
assert_eq!(
|
||||
with_feature_betas(headers(input), &request(fields)),
|
||||
headers(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn differently_cased_beta_header_is_replaced_by_one_sorted_header() {
|
||||
let merged = with_feature_betas(
|
||||
headers(&[("Anthropic-Beta", "interleaved-thinking-2025-05-14")]),
|
||||
&request(
|
||||
json!({"messages": [{"role": "system", "content": "env", "output_config": {"effort": "low"}}]}),
|
||||
),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::PER_TURN_CONTROL_2026_07_01
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_beta_header_casing_is_unioned_into_one_header() {
|
||||
let merged = with_feature_betas(
|
||||
headers(&[
|
||||
("anthropic-beta", "interleaved-thinking-2025-05-14"),
|
||||
("Anthropic-Beta", "web-search-2025-03-05"),
|
||||
]),
|
||||
&request(json!({"speed": "fast"})),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
beta::FAST_MODE_2026_02_01,
|
||||
"interleaved-thinking-2025-05-14",
|
||||
"web-search-2025-03-05"
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_client_betas_survive_alongside_the_added_one() {
|
||||
let client_betas = [
|
||||
"claude-code-20250219",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
"effort-2025-11-24",
|
||||
];
|
||||
let merged = with_feature_betas(
|
||||
headers(&[("anthropic-beta", &betas(&client_betas))]),
|
||||
&request(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
"claude-code-20250219",
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
"effort-2025-11-24",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_feature_merges_with_the_oauth_beta_sorted_and_last() {
|
||||
let oauth_headers = authenticate_with(&[], Some(OAUTH_TOKEN), &[]).unwrap();
|
||||
let all_features = request(json!({
|
||||
"compaction": {"enabled": true},
|
||||
"output_format": {"type": "json_schema"},
|
||||
"speed": "fast",
|
||||
"tools": [{"type": "advisor_20260301"}, {"type": "tool_search_tool_bm25_20251119"}],
|
||||
"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_thinking_20251015"}]},
|
||||
"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}],
|
||||
}));
|
||||
assert_eq!(
|
||||
with_feature_betas(oauth_headers, &all_features),
|
||||
headers(&[
|
||||
("authorization", OAUTH_BEARER),
|
||||
BROWSER_ACCESS,
|
||||
(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
beta::ADVANCED_TOOL_USE_2025_11_20,
|
||||
beta::ADVISOR_TOOL_2026_03_01,
|
||||
beta::COMPACT_2026_01_12,
|
||||
beta::COMPACT_2026_09_04,
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
beta::FAST_MODE_2026_02_01,
|
||||
ANTHROPIC_OAUTH_BETA_HEADER,
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
beta::STRUCTURED_OUTPUT,
|
||||
])
|
||||
),
|
||||
])
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,2 +1,5 @@
|
|||
pub mod handler;
|
||||
pub mod headers;
|
||||
pub mod streaming_iterator;
|
||||
pub mod thinking;
|
||||
pub mod transformation;
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,9 +1,28 @@
|
|||
use crate::base_llm::{
|
||||
anthropic_messages::transformation::BaseAnthropicMessagesConfig, chat::transformation::Error,
|
||||
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::{
|
||||
headers::{authenticate, with_feature_betas},
|
||||
thinking::{ThinkingBudgets, ThinkingContext, translate_thinking},
|
||||
};
|
||||
use crate::{
|
||||
anthropic::common_utils::{
|
||||
AnthropicModelCapabilities, has_advisor_tool, strip_advisor_blocks,
|
||||
strip_encrypted_reasoning_blocks,
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext,
|
||||
},
|
||||
chat::transformation::Error,
|
||||
},
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE";
|
||||
const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL";
|
||||
const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com";
|
||||
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
|
||||
|
|
@ -11,6 +30,26 @@ pub struct AnthropicMessagesConfig;
|
|||
|
||||
pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig;
|
||||
|
||||
impl MessagesTransformContext {
|
||||
pub fn new(capabilities: AnthropicModelCapabilities, drop_params: bool) -> Self {
|
||||
Self::with_lookup(capabilities, drop_params, &ProcessEnvironment)
|
||||
}
|
||||
|
||||
pub fn with_lookup(
|
||||
capabilities: AnthropicModelCapabilities,
|
||||
drop_params: bool,
|
||||
env: &impl Lookup,
|
||||
) -> Self {
|
||||
Self {
|
||||
thinking: ThinkingContext {
|
||||
capabilities,
|
||||
budgets: ThinkingBudgets::from_lookup(env),
|
||||
},
|
||||
drop_params,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
|
|
@ -21,6 +60,35 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
Ok(complete_anthropic_url(api_base, env_lookup))
|
||||
}
|
||||
|
||||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
if request.max_tokens.is_none() {
|
||||
return Err(Error::InvalidRequest(
|
||||
"max_tokens is required for Anthropic /v1/messages API".to_string(),
|
||||
));
|
||||
}
|
||||
let request = drop_unsupported_params(request, context)?;
|
||||
let request = translate_thinking(request, &context.thinking)?;
|
||||
let context_management = request
|
||||
.context_management
|
||||
.as_ref()
|
||||
.and_then(map_openai_context_management_to_anthropic)
|
||||
.or_else(|| request.context_management.clone());
|
||||
let messages = if has_advisor_tool(request.tools.as_deref()) {
|
||||
request.messages
|
||||
} else {
|
||||
strip_advisor_blocks(request.messages)
|
||||
};
|
||||
Ok(AnthropicMessagesRequest {
|
||||
messages: strip_encrypted_reasoning_blocks(messages),
|
||||
context_management,
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
|
|
@ -28,6 +96,113 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
) -> Result<String, Error> {
|
||||
resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from)
|
||||
}
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[
|
||||
ANTHROPIC_API_KEY_ENV,
|
||||
ANTHROPIC_AUTH_TOKEN_ENV,
|
||||
ANTHROPIC_API_BASE_ENV,
|
||||
ANTHROPIC_BASE_URL_ENV,
|
||||
]
|
||||
}
|
||||
|
||||
fn authenticate(
|
||||
&self,
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, Error> {
|
||||
authenticate(headers, api_key, env_lookup).map_err(Error::from)
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
with_feature_betas(headers, request)
|
||||
}
|
||||
}
|
||||
|
||||
fn unsupported_param(model: &str, param: &str, value: &str, hint: &str) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"{model} does not support {param}={value}. {hint}To drop unsupported params, set `litellm.drop_params = True`."
|
||||
))
|
||||
}
|
||||
|
||||
fn drop_unsupported_params(
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
let capabilities = &context.thinking.capabilities;
|
||||
let model = request.model.clone();
|
||||
let reject = |param: &str, value: String, hint: &str| -> Result<(), Error> {
|
||||
if context.drop_params {
|
||||
return Ok(());
|
||||
}
|
||||
Err(unsupported_param(&model, param, &value, hint))
|
||||
};
|
||||
let speed = match request.speed.as_deref() {
|
||||
Some(speed) if !capabilities.supports_speed => {
|
||||
reject("speed", format!("'{speed}'"), "")?;
|
||||
None
|
||||
}
|
||||
_ => request.speed.clone(),
|
||||
};
|
||||
if capabilities.supports_sampling_params {
|
||||
return Ok(AnthropicMessagesRequest { speed, ..request });
|
||||
}
|
||||
let temperature = match request.temperature {
|
||||
Some(temperature) if temperature != 1.0 => {
|
||||
reject(
|
||||
"temperature",
|
||||
json!(temperature).to_string(),
|
||||
"Only temperature=1 is supported. ",
|
||||
)?;
|
||||
None
|
||||
}
|
||||
temperature => temperature,
|
||||
};
|
||||
if let Some(top_p) = request.top_p {
|
||||
reject("top_p", json!(top_p).to_string(), "")?;
|
||||
}
|
||||
if let Some(top_k) = request.top_k {
|
||||
reject("top_k", json!(top_k).to_string(), "")?;
|
||||
}
|
||||
Ok(AnthropicMessagesRequest {
|
||||
speed,
|
||||
temperature,
|
||||
top_p: None,
|
||||
top_k: None,
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
pub fn map_openai_context_management_to_anthropic(context_management: &Value) -> Option<Value> {
|
||||
match context_management {
|
||||
Value::Object(edits) if edits.contains_key("edits") => Some(context_management.clone()),
|
||||
Value::Array(entries) => {
|
||||
let edits: Vec<Value> = entries
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.filter(|entry| entry.get("type").and_then(Value::as_str) == Some("compaction"))
|
||||
.map(|entry| {
|
||||
let trigger = entry.get("compact_threshold").and_then(Value::as_f64).map(
|
||||
|threshold| json!({"type": "input_tokens", "value": threshold as i64}),
|
||||
);
|
||||
let passthrough = entry
|
||||
.iter()
|
||||
.filter(|(key, _)| !matches!(key.as_str(), "type" | "compact_threshold"))
|
||||
.map(|(key, value)| (key.clone(), value.clone()));
|
||||
Value::Object(
|
||||
[("type".to_string(), json!("compact_20260112"))]
|
||||
.into_iter()
|
||||
.chain(trigger.map(|trigger| ("trigger".to_string(), trigger)))
|
||||
.chain(passthrough)
|
||||
.collect::<Map<String, Value>>(),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
(!edits.is_empty()).then(|| json!({"edits": edits}))
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
|
|
@ -64,70 +239,619 @@ pub fn resolve_anthropic_api_base(
|
|||
api_base: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
let env = |name: &str| env_lookup(name).filter(|value| !value.trim().is_empty());
|
||||
non_empty(api_base)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(ANTHROPIC_API_BASE_ENV).filter(|value| !value.trim().is_empty()))
|
||||
.or_else(|| env(ANTHROPIC_API_BASE_ENV))
|
||||
.or_else(|| env(ANTHROPIC_BASE_URL_ENV))
|
||||
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::process::Command;
|
||||
|
||||
use rstest::{fixture, rstest};
|
||||
|
||||
use super::*;
|
||||
use crate::anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta};
|
||||
|
||||
#[test]
|
||||
fn url_defaults_to_public_anthropic_endpoint() {
|
||||
type Env = &'static [(&'static str, &'static str)];
|
||||
|
||||
const BOTH_BASE_ENVS: Env = &[
|
||||
(ANTHROPIC_API_BASE_ENV, "https://api-base.example.com"),
|
||||
(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com"),
|
||||
];
|
||||
const API_KEY_ENV: Env = &[(ANTHROPIC_API_KEY_ENV, "sk-env")];
|
||||
const MISSING_API_KEY: &str =
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable";
|
||||
const LOW_BUDGET_ENV: &str = "DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET";
|
||||
const PROCESS_ENV_PROBE: &str = "LITELLM_MESSAGES_TRANSFORM_CONTEXT_PROBE";
|
||||
|
||||
fn merged(base: Value, fields: Value) -> Value {
|
||||
Value::Object(
|
||||
base.as_object()
|
||||
.unwrap()
|
||||
.clone()
|
||||
.into_iter()
|
||||
.chain(fields.as_object().unwrap().clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
fn body(fields: Value) -> Value {
|
||||
merged(
|
||||
json!({
|
||||
"model": "claude",
|
||||
"max_tokens": 1024,
|
||||
"messages": [{"role": "user", "content": "Hello"}]
|
||||
}),
|
||||
fields,
|
||||
)
|
||||
}
|
||||
|
||||
fn request(fields: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(body(fields)).unwrap()
|
||||
}
|
||||
|
||||
fn no_env(_: &str) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
fn env(vars: Env) -> impl Fn(&str) -> Option<String> {
|
||||
move |name| {
|
||||
vars.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn transform(
|
||||
fields: Value,
|
||||
capabilities: AnthropicModelCapabilities,
|
||||
drop_params: bool,
|
||||
) -> Result<Value, Error> {
|
||||
ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(
|
||||
request(fields),
|
||||
&MessagesTransformContext::with_lookup(capabilities, drop_params, &no_env),
|
||||
)
|
||||
.map(|transformed| serde_json::to_value(transformed).unwrap())
|
||||
}
|
||||
|
||||
fn invalid(message: &str) -> Result<Value, Error> {
|
||||
Err(Error::InvalidRequest(message.to_string()))
|
||||
}
|
||||
|
||||
fn advisor_history() -> Value {
|
||||
json!([
|
||||
{"role": "user", "content": "Build a worker pool."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "text", "text": "Let me consult the advisor."},
|
||||
{"type": "server_tool_use", "id": "srvtoolu_abc123", "name": "advisor", "input": {}},
|
||||
{"type": "advisor_tool_result", "tool_use_id": "srvtoolu_abc123", "content": {"type": "advisor_result", "text": "Use channels."}},
|
||||
{"type": "text", "text": "Here is the implementation."}
|
||||
]}
|
||||
])
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn unmapped() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities::default()
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn sampling_removed() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn fast_mode() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::alone(json!({"max_tokens": null}))]
|
||||
#[case::ahead_of_the_param_gate(json!({"max_tokens": null, "speed": "fast"}))]
|
||||
fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: AnthropicModelCapabilities) {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(None, &|_| None),
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
transform(fields, unmapped, false),
|
||||
invalid("max_tokens is required for Anthropic /v1/messages API")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::sampling_params_on_a_sampling_model(
|
||||
unmapped(),
|
||||
false,
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
|
||||
)]
|
||||
#[case::sampling_params_on_a_sampling_model_under_drop_params(
|
||||
unmapped(),
|
||||
true,
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
|
||||
)]
|
||||
#[case::unit_temperature_on_a_sampling_removed_model(
|
||||
sampling_removed(),
|
||||
false,
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
#[case::unit_temperature_on_a_sampling_removed_model_under_drop_params(
|
||||
sampling_removed(),
|
||||
true,
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
#[case::speed_on_a_fast_mode_model(fast_mode(), false, json!({"speed": "fast"}))]
|
||||
#[case::speed_on_a_fast_mode_model_under_drop_params(fast_mode(), true, json!({"speed": "fast"}))]
|
||||
#[case::native_context_management_edits(unmapped(), false, json!({"context_management": {"edits": [{
|
||||
"type": "clear_tool_uses_20250919",
|
||||
"trigger": {"type": "input_tokens", "value": 30000},
|
||||
"keep": {"type": "tool_uses", "value": 3},
|
||||
"clear_at_least": {"type": "input_tokens", "value": 5000},
|
||||
"exclude_tools": ["web_search"],
|
||||
"clear_tool_inputs": false
|
||||
}]}}))]
|
||||
#[case::first_party_billing_header_system_block(unmapped(), false, json!({"system": [
|
||||
{"type": "text", "text": "x-anthropic-billing-header: cc_version=1"},
|
||||
{"type": "text", "text": "real system prompt"}
|
||||
]}))]
|
||||
#[case::anthropic_signed_reasoning_history(unmapped(), false, json!({"messages": [
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": "EqQBCkYIAxgCIkA_anthropic_signed"},
|
||||
{"type": "redacted_thinking", "data": "EmwKAhgBEgy_anthropic_minted"},
|
||||
{"type": "text", "text": "The answer."}
|
||||
]}
|
||||
]}))]
|
||||
#[case::advisor_history_alongside_the_advisor_tool(unmapped(), false, json!({
|
||||
"messages": advisor_history(),
|
||||
"tools": [{"type": "advisor_20260301", "name": "advisor"}]
|
||||
}))]
|
||||
fn request_is_forwarded_unchanged(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] drop_params: bool,
|
||||
#[case] fields: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
transform(fields.clone(), capabilities, drop_params),
|
||||
Ok(body(fields))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::temperature(sampling_removed(), json!({"temperature": 0.3}), json!({}))]
|
||||
#[case::top_p(sampling_removed(), json!({"top_p": 0.9}), json!({}))]
|
||||
#[case::top_k(sampling_removed(), json!({"top_k": 40}), json!({}))]
|
||||
#[case::every_sampling_param_keeping_the_rest(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40, "stream": true}),
|
||||
json!({"stream": true})
|
||||
)]
|
||||
#[case::speed_on_a_sampling_model(
|
||||
unmapped(),
|
||||
json!({"speed": "fast", "temperature": 0.5}),
|
||||
json!({"temperature": 0.5})
|
||||
)]
|
||||
#[case::speed_on_a_sampling_removed_model(
|
||||
sampling_removed(),
|
||||
json!({"speed": "fast", "temperature": 1.0}),
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
fn removed_params_are_dropped_under_drop_params(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
assert_eq!(transform(fields, capabilities, true), Ok(body(expected)));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::temperature(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.3}),
|
||||
"claude does not support temperature=0.3. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::temperature_just_below_one(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.99}),
|
||||
"claude does not support temperature=0.99. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::whole_number_temperature_keeps_its_decimal(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 2.0}),
|
||||
"claude does not support temperature=2.0. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_p(
|
||||
sampling_removed(),
|
||||
json!({"top_p": 0.9}),
|
||||
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_k(
|
||||
sampling_removed(),
|
||||
json!({"top_k": 5}),
|
||||
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_k_next_to_unit_temperature(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 1.0, "top_k": 5}),
|
||||
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::temperature_ahead_of_top_k(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.5, "top_k": 5}),
|
||||
"claude does not support temperature=0.5. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_p_ahead_of_top_k(
|
||||
sampling_removed(),
|
||||
json!({"top_p": 0.9, "top_k": 5}),
|
||||
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::speed(
|
||||
unmapped(),
|
||||
json!({"speed": "fast"}),
|
||||
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::speed_ahead_of_sampling_params(
|
||||
sampling_removed(),
|
||||
json!({"speed": "fast", "temperature": 0.5}),
|
||||
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
fn removed_params_are_rejected_without_drop_params(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] message: &str,
|
||||
) {
|
||||
assert_eq!(transform(fields, capabilities, false), invalid(message));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::compaction_threshold(
|
||||
json!([{"type": "compaction", "compact_threshold": 200000}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]}))
|
||||
)]
|
||||
#[case::other_keys_pass_through(
|
||||
json!([{"type": "compaction", "compact_threshold": 150000, "instructions": "Focus on preserving code snippets"}]),
|
||||
Some(json!({"edits": [{
|
||||
"type": "compact_20260112",
|
||||
"trigger": {"type": "input_tokens", "value": 150000},
|
||||
"instructions": "Focus on preserving code snippets"
|
||||
}]}))
|
||||
)]
|
||||
#[case::float_threshold_is_truncated(
|
||||
json!([{"type": "compaction", "compact_threshold": 150000.9}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
|
||||
)]
|
||||
#[case::compaction_without_threshold(
|
||||
json!([{"type": "compaction"}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112"}]}))
|
||||
)]
|
||||
#[case::non_numeric_threshold_is_dropped(
|
||||
json!([{"type": "compaction", "compact_threshold": "150000"}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112"}]}))
|
||||
)]
|
||||
#[case::non_object_entries_are_skipped(
|
||||
json!([42, "compaction", null, [], {"type": "compaction", "compact_threshold": 1000}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}}]}))
|
||||
)]
|
||||
#[case::only_compaction_entries_are_mapped_in_order(
|
||||
json!([
|
||||
{"type": "compaction", "compact_threshold": 1000},
|
||||
{"type": "other", "compact_threshold": 5},
|
||||
{"type": "compaction", "instructions": "second"}
|
||||
]),
|
||||
Some(json!({"edits": [
|
||||
{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}},
|
||||
{"type": "compact_20260112", "instructions": "second"}
|
||||
]}))
|
||||
)]
|
||||
#[case::list_without_compaction(json!([{"type": "other"}]), None)]
|
||||
#[case::empty_list(json!([]), None)]
|
||||
#[case::anthropic_edits_pass_through(
|
||||
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
|
||||
)]
|
||||
#[case::object_without_edits(json!({"type": "compaction"}), None)]
|
||||
#[case::scalar(json!("compaction"), None)]
|
||||
fn openai_context_management_maps_to_anthropic_edits(
|
||||
#[case] context_management: Value,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(
|
||||
map_openai_context_management_to_anthropic(&context_management),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::openai_list_is_mapped(
|
||||
json!([{"type": "compaction", "compact_threshold": 200000}]),
|
||||
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]})
|
||||
)]
|
||||
#[case::unmappable_list_is_kept(json!([{"type": "other"}]), json!([{"type": "other"}]))]
|
||||
#[case::unmappable_object_is_kept(json!({"type": "other"}), json!({"type": "other"}))]
|
||||
fn context_management_reaches_the_wire(
|
||||
#[case] context_management: Value,
|
||||
#[case] expected: Value,
|
||||
unmapped: AnthropicModelCapabilities,
|
||||
) {
|
||||
assert_eq!(
|
||||
transform(
|
||||
json!({"context_management": context_management}),
|
||||
unmapped,
|
||||
false
|
||||
),
|
||||
Ok(body(json!({"context_management": expected})))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::without_tools(json!({}))]
|
||||
#[case::with_only_other_tools(json!({"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}]}))]
|
||||
fn advisor_history_is_stripped_without_the_advisor_tool(
|
||||
#[case] tools: Value,
|
||||
unmapped: AnthropicModelCapabilities,
|
||||
) {
|
||||
let stripped = json!([
|
||||
{"role": "user", "content": "Build a worker pool."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "text", "text": "Let me consult the advisor."},
|
||||
{"type": "text", "text": "Here is the implementation."}
|
||||
]}
|
||||
]);
|
||||
assert_eq!(
|
||||
transform(
|
||||
merged(tools.clone(), json!({"messages": advisor_history()})),
|
||||
unmapped,
|
||||
false
|
||||
),
|
||||
Ok(body(merged(tools, json!({"messages": stripped}))))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn bridge_minted_reasoning_is_stripped_from_the_wire(unmapped: AnthropicModelCapabilities) {
|
||||
let messages = json!([
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_1")},
|
||||
{"type": "redacted_thinking", "data": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_2")},
|
||||
{"type": "text", "text": "The answer."}
|
||||
]},
|
||||
{"role": "user", "content": "And the next one?"}
|
||||
]);
|
||||
assert_eq!(
|
||||
transform(json!({"messages": messages}), unmapped, false),
|
||||
Ok(body(json!({"messages": [
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "The answer."}]},
|
||||
{"role": "user", "content": "And the next one?"}
|
||||
]})))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_appends_messages_suffix_to_custom_base() {
|
||||
fn thinking_is_translated_with_the_context_budgets() {
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&env(&[(LOW_BUDGET_ENV, "2000")]),
|
||||
);
|
||||
let transformed = ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(
|
||||
request(json!({"max_tokens": 4096, "reasoning_effort": "low"})),
|
||||
&context,
|
||||
)
|
||||
.map(|transformed| serde_json::to_value(transformed).unwrap());
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some("https://proxy.internal"), &|_| None),
|
||||
"https://proxy.internal/v1/messages"
|
||||
transformed,
|
||||
Ok(body(json!({
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2000}
|
||||
})))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_leaves_complete_messages_endpoint_untouched() {
|
||||
fn new_reads_thinking_budgets_from_the_process_environment() {
|
||||
if std::env::var_os(PROCESS_ENV_PROBE).is_some() {
|
||||
assert_eq!(
|
||||
MessagesTransformContext::new(sampling_removed(), true),
|
||||
MessagesTransformContext {
|
||||
thinking: ThinkingContext {
|
||||
capabilities: sampling_removed(),
|
||||
budgets: ThinkingBudgets {
|
||||
low: 2000,
|
||||
..ThinkingBudgets::default()
|
||||
},
|
||||
},
|
||||
drop_params: true,
|
||||
}
|
||||
);
|
||||
return;
|
||||
}
|
||||
let (_, test_path) = concat!(
|
||||
module_path!(),
|
||||
"::new_reads_thinking_budgets_from_the_process_environment"
|
||||
)
|
||||
.split_once("::")
|
||||
.unwrap();
|
||||
let other_tiers = ["MINIMAL", "MEDIUM", "HIGH", "XHIGH", "MAX"]
|
||||
.map(|tier| format!("DEFAULT_REASONING_EFFORT_{tier}_THINKING_BUDGET"));
|
||||
let output = other_tiers
|
||||
.iter()
|
||||
.fold(
|
||||
Command::new(std::env::current_exe().unwrap()),
|
||||
|mut command, name| {
|
||||
command.env_remove(name);
|
||||
command
|
||||
},
|
||||
)
|
||||
.args([test_path, "--exact"])
|
||||
.env(PROCESS_ENV_PROBE, "1")
|
||||
.env(LOW_BUDGET_ENV, "2000")
|
||||
.output()
|
||||
.unwrap();
|
||||
let stdout = String::from_utf8_lossy(&output.stdout);
|
||||
assert!(
|
||||
output.status.success() && stdout.contains("1 passed"),
|
||||
"{stdout}{}",
|
||||
String::from_utf8_lossy(&output.stderr)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint_by_default(None, &[], "https://api.anthropic.com")]
|
||||
#[case::explicit_api_base_beats_env(
|
||||
Some("https://explicit.example.com"),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::explicit_api_base_is_trimmed(
|
||||
Some(" https://explicit.example.com "),
|
||||
&[],
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_falls_back_to_env(
|
||||
Some(" "),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://api-base.example.com"
|
||||
)]
|
||||
#[case::api_base_env_beats_base_url_env(None, BOTH_BASE_ENVS, "https://api-base.example.com")]
|
||||
#[case::base_url_env_without_api_base_env(
|
||||
None,
|
||||
&[(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_env_falls_back_to_base_url_env(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, " \t "), (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_envs_fall_back_to_public_endpoint(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, ""), (ANTHROPIC_BASE_URL_ENV, " ")],
|
||||
"https://api.anthropic.com"
|
||||
)]
|
||||
fn api_base_resolution(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(resolve_anthropic_api_base(api_base, &env(vars)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint(None, &[], "https://api.anthropic.com/v1/messages")]
|
||||
#[case::base_url_env(
|
||||
None,
|
||||
&[(ANTHROPIC_BASE_URL_ENV, "https://custom.example.com")],
|
||||
"https://custom.example.com/v1/messages"
|
||||
)]
|
||||
#[case::custom_base(Some("https://proxy.internal"), &[], "https://proxy.internal/v1/messages")]
|
||||
#[case::trailing_slash(Some("https://proxy.internal/"), &[], "https://proxy.internal/v1/messages")]
|
||||
#[case::complete_endpoint(
|
||||
Some("https://proxy.internal/v1/messages"),
|
||||
&[],
|
||||
"https://proxy.internal/v1/messages"
|
||||
)]
|
||||
#[case::complete_endpoint_with_trailing_slash(
|
||||
Some("https://proxy.internal/v1/messages/"),
|
||||
&[],
|
||||
"https://proxy.internal/v1/messages"
|
||||
)]
|
||||
fn complete_url_ends_in_the_messages_path(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some("https://proxy.internal/v1/messages"), &|_| None),
|
||||
"https://proxy.internal/v1/messages"
|
||||
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(api_base, "claude", &env(vars)),
|
||||
Ok(expected.to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::param_beats_env(Some("sk-param"), API_KEY_ENV, Ok("sk-param"))]
|
||||
#[case::param_is_trimmed(Some(" sk-param "), &[], Ok("sk-param"))]
|
||||
#[case::blank_param_falls_back_to_env(Some(" "), API_KEY_ENV, Ok("sk-env"))]
|
||||
#[case::env_without_param(None, API_KEY_ENV, Ok("sk-env"))]
|
||||
#[case::blank_env_is_missing(None, &[(ANTHROPIC_API_KEY_ENV, " ")], Err(MISSING_API_KEY))]
|
||||
#[case::nothing_is_missing(None, &[], Err(MISSING_API_KEY))]
|
||||
fn api_key_resolution(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: Result<&str, &str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(api_key, &env(vars)).map_err(|error| error.to_string()),
|
||||
expected.map(str::to_string).map_err(str::to_string)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_falls_back_to_env_base() {
|
||||
let with_env = |key: &str| {
|
||||
(key == ANTHROPIC_API_BASE_ENV).then(|| "https://env.anthropic".to_string())
|
||||
};
|
||||
fn config_reports_a_missing_key_as_an_auth_error() {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some(" "), &with_env),
|
||||
"https://env.anthropic/v1/messages"
|
||||
ANTHROPIC_MESSAGES_CONFIG.resolve_api_key(None, &no_env),
|
||||
Err(Error::Auth(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_key_prefers_param_then_env_then_errors() {
|
||||
fn config_authenticates_with_the_anthropic_auth_token() {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(Some("sk-param"), &|_| None).unwrap(),
|
||||
"sk-param"
|
||||
ANTHROPIC_MESSAGES_CONFIG.authenticate(
|
||||
vec![],
|
||||
None,
|
||||
&env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")])
|
||||
),
|
||||
Ok(headers(&[("authorization", "Bearer auth-token")]))
|
||||
);
|
||||
let with_env = |key: &str| (key == ANTHROPIC_API_KEY_ENV).then(|| "sk-env".to_string());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_requests_the_betas_the_request_features_need() {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(Some(" "), &with_env).unwrap(),
|
||||
"sk-env"
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(None, &|_| None)
|
||||
.expect_err("missing key")
|
||||
.to_string(),
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
|
||||
ANTHROPIC_MESSAGES_CONFIG.request_headers(
|
||||
headers(&[("x-api-key", "sk")]),
|
||||
&request(json!({"speed": "fast"}))
|
||||
),
|
||||
headers(&[
|
||||
("x-api-key", "sk"),
|
||||
("anthropic-beta", beta::FAST_MODE_2026_02_01)
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None, None)]
|
||||
#[case::blank(Some(" \t "), None)]
|
||||
#[case::padded(Some(" value "), Some("value"))]
|
||||
fn non_empty_trims_and_drops_blank_values(
|
||||
#[case] value: Option<&str>,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
assert_eq!(non_empty(value), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_strategy_and_default_headers_match_anthropic() {
|
||||
assert_eq!(
|
||||
|
|
@ -142,4 +866,26 @@ mod tests {
|
|||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn secret_names_cover_every_credential_and_base_lookup() {
|
||||
let requested = std::cell::RefCell::new(Vec::<String>::new());
|
||||
let record = |name: &str| -> Option<String> {
|
||||
requested.borrow_mut().push(name.to_string());
|
||||
None
|
||||
};
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
|
||||
let requested = requested.into_inner();
|
||||
assert!(!requested.is_empty());
|
||||
let undeclared: Vec<&String> = requested
|
||||
.iter()
|
||||
.filter(|name| {
|
||||
!ANTHROPIC_MESSAGES_CONFIG
|
||||
.secret_names()
|
||||
.contains(&name.as_str())
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(undeclared, Vec::<&String>::new());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
pub mod batches;
|
||||
pub mod chat;
|
||||
pub mod common_utils;
|
||||
pub mod count_tokens;
|
||||
pub mod experimental_pass_through;
|
||||
|
||||
|
|
|
|||
|
|
@ -4,14 +4,15 @@ use litellm_types::llms::anthropic_messages::{
|
|||
},
|
||||
anthropic_response::AnthropicMessagesResponse,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::transformation::{
|
||||
ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty,
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::{BaseAnthropicMessagesConfig, MessagesAuthStrategy},
|
||||
anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, Headers, MessagesAuthStrategy, MessagesTransformContext,
|
||||
},
|
||||
chat::transformation::Error,
|
||||
},
|
||||
};
|
||||
|
|
@ -21,7 +22,6 @@ const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE";
|
|||
const ANTHROPIC_PATH_SEGMENT: &str = "/anthropic";
|
||||
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
const SYSTEM_ROLE: &str = "system";
|
||||
const TEXT_BLOCK_TYPE: &str = "text";
|
||||
|
||||
pub struct AzureAnthropicMessagesConfig {
|
||||
anthropic: AnthropicMessagesConfig,
|
||||
|
|
@ -45,6 +45,7 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
let mut request = fold_system_role_messages(request);
|
||||
if let Some(system) = request.system.as_mut() {
|
||||
|
|
@ -54,7 +55,8 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
.messages
|
||||
.iter_mut()
|
||||
.for_each(strip_scope_from_message);
|
||||
self.anthropic.transform_anthropic_messages_request(request)
|
||||
self.anthropic
|
||||
.transform_anthropic_messages_request(request, context)
|
||||
}
|
||||
|
||||
fn transform_anthropic_messages_response(
|
||||
|
|
@ -74,6 +76,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
resolve_azure_api_key(api_key, env_lookup)
|
||||
}
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[AZURE_API_KEY_ENV, AZURE_API_BASE_ENV]
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
self.anthropic.auth_strategy()
|
||||
}
|
||||
|
|
@ -85,6 +91,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
self.anthropic.default_headers()
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
self.anthropic.request_headers(headers, request)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_azure_api_key(
|
||||
|
|
@ -143,17 +153,7 @@ fn strip_scope_from_message(message: &mut AnthropicMessage) {
|
|||
}
|
||||
|
||||
fn text_content_block(text: String) -> ContentBlock {
|
||||
let extra = Map::from_iter([
|
||||
(
|
||||
"type".to_string(),
|
||||
Value::String(TEXT_BLOCK_TYPE.to_string()),
|
||||
),
|
||||
("text".to_string(), Value::String(text)),
|
||||
]);
|
||||
ContentBlock {
|
||||
cache_control: None,
|
||||
extra,
|
||||
}
|
||||
ContentBlock::text(text)
|
||||
}
|
||||
|
||||
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
|
||||
|
|
@ -202,6 +202,7 @@ mod tests {
|
|||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
|
||||
fn request_from(value: serde_json::Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(value).expect("valid request")
|
||||
|
|
@ -346,7 +347,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -373,10 +374,13 @@ mod tests {
|
|||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}));
|
||||
let once = AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms");
|
||||
let twice = AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(once.clone())
|
||||
.transform_anthropic_messages_request(
|
||||
once.clone(),
|
||||
&MessagesTransformContext::default(),
|
||||
)
|
||||
.expect("request transforms");
|
||||
assert_eq!(once, twice);
|
||||
assert_eq!(to_value(once)["system"], json!("plain string system"));
|
||||
|
|
@ -408,9 +412,21 @@ mod tests {
|
|||
"inference_geo": "us",
|
||||
"litellm_metadata": {"trace": "abc"}
|
||||
});
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_legacy_thinking: true,
|
||||
supports_output_config: true,
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&|_: &str| None,
|
||||
);
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request_from(body.clone()))
|
||||
.transform_anthropic_messages_request(request_from(body.clone()), &context)
|
||||
.expect("request transforms"),
|
||||
);
|
||||
assert_eq!(transformed, body);
|
||||
|
|
@ -430,7 +446,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -460,7 +476,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -485,9 +501,21 @@ mod tests {
|
|||
{"role": "assistant", "content": "hello"}
|
||||
]
|
||||
});
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_legacy_thinking: true,
|
||||
supports_output_config: true,
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&|_: &str| None,
|
||||
);
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request_from(body.clone()))
|
||||
.transform_anthropic_messages_request(request_from(body.clone()), &context)
|
||||
.expect("request transforms"),
|
||||
);
|
||||
assert_eq!(transformed, body);
|
||||
|
|
@ -500,6 +528,57 @@ mod tests {
|
|||
assert!(err.is_data());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::compact_context_management_edit(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
|
||||
&[],
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "compact-2026-01-12")]
|
||||
)]
|
||||
#[case::forwarded_beta_merged_with_structured_output(
|
||||
json!({"output_config": {"format": {"type": "json_schema"}}}),
|
||||
&[("anthropic-beta", "web-search-2025-03-05")],
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "structured-outputs-2025-11-13,web-search-2025-03-05")]
|
||||
)]
|
||||
#[case::no_feature_needs_a_beta(json!({}), &[], &[("x-api-key", "k")])]
|
||||
fn request_headers_carry_the_anthropic_feature_betas(
|
||||
#[case] fields: serde_json::Value,
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let pairs = |pairs: &[(&str, &str)]| -> Vec<(String, String)> {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
};
|
||||
let serde_json::Value::Object(fields) = fields else {
|
||||
panic!("case fields are an object")
|
||||
};
|
||||
let request = request_from(serde_json::Value::Object(
|
||||
[
|
||||
("model".to_string(), json!("claude-sonnet")),
|
||||
("max_tokens".to_string(), json!(16)),
|
||||
(
|
||||
"messages".to_string(),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
),
|
||||
]
|
||||
.into_iter()
|
||||
.chain(fields)
|
||||
.collect(),
|
||||
));
|
||||
assert_eq!(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG.request_headers(
|
||||
pairs(&[("x-api-key", "k")])
|
||||
.into_iter()
|
||||
.chain(pairs(forwarded))
|
||||
.collect(),
|
||||
&request
|
||||
),
|
||||
pairs(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_response_passes_through() {
|
||||
let response: AnthropicMessagesResponse = serde_json::from_value(json!({
|
||||
|
|
@ -521,4 +600,26 @@ mod tests {
|
|||
assert_eq!(value["stop_sequence"], json!(null));
|
||||
assert_eq!(value["content"][0]["text"], json!("hello"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn secret_names_cover_every_credential_and_base_lookup() {
|
||||
let requested = std::cell::RefCell::new(Vec::<String>::new());
|
||||
let record = |name: &str| -> Option<String> {
|
||||
requested.borrow_mut().push(name.to_string());
|
||||
None
|
||||
};
|
||||
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &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
|
||||
.iter()
|
||||
.filter(|name| {
|
||||
!AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.secret_names()
|
||||
.contains(&name.as_str())
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(undeclared, Vec::<&String>::new());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,14 @@
|
|||
use litellm_http::request::{has_bearer_auth, has_header};
|
||||
use litellm_types::llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
};
|
||||
|
||||
use crate::base_llm::chat::transformation::Error;
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::thinking::ThinkingContext,
|
||||
base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
pub type Headers = Vec<(String, String)>;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum MessagesAuthStrategy {
|
||||
|
|
@ -19,6 +25,12 @@ impl MessagesAuthStrategy {
|
|||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
||||
pub struct MessagesTransformContext {
|
||||
pub thinking: ThinkingContext,
|
||||
pub drop_params: bool,
|
||||
}
|
||||
|
||||
pub trait BaseAnthropicMessagesConfig: Sync {
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
|
|
@ -30,6 +42,7 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
_context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
Ok(request)
|
||||
}
|
||||
|
|
@ -48,6 +61,8 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error>;
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str];
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
MessagesAuthStrategy::Header("x-api-key")
|
||||
}
|
||||
|
|
@ -56,10 +71,225 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
false
|
||||
}
|
||||
|
||||
fn authenticate(
|
||||
&self,
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, Error> {
|
||||
let strategy = self.auth_strategy();
|
||||
if has_header(&headers, strategy.header_name())
|
||||
|| (self.accepts_bearer_auth() && has_bearer_auth(&headers))
|
||||
{
|
||||
return Ok(headers);
|
||||
}
|
||||
let api_key = self.resolve_api_key(api_key, env_lookup)?;
|
||||
let auth_header = match strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
Ok(headers.into_iter().chain([auth_header]).collect())
|
||||
}
|
||||
|
||||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
&[
|
||||
("anthropic-version", "2023-06-01"),
|
||||
("content-type", "application/json"),
|
||||
]
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, _request: &AnthropicMessagesRequest) -> Headers {
|
||||
headers
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
const X_API_KEY: MessagesAuthStrategy = MessagesAuthStrategy::Header("x-api-key");
|
||||
|
||||
struct StubConfig {
|
||||
strategy: MessagesAuthStrategy,
|
||||
accepts_bearer: bool,
|
||||
}
|
||||
|
||||
impl BaseAnthropicMessagesConfig for StubConfig {
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[]
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
_api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
api_key
|
||||
.map(str::to_string)
|
||||
.ok_or(Error::MissingField("api_key"))
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
self.strategy
|
||||
}
|
||||
|
||||
fn accepts_bearer_auth(&self) -> bool {
|
||||
self.accepts_bearer
|
||||
}
|
||||
}
|
||||
|
||||
struct DefaultsConfig;
|
||||
|
||||
impl BaseAnthropicMessagesConfig for DefaultsConfig {
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[]
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
_api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
api_key
|
||||
.map(str::to_string)
|
||||
.ok_or(Error::MissingField("api_key"))
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_config_adds_its_key_next_to_a_forwarded_bearer() {
|
||||
assert_eq!(
|
||||
DefaultsConfig.authenticate(
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
Some("sk"),
|
||||
&|_| None
|
||||
),
|
||||
Ok(headers(&[
|
||||
("authorization", "Bearer forwarded"),
|
||||
("x-api-key", "sk")
|
||||
]))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_request_headers_are_the_given_headers() {
|
||||
let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "claude",
|
||||
"max_tokens": 16,
|
||||
"speed": "fast",
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
DefaultsConfig.request_headers(headers(&[("x-api-key", "sk")]), &request),
|
||||
headers(&[("x-api-key", "sk")])
|
||||
);
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::own_header_is_kept(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("x-api-key", "forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("x-api-key", "forwarded")]))
|
||||
)]
|
||||
#[case::own_header_in_any_casing_is_kept(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("X-Api-Key", "forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("X-Api-Key", "forwarded")]))
|
||||
)]
|
||||
#[case::accepted_bearer_is_kept(
|
||||
X_API_KEY,
|
||||
true,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("authorization", "Bearer forwarded")]))
|
||||
)]
|
||||
#[case::bearer_the_provider_does_not_accept_gets_the_key_too(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer forwarded"), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::blank_bearer_gets_the_key(
|
||||
X_API_KEY,
|
||||
true,
|
||||
headers(&[("authorization", "Bearer ")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer "), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::key_goes_in_the_provider_header(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("content-type", "application/json")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("content-type", "application/json"), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::key_goes_in_a_bearer(
|
||||
MessagesAuthStrategy::Bearer,
|
||||
false,
|
||||
headers(&[]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer sk")]))
|
||||
)]
|
||||
#[case::bearer_strategy_keeps_a_forwarded_authorization(
|
||||
MessagesAuthStrategy::Bearer,
|
||||
false,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("authorization", "Bearer forwarded")]))
|
||||
)]
|
||||
#[case::missing_key_is_an_error(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[]),
|
||||
None,
|
||||
Err(Error::MissingField("api_key"))
|
||||
)]
|
||||
fn default_authenticate_applies_the_key_unless_a_credential_is_forwarded(
|
||||
#[case] strategy: MessagesAuthStrategy,
|
||||
#[case] accepts_bearer: bool,
|
||||
#[case] forwarded: Headers,
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] expected: Result<Headers, Error>,
|
||||
) {
|
||||
let config = StubConfig {
|
||||
strategy,
|
||||
accepts_bearer,
|
||||
};
|
||||
assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,9 +2,11 @@ use bytes::Bytes;
|
|||
use litellm_core::messages::{
|
||||
Error,
|
||||
route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput},
|
||||
types::MessagesShaping,
|
||||
};
|
||||
use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use pyo3::{
|
||||
exceptions::{PyException, PyValueError},
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
|
|
@ -18,9 +20,10 @@ use crate::{
|
|||
marshal::{optional_timeout, python_timeout_seconds},
|
||||
};
|
||||
|
||||
/// The Anthropic Messages body fields a caller may pass besides `model` and `messages`,
|
||||
/// as `AnthropicMessagesRequestOptionalParams` declares them.
|
||||
const BODY_FIELDS: [&str; 20] = [
|
||||
const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host";
|
||||
const REQUEST_ERROR_MARKER: &str = "messages_request_error";
|
||||
|
||||
const BODY_FIELDS: [&str; 22] = [
|
||||
"max_tokens",
|
||||
"metadata",
|
||||
"stop_sequences",
|
||||
|
|
@ -35,14 +38,46 @@ const BODY_FIELDS: [&str; 20] = [
|
|||
"top_p",
|
||||
"mcp_servers",
|
||||
"context_management",
|
||||
"compaction",
|
||||
"container",
|
||||
"output_format",
|
||||
"speed",
|
||||
"output_config",
|
||||
"cache_control",
|
||||
"reasoning_effort",
|
||||
"safeguards",
|
||||
];
|
||||
|
||||
fn merge_headers(
|
||||
forwarded: Option<Map<String, Value>>,
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
) -> Option<Map<String, Value>> {
|
||||
let merged: Map<String, Value> = forwarded
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(extra_headers.into_iter().flatten())
|
||||
.collect();
|
||||
(!merged.is_empty()).then_some(merged)
|
||||
}
|
||||
|
||||
fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
||||
match error {
|
||||
Error::Transport(TransportError::Http { status, body }) => {
|
||||
let error = RustUpstreamError::new_err((status, body));
|
||||
error
|
||||
.value(py)
|
||||
.setattr("headers", Vec::<(String, String)>::new())?;
|
||||
Ok(error)
|
||||
}
|
||||
Error::InvalidRequest(message) => {
|
||||
let error = PyValueError::new_err(message);
|
||||
error.value(py).setattr(REQUEST_ERROR_MARKER, true)?;
|
||||
Ok(error)
|
||||
}
|
||||
other => Ok(messages_error_to_pyerr(other)),
|
||||
}
|
||||
}
|
||||
|
||||
/// The Python side of the Messages route: projects the prepared arguments and builds the
|
||||
/// public response, chunks and exceptions.
|
||||
pub(super) struct MessagesRouteHost {
|
||||
|
|
@ -84,19 +119,65 @@ impl MessagesRouteHost {
|
|||
.map(|value| python_timeout_seconds(py, value.unbind()))
|
||||
.transpose()?
|
||||
.flatten();
|
||||
let custom_llm_provider = string("custom_llm_provider")?;
|
||||
let shaping = self.shaping(py, &model, custom_llm_provider.as_deref(), arguments)?;
|
||||
Ok(MessagesCall {
|
||||
model,
|
||||
body,
|
||||
api_key: string("api_key")?,
|
||||
api_base: string("api_base")?,
|
||||
custom_llm_provider: string("custom_llm_provider")?,
|
||||
extra_headers: argument("extra_headers")?
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()?,
|
||||
extra_headers: self.merged_headers(py, arguments)?,
|
||||
provider_specific_header: self.provider_specific_header(py, arguments)?,
|
||||
custom_llm_provider,
|
||||
timeout: optional_timeout(timeout),
|
||||
shaping,
|
||||
})
|
||||
}
|
||||
|
||||
fn merged_headers(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Option<Map<String, Value>>> {
|
||||
let request = self.request.bind(py);
|
||||
let mapping = |name: &str| -> PyResult<Option<Map<String, Value>>> {
|
||||
lookup(arguments, request, name)?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()
|
||||
};
|
||||
Ok(merge_headers(
|
||||
mapping("headers")?,
|
||||
mapping("extra_headers")?,
|
||||
))
|
||||
}
|
||||
|
||||
fn provider_specific_header(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Option<ProviderSpecificHeaders>> {
|
||||
lookup(arguments, self.request.bind(py), "provider_specific_header")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn shaping(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<MessagesShaping> {
|
||||
let projected = py.import(ROUTE_HOST_MODULE)?.getattr("shaping")?.call1((
|
||||
model,
|
||||
custom_llm_provider,
|
||||
arguments,
|
||||
))?;
|
||||
from_py(&projected)
|
||||
}
|
||||
|
||||
fn provider(&self, py: Python<'_>) -> String {
|
||||
self.request
|
||||
.bind(py)
|
||||
|
|
@ -112,7 +193,7 @@ impl MessagesRouteHost {
|
|||
return error;
|
||||
}
|
||||
let mapped = py
|
||||
.import("litellm.rust_bridge.messages.route_host")
|
||||
.import(ROUTE_HOST_MODULE)
|
||||
.and_then(|module| module.getattr("map_failure"))
|
||||
.and_then(|map| map.call1((error.value(py), self.request.bind(py), self.provider(py))))
|
||||
.and_then(|mapped| {
|
||||
|
|
@ -148,7 +229,7 @@ impl RouteHost for MessagesRouteHost {
|
|||
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {
|
||||
match response {
|
||||
MessagesOutput::Message(message) => py
|
||||
.import("litellm.rust_bridge.messages.route_host")?
|
||||
.import(ROUTE_HOST_MODULE)?
|
||||
.getattr("response")?
|
||||
.call1((to_py(py, message.as_ref())?,))
|
||||
.map(Bound::unbind),
|
||||
|
|
@ -161,17 +242,12 @@ impl RouteHost for MessagesRouteHost {
|
|||
}
|
||||
|
||||
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
||||
let native = match error {
|
||||
Error::Transport(TransportError::Http { status, body }) => {
|
||||
let error = RustUpstreamError::new_err((status, body));
|
||||
error
|
||||
.value(py)
|
||||
.setattr("headers", Vec::<(String, String)>::new())?;
|
||||
error
|
||||
}
|
||||
other => messages_error_to_pyerr(other),
|
||||
};
|
||||
Ok(self.map_failure(py, native))
|
||||
if let Error::Secret(source) = &error
|
||||
&& let Some(original) = crate::secrets::python_error(py, source.source_error())
|
||||
{
|
||||
return Ok(original);
|
||||
}
|
||||
Ok(self.map_failure(py, native_error(py, error)?))
|
||||
}
|
||||
|
||||
fn host_error(error: &PyErr) -> Error {
|
||||
|
|
@ -184,3 +260,62 @@ impl RouteHost for MessagesRouteHost {
|
|||
visit.call(&self.request)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn map(value: Value) -> Map<String, Value> {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::extra_over_forwarded(
|
||||
Some(json!({"X-Priority": "forwarded", "X-Forwarded-Only": "keep"})),
|
||||
Some(json!({"X-Priority": "extra", "X-Extra-Only": "also-keep"})),
|
||||
Some(json!({"X-Priority": "extra", "X-Forwarded-Only": "keep", "X-Extra-Only": "also-keep"})),
|
||||
)]
|
||||
#[case::only_forwarded(Some(json!({"X-Forwarded": "yes"})), None, Some(json!({"X-Forwarded": "yes"})))]
|
||||
#[case::only_extra_headers(
|
||||
None,
|
||||
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
|
||||
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
|
||||
)]
|
||||
#[case::nothing(None, Some(json!({})), None)]
|
||||
fn headers_merge_forwarded_then_extra(
|
||||
#[case] forwarded: Option<Value>,
|
||||
#[case] extra_headers: Option<Value>,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(
|
||||
merge_headers(forwarded.map(map), extra_headers.map(map)),
|
||||
expected.map(map)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)]
|
||||
#[case::unresolvable_provider(Error::InvalidProvider("openai".into()), false)]
|
||||
#[case::upstream_failure(
|
||||
Error::Transport(TransportError::Http { status: 400, body: "bad".into() }),
|
||||
false,
|
||||
)]
|
||||
fn only_request_rejections_carry_the_request_error_marker(
|
||||
#[case] error: Error,
|
||||
#[case] marked: bool,
|
||||
) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let native = native_error(py, error).unwrap();
|
||||
let marker = native
|
||||
.value(py)
|
||||
.getattr_opt(REQUEST_ERROR_MARKER)
|
||||
.unwrap()
|
||||
.map(|value| value.extract::<bool>().unwrap());
|
||||
assert_eq!(marker.unwrap_or(false), marked);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -39,11 +39,12 @@ fn run_messages(
|
|||
"the Rust Messages route does not serve this provider",
|
||||
));
|
||||
}
|
||||
let secrets = crate::secrets::source(py)?;
|
||||
run_legacy_call(
|
||||
py,
|
||||
SURFACE,
|
||||
PublicCall::capture(&request, &args, &kwargs)?,
|
||||
crate::logger::LoggedMachine::new(messages_machine()),
|
||||
crate::logger::LoggedMachine::new(messages_machine(secrets)),
|
||||
MessagesRouteHost::new(request.unbind()),
|
||||
asynchronous,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,3 +8,6 @@ repository.workspace = true
|
|||
[dependencies]
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -17,12 +17,48 @@ pub enum MessageContent {
|
|||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ContentBlock {
|
||||
#[serde(rename = "type", default, skip_serializing_if = "Option::is_none")]
|
||||
pub block_type: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub text: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub thinking: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub signature: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub data: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_use_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub content: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub provider_specific_fields: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_control: Option<CacheControl>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
impl ContentBlock {
|
||||
pub fn text(text: impl Into<String>) -> Self {
|
||||
Self {
|
||||
block_type: Some("text".to_string()),
|
||||
text: Some(text.into()),
|
||||
..Self::default()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_type(&self, block_type: &str) -> bool {
|
||||
self.block_type.as_deref() == Some(block_type)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CacheControl {
|
||||
#[serde(rename = "type", skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -85,6 +121,126 @@ pub struct AnthropicMessagesRequest {
|
|||
pub speed: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub inference_geo: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning_effort: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub compaction: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
impl AnthropicMessage {
|
||||
pub fn blocks(&self) -> &[ContentBlock] {
|
||||
match &self.content {
|
||||
MessageContent::Blocks(blocks) => blocks,
|
||||
MessageContent::Text(_) => &[],
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_blocks(self, blocks: Vec<ContentBlock>) -> Self {
|
||||
Self {
|
||||
content: MessageContent::Blocks(blocks),
|
||||
..self
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn round_trip<T: serde::de::DeserializeOwned + Serialize>(value: &Value) -> Value {
|
||||
let parsed: T = serde_json::from_value(value.clone()).unwrap();
|
||||
serde_json::to_value(parsed).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::text(json!({"type": "text", "text": "hi"}))]
|
||||
#[case::text_with_citations_and_cache_control(json!({
|
||||
"type": "text",
|
||||
"text": "hi",
|
||||
"citations": [{"type": "char_location", "cited_text": "x"}],
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h", "scope": "global", "future": 1}
|
||||
}))]
|
||||
#[case::image(json!({"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "AA=="}}))]
|
||||
#[case::thinking(json!({"type": "thinking", "thinking": "hmm", "signature": "sig"}))]
|
||||
#[case::redacted_thinking(json!({"type": "redacted_thinking", "data": "opaque"}))]
|
||||
#[case::tool_use(json!({"type": "tool_use", "id": "toolu_1", "name": "f", "input": {"q": [1, null]}}))]
|
||||
#[case::tool_result_with_text(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": "ok", "is_error": false}))]
|
||||
#[case::tool_result_with_blocks(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": [{"type": "text", "text": "ok"}]}))]
|
||||
#[case::web_search_result_with_nulls(json!({
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": "srvtoolu_1",
|
||||
"content": [{"type": "web_search_result", "url": "u", "page_age": null, "encrypted_content": ""}]
|
||||
}))]
|
||||
#[case::provider_specific_fields(json!({"type": "tool_use", "id": "t", "name": "f", "input": {}, "provider_specific_fields": {"x": 1}}))]
|
||||
#[case::untyped(json!({"unknown": {"nested": true}}))]
|
||||
fn content_block_round_trips_unchanged(#[case] block: Value) {
|
||||
assert_eq!(round_trip::<ContentBlock>(&block), block);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_constructor_serializes_as_a_text_block() {
|
||||
assert_eq!(
|
||||
serde_json::to_value(ContentBlock::text("hello")).unwrap(),
|
||||
json!({"type": "text", "text": "hello"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::same_type(json!({"type": "tool_use"}), "tool_use", true)]
|
||||
#[case::other_type(json!({"type": "tool_result"}), "tool_use", false)]
|
||||
#[case::prefix_of_type(json!({"type": "tool_use"}), "tool", false)]
|
||||
#[case::no_type(json!({"text": "x"}), "text", false)]
|
||||
fn is_type_matches_the_exact_block_type(
|
||||
#[case] block: Value,
|
||||
#[case] block_type: &str,
|
||||
#[case] expected: bool,
|
||||
) {
|
||||
let block: ContentBlock = serde_json::from_value(block).unwrap();
|
||||
assert_eq!(block.is_type(block_type), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::string_content(json!({"role": "user", "content": "hi"}), vec![])]
|
||||
#[case::block_content(
|
||||
json!({"role": "user", "content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}]}),
|
||||
vec![ContentBlock::text("a"), ContentBlock::text("b")],
|
||||
)]
|
||||
fn message_blocks_list_only_block_content(
|
||||
#[case] message: Value,
|
||||
#[case] expected: Vec<ContentBlock>,
|
||||
) {
|
||||
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
|
||||
assert_eq!(message.blocks(), expected.as_slice());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::replaces_string_content(json!({"role": "assistant", "content": "old", "name": "kept"}))]
|
||||
#[case::replaces_block_content(json!({"role": "assistant", "content": [{"type": "text", "text": "old"}], "name": "kept"}))]
|
||||
fn with_blocks_replaces_content_and_keeps_the_rest(#[case] message: Value) {
|
||||
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(message.with_blocks(vec![ContentBlock::text("new")])).unwrap(),
|
||||
json!({"role": "assistant", "content": [{"type": "text", "text": "new"}], "name": "kept"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::minimal(json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]}))]
|
||||
#[case::reasoning_effort_compaction_and_unknown_fields(json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
|
||||
"max_tokens": 8,
|
||||
"reasoning_effort": "high",
|
||||
"compaction": {"type": "auto"},
|
||||
"safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1}}],
|
||||
"metadata": {"user_id": "u"}
|
||||
}))]
|
||||
fn request_round_trips_unchanged(#[case] request: Value) {
|
||||
assert_eq!(round_trip::<AnthropicMessagesRequest>(&request), request);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,8 +9,6 @@ pub struct AnthropicMessagesResponse {
|
|||
pub role: String,
|
||||
pub model: String,
|
||||
pub content: Vec<Value>,
|
||||
// Anthropic always includes stop_reason / stop_sequence, null until the turn
|
||||
// ends; serialize them even when None so callers see the same shape as Python.
|
||||
pub stop_reason: Option<String>,
|
||||
pub stop_sequence: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -20,3 +18,61 @@ pub struct AnthropicMessagesResponse {
|
|||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn response(
|
||||
stop_reason: Option<&str>,
|
||||
stop_sequence: Option<&str>,
|
||||
usage: Option<Value>,
|
||||
container: Option<Value>,
|
||||
) -> AnthropicMessagesResponse {
|
||||
AnthropicMessagesResponse {
|
||||
id: "msg_1".to_string(),
|
||||
message_type: "message".to_string(),
|
||||
role: "assistant".to_string(),
|
||||
model: "claude".to_string(),
|
||||
content: vec![],
|
||||
stop_reason: stop_reason.map(str::to_string),
|
||||
stop_sequence: stop_sequence.map(str::to_string),
|
||||
usage,
|
||||
container,
|
||||
extra: Map::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::turn_in_progress(None, None, json!(null), json!(null))]
|
||||
#[case::ended_on_end_turn(Some("end_turn"), None, json!("end_turn"), json!(null))]
|
||||
#[case::ended_on_stop_sequence(Some("stop_sequence"), Some("###"), json!("stop_sequence"), json!("###"))]
|
||||
fn stop_fields_are_always_serialized(
|
||||
#[case] stop_reason: Option<&str>,
|
||||
#[case] stop_sequence: Option<&str>,
|
||||
#[case] expected_reason: Value,
|
||||
#[case] expected_sequence: Value,
|
||||
) {
|
||||
let body: Value = serde_json::to_value(response(stop_reason, stop_sequence, None, None))
|
||||
.expect("serializable");
|
||||
assert_eq!(body.get("stop_reason"), Some(&expected_reason));
|
||||
assert_eq!(body.get("stop_sequence"), Some(&expected_sequence));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None, None)]
|
||||
#[case::present(Some(json!({"input_tokens": 1})), Some(json!({"id": "c_1"})))]
|
||||
fn usage_and_container_are_omitted_only_when_none(
|
||||
#[case] usage: Option<Value>,
|
||||
#[case] container: Option<Value>,
|
||||
) {
|
||||
let body: Value =
|
||||
serde_json::to_value(response(None, None, usage.clone(), container.clone()))
|
||||
.expect("serializable");
|
||||
assert_eq!(body.get("usage").cloned(), usage);
|
||||
assert_eq!(body.get("container").cloned(), container);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,21 @@ use serde_json::{Map, Value};
|
|||
|
||||
use crate::llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk};
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ProviderSpecificHeader {
|
||||
#[serde(default)]
|
||||
pub custom_llm_provider: String,
|
||||
#[serde(default)]
|
||||
pub extra_headers: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum ProviderSpecificHeaders {
|
||||
One(ProviderSpecificHeader),
|
||||
Many(Vec<ProviderSpecificHeader>),
|
||||
}
|
||||
|
||||
/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python
|
||||
/// path reports so cost tracking sees the same numbers on either path.
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
|
|
|
|||
|
|
@ -11,7 +11,10 @@ from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_
|
|||
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
|
||||
from litellm.llms.bedrock.batches.transformation import titan_embedding_usage_from_batch_output
|
||||
from litellm.llms.vertex_ai.batches.transformation import vertex_prompt_tokens_details
|
||||
from litellm.llms.vertex_ai.batches.transformation import (
|
||||
is_native_vertex_batch_output_row,
|
||||
native_vertex_batch_row_stats,
|
||||
)
|
||||
from litellm.types.llms.openai import Batch
|
||||
from litellm.types.utils import ModelInfo, Usage
|
||||
from litellm.utils import token_counter
|
||||
|
|
@ -31,6 +34,20 @@ class BatchCostUsageResult:
|
|||
|
||||
|
||||
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
|
||||
|
||||
|
||||
def _uses_native_vertex_output(
|
||||
custom_llm_provider: str,
|
||||
model_name: str | None,
|
||||
first_row: Mapping[str, object] | None,
|
||||
) -> bool:
|
||||
if custom_llm_provider != "vertex_ai":
|
||||
return False
|
||||
if model_name and getattr(litellm, "disable_vertex_batch_output_transformation", False):
|
||||
return True
|
||||
return first_row is not None and is_native_vertex_batch_output_row(first_row)
|
||||
|
||||
|
||||
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
|
||||
|
||||
|
||||
|
|
@ -66,12 +83,9 @@ async def calculate_batch_cost_and_usage(
|
|||
deployment-specific pricing (e.g. input_cost_per_token_batches)
|
||||
is used instead of the global cost map.
|
||||
"""
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
):
|
||||
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
|
||||
first_row: Final = file_content_dictionary[0] if file_content_dictionary else None
|
||||
if _uses_native_vertex_output(custom_llm_provider, model_name, first_row):
|
||||
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name, model_info=model_info)
|
||||
|
||||
return _aggregate_batch_cost_usage_models(
|
||||
entries=file_content_dictionary,
|
||||
|
|
@ -126,11 +140,11 @@ async def _handle_completed_batch(
|
|||
)
|
||||
|
||||
output_file_result: Final = (
|
||||
calculate_vertex_ai_batch_cost_and_usage(_get_file_content_as_dictionary(file_content), model_name)
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
calculate_vertex_ai_batch_cost_and_usage(
|
||||
_iter_batch_output_entries(file_content), model_name, model_info=model_info
|
||||
)
|
||||
if _uses_native_vertex_output(
|
||||
custom_llm_provider, model_name, next(_iter_batch_output_entries(file_content), None)
|
||||
)
|
||||
else _aggregate_batch_cost_usage_models(
|
||||
entries=_iter_batch_output_entries(file_content),
|
||||
|
|
@ -332,69 +346,36 @@ def _aggregate_batch_cost_usage_models(
|
|||
|
||||
|
||||
def calculate_vertex_ai_batch_cost_and_usage(
|
||||
vertex_ai_batch_responses: list[dict],
|
||||
vertex_ai_batch_responses: Iterable[dict],
|
||||
model_name: str | None = None,
|
||||
model_info: ModelInfo | None = None,
|
||||
) -> BatchCostUsageResult:
|
||||
"""
|
||||
Calculate both cost and usage from raw Vertex AI batch responses.
|
||||
|
||||
Used only when ``litellm.disable_vertex_batch_output_transformation = True``.
|
||||
In that case the GCS predictions.jsonl is returned as-is, with each line in
|
||||
the native Vertex format:
|
||||
|
||||
{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}}}
|
||||
|
||||
usageMetadata contains promptTokenCount, candidatesTokenCount, totalTokenCount.
|
||||
|
||||
A row with no ``response`` is counted as failed - the same signal already
|
||||
used to skip it from cost/usage aggregation, since Vertex batch prediction
|
||||
output doesn't establish a distinct error shape in this (non-default) path.
|
||||
Cost and usage of a native Vertex predictions.jsonl, one
|
||||
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
|
||||
generateContent row or `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
|
||||
embedding row per line. `model_name` (the deployment model) prices every row, else each row's own
|
||||
`modelVersion` does; a row without a usable response counts as failed.
|
||||
"""
|
||||
from litellm.cost_calculator import batch_cost_calculator
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
|
||||
total_prompt_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
|
||||
total_completion_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
|
||||
total_tokens = 0
|
||||
prompt_tokens = 0
|
||||
completion_tokens = 0
|
||||
successful_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
|
||||
failed_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
|
||||
actual_model_name: Final = model_name or "gemini-2.0-flash-001"
|
||||
|
||||
for response in vertex_ai_batch_responses:
|
||||
response_body = response.get("response")
|
||||
if response_body is None:
|
||||
failed_requests += 1
|
||||
continue
|
||||
successful_requests += 1
|
||||
|
||||
usage_metadata = response_body.get("usageMetadata", {})
|
||||
_prompt = usage_metadata.get("promptTokenCount", 0) or 0
|
||||
_completion = usage_metadata.get("candidatesTokenCount", 0) or 0
|
||||
_total = usage_metadata.get("totalTokenCount", 0) or (_prompt + _completion)
|
||||
|
||||
line_usage = Usage(
|
||||
prompt_tokens=_prompt,
|
||||
completion_tokens=_completion,
|
||||
total_tokens=_total,
|
||||
prompt_tokens_details=vertex_prompt_tokens_details(usage_metadata),
|
||||
row_stats: Final = tuple(
|
||||
native_vertex_batch_row_stats(
|
||||
row,
|
||||
model_name,
|
||||
model_info=model_info,
|
||||
calculate_usage=VertexGeminiConfig._calculate_usage,
|
||||
cost_calculator=batch_cost_calculator,
|
||||
)
|
||||
|
||||
try:
|
||||
p_cost, c_cost = batch_cost_calculator(
|
||||
usage=line_usage,
|
||||
model=actual_model_name,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
total_prompt_cost += p_cost
|
||||
total_completion_cost += c_cost
|
||||
except Exception as e:
|
||||
verbose_logger.debug("vertex_ai batch cost calculation error for line: %s", str(e))
|
||||
|
||||
prompt_tokens += _prompt
|
||||
completion_tokens += _completion
|
||||
total_tokens += _total
|
||||
|
||||
for row in vertex_ai_batch_responses
|
||||
)
|
||||
priced: Final = tuple(stats for stats in row_stats if stats is not None)
|
||||
total_prompt_cost: Final = sum(stats.prompt_cost for stats in priced)
|
||||
total_completion_cost: Final = sum(stats.completion_cost for stats in priced)
|
||||
prompt_tokens: Final = sum(stats.usage.prompt_tokens for stats in priced)
|
||||
completion_tokens: Final = sum(stats.usage.completion_tokens for stats in priced)
|
||||
total_tokens: Final = sum(stats.total_tokens for stats in priced)
|
||||
total_cost: Final = total_prompt_cost + total_completion_cost
|
||||
verbose_logger.info(
|
||||
"vertex_ai batch cost: cost=%s, prompt=%d, completion=%d, total=%d, successful=%d, failed=%d",
|
||||
|
|
@ -402,8 +383,8 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
prompt_tokens,
|
||||
completion_tokens,
|
||||
total_tokens,
|
||||
successful_requests,
|
||||
failed_requests,
|
||||
len(priced),
|
||||
len(row_stats) - len(priced),
|
||||
)
|
||||
|
||||
return BatchCostUsageResult(
|
||||
|
|
@ -413,9 +394,13 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
),
|
||||
models=[actual_model_name],
|
||||
successful_requests=successful_requests,
|
||||
failed_requests=failed_requests,
|
||||
models=(
|
||||
[model_name]
|
||||
if model_name
|
||||
else list(dict.fromkeys(stats.model for stats in priced if stats.model is not None))
|
||||
),
|
||||
successful_requests=len(priced),
|
||||
failed_requests=len(row_stats) - len(priced),
|
||||
prompt_cost=total_prompt_cost,
|
||||
completion_cost=total_completion_cost,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -176,6 +176,22 @@ def create_file(
|
|||
if logging_obj is None:
|
||||
raise ValueError("logging_obj is required")
|
||||
client: Final = kwargs.get("client")
|
||||
if litellm_params_dict.get("passthrough") is True and (
|
||||
custom_llm_provider != "vertex_ai" or purpose != "batch"
|
||||
):
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message=(
|
||||
"`passthrough=True` uploads the file bytes unchanged for a native Vertex AI batch, so it needs "
|
||||
f"custom_llm_provider='vertex_ai' and purpose='batch', got '{custom_llm_provider}' and '{purpose}'."
|
||||
),
|
||||
model="n/a",
|
||||
llm_provider=custom_llm_provider or "n/a",
|
||||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="passthrough needs a vertex_ai batch",
|
||||
request=httpx.Request(method="create_file", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
|
|||
|
|
@ -2778,6 +2778,13 @@ class PrometheusLogger(CustomLogger):
|
|||
- increment deployment failure responses metric
|
||||
- increment deployment total requests metric
|
||||
|
||||
Both counters also carry a model_group label. When a deployment was
|
||||
actually selected, model_group is the router-resolved value and is
|
||||
trusted as-is. On a pre-routing reject (no deployment selected), it
|
||||
is caller-supplied via litellm_params.metadata and is bounded with
|
||||
_bounded_requested_model_label the same way requested_model is, so an
|
||||
unrecognized value cannot mint unbounded label series.
|
||||
|
||||
Args:
|
||||
request_kwargs: dict
|
||||
|
||||
|
|
@ -2844,6 +2851,7 @@ class PrometheusLogger(CustomLogger):
|
|||
label_api_base = api_base
|
||||
label_api_provider = llm_provider
|
||||
label_requested_model = model_group or litellm_model_name
|
||||
label_model_group = model_group
|
||||
else:
|
||||
label_litellm_model_name = ""
|
||||
label_model_id = ""
|
||||
|
|
@ -2852,6 +2860,7 @@ class PrometheusLogger(CustomLogger):
|
|||
label_requested_model = (
|
||||
_bounded_requested_model_label(litellm_model_name or model_group, router_originated=True) or ""
|
||||
)
|
||||
label_model_group = _bounded_requested_model_label(model_group, router_originated=True)
|
||||
|
||||
enum_values: Final = UserAPIKeyLabelValues(
|
||||
litellm_model_name=label_litellm_model_name,
|
||||
|
|
@ -2861,6 +2870,7 @@ class PrometheusLogger(CustomLogger):
|
|||
exception_status=exception_status,
|
||||
exception_class=(self._get_exception_class_name(exception) if exception else None),
|
||||
requested_model=label_requested_model,
|
||||
model_group=label_model_group,
|
||||
hashed_api_key=hashed_api_key,
|
||||
api_key_alias=api_key_alias,
|
||||
user_email=user_email,
|
||||
|
|
@ -2912,9 +2922,21 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id: str | None,
|
||||
api_base: str | None,
|
||||
llm_provider: str | None,
|
||||
model_group: str | None,
|
||||
):
|
||||
"""
|
||||
Set the deployment TPM and RPM limits metrics
|
||||
|
||||
Args:
|
||||
model_info: the deployment's static model_info config (id, tpm, rpm, etc.)
|
||||
litellm_params: the deployment's litellm_params, as a tpm/rpm fallback source
|
||||
litellm_model_name: the resolved deployment model name
|
||||
model_id: the deployment's model_id
|
||||
api_base: the deployment's api_base
|
||||
llm_provider: the deployment's custom_llm_provider
|
||||
model_group: the router-resolved model_group the deployment belongs to,
|
||||
from the caller's already-resolved enum_values.model_group (trusted,
|
||||
not caller-supplied at this call site)
|
||||
"""
|
||||
tpm: Final = model_info.get("tpm") or litellm_params.get("tpm")
|
||||
rpm: Final = model_info.get("rpm") or litellm_params.get("rpm")
|
||||
|
|
@ -2927,6 +2949,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
api_provider=llm_provider,
|
||||
model_group=model_group,
|
||||
),
|
||||
)
|
||||
self.litellm_deployment_tpm_limit.labels(**_labels).set(tpm)
|
||||
|
|
@ -2939,6 +2962,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
api_provider=llm_provider,
|
||||
model_group=model_group,
|
||||
),
|
||||
)
|
||||
self.litellm_deployment_rpm_limit.labels(**_labels).set(rpm)
|
||||
|
|
@ -3058,6 +3082,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
llm_provider=llm_provider,
|
||||
model_group=enum_values.model_group,
|
||||
)
|
||||
|
||||
remaining_requests: int | None = None
|
||||
|
|
|
|||
|
|
@ -765,3 +765,30 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo
|
|||
RESPONSE_COST_HEADER: cost,
|
||||
}
|
||||
hidden_params["additional_headers"] = merged
|
||||
|
||||
|
||||
_HIDDEN_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
_PROVIDER_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])
|
||||
|
||||
|
||||
def set_provider_response_headers_in_hidden_params(
|
||||
response: _CarriesHiddenParams, headers: httpx.Headers | Mapping[str, str]
|
||||
) -> None:
|
||||
hidden_params: Final = response._hidden_params # pyright: ignore[reportPrivateUsage] # no public accessor
|
||||
existing_additional_headers: Final[object] = hidden_params.get("additional_headers")
|
||||
raw_headers: Final[dict[str, str]] = dict(headers) # mutable-ok: stored as the plain-dict hidden param
|
||||
additional_headers: Final[dict[str, object]] = { # mutable-ok: assigned into the plain-dict hidden params
|
||||
**process_response_headers(raw_headers),
|
||||
**(existing_additional_headers if isinstance(existing_additional_headers, Mapping) else _NO_HEADERS),
|
||||
}
|
||||
hidden_params["headers"] = raw_headers
|
||||
hidden_params["additional_headers"] = additional_headers
|
||||
|
||||
|
||||
def get_provider_response_headers_from_hidden_params(response: object) -> Mapping[str, str] | None:
|
||||
hidden_params: Final[object] = getattr(response, "_hidden_params", None)
|
||||
try:
|
||||
validated: Final = _HIDDEN_PARAMS_ADAPTER.validate_python(hidden_params)
|
||||
return _PROVIDER_HEADERS_ADAPTER.validate_python(validated.get("headers"))
|
||||
except ValidationError:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ from litellm.litellm_core_utils.classifier_logging import (
|
|||
is_classifier_call,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_provider_response_headers_from_hidden_params,
|
||||
is_expected_client_error,
|
||||
reconstruct_model_name,
|
||||
set_response_cost_in_hidden_params,
|
||||
|
|
@ -2353,6 +2354,15 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
return logging_result
|
||||
|
||||
def _surface_response_headers_from_result(self, logging_result: object) -> None:
|
||||
existing: Final[object] = self.model_call_details.get("response_headers")
|
||||
if existing is not None:
|
||||
return
|
||||
headers: Final = get_provider_response_headers_from_hidden_params(logging_result)
|
||||
if headers is None:
|
||||
return
|
||||
self.model_call_details["response_headers"] = headers
|
||||
|
||||
def _merge_hidden_params_from_response_into_metadata(self, logging_result: object) -> None:
|
||||
"""
|
||||
Copy response._hidden_params into litellm_params.metadata['hidden_params'].
|
||||
|
|
@ -2386,6 +2396,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
build_logging_payload: bool = True,
|
||||
):
|
||||
"""Resolve hidden params, compute response cost, and emit the standard logging payload."""
|
||||
self._surface_response_headers_from_result(logging_result)
|
||||
hidden_params: Final = getattr(logging_result, "_hidden_params", {})
|
||||
if hidden_params:
|
||||
if self.model_call_details.get("litellm_params") is not None:
|
||||
|
|
@ -2788,6 +2799,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if complete_streaming_response is not None:
|
||||
verbose_logger.debug("Logging Details LiteLLM-Success Call streaming complete")
|
||||
self.model_call_details["complete_streaming_response"] = complete_streaming_response
|
||||
self._surface_response_headers_from_result(complete_streaming_response)
|
||||
self.model_call_details["response_cost"] = self._response_cost_calculator(
|
||||
result=complete_streaming_response
|
||||
)
|
||||
|
|
@ -3302,6 +3314,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
print_verbose("Async success callbacks: Got a complete streaming response")
|
||||
|
||||
self.model_call_details["async_complete_streaming_response"] = complete_streaming_response
|
||||
self._surface_response_headers_from_result(complete_streaming_response)
|
||||
|
||||
try:
|
||||
if self.model_call_details.get("cache_hit", False) is True:
|
||||
|
|
@ -6362,12 +6375,15 @@ def _extract_response_obj_and_hidden_params(
|
|||
original_exception: Exception | None,
|
||||
) -> tuple[dict, dict | None]:
|
||||
"""Extract response_obj and hidden_params from init_response_obj."""
|
||||
hidden_params: dict | None = None
|
||||
hidden_params: dict | None = (
|
||||
getattr(init_response_obj, "_hidden_params", None)
|
||||
if isinstance(init_response_obj, BaseModel | HttpxBinaryResponseContent)
|
||||
else None
|
||||
)
|
||||
if init_response_obj is None:
|
||||
response_obj = {}
|
||||
elif isinstance(init_response_obj, BaseModel):
|
||||
response_obj = init_response_obj.model_dump()
|
||||
hidden_params = getattr(init_response_obj, "_hidden_params", None)
|
||||
elif isinstance(init_response_obj, dict):
|
||||
response_obj = init_response_obj
|
||||
elif isinstance(init_response_obj, HttpxBinaryResponseContent):
|
||||
|
|
|
|||
|
|
@ -1313,6 +1313,8 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
alt_model: Final = BedrockModelInfo.get_non_litellm_routing_model_name(model=model)
|
||||
if base_model in litellm.bedrock_converse_models or alt_model in litellm.bedrock_converse_models:
|
||||
return "converse"
|
||||
if _OPENAI_FAMILY_MODEL_RE.search(base_model):
|
||||
return "converse"
|
||||
return "invoke"
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
|||
SUBTITLE_RESPONSE_FORMATS,
|
||||
synthesize_subtitle_document,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
|
||||
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
|
||||
from litellm.litellm_core_utils.realtime_errors import (
|
||||
|
|
@ -1461,6 +1462,7 @@ class BaseLLMHTTPHandler:
|
|||
transformed: Final = provider_config.transform_audio_transcription_response(
|
||||
raw_response=response,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(transformed, response.headers)
|
||||
if not provider_config.supports_subtitle_synthesis:
|
||||
return transformed
|
||||
requested_format: Final = optional_params.get("response_format")
|
||||
|
|
@ -6960,11 +6962,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=image_edit_provider_config,
|
||||
)
|
||||
|
||||
return image_edit_provider_config.transform_image_edit_response(
|
||||
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
|
||||
return image_edit_response
|
||||
|
||||
async def async_image_edit_handler(
|
||||
self,
|
||||
|
|
@ -7059,11 +7063,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=image_edit_provider_config,
|
||||
)
|
||||
|
||||
return image_edit_provider_config.transform_image_edit_response(
|
||||
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
|
||||
return image_edit_response
|
||||
|
||||
def image_generation_handler(
|
||||
self,
|
||||
|
|
@ -7186,6 +7192,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
encoding=None,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(model_response, response.headers)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -7293,6 +7300,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
encoding=None,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(model_response, response.headers)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -12077,11 +12085,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=text_to_speech_provider_config,
|
||||
)
|
||||
|
||||
return text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
|
||||
return speech_response
|
||||
|
||||
async def async_text_to_speech_handler(
|
||||
self,
|
||||
|
|
@ -12176,11 +12186,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=text_to_speech_provider_config,
|
||||
)
|
||||
|
||||
return text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
|
||||
return speech_response
|
||||
|
||||
#########################################################
|
||||
########## SKILLS API HANDLERS ##########################
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm import LlmProviders
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import DEFAULT_MAX_RETRIES
|
||||
from litellm.files.types import FileContentStreamingResult
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
|
|
@ -1404,7 +1405,6 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
organization: str | None = None,
|
||||
headers: dict | None = None,
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
openai_aclient: Final = self._get_openai_client(
|
||||
is_async=True,
|
||||
|
|
@ -1428,8 +1428,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
)
|
||||
|
||||
request_data: Final = {**data, "extra_headers": headers} if headers else data
|
||||
response = await openai_aclient.images.generate(**request_data, timeout=timeout)
|
||||
stringified_response: Final = response.model_dump()
|
||||
raw_response: Final = await openai_aclient.images.with_raw_response.generate(
|
||||
**request_data, timeout=timeout
|
||||
)
|
||||
stringified_response: Final = raw_response.parse().model_dump()
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
input=prompt,
|
||||
|
|
@ -1437,11 +1439,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=stringified_response,
|
||||
)
|
||||
return convert_to_model_response_object(
|
||||
image_response: Final[ImageResponse] = convert_to_model_response_object(
|
||||
response_object=stringified_response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
|
||||
return image_response
|
||||
except Exception as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -1512,9 +1516,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
|
||||
## COMPLETION CALL
|
||||
request_data: Final = {**data, "extra_headers": headers} if headers else data
|
||||
_response: Final = openai_client.images.generate(**request_data, timeout=timeout)
|
||||
raw_response: Final = openai_client.images.with_raw_response.generate(**request_data, timeout=timeout)
|
||||
|
||||
response: Final = _response.model_dump()
|
||||
response: Final = raw_response.parse().model_dump()
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
input=prompt,
|
||||
|
|
@ -1522,11 +1526,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=response,
|
||||
)
|
||||
return convert_to_model_response_object(
|
||||
image_response: Final[ImageResponse] = convert_to_model_response_object(
|
||||
response_object=response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
|
||||
return image_response
|
||||
except OpenAIError as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -1609,7 +1615,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
input=input,
|
||||
**optional_params,
|
||||
)
|
||||
return HttpxBinaryResponseContent(response=response.response)
|
||||
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
|
||||
return speech_response
|
||||
|
||||
async def async_audio_speech(
|
||||
self,
|
||||
|
|
@ -1655,8 +1663,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
input=input,
|
||||
**optional_params,
|
||||
)
|
||||
|
||||
return HttpxBinaryResponseContent(response=response.response)
|
||||
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
|
||||
return speech_response
|
||||
|
||||
|
||||
class OpenAIFilesAPI(BaseLLM):
|
||||
|
|
|
|||
|
|
@ -4,11 +4,10 @@ import httpx
|
|||
from openai import AsyncOpenAI, OpenAI
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from aiohttp import ClientSession
|
||||
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
BaseAudioTranscriptionConfig,
|
||||
|
|
@ -31,11 +30,6 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
data: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
):
|
||||
"""
|
||||
Helper to:
|
||||
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
|
||||
- call openai_aclient.audio.transcriptions.create by default
|
||||
"""
|
||||
try:
|
||||
raw_response = await openai_aclient.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
|
|
@ -51,20 +45,11 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
data: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
):
|
||||
"""
|
||||
Helper to:
|
||||
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
|
||||
- call openai_aclient.audio.transcriptions.create by default
|
||||
"""
|
||||
try:
|
||||
if litellm.return_response_headers is True:
|
||||
raw_response = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
response = raw_response.parse()
|
||||
return headers, response
|
||||
else:
|
||||
response = openai_client.audio.transcriptions.create(**data, timeout=timeout)
|
||||
return None, response
|
||||
raw_response: Final = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
response: Final = raw_response.parse()
|
||||
return headers, response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
|
@ -133,11 +118,12 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
"complete_input_dict": data,
|
||||
},
|
||||
)
|
||||
_, response = self.make_sync_openai_audio_transcriptions_request(
|
||||
headers, response = self.make_sync_openai_audio_transcriptions_request(
|
||||
openai_client=openai_client,
|
||||
data=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
logging_obj.model_call_details["response_headers"] = headers
|
||||
|
||||
if isinstance(response, BaseModel):
|
||||
stringified_response = response.model_dump()
|
||||
|
|
@ -158,6 +144,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
hidden_params=hidden_params,
|
||||
response_type="audio_transcription",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(final_response, headers)
|
||||
return final_response
|
||||
|
||||
async def async_audio_transcriptions(
|
||||
|
|
@ -217,12 +204,14 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
actual_model: Final = data.get("model", "whisper-1")
|
||||
hidden_params: Final = {"model": actual_model, "custom_llm_provider": "openai"}
|
||||
|
||||
return convert_to_model_response_object(
|
||||
final_response: Final[TranscriptionResponse] = convert_to_model_response_object(
|
||||
response_object=stringified_response,
|
||||
model_response_object=model_response,
|
||||
hidden_params=hidden_params,
|
||||
response_type="audio_transcription",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(final_response, headers)
|
||||
return final_response
|
||||
except Exception as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
|
|||
|
|
@ -1,7 +1,11 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Final, Protocol
|
||||
from urllib.parse import unquote
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
VertexAIError,
|
||||
|
|
@ -9,35 +13,128 @@ from litellm.llms.vertex_ai.common_utils import (
|
|||
)
|
||||
from litellm.types.llms.openai import BatchJobStatus, CreateBatchRequest
|
||||
from litellm.types.llms.vertex_ai import *
|
||||
from litellm.types.utils import LiteLLMBatch, PromptTokensDetailsWrapper
|
||||
from litellm.types.llms.vertex_ai import GenerateContentResponseBody
|
||||
from litellm.types.utils import LiteLLMBatch, ModelInfo, Usage
|
||||
|
||||
_NATIVE_VERTEX_RESPONSE: Final = TypeAdapter(GenerateContentResponseBody)
|
||||
|
||||
|
||||
def vertex_prompt_tokens_details(
|
||||
usage_metadata: Mapping[str, object],
|
||||
) -> PromptTokensDetailsWrapper | None:
|
||||
raw_details: Final = usage_metadata.get("promptTokensDetails")
|
||||
if not isinstance(raw_details, list):
|
||||
return None
|
||||
def _int_field(mapping: Mapping[str, object], key: str) -> int:
|
||||
value: Final = mapping.get(key)
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
return int(value) if isinstance(value, str) and value.isdigit() else 0
|
||||
|
||||
def _normalize(detail: object) -> tuple[str, int] | None:
|
||||
if not isinstance(detail, Mapping):
|
||||
|
||||
def vertex_embedding_prompt_token_count(vertex_response: Mapping[str, object]) -> int:
|
||||
"""
|
||||
Prompt tokens billed for one Vertex Gemini Embedding batch row.
|
||||
|
||||
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
|
||||
a fallback.
|
||||
"""
|
||||
usage_metadata: Final = vertex_response.get("usageMetadata")
|
||||
if isinstance(usage_metadata, Mapping):
|
||||
return _int_field(usage_metadata, "promptTokenCount")
|
||||
return _int_field(vertex_response, "tokenCount")
|
||||
|
||||
|
||||
def is_vertex_embedding_batch_output_response(response_body: Mapping[str, object]) -> bool:
|
||||
return isinstance(response_body.get("embedding"), dict)
|
||||
|
||||
|
||||
def is_native_vertex_batch_output_row(row: Mapping[str, object]) -> bool:
|
||||
return isinstance(row.get("request"), dict)
|
||||
|
||||
|
||||
class NativeVertexBatchCostCalculator(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
usage: Usage,
|
||||
model: str,
|
||||
custom_llm_provider: str | None = None,
|
||||
model_info: ModelInfo | None = None,
|
||||
) -> tuple[float, float]: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeVertexBatchRowStats:
|
||||
usage: Usage
|
||||
total_tokens: int
|
||||
model: str | None
|
||||
prompt_cost: float
|
||||
completion_cost: float
|
||||
|
||||
|
||||
def _native_vertex_row_usage(
|
||||
response_body: Mapping[str, object],
|
||||
calculate_usage: Callable[[GenerateContentResponseBody], Usage],
|
||||
) -> Usage | None:
|
||||
if "usageMetadata" not in response_body:
|
||||
if not is_vertex_embedding_batch_output_response(response_body):
|
||||
return None
|
||||
modality: Final = detail.get("modality")
|
||||
token_count: Final = detail.get("tokenCount")
|
||||
if not isinstance(modality, str) or not isinstance(token_count, int):
|
||||
return None
|
||||
return modality.upper(), token_count
|
||||
|
||||
parsed_details: Final = tuple(_normalize(detail) for detail in raw_details)
|
||||
normalized: Final = tuple(detail for detail in parsed_details if detail is not None)
|
||||
if len(normalized) != len(parsed_details):
|
||||
prompt_tokens: Final = vertex_embedding_prompt_token_count(response_body)
|
||||
return Usage(prompt_tokens=prompt_tokens, completion_tokens=0, total_tokens=prompt_tokens)
|
||||
try:
|
||||
completion_response: Final = _NATIVE_VERTEX_RESPONSE.validate_python(response_body)
|
||||
except ValidationError as e:
|
||||
verbose_logger.debug("vertex_ai batch row response is not a GenerateContentResponse: %s", str(e))
|
||||
return None
|
||||
return calculate_usage(completion_response)
|
||||
|
||||
return PromptTokensDetailsWrapper(
|
||||
text_tokens=sum(token_count for modality, token_count in normalized if modality in ("TEXT", "DOCUMENT")),
|
||||
audio_tokens=sum(token_count for modality, token_count in normalized if modality == "AUDIO"),
|
||||
image_tokens=sum(token_count for modality, token_count in normalized if modality == "IMAGE"),
|
||||
video_tokens=sum(token_count for modality, token_count in normalized if modality == "VIDEO"),
|
||||
|
||||
def native_vertex_batch_row_stats(
|
||||
row: Mapping[str, object],
|
||||
model_name: str | None,
|
||||
*,
|
||||
model_info: ModelInfo | None,
|
||||
calculate_usage: Callable[[GenerateContentResponseBody], Usage],
|
||||
cost_calculator: NativeVertexBatchCostCalculator,
|
||||
) -> NativeVertexBatchRowStats | None:
|
||||
"""
|
||||
Usage and cost of one native Vertex predictions.jsonl row, a
|
||||
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
|
||||
generateContent object or a `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
|
||||
embedding object (an embedding row without `usageMetadata` is billed from its documented `tokenCount`).
|
||||
`model_name` (the deployment model) prices the row unless it is a wildcard, else its own `modelVersion`
|
||||
does, else the wildcard name so explicit deployment prices still apply; a row without a response, a
|
||||
generateContent row without `response.usageMetadata`, and a row whose response fails validation are
|
||||
None (failed).
|
||||
"""
|
||||
response_body: Final = row.get("response")
|
||||
if not isinstance(response_body, dict):
|
||||
return None
|
||||
usage: Final = _native_vertex_row_usage(response_body, calculate_usage)
|
||||
if usage is None:
|
||||
return None
|
||||
total_tokens: Final = usage.total_tokens or (usage.prompt_tokens + usage.completion_tokens)
|
||||
model_version: Final = response_body.get("modelVersion")
|
||||
deployment_model: Final = model_name if model_name and "*" not in model_name else None
|
||||
model: Final = deployment_model or (model_version if isinstance(model_version, str) else model_name)
|
||||
if model is None:
|
||||
verbose_logger.warning(
|
||||
"vertex_ai batch output row could not be costed, so it is billed at $0 and the rest of the batch "
|
||||
"is still billed: the row has no modelVersion and the batch has no deployment model"
|
||||
)
|
||||
return NativeVertexBatchRowStats(
|
||||
usage=usage, total_tokens=total_tokens, model=None, prompt_cost=0.0, completion_cost=0.0
|
||||
)
|
||||
try:
|
||||
prompt_cost, completion_cost = cost_calculator(
|
||||
usage=usage, model=model, custom_llm_provider="vertex_ai", model_info=model_info
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # one unpriceable row must not abort the batch's cost accounting
|
||||
verbose_logger.warning(
|
||||
"vertex_ai batch output row could not be costed, so it is billed at $0 and the rest of the batch "
|
||||
"is still billed. model=%s error=%s",
|
||||
model,
|
||||
str(e),
|
||||
)
|
||||
return NativeVertexBatchRowStats(
|
||||
usage=usage, total_tokens=total_tokens, model=model, prompt_cost=0.0, completion_cost=0.0
|
||||
)
|
||||
return NativeVertexBatchRowStats(
|
||||
usage=usage, total_tokens=total_tokens, model=model, prompt_cost=prompt_cost, completion_cost=completion_cost
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from collections.abc import AsyncGenerator, Callable, Iterable, Iterator, Mappin
|
|||
from contextlib import aclosing
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, TypedDict
|
||||
from typing import IO, Any, Final, TypedDict
|
||||
from urllib.parse import quote, unquote
|
||||
|
||||
import httpx
|
||||
|
|
@ -41,6 +41,7 @@ from litellm.llms.base_llm.files.transformation import (
|
|||
BaseFileUploadStream,
|
||||
LiteLLMLoggingObj,
|
||||
)
|
||||
from litellm.llms.vertex_ai.batches.transformation import vertex_embedding_prompt_token_count
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
_convert_vertex_datetime_to_openai_datetime,
|
||||
get_vertex_ai_fine_tuned_endpoint_id,
|
||||
|
|
@ -56,6 +57,7 @@ from litellm.types.files import StreamingMediaUploadConfig
|
|||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
CreateFileRequest,
|
||||
FileContent,
|
||||
FileTypes,
|
||||
HttpxBinaryResponseContent,
|
||||
OpenAICreateFileRequestOptionalParams,
|
||||
|
|
@ -87,6 +89,8 @@ _EMBED_REQUEST_FIELD_BY_GEMINI_PARAM: Final = (
|
|||
_VERTEX_BATCH_FANNED_OUT_KEY_PATTERN: Final = re.compile(r"(?P<custom_id>[^#]*)#(?P<index>\d+)/(?P<total>\d+)")
|
||||
_JSONL_NEWLINE: Final = b"\n"
|
||||
_BATCH_OUTPUT_FIRST_ROW_PEEK_LIMIT_BYTES: Final = 32 * 1024 * 1024
|
||||
_PASSTHROUGH_MANAGED_GCS_PREFIX: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}passthrough/"
|
||||
_RAW_UPLOAD_CHUNK_BYTES: Final = 1024 * 1024
|
||||
|
||||
|
||||
class _GcsObjectMetadataJson(TypedDict, total=False):
|
||||
|
|
@ -418,19 +422,6 @@ def _split_vertex_batch_key(vertex_output_row: Mapping[str, object]) -> tuple[st
|
|||
return unquote(match["custom_id"]), int(match["index"]), int(match["total"])
|
||||
|
||||
|
||||
def _embedding_prompt_token_count(vertex_response: _VertexEmbeddingResponse) -> int:
|
||||
"""
|
||||
Prompt tokens billed for one Vertex Gemini Embedding batch row.
|
||||
|
||||
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
|
||||
a fallback.
|
||||
"""
|
||||
usage_metadata = vertex_response.get("usageMetadata")
|
||||
if isinstance(usage_metadata, Mapping):
|
||||
return int(usage_metadata.get("promptTokenCount") or 0)
|
||||
return int(vertex_response.get("tokenCount") or 0)
|
||||
|
||||
|
||||
def _vertex_embeddings_rows_to_openai_batch_output_row(
|
||||
custom_id: str,
|
||||
vertex_output_rows: tuple[_VertexEmbeddingBatchRow, ...],
|
||||
|
|
@ -471,7 +462,7 @@ def _vertex_embeddings_rows_to_openai_batch_output_row(
|
|||
)
|
||||
|
||||
responses = tuple(row["response"] for row in vertex_output_rows)
|
||||
token_count = sum(_embedding_prompt_token_count(response) for response in responses)
|
||||
token_count = sum(vertex_embedding_prompt_token_count(response) for response in responses)
|
||||
body = EmbeddingResponse(
|
||||
model=model or "",
|
||||
data=[
|
||||
|
|
@ -528,6 +519,16 @@ def _model_from_managed_gcs_url(url: str) -> str | None:
|
|||
return match.group(1) if match else None
|
||||
|
||||
|
||||
def is_passthrough_managed_gcs_url(url: str) -> bool:
|
||||
decoded_url: Final = unquote(url)
|
||||
managed_prefix_start: Final = decoded_url.find(VERTEX_AI_MANAGED_GCS_PREFIX)
|
||||
return managed_prefix_start >= 0 and decoded_url.startswith(_PASSTHROUGH_MANAGED_GCS_PREFIX, managed_prefix_start)
|
||||
|
||||
|
||||
def is_passthrough_batch_upload(create_file_data: Mapping[str, object], litellm_params: Mapping[str, object]) -> bool:
|
||||
return create_file_data.get("purpose") == "batch" and litellm_params.get("passthrough") is True
|
||||
|
||||
|
||||
def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
Whether an OpenAI batch JSONL line targets the embeddings endpoint.
|
||||
|
|
@ -791,6 +792,58 @@ class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
|
|||
return self._iter_vertex_jsonl_chunks()
|
||||
|
||||
|
||||
def _read_chunk_as_bytes(handle: IO[bytes]) -> bytes:
|
||||
chunk: Final[bytes | str] = handle.read(_RAW_UPLOAD_CHUNK_BYTES)
|
||||
return chunk.encode("utf-8") if isinstance(chunk, str) else bytes(chunk)
|
||||
|
||||
|
||||
def _iter_raw_file_chunks(file_content: FileTypes) -> Iterator[bytes]:
|
||||
content: Final[FileContent | str] = file_content[1] if isinstance(file_content, tuple) else file_content
|
||||
if isinstance(content, (bytes, bytearray)):
|
||||
yield from (
|
||||
bytes(content[offset : offset + _RAW_UPLOAD_CHUNK_BYTES])
|
||||
for offset in range(0, len(content), _RAW_UPLOAD_CHUNK_BYTES)
|
||||
)
|
||||
return
|
||||
if isinstance(content, str):
|
||||
yield content.encode("utf-8")
|
||||
return
|
||||
if isinstance(content, PathLike):
|
||||
with open(str(content), "rb") as handle:
|
||||
yield from iter(lambda: handle.read(_RAW_UPLOAD_CHUNK_BYTES), b"")
|
||||
return
|
||||
if not hasattr(content, "read"):
|
||||
raise ValueError("Unsupported file content type")
|
||||
seek: Final = getattr(content, "seek", None)
|
||||
if seek is None:
|
||||
raise ValueError(
|
||||
"Batch upload file handle must be seekable; got a non-seekable "
|
||||
"stream. Pass bytes, a path, or a seekable handle."
|
||||
)
|
||||
seek(0)
|
||||
yield from iter(lambda: _read_chunk_as_bytes(content), b"")
|
||||
|
||||
|
||||
class _RawFileUploadStream(BaseFileUploadStream):
|
||||
def __init__(self, file_content: FileTypes) -> None:
|
||||
self._file_content = file_content
|
||||
|
||||
def iter_bytes(self) -> Iterator[bytes]:
|
||||
return _iter_raw_file_chunks(self._file_content)
|
||||
|
||||
|
||||
def _managed_batch_object_name(raw_model: str, *, passthrough: bool) -> str:
|
||||
endpoint_id: Final = get_vertex_ai_fine_tuned_endpoint_id(raw_model)
|
||||
model_path: Final = (
|
||||
f"endpoints/{endpoint_id}"
|
||||
if endpoint_id is not None
|
||||
else (raw_model if "publishers/google/models" in raw_model else f"publishers/google/models/{raw_model}")
|
||||
)
|
||||
safe_model_path: Final = sanitize_cloud_object_path(model_path, fallback="model")
|
||||
prefix: Final = _PASSTHROUGH_MANAGED_GCS_PREFIX if passthrough else VERTEX_AI_MANAGED_GCS_PREFIX
|
||||
return f"{prefix}{safe_model_path}/{uuid.uuid4()}"
|
||||
|
||||
|
||||
class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
||||
"""
|
||||
Config for VertexAI Files
|
||||
|
|
@ -848,23 +901,34 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
if deployment_model
|
||||
else openai_jsonl_content[0].get("body", {}).get("model", "")
|
||||
)
|
||||
endpoint_id: Final = get_vertex_ai_fine_tuned_endpoint_id(raw_model)
|
||||
model_path: Final = (
|
||||
f"endpoints/{endpoint_id}"
|
||||
if endpoint_id is not None
|
||||
else (raw_model if "publishers/google/models" in raw_model else f"publishers/google/models/{raw_model}")
|
||||
)
|
||||
safe_model_path: Final = sanitize_cloud_object_path(model_path, fallback="model")
|
||||
object_name: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
|
||||
return object_name
|
||||
return _managed_batch_object_name(raw_model, passthrough=False)
|
||||
|
||||
def get_object_name(self, file_data: FileTypes, purpose: str, deployment_model: str | None = None) -> str:
|
||||
def _get_passthrough_gcs_object_name(self, deployment_model: str | None) -> str:
|
||||
if not deployment_model:
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"Native Vertex batch passthrough uploads need the deployment model to name the GCS object, "
|
||||
"since native rows carry no model: pass `target_model_names` (proxy) or `model` (SDK)."
|
||||
),
|
||||
)
|
||||
return _managed_batch_object_name(deployment_model.removeprefix("vertex_ai/"), passthrough=True)
|
||||
|
||||
def get_object_name(
|
||||
self,
|
||||
file_data: FileTypes,
|
||||
purpose: str,
|
||||
deployment_model: str | None = None,
|
||||
passthrough: bool = False,
|
||||
) -> str:
|
||||
"""
|
||||
Get the object name for the request.
|
||||
|
||||
Reads only the first JSONL entry (streamed) for batch files, so a large
|
||||
upload is never materialized just to derive the GCS object name.
|
||||
"""
|
||||
if purpose == "batch" and passthrough:
|
||||
return self._get_passthrough_gcs_object_name(deployment_model)
|
||||
if purpose == "batch":
|
||||
## 1. If jsonl, derive the object name from the deployment model (or the first entry's)
|
||||
first_entry: Final = next(_iter_openai_jsonl_entries(file_data), None)
|
||||
|
|
@ -922,6 +986,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
file_data,
|
||||
purpose,
|
||||
deployment_model=configured_model if isinstance(configured_model, str) else None,
|
||||
passthrough=is_passthrough_batch_upload(data, litellm_params),
|
||||
)
|
||||
if object_prefix:
|
||||
object_name = f"{object_prefix}/{object_name}"
|
||||
|
|
@ -984,6 +1049,14 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
if file_data is None:
|
||||
raise ValueError("file is required")
|
||||
|
||||
if is_passthrough_batch_upload(create_file_data, litellm_params):
|
||||
return {
|
||||
"streaming_media_upload": StreamingMediaUploadConfig(
|
||||
body_stream=_RawFileUploadStream(file_data),
|
||||
content_type="application/json",
|
||||
)
|
||||
}
|
||||
|
||||
_, content_type = extract_file_metadata(file_data)
|
||||
if FilesAPIUtils.is_batch_jsonl_request(
|
||||
create_file_data=create_file_data,
|
||||
|
|
@ -1164,6 +1237,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
# transformation, e.g. if they consume raw `predictions.jsonl` directly.
|
||||
if getattr(litellm, "disable_vertex_batch_output_transformation", False):
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
if is_passthrough_managed_gcs_url(str(raw_response.request.url)):
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
|
||||
# Try to transform batch output if it's a JSONL file
|
||||
content: Final = raw_response.content
|
||||
|
|
@ -1209,7 +1284,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
Everything else is passed through unchanged, including a row that fails to
|
||||
transform mid-stream.
|
||||
"""
|
||||
if litellm.disable_vertex_batch_output_transformation:
|
||||
if litellm.disable_vertex_batch_output_transformation or is_passthrough_managed_gcs_url(request_url):
|
||||
return FileContentStreamingResult(stream_iterator=stream_iterator, headers=headers)
|
||||
|
||||
first_line, buffered = await _peek_first_jsonl_line(
|
||||
|
|
|
|||
|
|
@ -5442,7 +5442,7 @@
|
|||
"supports_web_search": false
|
||||
},
|
||||
"azure/gpt-4.1-nano": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.5e-08,
|
||||
"input_cost_per_token": 1e-07,
|
||||
"input_cost_per_token_batches": 5e-08,
|
||||
|
|
@ -5476,7 +5476,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-4.1-nano-2025-04-14": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.5e-08,
|
||||
"input_cost_per_token": 1e-07,
|
||||
"input_cost_per_token_batches": 5e-08,
|
||||
|
|
@ -5734,7 +5734,7 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"azure/gpt-audio-mini": {
|
||||
"deprecation_date": "2027-04-06",
|
||||
"deprecation_date": "2027-06-15",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_token": 6e-07,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -6057,7 +6057,7 @@
|
|||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-06-25",
|
||||
"deprecation_date": "2027-07-31",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
"input_cost_per_image_token": 5e-06,
|
||||
"input_cost_per_token": 4e-06,
|
||||
|
|
@ -6269,7 +6269,7 @@
|
|||
"supports_tool_choice": true
|
||||
},
|
||||
"azure/gpt-4o-transcribe": {
|
||||
"deprecation_date": "2026-10-15",
|
||||
"deprecation_date": "2026-12-31",
|
||||
"input_cost_per_audio_token": 2.5e-06,
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -6768,7 +6768,7 @@
|
|||
},
|
||||
"azure/gpt-5-chat": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"deprecation_date": "2026-05-13",
|
||||
"deprecation_date": "2026-06-29",
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 128000,
|
||||
|
|
@ -10796,7 +10796,7 @@
|
|||
"supports_web_search": false
|
||||
},
|
||||
"azure/us/gpt-4.1-nano-2025-04-14": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.8e-08,
|
||||
"input_cost_per_token": 1.1e-07,
|
||||
"input_cost_per_token_batches": 5.5e-08,
|
||||
|
|
@ -26312,7 +26312,11 @@
|
|||
},
|
||||
"gemini-3-pro-image-preview": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
|
|
@ -26322,6 +26326,7 @@
|
|||
"output_cost_per_image": 0.134,
|
||||
"output_cost_per_image_token": 0.00012,
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing",
|
||||
"supported_endpoints": [
|
||||
|
|
@ -26398,7 +26403,10 @@
|
|||
},
|
||||
"gemini-3.1-flash-image-preview": {
|
||||
"input_cost_per_image": 0.00056,
|
||||
"cache_read_input_token_cost": 5e-08,
|
||||
"cache_read_input_token_cost_batches": 2.5e-08,
|
||||
"input_cost_per_token": 5e-07,
|
||||
"input_cost_per_token_batches": 2.5e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
"max_output_tokens": 32768,
|
||||
|
|
@ -26407,6 +26415,7 @@
|
|||
"output_cost_per_image": 0.0672,
|
||||
"output_cost_per_image_token": 6e-05,
|
||||
"output_cost_per_token": 3e-06,
|
||||
"output_cost_per_token_batches": 1.5e-06,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
|
|
@ -41382,7 +41391,6 @@
|
|||
},
|
||||
"openrouter/deepseek/deepseek-v4-pro": {
|
||||
"input_cost_per_token": 9.1263e-07,
|
||||
"input_cost_per_token_cache_hit": 7.60525e-08,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 384000,
|
||||
|
|
@ -49485,7 +49493,11 @@
|
|||
},
|
||||
"vertex_ai/gemini-3-pro-image-preview": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
|
|
@ -49495,6 +49507,7 @@
|
|||
"output_cost_per_image": 0.134,
|
||||
"output_cost_per_image_token": 0.00012,
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"supports_reasoning": false,
|
||||
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
|
||||
|
|
@ -49523,7 +49536,10 @@
|
|||
},
|
||||
"vertex_ai/gemini-3.1-flash-image-preview": {
|
||||
"input_cost_per_image": 0.00056,
|
||||
"cache_read_input_token_cost": 5e-08,
|
||||
"cache_read_input_token_cost_batches": 2.5e-08,
|
||||
"input_cost_per_token": 5e-07,
|
||||
"input_cost_per_token_batches": 2.5e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
"max_output_tokens": 32768,
|
||||
|
|
@ -49532,6 +49548,7 @@
|
|||
"output_cost_per_image": 0.0672,
|
||||
"output_cost_per_image_token": 6e-05,
|
||||
"output_cost_per_token": 3e-06,
|
||||
"output_cost_per_token_batches": 1.5e-06,
|
||||
"supports_reasoning": false,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models"
|
||||
},
|
||||
|
|
@ -68548,6 +68565,7 @@
|
|||
"source": "https://api.together.ai/v1/models"
|
||||
},
|
||||
"vertex_ai/gemini-2.5-flash-native-audio": {
|
||||
"deprecation_date": "2026-12-13",
|
||||
"input_cost_per_audio_token": 3e-06,
|
||||
"input_cost_per_token": 5e-07,
|
||||
"litellm_provider": "vertex_ai",
|
||||
|
|
@ -69100,7 +69118,7 @@
|
|||
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
|
||||
},
|
||||
"azure/eu/gpt-4.1-nano": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.8e-08,
|
||||
"input_cost_per_token": 1.1e-07,
|
||||
"input_cost_per_token_batches": 5.5e-08,
|
||||
|
|
@ -69288,7 +69306,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 5.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 8.25e-05,
|
||||
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
|
||||
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'",
|
||||
"supports_reasoning": true
|
||||
},
|
||||
"azure/eu/gpt-6-luna": {
|
||||
"deprecation_date": "2028-03-11",
|
||||
|
|
@ -69535,7 +69554,7 @@
|
|||
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
|
||||
},
|
||||
"azure/us/gpt-4.1-nano": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.8e-08,
|
||||
"input_cost_per_token": 1.1e-07,
|
||||
"input_cost_per_token_batches": 5.5e-08,
|
||||
|
|
|
|||
|
|
@ -111,6 +111,13 @@ _MCP_DESTINATIONS_SCOPE_KEY: Final = "litellm_otel_request_destinations"
|
|||
_MCP_PROTOCOL_VERSION_HEADER: Final = b"mcp-protocol-version"
|
||||
|
||||
|
||||
def reject_disallowed_mcp_origin(request: StarletteRequest) -> None:
|
||||
from litellm.proxy.proxy_server import origins # noqa: PLC0415 # proxy imports this module during startup
|
||||
|
||||
if "*" not in origins and any(origin not in origins for origin in request.headers.getlist("origin")):
|
||||
raise HTTPException(status_code=403, detail="Invalid Origin header")
|
||||
|
||||
|
||||
def unsupported_protocol_version(scope: Scope) -> str | None:
|
||||
"""Return the unsupported ``MCP-Protocol-Version`` header value, if any.
|
||||
|
||||
|
|
@ -1931,6 +1938,7 @@ if MCP_AVAILABLE:
|
|||
async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
"""Handle MCP requests through StreamableHTTP."""
|
||||
try:
|
||||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
|
|
@ -2275,6 +2283,7 @@ if MCP_AVAILABLE:
|
|||
async def handle_sse_mcp(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
"""Handle MCP requests through SSE."""
|
||||
try:
|
||||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
|
|
|
|||
|
|
@ -307,6 +307,7 @@ class KeyManagementRoutes(str, enum.Enum):
|
|||
# team usage routes
|
||||
TEAM_DAILY_ACTIVITY = "/team/daily/activity"
|
||||
TEAM_DAILY_ACTIVITY_AGGREGATED = "/team/daily/activity/aggregated"
|
||||
TEAM_DAILY_ACTIVITY_AGGREGATED_SEARCH = "/team/daily/activity/aggregated/search"
|
||||
|
||||
# team spend-log viewing
|
||||
SPEND_LOGS = "/spend/logs"
|
||||
|
|
@ -673,6 +674,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED_SEARCH.value,
|
||||
KeyManagementRoutes.SPEND_LOGS.value,
|
||||
KeyManagementRoutes.SPEND_LOGS_V2.value,
|
||||
KeyManagementRoutes.KEY_RESET_SPEND.value,
|
||||
|
|
@ -699,6 +701,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/user/list",
|
||||
"/user/daily/activity",
|
||||
"/user/daily/activity/aggregated",
|
||||
"/user/daily/activity/aggregated/search",
|
||||
# team
|
||||
"/team/new",
|
||||
"/team/update",
|
||||
|
|
@ -716,6 +719,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/team/permissions_bulk_update",
|
||||
"/team/daily/activity",
|
||||
"/team/daily/activity/aggregated",
|
||||
"/team/daily/activity/aggregated/search",
|
||||
"/team/spend/by_user",
|
||||
# gateway request counts (SGR); deployment-wide, admin-only
|
||||
"/gateway/daily/activity",
|
||||
|
|
@ -886,6 +890,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/team/permissions_update",
|
||||
"/team/daily/activity",
|
||||
"/team/daily/activity/aggregated",
|
||||
"/team/daily/activity/aggregated/search",
|
||||
"/team/spend/by_user",
|
||||
"/team/{team_id}/members/me",
|
||||
# POST/GET the team's logging callbacks, and DELETE one of them. Every
|
||||
|
|
@ -901,6 +906,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/model/delete",
|
||||
"/user/daily/activity",
|
||||
"/user/daily/activity/aggregated",
|
||||
"/user/daily/activity/aggregated/search",
|
||||
# Endpoint restricts results to organizations the caller is ORG_ADMIN
|
||||
# of; a caller who administers none gets an empty result set.
|
||||
"/organization/daily/activity",
|
||||
|
|
@ -984,6 +990,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/user/daily/activity",
|
||||
"/team/daily/activity",
|
||||
"/team/daily/activity/aggregated",
|
||||
"/team/daily/activity/aggregated/search",
|
||||
"/tag/daily/activity",
|
||||
"/tag/list",
|
||||
"/audit",
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
|
|
@ -87,11 +88,13 @@ from litellm.repositories.verification_token_repository import (
|
|||
VerificationTokenRepository,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
DailySpendMetadata,
|
||||
SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
|
||||
BulkUpdateUserRequest,
|
||||
BulkUpdateUserResponse,
|
||||
KeyActivitySearchWhere,
|
||||
UserListResponse,
|
||||
UserSearchWhere,
|
||||
UserUpdateResult,
|
||||
|
|
@ -2991,6 +2994,27 @@ async def get_user_daily_activity(
|
|||
)
|
||||
|
||||
|
||||
def _resolve_user_daily_activity_entity_id(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
user_id: str | None,
|
||||
) -> str | None:
|
||||
is_admin: Final = _user_has_admin_view(user_api_key_dict)
|
||||
|
||||
if is_admin:
|
||||
return user_id
|
||||
|
||||
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
|
||||
effective_user_id: Final = user_id if user_id is not None else caller_user_id
|
||||
if effective_user_id != caller_user_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={ # mutable-ok: FastAPI detail payload shape
|
||||
"error": "Non-admin users can only view their own spend data."
|
||||
},
|
||||
)
|
||||
return effective_user_id
|
||||
|
||||
|
||||
@router.get(
|
||||
"/user/daily/activity/aggregated",
|
||||
tags=["Budget & Spend Tracking", "Internal User management"],
|
||||
|
|
@ -3057,20 +3081,7 @@ async def get_user_daily_activity_aggregated(
|
|||
)
|
||||
|
||||
try:
|
||||
is_admin: Final = _user_has_admin_view(user_api_key_dict)
|
||||
|
||||
if is_admin:
|
||||
entity_id = user_id # None means global view, otherwise filter by user
|
||||
else:
|
||||
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
|
||||
if user_id is None:
|
||||
user_id = caller_user_id
|
||||
if user_id != caller_user_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={"error": "Non-admin users can only view their own spend data."},
|
||||
)
|
||||
entity_id = user_id
|
||||
entity_id: Final = _resolve_user_daily_activity_entity_id(user_api_key_dict, user_id)
|
||||
|
||||
return await get_daily_activity_aggregated(
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -3094,3 +3105,117 @@ async def get_user_daily_activity_aggregated(
|
|||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": f"Failed to fetch analytics: {e}"},
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/user/daily/activity/aggregated/search",
|
||||
tags=["Budget & Spend Tracking", "Internal User management"], # mutable-ok: FastAPI route tags shape
|
||||
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI route dependencies shape
|
||||
response_model=SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def search_user_daily_activity_keys(
|
||||
search: str = fastapi.Query(
|
||||
...,
|
||||
min_length=1,
|
||||
description="Matches keys whose hash equals the value, or whose key alias or user ID contains it (case-insensitive)",
|
||||
),
|
||||
start_date: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Start date in YYYY-MM-DD format",
|
||||
),
|
||||
end_date: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="End date in YYYY-MM-DD format",
|
||||
),
|
||||
user_id: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter by specific user ID. Admins can filter by any user or omit for global view. Non-admins must provide their own user_id.",
|
||||
),
|
||||
timezone: int | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Timezone offset in minutes from UTC (e.g., 480 for PST). "
|
||||
"Matches JavaScript's Date.getTimezoneOffset() convention.",
|
||||
),
|
||||
include_current_utc_day: bool = fastapi.Query(
|
||||
default=False,
|
||||
description="When the range ends on the caller's current local day, extend it to "
|
||||
"today's UTC bucket so spend written after the caller's local midnight (in UTC "
|
||||
"terms) is included. Requires the timezone parameter. Historical ranges are "
|
||||
"never extended.",
|
||||
),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI dependency injection
|
||||
) -> SpendAnalyticsPaginatedResponse:
|
||||
"""
|
||||
Search verification tokens by exact token hash or by a case-insensitive substring of
|
||||
the key alias or owning user ID, then return the aggregated daily activity for the
|
||||
matches. Lets the Usage page surface keys that fell outside the top-spend subset
|
||||
the aggregated endpoint loads.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={ # mutable-ok: FastAPI detail payload shape
|
||||
"error": CommonProxyErrors.db_not_connected_error.value
|
||||
},
|
||||
)
|
||||
|
||||
if start_date is None or end_date is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail={"error": "Please provide start_date and end_date"}, # mutable-ok: FastAPI detail payload shape
|
||||
)
|
||||
|
||||
try:
|
||||
entity_id: Final = _resolve_user_daily_activity_entity_id(user_api_key_dict, user_id)
|
||||
|
||||
search_or: Final = (
|
||||
{"token": search}, # mutable-ok: prisma serializes where clauses, keep plain dicts
|
||||
{"key_alias": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
|
||||
{"user_id": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
|
||||
)
|
||||
where: Final[KeyActivitySearchWhere] = (
|
||||
{"OR": search_or} # mutable-ok: prisma where clause root
|
||||
if entity_id is None
|
||||
else {"user_id": entity_id, "OR": search_or} # mutable-ok: prisma where clause root
|
||||
)
|
||||
matched_keys: Final = await VerificationTokenRepository(prisma_client).table.find_many(
|
||||
where=where,
|
||||
take=USAGE_TOP_API_KEYS_LIMIT,
|
||||
order={"spend": "desc"}, # mutable-ok: prisma serializes order, keep it a plain dict
|
||||
)
|
||||
tokens: Final = [key.token for key in matched_keys] # mutable-ok: api_key filter union expects a list
|
||||
|
||||
if not tokens:
|
||||
return SpendAnalyticsPaginatedResponse(
|
||||
results=[], # mutable-ok: response model field shape
|
||||
metadata=DailySpendMetadata(
|
||||
api_key_limit=USAGE_TOP_API_KEYS_LIMIT,
|
||||
total_api_keys=0,
|
||||
),
|
||||
)
|
||||
|
||||
return await get_daily_activity_aggregated(
|
||||
prisma_client=prisma_client,
|
||||
table_name="litellm_dailyuserspend",
|
||||
entity_id_field="user_id",
|
||||
entity_id=entity_id,
|
||||
entity_metadata_field=None,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
model=None,
|
||||
api_key=tokens,
|
||||
timezone_offset_minutes=timezone,
|
||||
include_current_utc_day=include_current_utc_day,
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("/user/daily/activity/aggregated/search: Exception occured - %s", e)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": f"Failed to fetch analytics: {e}"}, # mutable-ok: FastAPI detail payload shape
|
||||
)
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ from typing_extensions import ReadOnly, TypedDict, assert_never
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -196,6 +197,7 @@ from litellm.repositories.verification_token_repository import (
|
|||
from litellm.router import Router
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
DailySpendMetadata,
|
||||
SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
|
|
@ -204,7 +206,9 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
BulkUpdateTeamMemberPermissionsResponse,
|
||||
GetTeamMemberPermissionsResponse,
|
||||
TeamIdSearchFilter,
|
||||
TeamIdSearchMatch,
|
||||
TeamKeyActivitySearchWhere,
|
||||
TeamListItem,
|
||||
TeamListResponse,
|
||||
TeamMemberAddResult,
|
||||
|
|
@ -6805,6 +6809,111 @@ async def get_team_daily_activity_aggregated(
|
|||
)
|
||||
|
||||
|
||||
def _team_key_search_where(*, search: str, scope: _TeamDailyActivityScope) -> TeamKeyActivitySearchWhere:
|
||||
"""Caller scoping lives inside the same Prisma where as the search term so `take`
|
||||
never trims visible matches in favour of keys the caller is not allowed to see."""
|
||||
search_or: Final = (
|
||||
{"token": search}, # mutable-ok: prisma where clause leaf
|
||||
{"key_alias": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
|
||||
{"user_id": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
|
||||
)
|
||||
own_keys: Final = tuple(scope.api_key_filter) if isinstance(scope.api_key_filter, list) else None
|
||||
team_filter: Final[TeamIdSearchFilter | None] = (
|
||||
{ # mutable-ok: prisma where clause leaf
|
||||
"in": tuple(scope.team_ids),
|
||||
"notIn": tuple(scope.exclude_team_ids),
|
||||
}
|
||||
if scope.team_ids is not None and scope.exclude_team_ids is not None
|
||||
else {"in": tuple(scope.team_ids)} # mutable-ok: prisma where clause leaf
|
||||
if scope.team_ids is not None
|
||||
else {"notIn": tuple(scope.exclude_team_ids)} # mutable-ok: prisma where clause leaf
|
||||
if scope.exclude_team_ids is not None
|
||||
else None
|
||||
)
|
||||
if team_filter is None and own_keys is None:
|
||||
return {"OR": search_or} # mutable-ok: prisma where clause root
|
||||
if team_filter is None and own_keys is not None:
|
||||
return {"token": {"in": own_keys}, "OR": search_or} # mutable-ok: prisma where clause root
|
||||
if team_filter is not None and own_keys is None:
|
||||
return {"team_id": team_filter, "OR": search_or} # mutable-ok: prisma where clause root
|
||||
assert team_filter is not None and own_keys is not None
|
||||
return { # mutable-ok: prisma where clause root
|
||||
"team_id": team_filter,
|
||||
"token": {"in": own_keys}, # mutable-ok: prisma where clause leaf
|
||||
"OR": search_or,
|
||||
}
|
||||
|
||||
|
||||
@router.get(
|
||||
"/team/daily/activity/aggregated/search",
|
||||
response_model=SpendAnalyticsPaginatedResponse,
|
||||
tags=["team management"], # mutable-ok: FastAPI route tags shape
|
||||
)
|
||||
async def search_team_daily_activity_keys(
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
search: str = fastapi.Query(
|
||||
...,
|
||||
min_length=1,
|
||||
description="Exact token hash, or a case-insensitive substring of the key alias or owning user id",
|
||||
),
|
||||
team_ids: str | None = None,
|
||||
start_date: str | None = None,
|
||||
end_date: str | None = None,
|
||||
exclude_team_ids: str | None = None,
|
||||
timezone: int | None = None,
|
||||
) -> SpendAnalyticsPaginatedResponse:
|
||||
"""Aggregated daily team activity for the keys matching `search`, across every key the caller may
|
||||
see rather than only the top USAGE_TOP_API_KEYS_LIMIT keys by spend."""
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
range_error: Final = _aggregated_date_range_error(start_date, end_date)
|
||||
if range_error is not None:
|
||||
raise _daily_activity_error(status_code=400, message=range_error)
|
||||
|
||||
scope: Final = await _resolve_team_daily_activity_scope(
|
||||
team_ids=team_ids,
|
||||
exclude_team_ids=exclude_team_ids,
|
||||
api_key=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
matched_keys: Final = await _tokens_db(prisma_client).find_many(
|
||||
where=_team_key_search_where(search=search, scope=scope),
|
||||
take=USAGE_TOP_API_KEYS_LIMIT,
|
||||
order={"spend": "desc"}, # mutable-ok: prisma serializes order, keep it a plain dict
|
||||
)
|
||||
tokens: Final = [key.token for key in matched_keys] # mutable-ok: get_daily_activity_aggregated takes list[str]
|
||||
if not tokens:
|
||||
return SpendAnalyticsPaginatedResponse(
|
||||
results=[], # mutable-ok: response model field shape
|
||||
metadata=DailySpendMetadata(api_key_limit=USAGE_TOP_API_KEYS_LIMIT, total_api_keys=0),
|
||||
)
|
||||
|
||||
return await get_daily_activity_aggregated(
|
||||
prisma_client=prisma_client,
|
||||
table_name="litellm_dailyteamspend",
|
||||
entity_id_field="team_id",
|
||||
entity_id=scope.team_ids,
|
||||
entity_metadata_field=scope.team_alias_metadata,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
model=None,
|
||||
api_key=tokens,
|
||||
exclude_entity_ids=scope.exclude_team_ids,
|
||||
timezone_offset_minutes=timezone,
|
||||
include_entity_breakdown=True,
|
||||
)
|
||||
|
||||
|
||||
def _team_user_spend_sql(*, team_count: int, restrict_to_user: bool) -> str:
|
||||
team_placeholders: Final = ", ".join(f"${i}" for i in range(3, 3 + team_count))
|
||||
user_clause: Final = f' AND sl."user" = ${3 + team_count}' if restrict_to_user else ""
|
||||
|
|
|
|||
|
|
@ -8,10 +8,25 @@ from typing_extensions import assert_never
|
|||
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
BATCH_LINE_REQUIRED_KEYS: Final = ("custom_id", "method", "url", "body")
|
||||
_MB: Final = 1024 * 1024
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BatchLineShape:
|
||||
required_keys: tuple[str, ...]
|
||||
hint: str
|
||||
|
||||
|
||||
BATCH_LINE_SHAPE: Final = BatchLineShape(
|
||||
required_keys=("custom_id", "method", "url", "body"),
|
||||
hint="Each line must be a JSON object with keys custom_id, method, url, body",
|
||||
)
|
||||
PASSTHROUGH_BATCH_LINE_SHAPE: Final = BatchLineShape(
|
||||
required_keys=("request",),
|
||||
hint="A passthrough upload takes native Vertex batch rows, so each line must be a JSON object with a request key",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BatchFileTooLarge:
|
||||
size_bytes: int
|
||||
|
|
@ -42,6 +57,7 @@ class BatchFileLineNotObject:
|
|||
class BatchFileMissingLineKey:
|
||||
line_number: int
|
||||
key: str
|
||||
line_shape: BatchLineShape = BATCH_LINE_SHAPE
|
||||
|
||||
|
||||
BatchFileValidationFailure = (
|
||||
|
|
@ -70,20 +86,20 @@ def _iter_lines(file_source: bytes | BinaryIO) -> Iterator[bytes]:
|
|||
return iter(file_source)
|
||||
|
||||
|
||||
def _check_line(line_number: int, raw_line: bytes) -> BatchFileValidationFailure | None:
|
||||
def _check_line(line_number: int, raw_line: bytes, line_shape: BatchLineShape) -> BatchFileValidationFailure | None:
|
||||
try:
|
||||
parsed: Final = json.loads(raw_line)
|
||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||
return BatchFileInvalidJsonLine(line_number=line_number)
|
||||
if not isinstance(parsed, dict):
|
||||
return BatchFileLineNotObject(line_number=line_number)
|
||||
missing: Final = next((key for key in BATCH_LINE_REQUIRED_KEYS if key not in parsed), None)
|
||||
missing: Final = next((key for key in line_shape.required_keys if key not in parsed), None)
|
||||
if missing is None:
|
||||
return None
|
||||
return BatchFileMissingLineKey(line_number=line_number, key=missing)
|
||||
return BatchFileMissingLineKey(line_number=line_number, key=missing, line_shape=line_shape)
|
||||
|
||||
|
||||
def _scan_lines(file_source: bytes | BinaryIO) -> BatchFileValidationFailure | None:
|
||||
def _scan_lines(file_source: bytes | BinaryIO, line_shape: BatchLineShape) -> BatchFileValidationFailure | None:
|
||||
content_lines: Final = (
|
||||
(line_number, raw_line)
|
||||
for line_number, raw_line in enumerate(_iter_lines(file_source), start=1)
|
||||
|
|
@ -96,7 +112,7 @@ def _scan_lines(file_source: bytes | BinaryIO) -> BatchFileValidationFailure | N
|
|||
(
|
||||
failure
|
||||
for line_number, raw_line in chain((first_line,), content_lines)
|
||||
for failure in (_check_line(line_number, raw_line),)
|
||||
for failure in (_check_line(line_number, raw_line, line_shape),)
|
||||
if failure is not None
|
||||
),
|
||||
None,
|
||||
|
|
@ -107,6 +123,7 @@ def check_batch_file_upload(
|
|||
filename: str | None,
|
||||
file_source: bytes | BinaryIO,
|
||||
max_batch_file_size_mb: int | None,
|
||||
line_shape: BatchLineShape = BATCH_LINE_SHAPE,
|
||||
) -> BatchFileValidationFailure | None:
|
||||
if filename is None or not filename.lower().endswith(".jsonl"):
|
||||
return BatchFileWrongExtension(filename=filename or "")
|
||||
|
|
@ -114,7 +131,7 @@ def check_batch_file_upload(
|
|||
size_bytes: Final = _file_size_bytes(file_source)
|
||||
if size_bytes > max_batch_file_size_mb * _MB:
|
||||
return BatchFileTooLarge(size_bytes=size_bytes, limit_mb=max_batch_file_size_mb)
|
||||
scan_failure: Final = _scan_lines(file_source)
|
||||
scan_failure: Final = _scan_lines(file_source, line_shape)
|
||||
if not isinstance(file_source, bytes):
|
||||
file_source.seek(0)
|
||||
return scan_failure
|
||||
|
|
@ -169,11 +186,11 @@ def raise_batch_file_validation_failure(failure: BatchFileValidationFailure) ->
|
|||
param="file",
|
||||
code=400,
|
||||
)
|
||||
case BatchFileMissingLineKey(line_number=line_number, key=key):
|
||||
case BatchFileMissingLineKey(line_number=line_number, key=key, line_shape=line_shape):
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"Missing required parameter: '{key}' (batch input file line {line_number}). "
|
||||
f"Each line must be a JSON object with keys {', '.join(BATCH_LINE_REQUIRED_KEYS)}. "
|
||||
f"{line_shape.hint}. "
|
||||
"The file was not forwarded to the provider."
|
||||
),
|
||||
type="invalid_request_error",
|
||||
|
|
|
|||
|
|
@ -57,6 +57,8 @@ from litellm.proxy.common_utils.openai_error_payload import (
|
|||
openai_error_type,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.batch_file_validation import (
|
||||
BATCH_LINE_SHAPE,
|
||||
PASSTHROUGH_BATCH_LINE_SHAPE,
|
||||
check_batch_file_upload,
|
||||
raise_batch_file_validation_failure,
|
||||
)
|
||||
|
|
@ -207,10 +209,91 @@ def get_files_provider_config(
|
|||
return None
|
||||
|
||||
|
||||
def _deployment_provider(llm_router: Router, model_id: str, team_id: str | None) -> str | None:
|
||||
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id)
|
||||
return None if credentials is None else credentials.get("custom_llm_provider")
|
||||
|
||||
|
||||
def _resolves_to_vertex_deployments_only(llm_router: Router | None, model_name: str, team_id: str | None) -> bool:
|
||||
if llm_router is None or _deployment_provider(llm_router, model_name, team_id) != "vertex_ai":
|
||||
return False
|
||||
return all(
|
||||
_deployment_provider(llm_router, str(deployment["model_info"]["id"]), team_id) == "vertex_ai"
|
||||
for deployment in llm_router.get_model_list(model_name=model_name, team_id=team_id) or ()
|
||||
if "id" in deployment.get("model_info", {})
|
||||
)
|
||||
|
||||
|
||||
def _validate_passthrough_upload(
|
||||
*,
|
||||
purpose: str,
|
||||
target_model_names: Sequence[str],
|
||||
model: str | None,
|
||||
target_storage: str | None,
|
||||
llm_router: Router | None,
|
||||
team_id: str | None,
|
||||
) -> None:
|
||||
if purpose != "batch":
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"`passthrough` uploads the file bytes unchanged for a native Vertex batch, "
|
||||
f"so purpose must be 'batch', got '{purpose}'."
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="passthrough",
|
||||
code=400,
|
||||
)
|
||||
if target_storage and target_storage != "default":
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"`passthrough` writes the native batch file to the Vertex AI deployment's GCS bucket, "
|
||||
f"so it cannot be combined with target_storage='{target_storage}'."
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="target_storage",
|
||||
code=400,
|
||||
)
|
||||
named_deployments: Final = (
|
||||
*(("target_model_names", name) for name in target_model_names),
|
||||
*((("model", model),) if model else ()),
|
||||
)
|
||||
if not named_deployments:
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"`passthrough` needs the Vertex AI deployment that will run the batch, "
|
||||
"since native rows carry no model: pass `target_model_names` or `model`."
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="target_model_names",
|
||||
code=400,
|
||||
)
|
||||
offending: Final = next(
|
||||
(
|
||||
(param, name)
|
||||
for param, name in named_deployments
|
||||
if not _resolves_to_vertex_deployments_only(llm_router, name, team_id)
|
||||
),
|
||||
None,
|
||||
)
|
||||
if offending is None:
|
||||
return
|
||||
param, name = offending
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"`passthrough` is only supported for Vertex AI deployments; '{name}' does not resolve "
|
||||
"to vertex_ai deployments only."
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param=param,
|
||||
code=400,
|
||||
)
|
||||
|
||||
|
||||
async def _scan_batch_upload(
|
||||
*,
|
||||
file_source: bytes | BinaryIO,
|
||||
purpose: str,
|
||||
passthrough: bool,
|
||||
request_metadata: Mapping[str, object],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
|
|
@ -222,6 +305,17 @@ async def _scan_batch_upload(
|
|||
or not proxy_logging_obj.has_pre_call_guardrails(request_metadata)
|
||||
):
|
||||
return None
|
||||
if passthrough:
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"Batch guardrails cannot scan native Vertex batch rows, so a `passthrough` upload is refused "
|
||||
"when the key, team, or request has pre-call guardrails configured. "
|
||||
"The file was not forwarded to the provider."
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="passthrough",
|
||||
code=400,
|
||||
)
|
||||
outcome: Final = await scan_batch_input_file(
|
||||
file_source=file_source,
|
||||
request_metadata=request_metadata,
|
||||
|
|
@ -458,6 +552,7 @@ async def create_file(
|
|||
custom_llm_provider: str = Form(default="openai"),
|
||||
file: UploadFile = File(...),
|
||||
litellm_metadata: str | None = Form(default=None),
|
||||
passthrough: bool = Form(default=False),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
|
|
@ -560,17 +655,28 @@ async def create_file(
|
|||
if blocked_extension_failure is not None:
|
||||
raise_upload_validation_failure(blocked_extension_failure)
|
||||
|
||||
if passthrough:
|
||||
_validate_passthrough_upload(
|
||||
purpose=purpose,
|
||||
target_model_names=target_model_names_list,
|
||||
model=model_param,
|
||||
target_storage=target_storage,
|
||||
llm_router=llm_router,
|
||||
team_id=user_api_key_dict.team_id,
|
||||
)
|
||||
|
||||
if purpose == "batch":
|
||||
batch_file_failure: Final = await asyncio.to_thread(
|
||||
check_batch_file_upload,
|
||||
file.filename,
|
||||
file_source,
|
||||
_MAX_BATCH_FILE_SIZE_MB_ADAPTER.validate_python(general_settings.get("max_batch_file_size_mb")),
|
||||
PASSTHROUGH_BATCH_LINE_SHAPE if passthrough else BATCH_LINE_SHAPE,
|
||||
)
|
||||
if batch_file_failure is not None:
|
||||
raise_batch_file_validation_failure(batch_file_failure)
|
||||
|
||||
data = {}
|
||||
data = {"passthrough": True} if passthrough else {}
|
||||
|
||||
# Parse expires_after if provided
|
||||
expires_after: FileExpiresAfter | None = None
|
||||
|
|
@ -673,6 +779,7 @@ async def create_file(
|
|||
scan_result: Final = await _scan_batch_upload(
|
||||
file_source=file_source,
|
||||
purpose=purpose,
|
||||
passthrough=passthrough,
|
||||
request_metadata=request_metadata,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
|
|
|
|||
|
|
@ -5980,6 +5980,7 @@ class Router:
|
|||
|
||||
replace_model_in_jsonl_bool: Final = should_replace_model_in_jsonl(
|
||||
purpose=purpose,
|
||||
passthrough=kwargs.get("passthrough") is True,
|
||||
)
|
||||
if replace_model_in_jsonl_bool:
|
||||
file = replace_model_in_jsonl(
|
||||
|
|
|
|||
|
|
@ -62,15 +62,15 @@ def parse_jsonl_with_embedded_newlines(content: str) -> list[dict]:
|
|||
|
||||
def should_replace_model_in_jsonl(
|
||||
purpose: OpenAIFilesPurpose,
|
||||
passthrough: bool = False,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if the model name should be replaced in the JSONL file for the deployment model name.
|
||||
|
||||
Azure raises an error on create batch if the model name for deployment is not in the .jsonl.
|
||||
A passthrough upload keeps the caller's bytes untouched, so its rows are never rewritten.
|
||||
"""
|
||||
if purpose == "batch":
|
||||
return True
|
||||
return False
|
||||
return purpose == "batch" and not passthrough
|
||||
|
||||
|
||||
def replace_model_in_jsonl(file_content: FileTypes, new_model_name: str) -> FileTypes:
|
||||
|
|
|
|||
|
|
@ -1,12 +1,50 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import cast # noqa: TID251 # narrows the normalized native payload to the public TypedDict
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Final, cast # noqa: TID251 # narrows the normalized native payload to the public TypedDict
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
|
||||
from litellm.llms.anthropic.experimental_pass_through.utils import is_reasoning_auto_summary_enabled
|
||||
from litellm.rust_bridge import failures
|
||||
from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
|
||||
|
||||
_DROP_PATHS: Final = TypeAdapter(list[object])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EffortTiers:
|
||||
minimal: bool
|
||||
low: bool
|
||||
medium: bool
|
||||
high: bool
|
||||
xhigh: bool
|
||||
max: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ModelCapabilities:
|
||||
supports_reasoning: bool
|
||||
supports_adaptive_thinking: bool
|
||||
thinking_always_on: bool
|
||||
supports_legacy_thinking: bool
|
||||
supports_output_config: bool
|
||||
supports_sampling_params: bool
|
||||
supports_speed: bool
|
||||
effort_tiers: EffortTiers
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MessagesShaping:
|
||||
capabilities: ModelCapabilities
|
||||
drop_params: bool
|
||||
reasoning_auto_summary: bool
|
||||
additional_drop_params: Sequence[str]
|
||||
|
||||
|
||||
def response(value: Mapping[str, object]) -> AnthropicMessagesResponse:
|
||||
return cast( # cast-ok: AnthropicMessagesResponse is a TypedDict over the normalized native payload
|
||||
|
|
@ -20,4 +58,72 @@ def arguments(request: LiteLLMMessagesRequest) -> Mapping[str, object]:
|
|||
|
||||
|
||||
def map_failure(error: Exception, request: LiteLLMMessagesRequest, request_provider: str) -> Exception:
|
||||
if getattr(error, "messages_request_error", False):
|
||||
return litellm.BadRequestError(
|
||||
message=str(error),
|
||||
model=request.model.removeprefix(f"{request_provider}/"),
|
||||
llm_provider=request_provider,
|
||||
)
|
||||
return failures.map_native_failure(error, request.model, request_provider, arguments(request), request.api_base)
|
||||
|
||||
|
||||
def _resolved_provider(model: str, custom_llm_provider: str | None) -> tuple[str, str]:
|
||||
try:
|
||||
resolved_model, provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
except Exception: # noqa: BLE001 # an unroutable model still shapes as a bare Anthropic id
|
||||
return model, custom_llm_provider or "anthropic"
|
||||
return resolved_model, provider
|
||||
|
||||
|
||||
def model_capabilities(model: str, custom_llm_provider: str | None) -> ModelCapabilities:
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
resolved_model, provider = _resolved_provider(model, custom_llm_provider)
|
||||
|
||||
def supports(flag: str) -> bool:
|
||||
return AnthropicModelInfo._supports_model_capability(model, flag, provider) # pyright: ignore[reportPrivateUsage] # same probes the Python transform runs; forking them would drift
|
||||
|
||||
def tier(level: str) -> bool:
|
||||
return AnthropicConfig._supports_effort_level(model, level, provider) # pyright: ignore[reportPrivateUsage] # same probe the Python transform runs
|
||||
|
||||
return ModelCapabilities(
|
||||
supports_reasoning=supports("supports_reasoning"),
|
||||
supports_adaptive_thinking=supports("supports_adaptive_thinking"),
|
||||
thinking_always_on=supports("thinking_always_on"),
|
||||
supports_legacy_thinking=supports("supports_legacy_thinking"),
|
||||
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
|
||||
effort_tiers=EffortTiers(
|
||||
minimal=tier("minimal"),
|
||||
low=tier("low"),
|
||||
medium=tier("medium"),
|
||||
high=tier("high"),
|
||||
xhigh=tier("xhigh"),
|
||||
max=tier("max"),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _drop_params(kwargs: Mapping[str, object]) -> bool:
|
||||
return bool(litellm.drop_params) or normalize_drop_params(kwargs.get("drop_params")) is True
|
||||
|
||||
|
||||
def _additional_drop_params(kwargs: Mapping[str, object]) -> tuple[str, ...]:
|
||||
try:
|
||||
configured: Final = _DROP_PATHS.validate_python(kwargs.get("additional_drop_params"))
|
||||
except ValidationError:
|
||||
return ()
|
||||
return tuple(path for path in configured if isinstance(path, str))
|
||||
|
||||
|
||||
def shaping(model: str, custom_llm_provider: str | None, kwargs: Mapping[str, object]) -> dict[str, object]:
|
||||
return asdict(
|
||||
MessagesShaping(
|
||||
capabilities=model_capabilities(model, custom_llm_provider),
|
||||
drop_params=_drop_params(kwargs),
|
||||
reasoning_auto_summary=is_reasoning_auto_summary_enabled(),
|
||||
additional_drop_params=_additional_drop_params(kwargs),
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -664,6 +664,7 @@ class PrometheusMetricLabels:
|
|||
]
|
||||
|
||||
litellm_deployment_tpm_limit = [
|
||||
UserAPIKeyLabelNames.MODEL_GROUP.value,
|
||||
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
|
||||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
UserAPIKeyLabelNames.API_BASE.value,
|
||||
|
|
@ -770,6 +771,7 @@ class PrometheusMetricLabels:
|
|||
|
||||
# Add deployment metrics
|
||||
litellm_deployment_failure_responses = [
|
||||
UserAPIKeyLabelNames.MODEL_GROUP.value,
|
||||
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
|
||||
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
|
||||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
|
|
@ -786,6 +788,7 @@ class PrometheusMetricLabels:
|
|||
]
|
||||
|
||||
litellm_deployment_total_requests = [
|
||||
UserAPIKeyLabelNames.MODEL_GROUP.value,
|
||||
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
|
||||
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
|
||||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ from collections.abc import Mapping, Sequence
|
|||
from typing import Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_UserTableWithKeyCount,
|
||||
|
|
@ -28,6 +28,16 @@ class UserSearchWhere(TypedDict):
|
|||
OR: ReadOnly[tuple[Mapping[Literal["user_id", "user_email"], InsensitiveContains], ...]]
|
||||
|
||||
|
||||
class KeyActivitySearchWhere(TypedDict):
|
||||
"""Prisma filter behind `/user/daily/activity/aggregated/search`: exact token hash, or key alias
|
||||
or user id containing the term, case-insensitive."""
|
||||
|
||||
user_id: NotRequired[ReadOnly[str]]
|
||||
OR: ReadOnly[
|
||||
tuple[Mapping[Literal["token"], str] | Mapping[Literal["key_alias", "user_id"], InsensitiveContains], ...]
|
||||
]
|
||||
|
||||
|
||||
class UserListResponse(BaseModel):
|
||||
"""
|
||||
Response model for the user list endpoint
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy._types import (
|
||||
KeyManagementRoutes,
|
||||
|
|
@ -11,10 +13,32 @@ from litellm.proxy._types import (
|
|||
MemberDeleteRequest,
|
||||
)
|
||||
from litellm.proxy.common_utils.timezone_utils import budget_duration_error
|
||||
from litellm.types.proxy.management_endpoints.internal_user_endpoints import InsensitiveContains
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import ResourceResponse
|
||||
|
||||
TeamIdSearchMatch = Literal["exact", "prefix"]
|
||||
|
||||
|
||||
TeamIdSearchFilter = TypedDict(
|
||||
"TeamIdSearchFilter",
|
||||
{ # mutable-ok: functional TypedDict field map
|
||||
"in": NotRequired[ReadOnly[Sequence[str]]],
|
||||
"notIn": NotRequired[ReadOnly[Sequence[str]]],
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class TeamKeyActivitySearchWhere(TypedDict):
|
||||
"""Prisma filter behind `/team/daily/activity/aggregated/search`: exact token hash, or key alias
|
||||
or user id containing the term, case-insensitive, narrowed to the teams and keys the caller may see."""
|
||||
|
||||
team_id: NotRequired[ReadOnly[TeamIdSearchFilter]]
|
||||
token: NotRequired[ReadOnly[Mapping[Literal["in"], Sequence[str]]]]
|
||||
OR: ReadOnly[
|
||||
tuple[Mapping[Literal["token"], str] | Mapping[Literal["key_alias", "user_id"], InsensitiveContains], ...]
|
||||
]
|
||||
|
||||
|
||||
MAX_BULK_TEAM_MEMBER_DELETES: Final = 500
|
||||
|
||||
MAX_BULK_TEAM_MEMBER_BUDGET_UPDATES: Final = 500
|
||||
|
|
|
|||
|
|
@ -5442,7 +5442,7 @@
|
|||
"supports_web_search": false
|
||||
},
|
||||
"azure/gpt-4.1-nano": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.5e-08,
|
||||
"input_cost_per_token": 1e-07,
|
||||
"input_cost_per_token_batches": 5e-08,
|
||||
|
|
@ -5476,7 +5476,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-4.1-nano-2025-04-14": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.5e-08,
|
||||
"input_cost_per_token": 1e-07,
|
||||
"input_cost_per_token_batches": 5e-08,
|
||||
|
|
@ -5734,7 +5734,7 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"azure/gpt-audio-mini": {
|
||||
"deprecation_date": "2027-04-06",
|
||||
"deprecation_date": "2027-06-15",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_token": 6e-07,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -6057,7 +6057,7 @@
|
|||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-06-25",
|
||||
"deprecation_date": "2027-07-31",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
"input_cost_per_image_token": 5e-06,
|
||||
"input_cost_per_token": 4e-06,
|
||||
|
|
@ -6269,7 +6269,7 @@
|
|||
"supports_tool_choice": true
|
||||
},
|
||||
"azure/gpt-4o-transcribe": {
|
||||
"deprecation_date": "2026-10-15",
|
||||
"deprecation_date": "2026-12-31",
|
||||
"input_cost_per_audio_token": 2.5e-06,
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -6768,7 +6768,7 @@
|
|||
},
|
||||
"azure/gpt-5-chat": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"deprecation_date": "2026-05-13",
|
||||
"deprecation_date": "2026-06-29",
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 128000,
|
||||
|
|
@ -10796,7 +10796,7 @@
|
|||
"supports_web_search": false
|
||||
},
|
||||
"azure/us/gpt-4.1-nano-2025-04-14": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.8e-08,
|
||||
"input_cost_per_token": 1.1e-07,
|
||||
"input_cost_per_token_batches": 5.5e-08,
|
||||
|
|
@ -26312,7 +26312,11 @@
|
|||
},
|
||||
"gemini-3-pro-image-preview": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
|
|
@ -26322,6 +26326,7 @@
|
|||
"output_cost_per_image": 0.134,
|
||||
"output_cost_per_image_token": 0.00012,
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing",
|
||||
"supported_endpoints": [
|
||||
|
|
@ -26398,7 +26403,10 @@
|
|||
},
|
||||
"gemini-3.1-flash-image-preview": {
|
||||
"input_cost_per_image": 0.00056,
|
||||
"cache_read_input_token_cost": 5e-08,
|
||||
"cache_read_input_token_cost_batches": 2.5e-08,
|
||||
"input_cost_per_token": 5e-07,
|
||||
"input_cost_per_token_batches": 2.5e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
"max_output_tokens": 32768,
|
||||
|
|
@ -26407,6 +26415,7 @@
|
|||
"output_cost_per_image": 0.0672,
|
||||
"output_cost_per_image_token": 6e-05,
|
||||
"output_cost_per_token": 3e-06,
|
||||
"output_cost_per_token_batches": 1.5e-06,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
|
|
@ -41382,7 +41391,6 @@
|
|||
},
|
||||
"openrouter/deepseek/deepseek-v4-pro": {
|
||||
"input_cost_per_token": 9.1263e-07,
|
||||
"input_cost_per_token_cache_hit": 7.60525e-08,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 384000,
|
||||
|
|
@ -49485,7 +49493,11 @@
|
|||
},
|
||||
"vertex_ai/gemini-3-pro-image-preview": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
|
|
@ -49495,6 +49507,7 @@
|
|||
"output_cost_per_image": 0.134,
|
||||
"output_cost_per_image_token": 0.00012,
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"supports_reasoning": false,
|
||||
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
|
||||
|
|
@ -49523,7 +49536,10 @@
|
|||
},
|
||||
"vertex_ai/gemini-3.1-flash-image-preview": {
|
||||
"input_cost_per_image": 0.00056,
|
||||
"cache_read_input_token_cost": 5e-08,
|
||||
"cache_read_input_token_cost_batches": 2.5e-08,
|
||||
"input_cost_per_token": 5e-07,
|
||||
"input_cost_per_token_batches": 2.5e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
"max_output_tokens": 32768,
|
||||
|
|
@ -49532,6 +49548,7 @@
|
|||
"output_cost_per_image": 0.0672,
|
||||
"output_cost_per_image_token": 6e-05,
|
||||
"output_cost_per_token": 3e-06,
|
||||
"output_cost_per_token_batches": 1.5e-06,
|
||||
"supports_reasoning": false,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models"
|
||||
},
|
||||
|
|
@ -68548,6 +68565,7 @@
|
|||
"source": "https://api.together.ai/v1/models"
|
||||
},
|
||||
"vertex_ai/gemini-2.5-flash-native-audio": {
|
||||
"deprecation_date": "2026-12-13",
|
||||
"input_cost_per_audio_token": 3e-06,
|
||||
"input_cost_per_token": 5e-07,
|
||||
"litellm_provider": "vertex_ai",
|
||||
|
|
@ -69100,7 +69118,7 @@
|
|||
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
|
||||
},
|
||||
"azure/eu/gpt-4.1-nano": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.8e-08,
|
||||
"input_cost_per_token": 1.1e-07,
|
||||
"input_cost_per_token_batches": 5.5e-08,
|
||||
|
|
@ -69288,7 +69306,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 5.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 8.25e-05,
|
||||
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
|
||||
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'",
|
||||
"supports_reasoning": true
|
||||
},
|
||||
"azure/eu/gpt-6-luna": {
|
||||
"deprecation_date": "2028-03-11",
|
||||
|
|
@ -69535,7 +69554,7 @@
|
|||
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
|
||||
},
|
||||
"azure/us/gpt-4.1-nano": {
|
||||
"deprecation_date": "2026-10-14",
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 2.8e-08,
|
||||
"input_cost_per_token": 1.1e-07,
|
||||
"input_cost_per_token_batches": 5.5e-08,
|
||||
|
|
|
|||
|
|
@ -28,10 +28,12 @@ GET /tag/user-agent/per-user-analytics
|
|||
GET /tag/wau
|
||||
GET /team/daily/activity
|
||||
GET /team/daily/activity/aggregated
|
||||
GET /team/daily/activity/aggregated/search
|
||||
GET /team/spend/by_user
|
||||
GET /team/spend/report
|
||||
GET /user/daily/activity
|
||||
GET /user/daily/activity/aggregated
|
||||
GET /user/daily/activity/aggregated/search
|
||||
GET /user/spend/report
|
||||
|
||||
# Admin UI helper endpoints; serve UI forms and caller-scoped views, not desired state
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ from e2e_http import (
|
|||
StreamingResponse,
|
||||
Success,
|
||||
UnknownApiError,
|
||||
proxy_error,
|
||||
require_successful_call,
|
||||
unwrap,
|
||||
)
|
||||
|
|
@ -1752,3 +1753,114 @@ class TestBatchTerminalState:
|
|||
assert (cost_row.total_tokens or 0) > 0, (
|
||||
f"batch cost row has no token usage: {cost_row.total_tokens!r}"
|
||||
)
|
||||
|
||||
|
||||
NATIVE_VERTEX_BATCH_ROWS: Final = b"".join(
|
||||
json.dumps(
|
||||
{
|
||||
"request": {
|
||||
"contents": [{"role": "user", "parts": [{"text": text}]}],
|
||||
"tools": [{"googleSearch": {"excludeDomains": ["example.com"]}}],
|
||||
}
|
||||
}
|
||||
).encode()
|
||||
+ b"\n"
|
||||
for text in ("What is the tallest building in the world?", "Who won the last FIFA World Cup?")
|
||||
)
|
||||
VERTEX_BATCH_PROVIDER: Final = next(p for p in PROVIDERS if p.name == "vertex_ai")
|
||||
|
||||
|
||||
class TestVertexNativePassthrough:
|
||||
"""`passthrough=true` on POST /v1/files uploads native Vertex batch JSONL byte for
|
||||
byte (no OpenAI-to-Vertex translation, so `googleSearch` tools and the grounding
|
||||
metadata they produce survive), and a batch created from that file is accepted.
|
||||
|
||||
Terminal-state assertions (native output rows with groundingMetadata, the spend
|
||||
row) are deliberately not here: retrieving a non-terminal batch books a $0 spend
|
||||
row that blocks the real-cost row, the same reason TestBatchTerminalState polls
|
||||
the list endpoint only. Those are proven by the PR's live curl proof instead.
|
||||
"""
|
||||
|
||||
@pytest.mark.covers(
|
||||
"llm.files.vertex.native_passthrough.nonstream.works",
|
||||
"llm.batches.vertex.native_passthrough.nonstream.works",
|
||||
exercised_on=["files", "batches"],
|
||||
)
|
||||
def test_native_jsonl_round_trips_untouched_and_starts_a_batch(
|
||||
self, client: BatchClient, resources: ResourceManager, batch_deployments: None
|
||||
) -> None:
|
||||
key = resources.key()
|
||||
file = unwrap(
|
||||
client.upload_file(
|
||||
content=NATIVE_VERTEX_BATCH_ROWS,
|
||||
form=FileUploadForm(
|
||||
purpose="batch", target_model_names=VERTEX_BATCH_PROVIDER.model, passthrough=True
|
||||
),
|
||||
key=key,
|
||||
)
|
||||
)
|
||||
resources.defer(lambda: cleanup_file(client, file.id, key=key))
|
||||
assert_file_object(file, provider="vertex_ai")
|
||||
assert is_managed_id(file.id), f"passthrough upload must return a managed file id, got {file.id!r}"
|
||||
assert file.bytes == len(NATIVE_VERTEX_BATCH_ROWS), (
|
||||
f"passthrough upload must report the caller's byte count, got {file.bytes}"
|
||||
)
|
||||
|
||||
downloaded = client.proxy.transport.download(
|
||||
f"/v1/files/{file.id}/content", headers=client.proxy.transport.bearer(key)
|
||||
)
|
||||
assert downloaded.status_code == 200, (
|
||||
f"file content must be 200, got {downloaded.status_code}: {downloaded.body[:300]}"
|
||||
)
|
||||
assert downloaded.body.encode() == NATIVE_VERTEX_BATCH_ROWS, (
|
||||
"passthrough file content must be the uploaded native rows byte for byte"
|
||||
)
|
||||
|
||||
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
|
||||
require_successful_call(created)
|
||||
batch = BatchObject.model_validate_json(created.body)
|
||||
resources.defer(lambda: cleanup_batch(client, batch.id, key=key, delete_output_files=True))
|
||||
assert is_managed_id(batch.id), f"passthrough batch must be LiteLLM-managed, got {batch.id!r}"
|
||||
assert batch.status in CREATED_BATCH_STATUSES, f"passthrough batch has non-transitional status {batch.status!r}"
|
||||
assert batch.input_file_id == file.id
|
||||
|
||||
@pytest.mark.covers("llm.files.vertex.native_passthrough_validation.nonstream.works", exercised_on=["files"])
|
||||
@pytest.mark.parametrize(
|
||||
"content, form, expected_param",
|
||||
[
|
||||
pytest.param(
|
||||
NATIVE_VERTEX_BATCH_ROWS,
|
||||
FileUploadForm(purpose="batch", passthrough=True),
|
||||
"target_model_names",
|
||||
id="no-target-model",
|
||||
),
|
||||
pytest.param(
|
||||
NATIVE_VERTEX_BATCH_ROWS,
|
||||
FileUploadForm(purpose="batch", target_model_names=OPENAI_BATCH_MODEL, passthrough=True),
|
||||
"target_model_names",
|
||||
id="non-vertex-target-model",
|
||||
),
|
||||
pytest.param(
|
||||
render_jsonl(VERTEX_BATCH_PROVIDER.raw_model),
|
||||
FileUploadForm(purpose="batch", target_model_names=VERTEX_BATCH_PROVIDER.model, passthrough=True),
|
||||
"request",
|
||||
id="openai-shaped-rows",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_passthrough_upload_is_rejected_outside_a_native_vertex_batch(
|
||||
self,
|
||||
content: bytes,
|
||||
form: FileUploadForm,
|
||||
expected_param: str,
|
||||
client: BatchClient,
|
||||
resources: ResourceManager,
|
||||
batch_deployments: None,
|
||||
) -> None:
|
||||
key = resources.key()
|
||||
result = client.upload_file(content=content, form=form, key=key)
|
||||
assert isinstance(result, UnknownApiError), f"expected a 400, got {result!r}"
|
||||
assert result.status_code == 400, f"expected 400, got {result.status_code}: {result.body[:300]}"
|
||||
error = proxy_error(result.body)
|
||||
assert error.param == expected_param, f"unexpected error param in {error!r}"
|
||||
assert "passthrough" in error.message
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@
|
|||
- {id: llm.batches.openai_provider_fallback.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py", rationale: "Provider-fallback raw-id scenario"}
|
||||
- {id: llm.batches.azure_openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Azure batches all scenarios"}
|
||||
- {id: llm.batches.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Vertex batches"}
|
||||
- {id: llm.batches.vertex.native_passthrough.nonstream.works, module: llm, tier: P1, subject_endpoint: batches, route: vertex, capability: native_passthrough, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py / LIT-4790", rationale: "A batch created from a passthrough-uploaded native Vertex JSONL file is accepted and starts on the deployment named at upload"}
|
||||
- {id: llm.batches.bedrock.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Bedrock batches (encoded/unified only)"}
|
||||
- {id: llm.batches.bedrock.assume_role.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: assume_role, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create under STS assume-role credentials"}
|
||||
- {id: llm.batches.bedrock.govcloud_partition.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: govcloud_partition, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create in the us-gov-west-1 partition"}
|
||||
|
|
@ -46,6 +47,8 @@
|
|||
- {id: llm.files.openai.passthrough.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_passthrough_e2e.py", rationale: "POST/DELETE /openai_passthrough/v1/files relay OpenAI's own file object; the dedicated prefix must not bind as a provider name on the /{provider}/v1/files route (GitHub issue #36086)"}
|
||||
- {id: llm.files.azure_openai.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:45", rationale: "Azure file upload managed backend"}
|
||||
- {id: llm.files.vertex.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:52", rationale: "Vertex file upload to GCS"}
|
||||
- {id: llm.files.vertex.native_passthrough.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: vertex, capability: native_passthrough, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py / LIT-4790", rationale: "POST /v1/files with passthrough=true ships native Vertex batch JSONL (googleSearch tools and all) to GCS untouched and GET /v1/files/{id}/content returns the same bytes"}
|
||||
- {id: llm.files.vertex.native_passthrough_validation.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: vertex, capability: input_validation, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py / LIT-4790", rationale: "passthrough=true without a Vertex target_model_names, or with OpenAI-shaped rows, is a 400 naming the offending field and nothing is uploaded"}
|
||||
- {id: llm.files.bedrock.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:59", rationale: "Bedrock file upload to S3"}
|
||||
- {id: llm.files.bedrock.govcloud_partition.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: govcloud_partition, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock file upload to an S3 bucket in the us-gov-west-1 partition"}
|
||||
- {id: llm.files.bedrock.split_s3_credentials.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: split_s3_credentials, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py / LIT-8297", rationale: "Bedrock file upload, content and delete sign S3 with s3_access_key_id / s3_secret_access_key when they differ from the aws_* identity"}
|
||||
|
|
|
|||
|
|
@ -74,6 +74,7 @@ LlmCapability = Literal[
|
|||
"mid_conversation_system",
|
||||
"multi_turn",
|
||||
"native_extensions",
|
||||
"native_passthrough",
|
||||
"pdf_input",
|
||||
"prompt_cache_1h",
|
||||
"prompt_cache_5m",
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ class FileUploadForm(BaseModel):
|
|||
purpose: str = "batch"
|
||||
target_model_names: str | None = None
|
||||
custom_llm_provider: str | None = None
|
||||
passthrough: bool | None = None
|
||||
|
||||
|
||||
# ---------- Result types ----------
|
||||
|
|
@ -376,12 +377,18 @@ class ProxyErrorDetail(BaseModel):
|
|||
message: str
|
||||
type: str
|
||||
code: str
|
||||
param: str | None = None
|
||||
|
||||
|
||||
class _ProxyErrorBody(BaseModel):
|
||||
error: ProxyErrorDetail
|
||||
|
||||
|
||||
def proxy_error(body: str) -> ProxyErrorDetail:
|
||||
"""The proxy's own error envelope (`{"error": {message, type, param, code}}`) parsed off a rejected call."""
|
||||
return _ProxyErrorBody.model_validate_json(body).error
|
||||
|
||||
|
||||
def relayed_provider_rate_limit(outcome: RateLimitedError) -> ProxyErrorDetail | None:
|
||||
"""The provider's own 429 as the proxy relayed it, or None when the 429 is the proxy's own."""
|
||||
if PROVIDER_RATE_LIMIT_MARKER not in outcome.body:
|
||||
|
|
|
|||
|
|
@ -8,7 +8,8 @@ caller can hand AWS support the request id behind a completion. Regional
|
|||
inference-profile ids are the deployment shape most Bedrock customers run; a
|
||||
v1.90.0 regression timed them out, and the Converse route keeps them covered in
|
||||
test_chat_completions_regression_e2e.py, so the invoke route carries its own
|
||||
rows here.
|
||||
rows here. The file also covers Bedrock-native OpenAI model ids taking the
|
||||
default (Converse) route with max_tokens.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -26,6 +27,7 @@ pytestmark = pytest.mark.e2e
|
|||
|
||||
CONVERSE_REGIONAL_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
INVOKE_REGIONAL_BACKEND = "bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
OPENAI_FAMILY_BACKEND = "bedrock/global.openai.gpt-6-sol"
|
||||
PROVIDER_HEADER_PREFIX = "llm_provider-"
|
||||
BEDROCK_REQUEST_ID_HEADER = "llm_provider-x-amzn-requestid"
|
||||
|
||||
|
|
@ -199,3 +201,18 @@ class TestBedrockInvokeRegionalModelIds:
|
|||
)
|
||||
|
||||
_assert_streamed_completion(result)
|
||||
|
||||
|
||||
class TestBedrockOpenAIFamilyDefaultRoute:
|
||||
@pytest.mark.covers("llm.chat_completions.bedrock_converse.basic.nonstream.works", exercised_on=[])
|
||||
def test_openai_family_model_id_completes_with_max_tokens(
|
||||
self, client: PassthroughClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model = _register_bedrock_model(
|
||||
client, resources, "e2e-bedrock-openai-family", OPENAI_FAMILY_BACKEND
|
||||
)
|
||||
key = resources.key()
|
||||
|
||||
response = unwrap(client.proxy.chat(key, ChatBody(model=model, messages=_prompt(), max_tokens=64)))
|
||||
|
||||
_assert_completion(response)
|
||||
|
|
|
|||
|
|
@ -57,3 +57,13 @@ export async function expectUnrestrictedDashboard(page: Page): Promise<void> {
|
|||
expect(info.ok(), `Read own user with dashboard session: HTTP ${info.status()}`).toBe(true);
|
||||
expect((await info.json()).user_id).toBe(session.user_id);
|
||||
}
|
||||
|
||||
export async function logInThroughLoginPage(page: Page, email: string, password: string): Promise<void> {
|
||||
await page.goto(`${rootPath()}/ui/login`);
|
||||
await page.getByPlaceholder("Enter your username").fill(email);
|
||||
await page.getByPlaceholder("Enter your password").fill(password);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
await page.waitForURL((url) => url.pathname.startsWith(`${rootPath()}/ui`) && !url.pathname.includes("/login"), {
|
||||
timeout: 30_000,
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,10 +1,14 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { ADMIN_STORAGE_PATH } from "../../constants";
|
||||
import { Role, users } from "../../fixtures/users";
|
||||
import { logInThroughLoginPage } from "../../helpers/userOnboarding";
|
||||
|
||||
test.describe("Logout", () => {
|
||||
test.use({ storageState: ADMIN_STORAGE_PATH });
|
||||
test.use({ storageState: { cookies: [], origins: [] } });
|
||||
|
||||
test("Clicking Logout clears the session and forces re-login on a protected page", async ({ page }) => {
|
||||
const admin = users[Role.ProxyAdmin];
|
||||
await logInThroughLoginPage(page, admin.email, admin.password);
|
||||
|
||||
await page.goto("/ui");
|
||||
// Scope to the sidebar; the top-bar breadcrumb also shows "Virtual Keys".
|
||||
await expect(page.getByRole("complementary").getByText("Virtual Keys")).toBeVisible({ timeout: 10_000 });
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { ADMIN_STORAGE_PATH } from "../../constants";
|
||||
import { Role, users } from "../../fixtures/users";
|
||||
import { logInThroughLoginPage } from "../../helpers/userOnboarding";
|
||||
|
||||
/**
|
||||
* Runs as part of the standard e2e suite: both `run_e2e.sh` and the CircleCI
|
||||
|
|
@ -16,9 +17,12 @@ const LOGOUT_URL = process.env.PROXY_LOGOUT_URL ?? "";
|
|||
test.skip(!LOGOUT_URL, "Requires PROXY_LOGOUT_URL env var");
|
||||
|
||||
test.describe("PROXY_LOGOUT_URL redirect", () => {
|
||||
test.use({ storageState: ADMIN_STORAGE_PATH });
|
||||
test.use({ storageState: { cookies: [], origins: [] } });
|
||||
|
||||
test("Logout clears the session and redirects to PROXY_LOGOUT_URL", async ({ page }) => {
|
||||
const admin = users[Role.ProxyAdmin];
|
||||
await logInThroughLoginPage(page, admin.email, admin.password);
|
||||
|
||||
const target = new URL(LOGOUT_URL);
|
||||
|
||||
// Stub the external logout destination so the assertion doesn't depend on
|
||||
|
|
@ -46,8 +50,8 @@ test.describe("PROXY_LOGOUT_URL redirect", () => {
|
|||
await expect(page.getByRole("complementary").getByText("Virtual Keys")).toBeVisible({ timeout: 15_000 });
|
||||
await settingsLoaded;
|
||||
|
||||
// Pre-condition: we start authenticated. The admin storage state carries a
|
||||
// `token` cookie, so a real logout has something to tear down.
|
||||
// Pre-condition: we start authenticated. The fresh login set a `token`
|
||||
// cookie, so a real logout has something to tear down.
|
||||
const tokensBefore = (await page.context().cookies()).filter((c) => c.name === "token");
|
||||
expect(tokensBefore.length, "should start logged in with a token cookie").toBeGreaterThan(0);
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
[
|
||||
"tests/e2e/ui/tests/integrationCritical/projectDetachment.spec.ts::project creation and explicit detachment preserve saved scope and restore serving",
|
||||
"tests/e2e/ui/tests/integrationCritical/teamGlobalGuardrailKillSwitch.spec.ts::proxy admin can enable the global guardrail kill switch from the models page team drill-in",
|
||||
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::per-user MCP env var stays updatable and clearable from the card after it is set",
|
||||
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::cancelling the clear confirmation keeps the stored value and sends no delete",
|
||||
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::pressing Enter on Update opens the credentials modal instead of the server editor",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,77 @@
|
|||
import {
|
||||
test,
|
||||
expect,
|
||||
APIRequestContext,
|
||||
Page as PlaywrightPage,
|
||||
} from "@playwright/test";
|
||||
import { randomUUID } from "node:crypto";
|
||||
|
||||
const master = process.env.LITELLM_MASTER_KEY ?? "sk-integration-master";
|
||||
const headers = { Authorization: `Bearer ${master}` };
|
||||
|
||||
async function createTeam(request: APIRequestContext): Promise<string> {
|
||||
const created = await request.post("/team/new", {
|
||||
headers,
|
||||
data: {
|
||||
team_alias: `int_kill_switch_${randomUUID().replace(/-/g, "").slice(0, 12)}`,
|
||||
},
|
||||
});
|
||||
expect(created.ok(), await created.text()).toBe(true);
|
||||
return (await created.json()).team_id as string;
|
||||
}
|
||||
|
||||
async function loginAsAdmin(page: PlaywrightPage): Promise<void> {
|
||||
await page.goto("/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill("admin");
|
||||
await page.getByPlaceholder("Enter your password").fill(master);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
await expect(page).toHaveURL(
|
||||
(url) => url.pathname.startsWith("/ui") && !url.pathname.includes("login"),
|
||||
);
|
||||
}
|
||||
|
||||
test("proxy admin can enable the global guardrail kill switch from the models page team drill-in", async ({
|
||||
page,
|
||||
request,
|
||||
}) => {
|
||||
const teamId = await createTeam(request);
|
||||
try {
|
||||
await loginAsAdmin(page);
|
||||
await page.goto(`/ui/models-and-endpoints?team=${teamId}`);
|
||||
await page.getByRole("tab", { name: "Settings" }).click();
|
||||
await page.getByRole("button", { name: /edit settings/i }).click();
|
||||
await expect(page.getByLabel(/Team Name/)).toBeVisible();
|
||||
const killSwitch = page.getByRole("switch", {
|
||||
name: /disable all global guardrails/i,
|
||||
});
|
||||
await expect(killSwitch).toBeVisible();
|
||||
await expect(killSwitch).not.toBeChecked();
|
||||
await killSwitch.click();
|
||||
await page.getByRole("button", { name: "Save Changes" }).click();
|
||||
await expect
|
||||
.poll(async () => {
|
||||
const response = await request.get(`/team/info?team_id=${teamId}`, {
|
||||
headers,
|
||||
});
|
||||
expect(response.ok(), await response.text()).toBe(true);
|
||||
const json = await response.json();
|
||||
return json.team_info?.metadata?.disable_global_guardrails;
|
||||
})
|
||||
.toBe(true);
|
||||
await page.reload();
|
||||
await page.getByRole("tab", { name: "Settings" }).click();
|
||||
await page.getByRole("button", { name: /edit settings/i }).click();
|
||||
await expect(page.getByLabel(/Team Name/)).toBeVisible();
|
||||
await expect(
|
||||
page.getByRole("switch", { name: /disable all global guardrails/i }),
|
||||
).toBeChecked();
|
||||
} finally {
|
||||
const removed = await request.post("/team/delete", {
|
||||
headers,
|
||||
data: { team_ids: [teamId] },
|
||||
});
|
||||
expect(removed.ok() || removed.status() === 404, await removed.text()).toBe(
|
||||
true,
|
||||
);
|
||||
}
|
||||
});
|
||||
|
|
@ -1031,6 +1031,7 @@ def test_set_llm_deployment_success_metrics(prometheus_logger):
|
|||
api_base="https://api.openai.com",
|
||||
api_provider="openai",
|
||||
requested_model="my_custom_model_group",
|
||||
model_group="my_custom_model_group",
|
||||
hashed_api_key=standard_logging_payload["metadata"]["user_api_key_hash"],
|
||||
api_key_alias=standard_logging_payload["metadata"]["user_api_key_alias"],
|
||||
team=standard_logging_payload["metadata"]["user_api_key_team_id"],
|
||||
|
|
@ -1047,6 +1048,7 @@ def test_set_llm_deployment_success_metrics(prometheus_logger):
|
|||
api_base="https://api.openai.com",
|
||||
api_provider="openai",
|
||||
requested_model="my_custom_model_group",
|
||||
model_group="my_custom_model_group",
|
||||
hashed_api_key=standard_logging_payload["metadata"]["user_api_key_hash"],
|
||||
api_key_alias=standard_logging_payload["metadata"]["user_api_key_alias"],
|
||||
team=standard_logging_payload["metadata"]["user_api_key_team_id"],
|
||||
|
|
|
|||
|
|
@ -250,6 +250,7 @@ async def test_azure_image_edit_litellm_sdk():
|
|||
self._json_data = json_data
|
||||
self.status_code = status_code
|
||||
self.text = json.dumps(json_data)
|
||||
self.headers = {}
|
||||
|
||||
def json(self):
|
||||
return self._json_data
|
||||
|
|
@ -370,6 +371,7 @@ async def test_openai_image_edit_cost_tracking():
|
|||
self._json_data = json_data
|
||||
self.status_code = status_code
|
||||
self.text = json.dumps(json_data)
|
||||
self.headers = {}
|
||||
|
||||
def json(self):
|
||||
return self._json_data
|
||||
|
|
@ -460,6 +462,7 @@ async def test_azure_image_edit_cost_tracking():
|
|||
self._json_data = json_data
|
||||
self.status_code = status_code
|
||||
self.text = json.dumps(json_data)
|
||||
self.headers = {}
|
||||
|
||||
def json(self):
|
||||
return self._json_data
|
||||
|
|
@ -737,6 +740,7 @@ async def test_image_edit_array_handling():
|
|||
self._json_data = json_data
|
||||
self.status_code = status_code
|
||||
self.text = json.dumps(json_data)
|
||||
self.headers = {}
|
||||
|
||||
def json(self):
|
||||
return self._json_data
|
||||
|
|
|
|||
|
|
@ -24,9 +24,14 @@ async def test_xinference_image_generation():
|
|||
def model_dump(self):
|
||||
return mock_openai_response
|
||||
|
||||
# Create a mock client with the images.generate method
|
||||
class MockRawResponse:
|
||||
headers = {}
|
||||
|
||||
def parse(self):
|
||||
return MockResponse()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.images.generate = AsyncMock(return_value=MockResponse())
|
||||
mock_client.images.with_raw_response.generate = AsyncMock(return_value=MockRawResponse())
|
||||
|
||||
# Capture the actual arguments sent to OpenAI client
|
||||
captured_args = None
|
||||
|
|
@ -36,9 +41,9 @@ async def test_xinference_image_generation():
|
|||
nonlocal captured_args, captured_kwargs
|
||||
captured_args = args
|
||||
captured_kwargs = kwargs
|
||||
return MockResponse()
|
||||
return MockRawResponse()
|
||||
|
||||
mock_client.images.generate.side_effect = capture_generate_call
|
||||
mock_client.images.with_raw_response.generate.side_effect = capture_generate_call
|
||||
|
||||
# Mock the _get_openai_client method to return our mock client
|
||||
with patch.object(
|
||||
|
|
@ -65,7 +70,7 @@ async def test_xinference_image_generation():
|
|||
assert response.data[0].url == "https://example.com/image.png"
|
||||
|
||||
# Validate that the OpenAI client was called with correct parameters
|
||||
mock_client.images.generate.assert_called_once()
|
||||
mock_client.images.with_raw_response.generate.assert_called_once()
|
||||
assert captured_kwargs is not None
|
||||
assert (
|
||||
captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large"
|
||||
|
|
@ -97,9 +102,14 @@ async def test_xinference_image_generation_with_response_format():
|
|||
def model_dump(self):
|
||||
return mock_openai_response
|
||||
|
||||
# Create a mock client with the images.generate method
|
||||
class MockRawResponse:
|
||||
headers = {}
|
||||
|
||||
def parse(self):
|
||||
return MockResponse()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.images.generate = AsyncMock(return_value=MockResponse())
|
||||
mock_client.images.with_raw_response.generate = AsyncMock(return_value=MockRawResponse())
|
||||
|
||||
# Capture the actual arguments sent to OpenAI client
|
||||
captured_args = None
|
||||
|
|
@ -109,9 +119,9 @@ async def test_xinference_image_generation_with_response_format():
|
|||
nonlocal captured_args, captured_kwargs
|
||||
captured_args = args
|
||||
captured_kwargs = kwargs
|
||||
return MockResponse()
|
||||
return MockRawResponse()
|
||||
|
||||
mock_client.images.generate.side_effect = capture_generate_call
|
||||
mock_client.images.with_raw_response.generate.side_effect = capture_generate_call
|
||||
|
||||
# Mock the _get_openai_client method to return our mock client
|
||||
with patch.object(
|
||||
|
|
@ -141,7 +151,7 @@ async def test_xinference_image_generation_with_response_format():
|
|||
assert response.data[0].b64_json is not None
|
||||
|
||||
# Validate that the OpenAI client was called with correct parameters
|
||||
mock_client.images.generate.assert_called_once()
|
||||
mock_client.images.with_raw_response.generate.assert_called_once()
|
||||
assert captured_kwargs is not None
|
||||
assert (
|
||||
captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large"
|
||||
|
|
|
|||
|
|
@ -195,7 +195,7 @@ def test_malformed_bodies_missing_users_and_foreign_servers_are_rejected(gateway
|
|||
path: Final = f"/v1/mcp/server/{identity}/user-env-vars"
|
||||
payload: Final[dict[str, JsonValue]] = {"values": {TOKEN: "x"}}
|
||||
malformed: Final[tuple[dict[str, JsonValue], ...]] = ({"values": {TOKEN: 7}}, {"values": ["a"]}, {})
|
||||
assert [gateway.request("POST", path, body, key=key).status_code for body in malformed] == [422, 422, 422]
|
||||
assert [gateway.request("POST", path, bad, key=key).status_code for bad in malformed] == [422, 422, 422]
|
||||
assert set_names(env_status(gateway, key, identity)) == {TOKEN: False}
|
||||
assert [gateway.client.request(method, path, json=payload).status_code for method in METHODS] == [401, 401, 401]
|
||||
no_user: Final = tuple(gateway.request(method, path, payload, key=userless) for method in METHODS)
|
||||
|
|
|
|||
|
|
@ -1,12 +1,9 @@
|
|||
import json
|
||||
import signal
|
||||
import socket
|
||||
import subprocess
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable, Generator
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -14,6 +11,7 @@ import httpx
|
|||
import psutil
|
||||
from integration._support.client import Gateway, eventually, object_value, string_value
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.redis_process import owned_redis
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from prometheus_client.parser import text_string_to_metric_families
|
||||
from test_cache_hit_guardrail_metrics import (
|
||||
|
|
@ -31,22 +29,6 @@ from test_cache_hit_guardrail_metrics import (
|
|||
BURST: Final = 10
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _redis(port: int) -> Generator[subprocess.Popen[bytes], None, None]:
|
||||
process: Final = subprocess.Popen(["redis-server", "--port", str(port), "--save", ""], stdout=subprocess.DEVNULL)
|
||||
try:
|
||||
yield process
|
||||
finally:
|
||||
process.kill()
|
||||
process.wait(timeout=10)
|
||||
|
||||
|
||||
def _free_port() -> int:
|
||||
with socket.socket() as reserve:
|
||||
reserve.bind(("127.0.0.1", 0))
|
||||
return reserve.getsockname()[1]
|
||||
|
||||
|
||||
def _deployment_id(candidate: Gateway, model_name: str) -> str:
|
||||
entries: Final = candidate.get("/model/info")["data"]
|
||||
assert isinstance(entries, list)
|
||||
|
|
@ -224,18 +206,16 @@ def test_stalled_guardrail_sink_recovers_and_counts(gateway: Gateway, tmp_path:
|
|||
|
||||
|
||||
def test_redis_outage_keeps_serving_in_memory_hits(gateway: Gateway, tmp_path: Path) -> None:
|
||||
"""X2: the redis cache keeps an in-memory shadow, so a redis kill does not stop cache-hit rejects."""
|
||||
"""X2: the redis cache keeps an in-memory shadow, so a redis outage does not stop cache-hit rejects."""
|
||||
marker: Final = uuid.uuid4().hex
|
||||
port: Final = _free_port()
|
||||
with _redis(port) as redis_one:
|
||||
with _rig(gateway, tmp_path, marker, env={"REDIS_HOST": "127.0.0.1", "REDIS_PORT": str(port)}) as rig:
|
||||
with owned_redis(tmp_path) as cache:
|
||||
with _rig(gateway, tmp_path, marker, env={"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)}) as rig:
|
||||
bodies: Final = _burst_bodies(rig, marker, None)[:BURST]
|
||||
_warm(rig, bodies)
|
||||
reject: Final = rig.candidate.request("POST", *bodies[0])
|
||||
assert reject.status_code == 400, reject.text
|
||||
warmed_hits: Final = rig.provider.received.qsize()
|
||||
redis_one.kill()
|
||||
redis_one.wait(timeout=10)
|
||||
cache.stop()
|
||||
outcomes: Final = _fire(rig, bodies[1:])
|
||||
assert all(status == 400 for status, _ in outcomes), outcomes
|
||||
assert rig.provider.received.qsize() == warmed_hits, (
|
||||
|
|
@ -243,13 +223,13 @@ def test_redis_outage_keeps_serving_in_memory_hits(gateway: Gateway, tmp_path: P
|
|||
warmed_hits,
|
||||
rig.provider.received.qsize(),
|
||||
)
|
||||
with _redis(port):
|
||||
recovered: Final = rig.candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
_chat_body(rig.model_name, "x2 rehit " + marker, rig.guardrail_name),
|
||||
)
|
||||
assert recovered.status_code == 400, recovered.text
|
||||
cache.start()
|
||||
recovered: Final = rig.candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
_chat_body(rig.model_name, "x2 rehit " + marker, rig.guardrail_name),
|
||||
)
|
||||
assert recovered.status_code == 400, recovered.text
|
||||
_expect_exactly_once(rig, (rig.model_name,), (rig.deployment_id,), 1 + len(bodies))
|
||||
|
||||
|
||||
|
|
|
|||
128
tests/integration/spend/test_team_daily_activity_key_search.py
Normal file
128
tests/integration/spend/test_team_daily_activity_key_search.py
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from hashlib import sha256
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
from pydantic import JsonValue
|
||||
|
||||
_SEARCH_PATH: Final = "/team/daily/activity/aggregated/search"
|
||||
|
||||
|
||||
def _range_around_today() -> dict[str, str]:
|
||||
today: Final = datetime.now(timezone.utc)
|
||||
return {
|
||||
"start_date": (today - timedelta(days=1)).strftime("%Y-%m-%d"),
|
||||
"end_date": (today + timedelta(days=1)).strftime("%Y-%m-%d"),
|
||||
"timezone": "0",
|
||||
}
|
||||
|
||||
|
||||
def _team_key_breakdown(body: dict[str, JsonValue], team: str) -> dict[str, JsonValue]:
|
||||
results: Final = body["results"]
|
||||
assert isinstance(results, list) and len(results) == 1, body
|
||||
entities: Final = object_value(object_value(object_value(results[0])["breakdown"])["entities"])
|
||||
return object_value(object_value(entities[team])["api_key_breakdown"])
|
||||
|
||||
|
||||
def test_team_key_search_returns_only_the_matching_key_spend_by_alias_and_by_hash(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
team: Final = scenario.team(models=[model])
|
||||
needle_alias: Final = f"needle-{uuid.uuid4().hex}"
|
||||
needle: Final = scenario.key(team_id=team, models=[model], key_alias=needle_alias)
|
||||
other: Final = scenario.key(team_id=team, models=[model], key_alias=f"other-{uuid.uuid4().hex}")
|
||||
needle_digest: Final = sha256(needle.encode()).hexdigest()
|
||||
other_digest: Final = sha256(other.encode()).hexdigest()
|
||||
for key in (needle, other):
|
||||
reply: Final = gateway.chat(model, key=key, text=f"key search {uuid.uuid4().hex}")
|
||||
assert object_value(reply["usage"])["total_tokens"] == 40, reply
|
||||
daily: Final = eventually(
|
||||
lambda: read_rows('SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
|
||||
lambda values: sorted(row["api_key"] for row in values) == sorted((needle_digest, other_digest)),
|
||||
seconds=70,
|
||||
)
|
||||
assert all(float(row["spend"]) == pytest.approx(0.06) for row in daily), daily
|
||||
for search in (needle_alias.upper(), needle_digest):
|
||||
response: Final = gateway.request(
|
||||
"GET", _SEARCH_PATH, params={"team_ids": team, "search": search, **_range_around_today()}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = object_value(response.json())
|
||||
assert object_value(body["metadata"])["total_spend"] == pytest.approx(0.06), response.text
|
||||
per_key: Final = _team_key_breakdown(body, team)
|
||||
assert set(per_key) == {needle_digest}, response.text
|
||||
assert object_value(object_value(per_key[needle_digest])["metrics"])["spend"] == pytest.approx(0.06)
|
||||
|
||||
|
||||
def test_team_key_search_is_scoped_to_the_teams_the_caller_belongs_to(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
team: Final = scenario.team(models=[model])
|
||||
needle_alias: Final = f"needle-{uuid.uuid4().hex}"
|
||||
needle: Final = scenario.key(team_id=team, models=[model], key_alias=needle_alias)
|
||||
needle_digest: Final = sha256(needle.encode()).hexdigest()
|
||||
reply: Final = gateway.chat(model, key=needle, text=f"key search {uuid.uuid4().hex}")
|
||||
assert object_value(reply["usage"])["total_tokens"] == 40, reply
|
||||
eventually(
|
||||
lambda: read_rows('SELECT api_key FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
|
||||
lambda values: [row["api_key"] for row in values] == [needle_digest],
|
||||
seconds=70,
|
||||
)
|
||||
outsider: Final = scenario.user(user_role="internal_user")
|
||||
outsider_team: Final = scenario.team(models=[model], members_with_roles=[{"user_id": outsider, "role": "user"}])
|
||||
outsider_key: Final = scenario.key(user_id=outsider, team_id=outsider_team, models=[model])
|
||||
params: Final = {"search": needle_alias, **_range_around_today()}
|
||||
admin_view: Final = gateway.request("GET", _SEARCH_PATH, params={"team_ids": team, **params})
|
||||
assert admin_view.status_code == 200, admin_view.text
|
||||
assert set(_team_key_breakdown(object_value(admin_view.json()), team)) == {needle_digest}, admin_view.text
|
||||
own_teams_view: Final = gateway.request("GET", _SEARCH_PATH, params=params, key=outsider_key)
|
||||
assert own_teams_view.status_code == 200, own_teams_view.text
|
||||
own_teams_body: Final = object_value(own_teams_view.json())
|
||||
assert own_teams_body["results"] == [], own_teams_view.text
|
||||
assert object_value(own_teams_body["metadata"])["total_api_keys"] == 0, own_teams_view.text
|
||||
foreign_team_view: Final = gateway.request(
|
||||
"GET", _SEARCH_PATH, params={"team_ids": team, **params}, key=outsider_key
|
||||
)
|
||||
assert foreign_team_view.status_code == 404, foreign_team_view.text
|
||||
|
||||
|
||||
def test_team_key_search_excludes_teams_inside_the_where(gateway: Gateway) -> None:
|
||||
"""The dashboard always sends exclude_team_ids; a matching key in an excluded
|
||||
team with higher spend must not consume a take slot nor appear in the result."""
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
team_keep: Final = scenario.team(models=[model])
|
||||
team_drop: Final = scenario.team(models=[model])
|
||||
shared_alias: Final = f"needle-{uuid.uuid4().hex}"
|
||||
keep: Final = scenario.key(team_id=team_keep, models=[model], key_alias=f"{shared_alias}-keep")
|
||||
drop: Final = scenario.key(team_id=team_drop, models=[model], key_alias=f"{shared_alias}-drop")
|
||||
keep_digest: Final = sha256(keep.encode()).hexdigest()
|
||||
drop_digest: Final = sha256(drop.encode()).hexdigest()
|
||||
for _ in range(2):
|
||||
reply: Final = gateway.chat(model, key=drop, text=f"key search {uuid.uuid4().hex}")
|
||||
assert object_value(reply["usage"])["total_tokens"] == 40, reply
|
||||
reply = gateway.chat(model, key=keep, text=f"key search {uuid.uuid4().hex}")
|
||||
assert object_value(reply["usage"])["total_tokens"] == 40, reply
|
||||
eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id IN (%s, %s)',
|
||||
(team_keep, team_drop),
|
||||
),
|
||||
lambda values: sorted(row["api_key"] for row in values) == sorted((keep_digest, drop_digest)),
|
||||
seconds=70,
|
||||
)
|
||||
response: Final = gateway.request(
|
||||
"GET",
|
||||
_SEARCH_PATH,
|
||||
params={"search": shared_alias, "exclude_team_ids": team_drop, **_range_around_today()},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = object_value(response.json())
|
||||
results: Final = body["results"]
|
||||
assert isinstance(results, list) and len(results) == 1, body
|
||||
entities: Final = object_value(object_value(object_value(results[0])["breakdown"])["entities"])
|
||||
assert set(entities) == {team_keep}, response.text
|
||||
assert set(_team_key_breakdown(body, team_keep)) == {keep_digest}, response.text
|
||||
|
|
@ -210,11 +210,14 @@ async def test_litellm_gateway_image_generation_direct(is_async):
|
|||
"created": 1,
|
||||
"data": [{"url": "https://example.com/image.png"}],
|
||||
}
|
||||
mock_raw_response = MagicMock()
|
||||
mock_raw_response.parse.return_value = mock_openai_response
|
||||
mock_raw_response.headers = {}
|
||||
|
||||
if is_async:
|
||||
# Mock the AsyncOpenAI client that gets created inside _get_openai_client
|
||||
mock_async_client = AsyncMock()
|
||||
mock_async_client.images.generate = AsyncMock(return_value=mock_openai_response)
|
||||
mock_async_client.images.with_raw_response.generate = AsyncMock(return_value=mock_raw_response)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.openai.openai.AsyncOpenAI", return_value=mock_async_client
|
||||
|
|
@ -234,14 +237,14 @@ async def test_litellm_gateway_image_generation_direct(is_async):
|
|||
assert constructor_kwargs["base_url"] == "http://my-proxy"
|
||||
|
||||
# Verify the AsyncOpenAI client was called correctly
|
||||
mock_async_client.images.generate.assert_awaited_once()
|
||||
call_kwargs = mock_async_client.images.generate.call_args.kwargs
|
||||
mock_async_client.images.with_raw_response.generate.assert_awaited_once()
|
||||
call_kwargs = mock_async_client.images.with_raw_response.generate.call_args.kwargs
|
||||
assert call_kwargs["model"] == "dall-e-3"
|
||||
assert call_kwargs["prompt"] == "A beautiful sunset over mountains"
|
||||
else:
|
||||
# Mock the sync OpenAI client that gets created inside _get_openai_client
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.images.generate.return_value = mock_openai_response
|
||||
mock_sync_client.images.with_raw_response.generate.return_value = mock_raw_response
|
||||
|
||||
with patch(
|
||||
"litellm.llms.openai.openai.OpenAI", return_value=mock_sync_client
|
||||
|
|
@ -260,8 +263,8 @@ async def test_litellm_gateway_image_generation_direct(is_async):
|
|||
assert constructor_kwargs["base_url"] == "http://my-proxy"
|
||||
|
||||
# Verify the OpenAI client was called correctly
|
||||
mock_sync_client.images.generate.assert_called_once()
|
||||
call_kwargs = mock_sync_client.images.generate.call_args.kwargs
|
||||
mock_sync_client.images.with_raw_response.generate.assert_called_once()
|
||||
call_kwargs = mock_sync_client.images.with_raw_response.generate.call_args.kwargs
|
||||
assert call_kwargs["model"] == "dall-e-3"
|
||||
assert call_kwargs["prompt"] == "A beautiful sunset over mountains"
|
||||
|
||||
|
|
@ -285,6 +288,7 @@ async def test_litellm_gateway_from_sdk_image_edit(is_async):
|
|||
self._json_data = json_data
|
||||
self.status_code = status_code
|
||||
self.text = json.dumps(json_data)
|
||||
self.headers = {}
|
||||
|
||||
def json(self):
|
||||
return self._json_data
|
||||
|
|
|
|||
|
|
@ -313,7 +313,7 @@ def test_openai_max_retries_0(mock_get_openai_client):
|
|||
def test_openai_image_generation_forwards_organization(mock_get_openai_client):
|
||||
"""Ensure organization flows to OpenAI client for image generation."""
|
||||
|
||||
class _DummyImages:
|
||||
class _DummyRawImages:
|
||||
def generate(self, **kwargs): # type: ignore
|
||||
class _Resp:
|
||||
def model_dump(self_inner): # minimal OpenAI ImagesResponse shape
|
||||
|
|
@ -327,7 +327,16 @@ def test_openai_image_generation_forwards_organization(mock_get_openai_client):
|
|||
},
|
||||
}
|
||||
|
||||
return _Resp()
|
||||
class _RawResp:
|
||||
headers = {}
|
||||
|
||||
def parse(self_inner):
|
||||
return _Resp()
|
||||
|
||||
return _RawResp()
|
||||
|
||||
class _DummyImages:
|
||||
with_raw_response = _DummyRawImages()
|
||||
|
||||
class _DummyClient:
|
||||
def __init__(self):
|
||||
|
|
|
|||
|
|
@ -1361,16 +1361,14 @@ def test_ollama_image():
|
|||
|
||||
from PIL import Image
|
||||
|
||||
sent_images = []
|
||||
|
||||
def mock_post(url, **kwargs):
|
||||
sent_images.append(json.loads(kwargs["data"])["images"])
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"Content-Type": "application/json"}
|
||||
data_json = json.loads(kwargs["data"])
|
||||
mock_response.json.return_value = {
|
||||
# return the image in the response so that it can be tested
|
||||
# against the original
|
||||
"response": data_json["images"]
|
||||
}
|
||||
mock_response.json.return_value = {"response": "a black pixel"}
|
||||
return mock_response
|
||||
|
||||
def make_b64image(format):
|
||||
|
|
@ -1399,9 +1397,10 @@ def test_ollama_image():
|
|||
|
||||
client = HTTPHandler()
|
||||
for test in tests:
|
||||
sent_images.clear()
|
||||
try:
|
||||
with patch.object(client, "post", side_effect=mock_post):
|
||||
response = completion(
|
||||
completion(
|
||||
model="ollama/llava",
|
||||
messages=[
|
||||
{
|
||||
|
|
@ -1417,14 +1416,14 @@ def test_ollama_image():
|
|||
],
|
||||
client=client,
|
||||
)
|
||||
(image_data,) = sent_images[0]
|
||||
if not test[1]:
|
||||
# the conversion process may not always generate the same image,
|
||||
# so just check for a JPEG image when a conversion was done.
|
||||
image_data = response["choices"][0]["message"]["content"][0]
|
||||
image = Image.open(io.BytesIO(base64.b64decode(image_data)))
|
||||
assert image.format == "JPEG"
|
||||
else:
|
||||
assert response["choices"][0]["message"]["content"][0] == test[1]
|
||||
assert image_data == test[1]
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
||||
|
|
|
|||
|
|
@ -5,8 +5,9 @@ from .actors import Actor
|
|||
pytestmark = pytest.mark.asyncio(loop_scope="session")
|
||||
|
||||
|
||||
# GET /team/daily/activity and its /aggregated variant (same shared scope
|
||||
# resolver, so the matrix must hold for both). A proxy admin (admin view) sees
|
||||
# GET /team/daily/activity, its /aggregated variant, and the key-search
|
||||
# variant (same shared scope resolver, so the matrix must hold for all
|
||||
# three). A proxy admin (admin view) sees
|
||||
# activity for any team. A non-admin is scoped to user_info.teams: a bare query
|
||||
# defaults to its own teams (200), and an explicit team_ids filter naming a
|
||||
# team it does not belong to is 404 (the VERIA-43 fix). Org admins have no
|
||||
|
|
@ -43,8 +44,12 @@ _DATES = "start_date=2024-01-01&end_date=2024-12-31"
|
|||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint",
|
||||
("/team/daily/activity", "/team/daily/activity/aggregated"),
|
||||
ids=("paginated", "aggregated"),
|
||||
(
|
||||
"/team/daily/activity",
|
||||
"/team/daily/activity/aggregated",
|
||||
"/team/daily/activity/aggregated/search",
|
||||
),
|
||||
ids=("paginated", "aggregated", "search"),
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"actor,team,expected_status",
|
||||
|
|
@ -54,7 +59,7 @@ _DATES = "start_date=2024-01-01&end_date=2024-12-31"
|
|||
async def test_team_daily_activity_matrix(
|
||||
actor: Actor, team: str, expected_status: int, endpoint: str, proxy_client, world
|
||||
):
|
||||
query = _DATES
|
||||
query = _DATES + ("&search=x" if endpoint.endswith("/search") else "")
|
||||
if team == "alpha":
|
||||
query += f"&team_ids={world.team_alpha_id}"
|
||||
elif team == "beta":
|
||||
|
|
|
|||
|
|
@ -141,6 +141,7 @@ def test_should_replace_model_in_jsonl():
|
|||
from litellm.router_utils.batch_utils import should_replace_model_in_jsonl
|
||||
|
||||
assert should_replace_model_in_jsonl(purpose="batch") is True
|
||||
assert should_replace_model_in_jsonl(purpose="batch", passthrough=True) is False
|
||||
assert should_replace_model_in_jsonl(purpose="test") is False
|
||||
assert should_replace_model_in_jsonl(purpose="user_data") is False
|
||||
|
||||
|
|
|
|||
|
|
@ -134,7 +134,7 @@ async def test_create_mcp_server_direct():
|
|||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
|
||||
) as mock_get_prisma,
|
||||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server",
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free",
|
||||
new_callable=mock.AsyncMock,
|
||||
) as mock_create,
|
||||
mock.patch(
|
||||
|
|
@ -345,7 +345,7 @@ async def test_create_mcp_server_invalid_alias():
|
|||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server"
|
||||
) as mock_get_server,
|
||||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server"
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server_if_identifier_free"
|
||||
) as mock_create,
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
|
|
|
|||
387
tests/test_litellm/batches/test_batch_utils.py
Normal file
387
tests/test_litellm/batches/test_batch_utils.py
Normal file
|
|
@ -0,0 +1,387 @@
|
|||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
import litellm.batches.batch_utils as bu
|
||||
from litellm.types.llms.openai import Batch
|
||||
|
||||
GROUNDED_USAGE_METADATA = {
|
||||
"promptTokenCount": 19,
|
||||
"candidatesTokenCount": 59,
|
||||
"thoughtsTokenCount": 406,
|
||||
"toolUsePromptTokenCount": 73,
|
||||
"totalTokenCount": 557,
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 19}],
|
||||
"candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 59}],
|
||||
"toolUsePromptTokensDetails": [{"modality": "TEXT", "tokenCount": 73}],
|
||||
"trafficType": "ON_DEMAND",
|
||||
}
|
||||
PASSTHROUGH_OUTPUT_URI = (
|
||||
"gs://litellm-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/u/"
|
||||
"predictions.jsonl"
|
||||
)
|
||||
UNGROUNDED_USAGE_METADATA = {
|
||||
"promptTokenCount": 20,
|
||||
"candidatesTokenCount": 48,
|
||||
"thoughtsTokenCount": 195,
|
||||
"toolUsePromptTokenCount": 73,
|
||||
"totalTokenCount": 336,
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 20}],
|
||||
"trafficType": "ON_DEMAND",
|
||||
}
|
||||
|
||||
|
||||
def _batch(output_file_id: str) -> Batch:
|
||||
return Batch(
|
||||
id="b",
|
||||
completion_window="24h",
|
||||
created_at=1,
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="f",
|
||||
object="batch",
|
||||
status="completed",
|
||||
output_file_id=output_file_id,
|
||||
)
|
||||
|
||||
|
||||
def _vertex_jsonl(rows: list[dict]) -> bytes:
|
||||
return "\n".join(json.dumps(row) for row in rows).encode()
|
||||
|
||||
|
||||
def _vertex_openai_row(custom_id: str, model: str, prompt_tokens: int, completion_tokens: int) -> dict:
|
||||
return {
|
||||
"id": f"batch_req_{custom_id}",
|
||||
"custom_id": custom_id,
|
||||
"response": {
|
||||
"status_code": 200,
|
||||
"request_id": custom_id,
|
||||
"body": {
|
||||
"id": f"chatcmpl-{custom_id}",
|
||||
"object": "chat.completion",
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
|
||||
"usage": {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": prompt_tokens + completion_tokens,
|
||||
},
|
||||
},
|
||||
},
|
||||
"error": None,
|
||||
}
|
||||
|
||||
|
||||
def _native_vertex_row(usage_metadata: dict, *, grounded: bool, model_version: str | None = "gemini-2.5-flash"):
|
||||
candidate = {"content": {"role": "model", "parts": [{"text": "ok"}]}, "finishReason": "STOP"}
|
||||
grounding = {"groundingMetadata": {"webSearchQueries": ["q"]}} if grounded else {}
|
||||
response = {"candidates": [{**candidate, **grounding}], "usageMetadata": usage_metadata}
|
||||
return {
|
||||
"request": {"contents": [{"role": "user", "parts": [{"text": "q"}]}], "tools": [{"googleSearch": {}}]},
|
||||
"status": "",
|
||||
"response": {**response, **({"modelVersion": model_version} if model_version else {})},
|
||||
"processed_time": "2026-09-23T19:02:00.000+00:00",
|
||||
}
|
||||
|
||||
|
||||
def _capture_cost_calls(monkeypatch, prompt_cost=0.5, completion_cost=0.25) -> list:
|
||||
import litellm.cost_calculator as cc
|
||||
|
||||
calls: list = []
|
||||
|
||||
def _calc(**kw):
|
||||
calls.append(kw)
|
||||
return (prompt_cost, completion_cost)
|
||||
|
||||
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
|
||||
return calls
|
||||
|
||||
|
||||
def test_vertex_native_cost_bills_embedding_rows(monkeypatch):
|
||||
monkeypatch.setitem(litellm.model_cost, "vertex_ai/gemini-embedding-2", {"input_cost_per_token_batches": 1e-7})
|
||||
rows = [
|
||||
{
|
||||
"key": "id_1",
|
||||
"status": "",
|
||||
"request": {"content": {"parts": [{"text": "hello world"}]}},
|
||||
"response": {"embedding": {"values": [0.1, 0.2]}, "usageMetadata": {"promptTokenCount": 2}},
|
||||
},
|
||||
{
|
||||
"key": "id_2",
|
||||
"status": "",
|
||||
"request": {"content": {"parts": [{"text": "hello"}]}},
|
||||
"response": {"embedding": {"values": [0.3]}, "tokenCount": "3"},
|
||||
},
|
||||
{"key": "id_3", "status": "INVALID_ARGUMENT", "request": {"content": {"parts": [{"text": ""}]}}},
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-embedding-2")
|
||||
|
||||
assert (result.successful_requests, result.failed_requests) == (2, 1)
|
||||
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (5, 0, 5)
|
||||
assert result.cost == pytest.approx(5 * 1e-7)
|
||||
assert result.models == ["gemini-embedding-2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_vertex_rows_route_to_vertex_cost_path_without_flag(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
|
||||
monkeypatch.setattr(
|
||||
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
|
||||
)
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False),
|
||||
]
|
||||
|
||||
result = await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
|
||||
)
|
||||
|
||||
assert result.cost == pytest.approx(1.5)
|
||||
assert (result.successful_requests, result.failed_requests) == (2, 0)
|
||||
assert result.models == ["gemini-2.5-flash"]
|
||||
assert {(call["model"], call["custom_llm_provider"]) for call in calls} == {("gemini-2.5-flash", "vertex_ai")}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_shaped_vertex_rows_keep_the_generic_path_without_flag(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
|
||||
monkeypatch.setattr(
|
||||
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
|
||||
)
|
||||
_capture_cost_calls(monkeypatch)
|
||||
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
|
||||
|
||||
result = await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
|
||||
)
|
||||
|
||||
assert result.successful_requests == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_vertex_rows_on_another_provider_keep_the_generic_path(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
|
||||
)
|
||||
_capture_cost_calls(monkeypatch)
|
||||
|
||||
result = await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert result.successful_requests == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_completed_batch_routes_native_rows_without_flag(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
|
||||
raw_rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)]
|
||||
|
||||
async def fake_fetch(batch, custom_llm_provider, litellm_params=None):
|
||||
return _vertex_jsonl(raw_rows)
|
||||
|
||||
monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch)
|
||||
monkeypatch.setattr(
|
||||
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
|
||||
)
|
||||
calls = _capture_cost_calls(monkeypatch, prompt_cost=0.7, completion_cost=0.3)
|
||||
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
|
||||
|
||||
result = await bu._handle_completed_batch(
|
||||
_batch(PASSTHROUGH_OUTPUT_URI),
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_name="gemini-2.5-flash",
|
||||
model_info=deployment_model_info,
|
||||
)
|
||||
|
||||
assert result.cost == pytest.approx(1.0)
|
||||
assert result.usage.total_tokens == 557
|
||||
assert [call["model_info"] for call in calls] == [deployment_model_info]
|
||||
|
||||
|
||||
def test_native_vertex_usage_is_billed_like_the_online_path(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
grounded = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)
|
||||
ungrounded = _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False)
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage([grounded, ungrounded], "gemini-2.5-flash")
|
||||
|
||||
grounded_usage, ungrounded_usage = (call["usage"] for call in calls)
|
||||
assert grounded_usage.prompt_tokens == 19
|
||||
assert grounded_usage.completion_tokens == 59 + 406
|
||||
assert grounded_usage.completion_tokens_details.reasoning_tokens == 406
|
||||
assert ungrounded_usage.prompt_tokens == 20 + 73
|
||||
assert ungrounded_usage.completion_tokens == 48 + 195
|
||||
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (
|
||||
19 + 93,
|
||||
465 + 243,
|
||||
557 + 336,
|
||||
)
|
||||
|
||||
|
||||
def test_native_vertex_rows_are_priced_by_model_version_without_a_model_name(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-pro"),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
|
||||
|
||||
assert [call["model"] for call in calls] == ["gemini-2.5-flash", "gemini-2.5-pro"]
|
||||
assert result.models == ["gemini-2.5-flash", "gemini-2.5-pro"]
|
||||
assert result.cost == pytest.approx(1.5)
|
||||
assert result.successful_requests == 3
|
||||
assert result.usage.total_tokens == 557 + 336 + 336
|
||||
|
||||
|
||||
def test_native_vertex_rows_without_usage_metadata_count_as_failed(monkeypatch):
|
||||
_capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
{"request": {"contents": []}, "status": "Error: bad request", "processed_time": "t"},
|
||||
{"request": {"contents": []}, "response": {"candidates": []}},
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
|
||||
|
||||
assert (result.successful_requests, result.failed_requests) == (1, 2)
|
||||
assert result.usage.total_tokens == 557
|
||||
|
||||
|
||||
def test_native_vertex_batch_whose_rows_all_failed_still_names_the_deployment_model(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [{"request": {"contents": []}, "status": "Error: quota exceeded", "processed_time": "t"}] * 2
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
|
||||
|
||||
assert result.models == ["gemini-2.5-flash"]
|
||||
assert (result.successful_requests, result.failed_requests, result.cost) == (0, 2, 0.0)
|
||||
assert calls == []
|
||||
|
||||
|
||||
def test_native_vertex_rows_are_priced_with_the_deployment_model_info(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
|
||||
|
||||
bu.calculate_vertex_ai_batch_cost_and_usage(
|
||||
[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
|
||||
"gemini-2.5-flash",
|
||||
model_info=deployment_model_info,
|
||||
)
|
||||
|
||||
assert [call["model_info"] for call in calls] == [deployment_model_info]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_vertex_rows_keep_the_deployment_model_info_through_the_batch_entrypoint(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
deployment_model_info = {"input_cost_per_token_batches": 1e-6}
|
||||
|
||||
await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_name="gemini-2.5-flash",
|
||||
model_info=deployment_model_info,
|
||||
)
|
||||
|
||||
assert [call["model_info"] for call in calls] == [deployment_model_info]
|
||||
|
||||
|
||||
def test_native_vertex_rows_are_priced_by_the_deployment_model_over_model_version(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-pro")]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
|
||||
|
||||
assert [call["model"] for call in calls] == ["gemini-2.5-flash"]
|
||||
assert result.models == ["gemini-2.5-flash"]
|
||||
|
||||
|
||||
def test_native_vertex_rows_that_fail_response_validation_count_as_failed(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
{"request": {"contents": []}, "response": {"candidates": "nope", "usageMetadata": GROUNDED_USAGE_METADATA}},
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
|
||||
|
||||
assert (result.successful_requests, result.failed_requests) == (1, 1)
|
||||
assert result.usage.total_tokens == 557
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("wildcard_model", ["*", "vertex_ai/*"])
|
||||
def test_native_vertex_rows_under_a_wildcard_deployment_are_priced_by_model_version(monkeypatch, wildcard_model):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, wildcard_model)
|
||||
|
||||
assert [call["model"] for call in calls] == ["gemini-2.5-flash", wildcard_model]
|
||||
assert result.cost == pytest.approx(1.5)
|
||||
assert (result.successful_requests, result.failed_requests) == (2, 0)
|
||||
assert result.usage.total_tokens == 557 + 336
|
||||
|
||||
|
||||
def test_native_vertex_row_without_model_version_under_a_wildcard_deployment_bills_its_explicit_prices():
|
||||
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
|
||||
with_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash")
|
||||
without_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version=None)
|
||||
|
||||
twin = bu.calculate_vertex_ai_batch_cost_and_usage([with_version], "vertex_ai/*", model_info=deployment_model_info)
|
||||
both = bu.calculate_vertex_ai_batch_cost_and_usage(
|
||||
[with_version, without_version], "vertex_ai/*", model_info=deployment_model_info
|
||||
)
|
||||
|
||||
assert twin.cost > 0
|
||||
assert both.cost == pytest.approx(2 * twin.cost)
|
||||
assert (both.successful_requests, both.failed_requests) == (2, 0)
|
||||
|
||||
|
||||
def test_native_vertex_row_the_cost_map_cannot_price_is_billed_at_zero_and_the_rest_still_bills(monkeypatch):
|
||||
import litellm.cost_calculator as cc
|
||||
|
||||
def _calc(**kw):
|
||||
if kw["model"] == "gemini-unpriced":
|
||||
raise ValueError("no pricing")
|
||||
return (0.5, 0.25)
|
||||
|
||||
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
|
||||
rows = [
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-unpriced"),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-flash"),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
|
||||
|
||||
assert result.cost == pytest.approx(0.75)
|
||||
assert (result.successful_requests, result.failed_requests) == (2, 0)
|
||||
assert result.usage.total_tokens == 557 + 336
|
||||
assert result.models == ["gemini-unpriced", "gemini-2.5-flash"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_known(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False)
|
||||
monkeypatch.setattr(
|
||||
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
|
||||
)
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
|
||||
|
||||
result = await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
|
||||
)
|
||||
|
||||
assert calls == []
|
||||
assert (result.successful_requests, result.failed_requests) == (0, 1)
|
||||
0
tests/test_litellm/files/__init__.py
Normal file
0
tests/test_litellm/files/__init__.py
Normal file
71
tests/test_litellm/files/test_main.py
Normal file
71
tests/test_litellm/files/test_main.py
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
from typing import Final
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
NATIVE_VERTEX_ROWS: Final = (
|
||||
b'{"request": {"contents": [{"role": "user", "parts": [{"text": "Who won the 2024 Tour de France?"}]}],'
|
||||
b' "tools": [{"googleSearch": {"excludeDomains": ["example.com"]}}]}}\n'
|
||||
b'{"request": {"contents": [{"role": "user", "parts": [{"text": "What is the tallest building in Tokyo?"}]}],'
|
||||
b' "tools": [{"googleSearch": {}}]}}\n'
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider, purpose",
|
||||
[("openai", "batch"), ("vertex_ai", "assistants")],
|
||||
ids=["non-vertex-provider", "non-batch-purpose"],
|
||||
)
|
||||
def test_create_file_passthrough_is_rejected_outside_a_vertex_batch(custom_llm_provider, purpose):
|
||||
with pytest.raises(litellm.BadRequestError) as exc_info:
|
||||
litellm.create_file(
|
||||
file=("batch.jsonl", b'{"request": {"contents": []}}\n', "application/jsonl"),
|
||||
purpose=purpose,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
passthrough=True,
|
||||
api_key="sk-test",
|
||||
api_base="http://127.0.0.1:9",
|
||||
)
|
||||
|
||||
assert "vertex_ai" in str(exc_info.value)
|
||||
assert "batch" in str(exc_info.value)
|
||||
|
||||
|
||||
def _gcs_upload_transport(uploads: list[httpx.Request]) -> httpx.MockTransport:
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
uploads.append(request)
|
||||
object_name: Final = parse_qs(urlparse(str(request.url)).query)["name"][0]
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": f"my-bucket/{object_name}/1758585600000000",
|
||||
"name": object_name,
|
||||
"size": str(len(request.read())),
|
||||
"timeCreated": "2026-09-23T00:00:00.000Z",
|
||||
},
|
||||
)
|
||||
|
||||
return httpx.MockTransport(respond)
|
||||
|
||||
|
||||
def test_create_file_passthrough_kwarg_ships_native_rows_byte_for_byte_under_the_passthrough_prefix():
|
||||
uploads: Final[list[httpx.Request]] = []
|
||||
file_object = litellm.create_file(
|
||||
file=("batch.jsonl", NATIVE_VERTEX_ROWS, "application/jsonl"),
|
||||
purpose="batch",
|
||||
custom_llm_provider="vertex_ai",
|
||||
passthrough=True,
|
||||
model="vertex_ai/gemini-2.5-flash",
|
||||
gcs_bucket_name="my-bucket",
|
||||
api_key="test-token",
|
||||
client=HTTPHandler(client=httpx.Client(transport=_gcs_upload_transport(uploads))),
|
||||
)
|
||||
(upload,) = uploads
|
||||
object_name: Final = parse_qs(urlparse(str(upload.url)).query)["name"][0]
|
||||
assert upload.read() == NATIVE_VERTEX_ROWS
|
||||
assert object_name.startswith("litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/")
|
||||
assert file_object.id == f"gs://my-bucket/{object_name}"
|
||||
|
|
@ -1207,14 +1207,18 @@ def test_speech_response_without_a_byte_count_produces_no_output() -> None:
|
|||
def test_speech_binary_response_is_logged_as_its_summary_not_dropped() -> None:
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.litellm_logging import _extract_response_obj_and_hidden_params
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
raw: Final = httpx.Response(200, headers={"content-type": "audio/mpeg"}, content=b"\x00" * 1234)
|
||||
response_obj, hidden_params = _extract_response_obj_and_hidden_params(HttpxBinaryResponseContent(raw), None)
|
||||
speech: Final = HttpxBinaryResponseContent(raw)
|
||||
set_provider_response_headers_in_hidden_params(speech, raw.headers)
|
||||
response_obj, hidden_params = _extract_response_obj_and_hidden_params(speech, None)
|
||||
|
||||
assert response_obj == {"object": "binary", "content_type": "audio/mpeg", "num_bytes": 1234}
|
||||
assert hidden_params is None
|
||||
assert hidden_params is not None
|
||||
assert hidden_params["headers"]["content-type"] == "audio/mpeg"
|
||||
|
||||
|
||||
def test_speech_binary_response_still_streaming_reports_the_bytes_downloaded_so_far() -> None:
|
||||
|
|
|
|||
|
|
@ -787,6 +787,176 @@ async def test_failure_hook_prefers_request_data_provider_over_exception_provide
|
|||
) == ["azure"]
|
||||
|
||||
|
||||
def test_model_group_in_deployment_metrics():
|
||||
"""
|
||||
Test that model_group label is present on the deployment-scoped metrics
|
||||
needed to build model-group dashboards (request counts, success/failure
|
||||
counts, tpm/rpm limits). These metrics previously only carried
|
||||
requested_model, litellm_model_name and model_id, none of which identify
|
||||
the model_group a pooled deployment belongs to.
|
||||
"""
|
||||
model_group_label = UserAPIKeyLabelNames.MODEL_GROUP.value
|
||||
|
||||
metrics_with_model_group = [
|
||||
"litellm_deployment_total_requests",
|
||||
"litellm_deployment_success_responses",
|
||||
"litellm_deployment_failure_responses",
|
||||
"litellm_deployment_tpm_limit",
|
||||
"litellm_deployment_rpm_limit",
|
||||
]
|
||||
|
||||
for metric_name in metrics_with_model_group:
|
||||
labels = PrometheusMetricLabels.get_labels(metric_name)
|
||||
assert (
|
||||
model_group_label in labels
|
||||
), f"Metric {metric_name} should contain model_group label"
|
||||
print(f"✅ {metric_name} contains model_group label")
|
||||
|
||||
|
||||
def test_model_group_value_flows_through_deployment_metrics_label_factory():
|
||||
"""
|
||||
The label being in the allow-list is necessary but not sufficient: the
|
||||
factory must also carry the value from the enum through to the emitted
|
||||
label. This would fail if the label were dropped from a metric's list or
|
||||
if the value plumbing regressed, which the allow-list assertion above
|
||||
cannot catch on its own.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.integrations.prometheus import (
|
||||
PrometheusLogger,
|
||||
UserAPIKeyLabelValues,
|
||||
prometheus_label_factory,
|
||||
)
|
||||
|
||||
prometheus_logger = MagicMock()
|
||||
prometheus_logger._cached_metric_labels = {}
|
||||
prometheus_logger.label_filters = {}
|
||||
prometheus_logger.get_labels_for_metric = (
|
||||
PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger)
|
||||
)
|
||||
|
||||
enum_values = UserAPIKeyLabelValues(
|
||||
model_group="example-model-group",
|
||||
litellm_model_name="gpt-4o-mini",
|
||||
requested_model="example-model-group",
|
||||
status_code="200",
|
||||
)
|
||||
|
||||
for metric_name in [
|
||||
"litellm_deployment_total_requests",
|
||||
"litellm_deployment_success_responses",
|
||||
"litellm_deployment_failure_responses",
|
||||
"litellm_deployment_tpm_limit",
|
||||
"litellm_deployment_rpm_limit",
|
||||
]:
|
||||
labels = prometheus_label_factory(
|
||||
supported_enum_labels=prometheus_logger.get_labels_for_metric(
|
||||
metric_name=metric_name
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
assert (
|
||||
labels.get("model_group") == "example-model-group"
|
||||
), f"{metric_name} should emit model_group=example-model-group, got {labels.get('model_group')!r}"
|
||||
|
||||
|
||||
def test_deployment_failure_metrics_emit_model_group_from_standard_logging_payload():
|
||||
"""
|
||||
End-to-end emit wiring for the failure path.
|
||||
|
||||
The label-list and factory tests above prove the label exists and that
|
||||
the factory carries a value handed to it, but neither drives the real
|
||||
set_llm_deployment_failure_metrics code path, so deleting the production
|
||||
model_group=model_group assignment there would still pass them. This
|
||||
calls it directly with a standard_logging_object carrying model_group and
|
||||
asserts the real litellm_deployment_failure_responses / _total_requests
|
||||
Counter series actually carry it.
|
||||
"""
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
logger = PrometheusLogger()
|
||||
logger.set_llm_deployment_failure_metrics(
|
||||
request_kwargs={
|
||||
"model": "gpt-4o-mini",
|
||||
"litellm_params": {"metadata": {}},
|
||||
"standard_logging_object": {
|
||||
"model_group": "example-model-group",
|
||||
"model_id": "model-123",
|
||||
"api_base": "https://api.openai.com",
|
||||
"request_tags": [],
|
||||
},
|
||||
"exception": Exception("boom"),
|
||||
}
|
||||
)
|
||||
|
||||
for metric in (
|
||||
logger.litellm_deployment_failure_responses,
|
||||
logger.litellm_deployment_total_requests,
|
||||
):
|
||||
index = metric._labelnames.index("model_group")
|
||||
values = {sample_key[index] for sample_key in metric._metrics}
|
||||
assert values == {"example-model-group"}, (
|
||||
f"expected model_group=example-model-group on {metric._name}, got {values}"
|
||||
)
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
def test_deployment_tpm_rpm_limit_metrics_emit_model_group_from_enum_values():
|
||||
"""
|
||||
End-to-end emit wiring for the tpm/rpm limit gauges.
|
||||
|
||||
_set_deployment_tpm_rpm_limit_metrics used to build its own
|
||||
UserAPIKeyLabelValues with no model_group parameter at all, dropping the
|
||||
value even though its only caller (set_llm_deployment_success_metrics)
|
||||
already had it on enum_values. This drives set_llm_deployment_success_metrics
|
||||
directly with a deployment that has tpm/rpm configured and asserts the real
|
||||
litellm_deployment_tpm_limit / litellm_deployment_rpm_limit Gauge series
|
||||
carry model_group; it fails if that plumbing is removed.
|
||||
"""
|
||||
import datetime
|
||||
|
||||
from litellm.integrations.prometheus import PrometheusLogger, UserAPIKeyLabelValues
|
||||
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
logger = PrometheusLogger()
|
||||
now = datetime.datetime.now()
|
||||
enum_values = UserAPIKeyLabelValues(
|
||||
model_group="example-model-group",
|
||||
litellm_model_name="gpt-4o-mini",
|
||||
requested_model="example-model-group",
|
||||
status_code="200",
|
||||
)
|
||||
logger.set_llm_deployment_success_metrics(
|
||||
request_kwargs={
|
||||
"model": "gpt-4o-mini",
|
||||
"litellm_params": {"metadata": {"model_info": {"id": "model-123", "tpm": 1000, "rpm": 10}}},
|
||||
"standard_logging_object": {
|
||||
"model_group": "example-model-group",
|
||||
"model_id": "model-123",
|
||||
"api_base": "https://api.openai.com",
|
||||
"hidden_params": {"additional_headers": None, "litellm_overhead_time_ms": None},
|
||||
},
|
||||
},
|
||||
start_time=now,
|
||||
end_time=now,
|
||||
enum_values=enum_values,
|
||||
)
|
||||
|
||||
for metric in (logger.litellm_deployment_tpm_limit, logger.litellm_deployment_rpm_limit):
|
||||
index = metric._labelnames.index("model_group")
|
||||
values = {sample_key[index] for sample_key in metric._metrics}
|
||||
assert values == {"example-model-group"}, (
|
||||
f"expected model_group=example-model-group on {metric._name}, got {values}"
|
||||
)
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_user_email_in_required_metrics()
|
||||
test_user_email_label_exists()
|
||||
|
|
|
|||
|
|
@ -2,22 +2,27 @@
|
|||
|
||||
import logging
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_FINISH_REASON_MAP,
|
||||
RESPONSE_COST_HEADER,
|
||||
bind_budget_reservation_to_callbacks,
|
||||
budget_reservation_from_metadata,
|
||||
drop_params_env_flag,
|
||||
drop_params_flag,
|
||||
get_or_create_metadata_bucket,
|
||||
get_provider_response_headers_from_hidden_params,
|
||||
map_finish_reason,
|
||||
normalize_drop_params,
|
||||
reconstruct_model_name,
|
||||
redact_nested_match_and_regex_keys,
|
||||
set_provider_response_headers_in_hidden_params,
|
||||
unbind_budget_reservation_from_callbacks,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.utils import ImageResponse, TranscriptionResponse
|
||||
|
||||
|
||||
class TestBudgetReservationBinding:
|
||||
|
|
@ -489,3 +494,66 @@ class TestIsExpectedClientError:
|
|||
category=RateLimitErrorCategory.VENDOR_RATE_LIMIT,
|
||||
)
|
||||
assert is_expected_client_error(vendor_limit) is False
|
||||
|
||||
|
||||
class TestProviderResponseHeadersInHiddenParams:
|
||||
def test_records_raw_headers_and_the_processed_additional_headers(self):
|
||||
response = ImageResponse()
|
||||
response._hidden_params = {"additional_headers": {RESPONSE_COST_HEADER: 0.04}}
|
||||
|
||||
set_provider_response_headers_in_hidden_params(
|
||||
response, httpx.Headers({"X-Request-Id": "req_img", "x-ratelimit-remaining-requests": "41"})
|
||||
)
|
||||
|
||||
assert response._hidden_params["headers"] == {
|
||||
"x-request-id": "req_img",
|
||||
"x-ratelimit-remaining-requests": "41",
|
||||
}
|
||||
additional_headers = response._hidden_params["additional_headers"]
|
||||
assert additional_headers["llm_provider-x-request-id"] == "req_img"
|
||||
assert additional_headers["x-ratelimit-remaining-requests"] == "41"
|
||||
assert additional_headers[RESPONSE_COST_HEADER] == 0.04
|
||||
|
||||
def test_litellm_owned_additional_headers_win_over_provider_headers(self):
|
||||
response = TranscriptionResponse(text="hi")
|
||||
response._hidden_params = {"additional_headers": {"llm_provider-x-request-id": "kept"}}
|
||||
|
||||
set_provider_response_headers_in_hidden_params(response, {"x-request-id": "provider"})
|
||||
|
||||
assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "kept"
|
||||
assert response._hidden_params["headers"] == {"x-request-id": "provider"}
|
||||
|
||||
def test_getter_returns_the_recorded_headers(self):
|
||||
response = ImageResponse()
|
||||
|
||||
set_provider_response_headers_in_hidden_params(response, {"x-request-id": "req_img"})
|
||||
|
||||
assert get_provider_response_headers_from_hidden_params(response) == {"x-request-id": "req_img"}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"hidden_params",
|
||||
[
|
||||
None,
|
||||
"headers",
|
||||
{"additional_headers": {}},
|
||||
{"headers": "x-request-id: req_img"},
|
||||
{"headers": {"x-request-id": 7}},
|
||||
],
|
||||
)
|
||||
def test_getter_returns_none_without_a_string_header_mapping(self, hidden_params):
|
||||
response = ImageResponse()
|
||||
response._hidden_params = hidden_params
|
||||
|
||||
assert get_provider_response_headers_from_hidden_params(response) is None
|
||||
|
||||
def test_getter_returns_none_for_an_object_without_hidden_params(self):
|
||||
assert get_provider_response_headers_from_hidden_params(object()) is None
|
||||
|
||||
def test_headers_never_leak_into_a_sibling_response(self):
|
||||
recorded = TranscriptionResponse()
|
||||
sibling = TranscriptionResponse()
|
||||
|
||||
set_provider_response_headers_in_hidden_params(recorded, {"x-request-id": "req_stt"})
|
||||
|
||||
assert get_provider_response_headers_from_hidden_params(sibling) is None
|
||||
assert "additional_headers" not in sibling._hidden_params
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from litellm.cost_calculator import ocr_batch_cost
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
_extract_response_obj_and_hidden_params,
|
||||
_get_status_fields,
|
||||
set_callbacks,
|
||||
)
|
||||
|
|
@ -32,6 +33,7 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.types.llms.openai import ResponseAPIUsage, ResponseCompletedEvent, ResponsesAPIResponse
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
ImageResponse,
|
||||
LiteLLMRealtimeStreamLoggingObject,
|
||||
ModelResponse,
|
||||
TextCompletionResponse,
|
||||
|
|
@ -8694,3 +8696,88 @@ async def test_async_failure_handler_delivers_failure_payload_to_custom_logger()
|
|||
assert "smoke-failure" in payload["error_str"]
|
||||
assert payload["model"] == "openai/gpt-5.6"
|
||||
assert events.empty()
|
||||
|
||||
|
||||
def _image_logging_obj() -> LitellmLogging:
|
||||
logging_obj = LitellmLogging(
|
||||
model="gpt-image-2",
|
||||
messages="a cat",
|
||||
stream=False,
|
||||
call_type="aimage_generation",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="response-headers-test",
|
||||
function_id="response-headers-test",
|
||||
)
|
||||
logging_obj.model_call_details["litellm_params"] = {"metadata": {}}
|
||||
logging_obj.optional_params = {}
|
||||
return logging_obj
|
||||
|
||||
|
||||
def _image_result_with_headers(request_id: str) -> ImageResponse:
|
||||
result = ImageResponse(created=1, data=[])
|
||||
result._hidden_params = {"headers": {"x-request-id": request_id}}
|
||||
return result
|
||||
|
||||
|
||||
def test_process_hidden_params_surfaces_response_headers_from_the_result():
|
||||
logging_obj = _image_logging_obj()
|
||||
|
||||
logging_obj._process_hidden_params_and_response_cost(
|
||||
_image_result_with_headers("req_img"), datetime.datetime.now(), datetime.datetime.now()
|
||||
)
|
||||
|
||||
assert logging_obj.model_call_details["response_headers"] == {"x-request-id": "req_img"}
|
||||
|
||||
|
||||
def test_process_hidden_params_keeps_handler_set_response_headers():
|
||||
logging_obj = _image_logging_obj()
|
||||
logging_obj.model_call_details["response_headers"] = {"x-request-id": "from-handler"}
|
||||
|
||||
logging_obj._process_hidden_params_and_response_cost(
|
||||
_image_result_with_headers("from-result"), datetime.datetime.now(), datetime.datetime.now()
|
||||
)
|
||||
|
||||
assert logging_obj.model_call_details["response_headers"] == {"x-request-id": "from-handler"}
|
||||
|
||||
|
||||
def _assembled_stream_result_with_headers() -> ModelResponse:
|
||||
result = _assembled_stream_result()
|
||||
result._hidden_params = {"headers": {"x-request-id": "req_stream"}}
|
||||
return result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_success_passes_result_headers_to_callback_kwargs():
|
||||
releasing = CustomLogger()
|
||||
releasing.async_log_success_event = AsyncMock()
|
||||
patcher, logging_obj = _streaming_logging_obj_with_callbacks([releasing])
|
||||
|
||||
with patcher:
|
||||
await logging_obj.async_success_handler(result=_assembled_stream_result_with_headers())
|
||||
|
||||
kwargs = releasing.async_log_success_event.await_args.kwargs["kwargs"]
|
||||
assert kwargs["response_headers"] == {"x-request-id": "req_stream"}
|
||||
|
||||
|
||||
def test_sync_streaming_success_passes_result_headers_to_callback_kwargs():
|
||||
releasing = CustomLogger()
|
||||
releasing.log_success_event = MagicMock()
|
||||
patcher, logging_obj = _streaming_logging_obj_with_callbacks([releasing])
|
||||
|
||||
with patcher:
|
||||
logging_obj.success_handler(result=_assembled_stream_result_with_headers())
|
||||
|
||||
kwargs = releasing.log_success_event.call_args.kwargs["kwargs"]
|
||||
assert kwargs["response_headers"] == {"x-request-id": "req_stream"}
|
||||
|
||||
|
||||
def test_extract_response_obj_and_hidden_params_reads_binary_content_hidden_params():
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent as LiteLLMBinaryResponseContent
|
||||
|
||||
result = LiteLLMBinaryResponseContent(response=httpx.Response(status_code=200, content=b"audio bytes"))
|
||||
result._hidden_params = {"headers": {"x-request-id": "req_tts"}}
|
||||
|
||||
response_obj, hidden_params = _extract_response_obj_and_hidden_params(result, None)
|
||||
|
||||
assert hidden_params == {"headers": {"x-request-id": "req_tts"}}
|
||||
assert response_obj["object"] == "binary"
|
||||
|
|
|
|||
|
|
@ -346,7 +346,7 @@ def test_route_prefix_matched_as_path_segment_not_substring():
|
|||
BedrockModelInfo.get_bedrock_route("bedrock_mantle/openai.gpt-5.5") != "mantle"
|
||||
)
|
||||
assert (
|
||||
BedrockModelInfo.get_bedrock_route("bedrock_mantle/openai.gpt-5.4") == "invoke"
|
||||
BedrockModelInfo.get_bedrock_route("bedrock_mantle/openai.gpt-5.4") == "converse"
|
||||
)
|
||||
assert (
|
||||
BedrockModelInfo._explicit_mantle_route("bedrock_mantle/openai.gpt-5.5")
|
||||
|
|
@ -996,3 +996,20 @@ def test_s3_static_key_pair_is_none_without_a_full_pair(partial_s3_pair):
|
|||
from litellm.llms.bedrock.common_utils import s3_static_key_pair
|
||||
|
||||
assert s3_static_key_pair({"aws_access_key_id": "bedrock-key", **partial_s3_pair}) is None
|
||||
|
||||
|
||||
def test_unmapped_openai_family_model_routes_to_converse():
|
||||
"""A Bedrock-native OpenAI model that is not in the cost map yet must not fall to the invoke route.
|
||||
|
||||
The invoke ``openai`` provider is the imported-model path and sends ``max_tokens``, which Bedrock
|
||||
rejects for these models; Converse maps it to ``inferenceConfig.maxTokens``.
|
||||
"""
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
|
||||
unmapped: Final = "bedrock/global.openai.gpt-99-unmapped"
|
||||
assert unmapped.removeprefix("bedrock/") not in litellm.bedrock_converse_models
|
||||
assert BedrockModelInfo.get_bedrock_route(unmapped) == "converse"
|
||||
imported: Final = "bedrock/openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/abc123"
|
||||
assert BedrockModelInfo.get_bedrock_route(imported) == "openai"
|
||||
|
|
|
|||
|
|
@ -26,6 +26,8 @@ from litellm.llms.base_llm.search.transformation import BaseSearchConfig, Search
|
|||
from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS
|
||||
from litellm.llms.brave.search.transformation import BraveSearchConfig
|
||||
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
||||
from litellm.llms.base_llm.image_generation.transformation import BaseImageGenerationConfig
|
||||
from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeechConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import (
|
||||
BaseLLMHTTPHandler,
|
||||
|
|
@ -41,7 +43,7 @@ from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_tran
|
|||
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
|
||||
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
|
||||
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent, ResponsesAPIResponse
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import ImageObject, ImageResponse, ModelResponse, TranscriptionResponse
|
||||
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
|
||||
|
|
@ -4166,3 +4168,202 @@ async def test_chat_completion_agentic_followup_does_not_repeat_request_params_f
|
|||
assert followup_calls[0]["temperature"] == 0.2
|
||||
assert followup_calls[0]["api_base"] == "https://a"
|
||||
assert followup_calls[0]["model"] == "openai/gpt-5"
|
||||
|
||||
|
||||
_UPSTREAM_HEADERS: Final = {"x-request-id": "req_upstream", "x-ratelimit-remaining-requests": "41"}
|
||||
|
||||
|
||||
def _assert_upstream_headers_recorded(response) -> None:
|
||||
assert response._hidden_params["headers"]["x-request-id"] == "req_upstream"
|
||||
assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_upstream"
|
||||
assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41"
|
||||
|
||||
|
||||
def _json_with_upstream_headers(payload: dict) -> httpx.MockTransport:
|
||||
return httpx.MockTransport(lambda request: httpx.Response(200, json=payload, headers=_UPSTREAM_HEADERS))
|
||||
|
||||
|
||||
def _binary_with_upstream_headers() -> httpx.MockTransport:
|
||||
return httpx.MockTransport(
|
||||
lambda request: httpx.Response(
|
||||
200, content=b"audio-bytes", headers={**_UPSTREAM_HEADERS, "content-type": "audio/mpeg"}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_audio_transcriptions_records_upstream_response_headers():
|
||||
client = HTTPHandler(client=httpx.Client(transport=_json_with_upstream_headers({"text": "transcribed"})))
|
||||
|
||||
response = BaseLLMHTTPHandler().audio_transcriptions(
|
||||
client=client,
|
||||
atranscription=False,
|
||||
**_json_transcription_call_kwargs(_JSONBodyAudioTranscriptionConfig()),
|
||||
)
|
||||
|
||||
_assert_upstream_headers_recorded(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_audio_transcriptions_records_upstream_response_headers():
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"text": "transcribed"}))
|
||||
|
||||
response = await BaseLLMHTTPHandler().async_audio_transcriptions(
|
||||
client=client,
|
||||
**_json_transcription_call_kwargs(_JSONBodyAudioTranscriptionConfig()),
|
||||
)
|
||||
|
||||
_assert_upstream_headers_recorded(response)
|
||||
|
||||
|
||||
def _image_edit_call_kwargs() -> dict:
|
||||
return {
|
||||
"model": "edit-model",
|
||||
"image": b"raw-image",
|
||||
"prompt": "add a hat",
|
||||
"image_edit_provider_config": _ImageEditRecordingConfig(),
|
||||
"image_edit_optional_request_params": {},
|
||||
"custom_llm_provider": "openai",
|
||||
"litellm_params": GenericLiteLLMParams(),
|
||||
"logging_obj": Mock(),
|
||||
"timeout": 10.0,
|
||||
}
|
||||
|
||||
|
||||
def test_image_edit_handler_records_upstream_response_headers():
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=_json_with_upstream_headers({"transformed_by": "sync"}))
|
||||
|
||||
response = BaseLLMHTTPHandler().image_edit_handler(client=client, **_image_edit_call_kwargs())
|
||||
|
||||
_assert_upstream_headers_recorded(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_image_edit_handler_records_upstream_response_headers():
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"transformed_by": "async"}))
|
||||
|
||||
response = await BaseLLMHTTPHandler().async_image_edit_handler(client=client, **_image_edit_call_kwargs())
|
||||
|
||||
_assert_upstream_headers_recorded(response)
|
||||
|
||||
|
||||
class _HeaderImageGenerationConfig(BaseImageGenerationConfig):
|
||||
def get_supported_openai_params(self, model):
|
||||
return []
|
||||
|
||||
def map_openai_params(self, non_default_params, optional_params, model, drop_params):
|
||||
return optional_params
|
||||
|
||||
def get_complete_url(self, api_base, api_key, model, optional_params, litellm_params, stream=None):
|
||||
return "https://images.example/v1/generations"
|
||||
|
||||
def transform_image_generation_request(self, model, prompt, optional_params, litellm_params, headers):
|
||||
return {"prompt": prompt}
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model,
|
||||
raw_response,
|
||||
model_response,
|
||||
logging_obj,
|
||||
request_data,
|
||||
optional_params,
|
||||
litellm_params,
|
||||
encoding,
|
||||
api_key=None,
|
||||
json_mode=None,
|
||||
):
|
||||
return ImageResponse(data=[ImageObject(b64_json=raw_response.json()["b64_json"])])
|
||||
|
||||
|
||||
def _image_generation_call_kwargs() -> dict:
|
||||
return {
|
||||
"model": "image-model",
|
||||
"prompt": "a cat",
|
||||
"image_generation_provider_config": _HeaderImageGenerationConfig(),
|
||||
"image_generation_optional_request_params": {},
|
||||
"custom_llm_provider": "openai",
|
||||
"litellm_params": {},
|
||||
"logging_obj": Mock(),
|
||||
"timeout": 10.0,
|
||||
}
|
||||
|
||||
|
||||
def test_image_generation_handler_records_upstream_response_headers():
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=_json_with_upstream_headers({"b64_json": "abc"}))
|
||||
|
||||
response = BaseLLMHTTPHandler().image_generation_handler(client=client, **_image_generation_call_kwargs())
|
||||
|
||||
assert response.data[0].b64_json == "abc"
|
||||
_assert_upstream_headers_recorded(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_image_generation_handler_records_upstream_response_headers():
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"b64_json": "abc"}))
|
||||
|
||||
response = await BaseLLMHTTPHandler().async_image_generation_handler(
|
||||
client=client, **_image_generation_call_kwargs()
|
||||
)
|
||||
|
||||
assert response.data[0].b64_json == "abc"
|
||||
_assert_upstream_headers_recorded(response)
|
||||
|
||||
|
||||
class _HeaderTextToSpeechConfig(BaseTextToSpeechConfig):
|
||||
def get_supported_openai_params(self, model):
|
||||
return []
|
||||
|
||||
def map_openai_params(self, model, optional_params, voice=None, drop_params=False, kwargs=None):
|
||||
return voice, optional_params
|
||||
|
||||
def validate_environment(self, headers, model, api_key=None, api_base=None):
|
||||
return {}
|
||||
|
||||
def get_complete_url(self, model, api_base, litellm_params):
|
||||
return "https://tts.example/v1/speech"
|
||||
|
||||
def transform_text_to_speech_request(self, model, input, voice, optional_params, litellm_params, headers):
|
||||
return {"dict_body": {"input": input}}
|
||||
|
||||
def transform_text_to_speech_response(self, model, raw_response, logging_obj):
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
|
||||
|
||||
def _text_to_speech_call_kwargs() -> dict:
|
||||
return {
|
||||
"model": "tts-model",
|
||||
"input": "hello",
|
||||
"voice": "alloy",
|
||||
"text_to_speech_provider_config": _HeaderTextToSpeechConfig(),
|
||||
"text_to_speech_optional_params": {},
|
||||
"custom_llm_provider": "openai",
|
||||
"litellm_params": {},
|
||||
"logging_obj": Mock(),
|
||||
"timeout": 10.0,
|
||||
}
|
||||
|
||||
|
||||
def test_text_to_speech_handler_records_upstream_response_headers():
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=_binary_with_upstream_headers())
|
||||
|
||||
response = BaseLLMHTTPHandler().text_to_speech_handler(client=client, **_text_to_speech_call_kwargs())
|
||||
|
||||
assert response.content == b"audio-bytes"
|
||||
_assert_upstream_headers_recorded(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_text_to_speech_handler_records_upstream_response_headers():
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=_binary_with_upstream_headers())
|
||||
|
||||
response = await BaseLLMHTTPHandler().async_text_to_speech_handler(client=client, **_text_to_speech_call_kwargs())
|
||||
|
||||
assert response.content == b"audio-bytes"
|
||||
_assert_upstream_headers_recorded(response)
|
||||
|
|
|
|||
|
|
@ -1,13 +1,15 @@
|
|||
import asyncio
|
||||
import json
|
||||
from typing import Final
|
||||
from unittest.mock import Mock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openai import AsyncOpenAI
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
||||
import litellm
|
||||
from litellm.llms.openai.openai import OpenAIChatCompletion
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -253,3 +255,98 @@ async def test_acompletion_streams_tool_call_arguments_over_injected_transport()
|
|||
assert tool_call.function.name == "get_weather"
|
||||
assert json.loads(tool_call.function.arguments) == {"city": "Paris"}
|
||||
assert rebuilt.choices[0].finish_reason == "tool_calls"
|
||||
|
||||
|
||||
_PROVIDER_HEADERS: Final = {"x-request-id": "req_openai", "x-ratelimit-remaining-requests": "41"}
|
||||
|
||||
|
||||
def _image_generation_transport() -> httpx.MockTransport:
|
||||
return httpx.MockTransport(
|
||||
lambda request: httpx.Response(
|
||||
200, json={"created": 1, "data": [{"b64_json": "abc"}]}, headers=_PROVIDER_HEADERS
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _speech_transport() -> httpx.MockTransport:
|
||||
return httpx.MockTransport(
|
||||
lambda request: httpx.Response(
|
||||
200, content=b"audio-bytes", headers={**_PROVIDER_HEADERS, "content-type": "audio/mpeg"}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _assert_provider_headers_recorded(response) -> None:
|
||||
assert response._hidden_params["headers"]["x-request-id"] == "req_openai"
|
||||
assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_openai"
|
||||
assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41"
|
||||
|
||||
|
||||
def _image_generation_kwargs() -> dict:
|
||||
return {
|
||||
"model": "gpt-image-2",
|
||||
"prompt": "a cat",
|
||||
"timeout": 10,
|
||||
"optional_params": {},
|
||||
"logging_obj": Mock(),
|
||||
"api_key": "transport-only",
|
||||
"model_response": ImageResponse(),
|
||||
}
|
||||
|
||||
|
||||
def test_image_generation_records_provider_response_headers():
|
||||
with httpx.Client(transport=_image_generation_transport()) as http_client:
|
||||
response = OpenAIChatCompletion().image_generation(
|
||||
client=OpenAI(api_key="transport-only", http_client=http_client), **_image_generation_kwargs()
|
||||
)
|
||||
|
||||
_assert_provider_headers_recorded(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aimage_generation_records_provider_response_headers():
|
||||
async with httpx.AsyncClient(transport=_image_generation_transport()) as http_client:
|
||||
response = await OpenAIChatCompletion().image_generation(
|
||||
client=AsyncOpenAI(api_key="transport-only", http_client=http_client),
|
||||
aimg_generation=True,
|
||||
**_image_generation_kwargs(),
|
||||
)
|
||||
|
||||
_assert_provider_headers_recorded(response)
|
||||
|
||||
|
||||
def _audio_speech_kwargs() -> dict:
|
||||
return {
|
||||
"model": "gpt-4o-mini-tts",
|
||||
"input": "hello",
|
||||
"voice": "alloy",
|
||||
"optional_params": {},
|
||||
"api_key": "transport-only",
|
||||
"api_base": None,
|
||||
"organization": None,
|
||||
"project": None,
|
||||
"max_retries": 0,
|
||||
"timeout": 10,
|
||||
"logging_obj": Mock(),
|
||||
}
|
||||
|
||||
|
||||
def test_audio_speech_records_provider_response_headers():
|
||||
with httpx.Client(transport=_speech_transport()) as http_client:
|
||||
response = OpenAIChatCompletion().audio_speech(
|
||||
client=OpenAI(api_key="transport-only", http_client=http_client), **_audio_speech_kwargs()
|
||||
)
|
||||
|
||||
_assert_provider_headers_recorded(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_audio_speech_records_provider_response_headers():
|
||||
async with httpx.AsyncClient(transport=_speech_transport()) as http_client:
|
||||
response = await OpenAIChatCompletion().audio_speech(
|
||||
client=AsyncOpenAI(api_key="transport-only", http_client=http_client),
|
||||
aspeech=True,
|
||||
**_audio_speech_kwargs(),
|
||||
)
|
||||
|
||||
_assert_provider_headers_recorded(response)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,71 @@
|
|||
from typing import Final
|
||||
from unittest.mock import Mock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
||||
from litellm.llms.openai.transcriptions.handler import OpenAIAudioTranscription
|
||||
from litellm.types.utils import TranscriptionResponse
|
||||
|
||||
_PROVIDER_HEADERS: Final = {"x-request-id": "req_stt", "x-ratelimit-remaining-requests": "41"}
|
||||
|
||||
|
||||
def _transcription_transport() -> httpx.MockTransport:
|
||||
return httpx.MockTransport(lambda request: httpx.Response(200, json={"text": "hello"}, headers=_PROVIDER_HEADERS))
|
||||
|
||||
|
||||
def _logging_obj() -> Mock:
|
||||
logging_obj = Mock()
|
||||
logging_obj.model_call_details = {}
|
||||
return logging_obj
|
||||
|
||||
|
||||
def _call_kwargs(logging_obj: Mock) -> dict:
|
||||
return {
|
||||
"model": "gpt-4o-mini-transcribe",
|
||||
"audio_file": ("audio.wav", b"riff-bytes", "audio/wav"),
|
||||
"optional_params": {},
|
||||
"litellm_params": {},
|
||||
"model_response": TranscriptionResponse(),
|
||||
"timeout": 10.0,
|
||||
"max_retries": 0,
|
||||
"logging_obj": logging_obj,
|
||||
"api_key": "transport-only",
|
||||
"api_base": None,
|
||||
}
|
||||
|
||||
|
||||
def _assert_headers_recorded(response: TranscriptionResponse, logging_obj: Mock) -> None:
|
||||
assert response.text == "hello"
|
||||
assert response._hidden_params["headers"]["x-request-id"] == "req_stt"
|
||||
assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_stt"
|
||||
assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41"
|
||||
assert logging_obj.model_call_details["response_headers"]["x-request-id"] == "req_stt"
|
||||
|
||||
|
||||
def test_audio_transcriptions_records_provider_response_headers():
|
||||
logging_obj = _logging_obj()
|
||||
|
||||
with httpx.Client(transport=_transcription_transport()) as http_client:
|
||||
response = OpenAIAudioTranscription().audio_transcriptions(
|
||||
client=OpenAI(api_key="transport-only", http_client=http_client),
|
||||
atranscription=False,
|
||||
**_call_kwargs(logging_obj),
|
||||
)
|
||||
|
||||
_assert_headers_recorded(response, logging_obj)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_audio_transcriptions_records_provider_response_headers():
|
||||
logging_obj = _logging_obj()
|
||||
|
||||
async with httpx.AsyncClient(transport=_transcription_transport()) as http_client:
|
||||
response = await OpenAIAudioTranscription().audio_transcriptions(
|
||||
client=AsyncOpenAI(api_key="transport-only", http_client=http_client),
|
||||
atranscription=True,
|
||||
**_call_kwargs(logging_obj),
|
||||
)
|
||||
|
||||
_assert_headers_recorded(response, logging_obj)
|
||||
|
|
@ -19,7 +19,6 @@ import pytest
|
|||
|
||||
from litellm.llms.vertex_ai.batches.transformation import ( # noqa: E402
|
||||
VertexAIBatchTransformation,
|
||||
vertex_prompt_tokens_details,
|
||||
)
|
||||
from litellm.llms.vertex_ai.common_utils import ( # noqa: E402
|
||||
VertexAIError,
|
||||
|
|
@ -41,27 +40,6 @@ ENDPOINT_INPUT_FILE = (
|
|||
)
|
||||
|
||||
|
||||
def test_vertex_prompt_tokens_details_rejects_malformed_details():
|
||||
assert vertex_prompt_tokens_details({"promptTokensDetails": [1]}) is None
|
||||
assert vertex_prompt_tokens_details({"promptTokensDetails": [{"modality": "AUDIO"}]}) is None
|
||||
assert (
|
||||
vertex_prompt_tokens_details(
|
||||
{
|
||||
"promptTokensDetails": [
|
||||
{"modality": "AUDIO", "tokenCount": 1},
|
||||
"malformed",
|
||||
]
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# transform_openai_batch_request_to_vertex_ai_batch_request
|
||||
# =========================================================================== #
|
||||
|
||||
|
||||
def test_transform_openai_request_builds_full_vertex_job():
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.batches.transformation.uuid.uuid4",
|
||||
|
|
@ -477,3 +455,19 @@ def test_list_response_none_jobs_treated_as_empty():
|
|||
out = T.transform_vertex_ai_batch_list_response_to_openai_list_response({"batchPredictionJobs": None})
|
||||
assert out["data"] == []
|
||||
assert out["first_id"] is None
|
||||
|
||||
|
||||
PASSTHROUGH_INPUT_FILE = (
|
||||
"gs://litellm-testing-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/uuid-1"
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_from_passthrough_gcs_file():
|
||||
assert T._get_model_from_gcs_file(PASSTHROUGH_INPUT_FILE) == "publishers/google/models/gemini-2.5-flash"
|
||||
|
||||
|
||||
def test_get_gcs_uri_prefix_keeps_passthrough_segment_so_output_lands_beside_input():
|
||||
assert (
|
||||
T._get_gcs_uri_prefix_from_file(PASSTHROUGH_INPUT_FILE)
|
||||
== "gs://litellm-testing-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash"
|
||||
)
|
||||
|
|
|
|||
0
tests/test_litellm/llms/vertex_ai/files/__init__.py
Normal file
0
tests/test_litellm/llms/vertex_ai/files/__init__.py
Normal file
310
tests/test_litellm/llms/vertex_ai/files/test_transformation.py
Normal file
310
tests/test_litellm/llms/vertex_ai/files/test_transformation.py
Normal file
|
|
@ -0,0 +1,310 @@
|
|||
import io
|
||||
import json
|
||||
import urllib.parse
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAIError
|
||||
from litellm.llms.vertex_ai.files.transformation import VertexAIFilesConfig, is_passthrough_managed_gcs_url
|
||||
|
||||
NATIVE_VERTEX_ROW = json.dumps(
|
||||
{
|
||||
"request": {
|
||||
"contents": [{"role": "user", "parts": [{"text": "What is the tallest building in the world?"}]}],
|
||||
"tools": [{"googleSearch": {"excludeDomains": ["example.com"]}}],
|
||||
}
|
||||
}
|
||||
).encode()
|
||||
NATIVE_VERTEX_JSONL = NATIVE_VERTEX_ROW + b"\n" + NATIVE_VERTEX_ROW + b"\n"
|
||||
OPENAI_BATCH_JSONL = (
|
||||
b'{"custom_id": "r1", "method": "POST", "url": "/v1/chat/completions",'
|
||||
b' "body": {"model": "gemini-2.5-flash", "messages": [{"role": "user", "content": "hi"}]}}\n'
|
||||
)
|
||||
PASSTHROUGH_OBJECT = (
|
||||
"litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/uuid-1/predictions.jsonl"
|
||||
)
|
||||
TRANSFORMED_OBJECT = "litellm-vertex-files/publishers/google/models/gemini-2.5-flash/uuid-1/predictions.jsonl"
|
||||
UPLOAD_CHUNK_BYTES = 1024 * 1024
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def config() -> VertexAIFilesConfig:
|
||||
return VertexAIFilesConfig()
|
||||
|
||||
|
||||
def _gcs_media_url(object_name: str) -> str:
|
||||
return (
|
||||
f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{urllib.parse.quote(object_name, safe='')}?alt=media"
|
||||
)
|
||||
|
||||
|
||||
def _native_output_jsonl() -> bytes:
|
||||
return (
|
||||
json.dumps(
|
||||
{
|
||||
"request": json.loads(NATIVE_VERTEX_ROW)["request"],
|
||||
"status": "",
|
||||
"response": {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"role": "model", "parts": [{"text": "The Burj Khalifa."}]},
|
||||
"finishReason": "STOP",
|
||||
"groundingMetadata": {"webSearchQueries": ["tallest building in the world"]},
|
||||
}
|
||||
],
|
||||
"modelVersion": "gemini-2.5-flash",
|
||||
"usageMetadata": {"promptTokenCount": 20, "candidatesTokenCount": 48, "totalTokenCount": 68},
|
||||
},
|
||||
"processed_time": "2026-09-23T19:02:00.000+00:00",
|
||||
}
|
||||
).encode()
|
||||
+ b"\n"
|
||||
)
|
||||
|
||||
|
||||
def _upload_chunks(config: VertexAIFilesConfig, file: object, litellm_params: dict) -> list[bytes]:
|
||||
body = config.transform_create_file_request(
|
||||
model="",
|
||||
create_file_data={"file": file, "purpose": "batch"},
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
return list(body["streaming_media_upload"]["body_stream"].iter_bytes())
|
||||
|
||||
|
||||
def _upload_body_bytes(config: VertexAIFilesConfig, file: object, litellm_params: dict) -> bytes:
|
||||
return b"".join(_upload_chunks(config, file, litellm_params))
|
||||
|
||||
|
||||
class TestPassthroughBatchUpload:
|
||||
"""`passthrough=True` on a batch upload ships the caller's native Vertex JSONL
|
||||
to GCS byte for byte, filed under a `passthrough/` object path so the batch
|
||||
output that lands beside it is recognized and returned untouched as well."""
|
||||
|
||||
def _upload_url(self, config, litellm_params, file, purpose="batch") -> str:
|
||||
return config.get_complete_file_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="",
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
data={"file": file, "purpose": purpose},
|
||||
)
|
||||
|
||||
def test_passthrough_object_is_filed_under_passthrough_prefix_named_by_deployment_model(self, config):
|
||||
url = self._upload_url(
|
||||
config,
|
||||
{"gcs_bucket_name": "my-bucket", "model": "vertex_ai/gemini-2.5-flash", "passthrough": True},
|
||||
("batch.jsonl", NATIVE_VERTEX_JSONL, "application/jsonl"),
|
||||
)
|
||||
object_name = parse_qs(urlparse(url).query)["name"][0]
|
||||
assert object_name.startswith("litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/")
|
||||
|
||||
def test_passthrough_upload_without_deployment_model_is_rejected(self, config):
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
self._upload_url(
|
||||
config,
|
||||
{"gcs_bucket_name": "my-bucket", "passthrough": True},
|
||||
("batch.jsonl", NATIVE_VERTEX_JSONL, "application/jsonl"),
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "target_model_names" in exc_info.value.message
|
||||
|
||||
def test_passthrough_flag_does_not_ship_a_non_batch_upload_raw(self, config):
|
||||
result = config.transform_create_file_request(
|
||||
model="",
|
||||
create_file_data={"file": ("notes.txt", b"plain text", "text/plain"), "purpose": "user_data"},
|
||||
optional_params={},
|
||||
litellm_params={"gcs_bucket_name": "my-bucket", "passthrough": True},
|
||||
)
|
||||
assert result == b"plain text"
|
||||
|
||||
def test_passthrough_flag_is_ignored_for_non_batch_purposes(self, config):
|
||||
url = self._upload_url(
|
||||
config,
|
||||
{"gcs_bucket_name": "my-bucket", "model": "vertex_ai/gemini-2.5-flash", "passthrough": True},
|
||||
("notes.txt", b"plain text", "text/plain"),
|
||||
purpose="user_data",
|
||||
)
|
||||
object_name = parse_qs(urlparse(url).query)["name"][0]
|
||||
assert object_name.startswith("litellm-vertex-files/uploads/")
|
||||
assert "passthrough" not in object_name
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"file",
|
||||
[
|
||||
("batch.jsonl", NATIVE_VERTEX_JSONL, "application/jsonl"),
|
||||
NATIVE_VERTEX_JSONL,
|
||||
("batch.jsonl", io.BytesIO(NATIVE_VERTEX_JSONL), "application/jsonl"),
|
||||
("batch.jsonl", NATIVE_VERTEX_JSONL.decode(), "application/jsonl"),
|
||||
],
|
||||
ids=["bytes-tuple", "bare-bytes", "handle-tuple", "text-tuple"],
|
||||
)
|
||||
def test_passthrough_upload_body_is_the_callers_bytes(self, config, file):
|
||||
body = config.transform_create_file_request(
|
||||
model="",
|
||||
create_file_data={"file": file, "purpose": "batch"},
|
||||
optional_params={},
|
||||
litellm_params={"passthrough": True},
|
||||
)
|
||||
stream = body["streaming_media_upload"]["body_stream"]
|
||||
assert b"".join(stream.iter_bytes()) == NATIVE_VERTEX_JSONL
|
||||
assert b"".join(stream.iter_bytes()) == NATIVE_VERTEX_JSONL
|
||||
assert body["streaming_media_upload"]["content_type"] == "application/json"
|
||||
|
||||
def test_passthrough_upload_streams_a_large_handle_in_bounded_chunks(self, config):
|
||||
content = NATIVE_VERTEX_ROW * (3 * UPLOAD_CHUNK_BYTES // len(NATIVE_VERTEX_ROW) + 1)
|
||||
chunks = _upload_chunks(
|
||||
config, ("batch.jsonl", io.BytesIO(content), "application/jsonl"), {"passthrough": True}
|
||||
)
|
||||
assert len(chunks) >= 3
|
||||
assert max(len(chunk) for chunk in chunks) <= UPLOAD_CHUNK_BYTES
|
||||
assert b"".join(chunks) == content
|
||||
|
||||
def test_passthrough_upload_streams_a_path_in_bounded_chunks(self, config, tmp_path: Path):
|
||||
content = NATIVE_VERTEX_ROW * (2 * UPLOAD_CHUNK_BYTES // len(NATIVE_VERTEX_ROW) + 1)
|
||||
batch_path = tmp_path / "batch.jsonl"
|
||||
batch_path.write_bytes(content)
|
||||
chunks = _upload_chunks(config, ("batch.jsonl", batch_path, "application/jsonl"), {"passthrough": True})
|
||||
assert len(chunks) >= 2
|
||||
assert max(len(chunk) for chunk in chunks) <= UPLOAD_CHUNK_BYTES
|
||||
assert b"".join(chunks) == content
|
||||
|
||||
def test_passthrough_upload_rejects_a_non_seekable_handle(self, config):
|
||||
class _Pipe:
|
||||
def read(self, size=-1):
|
||||
return b""
|
||||
|
||||
with pytest.raises(ValueError, match="seekable"):
|
||||
_upload_body_bytes(config, ("batch.jsonl", _Pipe(), "application/jsonl"), {"passthrough": True})
|
||||
|
||||
def test_passthrough_upload_rejects_content_that_is_neither_bytes_path_nor_handle(self, config):
|
||||
with pytest.raises(ValueError, match="Unsupported file content type"):
|
||||
_upload_body_bytes(config, ("batch.jsonl", 42, "application/jsonl"), {"passthrough": True})
|
||||
|
||||
def test_openai_rows_are_translated_unless_passthrough_is_set(self, config):
|
||||
file = ("batch.jsonl", OPENAI_BATCH_JSONL, "application/jsonl")
|
||||
translated = _upload_body_bytes(config, file, {})
|
||||
untouched = _upload_body_bytes(config, file, {"passthrough": True})
|
||||
assert untouched == OPENAI_BATCH_JSONL
|
||||
assert translated != OPENAI_BATCH_JSONL
|
||||
assert b'"contents"' in translated
|
||||
|
||||
def test_passthrough_output_content_is_returned_untouched(self, config):
|
||||
raw_jsonl = _native_output_jsonl()
|
||||
|
||||
def _download(object_name: str) -> bytes:
|
||||
raw_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=raw_jsonl,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request("GET", _gcs_media_url(object_name)),
|
||||
)
|
||||
result = config.transform_file_content_response(
|
||||
raw_response=raw_response, logging_obj=MagicMock(), litellm_params={}
|
||||
)
|
||||
return result.response.content
|
||||
|
||||
assert _download(PASSTHROUGH_OBJECT) == raw_jsonl
|
||||
assert _download(f"team-a/{PASSTHROUGH_OBJECT}") == raw_jsonl
|
||||
transformed = _download(TRANSFORMED_OBJECT)
|
||||
assert transformed != raw_jsonl
|
||||
assert json.loads(transformed.splitlines()[0])["response"]["body"]["choices"]
|
||||
nested = _download(f"litellm-vertex-files/{PASSTHROUGH_OBJECT}")
|
||||
assert nested != raw_jsonl
|
||||
assert json.loads(nested.splitlines()[0])["response"]["body"]["choices"]
|
||||
|
||||
def test_output_of_an_upload_whose_model_smuggles_the_passthrough_segment_is_still_transformed(self, config):
|
||||
smuggled_model = b"litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash"
|
||||
upload_url = self._upload_url(
|
||||
config,
|
||||
{"gcs_bucket_name": "my-bucket"},
|
||||
("batch.jsonl", OPENAI_BATCH_JSONL.replace(b"gemini-2.5-flash", smuggled_model), "application/jsonl"),
|
||||
)
|
||||
object_name = parse_qs(urlparse(upload_url).query)["name"][0]
|
||||
raw_jsonl = _native_output_jsonl()
|
||||
raw_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=raw_jsonl,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request("GET", _gcs_media_url(f"{object_name}/predictions.jsonl")),
|
||||
)
|
||||
|
||||
result = config.transform_file_content_response(
|
||||
raw_response=raw_response, logging_obj=MagicMock(), litellm_params={}
|
||||
)
|
||||
|
||||
assert object_name.startswith("litellm-vertex-files/litellm-vertex-files/passthrough/")
|
||||
assert json.loads(result.response.content.splitlines()[0])["response"]["body"]["choices"]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url, expected",
|
||||
[
|
||||
(f"gs://my-bucket/{PASSTHROUGH_OBJECT}", True),
|
||||
(f"gs://my-bucket/team-a/{PASSTHROUGH_OBJECT}", True),
|
||||
(f"gs://my-bucket/litellm-vertex-files/{PASSTHROUGH_OBJECT}", False),
|
||||
(_gcs_media_url(f"team-a/sub/{PASSTHROUGH_OBJECT}"), True),
|
||||
(_gcs_media_url(f"litellm-vertex-files/publishers/google/models/x/{PASSTHROUGH_OBJECT}"), False),
|
||||
(_gcs_media_url(TRANSFORMED_OBJECT), False),
|
||||
],
|
||||
ids=["gs", "gs-prefixed", "gs-smuggled", "https-prefixed", "https-model-path-smuggled", "https-transformed"],
|
||||
)
|
||||
def test_passthrough_detection_anchors_on_the_first_managed_segment(self, url, expected):
|
||||
assert is_passthrough_managed_gcs_url(url) is expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_passthrough_output_stream_is_returned_untouched(self, config):
|
||||
stream_iterator = object()
|
||||
headers = {"content-type": "application/octet-stream"}
|
||||
result = await config.transform_file_content_stream(
|
||||
stream_iterator=stream_iterator,
|
||||
headers=headers,
|
||||
request_url=_gcs_media_url(f"team-a/{PASSTHROUGH_OBJECT}"),
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params={},
|
||||
)
|
||||
assert result.stream_iterator is stream_iterator
|
||||
assert result.headers == headers
|
||||
|
||||
|
||||
class TestEmbeddingOutputTranslation:
|
||||
EMBEDDING_OBJECT = (
|
||||
"litellm-vertex-files/publishers/google/models/gemini-embedding-2/prediction-model-1/predictions.jsonl"
|
||||
)
|
||||
|
||||
def _transform(self, config: VertexAIFilesConfig, rows: list[dict]) -> list[dict]:
|
||||
raw_response = httpx.Response(
|
||||
status_code=200,
|
||||
content="\n".join(json.dumps(row) for row in rows).encode(),
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request("GET", _gcs_media_url(self.EMBEDDING_OBJECT)),
|
||||
)
|
||||
result = config.transform_file_content_response(
|
||||
raw_response=raw_response, logging_obj=MagicMock(), litellm_params={}
|
||||
)
|
||||
return [json.loads(line) for line in result.response.content.decode().splitlines()]
|
||||
|
||||
def test_embedding_rows_become_openai_batch_rows_billed_by_their_prompt_tokens(self, config):
|
||||
live_row = {
|
||||
"key": "request-1",
|
||||
"request": {"content": {"parts": [{"text": "hello world"}]}},
|
||||
"response": {"embedding": {"values": [-0.015, 0.024]}, "usageMetadata": {"promptTokenCount": 2}},
|
||||
}
|
||||
documented_row = {
|
||||
"key": "request-2",
|
||||
"request": {"content": {"parts": [{"text": "hello"}]}},
|
||||
"response": {"embedding": {"values": [0.5]}, "tokenCount": "3"},
|
||||
}
|
||||
|
||||
live, documented = self._transform(config, [live_row, documented_row])
|
||||
|
||||
assert (live["custom_id"], live["error"], live["response"]["status_code"]) == ("request-1", None, 200)
|
||||
assert live["response"]["body"]["model"] == "gemini-embedding-2"
|
||||
assert live["response"]["body"]["data"] == [{"embedding": [-0.015, 0.024], "index": 0, "object": "embedding"}]
|
||||
live_usage, documented_usage = (row["response"]["body"]["usage"] for row in (live, documented))
|
||||
assert (live_usage["prompt_tokens"], live_usage["total_tokens"]) == (2, 2)
|
||||
assert (documented_usage["prompt_tokens"], documented_usage["total_tokens"]) == (3, 3)
|
||||
|
|
@ -10201,6 +10201,64 @@ async def test_active_request_ctx_var_feeds_get_current_session(_mcp_request_ctx
|
|||
assert _get_current_session() is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("method", "path", "session_headers"),
|
||||
(
|
||||
("POST", "/mcp", ()),
|
||||
("GET", "/mcp", (("mcp-session-id", "existing-session"),)),
|
||||
("DELETE", "/mcp", (("mcp-session-id", "existing-session"),)),
|
||||
("POST", "/server/mcp", ()),
|
||||
("GET", "/sse", ()),
|
||||
("POST", "/sse/messages", ()),
|
||||
),
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
("allowed_origins", "origin_headers", "expected_status"),
|
||||
(
|
||||
(("https://allowed.example",), (("origin", "https://evil.example"),), 403),
|
||||
(("https://allowed.example",), (("origin", "https://allowed.example.evil.example"),), 403),
|
||||
(("https://allowed.example",), (("origin", "null"),), 403),
|
||||
(("https://allowed.example",), (("origin", ""),), 403),
|
||||
(
|
||||
("https://allowed.example",),
|
||||
(("origin", "https://allowed.example"), ("origin", "https://evil.example")),
|
||||
403,
|
||||
),
|
||||
(("https://allowed.example",), (("origin", "https://allowed.example"),), 401),
|
||||
(("https://allowed.example",), (), 401),
|
||||
(("*",), (("origin", "https://another.example"),), 401),
|
||||
),
|
||||
)
|
||||
async def test_mcp_origin_admission_precedes_authentication(
|
||||
method: str,
|
||||
path: str,
|
||||
session_headers: tuple[tuple[str, str], ...],
|
||||
allowed_origins: tuple[str, ...],
|
||||
origin_headers: tuple[tuple[str, str], ...],
|
||||
expected_status: int,
|
||||
) -> None:
|
||||
import httpx
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
authenticate: Final = AsyncMock(side_effect=HTTPException(status_code=401, detail="authentication required"))
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.origins", allowed_origins),
|
||||
patch.object(server, "extract_mcp_auth_context", authenticate),
|
||||
):
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client:
|
||||
response: Final = await client.request(method, path, headers=(*session_headers, *origin_headers))
|
||||
|
||||
assert response.status_code == expected_status
|
||||
if expected_status == 403:
|
||||
assert response.json() == {"detail": "Invalid Origin header"}
|
||||
authenticate.assert_not_awaited()
|
||||
else:
|
||||
assert response.json() == {"detail": "authentication required"}
|
||||
authenticate.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_active_request_ctx_var_feeds_auth_resolution_recording(_mcp_request_ctx) -> None:
|
||||
from starlette.requests import Request
|
||||
|
|
|
|||
|
|
@ -3517,6 +3517,7 @@ def test_internal_user_still_blocked_from_another_users_info():
|
|||
[
|
||||
"/user/daily/activity",
|
||||
"/user/daily/activity/aggregated",
|
||||
"/user/daily/activity/aggregated/search",
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -3599,6 +3600,55 @@ def test_user_daily_activity_aggregated_not_covered_by_prefix_match():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route",
|
||||
[
|
||||
"/team/daily/activity",
|
||||
"/team/daily/activity/aggregated",
|
||||
"/team/daily/activity/aggregated/search",
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"user_role",
|
||||
[
|
||||
LitellmUserRoles.INTERNAL_USER.value,
|
||||
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
|
||||
],
|
||||
)
|
||||
def test_team_daily_activity_routes_reachable_by_non_admin(route, user_role):
|
||||
"""The Team Usage dashboard calls all three team daily-activity routes, and
|
||||
each handler self-scopes to the caller's teams and own keys
|
||||
(_resolve_team_daily_activity_scope). self_managed_routes is the only list
|
||||
granting them to a non-admin, and check_route_access is exact-match, so each
|
||||
sub-path needs its own entry: dropping one 401s the dashboard before the
|
||||
handler ever runs.
|
||||
"""
|
||||
user_obj = LiteLLM_UserTable(
|
||||
user_id="test_user",
|
||||
user_email="test@example.com",
|
||||
user_role=user_role,
|
||||
)
|
||||
valid_token = UserAPIKeyAuth(user_id="test_user", user_role=user_role)
|
||||
request = MagicMock(spec=Request)
|
||||
request.query_params = {}
|
||||
|
||||
def outcome() -> str:
|
||||
try:
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=user_role,
|
||||
route=route,
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={},
|
||||
)
|
||||
except Exception as exc:
|
||||
return f"denied: {exc}"
|
||||
return "allowed"
|
||||
|
||||
assert outcome() == "allowed"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"user_role",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -2659,6 +2659,172 @@ async def test_get_user_daily_activity_aggregated_non_admin_cannot_view_other_us
|
|||
assert mock_get_daily_agg.call_args.kwargs["entity_id"] == "regular-user-123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_user_daily_activity_keys_passes_matched_tokens_to_aggregation(monkeypatch):
|
||||
"""The search endpoint resolves matching verification tokens by hash, alias, or
|
||||
user id, then aggregates daily spend for exactly those tokens. This is what lets
|
||||
the Usage page find keys outside the top-spend subset the aggregated endpoint caps."""
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
search_user_daily_activity_keys,
|
||||
)
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(token="tok-a"), SimpleNamespace(token="tok-b")]
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_get_daily_agg = AsyncMock(return_value=mock_response)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated",
|
||||
mock_get_daily_agg,
|
||||
)
|
||||
|
||||
admin_key_dict = UserAPIKeyAuth(
|
||||
user_id="admin-user-001",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
result = await search_user_daily_activity_keys(
|
||||
search="gamma",
|
||||
start_date="2025-02-01",
|
||||
end_date="2025-02-28",
|
||||
user_id=None,
|
||||
timezone=480,
|
||||
include_current_utc_day=False,
|
||||
user_api_key_dict=admin_key_dict,
|
||||
)
|
||||
|
||||
assert result is mock_response
|
||||
|
||||
find_many_kwargs = mock_prisma_client.db.litellm_verificationtoken.find_many.call_args.kwargs
|
||||
assert find_many_kwargs["take"] == USAGE_TOP_API_KEYS_LIMIT
|
||||
assert find_many_kwargs["where"]["OR"] == (
|
||||
{"token": "gamma"},
|
||||
{"key_alias": {"contains": "gamma", "mode": "insensitive"}},
|
||||
{"user_id": {"contains": "gamma", "mode": "insensitive"}},
|
||||
)
|
||||
assert "user_id" not in find_many_kwargs["where"]
|
||||
|
||||
mock_get_daily_agg.assert_called_once_with(
|
||||
prisma_client=mock_prisma_client,
|
||||
table_name="litellm_dailyuserspend",
|
||||
entity_id_field="user_id",
|
||||
entity_id=None,
|
||||
entity_metadata_field=None,
|
||||
start_date="2025-02-01",
|
||||
end_date="2025-02-28",
|
||||
model=None,
|
||||
api_key=["tok-a", "tok-b"],
|
||||
timezone_offset_minutes=480,
|
||||
include_current_utc_day=False,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_user_daily_activity_keys_no_match_returns_empty_without_aggregating(monkeypatch):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
search_user_daily_activity_keys,
|
||||
)
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
mock_get_daily_agg = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated",
|
||||
mock_get_daily_agg,
|
||||
)
|
||||
|
||||
admin_key_dict = UserAPIKeyAuth(
|
||||
user_id="admin-user-001",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
result = await search_user_daily_activity_keys(
|
||||
search="nothing-matches",
|
||||
start_date="2025-02-01",
|
||||
end_date="2025-02-28",
|
||||
user_id=None,
|
||||
timezone=None,
|
||||
include_current_utc_day=False,
|
||||
user_api_key_dict=admin_key_dict,
|
||||
)
|
||||
|
||||
assert result.results == []
|
||||
assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT
|
||||
assert result.metadata.total_api_keys == 0
|
||||
mock_get_daily_agg.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_user_daily_activity_keys_non_admin_scoped_to_caller(monkeypatch):
|
||||
"""Same scoping contract as the aggregated route: a non-admin with no user_id
|
||||
is scoped to their own rows, and any other user_id is a 403."""
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
search_user_daily_activity_keys,
|
||||
)
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[SimpleNamespace(token="tok-a")])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
non_admin_key_dict = UserAPIKeyAuth(
|
||||
user_id="user-1",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_get_daily_agg = AsyncMock(return_value=mock_response)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated",
|
||||
mock_get_daily_agg,
|
||||
)
|
||||
|
||||
result = await search_user_daily_activity_keys(
|
||||
search="gamma",
|
||||
start_date="2025-02-01",
|
||||
end_date="2025-02-28",
|
||||
user_id=None,
|
||||
timezone=None,
|
||||
include_current_utc_day=False,
|
||||
user_api_key_dict=non_admin_key_dict,
|
||||
)
|
||||
|
||||
assert result is mock_response
|
||||
assert mock_get_daily_agg.call_args.kwargs["entity_id"] == "user-1"
|
||||
find_many_kwargs = mock_prisma_client.db.litellm_verificationtoken.find_many.call_args.kwargs
|
||||
assert find_many_kwargs["where"]["user_id"] == "user-1"
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await search_user_daily_activity_keys(
|
||||
search="gamma",
|
||||
start_date="2025-02-01",
|
||||
end_date="2025-02-28",
|
||||
user_id="user-2",
|
||||
timezone=None,
|
||||
include_current_utc_day=False,
|
||||
user_api_key_dict=non_admin_key_dict,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Non-admin users can only view their own spend data" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -14646,6 +14646,218 @@ async def test_get_team_daily_activity_aggregated_rejects_bad_ranges(
|
|||
mock_aggregated.assert_not_called()
|
||||
|
||||
|
||||
def _key_search_team_setup(mock_db_client, user_id: str, team_id: str):
|
||||
mock_user_info = LiteLLM_UserTable(
|
||||
user_id=user_id,
|
||||
teams=[team_id],
|
||||
max_budget=1000.0,
|
||||
spend=0.0,
|
||||
user_email="test@example.com",
|
||||
user_role="internal_user",
|
||||
)
|
||||
mock_team = MagicMock(spec=LiteLLM_TeamTable)
|
||||
mock_team.team_id = team_id
|
||||
mock_team.team_alias = "Test Team"
|
||||
mock_team.members_with_roles = [Member(user_id=user_id, role="user")]
|
||||
mock_team.model_dump.return_value = {
|
||||
"team_id": team_id,
|
||||
"team_alias": "Test Team",
|
||||
"members_with_roles": [{"user_id": user_id, "role": "user"}],
|
||||
}
|
||||
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team])
|
||||
return mock_user_info
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_team_daily_activity_keys_scopes_where_before_take(mock_db_client):
|
||||
"""A member's search must put the team and own-key scoping inside the same
|
||||
Prisma where as the term, because `take` trims rows before Python sees them:
|
||||
scoped outside the where, the top-N slice could be spent entirely on keys
|
||||
the caller is not allowed to see."""
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
search_team_daily_activity_keys,
|
||||
)
|
||||
|
||||
user_id = "test_user_123"
|
||||
team_id = "test_team_456"
|
||||
user_api_key_dict = UserAPIKeyAuth(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
mock_user_info = _key_search_team_setup(mock_db_client, user_id, team_id)
|
||||
|
||||
user_key_1 = MagicMock()
|
||||
user_key_1.token = "user_key_1"
|
||||
matched = MagicMock()
|
||||
matched.token = "user_key_1"
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=[[user_key_1], [matched]])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_user_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_user_object:
|
||||
mock_get_user_object.return_value = mock_user_info
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity_aggregated",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_aggregated:
|
||||
mock_aggregated.return_value = MagicMock()
|
||||
|
||||
await search_team_daily_activity_keys(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
search="Needle",
|
||||
team_ids=team_id,
|
||||
start_date="2024-01-01",
|
||||
end_date="2024-01-31",
|
||||
exclude_team_ids=None,
|
||||
timezone=480,
|
||||
)
|
||||
|
||||
token_calls = mock_db_client.db.litellm_verificationtoken.find_many.call_args_list
|
||||
assert len(token_calls) == 2
|
||||
search_kwargs = token_calls[1][1]
|
||||
assert search_kwargs["where"] == {
|
||||
"team_id": {"in": (team_id,)},
|
||||
"token": {"in": ("user_key_1",)},
|
||||
"OR": (
|
||||
{"token": "Needle"},
|
||||
{"key_alias": {"contains": "Needle", "mode": "insensitive"}},
|
||||
{"user_id": {"contains": "Needle", "mode": "insensitive"}},
|
||||
),
|
||||
}
|
||||
assert search_kwargs["take"] == USAGE_TOP_API_KEYS_LIMIT
|
||||
assert search_kwargs["order"] == {"spend": "desc"}
|
||||
|
||||
call_kwargs = mock_aggregated.call_args[1]
|
||||
assert call_kwargs["api_key"] == ["user_key_1"]
|
||||
assert call_kwargs["entity_id"] == [team_id]
|
||||
assert call_kwargs["table_name"] == "litellm_dailyteamspend"
|
||||
assert call_kwargs["include_entity_breakdown"] is True
|
||||
assert call_kwargs["timezone_offset_minutes"] == 480
|
||||
assert call_kwargs["model"] is None
|
||||
assert call_kwargs["entity_metadata_field"] == {team_id: {"team_alias": "Test Team"}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_team_daily_activity_keys_admin_unscoped_where(mock_db_client):
|
||||
"""An admin's search has no caller scoping, so the where is the bare OR over
|
||||
token, key alias and user id; every matched hash is passed through to the
|
||||
aggregation."""
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
search_team_daily_activity_keys,
|
||||
)
|
||||
|
||||
match_1 = MagicMock()
|
||||
match_1.token = "h1"
|
||||
match_2 = MagicMock()
|
||||
match_2.token = "h2"
|
||||
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[match_1, match_2])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity_aggregated",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_aggregated:
|
||||
mock_aggregated.return_value = MagicMock()
|
||||
|
||||
await search_team_daily_activity_keys(
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
search="Needle",
|
||||
team_ids=None,
|
||||
start_date="2024-01-01",
|
||||
end_date="2024-01-31",
|
||||
exclude_team_ids=None,
|
||||
timezone=None,
|
||||
)
|
||||
|
||||
search_kwargs = mock_db_client.db.litellm_verificationtoken.find_many.call_args[1]
|
||||
assert search_kwargs["where"] == {
|
||||
"OR": (
|
||||
{"token": "Needle"},
|
||||
{"key_alias": {"contains": "Needle", "mode": "insensitive"}},
|
||||
{"user_id": {"contains": "Needle", "mode": "insensitive"}},
|
||||
)
|
||||
}
|
||||
assert search_kwargs["take"] == USAGE_TOP_API_KEYS_LIMIT
|
||||
assert mock_aggregated.call_args[1]["api_key"] == ["h1", "h2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_team_daily_activity_keys_no_match_returns_empty_without_aggregating(
|
||||
mock_db_client,
|
||||
):
|
||||
"""A term matching no key still owes the caller the standard metadata shape
|
||||
(api_key_limit, total_api_keys), and the aggregated query must not run."""
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
search_team_daily_activity_keys,
|
||||
)
|
||||
|
||||
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity_aggregated",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_aggregated:
|
||||
result = await search_team_daily_activity_keys(
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
search="Needle",
|
||||
team_ids=None,
|
||||
start_date="2024-01-01",
|
||||
end_date="2024-01-31",
|
||||
exclude_team_ids=None,
|
||||
timezone=None,
|
||||
)
|
||||
|
||||
assert result.results == []
|
||||
assert result.metadata.total_api_keys == 0
|
||||
assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT
|
||||
mock_aggregated.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_team_daily_activity_keys_excludes_teams_in_where(mock_db_client):
|
||||
"""The dashboard always sends exclude_team_ids=litellm-dashboard; if that
|
||||
filter stayed out of the where, matching keys in excluded teams could fill
|
||||
the take=N slice and push visible matches out."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
search_team_daily_activity_keys,
|
||||
)
|
||||
|
||||
matched = MagicMock()
|
||||
matched.token = "h1"
|
||||
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[matched])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity_aggregated",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_aggregated:
|
||||
mock_aggregated.return_value = MagicMock()
|
||||
|
||||
await search_team_daily_activity_keys(
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
search="Needle",
|
||||
team_ids=None,
|
||||
start_date="2024-01-01",
|
||||
end_date="2024-01-31",
|
||||
exclude_team_ids="litellm-dashboard",
|
||||
timezone=None,
|
||||
)
|
||||
|
||||
search_kwargs = mock_db_client.db.litellm_verificationtoken.find_many.call_args[1]
|
||||
assert search_kwargs["where"] == {
|
||||
"team_id": {"notIn": ("litellm-dashboard",)},
|
||||
"OR": (
|
||||
{"token": "Needle"},
|
||||
{"key_alias": {"contains": "Needle", "mode": "insensitive"}},
|
||||
{"user_id": {"contains": "Needle", "mode": "insensitive"}},
|
||||
),
|
||||
}
|
||||
assert mock_aggregated.call_args[1]["exclude_entity_ids"] == ["litellm-dashboard"]
|
||||
|
||||
|
||||
def _wire_new_team_prisma(mock_db_client):
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data
|
||||
mock_db_client.get_data = AsyncMock(return_value=None)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,8 @@ import pytest
|
|||
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.openai_files_endpoints.batch_file_validation import (
|
||||
BATCH_LINE_REQUIRED_KEYS,
|
||||
BATCH_LINE_SHAPE,
|
||||
PASSTHROUGH_BATCH_LINE_SHAPE,
|
||||
BatchFileEmpty,
|
||||
BatchFileInvalidJsonLine,
|
||||
BatchFileLineNotObject,
|
||||
|
|
@ -96,7 +97,7 @@ def test_non_object_line_rejected():
|
|||
assert check_batch_file_upload("batch.jsonl", content, None) == BatchFileLineNotObject(line_number=2)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("missing_key", BATCH_LINE_REQUIRED_KEYS)
|
||||
@pytest.mark.parametrize("missing_key", BATCH_LINE_SHAPE.required_keys)
|
||||
def test_missing_required_key_rejected(missing_key):
|
||||
import json
|
||||
|
||||
|
|
@ -174,3 +175,36 @@ def test_failures_map_to_openai_shaped_proxy_exceptions(failure, expected_code,
|
|||
assert exc_info.value.param == expected_param
|
||||
for fragment in expected_fragments:
|
||||
assert fragment in exc_info.value.message
|
||||
|
||||
|
||||
NATIVE_VERTEX_LINE = b'{"request": {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}}'
|
||||
|
||||
|
||||
def test_passthrough_keys_accept_native_vertex_rows():
|
||||
content = NATIVE_VERTEX_LINE + b"\n" + NATIVE_VERTEX_LINE + b"\n"
|
||||
assert check_batch_file_upload("batch.jsonl", content, None, PASSTHROUGH_BATCH_LINE_SHAPE) is None
|
||||
|
||||
|
||||
def test_passthrough_keys_reject_openai_rows():
|
||||
content = NATIVE_VERTEX_LINE + b"\n" + VALID_LINE + b"\n"
|
||||
assert check_batch_file_upload(
|
||||
"batch.jsonl", content, None, PASSTHROUGH_BATCH_LINE_SHAPE
|
||||
) == BatchFileMissingLineKey(line_number=2, key="request", line_shape=PASSTHROUGH_BATCH_LINE_SHAPE)
|
||||
|
||||
|
||||
def test_default_keys_still_reject_native_vertex_rows():
|
||||
assert check_batch_file_upload("batch.jsonl", NATIVE_VERTEX_LINE, None) == BatchFileMissingLineKey(
|
||||
line_number=1, key="custom_id"
|
||||
)
|
||||
|
||||
|
||||
def test_passthrough_missing_key_message_says_what_a_passthrough_upload_takes():
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
raise_batch_file_validation_failure(
|
||||
BatchFileMissingLineKey(line_number=3, key="request", line_shape=PASSTHROUGH_BATCH_LINE_SHAPE)
|
||||
)
|
||||
assert exc_info.value.param == "request"
|
||||
assert "line 3" in exc_info.value.message
|
||||
assert "passthrough upload takes native Vertex batch rows" in exc_info.value.message
|
||||
assert "with a request key." in exc_info.value.message
|
||||
assert "custom_id" not in exc_info.value.message
|
||||
|
|
|
|||
|
|
@ -5669,3 +5669,241 @@ def test_model_routed_file_retrieve_allows_key_with_model_grant(mocker: MockerFi
|
|||
assert response.status_code == 200, response.text
|
||||
assert captured_kwargs["api_key"] == "mistral-key"
|
||||
assert captured_kwargs["custom_llm_provider"] == "mistral"
|
||||
|
||||
|
||||
NATIVE_VERTEX_BATCH_LINE = (
|
||||
b'{"request": {"contents": [{"role": "user", "parts": [{"text": "What is the tallest building?"}]}],'
|
||||
b' "tools": [{"googleSearch": {"excludeDomains": ["example.com"]}}]}}\n'
|
||||
)
|
||||
|
||||
|
||||
def _passthrough_router() -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "vertex-batch",
|
||||
"litellm_params": {
|
||||
"model": "vertex_ai/gemini-2.5-flash",
|
||||
"vertex_project": "proj",
|
||||
"vertex_location": "us-central1",
|
||||
},
|
||||
"model_info": {"id": "vertex-batch-id"},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {"model": "openai/gpt-3.5-turbo", "api_key": "openai_api_key"},
|
||||
"model_info": {"id": "gpt-3.5-turbo-id"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _setup_passthrough_upload_endpoint(monkeypatch, llm_router: Router) -> list:
|
||||
"""Like _setup_batch_upload_endpoint, but reads the forwarded file bytes while the spool is open."""
|
||||
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
|
||||
|
||||
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
|
||||
|
||||
async def fake_route_create_file(**kwargs):
|
||||
upload_source = kwargs["_create_file_request"]["file"][1]
|
||||
upload_source.seek(0)
|
||||
forwarded_calls.append({**kwargs, "file_bytes": upload_source.read()})
|
||||
return OpenAIFileObject(
|
||||
id="dummy-id",
|
||||
object="file",
|
||||
bytes=0,
|
||||
created_at=1234567890,
|
||||
filename="batch.jsonl",
|
||||
purpose="batch",
|
||||
status="uploaded",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(fe, "route_create_file", fake_route_create_file)
|
||||
return forwarded_calls
|
||||
|
||||
|
||||
def _upload(content: bytes, form: dict):
|
||||
return client.post(
|
||||
"/v1/files",
|
||||
files={"file": ("batch.jsonl", content, "application/jsonl")},
|
||||
data=form,
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
|
||||
def test_create_file_passthrough_forwards_native_vertex_rows_untouched(monkeypatch):
|
||||
forwarded_calls = _setup_passthrough_upload_endpoint(monkeypatch, _passthrough_router())
|
||||
content = NATIVE_VERTEX_BATCH_LINE * 2
|
||||
|
||||
try:
|
||||
response = _upload(content, {"purpose": "batch", "target_model_names": "vertex-batch", "passthrough": "true"})
|
||||
finally:
|
||||
_teardown_batch_upload_endpoint()
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
(call,) = forwarded_calls
|
||||
assert call["_create_file_request"]["passthrough"] is True
|
||||
assert call["file_bytes"] == content
|
||||
assert call["target_model_names_list"] == ["vertex-batch"]
|
||||
|
||||
|
||||
def test_create_file_passthrough_rejects_rows_without_a_request(monkeypatch):
|
||||
forwarded_calls = _setup_passthrough_upload_endpoint(monkeypatch, _passthrough_router())
|
||||
|
||||
try:
|
||||
response = _upload(
|
||||
NATIVE_VERTEX_BATCH_LINE + VALID_BATCH_LINE,
|
||||
{"purpose": "batch", "target_model_names": "vertex-batch", "passthrough": "true"},
|
||||
)
|
||||
finally:
|
||||
_teardown_batch_upload_endpoint()
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
error = response.json()["error"]
|
||||
assert error["param"] == "request"
|
||||
assert "line 2" in error["message"]
|
||||
assert forwarded_calls == []
|
||||
|
||||
|
||||
def test_create_file_without_passthrough_still_rejects_native_vertex_rows(monkeypatch):
|
||||
forwarded_calls = _setup_passthrough_upload_endpoint(monkeypatch, _passthrough_router())
|
||||
|
||||
try:
|
||||
response = _upload(NATIVE_VERTEX_BATCH_LINE, {"purpose": "batch", "target_model_names": "vertex-batch"})
|
||||
finally:
|
||||
_teardown_batch_upload_endpoint()
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
assert response.json()["error"]["param"] == "custom_id"
|
||||
assert forwarded_calls == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"form, expected_param, expected_fragment",
|
||||
[
|
||||
({"purpose": "batch", "passthrough": "true"}, "target_model_names", "target_model_names"),
|
||||
(
|
||||
{"purpose": "batch", "target_model_names": "gpt-3.5-turbo", "passthrough": "true"},
|
||||
"target_model_names",
|
||||
"'gpt-3.5-turbo'",
|
||||
),
|
||||
(
|
||||
{"purpose": "batch", "target_model_names": "vertex-batch,gpt-3.5-turbo", "passthrough": "true"},
|
||||
"target_model_names",
|
||||
"'gpt-3.5-turbo'",
|
||||
),
|
||||
({"purpose": "user_data", "target_model_names": "vertex-batch", "passthrough": "true"}, "passthrough", "batch"),
|
||||
(
|
||||
{"purpose": "batch", "target_model_names": "vertex-batch", "passthrough": "true", "target_storage": "s3"},
|
||||
"target_storage",
|
||||
"'s3'",
|
||||
),
|
||||
(
|
||||
{"purpose": "batch", "model": "gpt-3.5-turbo", "passthrough": "true"},
|
||||
"model",
|
||||
"'gpt-3.5-turbo'",
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"no-model",
|
||||
"non-vertex-model",
|
||||
"mixed-models",
|
||||
"non-batch-purpose",
|
||||
"target-storage",
|
||||
"non-vertex-model-param",
|
||||
],
|
||||
)
|
||||
def test_create_file_passthrough_rejected_outside_a_vertex_batch(monkeypatch, form, expected_param, expected_fragment):
|
||||
forwarded_calls = _setup_passthrough_upload_endpoint(monkeypatch, _passthrough_router())
|
||||
|
||||
try:
|
||||
response = _upload(NATIVE_VERTEX_BATCH_LINE, form)
|
||||
finally:
|
||||
_teardown_batch_upload_endpoint()
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
error = response.json()["error"]
|
||||
assert error["type"] == "invalid_request_error"
|
||||
assert error["param"] == expected_param
|
||||
assert expected_fragment in error["message"]
|
||||
assert forwarded_calls == []
|
||||
|
||||
|
||||
def test_create_file_passthrough_accepts_the_model_param_as_the_deployment(monkeypatch):
|
||||
forwarded_calls = _setup_passthrough_upload_endpoint(monkeypatch, _passthrough_router())
|
||||
|
||||
try:
|
||||
response = _upload(
|
||||
NATIVE_VERTEX_BATCH_LINE, {"purpose": "batch", "model": "vertex-batch", "passthrough": "true"}
|
||||
)
|
||||
finally:
|
||||
_teardown_batch_upload_endpoint()
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
(call,) = forwarded_calls
|
||||
assert call["model"] == "vertex-batch"
|
||||
assert call["_create_file_request"]["passthrough"] is True
|
||||
|
||||
|
||||
def test_create_file_passthrough_rejects_a_model_group_with_a_non_vertex_deployment(monkeypatch):
|
||||
mixed_router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "vertex-batch",
|
||||
"litellm_params": {
|
||||
"model": "vertex_ai/gemini-2.5-flash",
|
||||
"vertex_project": "proj",
|
||||
"vertex_location": "us-central1",
|
||||
},
|
||||
"model_info": {"id": "vertex-batch-id"},
|
||||
},
|
||||
{
|
||||
"model_name": "vertex-batch",
|
||||
"litellm_params": {"model": "openai/gpt-4.1-mini", "api_key": "openai_api_key"},
|
||||
"model_info": {"id": "vertex-batch-openai-id"},
|
||||
},
|
||||
]
|
||||
)
|
||||
forwarded_calls = _setup_passthrough_upload_endpoint(monkeypatch, mixed_router)
|
||||
|
||||
try:
|
||||
response = _upload(
|
||||
NATIVE_VERTEX_BATCH_LINE, {"purpose": "batch", "target_model_names": "vertex-batch", "passthrough": "true"}
|
||||
)
|
||||
finally:
|
||||
_teardown_batch_upload_endpoint()
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
error = response.json()["error"]
|
||||
assert error["param"] == "target_model_names"
|
||||
assert "'vertex-batch'" in error["message"]
|
||||
assert forwarded_calls == []
|
||||
|
||||
|
||||
def test_create_file_passthrough_fails_closed_when_guardrails_would_scan_the_batch(monkeypatch):
|
||||
"""Batch guardrails read OpenAI-shaped rows, so a passthrough upload on a guardrailed
|
||||
key is refused rather than forwarded unscanned."""
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
class _Redactor(CustomGuardrail):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
return data
|
||||
|
||||
forwarded_calls = _setup_passthrough_upload_endpoint(monkeypatch, _passthrough_router())
|
||||
monkeypatch.setattr(litellm, "callbacks", [_Redactor(guardrail_name="g", default_on=True)])
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
try:
|
||||
response = _upload(
|
||||
NATIVE_VERTEX_BATCH_LINE, {"purpose": "batch", "target_model_names": "vertex-batch", "passthrough": "true"}
|
||||
)
|
||||
finally:
|
||||
_teardown_batch_upload_endpoint()
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
error = response.json()["error"]
|
||||
assert error["param"] == "passthrough"
|
||||
assert "guardrails" in error["message"]
|
||||
assert forwarded_calls == []
|
||||
|
|
|
|||
|
|
@ -5,3 +5,5 @@ Test what each side of the bridge does, not the rollout policy that picks a side
|
|||
Call each path directly with an explicit decision instead. The Python path is the implementation the dispatcher falls back to, e.g. `litellm.ocr.main.ocr`. The Rust path is the native binding, e.g. `NATIVE_OCR.load()` from `litellm/rust_bridge/ocr/entrypoints.py`, called with the request, args and kwargs that dispatch would hand it. When the native side reads a policy-derived setting such as `settings.secret_manager().native`, pin that field in the test instead of deriving it from the catalog. `ocr/test_secrets.py` shows the pattern
|
||||
|
||||
Rollout policy itself, meaning which rule matches and what `LITELLM_RUST` changes, belongs in `test_catalog.py`, `test_configuration.py` and `test_dispatch.py`, tested against rules the test builds rather than the shipped `catalog.RULES`
|
||||
|
||||
Before adding a test here, ask whether it checks something Rust cannot. A `route_host.py` module is the Python half of a native route: it projects Python-only state (the cost map, `litellm.*` settings, request kwargs) into the plain values the Rust side consumes, and maps native failures back onto public exceptions. Those projections are what belongs here, because a wrong key or an ignored provider prefix ships the wrong value to Rust and no Rust test sees it. `messages/test_route_host.py` shows the shape. Behavior that lives in Rust (a request transform given its inputs, header assembly, stream relay) is tested in the crate, and the route end to end is tested against a recording server in `tests/test_litellm_rust/`. A test that only re-checks a Python helper the route host happens to call is a duplicate of that helper's own test and should not be added
|
||||
|
|
|
|||
112
tests/test_litellm/rust_bridge/messages/test_route_host.py
Normal file
112
tests/test_litellm/rust_bridge/messages/test_route_host.py
Normal file
|
|
@ -0,0 +1,112 @@
|
|||
from dataclasses import astuple
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge.messages import route_host
|
||||
|
||||
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
|
||||
|
||||
|
||||
def _flag_model(monkeypatch: pytest.MonkeyPatch, name: str, **flags: bool) -> None:
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
name,
|
||||
{
|
||||
"litellm_provider": "anthropic",
|
||||
"mode": "chat",
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": 0,
|
||||
**flags,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_flag_model(
|
||||
monkeypatch,
|
||||
"claude-test-adaptive",
|
||||
supports_reasoning=True,
|
||||
supports_adaptive_thinking=True,
|
||||
supports_output_config=True,
|
||||
supports_xhigh_reasoning_effort=True,
|
||||
supports_sampling_params=False,
|
||||
)
|
||||
|
||||
capabilities: Final = route_host.model_capabilities("anthropic/claude-test-adaptive", None)
|
||||
|
||||
assert capabilities.supports_adaptive_thinking
|
||||
assert capabilities.supports_output_config
|
||||
assert not capabilities.supports_legacy_thinking
|
||||
assert not capabilities.supports_sampling_params
|
||||
assert capabilities.effort_tiers.xhigh
|
||||
assert not capabilities.effort_tiers.max
|
||||
|
||||
|
||||
def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> None:
|
||||
capabilities: Final = route_host.model_capabilities("anthropic/not-a-real-model", None)
|
||||
|
||||
assert capabilities.supports_sampling_params
|
||||
assert not capabilities.supports_reasoning
|
||||
assert not capabilities.supports_adaptive_thinking
|
||||
assert not any(astuple(capabilities.effort_tiers))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("global_flag", "kwargs", "expected"),
|
||||
[
|
||||
(False, {}, False),
|
||||
(True, {}, True),
|
||||
(False, {"drop_params": "true"}, True),
|
||||
(False, {"drop_params": "nonsense"}, False),
|
||||
(False, {"drop_params": False}, False),
|
||||
],
|
||||
)
|
||||
def test_drop_params_merges_the_global_flag_with_the_request(
|
||||
monkeypatch: pytest.MonkeyPatch, global_flag: bool, kwargs: dict[str, object], expected: bool
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "drop_params", global_flag)
|
||||
|
||||
assert route_host.shaping("anthropic/not-a-real-model", None, kwargs)["drop_params"] is expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("configured", "expected"),
|
||||
[
|
||||
(["tools[*].input_examples", 3, "metadata.user_id"], ("tools[*].input_examples", "metadata.user_id")),
|
||||
("tools", ()),
|
||||
(None, ()),
|
||||
],
|
||||
)
|
||||
def test_additional_drop_params_keep_only_string_paths(configured: object, expected: tuple[str, ...]) -> None:
|
||||
shaping: Final = route_host.shaping("anthropic/not-a-real-model", None, {"additional_drop_params": configured})
|
||||
|
||||
assert shaping["additional_drop_params"] == expected
|
||||
|
||||
|
||||
def test_native_request_rejections_map_to_the_public_400() -> None:
|
||||
from types import MappingProxyType
|
||||
|
||||
from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
|
||||
|
||||
request: Final = LiteLLMMessagesRequest(
|
||||
model="anthropic/claude-sonnet-5",
|
||||
messages=(),
|
||||
max_tokens=8,
|
||||
stream=None,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
kwargs=MappingProxyType({}),
|
||||
)
|
||||
rejected: Final = ValueError("claude-sonnet-5 does not support top_k=5")
|
||||
rejected.messages_request_error = True # pyright: ignore[reportAttributeAccessIssue] # marker the native host sets
|
||||
|
||||
mapped: Final = route_host.map_failure(rejected, request, "anthropic")
|
||||
|
||||
assert isinstance(mapped, litellm.BadRequestError)
|
||||
assert mapped.status_code == 400
|
||||
assert "does not support top_k=5" in mapped.message
|
||||
assert mapped.model == "claude-sonnet-5"
|
||||
assert not isinstance(route_host.map_failure(ValueError("plain"), request, "anthropic"), litellm.BadRequestError)
|
||||
111
tests/test_litellm/rust_bridge/messages/test_secrets.py
Normal file
111
tests/test_litellm/rust_bridge/messages/test_secrets.py
Normal file
|
|
@ -0,0 +1,111 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Mapping
|
||||
from dataclasses import replace
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol, cast # noqa: TID251 # narrows the parametrized path to its protocol
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_secret_manager import CustomSecretManager
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.handler import anthropic_messages
|
||||
from litellm.rust_bridge import settings
|
||||
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, NATIVE_MESSAGES, LiteLLMMessagesRequest
|
||||
from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem
|
||||
from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service
|
||||
from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_MODEL, MESSAGES_RESPONSE
|
||||
|
||||
pytest.importorskip("litellm.rust_bridge._native")
|
||||
|
||||
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
|
||||
|
||||
|
||||
class Messages(Protocol):
|
||||
def __call__(self) -> Awaitable[object]: ...
|
||||
|
||||
|
||||
class _ManagedSecrets(CustomSecretManager):
|
||||
def __init__(self, values: Mapping[str, str]) -> None:
|
||||
super().__init__(secret_manager_name="rust_bridge_messages_test")
|
||||
self.values: Final = values
|
||||
|
||||
async def async_read_secret(
|
||||
self,
|
||||
secret_name: str,
|
||||
optional_params: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> str | None:
|
||||
raise AssertionError("get_secret reads custom managers synchronously")
|
||||
|
||||
def sync_read_secret(
|
||||
self,
|
||||
secret_name: str,
|
||||
optional_params: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> str | None:
|
||||
return self.values.get(secret_name)
|
||||
|
||||
|
||||
def _native_request() -> LiteLLMMessagesRequest:
|
||||
return LiteLLMMessagesRequest(
|
||||
model=MESSAGES_MODEL,
|
||||
messages=MESSAGES,
|
||||
max_tokens=8,
|
||||
stream=None,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
kwargs=MappingProxyType({}),
|
||||
)
|
||||
|
||||
|
||||
def _public_kwargs() -> dict[str, object]:
|
||||
return {"model": MESSAGES_MODEL, "messages": [dict(message) for message in MESSAGES], "max_tokens": 8}
|
||||
|
||||
|
||||
async def _python_messages() -> object:
|
||||
return await anthropic_messages(**_public_kwargs())
|
||||
|
||||
|
||||
async def _rust_messages() -> object:
|
||||
route: Final = NATIVE_MESSAGES.load()
|
||||
assert route is not None
|
||||
return route(_native_request(), (), _public_kwargs())
|
||||
|
||||
|
||||
async def _rust_amessages() -> object:
|
||||
route: Final = NATIVE_AMESSAGES.load()
|
||||
assert route is not None
|
||||
return await route(_native_request(), (), _public_kwargs())
|
||||
|
||||
|
||||
@pytest.fixture(
|
||||
params=(_python_messages, _rust_messages, _rust_amessages), ids=("python-async", "rust-sync", "rust-async")
|
||||
)
|
||||
def messages(request: pytest.FixtureRequest) -> Messages:
|
||||
return cast(Messages, request.param)
|
||||
|
||||
|
||||
async def test_secret_manager_supplies_the_anthropic_key_and_base(
|
||||
monkeypatch: pytest.MonkeyPatch, messages: Messages
|
||||
) -> None:
|
||||
for name in ("ANTHROPIC_API_KEY", "ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
with recording_service() as server:
|
||||
server.default_response = ResponseSpec(body=MESSAGES_RESPONSE)
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"secret_manager_client",
|
||||
_ManagedSecrets({"ANTHROPIC_API_KEY": "vault-key", "ANTHROPIC_BASE_URL": server.base_url}),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.CUSTOM)
|
||||
monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(access_mode="read_only"))
|
||||
configured: Final = settings.secret_manager
|
||||
monkeypatch.setattr(settings, "secret_manager", lambda: replace(configured(), native=True))
|
||||
|
||||
await messages()
|
||||
|
||||
assert len(server.requests) == 1
|
||||
assert server.requests[0].headers["x-api-key"] == "vault-key"
|
||||
|
|
@ -46,9 +46,9 @@ class _FakeSpeech:
|
|||
)()
|
||||
|
||||
|
||||
class _FakeImages:
|
||||
class _FakeRawImages:
|
||||
async def generate(self, **kwargs: Any) -> Any:
|
||||
return type(
|
||||
parsed: Final = type(
|
||||
"_Images",
|
||||
(),
|
||||
{
|
||||
|
|
@ -58,6 +58,16 @@ class _FakeImages:
|
|||
}
|
||||
},
|
||||
)()
|
||||
return type(
|
||||
"_RawImages",
|
||||
(),
|
||||
{"parse": lambda self: parsed, "headers": httpx.Headers({"x-request-id": "req-image"})},
|
||||
)()
|
||||
|
||||
|
||||
class _FakeImages:
|
||||
def __init__(self) -> None:
|
||||
self.with_raw_response = _FakeRawImages()
|
||||
|
||||
|
||||
class _FakeModerations:
|
||||
|
|
|
|||
|
|
@ -528,6 +528,39 @@ async def test_async_router_acreate_file_with_jsonl():
|
|||
assert first_call_content == non_jsonl_content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_router_acreate_file_passthrough_keeps_the_file_and_forwards_the_flag():
|
||||
"""A passthrough batch upload must reach the provider byte for byte: the router
|
||||
neither rewrites body.model to the deployment model nor drops the flag."""
|
||||
from io import BytesIO
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
jsonl_content = b'{"custom_id": "r1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "vertex-batch"}}\n'
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "vertex-batch",
|
||||
"litellm_params": {"model": "vertex_ai/gemini-2.5-flash", "vertex_project": "p"},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file:
|
||||
await router.acreate_file(
|
||||
model="vertex-batch", purpose="batch", file=BytesIO(jsonl_content), passthrough=True
|
||||
)
|
||||
forwarded = mock_acreate_file.call_args.kwargs
|
||||
assert forwarded["passthrough"] is True
|
||||
forwarded["file"].seek(0)
|
||||
assert forwarded["file"].read() == jsonl_content
|
||||
|
||||
mock_acreate_file.reset_mock()
|
||||
await router.acreate_file(model="vertex-batch", purpose="batch", file=BytesIO(jsonl_content))
|
||||
rewritten = mock_acreate_file.call_args.kwargs["file"]
|
||||
rewritten.seek(0)
|
||||
assert b'"gemini-2.5-flash"' in rewritten.read()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_router_acreate_file_does_not_fall_back_across_model_groups():
|
||||
"""A file created for batches only exists under the credentials of the model group
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue