mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
chore: merge main into model info discovery branch
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
39286245b5
123 changed files with 9046 additions and 1152 deletions
101
litellm-rust/Cargo.lock
generated
101
litellm-rust/Cargo.lock
generated
|
|
@ -70,6 +70,12 @@ dependencies = [
|
|||
"rustversion",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "arcstr"
|
||||
version = "1.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "03918c3dbd7701a85c6b9887732e2921175f26c350b4563841d0958c21d57e6d"
|
||||
|
||||
[[package]]
|
||||
name = "async-compression"
|
||||
version = "0.4.46"
|
||||
|
|
@ -1837,6 +1843,12 @@ version = "0.2.186"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66"
|
||||
|
||||
[[package]]
|
||||
name = "linux-raw-sys"
|
||||
version = "0.12.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53"
|
||||
|
||||
[[package]]
|
||||
name = "litellm-auth"
|
||||
version = "0.1.0"
|
||||
|
|
@ -1915,6 +1927,17 @@ dependencies = [
|
|||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-cache-redis"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"litellm-cache",
|
||||
"redis",
|
||||
"redis-test",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-core"
|
||||
version = "0.1.0"
|
||||
|
|
@ -2140,6 +2163,16 @@ dependencies = [
|
|||
"minimal-lexical",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-bigint"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "93e7820bc0a80a0238e650327316f929ba18d5be054b647490a3a6a339f3e7c0"
|
||||
dependencies = [
|
||||
"num-integer",
|
||||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-conv"
|
||||
version = "0.2.2"
|
||||
|
|
@ -2656,6 +2689,36 @@ dependencies = [
|
|||
"crossbeam-utils",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "redis"
|
||||
version = "1.7.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2acbc41a996f7652b2ddd9dfd98cc4ff602cfd742ae35382f07f608405ab50ed"
|
||||
dependencies = [
|
||||
"arcstr",
|
||||
"combine",
|
||||
"itoa",
|
||||
"num-bigint",
|
||||
"percent-encoding",
|
||||
"ryu",
|
||||
"sha1_smol",
|
||||
"socket2 0.6.5",
|
||||
"url",
|
||||
"xxhash-rust",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "redis-test"
|
||||
version = "1.0.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "804d36862e4323b69f96440cbb13c9894fc90176abdeaf91264e21d5d77f6aca"
|
||||
dependencies = [
|
||||
"rand 0.9.5",
|
||||
"redis",
|
||||
"socket2 0.6.5",
|
||||
"tempfile",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "redox_syscall"
|
||||
version = "0.5.18"
|
||||
|
|
@ -2846,6 +2909,19 @@ dependencies = [
|
|||
"semver",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustix"
|
||||
version = "1.1.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "891efababe418670775f199f0d233d84843c227a0949a883ce15b37c78d6629d"
|
||||
dependencies = [
|
||||
"bitflags",
|
||||
"errno",
|
||||
"libc",
|
||||
"linux-raw-sys",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustls"
|
||||
version = "0.21.12"
|
||||
|
|
@ -3096,6 +3172,12 @@ dependencies = [
|
|||
"digest 0.10.7",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sha1_smol"
|
||||
version = "1.0.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "bbfa15b3dddfee50a0fff136974b3e1bde555604ba463834a7eb7deb6417705d"
|
||||
|
||||
[[package]]
|
||||
name = "sha2"
|
||||
version = "0.10.9"
|
||||
|
|
@ -3299,6 +3381,19 @@ version = "0.13.5"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca"
|
||||
|
||||
[[package]]
|
||||
name = "tempfile"
|
||||
version = "3.27.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd"
|
||||
dependencies = [
|
||||
"fastrand",
|
||||
"getrandom 0.4.3",
|
||||
"once_cell",
|
||||
"rustix",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "thiserror"
|
||||
version = "1.0.69"
|
||||
|
|
@ -4182,6 +4277,12 @@ version = "0.13.6"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "66fee0b777b0f5ac1c69bb06d361268faafa61cd4682ae064a171c16c433e9e4"
|
||||
|
||||
[[package]]
|
||||
name = "xxhash-rust"
|
||||
version = "0.8.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "aee1b19627c7c60102ab80d3a9cbe18de90bfe03bfa6c3715447681f0e8c8af6"
|
||||
|
||||
[[package]]
|
||||
name = "yoke"
|
||||
version = "0.8.3"
|
||||
|
|
|
|||
15
litellm-rust/crates/cache-redis/Cargo.toml
Normal file
15
litellm-rust/crates/cache-redis/Cargo.toml
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
[package]
|
||||
name = "litellm-cache-redis"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-cache.workspace = true
|
||||
redis = "1.7.0"
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
redis-test = "1.0.4"
|
||||
315
litellm-rust/crates/cache-redis/src/cache.rs
Normal file
315
litellm-rust/crates/cache-redis/src/cache.rs
Normal file
|
|
@ -0,0 +1,315 @@
|
|||
use std::sync::{Arc, Mutex, MutexGuard};
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_cache::{
|
||||
BaseCache, CacheConnectionResult, CacheConnectionStatus, CacheEntry, CacheFuture, CacheKwargs,
|
||||
Error,
|
||||
};
|
||||
use redis::Commands;
|
||||
|
||||
const DEFAULT_TTL: Duration = Duration::from_secs(600);
|
||||
const KEY_PREFIX: &str = "litellm-cache:";
|
||||
|
||||
pub struct RedisCache<C = redis::Connection> {
|
||||
connection: Arc<Mutex<C>>,
|
||||
default_ttl: Duration,
|
||||
}
|
||||
|
||||
impl RedisCache<redis::Connection> {
|
||||
pub fn new(url: &str, default_ttl: Option<Duration>) -> Result<Self, Error> {
|
||||
let client = redis::Client::open(url).map_err(|_| Error::Unavailable)?;
|
||||
let connection = client.get_connection().map_err(|_| Error::Unavailable)?;
|
||||
Ok(Self::with_connection(connection, default_ttl))
|
||||
}
|
||||
}
|
||||
|
||||
impl<C> RedisCache<C>
|
||||
where
|
||||
C: redis::ConnectionLike + Send + 'static,
|
||||
{
|
||||
fn with_connection(connection: C, default_ttl: Option<Duration>) -> Self {
|
||||
Self {
|
||||
connection: Arc::new(Mutex::new(connection)),
|
||||
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
|
||||
}
|
||||
}
|
||||
|
||||
fn connection(&self) -> Result<MutexGuard<'_, C>, Error> {
|
||||
self.connection.lock().map_err(|_| Error::Unavailable)
|
||||
}
|
||||
|
||||
fn namespaced_key(key: &str) -> String {
|
||||
format!("{KEY_PREFIX}{key}")
|
||||
}
|
||||
|
||||
fn namespaced_pattern() -> &'static str {
|
||||
const PATTERN: &str = "litellm-cache:*";
|
||||
PATTERN
|
||||
}
|
||||
|
||||
fn encode(value: &CacheEntry) -> Result<Vec<u8>, Error> {
|
||||
serde_json::to_vec(value).map_err(|_| Error::InvalidEntry)
|
||||
}
|
||||
|
||||
fn decode(value: Vec<u8>) -> Result<CacheEntry, Error> {
|
||||
serde_json::from_slice(&value).map_err(|_| Error::InvalidEntry)
|
||||
}
|
||||
|
||||
fn ttl_seconds(ttl: Duration) -> u64 {
|
||||
ttl.as_secs()
|
||||
.saturating_add(u64::from(ttl.subsec_nanos() > 0))
|
||||
.max(1)
|
||||
}
|
||||
|
||||
fn run_blocking<T, F>(connection: Arc<Mutex<C>>, operation: F) -> CacheFuture<'static, T>
|
||||
where
|
||||
T: Send + 'static,
|
||||
F: FnOnce(&mut C) -> Result<T, Error> + Send + 'static,
|
||||
{
|
||||
Box::pin(async move {
|
||||
tokio::task::spawn_blocking(move || {
|
||||
let mut connection = connection.lock().map_err(|_| Error::Unavailable)?;
|
||||
operation(&mut connection)
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Unavailable)?
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl<C> BaseCache for RedisCache<C>
|
||||
where
|
||||
C: redis::ConnectionLike + Send + 'static,
|
||||
{
|
||||
type Value = CacheEntry;
|
||||
|
||||
fn default_ttl(&self) -> Duration {
|
||||
self.default_ttl
|
||||
}
|
||||
|
||||
fn set_cache(&self, key: &str, value: Self::Value, kwargs: CacheKwargs) -> Result<(), Error> {
|
||||
let payload = Self::encode(&value)?;
|
||||
let ttl = Self::ttl_seconds(self.get_ttl(&kwargs));
|
||||
self.connection()?
|
||||
.set_ex::<_, _, ()>(Self::namespaced_key(key), payload, ttl)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
}
|
||||
|
||||
fn get_cache(&self, key: &str, _: &CacheKwargs) -> Result<Option<Self::Value>, Error> {
|
||||
self.connection()?
|
||||
.get::<_, Option<Vec<u8>>>(Self::namespaced_key(key))
|
||||
.map_err(|_| Error::Unavailable)?
|
||||
.map(Self::decode)
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn delete_cache(&self, key: &str) -> Result<(), Error> {
|
||||
self.connection()?
|
||||
.del::<_, ()>(Self::namespaced_key(key))
|
||||
.map_err(|_| Error::Unavailable)
|
||||
}
|
||||
|
||||
fn flush_cache(&self) -> Result<(), Error> {
|
||||
let mut connection = self.connection()?;
|
||||
let keys = connection
|
||||
.scan_match(Self::namespaced_pattern())
|
||||
.map_err(|_| Error::Unavailable)?
|
||||
.collect::<redis::RedisResult<Vec<String>>>()
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
if keys.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
connection
|
||||
.del::<_, usize>(keys)
|
||||
.map(|_| ())
|
||||
.map_err(|_| Error::Unavailable)
|
||||
}
|
||||
|
||||
fn async_set_cache<'a>(
|
||||
&'a self,
|
||||
key: &'a str,
|
||||
value: Self::Value,
|
||||
kwargs: CacheKwargs,
|
||||
) -> CacheFuture<'a, ()> {
|
||||
let payload = Self::encode(&value);
|
||||
let key = Self::namespaced_key(key);
|
||||
let ttl = Self::ttl_seconds(self.get_ttl(&kwargs));
|
||||
Self::run_blocking(Arc::clone(&self.connection), move |connection| {
|
||||
connection
|
||||
.set_ex::<_, _, ()>(key, payload?, ttl)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
}
|
||||
|
||||
fn async_get_cache<'a>(
|
||||
&'a self,
|
||||
key: &'a str,
|
||||
_: &'a CacheKwargs,
|
||||
) -> CacheFuture<'a, Option<Self::Value>> {
|
||||
let key = Self::namespaced_key(key);
|
||||
Box::pin(async move {
|
||||
Self::run_blocking(Arc::clone(&self.connection), move |connection| {
|
||||
connection
|
||||
.get::<_, Option<Vec<u8>>>(key)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
.await?
|
||||
.map(Self::decode)
|
||||
.transpose()
|
||||
})
|
||||
}
|
||||
|
||||
fn async_set_cache_pipeline<'a>(
|
||||
&'a self,
|
||||
cache_list: Vec<(String, Self::Value)>,
|
||||
kwargs: CacheKwargs,
|
||||
) -> CacheFuture<'a, ()> {
|
||||
let entries = cache_list
|
||||
.into_iter()
|
||||
.map(|(key, value)| {
|
||||
Self::encode(&value).map(|payload| (Self::namespaced_key(&key), payload))
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>();
|
||||
let ttl = Self::ttl_seconds(self.get_ttl(&kwargs));
|
||||
Self::run_blocking(Arc::clone(&self.connection), move |connection| {
|
||||
for (key, payload) in entries? {
|
||||
connection
|
||||
.set_ex::<_, _, ()>(key, payload, ttl)
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
}
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
|
||||
fn async_delete_cache<'a>(&'a self, key: &'a str) -> CacheFuture<'a, ()> {
|
||||
let key = Self::namespaced_key(key);
|
||||
Self::run_blocking(Arc::clone(&self.connection), move |connection| {
|
||||
connection.del::<_, ()>(key).map_err(|_| Error::Unavailable)
|
||||
})
|
||||
}
|
||||
|
||||
fn disconnect(&self) -> CacheFuture<'_, ()> {
|
||||
Box::pin(async { Ok(()) })
|
||||
}
|
||||
|
||||
fn test_connection(&self) -> CacheFuture<'_, CacheConnectionResult> {
|
||||
Box::pin(async move {
|
||||
Self::run_blocking(Arc::clone(&self.connection), |connection| {
|
||||
redis::cmd("PING")
|
||||
.query::<String>(connection)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
.await?;
|
||||
Ok(CacheConnectionResult {
|
||||
status: CacheConnectionStatus::Success,
|
||||
message: "Redis cache connection test successful".into(),
|
||||
error: None,
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::RedisCache;
|
||||
use litellm_cache::{BaseCache, CacheEntry, CacheKwargs};
|
||||
use redis_test::{MockCmd, MockRedisConnection};
|
||||
use serde_json::json;
|
||||
use std::time::Duration;
|
||||
|
||||
fn entry() -> CacheEntry {
|
||||
CacheEntry {
|
||||
timestamp: 123.0,
|
||||
response: json!({"choices": [{"text": "cached"}]}),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cache_entries_round_trip_through_json() {
|
||||
let entry = entry();
|
||||
let encoded = RedisCache::<redis::Connection>::encode(&entry).unwrap();
|
||||
assert_eq!(
|
||||
RedisCache::<redis::Connection>::decode(encoded).unwrap(),
|
||||
entry
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_json_is_rejected() {
|
||||
assert!(RedisCache::<redis::Connection>::decode(b"not json".to_vec()).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ttl_seconds_rounds_up_and_keeps_expiration_positive() {
|
||||
assert_eq!(
|
||||
RedisCache::<redis::Connection>::ttl_seconds(Duration::ZERO),
|
||||
1
|
||||
);
|
||||
assert_eq!(
|
||||
RedisCache::<redis::Connection>::ttl_seconds(Duration::from_millis(1500)),
|
||||
2
|
||||
);
|
||||
assert_eq!(
|
||||
RedisCache::<redis::Connection>::ttl_seconds(Duration::from_secs(15)),
|
||||
15
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn redis_commands_round_trip_entries_and_delete_only_namespaced_keys() {
|
||||
let value = entry();
|
||||
let payload = RedisCache::<redis::Connection>::encode(&value).unwrap();
|
||||
let connection = MockRedisConnection::new([
|
||||
MockCmd::new(
|
||||
redis::cmd("SETEX")
|
||||
.arg("litellm-cache:key")
|
||||
.arg(600)
|
||||
.arg(payload.clone()),
|
||||
Ok("OK"),
|
||||
),
|
||||
MockCmd::new(redis::cmd("GET").arg("litellm-cache:key"), Ok(payload)),
|
||||
MockCmd::new(redis::cmd("DEL").arg("litellm-cache:key"), Ok(1u32)),
|
||||
])
|
||||
.assert_all_commands_consumed();
|
||||
let cache = RedisCache::with_connection(connection, None);
|
||||
|
||||
cache
|
||||
.set_cache("key", value.clone(), CacheKwargs::default())
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
cache.get_cache("key", &CacheKwargs::default()).unwrap(),
|
||||
Some(value)
|
||||
);
|
||||
cache.delete_cache("key").unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn flush_scans_and_deletes_only_cache_keys() {
|
||||
let connection = MockRedisConnection::new([
|
||||
MockCmd::new(
|
||||
redis::cmd("SCAN")
|
||||
.cursor_arg(0)
|
||||
.arg("MATCH")
|
||||
.arg("litellm-cache:*"),
|
||||
Ok(redis_test::redis_value!(["0", ["litellm-cache:key"]])),
|
||||
),
|
||||
MockCmd::new(redis::cmd("DEL").arg("litellm-cache:key"), Ok(1u32)),
|
||||
])
|
||||
.assert_all_commands_consumed();
|
||||
let cache = RedisCache::with_connection(connection, None);
|
||||
|
||||
cache.flush_cache().unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_connection_runs_ping_off_executor() {
|
||||
let connection = MockRedisConnection::new([MockCmd::new(redis::cmd("PING"), Ok("PONG"))])
|
||||
.assert_all_commands_consumed();
|
||||
let cache = RedisCache::with_connection(connection, None);
|
||||
|
||||
assert_eq!(
|
||||
cache.test_connection().await.unwrap().status,
|
||||
litellm_cache::CacheConnectionStatus::Success
|
||||
);
|
||||
}
|
||||
}
|
||||
3
litellm-rust/crates/cache-redis/src/lib.rs
Normal file
3
litellm-rust/crates/cache-redis/src/lib.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
mod cache;
|
||||
|
||||
pub use cache::RedisCache;
|
||||
6
litellm-rust/crates/cache-redis/tests/cache.rs
Normal file
6
litellm-rust/crates/cache-redis/tests/cache.rs
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
use litellm_cache_redis::RedisCache;
|
||||
|
||||
#[test]
|
||||
fn constructor_rejects_invalid_urls() {
|
||||
assert!(RedisCache::new("not a redis url", None).is_err());
|
||||
}
|
||||
|
|
@ -1818,6 +1818,9 @@ if TYPE_CHECKING:
|
|||
from .llms.azure.responses.o_series_transformation import (
|
||||
AzureOpenAIOSeriesResponsesAPIConfig as AzureOpenAIOSeriesResponsesAPIConfig,
|
||||
)
|
||||
from .llms.azure_ai.responses.transformation import (
|
||||
AzureAIResponsesAPIConfig as AzureAIResponsesAPIConfig,
|
||||
)
|
||||
from .llms.xai.responses.transformation import (
|
||||
XAIResponsesAPIConfig as XAIResponsesAPIConfig,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -234,6 +234,7 @@ LLM_CONFIG_NAMES: Final = (
|
|||
"OpenAIResponsesAPIConfig",
|
||||
"AzureOpenAIResponsesAPIConfig",
|
||||
"AzureOpenAIOSeriesResponsesAPIConfig",
|
||||
"AzureAIResponsesAPIConfig",
|
||||
"XAIResponsesAPIConfig",
|
||||
"LiteLLMProxyResponsesAPIConfig",
|
||||
"HostedVLLMResponsesAPIConfig",
|
||||
|
|
@ -946,6 +947,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
|
|||
".llms.azure.responses.o_series_transformation",
|
||||
"AzureOpenAIOSeriesResponsesAPIConfig",
|
||||
),
|
||||
"AzureAIResponsesAPIConfig": (
|
||||
".llms.azure_ai.responses.transformation",
|
||||
"AzureAIResponsesAPIConfig",
|
||||
),
|
||||
"XAIResponsesAPIConfig": (
|
||||
".llms.xai.responses.transformation",
|
||||
"XAIResponsesAPIConfig",
|
||||
|
|
|
|||
|
|
@ -346,13 +346,17 @@ class MCPClient:
|
|||
self.update_auth_value(auth_value)
|
||||
|
||||
async def discovery_auth_fingerprint(self) -> str:
|
||||
return self._hash_discovery_auth(await self.prepare_request_auth())
|
||||
|
||||
async def prepare_request_auth(self) -> httpx.Request:
|
||||
"""Preview the authenticated request without sending it, closing the auth flow afterwards."""
|
||||
request: Final = httpx.Request("POST", self.server_url or "http://localhost/", headers=self._get_auth_headers())
|
||||
if self._resolved_auth is None:
|
||||
return self._hash_discovery_auth(request)
|
||||
return request
|
||||
flow: Final = self._resolved_auth.async_auth_flow(request)
|
||||
try:
|
||||
authenticated: Final = await flow.__anext__()
|
||||
return self._hash_discovery_auth(authenticated)
|
||||
return authenticated
|
||||
finally:
|
||||
await flow.aclose()
|
||||
|
||||
|
|
|
|||
|
|
@ -148,6 +148,7 @@ class _ToolCallChunk(TypedDict):
|
|||
class _UsageBearingChunk(TypedDict, total=False):
|
||||
usage: Usage | None
|
||||
_hidden_params: Mapping[str, str]
|
||||
choices: ReadOnly[Sequence[StreamingChoices | Mapping[str, object]]]
|
||||
|
||||
|
||||
class _UsageSummary(TypedDict):
|
||||
|
|
@ -921,21 +922,22 @@ class ChunkProcessor:
|
|||
|
||||
prompt_tokens_details = attach_cache_creation_token_details(prompt_tokens_details, cache_creation_token_details)
|
||||
|
||||
completion_tokens = self._reset_anthropic_cursor_completion_tokens(
|
||||
recovered_completion_tokens: Final = self._reset_anthropic_cursor_completion_tokens(
|
||||
chunks=chunks,
|
||||
completion_tokens=completion_tokens,
|
||||
completion_usage_updates=completion_usage_updates,
|
||||
)
|
||||
cursor_was_reset: Final = recovered_completion_tokens != completion_tokens
|
||||
|
||||
return UsagePerChunk(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
completion_tokens=recovered_completion_tokens,
|
||||
cache_creation_input_tokens=cache_creation_input_tokens,
|
||||
cache_read_input_tokens=cache_read_input_tokens,
|
||||
server_tool_use=server_tool_use,
|
||||
web_search_requests=web_search_requests,
|
||||
google_maps_grounding_requests=google_maps_grounding_requests,
|
||||
completion_tokens_details=completion_tokens_details,
|
||||
completion_tokens_details=None if cursor_was_reset else completion_tokens_details,
|
||||
prompt_tokens_details=prompt_tokens_details,
|
||||
cost=cost,
|
||||
inference_geo=self._last_provider_pricing_field(chunks, "inference_geo"),
|
||||
|
|
@ -960,6 +962,30 @@ class ChunkProcessor:
|
|||
]
|
||||
return values[-1] if values else None
|
||||
|
||||
@staticmethod
|
||||
def _finish_reason_of_choice(choice: object) -> str | None:
|
||||
match choice:
|
||||
case StreamingChoices(finish_reason=reason) | Choices(finish_reason=reason):
|
||||
return reason
|
||||
case {"finish_reason": str() as reason}:
|
||||
return reason
|
||||
case _:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _chunk_choices(chunk: "_UsageBearingChunk | ModelResponse | ModelResponseStream") -> Sequence[object]:
|
||||
if isinstance(chunk, dict):
|
||||
return chunk.get("choices", ())
|
||||
return getattr(chunk, "choices", ())
|
||||
|
||||
@staticmethod
|
||||
def _saw_finish_reason(chunks: Sequence["_UsageBearingChunk | ModelResponse"]) -> bool:
|
||||
return any(
|
||||
ChunkProcessor._finish_reason_of_choice(choice) is not None
|
||||
for chunk in chunks
|
||||
for choice in ChunkProcessor._chunk_choices(chunk)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _reset_anthropic_cursor_completion_tokens(
|
||||
chunks: Sequence["_UsageBearingChunk | ModelResponse"],
|
||||
|
|
@ -970,18 +996,18 @@ class ChunkProcessor:
|
|||
|
||||
See the ``completion_usage_updates`` comment in
|
||||
``_calculate_usage_per_chunk``. The accumulated value is NOT a stale
|
||||
cursor when either it is > 1 (definitely not a placeholder) or we saw
|
||||
>= 2 completion-bearing usage events (positive evidence ``message_delta``
|
||||
arrived). Otherwise — the only completion update we ever saw was the
|
||||
Anthropic ``message_start`` cursor (=1) — reset to 0 so
|
||||
``calculate_usage()``'s ``or token_counter(text=...)`` fallback estimates
|
||||
from the actually-received completion text instead of trusting the
|
||||
placeholder. Gated on ``custom_llm_provider == "anthropic"`` so the
|
||||
heuristic (which encodes Anthropic's specific message_start SSE shape)
|
||||
does not silently affect other providers that may legitimately report
|
||||
``completion_tokens=1`` from a single usage event.
|
||||
cursor when we saw >= 2 completion-bearing usage events or any chunk
|
||||
carried a ``finish_reason`` (positive evidence ``message_delta``
|
||||
arrived). Otherwise the only completion update we ever saw was the
|
||||
Anthropic ``message_start`` cursor, a small placeholder whose magnitude
|
||||
varies per request (1 and 8 both observed live), so reset to 0 and let
|
||||
``calculate_usage()``'s ``or token_counter(...)`` fallback estimate from
|
||||
the actually-received text and reasoning instead. Gated on
|
||||
``custom_llm_provider == "anthropic"`` so the heuristic (which encodes
|
||||
Anthropic's specific message_start SSE shape) does not silently affect
|
||||
other providers that legitimately report usage from a single event.
|
||||
"""
|
||||
saw_non_cursor_completion: Final = completion_tokens > 1 or completion_usage_updates >= 2
|
||||
saw_non_cursor_completion: Final = completion_usage_updates >= 2 or ChunkProcessor._saw_finish_reason(chunks)
|
||||
if saw_non_cursor_completion:
|
||||
return completion_tokens
|
||||
|
||||
|
|
@ -995,7 +1021,7 @@ class ChunkProcessor:
|
|||
if isinstance(hp, dict):
|
||||
custom_llm_provider = hp.get("custom_llm_provider")
|
||||
|
||||
if custom_llm_provider == "anthropic" and completion_tokens == 1:
|
||||
if custom_llm_provider == "anthropic":
|
||||
return 0
|
||||
return completion_tokens
|
||||
|
||||
|
|
@ -1039,10 +1065,13 @@ class ChunkProcessor:
|
|||
returned_usage.prompt_tokens = 0
|
||||
returned_usage.completion_tokens = (
|
||||
completion_tokens
|
||||
or token_counter(
|
||||
model=model,
|
||||
text=completion_output,
|
||||
count_response_tokens=True, # count_response_tokens is a Flag to tell token counter this is a response, No need to add extra tokens we do for input messages
|
||||
or (
|
||||
token_counter(
|
||||
model=model,
|
||||
text=completion_output,
|
||||
count_response_tokens=True, # count_response_tokens is a Flag to tell token counter this is a response, No need to add extra tokens we do for input messages
|
||||
)
|
||||
+ (reasoning_tokens or 0)
|
||||
)
|
||||
)
|
||||
returned_usage.total_tokens = returned_usage.prompt_tokens + returned_usage.completion_tokens
|
||||
|
|
@ -1066,15 +1095,16 @@ class ChunkProcessor:
|
|||
returned_usage.completion_tokens_details = completion_tokens_details
|
||||
|
||||
if reasoning_tokens is not None:
|
||||
capped_reasoning_tokens: Final = min(max(0, reasoning_tokens), returned_usage.completion_tokens)
|
||||
if returned_usage.completion_tokens_details is None:
|
||||
returned_usage.completion_tokens_details = CompletionTokensDetailsWrapper(
|
||||
reasoning_tokens=reasoning_tokens
|
||||
reasoning_tokens=capped_reasoning_tokens,
|
||||
text_tokens=returned_usage.completion_tokens - capped_reasoning_tokens,
|
||||
)
|
||||
elif (
|
||||
returned_usage.completion_tokens_details is not None
|
||||
and returned_usage.completion_tokens_details.reasoning_tokens is None
|
||||
):
|
||||
capped_reasoning_tokens: Final = min(max(0, reasoning_tokens), returned_usage.completion_tokens)
|
||||
returned_usage.completion_tokens_details.reasoning_tokens = capped_reasoning_tokens
|
||||
if returned_usage.completion_tokens_details.text_tokens is None:
|
||||
returned_usage.completion_tokens_details.text_tokens = (
|
||||
|
|
|
|||
|
|
@ -49,12 +49,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
return BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params)
|
||||
|
||||
def get_stripped_model_name(self, model: str) -> str:
|
||||
# if "responses/" is in the model name, remove it
|
||||
if "responses/" in model:
|
||||
model = model.replace("responses/", "")
|
||||
if "o_series" in model:
|
||||
model = model.replace("o_series/", "")
|
||||
return model
|
||||
return model.replace("responses/", "").replace("o_series/", "").replace("azure_ai/", "")
|
||||
|
||||
def _handle_reasoning_item(self, item: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
AzureAIApiKeyHeader = Literal["Authorization", "api-key", "Api-Key", "Ocp-Apim-Subscription-Key"]
|
||||
AZURE_OPENAI_V1_HOST_SUFFIXES: Final = (".services.ai.azure.com", ".openai.azure.com")
|
||||
|
||||
|
||||
def is_foundry_model_inference_base(api_base: str) -> bool:
|
||||
|
|
@ -19,11 +20,13 @@ def is_foundry_model_inference_base(api_base: str) -> bool:
|
|||
return "/openai/deployments" not in parsed.path
|
||||
|
||||
|
||||
def api_key_header_for_base(api_base: str | None) -> AzureAIApiKeyHeader:
|
||||
def is_azure_openai_v1_host(api_base: str | None) -> bool:
|
||||
host: Final = urlparse(api_base).hostname if api_base else None
|
||||
if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
|
||||
return "api-key"
|
||||
return "Authorization"
|
||||
return host is not None and host.endswith(AZURE_OPENAI_V1_HOST_SUFFIXES)
|
||||
|
||||
|
||||
def api_key_header_for_base(api_base: str | None) -> AzureAIApiKeyHeader:
|
||||
return "api-key" if is_azure_openai_v1_host(api_base) else "Authorization"
|
||||
|
||||
|
||||
def get_azure_ai_entra_token(litellm_params: Mapping[str, object] | None = None) -> str | None:
|
||||
|
|
@ -70,6 +73,17 @@ def get_azure_ai_auth_headers(
|
|||
AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY: Final = "azure_model_router_selected_model"
|
||||
|
||||
|
||||
def azure_ai_supports_native_responses(model: str | None, api_base: str | None) -> bool:
|
||||
resolved_base: Final = AzureFoundryModelInfo.get_api_base(api_base)
|
||||
if resolved_base is not None and not is_azure_openai_v1_host(resolved_base):
|
||||
return False
|
||||
if model is None:
|
||||
return True
|
||||
if "claude" in model.lower():
|
||||
return False
|
||||
return AzureFoundryModelInfo.get_azure_ai_route(model) == "default"
|
||||
|
||||
|
||||
class AzureFoundryModelInfo(BaseLLMModelInfo):
|
||||
"""Model info for Azure AI / Azure Foundry models."""
|
||||
|
||||
|
|
|
|||
0
litellm/llms/azure_ai/responses/__init__.py
Normal file
0
litellm/llms/azure_ai/responses/__init__.py
Normal file
53
litellm/llms/azure_ai/responses/transformation.py
Normal file
53
litellm/llms/azure_ai/responses/transformation.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.azure.responses.transformation import AzureOpenAIResponsesAPIConfig
|
||||
from litellm.llms.azure_ai.common_utils import (
|
||||
AzureFoundryModelInfo,
|
||||
api_key_header_for_base,
|
||||
get_azure_ai_auth_headers,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
_PROJECT_PATH_PREFIX: Final = ("api", "projects")
|
||||
_RESPONSES_PATH: Final = ("openai", "v1", "responses")
|
||||
|
||||
|
||||
def _responses_url(api_base: str) -> str:
|
||||
base_url: Final = httpx.URL(api_base)
|
||||
segments: Final = tuple(segment for segment in base_url.path.split("/") if segment)
|
||||
project_root: Final = segments[:3] if segments[:2] == _PROJECT_PATH_PREFIX else ()
|
||||
return str(base_url.copy_with(path="/" + "/".join((*project_root, *_RESPONSES_PATH)), query=None))
|
||||
|
||||
|
||||
class AzureAIResponsesAPIConfig(AzureOpenAIResponsesAPIConfig):
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.AZURE_AI
|
||||
|
||||
def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict:
|
||||
params: Final = litellm_params or GenericLiteLLMParams()
|
||||
auth_headers: Final = get_azure_ai_auth_headers(
|
||||
api_key=AzureFoundryModelInfo.get_api_key(params.api_key),
|
||||
litellm_params=params.model_dump(),
|
||||
api_key_header=api_key_header_for_base(AzureFoundryModelInfo.get_api_base(params.api_base)),
|
||||
)
|
||||
return { # mutable-ok: the handler updates the returned headers in place per the dict contract
|
||||
**headers,
|
||||
**auth_headers,
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
return False
|
||||
|
||||
def get_complete_url(self, api_base: str | None, litellm_params: dict) -> str:
|
||||
resolved_base: Final = AzureFoundryModelInfo.get_api_base(api_base)
|
||||
if resolved_base is None:
|
||||
raise ValueError(
|
||||
"api_base is required for the Azure AI Foundry Responses API. "
|
||||
"Set the api_base parameter or the AZURE_AI_API_BASE environment variable."
|
||||
)
|
||||
return _responses_url(resolved_base)
|
||||
|
|
@ -12,6 +12,12 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
|||
|
||||
|
||||
class DashScopeChatConfig(OpenAIGPTConfig):
|
||||
def get_supported_openai_params(self, model: str) -> list[str]: # mutable-ok: base class contract returns a list
|
||||
return [ # mutable-ok: base class contract returns a list
|
||||
*super().get_supported_openai_params(model=model),
|
||||
"reasoning_effort",
|
||||
]
|
||||
|
||||
def remove_cache_control_flag_from_messages_and_tools(
|
||||
self,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -4,9 +4,9 @@ from typing import Final
|
|||
|
||||
import litellm
|
||||
from litellm.utils import (
|
||||
_is_explicitly_disabled_factory,
|
||||
_supports_factory,
|
||||
declared_value_factory,
|
||||
is_explicitly_disabled_factory,
|
||||
)
|
||||
|
||||
from .gpt_transformation import OpenAIGPTConfig
|
||||
|
|
@ -192,7 +192,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
|
|||
|
||||
Use this for opt-out checks where unknown models should be allowed through.
|
||||
"""
|
||||
return _is_explicitly_disabled_factory(
|
||||
return is_explicitly_disabled_factory(
|
||||
model=cls._model_map_lookup_name(model),
|
||||
custom_llm_provider=None,
|
||||
key=f"supports_{level}_reasoning_effort",
|
||||
|
|
|
|||
|
|
@ -79,6 +79,7 @@ from litellm.utils import (
|
|||
CustomStreamWrapper,
|
||||
ModelResponse,
|
||||
is_base64_encoded,
|
||||
is_explicitly_disabled_factory,
|
||||
supports_reasoning,
|
||||
)
|
||||
|
||||
|
|
@ -866,6 +867,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
else:
|
||||
raise _unsupported_reasoning_effort(reasoning_effort)
|
||||
|
||||
@staticmethod
|
||||
def _supports_minimal_thinking_level(model: str) -> bool:
|
||||
lowered: Final = model.lower()
|
||||
is_gemini3flash: Final = "gemini-3" in lowered and "flash" in lowered
|
||||
return is_gemini3flash and not is_explicitly_disabled_factory(
|
||||
model=model, custom_llm_provider=None, key="supports_minimal_reasoning_effort"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _map_reasoning_effort_to_thinking_level(
|
||||
reasoning_effort: str,
|
||||
|
|
@ -880,13 +889,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
Returns:
|
||||
GeminiThinkingConfig with thinkingLevel and includeThoughts
|
||||
"""
|
||||
# Check if this is gemini-3-flash which supports MINIMAL thinking level
|
||||
# Covers gemini-3-flash, gemini-3-flash-preview, gemini-3.1-flash, gemini-3.1-flash-lite-preview,
|
||||
# gemini-3.5-flash, and any future 3.x-flash variants.
|
||||
is_gemini3flash: Final = model and ("flash" in model.lower() and "gemini-3" in model.lower())
|
||||
supports_minimal: Final = bool(model) and VertexGeminiConfig._supports_minimal_thinking_level(model)
|
||||
is_gemini31pro: Final = model and ("gemini-3.1-pro-preview" in model.lower())
|
||||
if reasoning_effort == "minimal":
|
||||
if is_gemini3flash:
|
||||
if supports_minimal:
|
||||
return {"thinkingLevel": "minimal", "includeThoughts": True}
|
||||
else:
|
||||
return {"thinkingLevel": "low", "includeThoughts": True}
|
||||
|
|
@ -899,18 +906,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return {"thinkingLevel": "high", "includeThoughts": True}
|
||||
elif reasoning_effort == "high":
|
||||
return {"thinkingLevel": "high", "includeThoughts": True}
|
||||
elif reasoning_effort == "disable":
|
||||
# Gemini 3 cannot fully disable thinking, so we use "minimal" for gemini-3-flash-preview, "low" for others
|
||||
if is_gemini3flash:
|
||||
return {"thinkingLevel": "minimal", "includeThoughts": False}
|
||||
else:
|
||||
return {"thinkingLevel": "low", "includeThoughts": False}
|
||||
elif reasoning_effort == "none":
|
||||
# For gemini-3-flash-preview, use "minimal" instead of "low"
|
||||
if is_gemini3flash:
|
||||
return {"thinkingLevel": "minimal", "includeThoughts": False}
|
||||
else:
|
||||
return {"thinkingLevel": "low", "includeThoughts": False}
|
||||
elif reasoning_effort in ("disable", "none"):
|
||||
return {
|
||||
"thinkingLevel": "minimal" if supports_minimal else "low",
|
||||
"includeThoughts": False,
|
||||
}
|
||||
else:
|
||||
raise _unsupported_reasoning_effort(reasoning_effort)
|
||||
|
||||
|
|
@ -977,8 +977,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
params["includeThoughts"] = True
|
||||
# Follow provider defaults unless explicitly opted into legacy behavior.
|
||||
if litellm.enable_gemini_default_thinking_level_low is True:
|
||||
is_gemini3flash: Final = "gemini-3" in model.lower() and "flash" in model.lower()
|
||||
params["thinkingLevel"] = "minimal" if is_gemini3flash else "low"
|
||||
params["thinkingLevel"] = (
|
||||
"minimal" if VertexGeminiConfig._supports_minimal_thinking_level(model) else "low"
|
||||
)
|
||||
else:
|
||||
# Thinking disabled
|
||||
params["includeThoughts"] = False
|
||||
|
|
|
|||
|
|
@ -26105,6 +26105,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -26162,6 +26163,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -28111,6 +28113,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -28170,6 +28173,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -28592,6 +28596,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -28649,6 +28654,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
|
|||
|
|
@ -102,6 +102,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
|||
UpstreamCredentialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
prepare_mcp_client,
|
||||
raise_public,
|
||||
raise_token_exchange_challenge,
|
||||
raise_user_oauth_challenge,
|
||||
|
|
@ -2804,6 +2805,8 @@ class MCPServerManager:
|
|||
headers=headers,
|
||||
server_label=server.name or server.server_name or server.alias or server.server_id,
|
||||
relays_upstream_auth=server.is_client_forwarded_token,
|
||||
auth_type=server.auth_type,
|
||||
upstream_token_header=server.upstream_token_header,
|
||||
)
|
||||
tool_func.__name__ = prefixed_tool_name
|
||||
tool_func.__doc__ = description
|
||||
|
|
@ -4259,15 +4262,20 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
return MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
extra_headers=extra_headers,
|
||||
resolved_auth=resolved_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
return await prepare_mcp_client(
|
||||
resolved_server,
|
||||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
timeout=(
|
||||
resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
),
|
||||
extra_headers=extra_headers,
|
||||
resolved_auth=resolved_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
),
|
||||
)
|
||||
|
||||
# Create SigV4 auth if configured
|
||||
|
|
@ -4297,17 +4305,20 @@ class MCPServerManager:
|
|||
else AuthResolution.no_auth
|
||||
)
|
||||
record_auth_resolution(server.server_id, legacy_source)
|
||||
return MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
return await prepare_mcp_client(
|
||||
resolved_server,
|
||||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
),
|
||||
)
|
||||
|
||||
async def _get_tools_from_server(
|
||||
|
|
|
|||
|
|
@ -54,7 +54,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
from litellm.types.mcp import credential_redirect_hook, custom_credential_slot
|
||||
from litellm.types.mcp import MCPAuthType, credential_redirect_hook, custom_credential_slot
|
||||
|
||||
|
||||
class _OpenAPIJSONSchema(TypedDict, total=False):
|
||||
|
|
@ -471,6 +471,8 @@ def create_tool_function(
|
|||
headers: dict[str, str] | None = None,
|
||||
server_label: str | None = None,
|
||||
relays_upstream_auth: bool = False,
|
||||
auth_type: MCPAuthType = None,
|
||||
upstream_token_header: str | None = None,
|
||||
):
|
||||
"""Create a tool function for an OpenAPI operation.
|
||||
|
||||
|
|
@ -503,6 +505,18 @@ def create_tool_function(
|
|||
by using **kwargs instead of named parameters.
|
||||
"""
|
||||
effective_headers: Final = _merge_openapi_tool_request_headers(headers)
|
||||
if auth_type is not None:
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_public,
|
||||
validate_static_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok
|
||||
|
||||
match validate_static_credential(auth_type, effective_headers, upstream_token_header):
|
||||
case Error(error):
|
||||
raise_public(error)
|
||||
case Ok():
|
||||
pass
|
||||
|
||||
# Build URL from base_url and path
|
||||
url = base_url + path
|
||||
|
|
|
|||
|
|
@ -13,15 +13,17 @@ from __future__ import annotations
|
|||
|
||||
import base64
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Final, Literal, NoReturn
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import SecretStr
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.experimental_mcp_client.client import strip_auth_scheme, to_basic_credentials
|
||||
from litellm.experimental_mcp_client.client import MCPClient, strip_auth_scheme, to_basic_credentials
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import MCPServerURLCredentialsError
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
ApiKeyConfig,
|
||||
|
|
@ -39,7 +41,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
Subject,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth, MCPAuthType, MCPTransport
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -79,7 +81,7 @@ def to_server_spec(server: MCPServer) -> ServerSpec | None:
|
|||
|
||||
BYOK is the per-user source of the ``api_key`` mode; its scheme rides on ``auth_type`` just
|
||||
like a shared key, but the value is per-user and not migrated yet, so a BYOK server defers
|
||||
to v1 regardless of ``auth_type`` (this guard is the seam the BYOK arm replaces later).
|
||||
to v1 for its static schemes. Declared OBO always stays with the exchange arm.
|
||||
|
||||
Dispatches on the declared ``auth_type``. The match is exhaustive over ``MCPAuthType`` with
|
||||
an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is
|
||||
|
|
@ -90,8 +92,8 @@ def to_server_spec(server: MCPServer) -> ServerSpec | None:
|
|||
modes ``true_passthrough`` / ``oauth_delegate`` (``PassthroughConfig``); delegated/passthrough
|
||||
oauth2 and SigV4 return None and stay on v1.
|
||||
"""
|
||||
if server.is_byok:
|
||||
return None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type)
|
||||
if server.is_byok and server.auth_type != MCPAuth.oauth2_token_exchange:
|
||||
return None # per-user BYOK source not migrated yet -> defer to v1
|
||||
resource: Final = server.url or server.server_id
|
||||
auth_type: Final = server.auth_type
|
||||
match auth_type:
|
||||
|
|
@ -165,21 +167,9 @@ def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec:
|
|||
)
|
||||
|
||||
|
||||
def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None:
|
||||
"""Build a token_exchange (OBO) spec, or defer (None) when it is not OBO-configured.
|
||||
|
||||
An OBO server with ``client_id``/``client_secret`` is owned by the v2 arm even if the
|
||||
``token_exchange_endpoint``/``token_url`` is absent: a missing endpoint then fails closed (412) at
|
||||
the exchanger rather than silently deferring to v1 and connecting unauthenticated, since the
|
||||
gateway must not guess the IdP or fall back to a weaker source. Without client credentials there is
|
||||
nothing to own, so the server stays on v1 (parity-safe). ``profile`` selects the wire dialect
|
||||
(``rfc8693`` default, ``entra_obo`` for Microsoft Entra On-Behalf-Of); an unrecognized value
|
||||
normalizes to ``rfc8693`` so a bad config value cannot crash spec-building. ``audience`` is
|
||||
forwarded only when the operator set it; a missing one is omitted, not derived.
|
||||
"""
|
||||
def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec:
|
||||
"""Keep declared OBO owned by the resolver, including incomplete client configuration."""
|
||||
endpoint: Final = server.token_exchange_endpoint or server.effective_token_url
|
||||
if not server.client_id or not server.client_secret:
|
||||
return None
|
||||
profile: Final[Literal["rfc8693", "entra_obo"]] = (
|
||||
"entra_obo" if server.token_exchange_profile == "entra_obo" else "rfc8693"
|
||||
)
|
||||
|
|
@ -193,7 +183,7 @@ def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None:
|
|||
token_exchange_endpoint=endpoint,
|
||||
audience=server.audience,
|
||||
client_id=server.client_id,
|
||||
client_secret=SecretStr(server.client_secret),
|
||||
client_secret=SecretStr(server.client_secret) if server.client_secret else None,
|
||||
token_endpoint_auth_method=server.token_endpoint_auth_method,
|
||||
scopes=tuple(server.scopes or ()),
|
||||
),
|
||||
|
|
@ -397,3 +387,69 @@ def raise_token_exchange_challenge(
|
|||
detail="Unauthorized",
|
||||
headers={"WWW-Authenticate": www_authenticate},
|
||||
)
|
||||
|
||||
|
||||
_STATIC_MODES: Final = frozenset(
|
||||
(MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic, MCPAuth.token, MCPAuth.authorization)
|
||||
)
|
||||
|
||||
|
||||
def _usable_credential_value(auth_type: MCPAuthType, name: str, value: str) -> bool:
|
||||
if not value:
|
||||
return False
|
||||
if auth_type == MCPAuth.api_key and name != "authorization":
|
||||
return True
|
||||
if value.lower() in ("bearer", "basic", "token", "apikey"):
|
||||
return False
|
||||
if auth_type == MCPAuth.api_key:
|
||||
api_scheme: Final = value.split(None, 1)[0]
|
||||
if api_scheme.lower() in ("bearer", "token", "apikey"):
|
||||
api_credential: Final = strip_auth_scheme(value, api_scheme).strip()
|
||||
return api_credential.lower() != api_scheme.lower()
|
||||
if auth_type in (MCPAuth.bearer_token, MCPAuth.token):
|
||||
scheme: Final = "Bearer" if auth_type == MCPAuth.bearer_token else "token"
|
||||
credential: Final = strip_auth_scheme(value, scheme).strip()
|
||||
return bool(credential) and credential.lower() != scheme.lower()
|
||||
if auth_type == MCPAuth.basic:
|
||||
parts: Final = value.split(None, 1)
|
||||
if len(parts) != 2 or parts[0].lower() != "basic":
|
||||
return False
|
||||
try:
|
||||
decoded: Final = base64.b64decode(parts[1], validate=True).strip()
|
||||
return b":" in decoded
|
||||
except ValueError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def validate_static_credential(
|
||||
auth_type: MCPAuthType,
|
||||
headers: Mapping[str, str],
|
||||
upstream_token_header: str | None = None,
|
||||
) -> Result[None, CredError]:
|
||||
if auth_type not in _STATIC_MODES:
|
||||
return Ok(None)
|
||||
default_slot: Final = "X-API-Key" if auth_type == MCPAuth.api_key else "Authorization"
|
||||
slots: Final = frozenset(
|
||||
name.lower()
|
||||
for name in (
|
||||
upstream_token_header or default_slot,
|
||||
default_slot,
|
||||
"Authorization",
|
||||
)
|
||||
)
|
||||
values: Final = tuple((name.lower(), value.strip()) for name, value in headers.items() if name.lower() in slots)
|
||||
if any(_usable_credential_value(auth_type, name, value) for name, value in values):
|
||||
return Ok(None)
|
||||
return Error(CredError.of_misconfigured(f"{auth_type} requires a usable upstream credential"))
|
||||
|
||||
|
||||
async def prepare_mcp_client(server: MCPServer, client: MCPClient) -> MCPClient:
|
||||
if server.auth_type not in _STATIC_MODES or client.transport_type == MCPTransport.stdio:
|
||||
return client
|
||||
request: Final = await client.prepare_request_auth()
|
||||
match validate_static_credential(server.auth_type, request.headers, server.upstream_token_header):
|
||||
case Error(error):
|
||||
raise_public(error)
|
||||
case Ok():
|
||||
return client
|
||||
|
|
|
|||
|
|
@ -844,6 +844,9 @@ class LiteLLMRoutes(enum.Enum):
|
|||
)
|
||||
|
||||
self_managed_routes = [
|
||||
# update_team resolves proxy/org/team admin itself and filters team admins
|
||||
# through the team_admin_editable_team_fields setting
|
||||
"/team/update",
|
||||
"/team/member_add",
|
||||
"/team/member_delete",
|
||||
"/management/v1/teams/{team_id}/members/bulk_delete",
|
||||
|
|
@ -4467,6 +4470,29 @@ class TeamInfoMember(Member):
|
|||
user_alias: str | None = None
|
||||
|
||||
|
||||
class TeamEditUnrestricted(BaseModel):
|
||||
kind: Literal["unrestricted"] = "unrestricted"
|
||||
|
||||
|
||||
class TeamEditAsTeamAdmin(BaseModel):
|
||||
kind: Literal["team_admin"] = "team_admin"
|
||||
editable_fields: tuple[str, ...]
|
||||
|
||||
|
||||
class TeamEditAsTeamAdminDisabled(BaseModel):
|
||||
kind: Literal["team_admin_disabled"] = "team_admin_disabled"
|
||||
|
||||
|
||||
class TeamEditNone(BaseModel):
|
||||
kind: Literal["none"] = "none"
|
||||
|
||||
|
||||
TeamEditAccess = Annotated[
|
||||
TeamEditUnrestricted | TeamEditAsTeamAdmin | TeamEditAsTeamAdminDisabled | TeamEditNone,
|
||||
Field(discriminator="kind"),
|
||||
]
|
||||
|
||||
|
||||
class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable):
|
||||
members_with_roles: tuple[TeamInfoMember, ...] = ()
|
||||
team_member_budget_table: LiteLLM_BudgetTableFull | None = None
|
||||
|
|
@ -4479,6 +4505,7 @@ class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable):
|
|||
# None = no org or not a manager; [] or ["all-proxy-models"] = no ceiling.
|
||||
organization_models: list[str] | None = None
|
||||
model_max_budget_usage: Mapping[str, Mapping[str, object]] | None = None
|
||||
caller_edit_access: TeamEditAccess = Field(default_factory=TeamEditNone)
|
||||
|
||||
|
||||
class TeamInfoResponseObject(TypedDict):
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@ from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
|||
from fastapi.responses import JSONResponse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.anthropic_interface.exceptions import AnthropicErrorResponse, AnthropicExceptionMapping
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
from litellm.llms.anthropic.experimental_pass_through.context_management import (
|
||||
|
|
@ -22,13 +21,16 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
create_response,
|
||||
log_llm_api_exception,
|
||||
proxy_exception_from_http_exception,
|
||||
resolve_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
with_litellm_call_id,
|
||||
)
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
|
|
@ -218,10 +220,12 @@ async def anthropic_response(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=base_llm_response_processor.data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.anthropic_response(): Exception occured - %s", e)
|
||||
log_llm_api_exception(e, base_llm_response_processor.litellm_call_id)
|
||||
|
||||
if isinstance(e, ProxyException):
|
||||
return _anthropic_error_json_response(e, request)
|
||||
return _anthropic_error_json_response(
|
||||
with_litellm_call_id(e, base_llm_response_processor.litellm_call_id), request
|
||||
)
|
||||
|
||||
# Extract model_id from request metadata (same as success path)
|
||||
litellm_metadata: Final = data.get("litellm_metadata", {}) or {}
|
||||
|
|
@ -231,7 +235,7 @@ async def anthropic_response(
|
|||
# Get headers
|
||||
headers: Final = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_id=data.get("litellm_call_id", ""),
|
||||
call_id=base_llm_response_processor.litellm_call_id,
|
||||
model_id=model_id,
|
||||
version=version,
|
||||
response_cost=0,
|
||||
|
|
@ -288,6 +292,7 @@ async def count_tokens(
|
|||
"""
|
||||
from litellm.proxy.proxy_server import token_counter as internal_token_counter
|
||||
|
||||
litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
|
||||
try:
|
||||
request_data: Final = await _read_request_body(request=request)
|
||||
data: Final[dict] = {**request_data}
|
||||
|
|
@ -339,7 +344,7 @@ async def count_tokens(
|
|||
detail=detail,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("litellm.proxy.anthropic_endpoints.count_tokens(): Exception occurred - %s", e)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
raise HTTPException(status_code=500, detail={"error": f"Internal server error: {e}"})
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@
|
|||
import asyncio
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response
|
||||
|
|
@ -17,7 +18,11 @@ from litellm.batches.main import CancelBatchRequest, RetrieveBatchRequest
|
|||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.batches_endpoints.common_utils import validate_batch_list_limit
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
log_llm_api_exception,
|
||||
request_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.common_utils.callback_utils import sanitize_openai_provider_metadata
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.common_utils.openai_endpoint_utils import (
|
||||
|
|
@ -383,8 +388,9 @@ async def create_batch(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.create_batch(): Exception occured - %s", e)
|
||||
raise handle_exception_on_proxy(e)
|
||||
litellm_call_id: Final = request_litellm_call_id(data)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
raise handle_exception_on_proxy(e, litellm_call_id)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -674,8 +680,9 @@ async def retrieve_batch(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.retrieve_batch(): Exception occured - %s", e)
|
||||
raise handle_exception_on_proxy(e)
|
||||
litellm_call_id: Final = request_litellm_call_id(data)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
raise handle_exception_on_proxy(e, litellm_call_id)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -725,6 +732,7 @@ async def list_batches(
|
|||
)
|
||||
|
||||
verbose_proxy_logger.debug("GET /v1/batches after=%s limit=%s", after, limit)
|
||||
data: Mapping[str, object] = MappingProxyType({})
|
||||
try:
|
||||
if llm_router is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -854,10 +862,11 @@ async def list_batches(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=e,
|
||||
request_data={"after": after, "limit": limit},
|
||||
request_data={**data, "after": after, "limit": limit},
|
||||
)
|
||||
verbose_proxy_logger.error("litellm.proxy.proxy_server.retrieve_batch(): Exception occured - %s", e)
|
||||
raise handle_exception_on_proxy(e)
|
||||
litellm_call_id: Final = request_litellm_call_id(data)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
raise handle_exception_on_proxy(e, litellm_call_id)
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
@ -1079,8 +1088,9 @@ async def cancel_batch(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.create_batch(): Exception occured - %s", e)
|
||||
raise handle_exception_on_proxy(e)
|
||||
litellm_call_id: Final = request_litellm_call_id(data)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
raise handle_exception_on_proxy(e, litellm_call_id)
|
||||
|
||||
|
||||
######################################################################
|
||||
|
|
|
|||
|
|
@ -7,7 +7,18 @@ from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequen
|
|||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Protocol, TypeAlias, TypeVar, overload
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Final,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
Protocol,
|
||||
TypeAlias,
|
||||
TypeVar,
|
||||
overload,
|
||||
runtime_checkable,
|
||||
)
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
|
|
@ -34,7 +45,11 @@ from litellm.constants import (
|
|||
UNSAFE_PROXY_RESPONSE_HEADERS,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket, is_expected_client_error
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_or_create_metadata_bucket,
|
||||
independent_snapshot,
|
||||
is_expected_client_error,
|
||||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
|
||||
from litellm.litellm_core_utils.get_supported_openai_params import (
|
||||
get_supported_openai_params,
|
||||
|
|
@ -1452,7 +1467,19 @@ def _has_attribute_error_in_chain(exc: Exception) -> bool:
|
|||
_CLIENT_DISCONNECT_DETAIL: Final = "Client disconnected the request"
|
||||
|
||||
|
||||
def _log_llm_api_exception(e: Exception, litellm_call_id: str | None) -> None:
|
||||
@runtime_checkable
|
||||
class _CarriesLitellmCallId(Protocol):
|
||||
litellm_call_id: str | None
|
||||
|
||||
|
||||
def request_litellm_call_id(data: Mapping[str, object]) -> str | None:
|
||||
logging_obj: Final = data.get("litellm_logging_obj")
|
||||
logged_id: Final = logging_obj.litellm_call_id if isinstance(logging_obj, _CarriesLitellmCallId) else None
|
||||
call_id: Final = logged_id or data.get("litellm_call_id")
|
||||
return call_id if isinstance(call_id, str) else None
|
||||
|
||||
|
||||
def log_llm_api_exception(e: Exception, litellm_call_id: str | None) -> None:
|
||||
if getattr(e, "status_code", None) == 499 and getattr(e, "detail", None) == _CLIENT_DISCONNECT_DETAIL:
|
||||
verbose_proxy_logger.info(
|
||||
"litellm.proxy.proxy_server._handle_llm_api_exception(): client disconnected, "
|
||||
|
|
@ -1532,6 +1559,10 @@ class ProxyBaseLLMRequestProcessing:
|
|||
def __init__(self, data: dict):
|
||||
self.data = data
|
||||
|
||||
@property
|
||||
def litellm_call_id(self) -> str | None:
|
||||
return request_litellm_call_id(self.data)
|
||||
|
||||
@staticmethod
|
||||
def _merge_passthrough_streaming_headers(
|
||||
response_headers: httpx.Headers | dict | None,
|
||||
|
|
@ -2062,6 +2093,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
) -> tuple[dict, LiteLLMLoggingObj]:
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
configured_fallbacks: Final = (
|
||||
self._configured_fallbacks(llm_router=llm_router, user_api_key_dict=user_api_key_dict)
|
||||
if llm_router is not None and not self.data.get("disable_fallbacks")
|
||||
else None
|
||||
)
|
||||
pristine: Final = independent_snapshot(self.data) if configured_fallbacks else None
|
||||
|
||||
try:
|
||||
return await self.common_processing_pre_call_logic(
|
||||
request=request,
|
||||
|
|
@ -2080,14 +2118,19 @@ class ProxyBaseLLMRequestProcessing:
|
|||
llm_router=llm_router,
|
||||
)
|
||||
except ProxyRateLimitError as original_exc:
|
||||
original_model: Final = self.data.get("model")
|
||||
if not original_model or not llm_router or self.data.get("disable_fallbacks"):
|
||||
rate_limited_data: Final = self.data
|
||||
original_model: Final = rate_limited_data.get("model")
|
||||
if (
|
||||
pristine is None
|
||||
or not configured_fallbacks
|
||||
or rate_limited_data.get("disable_fallbacks")
|
||||
or not isinstance(original_model, str)
|
||||
):
|
||||
raise
|
||||
|
||||
fallback_models: Final = self._resolve_fallback_models(
|
||||
model=original_model,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
fallbacks=configured_fallbacks,
|
||||
)
|
||||
if not fallback_models:
|
||||
raise
|
||||
|
|
@ -2102,6 +2145,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
for fallback_model in fallback_models:
|
||||
if fallback_model == original_model:
|
||||
continue
|
||||
self.data = independent_snapshot(pristine)
|
||||
self.data["model"] = fallback_model
|
||||
try:
|
||||
return await self.common_processing_pre_call_logic(
|
||||
|
|
@ -2123,39 +2167,30 @@ class ProxyBaseLLMRequestProcessing:
|
|||
except ProxyRateLimitError:
|
||||
continue
|
||||
except BaseException:
|
||||
self.data["model"] = original_model
|
||||
self.data = rate_limited_data
|
||||
raise
|
||||
|
||||
self.data["model"] = original_model
|
||||
self.data = rate_limited_data
|
||||
raise original_exc
|
||||
|
||||
def _resolve_fallback_models(
|
||||
self,
|
||||
model: str,
|
||||
llm_router: Router,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> list | None:
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
|
||||
fallbacks = None
|
||||
|
||||
@staticmethod
|
||||
def _configured_fallbacks(llm_router: Router, user_api_key_dict: UserAPIKeyAuth) -> list | None:
|
||||
key_router_settings: Final = user_api_key_dict.router_settings
|
||||
if isinstance(key_router_settings, dict) and "fallbacks" in key_router_settings:
|
||||
fallbacks = key_router_settings["fallbacks"]
|
||||
key_fallbacks: Final = key_router_settings.get("fallbacks") if isinstance(key_router_settings, dict) else None
|
||||
fallbacks: Final = key_fallbacks if key_fallbacks is not None else llm_router.fallbacks
|
||||
return fallbacks if isinstance(fallbacks, list) and fallbacks else None
|
||||
|
||||
if fallbacks is None:
|
||||
fallbacks = llm_router.fallbacks
|
||||
|
||||
if not fallbacks:
|
||||
return None
|
||||
@staticmethod
|
||||
def _resolve_fallback_models(model: str, fallbacks: list) -> list | None:
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
|
||||
fallback_model_group, generic_fallback_idx = get_fallback_model_group(
|
||||
fallbacks=fallbacks,
|
||||
model_group=model,
|
||||
)
|
||||
if fallback_model_group is None and generic_fallback_idx is not None:
|
||||
fallback_model_group = fallbacks[generic_fallback_idx]["*"]
|
||||
return fallback_model_group
|
||||
if fallback_model_group is not None:
|
||||
return fallback_model_group
|
||||
return fallbacks[generic_fallback_idx]["*"] if generic_fallback_idx is not None else None
|
||||
|
||||
@staticmethod
|
||||
def _get_model_id_from_response(hidden_params: Mapping[str, object], data: Mapping[str, object]) -> str:
|
||||
|
|
@ -3429,11 +3464,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
version: str | None = None,
|
||||
):
|
||||
"""Raises ProxyException (OpenAI API compatible) if an exception is raised"""
|
||||
logging_obj: Final[LiteLLMLoggingObj | None] = self.data.get("litellm_logging_obj", None)
|
||||
_log_llm_api_exception(
|
||||
e,
|
||||
(logging_obj.litellm_call_id if logging_obj is not None else None) or self.data.get("litellm_call_id"),
|
||||
)
|
||||
log_llm_api_exception(e, self.litellm_call_id)
|
||||
# Allow callbacks to transform the error response
|
||||
transformed_exception: Final = await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -3463,9 +3494,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
custom_headers: Final = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_id=(
|
||||
_litellm_logging_obj.litellm_call_id if _litellm_logging_obj else self.data.get("litellm_call_id")
|
||||
),
|
||||
call_id=self.litellm_call_id,
|
||||
model_id=model_id,
|
||||
version=version,
|
||||
response_cost=0,
|
||||
|
|
|
|||
|
|
@ -9,6 +9,9 @@ from typing import Final
|
|||
from fastapi import status
|
||||
|
||||
from litellm.constants import STRINGIFIED_NONE
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
LITELLM_CALL_ID_HEADER: Final = "x-litellm-call-id"
|
||||
|
||||
_OPENAI_ERROR_TYPE_BY_STATUS: Final[Mapping[int, str]] = MappingProxyType(
|
||||
{
|
||||
|
|
@ -52,3 +55,23 @@ def openai_error_param(exc: object) -> str | None:
|
|||
serializes as JSON ``null``."""
|
||||
carried: Final = attribute_of(exc, "param")
|
||||
return carried if isinstance(carried, str) and carried != STRINGIFIED_NONE else None
|
||||
|
||||
|
||||
def litellm_call_id_headers(litellm_call_id: str | None) -> dict[str, str] | None: # mutable-ok: ProxyException.headers
|
||||
if litellm_call_id is None:
|
||||
return None
|
||||
return {LITELLM_CALL_ID_HEADER: litellm_call_id} # mutable-ok: ProxyException mutates its headers dict
|
||||
|
||||
|
||||
def with_litellm_call_id(exc: ProxyException, litellm_call_id: str | None) -> ProxyException:
|
||||
"""The same error object, answering with ``x-litellm-call-id`` when it was raised without one."""
|
||||
if litellm_call_id is not None:
|
||||
exc.headers.setdefault(LITELLM_CALL_ID_HEADER, litellm_call_id)
|
||||
return exc
|
||||
|
||||
|
||||
def headers_with_litellm_call_id(headers: Mapping[str, str] | None, litellm_call_id: str) -> Mapping[str, str]:
|
||||
"""``headers`` plus ``x-litellm-call-id``, keeping the value they already carry under that name."""
|
||||
if headers is None:
|
||||
return MappingProxyType({LITELLM_CALL_ID_HEADER: litellm_call_id})
|
||||
return MappingProxyType({LITELLM_CALL_ID_HEADER: litellm_call_id, **headers})
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import asyncio
|
||||
import io
|
||||
import traceback
|
||||
from collections.abc import Sequence
|
||||
from typing import Final, get_type_hints
|
||||
|
||||
|
|
@ -9,19 +8,23 @@ from fastapi import APIRouter, Depends, File, HTTPException, Request, Response,
|
|||
from fastapi.responses import ORJSONResponse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_str_from_messages,
|
||||
)
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
log_llm_api_exception,
|
||||
resolve_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
coerce_numeric_form_fields,
|
||||
numeric_form_fields,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
litellm_call_id_headers,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
|
|
@ -91,11 +94,12 @@ async def image_generation(
|
|||
version,
|
||||
)
|
||||
|
||||
data = {}
|
||||
litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
|
||||
data = {"litellm_call_id": litellm_call_id}
|
||||
try:
|
||||
# Use orjson to parse JSON data, orjson speeds up requests significantly
|
||||
body: Final = await request.body()
|
||||
data = orjson.loads(body)
|
||||
data = orjson.loads(body) | data
|
||||
|
||||
# Include original request and headers in the data
|
||||
data = await add_litellm_data_to_request(
|
||||
|
|
@ -153,9 +157,7 @@ async def image_generation(
|
|||
response = await llm_call
|
||||
|
||||
### ALERTING ###
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(litellm_call_id=data.get("litellm_call_id", ""), status="success")
|
||||
)
|
||||
asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
|
||||
|
||||
### CALL HOOKS ### - modify outgoing data (guardrails, otel, etc.)
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
|
|
@ -168,7 +170,7 @@ async def image_generation(
|
|||
cache_key: Final = hidden_params.get("cache_key", None) or ""
|
||||
api_base: Final = hidden_params.get("api_base", None) or ""
|
||||
response_cost: Final = hidden_params.get("response_cost", None) or ""
|
||||
litellm_call_id: Final = hidden_params.get("litellm_call_id", None) or ""
|
||||
response_call_id: Final = hidden_params.get("litellm_call_id", None) or ""
|
||||
|
||||
fastapi_response.headers.update(
|
||||
ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
|
|
@ -179,7 +181,7 @@ async def image_generation(
|
|||
version=version,
|
||||
response_cost=response_cost,
|
||||
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
|
||||
call_id=litellm_call_id,
|
||||
call_id=response_call_id,
|
||||
request_data=data,
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
|
|
@ -200,13 +202,13 @@ async def image_generation(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.error("litellm.proxy.proxy_server.image_generation(): Exception occured - %s", e)
|
||||
verbose_proxy_logger.debug(traceback.format_exc())
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
|
|
@ -215,6 +217,7 @@ async def image_generation(
|
|||
message=getattr(e, "message", error_msg),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
openai_code=getattr(e, "code", None),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,191 @@
|
|||
"""Proxy-wide allow-list of team-settings fields a team admin may change on /team/update."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.models.team import LiteLLM_TeamTable
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
UpdateTeamRequest,
|
||||
)
|
||||
|
||||
TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING: Final = "team_admin_editable_team_fields"
|
||||
|
||||
# TODO(LIT-5722): add the remaining team settings one per PR, each with its value-diff tests and dashboard field
|
||||
SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS: Final[frozenset[str]] = frozenset({"tpm_limit"})
|
||||
|
||||
_FIELD_LIST: Final = TypeAdapter(list[str])
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
|
||||
_EMPTY: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
_METADATA_FOLDED_FIELDS: Final[frozenset[str]] = frozenset(
|
||||
(*LiteLLM_ManagementEndpoint_MetadataFields, *LiteLLM_ManagementEndpoint_MetadataFields_Premium)
|
||||
)
|
||||
_SYSTEM_MANAGED_METADATA_KEYS: Final[frozenset[str]] = frozenset({"team_member_budget_id"})
|
||||
_NOT_COLUMNS: Final[frozenset[str]] = frozenset({"team_id", "metadata"})
|
||||
_SETTINGS_LOCATION: Final = "Settings > UI > Team admin editable fields"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TeamAdminEditAllowed:
|
||||
request: UpdateTeamRequest
|
||||
kind: Literal["allowed"] = "allowed"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TeamAdminEditingDisabled:
|
||||
kind: Literal["disabled"] = "disabled"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TeamAdminFieldNotPermitted:
|
||||
field: str
|
||||
kind: Literal["field_not_permitted"] = "field_not_permitted"
|
||||
|
||||
|
||||
TeamAdminEditVerdict: TypeAlias = TeamAdminEditAllowed | TeamAdminEditingDisabled | TeamAdminFieldNotPermitted
|
||||
|
||||
|
||||
def resolve_team_admin_editable_fields(
|
||||
general_settings: Mapping[str, object],
|
||||
supported: frozenset[str],
|
||||
) -> frozenset[str]:
|
||||
raw: Final = general_settings.get(TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING)
|
||||
if raw is None:
|
||||
return frozenset()
|
||||
try:
|
||||
configured: Final = frozenset(_FIELD_LIST.validate_python(raw))
|
||||
except ValidationError:
|
||||
verbose_proxy_logger.warning(
|
||||
"%s must be a list of field names; ignoring %r", TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING, raw
|
||||
)
|
||||
return frozenset()
|
||||
unsupported: Final = configured - supported
|
||||
if unsupported:
|
||||
verbose_proxy_logger.warning(
|
||||
"%s ignores unsupported field(s) %s; supported: %s",
|
||||
TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING,
|
||||
sorted(unsupported),
|
||||
sorted(supported),
|
||||
)
|
||||
return configured & supported
|
||||
|
||||
|
||||
def _as_object(value: object) -> Mapping[str, object]:
|
||||
try:
|
||||
return _JSON_OBJECT.validate_json(value) if isinstance(value, str) else _JSON_OBJECT.validate_python(value)
|
||||
except ValidationError:
|
||||
return _EMPTY
|
||||
|
||||
|
||||
def _stored_metadata(existing: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return _as_object(existing.get("metadata"))
|
||||
|
||||
|
||||
def _submitted_metadata(
|
||||
data: UpdateTeamRequest, submitted: Mapping[str, object], existing: Mapping[str, object]
|
||||
) -> Mapping[str, object]:
|
||||
"""Metadata as it would be stored: the caller's dict (or the stored one) with top-level folded fields laid over."""
|
||||
base: Final = (
|
||||
_as_object(submitted.get("metadata")) if "metadata" in data.model_fields_set else _stored_metadata(existing)
|
||||
)
|
||||
folded: Final = data.model_fields_set & _METADATA_FOLDED_FIELDS
|
||||
return MappingProxyType({key: submitted[key] if key in folded else base[key] for key in base.keys() | folded})
|
||||
|
||||
|
||||
def _metadata_changes(
|
||||
data: UpdateTeamRequest, submitted: Mapping[str, object], existing: Mapping[str, object]
|
||||
) -> frozenset[str]:
|
||||
merged: Final = _submitted_metadata(data, submitted, existing)
|
||||
stored: Final = _stored_metadata(existing)
|
||||
return frozenset(
|
||||
key if key in _METADATA_FOLDED_FIELDS else "metadata"
|
||||
for key in (merged.keys() | stored.keys()) - _SYSTEM_MANAGED_METADATA_KEYS
|
||||
if merged.get(key) != stored.get(key)
|
||||
)
|
||||
|
||||
|
||||
def _stored_model_aliases(existing_row: LiteLLM_TeamTable) -> Mapping[str, object]:
|
||||
table: Final = existing_row.litellm_model_table
|
||||
return _as_object(_JSON_OBJECT.validate_json(table.model_dump_json()).get("model_aliases")) if table else _EMPTY
|
||||
|
||||
|
||||
def _column_changed(
|
||||
field: str, submitted: Mapping[str, object], existing: Mapping[str, object], existing_row: LiteLLM_TeamTable
|
||||
) -> bool:
|
||||
if field == "model_aliases":
|
||||
return _as_object(submitted.get(field)) != _stored_model_aliases(existing_row)
|
||||
if field in LiteLLM_TeamTable.model_fields:
|
||||
return submitted.get(field) != existing.get(field)
|
||||
return True
|
||||
|
||||
|
||||
def changed_team_fields(data: UpdateTeamRequest, existing_row: LiteLLM_TeamTable) -> frozenset[str]:
|
||||
"""Logical field names whose stored value the request would change.
|
||||
|
||||
Request and stored row are compared as JSON values so both sides share one representation. Fields the
|
||||
server folds into metadata are attributed to their own name whether they arrive top-level or inside
|
||||
``metadata``; anything else in ``metadata`` is attributed to ``metadata``. Fields with no stored
|
||||
counterpart on the team row count as changed whenever they are sent.
|
||||
"""
|
||||
submitted: Final = _JSON_OBJECT.validate_json(data.model_dump_json(exclude_unset=True))
|
||||
existing: Final = _JSON_OBJECT.validate_json(existing_row.model_dump_json())
|
||||
column_fields: Final = frozenset(data.model_fields_set) - _NOT_COLUMNS - _METADATA_FOLDED_FIELDS
|
||||
column_changes: Final = frozenset(
|
||||
field for field in column_fields if _column_changed(field, submitted, existing, existing_row)
|
||||
)
|
||||
return column_changes | _metadata_changes(data, submitted, existing)
|
||||
|
||||
|
||||
def _only_changes(data: UpdateTeamRequest, changed: frozenset[str]) -> UpdateTeamRequest:
|
||||
"""The request without the values it resends unchanged, which would otherwise still trigger derived writes
|
||||
such as a resent budget_duration pushing budget_reset_at back."""
|
||||
sent: Final = frozenset(data.model_fields_set)
|
||||
via_metadata: Final = frozenset({"metadata"}) if changed - sent else frozenset()
|
||||
kept: Final = frozenset({"team_id"}) | (changed & sent) | via_metadata
|
||||
return UpdateTeamRequest.model_validate(data.model_dump(include=MappingProxyType({field: True for field in kept})))
|
||||
|
||||
|
||||
def team_admin_edit_verdict(
|
||||
data: UpdateTeamRequest,
|
||||
existing: LiteLLM_TeamTable,
|
||||
permitted: frozenset[str],
|
||||
) -> TeamAdminEditVerdict:
|
||||
if not permitted:
|
||||
return TeamAdminEditingDisabled()
|
||||
changed: Final = changed_team_fields(data, existing)
|
||||
blocked: Final = sorted(changed - permitted)
|
||||
if blocked:
|
||||
return TeamAdminFieldNotPermitted(field=blocked[0])
|
||||
return TeamAdminEditAllowed(request=_only_changes(data, changed))
|
||||
|
||||
|
||||
def team_admin_request_or_raise(verdict: TeamAdminEditVerdict) -> UpdateTeamRequest:
|
||||
match verdict:
|
||||
case TeamAdminEditAllowed(request=request):
|
||||
return request
|
||||
case TeamAdminEditingDisabled():
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=(
|
||||
"Team admins on this proxy cannot edit team settings. "
|
||||
f"Ask a proxy admin to enable fields under {_SETTINGS_LOCATION}."
|
||||
),
|
||||
)
|
||||
case TeamAdminFieldNotPermitted(field=field):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=(
|
||||
f"Team admins on this proxy do not have permission to update '{field}'. "
|
||||
f"Ask a proxy admin to add it under {_SETTINGS_LOCATION}."
|
||||
),
|
||||
)
|
||||
case _:
|
||||
assert_never(verdict)
|
||||
|
|
@ -18,12 +18,23 @@ from collections.abc import Iterable, Mapping, Sequence
|
|||
from collections.abc import Set as AbstractSet
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Annotated, Final, NamedTuple, NoReturn, Protocol, TypeVar, cast
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Annotated,
|
||||
Final,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
NoReturn,
|
||||
Protocol,
|
||||
TypeAlias,
|
||||
TypeVar,
|
||||
cast,
|
||||
)
|
||||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||
from pydantic import BaseModel, JsonValue, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict, assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -62,6 +73,11 @@ from litellm.proxy._types import (
|
|||
SpecialProxyStrings,
|
||||
TeamAccessGroupModelGrant,
|
||||
TeamAddMemberResponse,
|
||||
TeamEditAccess,
|
||||
TeamEditAsTeamAdmin,
|
||||
TeamEditAsTeamAdminDisabled,
|
||||
TeamEditNone,
|
||||
TeamEditUnrestricted,
|
||||
TeamInfoMember,
|
||||
TeamInfoResponseObject,
|
||||
TeamInfoResponseObjectTeamTable,
|
||||
|
|
@ -122,6 +138,12 @@ from litellm.proxy.management_endpoints.router_weights import validate_router_se
|
|||
from litellm.proxy.management_endpoints.tag_management_endpoints import (
|
||||
get_daily_activity,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_admin_field_permissions import (
|
||||
SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS,
|
||||
resolve_team_admin_editable_fields,
|
||||
team_admin_edit_verdict,
|
||||
team_admin_request_or_raise,
|
||||
)
|
||||
from litellm.proxy.management_helpers.access_group_team_sync import (
|
||||
TEAM_ADVISORY_LOCK_SQL,
|
||||
AccessGroupSyncTx,
|
||||
|
|
@ -439,32 +461,70 @@ async def _refresh_cached_team(
|
|||
)
|
||||
|
||||
|
||||
async def _can_manage_team(
|
||||
TeamAccessRole: TypeAlias = Literal["proxy_admin", "org_admin", "team_admin"]
|
||||
|
||||
|
||||
def _raise_team_access_denied() -> NoReturn:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="You do not have access to this team",
|
||||
)
|
||||
|
||||
|
||||
async def _resolve_team_access(
|
||||
team_obj: LiteLLM_TeamTable,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> bool:
|
||||
"""True for a proxy admin, an admin of this team, or an org admin for the team's organization."""
|
||||
) -> TeamAccessRole | None:
|
||||
"""Strongest role the caller holds over ``team_obj``, or None when they hold none.
|
||||
|
||||
Org admin outranks team admin so a caller holding both keeps unrestricted edits.
|
||||
"""
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
return True
|
||||
return "proxy_admin"
|
||||
|
||||
if await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team_obj):
|
||||
return "org_admin"
|
||||
|
||||
if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj):
|
||||
return True
|
||||
return "team_admin"
|
||||
|
||||
return await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team_obj)
|
||||
return None
|
||||
|
||||
|
||||
async def _verify_team_access(
|
||||
team_obj: LiteLLM_TeamTable,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
"""Raise HTTPException(403) unless the caller can manage the given team."""
|
||||
if await _can_manage_team(team_obj=team_obj, user_api_key_dict=user_api_key_dict):
|
||||
return
|
||||
"""Raise 403 unless the caller is a proxy admin, an org admin for the team's org, or a team admin."""
|
||||
if await _resolve_team_access(team_obj=team_obj, user_api_key_dict=user_api_key_dict) is None:
|
||||
_raise_team_access_denied()
|
||||
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="You do not have access to this team",
|
||||
)
|
||||
|
||||
_GENERAL_SETTINGS: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _general_settings() -> Mapping[str, object]:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
return _GENERAL_SETTINGS.validate_python(general_settings)
|
||||
|
||||
|
||||
def _caller_edit_access(role: TeamAccessRole | None, general_settings: Mapping[str, object]) -> TeamEditAccess:
|
||||
"""What the caller may change on /team/update, reported on /team/info so the dashboard never re-derives it."""
|
||||
match role:
|
||||
case "proxy_admin" | "org_admin":
|
||||
return TeamEditUnrestricted()
|
||||
case "team_admin":
|
||||
permitted: Final = resolve_team_admin_editable_fields(
|
||||
general_settings, SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS
|
||||
)
|
||||
if not permitted:
|
||||
return TeamEditAsTeamAdminDisabled()
|
||||
return TeamEditAsTeamAdmin(editable_fields=tuple(sorted(permitted)))
|
||||
case None:
|
||||
return TeamEditNone()
|
||||
case _:
|
||||
assert_never(role)
|
||||
|
||||
|
||||
class TeamMemberBudgetHandler:
|
||||
|
|
@ -2144,16 +2204,29 @@ async def update_team(
|
|||
)
|
||||
|
||||
if existing_team_row is None:
|
||||
# Non-proxy-admins get the same 403 as an access denial so /team/update
|
||||
# cannot be used to probe which team ids exist
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
_raise_team_access_denied()
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Team not found, passed team_id={data.team_id}"},
|
||||
)
|
||||
|
||||
# Verify caller has access to manage this team
|
||||
await _verify_team_access(
|
||||
team_obj=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
existing_team: Final = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump())
|
||||
access_role: Final = await _resolve_team_access(team_obj=existing_team, user_api_key_dict=user_api_key_dict)
|
||||
if access_role is None:
|
||||
_raise_team_access_denied()
|
||||
if access_role == "team_admin":
|
||||
data = team_admin_request_or_raise( # rebind-ok: resent values must not reach the derived writes below
|
||||
team_admin_edit_verdict(
|
||||
data=data,
|
||||
existing=existing_team,
|
||||
permitted=resolve_team_admin_editable_fields(
|
||||
_general_settings(), SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
await validate_router_settings_weights(
|
||||
data.router_settings,
|
||||
|
|
@ -2257,6 +2330,7 @@ async def update_team(
|
|||
org_id=org_id_to_check,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
prisma_client=prisma_client,
|
||||
include_budget_table=True,
|
||||
)
|
||||
if org_table is not None:
|
||||
await _check_org_team_limits(
|
||||
|
|
@ -4583,10 +4657,9 @@ async def team_info(
|
|||
)
|
||||
team_table: Final = LiteLLM_TeamTable.model_validate(team_info.model_dump())
|
||||
await validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_table)
|
||||
access_role: Final = await _resolve_team_access(team_obj=team_table, user_api_key_dict=user_api_key_dict)
|
||||
organization_models: Final[list[str] | None] = (
|
||||
_parent_organization_models(team_info)
|
||||
if await _can_manage_team(team_obj=team_table, user_api_key_dict=user_api_key_dict)
|
||||
else None
|
||||
_parent_organization_models(team_info) if access_role is not None else None
|
||||
)
|
||||
|
||||
## GET ALL KEYS ##
|
||||
|
|
@ -4655,6 +4728,7 @@ async def team_info(
|
|||
model_max_budget=resolved_team_info.model_max_budget,
|
||||
cache=model_max_budget_limiter.dual_cache,
|
||||
),
|
||||
"caller_edit_access": _caller_edit_access(access_role, _general_settings()),
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1180,9 +1180,8 @@ async def bedrock_proxy_route(
|
|||
endpoint_func: Final = create_pass_through_route(
|
||||
endpoint=endpoint,
|
||||
target=str(prepped.url),
|
||||
custom_headers=prepped.headers,
|
||||
custom_headers=_upstream_headers_for_bedrock_agent_runtime_route(request, user_api_key_dict, prepped.headers),
|
||||
is_streaming_request=is_streaming_request,
|
||||
_forward_headers=True,
|
||||
) # dynamically construct pass-through endpoint based on incoming path
|
||||
setattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, data)
|
||||
# SigV4 signs an exact payload; pass-through must send prepped.body, not json.dumps
|
||||
|
|
@ -2001,6 +2000,9 @@ _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS: Final = frozenset({"authorization", "x-a
|
|||
_HEADERS_NEVER_FORWARDED_TO_ANTHROPIC: Final = frozenset({"content-length", "host", "accept-encoding"}) | (
|
||||
SpecialHeaders.litellm_credential_header_names() - _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS
|
||||
)
|
||||
_HEADERS_NEVER_FORWARDED_TO_BEDROCK: Final = (
|
||||
frozenset({"content-length", "host", "accept-encoding"}) | SpecialHeaders.litellm_credential_header_names()
|
||||
)
|
||||
|
||||
|
||||
_MAPPED_ROUTE_CALLER_KEY_HEADER: Final = "litellm_user_api_key"
|
||||
|
|
@ -2099,6 +2101,17 @@ def _upstream_headers_for_anthropic_route(
|
|||
return MappingProxyType({**caller_headers, **(proxy_auth_header or {})})
|
||||
|
||||
|
||||
def _upstream_headers_for_bedrock_agent_runtime_route(
|
||||
request: Request, user_api_key_dict: UserAPIKeyAuth, signed_headers: Mapping[str, object]
|
||||
) -> Mapping[str, object]:
|
||||
caller_headers: Final = _caller_headers_without_litellm_secrets(
|
||||
request,
|
||||
user_api_key_dict,
|
||||
_HEADERS_NEVER_FORWARDED_TO_BEDROCK | frozenset(name.lower() for name in signed_headers),
|
||||
)
|
||||
return MappingProxyType({**caller_headers, **signed_headers})
|
||||
|
||||
|
||||
async def _prepare_vertex_auth_headers(
|
||||
request: Request,
|
||||
vertex_credentials: VertexPassThroughCredentials | None,
|
||||
|
|
|
|||
|
|
@ -72,7 +72,9 @@ from litellm.proxy.auth.auth_utils import request_dispatched_to_pass_through_end
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
log_llm_api_exception,
|
||||
open_sse_before_first_byte,
|
||||
resolve_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
|
|
@ -80,6 +82,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
litellm_call_id_headers,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
|
|
@ -196,14 +199,15 @@ async def chat_completion_pass_through_endpoint(
|
|||
version,
|
||||
)
|
||||
|
||||
data = {}
|
||||
litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
|
||||
data = {"litellm_call_id": litellm_call_id}
|
||||
try:
|
||||
body: Final = await request.body()
|
||||
body_str: Final = body.decode()
|
||||
try:
|
||||
data = ast.literal_eval(body_str)
|
||||
data = ast.literal_eval(body_str) | data
|
||||
except Exception:
|
||||
data = json.loads(body_str)
|
||||
data = json.loads(body_str) | data
|
||||
|
||||
data["adapter_id"] = adapter_id
|
||||
|
||||
|
|
@ -290,9 +294,7 @@ async def chat_completion_pass_through_endpoint(
|
|||
response_cost: Final = hidden_params.get("response_cost", None) or ""
|
||||
|
||||
### ALERTING ###
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(litellm_call_id=data.get("litellm_call_id", ""), status="success")
|
||||
)
|
||||
asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
|
||||
|
||||
verbose_proxy_logger.debug("final response: %s", response)
|
||||
|
||||
|
|
@ -313,12 +315,13 @@ async def chat_completion_pass_through_endpoint(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.completion(): Exception occured - %s", e)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -351,7 +351,10 @@ from litellm.proxy.common_request_processing import (
|
|||
_is_azure_model_router_request,
|
||||
_should_return_raw_model_name,
|
||||
create_response,
|
||||
log_llm_api_exception,
|
||||
open_sse_before_first_byte,
|
||||
request_litellm_call_id,
|
||||
resolve_litellm_call_id,
|
||||
ttft_keepalive_interval,
|
||||
)
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
|
||||
|
|
@ -390,6 +393,11 @@ from litellm.proxy.common_utils.model_listing_utils import (
|
|||
from litellm.proxy.common_utils.openai_endpoint_utils import (
|
||||
remove_sensitive_info_from_deployment,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
headers_with_litellm_call_id,
|
||||
litellm_call_id_headers,
|
||||
with_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.common_utils.periodic_reload_schedule import (
|
||||
MODEL_COST_MAP_RELOAD_PARAM_NAME,
|
||||
clear_reload_interval,
|
||||
|
|
@ -682,6 +690,9 @@ from litellm.proxy.types_utils.utils import get_instance_fn
|
|||
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
|
||||
router as ui_crud_endpoints_router,
|
||||
)
|
||||
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
|
||||
sync_ui_settings_to_general_settings,
|
||||
)
|
||||
from litellm.proxy.ui_crud_endpoints.user_banner_endpoints import (
|
||||
router as user_banner_endpoints_router,
|
||||
)
|
||||
|
|
@ -1764,10 +1775,6 @@ class _SSOConfigRow(Protocol):
|
|||
sso_settings: MutableMapping[str, object]
|
||||
|
||||
|
||||
class _UISettingsRow(Protocol):
|
||||
ui_settings: Mapping[str, object] | str | None
|
||||
|
||||
|
||||
class _InvitationLinkRow(Protocol):
|
||||
user_id: str
|
||||
expires_at: datetime
|
||||
|
|
@ -7423,7 +7430,12 @@ class ProxyConfig:
|
|||
Returns what the reconcile saw, captured before the lock is released so a
|
||||
caller's verdict cannot be corrupted by the next reconcile's own in-flight
|
||||
window. See ReconcileOutcome.
|
||||
|
||||
Also re-reads the UI settings that back runtime flags. That runs before the lock, so a
|
||||
setting written through one pod reaches the others without waiting on a model reconcile.
|
||||
"""
|
||||
await sync_ui_settings_to_general_settings(prisma_client)
|
||||
|
||||
async with MODEL_RECONCILE_LOCK:
|
||||
return await self._add_deployment_locked(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj)
|
||||
|
||||
|
|
@ -9670,35 +9682,12 @@ class ProxyStartupEvent:
|
|||
|
||||
@classmethod
|
||||
async def _sync_ui_settings_to_general_settings(cls):
|
||||
"""
|
||||
Load persisted UI settings from the database and sync runtime flags
|
||||
into general_settings so they take effect immediately after startup.
|
||||
"""
|
||||
try:
|
||||
import json
|
||||
|
||||
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
|
||||
_RUNTIME_GENERAL_SETTINGS_FLAGS,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
db_record: Final[_UISettingsRow | None] = cast( # cast-ok: prisma Json stub is `str`, runtime is a dict
|
||||
"_UISettingsRow | None",
|
||||
await UISettingsRepository(prisma_client).table.find_unique(where={"id": "ui_settings"}),
|
||||
)
|
||||
if db_record and db_record.ui_settings:
|
||||
raw: Final = db_record.ui_settings
|
||||
ui_settings: Final = json.loads(raw) if isinstance(raw, str) else dict(raw)
|
||||
flags_to_sync: Final = {k: ui_settings[k] for k in _RUNTIME_GENERAL_SETTINGS_FLAGS if k in ui_settings}
|
||||
if flags_to_sync:
|
||||
general_settings.update(flags_to_sync)
|
||||
verbose_proxy_logger.info(
|
||||
"Synced UI settings to general_settings on startup: %s",
|
||||
list(flags_to_sync.keys()),
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("UI settings sync on startup skipped or failed: %s", e)
|
||||
"""Apply the persisted UI settings to general_settings before this pod serves traffic."""
|
||||
if prisma_client is None:
|
||||
return
|
||||
applied: Final = await sync_ui_settings_to_general_settings(prisma_client)
|
||||
if applied:
|
||||
verbose_proxy_logger.info("Synced UI settings to general_settings on startup: %s", list(applied))
|
||||
|
||||
@classmethod
|
||||
async def _load_heuristic_v1_tuning_baselines(
|
||||
|
|
@ -11332,12 +11321,14 @@ async def completion(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.completion(): Exception occured - %s", e)
|
||||
litellm_call_id: Final = request_litellm_call_id(data)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
openai_code=getattr(e, "code", None),
|
||||
code=getattr(e, "status_code", 500),
|
||||
)
|
||||
|
|
@ -11494,11 +11485,12 @@ async def moderations(
|
|||
```
|
||||
"""
|
||||
global proxy_logging_obj
|
||||
data: dict = {}
|
||||
litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
|
||||
data: dict = {"litellm_call_id": litellm_call_id}
|
||||
try:
|
||||
# Use orjson to parse JSON data, orjson speeds up requests significantly
|
||||
body: Final = await request.body()
|
||||
data = orjson.loads(body)
|
||||
data = orjson.loads(body) | data
|
||||
|
||||
# Include original request and headers in the data
|
||||
data = await add_litellm_data_to_request(
|
||||
|
|
@ -11535,9 +11527,7 @@ async def moderations(
|
|||
response: Final = await llm_call
|
||||
|
||||
### ALERTING ###
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(litellm_call_id=data.get("litellm_call_id", ""), status="success")
|
||||
)
|
||||
asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
|
||||
|
||||
### RESPONSE HEADERS ###
|
||||
hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
|
||||
|
|
@ -11563,14 +11553,15 @@ async def moderations(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.moderations(): Exception occured - %s", e)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
if isinstance(e, ProxyException):
|
||||
raise
|
||||
raise with_litellm_call_id(e, litellm_call_id)
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
|
|
@ -11579,6 +11570,7 @@ async def moderations(
|
|||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
code=getattr(e, "status_code", 500),
|
||||
)
|
||||
|
||||
|
|
@ -11616,11 +11608,12 @@ async def audio_speech(
|
|||
https://platform.openai.com/docs/api-reference/audio/createSpeech
|
||||
"""
|
||||
global proxy_logging_obj
|
||||
data: dict = {}
|
||||
litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
|
||||
data: dict = {"litellm_call_id": litellm_call_id}
|
||||
try:
|
||||
# Use orjson to parse JSON data, orjson speeds up requests significantly
|
||||
body: Final = await request.body()
|
||||
data = orjson.loads(body)
|
||||
data = orjson.loads(body) | data
|
||||
|
||||
# Include original request and headers in the data
|
||||
data = await add_litellm_data_to_request(
|
||||
|
|
@ -11653,9 +11646,7 @@ async def audio_speech(
|
|||
response: Final = await llm_call
|
||||
|
||||
### ALERTING ###
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(litellm_call_id=data.get("litellm_call_id", ""), status="success")
|
||||
)
|
||||
asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
|
||||
|
||||
### RESPONSE HEADERS ###
|
||||
hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
|
||||
|
|
@ -11663,7 +11654,7 @@ async def audio_speech(
|
|||
cache_key: Final = hidden_params.get("cache_key", None) or ""
|
||||
api_base: Final = hidden_params.get("api_base", None) or ""
|
||||
response_cost: Final = hidden_params.get("response_cost", None) or ""
|
||||
litellm_call_id: Final = hidden_params.get("litellm_call_id", None) or ""
|
||||
response_call_id: Final = hidden_params.get("litellm_call_id", None) or ""
|
||||
|
||||
custom_headers: Final = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -11674,7 +11665,7 @@ async def audio_speech(
|
|||
response_cost=response_cost,
|
||||
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
|
||||
fastest_response_batch_completion=None,
|
||||
call_id=litellm_call_id,
|
||||
call_id=response_call_id,
|
||||
request_data=data,
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
|
|
@ -11710,14 +11701,20 @@ async def audio_speech(
|
|||
original_exception=e,
|
||||
request_data=data,
|
||||
)
|
||||
verbose_proxy_logger.error("litellm.proxy.proxy_server.audio_speech(): Exception occured - %s", e)
|
||||
verbose_proxy_logger.debug(traceback.format_exc())
|
||||
if isinstance(e, (ProxyException, HTTPException)):
|
||||
raise e
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
if isinstance(e, ProxyException):
|
||||
raise with_litellm_call_id(e, litellm_call_id)
|
||||
if isinstance(e, HTTPException):
|
||||
raise HTTPException(
|
||||
status_code=e.status_code,
|
||||
detail=e.detail,
|
||||
headers=headers_with_litellm_call_id(e.headers, litellm_call_id),
|
||||
)
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", f"{e}"),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
openai_code=getattr(e, "code", None),
|
||||
code=getattr(e, "status_code", 500),
|
||||
)
|
||||
|
|
@ -11745,11 +11742,12 @@ async def audio_transcriptions(
|
|||
https://platform.openai.com/docs/api-reference/audio/createTranscription?lang=curl
|
||||
"""
|
||||
global proxy_logging_obj
|
||||
data: dict = {}
|
||||
litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
|
||||
data: dict = {"litellm_call_id": litellm_call_id}
|
||||
try:
|
||||
# Use orjson to parse JSON data, orjson speeds up requests significantly
|
||||
form_data: Final = await get_form_data(request)
|
||||
data = {key: value for key, value in form_data.items() if key != "file"}
|
||||
data = {key: value for key, value in form_data.items() if key != "file"} | data
|
||||
|
||||
# Include original request and headers in the data
|
||||
data = await add_litellm_data_to_request(
|
||||
|
|
@ -11816,9 +11814,7 @@ async def audio_transcriptions(
|
|||
file_object.close() # close the file read in by io library
|
||||
|
||||
### ALERTING ###
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(litellm_call_id=data.get("litellm_call_id", ""), status="success")
|
||||
)
|
||||
asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
|
||||
|
||||
### RESPONSE HEADERS ###
|
||||
hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
|
||||
|
|
@ -11826,7 +11822,7 @@ async def audio_transcriptions(
|
|||
cache_key: Final = hidden_params.get("cache_key", None) or ""
|
||||
api_base: Final = hidden_params.get("api_base", None) or ""
|
||||
response_cost: Final = hidden_params.get("response_cost", None) or ""
|
||||
litellm_call_id: Final = hidden_params.get("litellm_call_id", None) or ""
|
||||
response_call_id: Final = hidden_params.get("litellm_call_id", None) or ""
|
||||
additional_headers: Final[dict] = hidden_params.get("additional_headers", {}) or {}
|
||||
|
||||
fastapi_response.headers.update(
|
||||
|
|
@ -11838,7 +11834,7 @@ async def audio_transcriptions(
|
|||
version=version,
|
||||
response_cost=response_cost,
|
||||
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
|
||||
call_id=litellm_call_id,
|
||||
call_id=response_call_id,
|
||||
request_data=data,
|
||||
hidden_params=hidden_params,
|
||||
**additional_headers,
|
||||
|
|
@ -11860,12 +11856,13 @@ async def audio_transcriptions(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.audio_transcription(): Exception occured - %s", e)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e.detail)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
|
|
@ -11874,6 +11871,7 @@ async def audio_transcriptions(
|
|||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
openai_code=getattr(e, "code", None),
|
||||
code=getattr(e, "status_code", 500),
|
||||
)
|
||||
|
|
@ -12892,7 +12890,6 @@ from litellm.repositories.table_repositories import (
|
|||
InvitationLinkRepository,
|
||||
PromptRepository,
|
||||
SSOConfigRepository,
|
||||
UISettingsRepository,
|
||||
)
|
||||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
|
|
@ -15570,18 +15567,34 @@ async def model_group_info(
|
|||
from litellm.proxy.utils import get_available_models_for_user
|
||||
|
||||
# Get available models for the user
|
||||
all_models_str: Final = await get_available_models_for_user(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
user_model=user_model,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id=None,
|
||||
include_model_access_groups=False,
|
||||
only_model_access_groups=False,
|
||||
return_wildcard_routes=False,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
is_proxy_admin: Final = user_api_key_dict.user_role in (
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
)
|
||||
all_models_str: Final = (
|
||||
get_complete_model_list(
|
||||
key_models=(),
|
||||
team_models=(),
|
||||
proxy_model_list=llm_router.get_model_names(),
|
||||
user_model=user_model,
|
||||
infer_model_from_keys=general_settings.get("infer_model_from_keys", False),
|
||||
return_wildcard_routes=False,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
if is_proxy_admin
|
||||
else await get_available_models_for_user(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
user_model=user_model,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id=None,
|
||||
include_model_access_groups=False,
|
||||
only_model_access_groups=False,
|
||||
return_wildcard_routes=False,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
)
|
||||
model_groups: list[ModelGroupInfoProxy] = _get_model_group_info(
|
||||
llm_router=llm_router, all_models_str=all_models_str, model_group=model_group
|
||||
|
|
|
|||
|
|
@ -7,12 +7,16 @@ import orjson
|
|||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from fastapi.responses import ORJSONResponse
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
log_llm_api_exception,
|
||||
resolve_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
litellm_call_id_headers,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
|
|
@ -54,10 +58,11 @@ async def rerank(
|
|||
version,
|
||||
)
|
||||
|
||||
data = {}
|
||||
litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
|
||||
data = {"litellm_call_id": litellm_call_id}
|
||||
try:
|
||||
body: Final = await request.body()
|
||||
data = orjson.loads(body)
|
||||
data = orjson.loads(body) | data
|
||||
|
||||
# Include original request and headers in the data
|
||||
data = await add_litellm_data_to_request(
|
||||
|
|
@ -82,9 +87,7 @@ async def rerank(
|
|||
response: Final = await llm_call
|
||||
|
||||
### ALERTING ###
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(litellm_call_id=data.get("litellm_call_id", ""), status="success")
|
||||
)
|
||||
asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
|
||||
|
||||
### RESPONSE HEADERS ###
|
||||
hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
|
||||
|
|
@ -95,7 +98,7 @@ async def rerank(
|
|||
fastapi_response.headers.update(
|
||||
ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_id=hidden_params.get("litellm_call_id", None) or data.get("litellm_call_id", None),
|
||||
call_id=hidden_params.get("litellm_call_id", None) or litellm_call_id,
|
||||
model_id=model_id,
|
||||
cache_key=cache_key,
|
||||
api_base=api_base,
|
||||
|
|
@ -113,12 +116,13 @@ async def rerank(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.error("litellm.proxy.proxy_server.rerank(): Exception occured - %s", e)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
|
|
@ -127,5 +131,6 @@ async def rerank(
|
|||
message=getattr(e, "message", error_msg),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
headers=litellm_call_id_headers(litellm_call_id),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ from typing import (
|
|||
from urllib.parse import urlparse
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, File, HTTPException, UploadFile
|
||||
from pydantic import ConfigDict, JsonValue, ValidationError, create_model
|
||||
from pydantic import ConfigDict, JsonValue, TypeAdapter, ValidationError, create_model
|
||||
from pydantic.fields import FieldInfo
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
|
|
@ -29,6 +29,10 @@ from litellm.proxy.config_resolvers.sso import (
|
|||
SSO_SECRET_FIELDS,
|
||||
resolve_sso_config,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_admin_field_permissions import (
|
||||
SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS,
|
||||
TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
|
||||
from litellm.proxy.utils import invalidate_config_param
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
|
|
@ -212,6 +216,9 @@ class UIThemeSettingsResponse(SettingsResponse):
|
|||
"""Response model for UI theme settings"""
|
||||
|
||||
|
||||
_TEAM_ADMIN_FIELD_ENUM: Final = tuple(sorted(SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS))
|
||||
|
||||
|
||||
class UISettings(BaseModel):
|
||||
"""Configuration for UI-specific flags"""
|
||||
|
||||
|
|
@ -304,6 +311,18 @@ class UISettings(BaseModel):
|
|||
description="If true, shows the Chat page in the UI sidebar, letting users chat with an LLM and connect their own MCP server credentials via OAuth.",
|
||||
)
|
||||
|
||||
team_admin_editable_team_fields: Sequence[str] = Field(
|
||||
default=(),
|
||||
description=(
|
||||
"Team settings fields a team admin may change on the teams they administer. "
|
||||
"Empty means team admins cannot edit team settings at all. "
|
||||
"Proxy admins and org admins are not affected."
|
||||
),
|
||||
json_schema_extra={ # mutable-ok: pydantic only merges json_schema_extra when it is a plain dict
|
||||
"items": {"type": "string", "enum": [*_TEAM_ADMIN_FIELD_ENUM]}, # mutable-ok: nested in the dict above
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class UISettingsResponse(SettingsResponse):
|
||||
"""Response model for UI settings"""
|
||||
|
|
@ -326,6 +345,7 @@ ALLOWED_UI_SETTINGS_FIELDS: Final = {
|
|||
"disable_custom_api_keys",
|
||||
"disable_key_generate_for_org_admin",
|
||||
"enable_chat_ui",
|
||||
TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING,
|
||||
}
|
||||
|
||||
ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: Final = "enable_ptu_cost_attribution"
|
||||
|
|
@ -360,6 +380,7 @@ _RUNTIME_GENERAL_SETTINGS_FLAGS: Final = [
|
|||
"disable_vector_stores_for_internal_users",
|
||||
"allow_vector_stores_for_team_admins",
|
||||
"disable_key_generate_for_org_admin",
|
||||
TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING,
|
||||
]
|
||||
|
||||
# Extension point: packages outside OSS (e.g. litellm_enterprise) can
|
||||
|
|
@ -1457,6 +1478,42 @@ async def get_ui_settings_cached() -> dict[str, JsonValue]:
|
|||
return ui_settings
|
||||
|
||||
|
||||
_UI_SETTINGS_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
|
||||
|
||||
def apply_runtime_general_settings_flags(ui_settings: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
|
||||
"""Copy the UI settings that gate runtime behavior into ``general_settings``. Returns what was applied."""
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
flags: Final = {k: ui_settings[k] for k in _RUNTIME_GENERAL_SETTINGS_FLAGS if k in ui_settings}
|
||||
if flags:
|
||||
general_settings.update(flags)
|
||||
return MappingProxyType(flags)
|
||||
|
||||
|
||||
async def sync_ui_settings_to_general_settings(prisma_client: object) -> Mapping[str, JsonValue]:
|
||||
"""Re-read the persisted UI settings and apply the runtime flags to ``general_settings``.
|
||||
|
||||
Runs on startup and on every periodic config reload: the PATCH handler only updates the pod
|
||||
that served it, so every other pod needs its own read to pick up a change without a restart.
|
||||
Never raises. A read that fails leaves this pod on the flags it already had.
|
||||
"""
|
||||
try:
|
||||
db_record: Final = await _ui_settings_db(UISettingsRepository(prisma_client)).find_unique(
|
||||
where={"id": "ui_settings"}
|
||||
)
|
||||
stored: Final = (db_record.ui_settings if db_record else None) or "{}"
|
||||
parsed: Final = (
|
||||
_UI_SETTINGS_OBJECT.validate_json(stored)
|
||||
if isinstance(stored, str)
|
||||
else _UI_SETTINGS_OBJECT.validate_python(stored)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("Could not refresh UI settings from the database: %s", e)
|
||||
return MappingProxyType({})
|
||||
return apply_runtime_general_settings_flags(parsed)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/get/ui_settings",
|
||||
tags=["UI Settings"],
|
||||
|
|
@ -1485,13 +1542,7 @@ async def get_ui_settings():
|
|||
# Sanitize any unexpected keys from persisted config before returning
|
||||
ui_settings: Final = {k: v for k, v in parsed.items() if k in ALLOWED_UI_SETTINGS_FIELDS}
|
||||
|
||||
# Sync runtime flags into general_settings so the proxy picks them up
|
||||
# at runtime (covers server restart scenarios).
|
||||
_flags_to_sync: Final = {k: ui_settings[k] for k in _RUNTIME_GENERAL_SETTINGS_FLAGS if k in ui_settings}
|
||||
if _flags_to_sync:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
general_settings.update(_flags_to_sync)
|
||||
apply_runtime_general_settings_flags(ui_settings)
|
||||
|
||||
# Refresh DualCache so other code paths (e.g. /user/filter/ui) see fresh values
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
|
@ -1571,6 +1622,20 @@ async def update_ui_settings(
|
|||
except ValidationError as e:
|
||||
raise HTTPException(status_code=422, detail=e.errors())
|
||||
|
||||
unsupported_team_fields: Final = sorted(
|
||||
frozenset(settings.team_admin_editable_team_fields) - SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS
|
||||
)
|
||||
if unsupported_team_fields:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={ # mutable-ok: HTTPException detail must be a plain dict for FastAPI JSON serialization
|
||||
"error": (
|
||||
f"{TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING} does not support {unsupported_team_fields}. "
|
||||
f"Supported fields: {sorted(SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS)}."
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
# Only include fields the caller actually sent (not Pydantic defaults).
|
||||
settings_dict: Final[Mapping[str, JsonValue]] = settings.model_dump(exclude_unset=True)
|
||||
|
||||
|
|
@ -1616,13 +1681,7 @@ async def update_ui_settings(
|
|||
},
|
||||
)
|
||||
|
||||
# Sync runtime flags to general_settings so the proxy picks them up
|
||||
# at runtime (general_settings is checked in pre-call utils).
|
||||
_flags_to_sync: Final = {k: ui_settings[k] for k in _RUNTIME_GENERAL_SETTINGS_FLAGS if k in ui_settings}
|
||||
if _flags_to_sync:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
general_settings.update(_flags_to_sync)
|
||||
apply_runtime_general_settings_flags(ui_settings)
|
||||
|
||||
# Invalidate + set DualCache so subsequent reads see the new values immediately
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
|
|
|||
|
|
@ -38,7 +38,11 @@ from litellm.proxy._types import (
|
|||
SpendLogsMetadata,
|
||||
SpendLogsPayload,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import openai_error_param
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
litellm_call_id_headers,
|
||||
openai_error_param,
|
||||
with_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.proxy.model_listing import ModelInfoResponse
|
||||
|
|
@ -3050,7 +3054,7 @@ class ProxyLogging:
|
|||
if litellm_logging_obj is None:
|
||||
from litellm._uuid import uuid
|
||||
|
||||
request_data["litellm_call_id"] = str(uuid.uuid4())
|
||||
request_data.setdefault("litellm_call_id", str(uuid.uuid4()))
|
||||
user_api_key_logged_metadata: Final = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(
|
||||
user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
|
@ -7659,7 +7663,7 @@ def _recreate_writer_on_read_only_transaction(prisma_client: "PrismaClient | Non
|
|||
asyncio.create_task(prisma_client.recreate_read_only_writer(reason="postgres_read_only_transaction"))
|
||||
|
||||
|
||||
def handle_exception_on_proxy(e: Exception) -> ProxyException:
|
||||
def handle_exception_on_proxy(e: Exception, litellm_call_id: str | None = None) -> ProxyException:
|
||||
"""
|
||||
Returns an Exception as ProxyException, this ensures all exceptions are OpenAI API compatible
|
||||
"""
|
||||
|
|
@ -7671,20 +7675,23 @@ def handle_exception_on_proxy(e: Exception) -> ProxyException:
|
|||
|
||||
_recreate_writer_on_read_only_transaction(prisma_client)
|
||||
|
||||
headers: Final = litellm_call_id_headers(litellm_call_id)
|
||||
if isinstance(e, HTTPException):
|
||||
return ProxyException(
|
||||
message=getattr(e, "detail", f"error({e})"),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
param=openai_error_param(e),
|
||||
headers=headers,
|
||||
code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR),
|
||||
)
|
||||
elif isinstance(e, ProxyException):
|
||||
return e
|
||||
return with_litellm_call_id(e, litellm_call_id)
|
||||
_status_code: Final = getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR)
|
||||
return ProxyException(
|
||||
message=str(e),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
param=openai_error_param(e),
|
||||
headers=headers,
|
||||
code=_status_code,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -482,18 +482,27 @@ class _AsyncPromptManagementOutcome:
|
|||
|
||||
|
||||
def _resolve_responses_api_provider_config(
|
||||
model: str, custom_llm_provider: str, model_info: object
|
||||
model: str, custom_llm_provider: str, model_info: object, api_base: str | None
|
||||
) -> BaseResponsesAPIConfig | None:
|
||||
provider_config: Final = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model, provider=custom_llm_provider
|
||||
model=model, provider=custom_llm_provider, api_base=api_base
|
||||
)
|
||||
if provider_config is not None or not _deployment_passes_through_responses(model_info):
|
||||
return provider_config
|
||||
return OpenAILikeResponsesConfig()
|
||||
|
||||
|
||||
def _api_base_kwarg(kwargs: Mapping[str, object]) -> str | None:
|
||||
api_base: Final = kwargs.get("api_base")
|
||||
return api_base if isinstance(api_base, str) else None
|
||||
|
||||
|
||||
def _will_bridge_to_chat_completions(
|
||||
model: str, custom_llm_provider: str | None, use_chat_completions_api: bool, model_info: object
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
use_chat_completions_api: bool,
|
||||
model_info: object,
|
||||
api_base: str | None,
|
||||
) -> bool:
|
||||
"""``_bridges_to_chat_completions`` for callers running before the provider config is resolved.
|
||||
|
||||
|
|
@ -507,7 +516,7 @@ def _will_bridge_to_chat_completions(
|
|||
if custom_llm_provider is None:
|
||||
return True
|
||||
return _bridges_to_chat_completions(
|
||||
_resolve_responses_api_provider_config(normalized_model[0], custom_llm_provider, model_info),
|
||||
_resolve_responses_api_provider_config(normalized_model[0], custom_llm_provider, model_info, api_base),
|
||||
use_chat_completions_api or normalized_model[1],
|
||||
)
|
||||
|
||||
|
|
@ -618,6 +627,7 @@ async def aresponses(
|
|||
custom_llm_provider,
|
||||
bool(kwargs.get("use_chat_completions_api")),
|
||||
kwargs.get("model_info"),
|
||||
_api_base_kwarg(kwargs),
|
||||
),
|
||||
):
|
||||
(
|
||||
|
|
@ -783,7 +793,11 @@ def _apply_prompt_management_to_responses_call(
|
|||
with _prompt_management_sees_a_provisional_message_list(
|
||||
kwargs,
|
||||
bridged=_will_bridge_to_chat_completions(
|
||||
model, custom_llm_provider, use_chat_completions_api, kwargs.get("model_info")
|
||||
model,
|
||||
custom_llm_provider,
|
||||
use_chat_completions_api,
|
||||
kwargs.get("model_info"),
|
||||
_api_base_kwarg(kwargs),
|
||||
),
|
||||
):
|
||||
(
|
||||
|
|
@ -1237,7 +1251,7 @@ def responses(
|
|||
responses_api_provider_config = None
|
||||
else:
|
||||
responses_api_provider_config = _resolve_responses_api_provider_config(
|
||||
model, custom_llm_provider, deployment_model_info
|
||||
model, custom_llm_provider, deployment_model_info, litellm_params.api_base
|
||||
)
|
||||
|
||||
if (
|
||||
|
|
@ -1496,6 +1510,7 @@ def delete_responses(
|
|||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
api_base=litellm_params.api_base,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -1667,6 +1682,7 @@ def get_responses(
|
|||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
api_base=litellm_params.api_base,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -1811,6 +1827,7 @@ def list_input_items(
|
|||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
api_base=litellm_params.api_base,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -1960,6 +1977,7 @@ def cancel_responses(
|
|||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
api_base=litellm_params.api_base,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -2132,6 +2150,7 @@ def compact_responses(
|
|||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=custom_llm_provider,
|
||||
api_base=litellm_params.api_base,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -2270,14 +2289,15 @@ async def _aresponses_websocket(
|
|||
custom_llm_provider=_custom_llm_provider,
|
||||
)
|
||||
|
||||
resolved_api_base: Final = dynamic_api_base or litellm_params.api_base or litellm.api_base or None
|
||||
responses_api_provider_config: BaseResponsesAPIConfig | None = None
|
||||
if _custom_llm_provider is not None:
|
||||
responses_api_provider_config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=resolved_model,
|
||||
provider=litellm.LlmProviders(_custom_llm_provider),
|
||||
api_base=resolved_api_base,
|
||||
)
|
||||
|
||||
resolved_api_base: Final = dynamic_api_base or litellm_params.api_base or litellm.api_base or None
|
||||
resolved_api_key: Final = (
|
||||
dynamic_api_key
|
||||
or litellm_params.api_key
|
||||
|
|
|
|||
|
|
@ -2689,7 +2689,7 @@ def declared_value_factory(model: str, custom_llm_provider: str | None, key: str
|
|||
"""Return a string value the model map declares for *key*, or ``None`` when it says nothing.
|
||||
|
||||
The string-valued sibling of :func:`_supports_factory` and
|
||||
:func:`_is_explicitly_disabled_factory`, public where those two are not because it is read
|
||||
:func:`is_explicitly_disabled_factory`, public like the latter because both are read
|
||||
from the provider configs rather than from this module, sharing their
|
||||
``get_llm_provider`` -> ``_get_model_info_helper`` chain and their unprefixed-twin
|
||||
fallback (#20885), so a provider-prefixed entry that omits the key still answers
|
||||
|
|
@ -2725,7 +2725,7 @@ def declared_value_factory(model: str, custom_llm_provider: str | None, key: str
|
|||
return None
|
||||
|
||||
|
||||
def _is_explicitly_disabled_factory(model: str, custom_llm_provider: str | None, key: str) -> bool:
|
||||
def is_explicitly_disabled_factory(model: str, custom_llm_provider: str | None, key: str) -> bool:
|
||||
"""Return True only when the model map explicitly sets *key* to ``False``.
|
||||
|
||||
This is the opt-out mirror of :func:`_supports_factory`. Where
|
||||
|
|
@ -2844,7 +2844,7 @@ def is_vision_explicitly_disabled(model: str, custom_llm_provider: str | None =
|
|||
The opt-out mirror of :func:`supports_vision`: a missing declaration reads as not
|
||||
disabled, so unknown or newly added models stay eligible for image routing.
|
||||
"""
|
||||
return _is_explicitly_disabled_factory(model, custom_llm_provider, "supports_vision")
|
||||
return is_explicitly_disabled_factory(model, custom_llm_provider, "supports_vision")
|
||||
|
||||
|
||||
def supports_vision(model: str, custom_llm_provider: str | None = None) -> bool:
|
||||
|
|
@ -8746,6 +8746,7 @@ class ProviderConfigManager:
|
|||
def get_provider_responses_api_config(
|
||||
provider: LlmProviders | str,
|
||||
model: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> BaseResponsesAPIConfig | None:
|
||||
from litellm.llms.openai_like.dynamic_config import (
|
||||
create_responses_config_class,
|
||||
|
|
@ -8767,7 +8768,7 @@ class ProviderConfigManager:
|
|||
pass
|
||||
|
||||
# Check Python classes first (custom overrides take priority)
|
||||
result: Final = ProviderConfigManager._get_python_responses_api_config(provider_enum, model)
|
||||
result: Final = ProviderConfigManager._get_python_responses_api_config(provider_enum, model, api_base)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
|
|
@ -8783,6 +8784,7 @@ class ProviderConfigManager:
|
|||
def _get_python_responses_api_config(
|
||||
provider: LlmProviders | None,
|
||||
model: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> BaseResponsesAPIConfig | None:
|
||||
"""Check for Python-class-based responses API configs (custom overrides)."""
|
||||
if provider is None:
|
||||
|
|
@ -8801,6 +8803,14 @@ class ProviderConfigManager:
|
|||
return litellm.AzureOpenAIOSeriesResponsesAPIConfig()
|
||||
else:
|
||||
return litellm.AzureOpenAIResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.AZURE_AI == provider:
|
||||
from litellm.llms.azure_ai.common_utils import (
|
||||
azure_ai_supports_native_responses,
|
||||
)
|
||||
|
||||
if azure_ai_supports_native_responses(model, api_base):
|
||||
return litellm.AzureAIResponsesAPIConfig()
|
||||
return None
|
||||
elif litellm.LlmProviders.XAI == provider:
|
||||
return litellm.XAIResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
|
||||
|
|
|
|||
|
|
@ -26105,6 +26105,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -26162,6 +26163,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -28111,6 +28113,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -28170,6 +28173,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -28592,6 +28596,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
@ -28649,6 +28654,7 @@
|
|||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
import os
|
||||
from datetime import date
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
|
||||
def _skip_live_prompt_caching_test():
|
||||
|
|
@ -8,3 +10,55 @@ def _skip_live_prompt_caching_test():
|
|||
pytest.skip("Live prompt-caching E2E tests are opt-in")
|
||||
if os.environ.get("CASSETTE_REDIS_URL"):
|
||||
pytest.skip("Live prompt-caching E2E tests cannot run under VCR replay")
|
||||
|
||||
|
||||
|
||||
class TogetherCostEntry(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
litellm_provider: str | None = None
|
||||
mode: str | None = None
|
||||
deprecation_date: str | None = None
|
||||
input_cost_per_token: float | None = None
|
||||
output_cost_per_token: float | None = None
|
||||
supports_function_calling: bool | None = None
|
||||
supports_response_schema: bool | None = None
|
||||
|
||||
|
||||
def cheapest_together_chat_model(
|
||||
*, function_calling: bool = False, response_schema: bool = False
|
||||
) -> str:
|
||||
import litellm
|
||||
|
||||
today = date.today().isoformat()
|
||||
|
||||
def qualifies(name: str, entry: TogetherCostEntry) -> bool:
|
||||
return (
|
||||
name.startswith("together_ai/")
|
||||
and entry.litellm_provider == "together_ai"
|
||||
and entry.mode == "chat"
|
||||
and (entry.deprecation_date is None or entry.deprecation_date > today)
|
||||
and (entry.input_cost_per_token or 0.0) > 0
|
||||
and (entry.output_cost_per_token or 0.0) > 0
|
||||
and (not function_calling or bool(entry.supports_function_calling))
|
||||
and (not response_schema or bool(entry.supports_response_schema))
|
||||
)
|
||||
|
||||
registry: dict[str, TogetherCostEntry] = {
|
||||
name: TogetherCostEntry.model_validate(raw)
|
||||
for name, raw in litellm.model_cost.items()
|
||||
if isinstance(raw, dict) and name.startswith("together_ai/")
|
||||
}
|
||||
candidates = sorted(
|
||||
(name for name, entry in registry.items() if qualifies(name, entry)),
|
||||
key=lambda name: (
|
||||
registry[name].input_cost_per_token or 0.0,
|
||||
registry[name].output_cost_per_token or 0.0,
|
||||
name,
|
||||
),
|
||||
)
|
||||
assert candidates, (
|
||||
"no live together_ai chat model in the cost map satisfies "
|
||||
f"function_calling={function_calling} response_schema={response_schema}"
|
||||
)
|
||||
return candidates[0]
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,11 +1,41 @@
|
|||
# Shared provider-response cache
|
||||
|
||||
`E2E_PROVIDER_CACHE=1` enables automatic response reuse in the live E2E mode. Standard OpenAI and Anthropic model registrations use the provider edge. Existing custom API bases, named credentials, mocked models and realtime WebSocket deployments keep their existing routing. Other provider protocols remain live
|
||||
`E2E_PROVIDER_CACHE=1` enables automatic response reuse in the live E2E mode. Standard OpenAI and Anthropic model registrations use the provider edge, as do Anthropic-on-Bedrock registrations that carry no AWS identity of their own. Existing custom API bases, named credentials, mocked models and realtime WebSocket deployments keep their existing routing. Other provider protocols remain live
|
||||
|
||||
The edge caches complete successful POST responses for `/v1/chat/completions` and `/v1/messages`, including streams. Unsupported endpoints pass through. It matches the method, original URL, effective outbound headers (including authentication and HTTP-library defaults), body presence and exact body bytes using a full keyed digest. It sends the same prepared request used for matching. No prompts, random markers, JSON values or credentials are normalized away. Provider `Set-Cookie` headers are dropped before validation and never recorded: the edge already withholds them from the proxy, and OpenAI responses always carry Cloudflare bot-management cookies
|
||||
The edge caches complete successful POST responses for `/v1/chat/completions`, `/v1/messages`, `/v1/embeddings` and `/v1/responses` on the OpenAI and Anthropic mounts, SSE streams included, and for `/model/{id}/converse` and `/model/{id}/invoke` on a Bedrock mount. Unsupported endpoints pass through. Each endpoint family has its own completeness rule, so a truncated embedding or a Responses run that never reached `response.completed` is not stored
|
||||
|
||||
Bedrock's streaming endpoints, `converse-stream` and `invoke-with-response-stream`, cache too. AWS frames those as binary `vnd.amazon.eventstream` rather than SSE, so botocore's own parser reads the frames and validates both CRCs, and each endpoint is then held to its terminal grammar. That matters more than the endpoint count suggests: the Claude Code compat cells drive the real CLI, which always streams, so streaming is most of the suite's Bedrock traffic
|
||||
|
||||
Two details of that rule are worth knowing before changing it. A ConverseStream ends with `metadata`, not with `messageStop`, and the `metadata` frame is what carries the token usage litellm prices the call from, so the rule requires it: a stream cut between the two still names a stop reason but would replay as a free call. And a dropped connection is invisible to the parser, which yields the frames it did receive and silently discards a trailing partial one, so the body is also checked against the frame lengths it declares. A stream cut one byte short parses clean and has to be caught that way
|
||||
|
||||
## Request identity
|
||||
|
||||
A recording belongs to one test. The key is a keyed digest over the test's node id, the method, the URL, the effective outbound headers (including authentication and HTTP-library defaults), body presence and the body bytes, with one normalization: a 12-hex-digit run, the shape `unique_marker()` mints, is replaced by a placeholder in both the URL and a UTF-8 body. Nothing else is normalized away. No prompts, JSON values or credentials are rewritten, and the rule is the one `fixture_canonical.py` already applies for record/replay, so there is a single definition of what a marker is
|
||||
|
||||
Requests that differ only by their markers therefore share a canonical identity, which is what makes the cache reusable across builds: every e2e test salts its prompt afresh, so an exact-byte key would miss on every call. Within one test, calls that share a canonical identity are still recorded and replayed separately, by a FIFO slot index appended to the key. That matters because a replayed response carries the recorded provider response id, `LiteLLM_SpendLogs.request_id` is that id, and one shared recording answering two calls would collapse two spend rows into one
|
||||
|
||||
Two different tests never share a recording, and a provider call made outside any test (fixtures, session setup) is never cached, because the identity has no test node id to bind to
|
||||
|
||||
A client that varies its own request between runs defeats that identity without breaking any rule, and the Claude Code compat cells did. The CLI sends a device id and a session id in `metadata.user_id`, and its system prompt names both its memory directory and its working directory, adding the branch and recent commits when that directory is a git repository. Driven with a fresh HOME and the checkout as its working directory, every cell sent different bytes every build. The fix belongs in the driver rather than here: `claude_code/cli_driver.py` pins the config directory, the working directory and both identifiers, which is why the cache needs no rule for any of it. Normalizing them instead would have hidden a real defect class, since a rule cannot tell a client's own churn from a value a test means to assert on
|
||||
|
||||
Provider `Set-Cookie` headers are dropped before validation and never recorded: the edge already withholds them from the proxy, and OpenAI responses always carry Cloudflare bot-management cookies
|
||||
|
||||
An eligible miss calls the provider. A complete successful response is stored immediately even if a later test assertion fails. Provider errors, malformed responses, truncated streams and cancelled captures are not stored. Cache reads, writes and lease failures fall through to normal provider behavior; they introduce no provider retry. An already-started response cannot be restarted after a delivery failure
|
||||
|
||||
## Bedrock
|
||||
|
||||
Bedrock could not be mounted before because SigV4 signs the `Host` header, so a rewritten `api_base` failed signature verification at the provider. The edge now re-signs: it drops the proxy's signature headers, signs the upstream request with the run pod's own AWS identity from its EKS Pod Identity association, and forwards that. The signature headers are excluded from the key, since `x-amz-date` is a timestamp and keying on it would make every Bedrock call a permanent miss
|
||||
|
||||
Almost every Bedrock deployment in the suite declares its region as `os.environ/AWS_REGION`, which only the proxy can resolve, and the run pod does not share that environment. A `us.` inference profile fans out across the US regions and is reachable from any of them, so those route to the default mount whatever the proxy resolved. A model that is not cross-region and declares its region that way keeps its direct path rather than being sent to a region it may not exist in.
|
||||
|
||||
Only deployments that carry no AWS identity of their own route to the edge. A deployment with `aws_role_name`, `aws_access_key_id`, an `api_base` or an `aws_bedrock_runtime_endpoint` keeps its direct path, because re-signing it would quietly replace the very credential chain that test exists to prove
|
||||
|
||||
Which models route is an explicit allowlist in `provider_cache_routing.py`, mirroring the runner role's IAM policy, which names its models one by one. That coupling is deliberate: the edge re-signs with the run pod's identity, so a model the role cannot invoke comes back 403 from Bedrock rather than falling back. An unlisted model keeps its direct path and loses only caching, so adding a Bedrock model to the suite can never turn it red. Adding one to the edge is a policy edit in litellm-ops plus a line here
|
||||
|
||||
Vertex and Gemini are not mounted, for different reasons. litellm grafts the default Vertex path onto an `api_base` only when that `api_base` has no path of its own, so a path-prefixed Vertex mount instead becomes `{api_base}:{endpoint}`, dropping project, location and model. Vertex needs a root-mounted edge on its own port, or a change in litellm
|
||||
|
||||
Gemini reaches a path-prefixed mount perfectly well and was mounted for one build, then backed out, because litellm's two Gemini endpoints disagree about what `api_base` means. Chat composes `{api_base}/models/{model}:{endpoint}` and defaults `api_base` to `https://generativelanguage.googleapis.com/v1beta`, so the version has to be inside it. File upload composes `{api_base}/upload/v1beta/files` and defaults to the host root, so the version has to be outside it. One `api_base` cannot satisfy both, and a deployment gives no signal at registration time about which it will be used for, so mounting Gemini turned `TestGeminiFiles::test_gemini_file_upload` red in build 227. Anyone pointing litellm's Gemini provider at an AI gateway or a corporate proxy hits the same thing; it is a litellm bug rather than a cache limitation, and mounting Gemini is one line once it is fixed
|
||||
|
||||
Recordings are shared across workers and builds through dedicated Redis, separate from the candidate's own cache. They expire 86,400 seconds after capture starts, based on Redis time. Reads never extend expiry. There is no scheduled recapture: the next miss calls the provider again. Bounded coordination reduces duplicate concurrent calls, but slow or failed captures may lead to extra live calls after the wait expires
|
||||
|
||||
## Configuration
|
||||
|
|
@ -18,16 +48,18 @@ The trusted runner receives:
|
|||
- `E2E_PROVIDER_CACHE_NAMESPACE`: shared environment namespace, independent of build and candidate revision
|
||||
- `E2E_PROVIDER_CACHE_METRICS_DIR`: optional per-process counter artifact directory
|
||||
|
||||
Do not give cache credentials to candidate deployments. Counter artifacts contain no recorded payloads or credentials. Hits count shared-cache responses; upstream attempts count actual forwards from the edge. Existing application-cache observations still count requests arriving at the edge, including shared-cache hits
|
||||
Do not give cache credentials to candidate deployments. Counter artifacts contain no recorded payloads or credentials. Hits count shared-cache responses; upstream attempts count actual forwards from the edge. A rejection also counts its reason, one of `rejected_cut_short` (the consumer walked away mid-capture), `rejected_error_status` (the provider answered, with an error), `rejected_incomplete` (the body arrived whole with a success status and failed its endpoint's rule) or `rejected_unreachable` (the provider could not be reached at all). A mount whose rejections are nearly all of one kind is a different problem from one whose rejections are nearly all of another, and the flat count cannot tell them apart. Every counter is emitted twice, once as a flat total and once under `mount:{mount}:`, so a hit rate can be read per provider rather than only in aggregate. Existing application-cache observations still count requests arriving at the edge, including shared-cache hits
|
||||
|
||||
Tests that require real provider timing, limits or state use `@pytest.mark.provider_live`. The marker keeps newly registered models on live routes without weakening their assertions. The provider prompt-caching tests carry it because a replayed priming response reports cache creation rather than a cache read. Ordinary assertion failures still fail E2E. The shared cache does not modify provider response IDs or make the proxy aware of replay
|
||||
Tests that require real provider timing, limits or state use `@pytest.mark.provider_live`. The marker keeps newly registered models on live routes without weakening their assertions. The provider prompt-caching tests carry it because a replayed priming response reports cache creation rather than a cache read.
|
||||
|
||||
One more class needs it, and it is the cost of normalizing the marker. A test that mints a fresh marker, sends it, and then asserts the provider's answer contains that exact value is asserting on the marker rather than using it as a salt. The key treats two such requests as the same identity, so a stale recording matches and answers with the marker from the run that recorded it. `TestOpenAIMessagesToolContinuation` is the one in the suite today: it sends a freshly minted receipt through a tool result and asserts the model echoes it back verbatim. If you add a test that asserts a provider echoed your own unique value, it belongs on the live path. Ordinary assertion failures still fail E2E. The shared cache does not modify provider response IDs or make the proxy aware of replay
|
||||
|
||||
## Recorded response semantics
|
||||
|
||||
Replay preserves the original response ID, usage and end-to-end headers. The proxy can therefore deduplicate repeated provider IDs when storing spend-log rows, just as it does when a live upstream returns the same ID twice. One spend-log row per invocation is not guaranteed for identical recorded responses. Existing spend reconciliation requests use distinct prompt markers and retain their distinct-ID and row-count assertions; accounting tests are not automatically excluded from caching
|
||||
Replay preserves the original response ID, usage and end-to-end headers. The proxy can therefore deduplicate repeated provider IDs when storing spend-log rows, just as it does when a live upstream returns the same ID twice. One spend-log row per invocation is not guaranteed for identical recorded responses. Spend reconciliation keeps its distinct-ID and row-count assertions: its prompts differ by an index as well as a marker, so they stay distinct once markers are normalized, and calls that are canonically equal within one test take separate FIFO slots and separate recordings anyway. Accounting tests are not automatically excluded from caching
|
||||
|
||||
Provider remaining-quota headers describe the captured response. Metrics derived from them are historical on a cache hit, not a measurement of current provider capacity. Gateway-generated API-key quota headers are a separate contract. A test of fresh provider quota or timing must use the live-provider policy; replay can still exercise how the proxy processes the recorded headers
|
||||
|
||||
## Qualification
|
||||
|
||||
`tests/code_coverage_tests/test_provider_cache.py` exercises local HTTP providers and disposable real Redis. CI runs these checks with the existing provider-edge and replay harness tests. These component checks do not establish Buildkite deployment, full-suite cross-build reuse or a genuine 24-hour expiry observation; those require separate runtime evidence
|
||||
`tests/code_coverage_tests/test_provider_cache.py` exercises local HTTP providers and disposable real Redis, including the marker-canonical key, the FIFO slot index, per-test isolation, SigV4 re-signing against a local upstream, and each endpoint's completeness rule. CI runs these checks with the existing provider-edge and replay harness tests. These component checks do not establish Buildkite deployment, full-suite cross-build reuse or a genuine 24-hour expiry observation; those require separate runtime evidence
|
||||
|
|
|
|||
|
|
@ -0,0 +1,162 @@
|
|||
"""The CLI must send the same request bytes from one build to the next.
|
||||
|
||||
Markerless harness test: it drives the real `claude` binary against a local
|
||||
stub instead of a proxy, so it carries no `e2e` marker. The binary is a
|
||||
prerequisite of this whole suite, so a missing one is a failure rather than a
|
||||
skip.
|
||||
|
||||
Two builds differ in ways the driver does not control: a fresh pod, so no CLI
|
||||
state survives, and a different candidate checked out at a different commit.
|
||||
Both used to reach the request body, through the memory path the system prompt
|
||||
names and through the git block the CLI adds for its working directory, so the
|
||||
shared provider cache missed on every Claude Code cell. This replays those two
|
||||
differences across a pair of invocations and holds the bytes equal.
|
||||
|
||||
A pinned session id is what makes the second test necessary. The matrix runs
|
||||
its cells across xdist workers, and the CLI refuses to start a session id that
|
||||
another live process already holds, so pinning one without also opting out of
|
||||
session persistence turns most of a parallel run red.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import threading
|
||||
from collections import Counter
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from typing import List, Tuple
|
||||
|
||||
import pytest
|
||||
|
||||
from claude_code.cli_driver import _FIXED_CLI_USER_ID, _seed_cli_identity, _stable_cli_state, run_claude
|
||||
from claude_code.rate_limiter import RateLimiter
|
||||
|
||||
pytestmark = pytest.mark.cli_determinism
|
||||
|
||||
_STUB_REPLY = {
|
||||
"id": "msg_stub",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-haiku-4-5",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 10, "output_tokens": 2},
|
||||
}
|
||||
|
||||
|
||||
def _make_repo(root: Path, subject: str) -> Path:
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
identity = {"NAME": "t", "EMAIL": "t@e2e"}
|
||||
env = dict(
|
||||
os.environ,
|
||||
**{f"GIT_{role}_{key}": value for role in ("AUTHOR", "COMMITTER") for key, value in identity.items()},
|
||||
)
|
||||
(root / "file.txt").write_text(subject, encoding="utf-8")
|
||||
for args in (["init", "-q"], ["add", "."], ["commit", "-q", "-m", subject]):
|
||||
subprocess.run(["git", *args], cwd=root, env=env, check=True, capture_output=True)
|
||||
return root
|
||||
|
||||
|
||||
@pytest.fixture(name="captured")
|
||||
def _captured() -> Tuple[str, List[bytes]]:
|
||||
bodies: List[bytes] = []
|
||||
lock = threading.Lock()
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_POST(self) -> None:
|
||||
raw = self.rfile.read(int(self.headers.get("content-length") or 0))
|
||||
if "count_tokens" not in self.path:
|
||||
with lock:
|
||||
bodies.append(raw)
|
||||
payload = json.dumps({"input_tokens": 10} if "count_tokens" in self.path else _STUB_REPLY).encode()
|
||||
self.send_response(200)
|
||||
self.send_header("content-type", "application/json")
|
||||
self.send_header("content-length", str(len(payload)))
|
||||
self.end_headers()
|
||||
self.wfile.write(payload)
|
||||
|
||||
def log_message(self, *_args: object) -> None:
|
||||
return
|
||||
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
|
||||
threading.Thread(target=server.serve_forever, daemon=True).start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{server.server_address[1]}", bodies
|
||||
finally:
|
||||
server.shutdown()
|
||||
|
||||
|
||||
def test_two_builds_send_the_same_request_bytes(captured: Tuple[str, List[bytes]], tmp_path: Path) -> None:
|
||||
base_url, bodies = captured
|
||||
limiter = RateLimiter(state_dir=tmp_path / "limiter")
|
||||
checkouts = (_make_repo(tmp_path / "build-1", "first"), _make_repo(tmp_path / "build-2", "second"))
|
||||
origin = Path.cwd()
|
||||
|
||||
sent = []
|
||||
for checkout in checkouts:
|
||||
shutil.rmtree(Path(_stable_cli_state()[0]).parent, ignore_errors=True)
|
||||
os.chdir(checkout)
|
||||
try:
|
||||
before = len(bodies)
|
||||
run_claude(
|
||||
prompt="say ok",
|
||||
model="claude-haiku-4-5",
|
||||
base_url=base_url,
|
||||
api_key="stub",
|
||||
extra_env={"CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC": "1"},
|
||||
rate_limiter=limiter,
|
||||
)
|
||||
sent.append(bodies[before:])
|
||||
finally:
|
||||
os.chdir(origin)
|
||||
|
||||
assert sent[0], "the CLI sent no request to the stub, so there is nothing to compare"
|
||||
assert sent[0] == sent[1]
|
||||
|
||||
|
||||
def test_concurrent_cells_do_not_collide_on_the_pinned_session(
|
||||
captured: Tuple[str, List[bytes]], tmp_path: Path
|
||||
) -> None:
|
||||
base_url, bodies = captured
|
||||
limiter = RateLimiter(state_dir=tmp_path / "limiter")
|
||||
|
||||
def one(_index: int) -> int:
|
||||
return run_claude(
|
||||
prompt="say ok",
|
||||
model="claude-haiku-4-5",
|
||||
base_url=base_url,
|
||||
api_key="stub",
|
||||
extra_env={"CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC": "1"},
|
||||
rate_limiter=limiter,
|
||||
).exit_code
|
||||
|
||||
with ThreadPoolExecutor(max_workers=4) as pool:
|
||||
codes = list(pool.map(one, range(4)))
|
||||
|
||||
assert codes == [0, 0, 0, 0]
|
||||
assert bodies, "the CLI sent no request to the stub, so there is nothing to compare"
|
||||
assert set(Counter(bodies).values()) == {4}
|
||||
|
||||
|
||||
def test_seeding_the_device_id_survives_threads_racing_on_the_same_directory(tmp_path: Path) -> None:
|
||||
"""`run_claude_models_parallel` drives several models from one process, so the
|
||||
seed's staged file has to be unique per thread and not merely per process."""
|
||||
config_dir = tmp_path / "config"
|
||||
config_dir.mkdir()
|
||||
seeded = config_dir / ".claude.json"
|
||||
|
||||
for _round in range(20):
|
||||
seeded.unlink(missing_ok=True)
|
||||
with ThreadPoolExecutor(max_workers=16) as pool:
|
||||
for outcome in [pool.submit(_seed_cli_identity, str(config_dir)) for _ in range(16)]:
|
||||
outcome.result()
|
||||
|
||||
assert json.loads(seeded.read_text(encoding="utf-8"))["userID"] == _FIXED_CLI_USER_ID
|
||||
assert sorted(entry.name for entry in config_dir.iterdir()) == [".claude.json"]
|
||||
|
|
@ -132,6 +132,62 @@ def _make_isolated_home() -> str:
|
|||
return tempfile.mkdtemp(prefix="claude-cli-home-")
|
||||
|
||||
|
||||
_FIXED_CLI_USER_ID = "0" * 64
|
||||
_FIXED_CLI_SESSION_ID = "00000000-0000-4000-8000-000000000000"
|
||||
|
||||
|
||||
def _seed_cli_identity(config_dir: str) -> None:
|
||||
"""Pin the device id the CLI would otherwise mint per config directory.
|
||||
|
||||
It mints 32 random bytes on first run, writes them to `.claude.json` as
|
||||
`userID`, and sends them in `metadata.user_id` forever after, so the value
|
||||
is stable for exactly as long as that file lives. Pinning it, and the
|
||||
session id passed beside it, costs nothing: both feed abuse detection
|
||||
rather than quota, caching or continuity.
|
||||
|
||||
The staged name has to be unique per *thread*, not per process:
|
||||
`run_claude_models_parallel` drives several models from one process, so a
|
||||
pid-suffixed name lets one thread rename the file another is still
|
||||
writing, and the loser dies on a missing path."""
|
||||
path = os.path.join(config_dir, ".claude.json")
|
||||
try:
|
||||
with open(path, encoding="utf-8") as handle:
|
||||
if json.load(handle).get("userID") == _FIXED_CLI_USER_ID:
|
||||
return
|
||||
except (OSError, ValueError):
|
||||
pass
|
||||
handle_fd, staged = tempfile.mkstemp(dir=config_dir, prefix=".claude.json.")
|
||||
with os.fdopen(handle_fd, "w", encoding="utf-8") as handle:
|
||||
json.dump({"userID": _FIXED_CLI_USER_ID}, handle)
|
||||
os.replace(staged, path)
|
||||
|
||||
|
||||
def _stable_cli_state() -> Tuple[str, str]:
|
||||
"""Config directory and working directory for the CLI, at fixed paths.
|
||||
|
||||
Both reach the request body. The memory directory the system prompt
|
||||
names is `$CLAUDE_CONFIG_DIR/projects/<cwd slug>/memory`, and a working
|
||||
directory inside a git repository also contributes its branch and recent
|
||||
commits. So a per-invocation config directory rewrites every body, and
|
||||
inheriting the checkout rewrites every body once per candidate, which is
|
||||
why the shared provider cache could never serve a Claude Code cell.
|
||||
Pinning both makes the bodies repeatable across builds.
|
||||
|
||||
This narrows what survives rather than widening it: HOME stays fresh and
|
||||
empty per invocation, so the isolation `_make_isolated_home` describes is
|
||||
unchanged, and the CLI's own state no longer outlives the pod either. The
|
||||
working directory is deliberately not the checkout, so a model-directed
|
||||
`Read` sees an empty directory instead of the repository.
|
||||
"""
|
||||
root = os.path.join(tempfile.gettempdir(), f"litellm-e2e-claude-{os.getuid()}")
|
||||
config_dir = os.path.join(root, "config")
|
||||
workspace = os.path.join(root, "workspace")
|
||||
for path in (root, config_dir, workspace):
|
||||
os.makedirs(path, mode=0o700, exist_ok=True)
|
||||
_seed_cli_identity(config_dir)
|
||||
return config_dir, workspace
|
||||
|
||||
|
||||
class ClaudeCLIError(RuntimeError):
|
||||
"""Raised when the `claude` CLI cannot be invoked or returns a fatal error."""
|
||||
|
||||
|
|
@ -222,6 +278,9 @@ def run_claude(
|
|||
"--verbose",
|
||||
"--model",
|
||||
model,
|
||||
"--session-id",
|
||||
_FIXED_CLI_SESSION_ID,
|
||||
"--no-session-persistence",
|
||||
]
|
||||
if extra_args:
|
||||
cmd.extend(extra_args)
|
||||
|
|
@ -244,6 +303,8 @@ def run_claude(
|
|||
# regardless of how the subprocess exits.
|
||||
isolated_home = _make_isolated_home()
|
||||
env["HOME"] = isolated_home
|
||||
config_dir, workspace = _stable_cli_state()
|
||||
env["CLAUDE_CONFIG_DIR"] = config_dir
|
||||
if extra_env:
|
||||
env.update(extra_env)
|
||||
|
||||
|
|
@ -262,6 +323,7 @@ def run_claude(
|
|||
completed = run_fn(
|
||||
cmd,
|
||||
env=env,
|
||||
cwd=workspace,
|
||||
input=stdin_input,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from typing import Final
|
|||
import pytest
|
||||
import requests
|
||||
from e2e_config import (
|
||||
CLI_DETERMINISM_OPT_IN_ENV,
|
||||
CONTROL_PLANE_BASE_URL,
|
||||
FIXTURE_DIR,
|
||||
FIXTURE_MODE_RAW,
|
||||
|
|
@ -53,6 +54,7 @@ OPT_IN_MARKERS: Final = MappingProxyType(
|
|||
"managed_files": MANAGED_FILES_OPT_IN_ENV,
|
||||
"prompt_caching_stack": PROMPT_CACHING_OPT_IN_ENV,
|
||||
"redis_chaos": REDIS_CHAOS_OPT_IN_ENV,
|
||||
"cli_determinism": CLI_DETERMINISM_OPT_IN_ENV,
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -85,7 +87,11 @@ def jwt_identity(idp: Keycloak, resources: ResourceManager, proxy: ProxyClient)
|
|||
|
||||
|
||||
def pytest_configure(config: pytest.Config) -> None:
|
||||
config.addinivalue_line("markers", "provider_live: requires actual provider timing, limits or state; bypass shared cache")
|
||||
config.addinivalue_line(
|
||||
"markers",
|
||||
"provider_live: requires actual provider timing, limits, state, or a response that echoes this"
|
||||
" run's own unique value; bypass shared cache",
|
||||
)
|
||||
config.addinivalue_line(
|
||||
"markers",
|
||||
"e2e: live test that requires a running proxy and real provider keys",
|
||||
|
|
@ -116,6 +122,11 @@ def pytest_configure(config: pytest.Config) -> None:
|
|||
"prompt_caching_stack: needs a proxy running with router_settings.optional_pre_call_checks including "
|
||||
"prompt_caching; deselected unless E2E_PROMPT_CACHING_STACK is set",
|
||||
)
|
||||
config.addinivalue_line(
|
||||
"markers",
|
||||
"cli_determinism: drives the real claude CLI for several seconds, which widens the window in which "
|
||||
"another test's in-flight upstream call is attributed to it; deselected unless E2E_CLI_DETERMINISM is set",
|
||||
)
|
||||
config.addinivalue_line(
|
||||
"markers",
|
||||
"redis_chaos: load test that pauses the proxy's Redis outright mid-run; needs a proxy booted from "
|
||||
|
|
|
|||
|
|
@ -30,6 +30,9 @@
|
|||
- {id: mgmt.key.health.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:4292", rationale: "Key health endpoint"}
|
||||
- {id: mgmt.key.bulk_update.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:2677", rationale: "Batch key updates"}
|
||||
- {id: mgmt.team.update.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py:1582", rationale: "Metadata/budget updates persist"}
|
||||
- {id: mgmt.team.update.team_admin_forbidden_until_enabled, module: mgmt, tier: P0, surface: api, assertions: [team_admin_forbidden_until_enabled], source: "team_admin_field_permissions.py:156", rationale: "With no team admin editable fields enabled, a team admin's /team/update is 403 and /team/info reports editing disabled"}
|
||||
- {id: mgmt.team.update.team_admin_limited_to_enabled_fields, module: mgmt, tier: P0, surface: api, assertions: [team_admin_limited_to_enabled_fields], source: "team_admin_field_permissions.py:156", rationale: "A team admin may change only the enabled fields; a request that also changes any other field is 403 and writes nothing"}
|
||||
- {id: mgmt.team.update.team_admin_resend_keeps_budget_reset, module: mgmt, tier: P1, surface: api, assertions: [team_admin_resend_keeps_budget_reset], source: "team_admin_field_permissions.py:147", fail_before_fix: proven, rationale: "A team admin resending unchanged budget settings with an enabled field must not push the team's budget reset times back"}
|
||||
- {id: mgmt.team.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py:1750", rationale: "Deletion prevents key access"}
|
||||
- {id: mgmt.team.block.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py", rationale: "Block suspends all members"}
|
||||
- {id: mgmt.team.info.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "team_endpoints.py:2244", rationale: "Metadata+members+budgets"}
|
||||
|
|
|
|||
|
|
@ -145,6 +145,7 @@ WEEKLY_ANOMALY_OPT_IN_ENV = "E2E_WEEKLY_ANOMALY"
|
|||
MANAGED_FILES_OPT_IN_ENV = "E2E_MANAGED_FILES_STACK"
|
||||
PROMPT_CACHING_OPT_IN_ENV = "E2E_PROMPT_CACHING_STACK"
|
||||
REDIS_CHAOS_OPT_IN_ENV = "E2E_REDIS_CHAOS"
|
||||
CLI_DETERMINISM_OPT_IN_ENV = "E2E_CLI_DETERMINISM"
|
||||
ANOMALY_SESSIONS = int(os.environ.get("E2E_ANOMALY_SESSIONS", "6"))
|
||||
ANOMALY_TURNS_PER_SESSION = int(os.environ.get("E2E_ANOMALY_TURNS_PER_SESSION", "6"))
|
||||
ANOMALY_TURN_ATTEMPTS = int(os.environ.get("E2E_ANOMALY_TURN_ATTEMPTS", "3"))
|
||||
|
|
|
|||
|
|
@ -51,6 +51,9 @@ SECRET_FIELD_SUFFIXES: Final[tuple[str, ...]] = (
|
|||
)
|
||||
SECRET_PLACEHOLDER: Final = "<secret>"
|
||||
|
||||
MARKER_PATTERN: Final = re.compile(r"(?<![0-9a-fA-F])[0-9a-f]{12}(?![0-9a-fA-F])")
|
||||
MARKER_PLACEHOLDER: Final = "<marker>"
|
||||
|
||||
PLACEHOLDER_RULES: Final[tuple[tuple[re.Pattern[str], str], ...]] = (
|
||||
(re.compile(r"(?<![0-9a-fA-F])[0-9a-f]{64}(?![0-9a-fA-F])"), "<sha256>"),
|
||||
(
|
||||
|
|
@ -67,7 +70,7 @@ PLACEHOLDER_RULES: Final[tuple[tuple[re.Pattern[str], str], ...]] = (
|
|||
re.compile(r"\b(?:chatcmpl|msgbatch|msg|resp|batch|call|req|ftjob|gen|file)[-_][A-Za-z0-9]{8,}\b"),
|
||||
"<id>",
|
||||
),
|
||||
(re.compile(r"(?<![0-9a-fA-F])[0-9a-f]{12}(?![0-9a-fA-F])"), "<marker>"),
|
||||
(MARKER_PATTERN, MARKER_PLACEHOLDER),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -372,6 +372,7 @@ def _request_tool(
|
|||
|
||||
|
||||
class TestOpenAIMessagesToolContinuation:
|
||||
@pytest.mark.provider_live
|
||||
@pytest.mark.parametrize("stream", [True, False], ids=["stream", "nonstream"])
|
||||
def test_required_tool_arguments_and_correlated_result(
|
||||
self, endpoints_client: EndpointsClient, resources: ResourceManager, stream: bool
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Live e2e: the /team/* management routes' block, membership, and admin-only
|
||||
contract.
|
||||
contract, plus the team settings a team admin may change on /team/update once a
|
||||
proxy admin enables them under Settings > UI > Team admin editable fields.
|
||||
|
||||
Each test creates its team/user/key resources under unique names (deleted on
|
||||
teardown) and asserts both halves of the contract: the recorded state (the info
|
||||
|
|
@ -8,21 +9,25 @@ Team writes reach the read path once their db/cache entry propagates, so the
|
|||
read-backs poll to a deadline instead of asserting once.
|
||||
|
||||
Everything the shared harness does not already model lives here: the local
|
||||
request/response models for /team/block, /team/member_update, and the
|
||||
/team/info fields (blocked flag and per-member budget) these tests assert on.
|
||||
request/response models for /team/block, /team/member_update, the partial
|
||||
/team/update, the UI settings allow-list, and the /team/info fields (blocked
|
||||
flag, limits, budgets, per-member budget, the caller's edit access) these tests
|
||||
assert on.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from typing import Literal
|
||||
from collections.abc import Callable, Generator
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import NoBody, StreamingResponse, unwrap
|
||||
from e2e_config import settle_propagation, unique_marker
|
||||
from e2e_http import NoBody, PartialBody, StreamingResponse, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from management_client import ManagementClient
|
||||
from models import (
|
||||
|
|
@ -39,6 +44,8 @@ pytestmark = pytest.mark.e2e
|
|||
|
||||
TeamRole = Literal["admin", "user"]
|
||||
|
||||
_TEAM_TPM_LIMIT: Final = 1000
|
||||
|
||||
|
||||
class TeamBlockBody(BaseModel):
|
||||
team_id: str
|
||||
|
|
@ -66,11 +73,37 @@ class TeamMembership(BaseModel):
|
|||
litellm_budget_table: MemberBudgetTable | None = None
|
||||
|
||||
|
||||
class TeamInfoData(BaseModel):
|
||||
class CallerEditAccess(BaseModel):
|
||||
kind: Literal["unrestricted", "team_admin", "team_admin_disabled", "none"]
|
||||
editable_fields: list[str] = []
|
||||
|
||||
|
||||
class BudgetWindow(BaseModel):
|
||||
budget_duration: str
|
||||
max_budget: float
|
||||
reset_at: str | None = None
|
||||
|
||||
|
||||
class TeamCustomMetadata(BaseModel):
|
||||
cost_center: str | None = None
|
||||
|
||||
|
||||
class TeamSettings(BaseModel):
|
||||
team_alias: str | None = None
|
||||
models: list[str] = []
|
||||
tpm_limit: int | None = None
|
||||
rpm_limit: int | None = None
|
||||
max_budget: float | None = None
|
||||
budget_duration: str | None = None
|
||||
budget_limits: list[BudgetWindow] | None = None
|
||||
metadata: TeamCustomMetadata | None = None
|
||||
|
||||
|
||||
class TeamInfoData(TeamSettings):
|
||||
blocked: bool | None = None
|
||||
members_with_roles: list[MemberRoleEntry] = []
|
||||
budget_reset_at: datetime | None = None
|
||||
caller_edit_access: CallerEditAccess | None = None
|
||||
|
||||
|
||||
class TeamInfoRead(BaseModel):
|
||||
|
|
@ -79,6 +112,27 @@ class TeamInfoRead(BaseModel):
|
|||
team_memberships: list[TeamMembership] = []
|
||||
|
||||
|
||||
class TeamWithAdminNewBody(TeamNewBody):
|
||||
tpm_limit: int
|
||||
members_with_roles: list[TeamMemberEntry]
|
||||
|
||||
|
||||
class TeamSettingsChange(PartialBody, TeamSettings):
|
||||
pass
|
||||
|
||||
|
||||
class TeamSettingsUpdate(TeamSettingsChange):
|
||||
team_id: str
|
||||
|
||||
|
||||
class TeamAdminEditableFields(BaseModel):
|
||||
team_admin_editable_team_fields: list[str] = []
|
||||
|
||||
|
||||
class UiSettingsRead(BaseModel):
|
||||
values: TeamAdminEditableFields
|
||||
|
||||
|
||||
def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T:
|
||||
deadline = time.monotonic() + client.proxy.poll_timeout
|
||||
while time.monotonic() < deadline:
|
||||
|
|
@ -107,17 +161,27 @@ def _generate_key(client: ManagementClient, resources: ResourceManager, body: Ke
|
|||
return key
|
||||
|
||||
|
||||
def _read_team(client: ManagementClient, team_id: str) -> TeamInfoRead:
|
||||
def _read_team(client: ManagementClient, team_id: str, caller_key: str | None = None) -> TeamInfoRead:
|
||||
return unwrap(
|
||||
client.proxy.transport.get(
|
||||
"/team/info",
|
||||
headers=client.proxy.transport.master,
|
||||
headers=client.proxy.transport.master if caller_key is None else client.proxy.transport.bearer(caller_key),
|
||||
params=TeamInfoParams(team_id=team_id),
|
||||
response_type=TeamInfoRead,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _poll_team(
|
||||
client: ManagementClient, team_id: str, ready: Callable[[TeamInfoData], bool], failure: str
|
||||
) -> TeamInfoData:
|
||||
def read() -> TeamInfoData | None:
|
||||
info = _read_team(client, team_id).team_info
|
||||
return info if ready(info) else None
|
||||
|
||||
return _poll(client, read, failure)
|
||||
|
||||
|
||||
def _set_blocked(client: ManagementClient, team_id: str, *, blocked: bool) -> None:
|
||||
_ = unwrap(
|
||||
client.proxy.transport.post(
|
||||
|
|
@ -301,3 +365,218 @@ class TestTeamManagementRoutes:
|
|||
client.add_team_member(team_id, member_id)
|
||||
member_key = _generate_key(client, resources, KeyGenerateBody(user_id=member_id, team_id=team_id))
|
||||
return member_id, other_id, member_key, team_id
|
||||
|
||||
|
||||
def _team_admin_editable_fields(client: ManagementClient) -> list[str]:
|
||||
return unwrap(
|
||||
client.proxy.transport.get(
|
||||
"/get/ui_settings",
|
||||
headers=client.proxy.transport.master,
|
||||
params=NoBody(),
|
||||
response_type=UiSettingsRead,
|
||||
)
|
||||
).values.team_admin_editable_team_fields
|
||||
|
||||
|
||||
def _set_team_admin_editable_fields(client: ManagementClient, fields: list[str]) -> None:
|
||||
_ = unwrap(
|
||||
client.proxy.transport.patch(
|
||||
"/update/ui_settings",
|
||||
headers=client.proxy.transport.master,
|
||||
json=TeamAdminEditableFields(team_admin_editable_team_fields=fields),
|
||||
response_type=NoBody,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _team_admins_may_edit(client: ManagementClient, fields: list[str]) -> Generator[None]:
|
||||
"""The allow-list is proxy-wide, so restore whatever was there. Other replicas pick a change up on their
|
||||
config reload, which the wait covers before any team admin call lands on one of them."""
|
||||
original = _team_admin_editable_fields(client)
|
||||
_set_team_admin_editable_fields(client, fields)
|
||||
settle_propagation(time.monotonic())
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_set_team_admin_editable_fields(client, original)
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def no_team_admin_editable_fields(client: ManagementClient) -> Generator[None]:
|
||||
with _team_admins_may_edit(client, []):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def tpm_limit_editable_by_team_admins(client: ManagementClient) -> Generator[None]:
|
||||
with _team_admins_may_edit(client, ["tpm_limit"]):
|
||||
yield
|
||||
|
||||
|
||||
def _team_with_admin(client: ManagementClient, resources: ResourceManager) -> tuple[str, str]:
|
||||
"""A team with a tpm_limit, and the key of a user who is an admin of that team."""
|
||||
admin_id = _create_user(client, resources, f"e2e-team-admin-{unique_marker()}@example.com")
|
||||
team_id = client.create_team(
|
||||
TeamWithAdminNewBody(
|
||||
team_alias=f"e2e-team-admin-{unique_marker()}",
|
||||
tpm_limit=_TEAM_TPM_LIMIT,
|
||||
members_with_roles=[TeamMemberEntry(role="admin", user_id=admin_id)],
|
||||
)
|
||||
)
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
return team_id, _generate_key(client, resources, KeyGenerateBody(user_id=admin_id))
|
||||
|
||||
|
||||
def _update_team_as(client: ManagementClient, caller_key: str, body: TeamSettingsUpdate) -> StreamingResponse:
|
||||
return client.proxy.transport.send("/team/update", headers=client.proxy.transport.bearer(caller_key), json=body)
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("no_team_admin_editable_fields")
|
||||
class TestTeamAdminWithNoEditableFields:
|
||||
"""No proxy admin has enabled a team field for team admins, which is how every proxy starts."""
|
||||
|
||||
@pytest.mark.covers("mgmt.team.update.team_admin_forbidden_until_enabled")
|
||||
def test_team_admin_cannot_change_any_team_setting(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
team_id, admin_key = _team_with_admin(client, resources)
|
||||
access = _read_team(client, team_id, admin_key).team_info.caller_edit_access
|
||||
assert access == CallerEditAccess(kind="team_admin_disabled"), (
|
||||
f"/team/info should tell the team admin that editing is disabled, got {access}"
|
||||
)
|
||||
|
||||
outcome = _update_team_as(client, admin_key, TeamSettingsUpdate(team_id=team_id, tpm_limit=5000))
|
||||
|
||||
assert outcome.status_code == 403, (
|
||||
f"/team/update by a team admin must be 403 while nothing is enabled, got {outcome.status_code}: "
|
||||
f"{outcome.body[:300]}"
|
||||
)
|
||||
assert "cannot edit team settings" in outcome.body, f"403 body should say why, got: {outcome.body[:300]}"
|
||||
tpm_limit = _read_team(client, team_id).team_info.tpm_limit
|
||||
assert tpm_limit == _TEAM_TPM_LIMIT, f"the refused update still changed tpm_limit to {tpm_limit}"
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("tpm_limit_editable_by_team_admins")
|
||||
class TestTeamAdminWithTpmLimitEnabled:
|
||||
"""A proxy admin has enabled tpm_limit, so a team admin may change that setting and no other."""
|
||||
|
||||
@pytest.mark.covers("mgmt.team.update.team_admin_limited_to_enabled_fields")
|
||||
def test_team_admin_saves_the_settings_form_with_a_new_tpm_limit(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
team_id, admin_key = _team_with_admin(client, resources)
|
||||
access = _read_team(client, team_id, admin_key).team_info.caller_edit_access
|
||||
assert access == CallerEditAccess(kind="team_admin", editable_fields=["tpm_limit"]), (
|
||||
f"/team/info should list tpm_limit as the team admin's only editable field, got {access}"
|
||||
)
|
||||
before = _read_team(client, team_id).team_info
|
||||
|
||||
outcome = _update_team_as(
|
||||
client,
|
||||
admin_key,
|
||||
TeamSettingsUpdate(team_id=team_id, team_alias=before.team_alias, models=before.models, tpm_limit=5000),
|
||||
)
|
||||
|
||||
assert outcome.status_code == 200, (
|
||||
f"a team admin resending the form with only tpm_limit changed must succeed, got {outcome.status_code}: "
|
||||
f"{outcome.body[:300]}"
|
||||
)
|
||||
after = _poll_team(
|
||||
client, team_id, lambda info: info.tpm_limit == 5000, "/team/info never reflected tpm_limit=5000"
|
||||
)
|
||||
assert after.model_copy(update={"tpm_limit": _TEAM_TPM_LIMIT}) == before, (
|
||||
f"the update changed more than tpm_limit: before {before}, after {after}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("mgmt.team.update.team_admin_limited_to_enabled_fields")
|
||||
@pytest.mark.parametrize(
|
||||
"change",
|
||||
[
|
||||
pytest.param(TeamSettingsChange(rpm_limit=10), id="rpm_limit"),
|
||||
pytest.param(TeamSettingsChange(max_budget=0.5), id="max_budget"),
|
||||
pytest.param(TeamSettingsChange(team_alias="renamed-by-team-admin"), id="team_alias"),
|
||||
pytest.param(TeamSettingsChange(models=["gemini-2.5-flash"]), id="models"),
|
||||
pytest.param(TeamSettingsChange(budget_duration="1d"), id="budget_duration"),
|
||||
pytest.param(TeamSettingsChange(metadata=TeamCustomMetadata(cost_center="team-admin")), id="metadata"),
|
||||
],
|
||||
)
|
||||
def test_team_admin_cannot_change_a_setting_that_is_not_enabled(
|
||||
self, client: ManagementClient, resources: ResourceManager, change: TeamSettingsChange
|
||||
) -> None:
|
||||
(field,) = change.model_fields_set
|
||||
team_id, admin_key = _team_with_admin(client, resources)
|
||||
before = _read_team(client, team_id).team_info
|
||||
|
||||
outcome = _update_team_as(
|
||||
client,
|
||||
admin_key,
|
||||
TeamSettingsUpdate.model_validate(
|
||||
{**change.model_dump(exclude_unset=True), "team_id": team_id, "tpm_limit": 5000}
|
||||
),
|
||||
)
|
||||
|
||||
assert outcome.status_code == 403, (
|
||||
f"a team admin changing {field} must be 403, got {outcome.status_code}: {outcome.body[:300]}"
|
||||
)
|
||||
assert f"'{field}'" in outcome.body, f"403 body should name {field}, got: {outcome.body[:300]}"
|
||||
after = _read_team(client, team_id).team_info
|
||||
assert after == before, (
|
||||
f"the refused update still wrote to the team, the enabled tpm_limit included: before {before}, "
|
||||
f"after {after}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("mgmt.team.update.team_admin_resend_keeps_budget_reset")
|
||||
def test_team_admin_resending_the_budget_settings_keeps_the_next_budget_reset(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
"""A 120s budget resets at the start of the minute after next. Resending it once the next minute has
|
||||
started would push that reset a minute later, while the stored reset is still a minute out, so the
|
||||
proxy's budget reset job cannot be what moves it."""
|
||||
team_id, admin_key = _team_with_admin(client, resources)
|
||||
_ = unwrap(
|
||||
client.proxy.transport.post(
|
||||
"/team/update",
|
||||
headers=client.proxy.transport.master,
|
||||
json=TeamSettingsUpdate(
|
||||
team_id=team_id,
|
||||
budget_duration="120s",
|
||||
budget_limits=[BudgetWindow(budget_duration="120s", max_budget=5.0)],
|
||||
),
|
||||
response_type=NoBody,
|
||||
)
|
||||
)
|
||||
budgeted = _poll_team(
|
||||
client,
|
||||
team_id,
|
||||
lambda info: info.budget_reset_at is not None and bool(info.budget_limits),
|
||||
"/team/info never reflected the 120s budget the proxy admin set",
|
||||
)
|
||||
assert budgeted.budget_reset_at is not None
|
||||
next_minute = budgeted.budget_reset_at - timedelta(seconds=58)
|
||||
time.sleep(max(0.0, (next_minute - datetime.now(UTC)).total_seconds()))
|
||||
|
||||
outcome = _update_team_as(
|
||||
client,
|
||||
admin_key,
|
||||
TeamSettingsUpdate(
|
||||
team_id=team_id,
|
||||
tpm_limit=5000,
|
||||
budget_duration=budgeted.budget_duration,
|
||||
budget_limits=budgeted.budget_limits,
|
||||
),
|
||||
)
|
||||
|
||||
assert outcome.status_code == 200, (
|
||||
f"resending unchanged budget settings with a new tpm_limit must succeed, got {outcome.status_code}: "
|
||||
f"{outcome.body[:300]}"
|
||||
)
|
||||
after = _poll_team(
|
||||
client, team_id, lambda info: info.tpm_limit == 5000, "/team/info never reflected tpm_limit=5000"
|
||||
)
|
||||
assert after.budget_reset_at == budgeted.budget_reset_at, (
|
||||
f"the team admin pushed the budget reset from {budgeted.budget_reset_at} to {after.budget_reset_at}"
|
||||
)
|
||||
assert after.budget_limits == budgeted.budget_limits, (
|
||||
f"the team admin pushed the budget window resets from {budgeted.budget_limits} to {after.budget_limits}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -951,6 +951,7 @@ class LiteLLMParamsBody(BaseModel):
|
|||
aws_access_key_id: str | None = None
|
||||
aws_secret_access_key: str | None = None
|
||||
aws_region_name: str | None = None
|
||||
aws_bedrock_runtime_endpoint: str | None = None
|
||||
vertex_project: str | None = None
|
||||
vertex_location: str | None = None
|
||||
vertex_credentials: str | None = None
|
||||
|
|
|
|||
|
|
@ -4,14 +4,18 @@ import base64
|
|||
import hashlib
|
||||
import hmac
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Generator, Mapping
|
||||
from contextlib import closing
|
||||
from dataclasses import dataclass, field
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, Protocol
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from botocore.eventstream import EventStreamBuffer, ParserError
|
||||
from e2e_http import (
|
||||
NetworkError,
|
||||
StreamChunk,
|
||||
|
|
@ -23,12 +27,36 @@ from e2e_http import (
|
|||
prepare_forward,
|
||||
primed_steps,
|
||||
)
|
||||
from fixture_canonical import MARKER_PATTERN, MARKER_PLACEHOLDER
|
||||
from fixture_mode import SESSION_TEST_KEY, current_test_key
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
LIFETIME_SECONDS: Final = 86_400
|
||||
MAX_REQUEST_BYTES: Final = 256 * 1024
|
||||
MAX_RESPONSE_BYTES: Final = 8 * 1024 * 1024
|
||||
UNRECORDED_RESPONSE_HEADERS: Final = frozenset({"set-cookie"})
|
||||
SIGNATURE_HEADERS: Final = frozenset(
|
||||
{"authorization", "x-amz-date", "x-amz-security-token", "x-amz-content-sha256"}
|
||||
)
|
||||
BEDROCK_MOUNT_PREFIX: Final = "bedrock"
|
||||
BEDROCK_CONVERSE_SUFFIX: Final = "/converse"
|
||||
BEDROCK_INVOKE_SUFFIX: Final = "/invoke"
|
||||
BEDROCK_CONVERSE_STREAM_SUFFIX: Final = "/converse-stream"
|
||||
BEDROCK_INVOKE_STREAM_SUFFIX: Final = "/invoke-with-response-stream"
|
||||
BEDROCK_SUFFIXES: Final = (
|
||||
BEDROCK_CONVERSE_SUFFIX,
|
||||
BEDROCK_INVOKE_SUFFIX,
|
||||
BEDROCK_CONVERSE_STREAM_SUFFIX,
|
||||
BEDROCK_INVOKE_STREAM_SUFFIX,
|
||||
)
|
||||
EVENTSTREAM_PRELUDE_BYTES: Final = 4
|
||||
CUT_SHORT: Final = "cut_short"
|
||||
INCOMPLETE: Final = "incomplete"
|
||||
UNREACHABLE: Final = "unreachable"
|
||||
ERROR_STATUS: Final = "error_status"
|
||||
EVENT_TYPE_HEADER: Final = ":event-type"
|
||||
EVENTSTREAM_HEADERS: Final[TypeAdapter[dict[str, str]]] = TypeAdapter(dict[str, str])
|
||||
OPENAI_JSON_PATHS: Final = frozenset({"/v1/chat/completions", "/v1/messages", "/v1/embeddings", "/v1/responses"})
|
||||
JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
|
||||
|
||||
|
|
@ -56,6 +84,24 @@ class CacheUnavailable:
|
|||
|
||||
|
||||
type CacheLookup = CacheHit | CaptureLease | CacheBusy | CacheUnavailable
|
||||
type RequestSigner = Callable[[str, str, Mapping[str, str], bytes | None], dict[str, str]]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MountPolicy:
|
||||
"""What a mount needs beyond plain forwarding.
|
||||
|
||||
``sign`` mints a fresh credential over the upstream URL, for providers whose
|
||||
auth covers the Host the edge rewrote. ``unkeyed_headers`` names headers that
|
||||
must stay out of the cache key because they change on every call and would
|
||||
otherwise make the mount a permanent miss: a minted signature, or an OAuth
|
||||
token the provider rotates. Naming one costs the guarantee that a recording
|
||||
can never cross credentials, so a mount with a rotating token relies on the
|
||||
environment holding one identity for that provider. Mounts with a static API
|
||||
key name nothing here and keep the guarantee whole."""
|
||||
|
||||
sign: RequestSigner | None = None
|
||||
unkeyed_headers: frozenset[str] = frozenset()
|
||||
|
||||
|
||||
class ResponseStore(Protocol):
|
||||
|
|
@ -83,28 +129,51 @@ class SignedResponse(BaseModel):
|
|||
signature: str
|
||||
|
||||
|
||||
def exact_key(secret: bytes, method: str, url: str, headers: Mapping[str, str], body: bytes | None) -> str:
|
||||
def canonical_text(value: str) -> str:
|
||||
return MARKER_PATTERN.sub(MARKER_PLACEHOLDER, value)
|
||||
|
||||
|
||||
def canonical_body(body: bytes) -> bytes:
|
||||
try:
|
||||
return canonical_text(body.decode("utf-8")).encode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
return body
|
||||
|
||||
|
||||
def request_identity(
|
||||
secret: bytes, test_key: str, method: str, url: str, headers: Mapping[str, str], body: bytes | None,
|
||||
) -> str:
|
||||
fields: Final = (
|
||||
b"provider-cache-exact-v1", method.encode(), url.encode(),
|
||||
b"provider-cache-canonical-v2", test_key.encode(), method.encode(), canonical_text(url).encode(),
|
||||
*(part.encode() for pair in sorted(headers.items()) for part in pair),
|
||||
b"no-body" if body is None else b"body", b"" if body is None else body,
|
||||
b"no-body" if body is None else b"body", b"" if body is None else canonical_body(body),
|
||||
)
|
||||
encoded: Final = b"".join(len(part).to_bytes(8, "big") + part for part in fields)
|
||||
return hmac.new(secret, encoded, hashlib.sha256).hexdigest()
|
||||
|
||||
|
||||
def cacheable_endpoint(method: str, url: str, body: bytes | None) -> bool:
|
||||
return (
|
||||
method == "POST"
|
||||
and urlsplit(url).path in {"/v1/chat/completions", "/v1/messages"}
|
||||
and body is not None
|
||||
and len(body) <= MAX_REQUEST_BYTES
|
||||
)
|
||||
def slotted_key(secret: bytes, identity: str, slot: int) -> str:
|
||||
return hmac.new(secret, f"{identity}:{slot}".encode(), hashlib.sha256).hexdigest()
|
||||
|
||||
|
||||
def successful_response(url: str, status: int, headers: Mapping[str, str], body: bytes) -> bool:
|
||||
def is_bedrock(mount: str) -> bool:
|
||||
return mount.partition("/")[0] == BEDROCK_MOUNT_PREFIX
|
||||
|
||||
|
||||
def cacheable_endpoint(mount: str, method: str, url: str, body: bytes | None) -> bool:
|
||||
if method != "POST" or body is None or len(body) > MAX_REQUEST_BYTES:
|
||||
return False
|
||||
path: Final = urlsplit(url).path
|
||||
if is_bedrock(mount):
|
||||
return path.startswith("/model/") and path.endswith(BEDROCK_SUFFIXES)
|
||||
return path in OPENAI_JSON_PATHS
|
||||
|
||||
|
||||
def successful_response(mount: str, url: str, status: int, headers: Mapping[str, str], body: bytes) -> bool:
|
||||
if not 200 <= status < 300 or len(body) > MAX_RESPONSE_BYTES:
|
||||
return False
|
||||
if is_bedrock(mount):
|
||||
return complete_bedrock_response(url, body)
|
||||
streaming: Final = "text/event-stream" in headers.get("content-type", "").lower()
|
||||
if streaming:
|
||||
try:
|
||||
|
|
@ -118,28 +187,33 @@ def successful_response(url: str, status: int, headers: Mapping[str, str], body:
|
|||
values: Final = tuple(JSON_VALUE.validate_json(event) for event in events if event != "[DONE]")
|
||||
except (UnicodeDecodeError, ValidationError):
|
||||
return False
|
||||
if not values or any(not isinstance(value, dict) or "error" in value or value.get("type") == "error" for value in values):
|
||||
if not values or any(
|
||||
not isinstance(value, dict) or value.get("error") is not None or value.get("type") == "error"
|
||||
for value in values
|
||||
):
|
||||
return False
|
||||
if urlsplit(url).path == "/v1/responses":
|
||||
return complete_responses_stream(values)
|
||||
if urlsplit(url).path == "/v1/chat/completions":
|
||||
return events[-1] == "[DONE]" and "[DONE]" not in events[:-1] and complete_chat_stream(values)
|
||||
return (
|
||||
"[DONE]" not in events
|
||||
and isinstance(values[0], dict) and values[0].get("type") == "message_start"
|
||||
and isinstance(values[-1], dict) and values[-1].get("type") == "message_stop"
|
||||
and any(
|
||||
isinstance(value, dict) and value.get("type") == "message_delta"
|
||||
and isinstance(delta := value.get("delta"), dict) and isinstance(delta.get("stop_reason"), str)
|
||||
for value in values
|
||||
)
|
||||
)
|
||||
return "[DONE]" not in events and complete_anthropic_stream(values)
|
||||
try:
|
||||
value: Final = JSON_VALUE.validate_json(body)
|
||||
except ValidationError:
|
||||
return False
|
||||
if not isinstance(value, dict) or "error" in value:
|
||||
if not isinstance(value, dict) or value.get("error") is not None:
|
||||
return False
|
||||
if urlsplit(url).path == "/v1/messages":
|
||||
path: Final = urlsplit(url).path
|
||||
if path == "/v1/messages":
|
||||
return value.get("type") == "message" and isinstance(value.get("content"), list) and isinstance(value.get("stop_reason"), str)
|
||||
if path == "/v1/embeddings":
|
||||
data: Final = value.get("data")
|
||||
return isinstance(data, list) and bool(data) and isinstance(value.get("usage"), dict) and all(
|
||||
isinstance(item, dict) and isinstance(item.get("embedding"), list) and bool(item["embedding"])
|
||||
for item in data
|
||||
)
|
||||
if path == "/v1/responses":
|
||||
return value.get("object") == "response" and value.get("status") == "completed"
|
||||
choices: Final = value.get("choices")
|
||||
return isinstance(choices, list) and bool(choices) and all(
|
||||
isinstance(choice, dict) and isinstance(choice.get("message"), dict) and isinstance(choice.get("finish_reason"), str)
|
||||
|
|
@ -147,6 +221,144 @@ def successful_response(url: str, status: int, headers: Mapping[str, str], body:
|
|||
)
|
||||
|
||||
|
||||
def complete_bedrock_response(url: str, body: bytes) -> bool:
|
||||
"""Converse answers with ``output`` plus a ``stopReason``; InvokeModel on an
|
||||
Anthropic model answers the Anthropic message shape. Either way a truncated
|
||||
or error body is missing the terminator field, which is what makes it safe to
|
||||
record."""
|
||||
path: Final = urlsplit(url).path
|
||||
if path.endswith(BEDROCK_CONVERSE_STREAM_SUFFIX):
|
||||
return complete_converse_stream(body)
|
||||
if path.endswith(BEDROCK_INVOKE_STREAM_SUFFIX):
|
||||
return complete_invoke_stream(body)
|
||||
try:
|
||||
value: Final = JSON_VALUE.validate_json(body)
|
||||
except ValidationError:
|
||||
return False
|
||||
if not isinstance(value, dict) or "message" in value:
|
||||
return False
|
||||
if path.endswith(BEDROCK_CONVERSE_SUFFIX):
|
||||
return isinstance(value.get("output"), dict) and isinstance(value.get("stopReason"), str)
|
||||
return (
|
||||
value.get("type") == "message"
|
||||
and isinstance(value.get("content"), list)
|
||||
and isinstance(value.get("stop_reason"), str)
|
||||
)
|
||||
|
||||
|
||||
def whole_eventstream_messages(body: bytes) -> bool:
|
||||
"""Whether the body is exactly a whole number of eventstream messages.
|
||||
|
||||
A dropped connection is the failure this catches, and it has to be caught
|
||||
here: botocore yields the messages it did receive and silently discards a
|
||||
trailing partial one, so a stream cut a single byte short parses clean. Each
|
||||
message declares its own total length in its first four bytes, so walking
|
||||
those is enough to tell a complete body from a cut one."""
|
||||
offset = 0 # rebind-ok: a cursor walking the declared frame lengths
|
||||
while offset + EVENTSTREAM_PRELUDE_BYTES <= len(body):
|
||||
total: int = int.from_bytes(body[offset : offset + EVENTSTREAM_PRELUDE_BYTES], "big")
|
||||
if total <= 0 or offset + total > len(body):
|
||||
return False
|
||||
offset += total
|
||||
return offset == len(body)
|
||||
|
||||
|
||||
def eventstream_events(body: bytes) -> tuple[tuple[str, JsonValue], ...] | None:
|
||||
"""The stream's (event type, decoded payload) pairs, or None if it is not a
|
||||
complete, uncorrupted stream.
|
||||
|
||||
botocore validates both CRCs and raises ``ParserError`` rather than decoding
|
||||
corruption into something plausible. A failure that began after Bedrock had
|
||||
already answered 200 arrives as an ``exception`` frame in place of the
|
||||
terminator, so it is the terminator rules below that reject it and this does
|
||||
not need to inspect ``:message-type`` as well."""
|
||||
if not body or not whole_eventstream_messages(body):
|
||||
return None
|
||||
buffer: Final = EventStreamBuffer()
|
||||
buffer.add_data(body)
|
||||
try:
|
||||
return tuple(
|
||||
(event_type(event.headers), JSON_VALUE.validate_json(event.payload))
|
||||
for event in buffer
|
||||
)
|
||||
except (ParserError, ValidationError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def event_type(headers: object) -> str:
|
||||
"""botocore's eventstream headers come back untyped, so the one header this
|
||||
reads is validated into a string rather than trusted."""
|
||||
parsed: Final = EVENTSTREAM_HEADERS.validate_python(headers)
|
||||
return parsed.get(EVENT_TYPE_HEADER, "")
|
||||
|
||||
|
||||
def complete_converse_stream(body: bytes) -> bool:
|
||||
"""ConverseStream ends with ``metadata``, not with ``messageStop``.
|
||||
|
||||
Requiring the metadata frame rather than the stop frame is deliberate: it
|
||||
carries the token usage litellm prices the call from, so a stream cut between
|
||||
the two still names a stop reason but would replay as a free call."""
|
||||
events: Final = eventstream_events(body)
|
||||
if not events or events[-1][0] != "metadata":
|
||||
return False
|
||||
return any(
|
||||
event_type == "messageStop" and isinstance(payload, dict) and isinstance(payload.get("stopReason"), str)
|
||||
for event_type, payload in events
|
||||
)
|
||||
|
||||
|
||||
def complete_invoke_stream(body: bytes) -> bool:
|
||||
"""InvokeModelWithResponseStream wraps the ordinary Anthropic event grammar
|
||||
in ``chunk`` frames, one base64 payload each, so it is held to the same
|
||||
terminator rule as the Anthropic SSE path. A frame Bedrock sends instead of a
|
||||
chunk, an exception among them, carries no such payload and fails the rule
|
||||
without the frame type needing to be read."""
|
||||
events: Final = eventstream_events(body)
|
||||
if not events:
|
||||
return False
|
||||
values: Final = tuple(invoke_chunk_value(payload) for _, payload in events)
|
||||
return all(value is not None for value in values) and complete_anthropic_stream(values)
|
||||
|
||||
|
||||
def invoke_chunk_value(payload: JsonValue) -> JsonValue | None:
|
||||
"""The Anthropic event inside one ``chunk`` frame, or None for a frame that
|
||||
carries no readable one."""
|
||||
if not isinstance(payload, dict) or not isinstance(encoded := payload.get("bytes"), str):
|
||||
return None
|
||||
try:
|
||||
return JSON_VALUE.validate_json(base64.b64decode(encoded, validate=True))
|
||||
except (ValidationError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def complete_anthropic_stream(values: tuple[JsonValue, ...]) -> bool:
|
||||
"""The Anthropic event grammar, shared by the SSE mounts and by Bedrock's
|
||||
invoke stream, which carries the same events inside eventstream frames. A
|
||||
``message_delta`` naming a stop reason is what separates a finished turn from
|
||||
one the connection cut short."""
|
||||
if not values:
|
||||
return False
|
||||
first: Final = values[0]
|
||||
last: Final = values[-1]
|
||||
return (
|
||||
isinstance(first, dict) and first.get("type") == "message_start"
|
||||
and isinstance(last, dict) and last.get("type") == "message_stop"
|
||||
and any(
|
||||
isinstance(value, dict) and value.get("type") == "message_delta"
|
||||
and isinstance(delta := value.get("delta"), dict) and isinstance(delta.get("stop_reason"), str)
|
||||
for value in values
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def complete_responses_stream(values: tuple[JsonValue, ...]) -> bool:
|
||||
"""The Responses API streams typed events and ends with ``response.completed``.
|
||||
A run that failed, was cancelled, or ran out of tokens ends with a different
|
||||
terminal event, so requiring that one keeps a half-finished response out."""
|
||||
last: Final = values[-1]
|
||||
return isinstance(last, dict) and last.get("type") == "response.completed"
|
||||
|
||||
|
||||
def complete_chat_stream(values: tuple[JsonValue, ...]) -> bool:
|
||||
if any(not isinstance(value, dict) or not isinstance(value.get("choices"), list) for value in values):
|
||||
return False
|
||||
|
|
@ -172,7 +384,7 @@ def encode_response(secret: bytes, response: CachedResponse) -> bytes:
|
|||
return SignedResponse(response=raw, signature=hmac.new(secret, raw.encode(), hashlib.sha256).hexdigest()).model_dump_json().encode()
|
||||
|
||||
|
||||
def decode_response(secret: bytes, key: str, payload: bytes, url: str) -> CachedResponse | None:
|
||||
def decode_response(secret: bytes, key: str, payload: bytes, mount: str, url: str) -> CachedResponse | None:
|
||||
if len(payload) > 2 * MAX_RESPONSE_BYTES:
|
||||
return None
|
||||
try:
|
||||
|
|
@ -183,11 +395,58 @@ def decode_response(secret: bytes, key: str, payload: bytes, url: str) -> Cached
|
|||
chunks: Final = tuple(base64.b64decode(chunk, validate=True) for chunk in response.chunks)
|
||||
except (ValidationError, ValueError):
|
||||
return None
|
||||
if response.request_key != key or not successful_response(url, response.status_code, response.headers, b"".join(chunks)):
|
||||
if response.request_key != key or not successful_response(
|
||||
mount, url, response.status_code, response.headers, b"".join(chunks)
|
||||
):
|
||||
return None
|
||||
return response
|
||||
|
||||
|
||||
def component_digests(
|
||||
test_key: str, method: str, url: str, headers: Mapping[str, str], body: bytes | None,
|
||||
) -> dict[str, str]:
|
||||
"""Per-component digests of everything the key covers.
|
||||
|
||||
A mount whose corpus never converges is a mount where one of these moves
|
||||
between builds, and the flat key cannot say which. Values are digested, so
|
||||
no payload or credential is written, and a JSON body contributes one digest
|
||||
per top-level field so the field that moved can be named."""
|
||||
parts: dict[str, str] = { # rebind-ok: a report assembled from three differently shaped sources
|
||||
"test_key": test_key,
|
||||
"method": method,
|
||||
"url": short_digest(canonical_text(url).encode()),
|
||||
}
|
||||
for name, value in sorted(headers.items()):
|
||||
parts[f"header:{name.lower()}"] = short_digest(value.encode())
|
||||
canonical: Final = b"" if body is None else canonical_body(body)
|
||||
parts["body"] = short_digest(canonical)
|
||||
try:
|
||||
parsed: Final = JSON_VALUE.validate_json(canonical)
|
||||
except ValidationError:
|
||||
return parts
|
||||
if isinstance(parsed, dict):
|
||||
for name, value in sorted(parsed.items()):
|
||||
parts[f"body:{name}"] = short_digest(json.dumps(value, sort_keys=True).encode())
|
||||
return parts
|
||||
|
||||
|
||||
def short_digest(value: bytes) -> str:
|
||||
return hashlib.sha256(value).hexdigest()[:16]
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class KeyProbe:
|
||||
"""Every keyed request's components, when a metrics directory is configured."""
|
||||
|
||||
rows: tuple[tuple[tuple[str, str], ...], ...] = ()
|
||||
lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
|
||||
def observe(self, mount: str, outcome: str, parts: Mapping[str, str]) -> None:
|
||||
row: Final = tuple({"mount": mount, "outcome": outcome, **parts}.items())
|
||||
with self.lock:
|
||||
self.rows = (*self.rows, row)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class CacheCounters:
|
||||
counts: tuple[tuple[str, int], ...] = ()
|
||||
|
|
@ -199,6 +458,24 @@ class CacheCounters:
|
|||
self.counts = tuple((current | {name: current.get(name, 0) + 1}).items())
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SlotCounter:
|
||||
"""FIFO position of a request among the canonically identical ones its test
|
||||
has already sent. Two calls in one test that differ only by ``unique_marker``
|
||||
canonicalize the same, so without this they would share one recording and the
|
||||
second would replay the first's provider response id."""
|
||||
|
||||
counts: tuple[tuple[str, int], ...] = ()
|
||||
lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
|
||||
def take(self, identity: str) -> int:
|
||||
with self.lock:
|
||||
current: Final = dict(self.counts)
|
||||
taken: Final = current.get(identity, 0)
|
||||
self.counts = tuple((current | {identity: taken + 1}).items())
|
||||
return taken
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ResponseCapture:
|
||||
buffer: io.BytesIO = field(default_factory=io.BytesIO)
|
||||
|
|
@ -226,14 +503,21 @@ def response_steps(response: CachedResponse) -> Generator[StreamStep, None, None
|
|||
yield StreamChunk(base64.b64decode(chunk, validate=True))
|
||||
|
||||
|
||||
NO_POLICIES: Final[Mapping[str, MountPolicy]] = MappingProxyType({})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CacheEdge:
|
||||
store: ResponseStore
|
||||
secret: bytes = field(repr=False)
|
||||
counters: CacheCounters = field(default_factory=CacheCounters)
|
||||
probe: KeyProbe = field(default_factory=KeyProbe)
|
||||
slots: SlotCounter = field(default_factory=SlotCounter)
|
||||
policies: Mapping[str, MountPolicy] = NO_POLICIES
|
||||
wait_seconds: float = 2.0
|
||||
clock: Callable[[], float] = time.monotonic
|
||||
sleep: Callable[[float], None] = time.sleep
|
||||
test_key: Callable[[], str] = current_test_key
|
||||
|
||||
def lookup(self, key: str) -> CacheLookup:
|
||||
deadline: Final = self.clock() + self.wait_seconds
|
||||
|
|
@ -241,59 +525,122 @@ class CacheEdge:
|
|||
self.sleep(min(0.05, max(0, deadline - self.clock())))
|
||||
return result
|
||||
|
||||
def forward(self, method: str, url: str, headers: dict[str, str], body: bytes | None, timeout: float) -> StreamHead | NetworkError:
|
||||
if not cacheable_endpoint(method, url, body):
|
||||
self.counters.increment("bypass")
|
||||
self.counters.increment("upstream_attempts")
|
||||
return forward_stream(method, url, headers=headers, body=body, timeout=timeout)
|
||||
prepared: Final = prepare_forward(method, url, headers, body)
|
||||
def count(self, mount: str, name: str) -> None:
|
||||
self.counters.increment(name)
|
||||
self.counters.increment(f"mount:{mount}:{name}")
|
||||
|
||||
def record_key(
|
||||
self, mount: str, outcome: str, test_key: str, method: str, url: str,
|
||||
headers: Mapping[str, str], body: bytes | None,
|
||||
) -> None:
|
||||
if not os.environ.get("E2E_PROVIDER_CACHE_METRICS_DIR"):
|
||||
return
|
||||
self.probe.observe(mount, outcome, component_digests(test_key, method, url, headers, body))
|
||||
|
||||
def outbound(self, mount: str, method: str, url: str, headers: dict[str, str], body: bytes | None) -> dict[str, str]:
|
||||
"""The headers actually sent upstream. A signing mount gets a signature
|
||||
minted over the upstream URL, because the edge rewrote the Host the proxy
|
||||
signed and the provider verifies it."""
|
||||
signer: Final = self.policies.get(mount, MountPolicy()).sign
|
||||
return headers if signer is None else signer(method, url, headers, body)
|
||||
|
||||
def keyed(self, mount: str, headers: Mapping[str, str]) -> Mapping[str, str]:
|
||||
"""Headers the cache key is built from. A mount keeps its credentials in
|
||||
the key unless its policy names them unkeyed, so by default one account
|
||||
can never read another's recording."""
|
||||
unkeyed: Final = self.policies.get(mount, MountPolicy()).unkeyed_headers
|
||||
if not unkeyed:
|
||||
return headers
|
||||
return {name: value for name, value in headers.items() if name.lower() not in unkeyed}
|
||||
|
||||
def forward(
|
||||
self, mount: str, method: str, url: str, headers: dict[str, str], body: bytes | None, timeout: float,
|
||||
) -> StreamHead | NetworkError:
|
||||
test_key: Final = self.test_key()
|
||||
if test_key == SESSION_TEST_KEY or not cacheable_endpoint(mount, method, url, body):
|
||||
self.count(mount, "bypass")
|
||||
self.count(mount, "upstream_attempts")
|
||||
return forward_stream(
|
||||
method, url, headers=self.outbound(mount, method, url, headers, body), body=body, timeout=timeout,
|
||||
)
|
||||
prepared: Final = prepare_forward(method, url, self.outbound(mount, method, url, headers, body), body)
|
||||
if isinstance(prepared, NetworkError):
|
||||
self.counters.increment("rejected")
|
||||
self.reject(mount, UNREACHABLE)
|
||||
return prepared
|
||||
key: Final = exact_key(self.secret, method, url, prepared.headers, body)
|
||||
keyed_headers: Final = self.keyed(mount, prepared.headers)
|
||||
identity: Final = request_identity(self.secret, test_key, method, url, keyed_headers, body)
|
||||
key: Final = slotted_key(self.secret, identity, self.slots.take(identity))
|
||||
found: Final = self.lookup(key)
|
||||
if isinstance(found, CacheHit):
|
||||
response: Final = decode_response(self.secret, key, found.payload, url)
|
||||
response: Final = decode_response(self.secret, key, found.payload, mount, url)
|
||||
if response is not None and self.clock() < found.valid_until:
|
||||
self.counters.increment("hits")
|
||||
self.count(mount, "hits")
|
||||
self.record_key(mount, "hit", test_key, method, url, keyed_headers, body)
|
||||
return StreamHead(response.status_code, response.headers, response_steps(response))
|
||||
self.counters.increment("corrupt" if response is None else "expired")
|
||||
self.count(mount, "corrupt" if response is None else "expired")
|
||||
self.store.discard(key, found.payload)
|
||||
capture_slot: Final = self.lookup(key) if isinstance(found, CacheHit) else found
|
||||
self.counters.increment("misses")
|
||||
self.count(mount, "misses")
|
||||
self.record_key(mount, "miss", test_key, method, url, keyed_headers, body)
|
||||
if isinstance(capture_slot, CacheUnavailable):
|
||||
self.counters.increment("cache_errors")
|
||||
self.counters.increment("upstream_attempts")
|
||||
self.count(mount, "cache_errors")
|
||||
self.count(mount, "upstream_attempts")
|
||||
head: Final = forward_prepared_stream(prepared, timeout)
|
||||
if not isinstance(capture_slot, CaptureLease):
|
||||
return head
|
||||
if isinstance(head, NetworkError):
|
||||
self.store.release(key, capture_slot)
|
||||
self.counters.increment("rejected")
|
||||
self.reject(mount, UNREACHABLE)
|
||||
return head
|
||||
return StreamHead(head.status_code, head.headers, primed_steps(self.capture(key, capture_slot, url, head)))
|
||||
return StreamHead(
|
||||
head.status_code, head.headers, primed_steps(self.capture(mount, key, capture_slot, url, head)),
|
||||
)
|
||||
|
||||
def capture(self, key: str, lease: CaptureLease, url: str, head: StreamHead) -> Generator[StreamStep, None, None]:
|
||||
def capture(
|
||||
self, mount: str, key: str, lease: CaptureLease, url: str, head: StreamHead,
|
||||
) -> Generator[StreamStep, None, None]:
|
||||
capture: Final = ResponseCapture()
|
||||
reason = CUT_SHORT # rebind-ok: a consumer that walks away never reaches the settle call below
|
||||
try:
|
||||
with closing(head.steps):
|
||||
yield StreamChunk(b"")
|
||||
for step in head.steps:
|
||||
yield step
|
||||
capture.observe(step)
|
||||
chunks: Final = capture.chunks() if capture.eligible else ()
|
||||
headers: Final = {
|
||||
name: value for name, value in head.headers.items() if name.lower() not in UNRECORDED_RESPONSE_HEADERS
|
||||
}
|
||||
if not capture.eligible or not successful_response(url, head.status_code, headers, b"".join(chunks)):
|
||||
self.counters.increment("rejected")
|
||||
return
|
||||
response: Final = CachedResponse(
|
||||
request_key=key, status_code=head.status_code, headers=headers,
|
||||
chunks=tuple(base64.b64encode(chunk).decode("ascii") for chunk in chunks),
|
||||
)
|
||||
published: Final = self.store.publish(key, lease, encode_response(self.secret, response))
|
||||
self.counters.increment("writes" if published else "write_failures")
|
||||
reason = self.settle(mount, key, lease, url, head, capture)
|
||||
finally:
|
||||
self.reject(mount, reason)
|
||||
self.store.release(key, lease)
|
||||
capture.buffer.close()
|
||||
|
||||
def settle(
|
||||
self, mount: str, key: str, lease: CaptureLease, url: str, head: StreamHead, capture: ResponseCapture,
|
||||
) -> str | None:
|
||||
"""None once the response is stored, otherwise the reason it was not."""
|
||||
if not capture.eligible:
|
||||
return CUT_SHORT
|
||||
headers: Final = {
|
||||
name: value for name, value in head.headers.items() if name.lower() not in UNRECORDED_RESPONSE_HEADERS
|
||||
}
|
||||
if not 200 <= head.status_code < 300:
|
||||
return ERROR_STATUS
|
||||
chunks: Final = capture.chunks()
|
||||
if not successful_response(mount, url, head.status_code, headers, b"".join(chunks)):
|
||||
return INCOMPLETE
|
||||
response: Final = CachedResponse(
|
||||
request_key=key, status_code=head.status_code, headers=headers,
|
||||
chunks=tuple(base64.b64encode(chunk).decode("ascii") for chunk in chunks),
|
||||
)
|
||||
published: Final = self.store.publish(key, lease, encode_response(self.secret, response))
|
||||
self.count(mount, "writes" if published else "write_failures")
|
||||
return None
|
||||
|
||||
def reject(self, mount: str, reason: str | None) -> None:
|
||||
"""A flat rejection count cannot separate a connection that went away from
|
||||
a body the provider finished sending and the rules turned down, and the two
|
||||
have opposite fixes. A mount whose rejections are nearly all one or the
|
||||
other is a different problem, so the report has to be able to say which."""
|
||||
if reason is None:
|
||||
return
|
||||
self.count(mount, "rejected")
|
||||
self.count(mount, f"rejected_{reason}")
|
||||
|
|
|
|||
|
|
@ -134,6 +134,10 @@ def write_metrics(cache: CacheEdge) -> None:
|
|||
root: Final = Path(directory)
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
(root / f"{os.getpid()}.json").write_text(report + "\n")
|
||||
if cache.probe.rows:
|
||||
(root / f"keys-{os.getpid()}.json").write_text(
|
||||
json.dumps([dict(row) for row in cache.probe.rows]) + "\n"
|
||||
)
|
||||
except OSError:
|
||||
logging.getLogger(__name__).warning("provider cache metrics artifact unavailable")
|
||||
logging.getLogger(__name__).info("%s", report)
|
||||
|
|
|
|||
|
|
@ -8,14 +8,80 @@ from models import LiteLLMParamsBody, ModelMode
|
|||
|
||||
LIVE_PROVIDER_REQUIRED: Final[ContextVar[bool]] = ContextVar("live_provider_required", default=False)
|
||||
|
||||
DEFAULT_BEDROCK_REGION: Final = "us-east-1"
|
||||
BEDROCK_CROSS_REGION_PREFIX: Final = "us."
|
||||
BEDROCK_EDGE_MODELS: Final = frozenset(
|
||||
{
|
||||
"us.anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
|
||||
"us.anthropic.claude-sonnet-5",
|
||||
"us.anthropic.claude-opus-4-7",
|
||||
}
|
||||
)
|
||||
ENV_REFERENCE_PREFIX: Final = "os.environ/"
|
||||
|
||||
|
||||
def bedrock_region(declared: str | None) -> str:
|
||||
"""The region whose edge mount a deployment belongs to.
|
||||
|
||||
Most Bedrock deployments declare `os.environ/AWS_REGION`, which only the
|
||||
proxy can resolve from its own environment; the run pod does not share it.
|
||||
Answering those with the default mount is correct because every model on the
|
||||
edge allowlist is a `us.` inference profile, which fans out across the US
|
||||
regions and is reachable from any of them. That invariant is enforced on the
|
||||
allowlist itself rather than re-checked per call."""
|
||||
if declared is None or declared.startswith(ENV_REFERENCE_PREFIX):
|
||||
return DEFAULT_BEDROCK_REGION
|
||||
return declared
|
||||
|
||||
|
||||
def bedrock_mount(params: LiteLLMParamsBody) -> str | None:
|
||||
"""The edge mount a Bedrock deployment belongs to, or None.
|
||||
|
||||
The allowlist mirrors the runner role's IAM policy, which names its models
|
||||
one by one. A model outside it would be re-signed with an identity that
|
||||
cannot invoke it and come back 403 from Bedrock, so an unlisted model keeps
|
||||
its direct path and loses only caching. Adding a model is a policy edit in
|
||||
litellm-ops and a line here."""
|
||||
route: Final = params.model.partition("/")[2]
|
||||
model: Final = route.partition("/")[2] or route
|
||||
if model not in BEDROCK_EDGE_MODELS:
|
||||
return None
|
||||
return f"bedrock/{bedrock_region(params.aws_region_name)}"
|
||||
|
||||
|
||||
def route_bedrock(
|
||||
params: LiteLLMParamsBody, base_for: Callable[[str], str | None], mode: ModelMode | None,
|
||||
) -> LiteLLMParamsBody:
|
||||
"""Deployments that carry their own AWS identity stay off the edge. The edge
|
||||
re-signs with the run pod's role, so routing an `aws_role_name` deployment
|
||||
would quietly replace the very assume-role chain that test exists to prove."""
|
||||
if mode is not None or params.aws_role_name is not None or params.aws_access_key_id is not None:
|
||||
return params
|
||||
if params.api_base is not None or params.aws_bedrock_runtime_endpoint is not None:
|
||||
return params
|
||||
mount: Final = bedrock_mount(params)
|
||||
if mount is None:
|
||||
return params
|
||||
base: Final = base_for(mount)
|
||||
if base is None:
|
||||
return params
|
||||
return params.model_copy(update={"aws_bedrock_runtime_endpoint": base})
|
||||
|
||||
|
||||
def route_cache_model(
|
||||
params: LiteLLMParamsBody, base_for: Callable[[str], str | None], *, enabled: bool, mode: ModelMode | None = None,
|
||||
) -> LiteLLMParamsBody:
|
||||
if not enabled or mode == "realtime" or LIVE_PROVIDER_REQUIRED.get() or params.api_base is not None or params.mock_response is not None:
|
||||
if not enabled or LIVE_PROVIDER_REQUIRED.get() or params.mock_response is not None:
|
||||
return params
|
||||
if params.litellm_credential_name is not None:
|
||||
return params
|
||||
provider: Final = params.model.partition("/")[0]
|
||||
if provider not in {"openai", "anthropic"} or params.litellm_credential_name is not None:
|
||||
if provider == "bedrock":
|
||||
return route_bedrock(params, base_for, mode)
|
||||
if mode == "realtime" or params.api_base is not None:
|
||||
return params
|
||||
if provider not in {"openai", "anthropic"}:
|
||||
return params
|
||||
base: Final = base_for(provider)
|
||||
if base is None:
|
||||
|
|
|
|||
|
|
@ -48,7 +48,7 @@ import threading
|
|||
from collections import deque
|
||||
from collections.abc import Generator, Mapping, Sequence
|
||||
from contextlib import closing, contextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from dataclasses import dataclass, field, replace
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from itertools import islice
|
||||
from pathlib import Path
|
||||
|
|
@ -94,17 +94,41 @@ from fixture_mode import (
|
|||
parse_fixture_mode,
|
||||
)
|
||||
from fixture_profile import IneligibleRequest, MatchProfile, match_profile, strict_identity
|
||||
from provider_cache import CacheEdge
|
||||
from provider_cache import SIGNATURE_HEADERS, CacheEdge, MountPolicy, is_bedrock
|
||||
from provider_cache_routing import LIVE_PROVIDER_REQUIRED
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
BEDROCK_REGIONS: Final[tuple[str, ...]] = ("us-east-1",)
|
||||
|
||||
EDGE_MOUNTS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"openai": "https://api.openai.com",
|
||||
"anthropic": "https://api.anthropic.com",
|
||||
**{
|
||||
f"bedrock/{region}": f"https://bedrock-runtime.{region}.amazonaws.com"
|
||||
for region in BEDROCK_REGIONS
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ResolvedMount:
|
||||
mount: str
|
||||
upstream_base: str
|
||||
upstream_path: str
|
||||
|
||||
|
||||
def resolve_mount(path: str, mounts: Mapping[str, str]) -> ResolvedMount | None:
|
||||
"""Longest mount prefix wins, so a region-qualified mount such as
|
||||
``bedrock/us-east-1`` resolves whole instead of leaving the region as the
|
||||
first segment of the upstream path."""
|
||||
trimmed: Final = path.lstrip("/")
|
||||
for mount in sorted(mounts, key=len, reverse=True):
|
||||
if trimmed == mount or trimmed.startswith(f"{mount}/"):
|
||||
return ResolvedMount(mount, mounts[mount], trimmed[len(mount):].lstrip("/"))
|
||||
return None
|
||||
|
||||
REPLAY_MISS_STATUS: Final = 599
|
||||
|
||||
_HOP_BY_HOP_HEADERS: Final[frozenset[str]] = frozenset(
|
||||
|
|
@ -754,14 +778,14 @@ def _handle_record(
|
|||
|
||||
def _handle_live(
|
||||
method: str, url: str, headers: Mapping[str, str], body: bytes | None, timeout: float,
|
||||
cache: CacheEdge | None = None,
|
||||
cache: CacheEdge | None = None, mount: str = "",
|
||||
) -> EdgeOutcome:
|
||||
forwarded: Final = {
|
||||
name: value for name, value in headers.items() if name.lower() not in _REQUEST_DROPPED_HEADERS
|
||||
}
|
||||
head: Final = (
|
||||
forward_stream(method, url, headers=forwarded, body=body, timeout=timeout)
|
||||
if cache is None else cache.forward(method, url, forwarded, body, timeout)
|
||||
if cache is None else cache.forward(mount, method, url, forwarded, body, timeout)
|
||||
)
|
||||
match head:
|
||||
case NetworkError(message=message):
|
||||
|
|
@ -796,10 +820,13 @@ def handle_edge_request(
|
|||
prefix, then record (forward + persist) or replay (serve from the bundle).
|
||||
Socket-free so unit tests exercise every branch without a server."""
|
||||
split: Final = urlsplit(raw_path)
|
||||
mount, _, upstream_path = split.path.lstrip("/").partition("/")
|
||||
upstream_base: Final = mounts.get(mount)
|
||||
if upstream_base is None:
|
||||
return _text_reply(404, f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(mounts))}")
|
||||
resolved: Final = resolve_mount(split.path, mounts)
|
||||
if resolved is None:
|
||||
unknown: Final = split.path.lstrip("/").partition("/")[0]
|
||||
return _text_reply(404, f"unknown provider mount {unknown!r}; known mounts: {', '.join(sorted(mounts))}")
|
||||
mount: Final = resolved.mount
|
||||
upstream_base: Final = resolved.upstream_base
|
||||
upstream_path: Final = resolved.upstream_path
|
||||
profile: Final = (
|
||||
backend.recorder.profile
|
||||
if isinstance(backend, RecordEdge)
|
||||
|
|
@ -830,7 +857,8 @@ def handle_edge_request(
|
|||
match backend:
|
||||
case CacheEdge():
|
||||
return _handle_live(
|
||||
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout, backend,
|
||||
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout,
|
||||
backend, mount,
|
||||
)
|
||||
case LiveEdge():
|
||||
return _handle_live(
|
||||
|
|
@ -891,7 +919,7 @@ class _EdgeHandler(BaseHTTPRequestHandler):
|
|||
)
|
||||
if isinstance(edge_server.backend, CacheEdge) and duplicate_headers:
|
||||
edge_server.backend.counters.increment("duplicate_header_bypass")
|
||||
if urlsplit(self.path).path.lstrip("/").partition("/")[0] in edge_server.mounts:
|
||||
if resolve_mount(urlsplit(self.path).path, edge_server.mounts) is not None:
|
||||
edge_server.backend.counters.increment("upstream_attempts")
|
||||
outcome: Final = handle_edge_request(
|
||||
selected_backend,
|
||||
|
|
@ -1079,6 +1107,8 @@ def provider_edge_api_base(
|
|||
return _shared_cache_edge(bind_host, advertise_host, forward_timeout).api_base(mount)
|
||||
return None
|
||||
case "record" | "replay":
|
||||
if is_bedrock(mount):
|
||||
return None
|
||||
if mount not in EDGE_MOUNTS:
|
||||
raise ValueError(f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(EDGE_MOUNTS))}")
|
||||
return _shared_edge(mode, bundle_dir, bind_host, advertise_host, forward_timeout, match_profile()).api_base(
|
||||
|
|
@ -1108,7 +1138,22 @@ def configured_cache_backend() -> CacheEdge | None:
|
|||
return None
|
||||
from provider_cache_redis import configured_cache
|
||||
|
||||
return configured_cache()
|
||||
cache: Final = configured_cache()
|
||||
return None if cache is None else replace(cache, policies=bedrock_policies())
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def bedrock_policies() -> Mapping[str, MountPolicy]:
|
||||
"""One policy per mounted Bedrock region, built lazily so a run that never
|
||||
mounts Bedrock neither imports botocore nor resolves an AWS identity."""
|
||||
from provider_edge_bedrock import bedrock_signer
|
||||
|
||||
return MappingProxyType(
|
||||
{
|
||||
f"bedrock/{region}": MountPolicy(sign=bedrock_signer(region), unkeyed_headers=SIGNATURE_HEADERS)
|
||||
for region in BEDROCK_REGIONS
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=8)
|
||||
|
|
|
|||
72
tests/e2e/provider_edge_bedrock.py
Normal file
72
tests/e2e/provider_edge_bedrock.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
"""SigV4 re-signing for Bedrock traffic routed through the provider edge.
|
||||
|
||||
Bedrock is the one provider the edge could never mount. SigV4 signs the Host
|
||||
header, so rewriting ``api_base`` to point at the edge invalidates the proxy's
|
||||
signature and Bedrock rejects the call before it reaches a model. The edge
|
||||
therefore has to drop the proxy's signature and mint its own over the upstream
|
||||
URL it is actually about to call.
|
||||
|
||||
The identity it signs with is the run pod's own, from the EKS Pod Identity
|
||||
association on ServiceAccount ``buildkite-e2e-run``. That role carries Bedrock
|
||||
invoke and converse on an allowlist of the Anthropic models the suite registers
|
||||
and nothing else, so a re-signed call can reach exactly the models the suite
|
||||
already uses. The proxy's own Bedrock credentials are not involved in a routed
|
||||
deployment, which is why ``aws_role_name`` deployments stay off the edge: their
|
||||
whole point is to prove the product's assume-role chain.
|
||||
|
||||
Signature headers are excluded from the cache key by the caller, and they have
|
||||
to be: ``x-amz-date`` is a timestamp, so keying on it would make every Bedrock
|
||||
request a permanent miss.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
from botocore.auth import SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
from botocore.credentials import Credentials
|
||||
from botocore.session import Session
|
||||
from provider_cache import SIGNATURE_HEADERS
|
||||
|
||||
BEDROCK_SERVICE: Final = "bedrock"
|
||||
|
||||
|
||||
class MissingAwsCredentials(RuntimeError):
|
||||
"""No AWS identity is resolvable, so the edge cannot sign for Bedrock."""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BedrockSigner:
|
||||
region: str
|
||||
credentials: Callable[[], Credentials]
|
||||
|
||||
def __call__(self, method: str, url: str, headers: Mapping[str, str], body: bytes | None) -> dict[str, str]:
|
||||
unsigned: Final = {
|
||||
name: value for name, value in headers.items() if name.lower() not in SIGNATURE_HEADERS
|
||||
}
|
||||
request: Final = AWSRequest(method=method, url=url, headers=unsigned, data=body or b"")
|
||||
SigV4Auth(self.credentials(), BEDROCK_SERVICE, self.region).add_auth(request)
|
||||
return dict(request.headers)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def pod_credentials() -> Credentials:
|
||||
"""The run pod's own identity, resolved once per process through botocore's
|
||||
ordinary chain, which reaches Pod Identity at the ``container-role`` link."""
|
||||
resolved: Final = Session().get_credentials()
|
||||
if resolved is None: # pyright: ignore[reportUnnecessaryComparison] # stubs miss the empty-chain None
|
||||
raise MissingAwsCredentials(
|
||||
"the provider edge is mounted for Bedrock but no AWS credentials resolve; "
|
||||
"the run pod gets them from the Pod Identity association on buildkite-e2e-run"
|
||||
)
|
||||
return resolved
|
||||
|
||||
|
||||
def bedrock_signer(region: str, credentials: Callable[[], Credentials] = pod_credentials) -> BedrockSigner:
|
||||
"""Credentials are resolved on the first signed request, not here, so a run
|
||||
that mounts Bedrock but never calls it needs no AWS identity at all."""
|
||||
return BedrockSigner(region, credentials)
|
||||
|
|
@ -10,4 +10,5 @@ markers =
|
|||
weekly: real-provider anomaly load test that spends real money; deselected unless E2E_WEEKLY_ANOMALY is set
|
||||
managed_files: needs a proxy running with require_managed_files enabled; deselected unless E2E_MANAGED_FILES_STACK is set
|
||||
prompt_caching_stack: needs a proxy running with router_settings.optional_pre_call_checks including prompt_caching; deselected unless E2E_PROMPT_CACHING_STACK is set
|
||||
cli_determinism: drives the real claude CLI for several seconds, which widens the window in which another test's in-flight upstream call is attributed to it; deselected unless E2E_CLI_DETERMINISM is set
|
||||
redis_chaos: load test that pauses the proxy's Redis outright mid-run; needs a proxy booted from gateway/redis_chaos_ci_config.yml on the same host, and is deselected unless E2E_REDIS_CHAOS is set
|
||||
|
|
|
|||
|
|
@ -1279,15 +1279,30 @@ class TestApiBaseSeam:
|
|||
)
|
||||
|
||||
def test_unknown_mount_raises_naming_the_known_mounts(self, tmp_path: Path) -> None:
|
||||
with pytest.raises(ValueError, match="unknown provider mount 'bedrock'"):
|
||||
with pytest.raises(ValueError, match="unknown provider mount 'cohere'"):
|
||||
provider_edge_api_base(
|
||||
"bedrock",
|
||||
"cohere",
|
||||
mode_raw="record",
|
||||
bundle_dir=tmp_path / "bundle",
|
||||
bind_host="127.0.0.1",
|
||||
advertise_host="127.0.0.1",
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("mode_raw", ["record", "replay"])
|
||||
def test_bedrock_never_wires_a_bundle_because_the_edge_cannot_sign_into_one(
|
||||
self, tmp_path: Path, mode_raw: str,
|
||||
) -> None:
|
||||
"""Record and replay serve from a bundle without re-signing, so a Bedrock
|
||||
deployment pointed at that edge would send the proxy's signature over a
|
||||
rewritten Host. It keeps its direct route in both modes."""
|
||||
assert provider_edge_api_base(
|
||||
"bedrock/us-east-1",
|
||||
mode_raw=mode_raw,
|
||||
bundle_dir=tmp_path / "bundle",
|
||||
bind_host="127.0.0.1",
|
||||
advertise_host="127.0.0.1",
|
||||
) is None
|
||||
|
||||
def test_record_mode_boots_one_shared_edge_and_prepares_the_bundle(self, tmp_path: Path) -> None:
|
||||
root = tmp_path / "bundle"
|
||||
first = provider_edge_api_base(
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ Test TogetherAI LLM
|
|||
"""
|
||||
|
||||
from base_llm_unit_tests import BaseLLMChatTest
|
||||
from tests._live_test_helpers import cheapest_together_chat_model
|
||||
import json
|
||||
import os
|
||||
from datetime import datetime
|
||||
|
|
@ -16,7 +17,11 @@ import pytest
|
|||
class TestTogetherAI(BaseLLMChatTest):
|
||||
def get_base_completion_call_args(self) -> dict:
|
||||
litellm.set_verbose = True
|
||||
return {"model": "together_ai/openai/gpt-oss-20b"}
|
||||
return {
|
||||
"model": cheapest_together_chat_model(
|
||||
function_calling=True, response_schema=True
|
||||
)
|
||||
}
|
||||
|
||||
def test_tool_call_no_arguments(self, tool_call_no_arguments):
|
||||
"""Test that tool calls with no arguments is translated correctly. Relevant issue: https://github.com/BerriAI/litellm/issues/6833"""
|
||||
|
|
|
|||
|
|
@ -57,23 +57,6 @@ def test_response_model_none():
|
|||
assert isinstance(x, litellm.ModelResponse)
|
||||
|
||||
|
||||
def test_completion_custom_provider_model_name():
|
||||
try:
|
||||
litellm.cache = None
|
||||
response = completion(
|
||||
model="together_ai/openai/gpt-oss-20b",
|
||||
messages=messages,
|
||||
logger_fn=logger_fn,
|
||||
)
|
||||
# Add assertions here to check the-response
|
||||
print(response)
|
||||
print(response["choices"][0]["finish_reason"])
|
||||
except litellm.Timeout as e:
|
||||
pass
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
||||
|
||||
def _openai_mock_response(*args, **kwargs) -> litellm.ModelResponse:
|
||||
new_response = MagicMock()
|
||||
new_response.headers = {"hello": "world"}
|
||||
|
|
@ -2803,41 +2786,6 @@ def test_completion_together_ai_llama():
|
|||
|
||||
|
||||
# test_completion_together_ai()
|
||||
def test_customprompt_together_ai():
|
||||
try:
|
||||
litellm.set_verbose = False
|
||||
litellm.num_retries = 0
|
||||
print("in test_customprompt_together_ai")
|
||||
print(litellm.success_callback)
|
||||
print(litellm._async_success_callback)
|
||||
response = completion(
|
||||
model="together_ai/openai/gpt-oss-20b",
|
||||
messages=messages,
|
||||
roles={
|
||||
"system": {
|
||||
"pre_message": "<|im_start|>system\n",
|
||||
"post_message": "<|im_end|>",
|
||||
},
|
||||
"assistant": {
|
||||
"pre_message": "<|im_start|>assistant\n",
|
||||
"post_message": "<|im_end|>",
|
||||
},
|
||||
"user": {
|
||||
"pre_message": "<|im_start|>user\n",
|
||||
"post_message": "<|im_end|>",
|
||||
},
|
||||
},
|
||||
)
|
||||
print(response)
|
||||
except litellm.exceptions.Timeout as e:
|
||||
print(f"Timeout Error")
|
||||
pass
|
||||
except Exception as e:
|
||||
print(f"ERROR TYPE {type(e)}")
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
||||
|
||||
# test_customprompt_together_ai()
|
||||
|
||||
|
||||
def response_format_tests(response: litellm.ModelResponse):
|
||||
|
|
@ -3644,28 +3592,6 @@ async def test_acompletion_stream_watsonx():
|
|||
# test_maritalk()
|
||||
|
||||
|
||||
def test_completion_together_ai_stream():
|
||||
litellm.set_verbose = True
|
||||
user_message = "Write 1pg about YC & litellm"
|
||||
messages = [{"content": user_message, "role": "user"}]
|
||||
try:
|
||||
response = completion(
|
||||
model="together_ai/openai/gpt-oss-20b",
|
||||
messages=messages,
|
||||
stream=True,
|
||||
max_tokens=5,
|
||||
)
|
||||
print(response)
|
||||
for chunk in response:
|
||||
print(chunk)
|
||||
# print(string_response)
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
||||
|
||||
# test_completion_together_ai_stream()
|
||||
|
||||
|
||||
def test_moderation():
|
||||
response = litellm.moderation(input="i'm ishaan cto of litellm")
|
||||
print(response)
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from unittest.mock import MagicMock, patch
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from tests._live_test_helpers import cheapest_together_chat_model
|
||||
from litellm import (
|
||||
RateLimitError,
|
||||
TextCompletionResponse,
|
||||
|
|
@ -4030,7 +4031,7 @@ def test_async_text_completion_together_ai():
|
|||
async def test_get_response():
|
||||
try:
|
||||
response = await litellm.atext_completion(
|
||||
model="together_ai/openai/gpt-oss-20b",
|
||||
model=cheapest_together_chat_model(),
|
||||
prompt="good morning",
|
||||
max_tokens=10,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -10,12 +10,11 @@ Pins the five helpers
|
|||
|
||||
Driven through /team/new + /team/update.
|
||||
|
||||
Structural finding, updated: /team/new loads the org via `get_org_object`
|
||||
WITH `include_budget_table=True`, so the org max_budget / org tpm / org rpm
|
||||
guards inside `_check_org_team_limits` are live there and are pinned as
|
||||
enforced below. /team/update still loads the org without the budget
|
||||
relation, so its budget guards remain no-ops. The `models` subset guard IS
|
||||
reachable on both because it reads `org_table.models` directly. The
|
||||
Structural finding, updated: /team/new and /team/update both load the org
|
||||
via `get_org_object` WITH `include_budget_table=True`, so the org max_budget /
|
||||
org tpm / org rpm guards inside `_check_org_team_limits` are live on both and
|
||||
are pinned as enforced below. The `models` subset guard reads
|
||||
`org_table.models` directly. The
|
||||
`_check_user_team_limits` guards reach all branches through
|
||||
`user_api_key_dict`, no relation include needed.
|
||||
"""
|
||||
|
|
@ -139,9 +138,8 @@ async def test_check_org_team_limits_models_subset(
|
|||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _check_org_team_limits — budget / tpm / rpm live on /team/new since its
|
||||
# get_org_object call passes include_budget_table=True. (/team/update still
|
||||
# loads the org without the budget relation, so its guards remain no-ops.)
|
||||
# _check_org_team_limits — budget / tpm / rpm live on /team/new and
|
||||
# /team/update since both get_org_object calls pass include_budget_table=True.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_ORG_BUDGET_ENFORCED_SCENARIOS = [
|
||||
|
|
@ -216,6 +214,35 @@ async def test_check_org_team_limits_budget_enforced(
|
|||
assert len(rows) == (1 if expected_status == 200 else 0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"org_budget,body_extras,expected_status",
|
||||
[(b, c, d) for (_id, b, c, d) in _ORG_BUDGET_ENFORCED_SCENARIOS],
|
||||
ids=[s[0] for s in _ORG_BUDGET_ENFORCED_SCENARIOS],
|
||||
)
|
||||
async def test_check_org_team_limits_budget_enforced_on_update(
|
||||
org_budget,
|
||||
body_extras: Dict[str, Any],
|
||||
expected_status: int,
|
||||
proxy_client,
|
||||
prisma,
|
||||
scratch,
|
||||
world,
|
||||
):
|
||||
org_id = await create_scratch_org(prisma, scratch.prefix, **org_budget)
|
||||
team_id = await create_scratch_team(prisma, scratch.tag("team"), organization_id=org_id)
|
||||
seeder = world.keys[Actor.PROXY_ADMIN].cleartext
|
||||
resp = await proxy_client.post(
|
||||
"/team/update",
|
||||
headers={"Authorization": f"Bearer {seeder}"},
|
||||
json={"team_id": team_id, **body_extras},
|
||||
)
|
||||
assert resp.status_code == expected_status, f"{body_extras!r} → {resp.status_code}: {resp.text}"
|
||||
row = await prisma.db.litellm_teamtable.find_unique(where={"team_id": team_id})
|
||||
assert row is not None
|
||||
persisted = {field: getattr(row, field) for field in body_extras}
|
||||
assert (persisted == body_extras) == (expected_status == 200)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _check_user_team_limits — fires for standalone (no-org) teams created by
|
||||
# a non-admin caller. Each guard reads from user_api_key_dict / user_obj.
|
||||
|
|
@ -310,65 +337,40 @@ async def test_check_user_team_limits(
|
|||
# /team/update path — budget authority.
|
||||
#
|
||||
# The caller's PERSONAL limits are never applied on update (that compared the
|
||||
# wrong thing). But raising a team's spend ceiling is reserved for proxy admins:
|
||||
# a team admin may keep or LOWER the budget, only a proxy admin may RAISE it.
|
||||
# _check_user_team_limits() only runs on /team/new.
|
||||
# wrong thing). Raising a team's spend ceiling is reserved for proxy admins.
|
||||
# max_budget is not on the team-admin allow-list yet (LIT-5722), so a team
|
||||
# admin is refused in either direction; the raise-only guard underneath the
|
||||
# allow-list is pinned in the unit tests. _check_user_team_limits() only runs
|
||||
# on /team/new.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_team_admin_raise_budget_blocked(proxy_client, prisma, scratch):
|
||||
"""A team admin cannot raise the team's budget; the block is NOT based on
|
||||
their personal budget (which here is higher than the requested value)."""
|
||||
caller_cleartext = await _seed_scratch_actor_with_caps(
|
||||
prisma,
|
||||
scratch.prefix,
|
||||
max_budget=100000.0, # generous personal budget; must not matter
|
||||
)
|
||||
creator_user_id = f"{scratch.prefix}-team-creator"
|
||||
@pytest.mark.parametrize(
|
||||
"personal_budget,requested_budget",
|
||||
[(100000.0, 999.0), (10.0, 300.0)],
|
||||
ids=["raise_with_generous_personal_budget", "lower_with_tiny_personal_budget"],
|
||||
)
|
||||
async def test_team_admin_cannot_change_budget_while_max_budget_is_not_editable(
|
||||
proxy_client, prisma, scratch, personal_budget: float, requested_budget: float
|
||||
):
|
||||
caller_cleartext = await _seed_scratch_actor_with_caps(prisma, scratch.prefix, max_budget=personal_budget)
|
||||
team_id = await create_scratch_team(
|
||||
prisma,
|
||||
team_id=scratch.tag("team"),
|
||||
admin_user_ids=[creator_user_id],
|
||||
max_budget=50.0,
|
||||
)
|
||||
# Raise the team budget 50 -> 999 as a team admin.
|
||||
resp = await proxy_client.post(
|
||||
"/team/update",
|
||||
headers={"Authorization": f"Bearer {caller_cleartext}"},
|
||||
json={"team_id": team_id, "max_budget": 999.0},
|
||||
)
|
||||
assert resp.status_code == 403, resp.text
|
||||
|
||||
row = await prisma.db.litellm_teamtable.find_unique(where={"team_id": team_id})
|
||||
assert row is not None
|
||||
assert row.max_budget == 50.0, "team budget must not change on a blocked raise"
|
||||
|
||||
|
||||
async def test_team_admin_lower_budget_allowed(proxy_client, prisma, scratch):
|
||||
"""A team admin may freely lower (or keep) the team's budget."""
|
||||
caller_cleartext = await _seed_scratch_actor_with_caps(
|
||||
prisma,
|
||||
scratch.prefix,
|
||||
max_budget=10.0, # below both the old and new team budget; must not matter
|
||||
)
|
||||
creator_user_id = f"{scratch.prefix}-team-creator"
|
||||
team_id = await create_scratch_team(
|
||||
prisma,
|
||||
team_id=scratch.tag("team"),
|
||||
admin_user_ids=[creator_user_id],
|
||||
admin_user_ids=[f"{scratch.prefix}-team-creator"],
|
||||
max_budget=500.0,
|
||||
)
|
||||
# Lower the team budget 500 -> 300 as a team admin.
|
||||
resp = await proxy_client.post(
|
||||
"/team/update",
|
||||
headers={"Authorization": f"Bearer {caller_cleartext}"},
|
||||
json={"team_id": team_id, "max_budget": 300.0},
|
||||
json={"team_id": team_id, "max_budget": requested_budget},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert resp.status_code == 403, resp.text
|
||||
assert "Team admin editable fields" in resp.text, resp.text
|
||||
|
||||
row = await prisma.db.litellm_teamtable.find_unique(where={"team_id": team_id})
|
||||
assert row is not None
|
||||
assert row.max_budget == 300.0, "team admin should be able to lower the budget"
|
||||
assert row.max_budget == 500.0, "a refused update must leave the team budget unchanged"
|
||||
|
||||
|
||||
async def test_proxy_admin_raise_budget_allowed(proxy_client, prisma, scratch):
|
||||
|
|
|
|||
|
|
@ -9,31 +9,31 @@ pytestmark = pytest.mark.asyncio(loop_scope="session")
|
|||
|
||||
|
||||
# POST /team/update — actor x team-shape matrix (shapes built by _seed_target).
|
||||
# Each request carries the team's own organization_id so a non-proxy-admin can
|
||||
# reach the org-scoped branch of the route-permission gate (401 on denial),
|
||||
# which fronts the handler's _verify_team_access. Only PROXY_ADMIN and an
|
||||
# ORG_ADMIN of the team's org pass: an internal_user team admin is filtered by
|
||||
# the route gate before _verify_team_access's team-admin branch is reached.
|
||||
# The route is self-managed (LIT-5722), so every authenticated caller reaches
|
||||
# update_team and denials are the handler's 403, never the route gate's 401.
|
||||
# Only PROXY_ADMIN and an ORG_ADMIN of the team's org pass: a team admin is
|
||||
# admitted by _resolve_team_access but then refused because no team field is
|
||||
# enabled for team admins (team_admin_editable_team_fields defaults to empty).
|
||||
MARKER_ALIAS = "behavior-pin-update-marker-alias"
|
||||
|
||||
_MATRIX = [
|
||||
("alpha/proxy_admin", Actor.PROXY_ADMIN, "alpha", 200),
|
||||
("alpha/org_admin", Actor.ORG_ADMIN, "alpha", 200),
|
||||
("alpha/team_admin", Actor.TEAM_ADMIN, "alpha", 401),
|
||||
("alpha/internal_user", Actor.INTERNAL_USER, "alpha", 401),
|
||||
("alpha/owner", Actor.OWNER, "alpha", 401),
|
||||
("alpha/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "alpha", 401),
|
||||
("alpha/cross_org_user", Actor.CROSS_ORG_USER, "alpha", 401),
|
||||
("alpha/service_account", Actor.SERVICE_ACCOUNT, "alpha", 401),
|
||||
("alpha/org_b_admin", Actor.ORG_B_ADMIN, "alpha", 401),
|
||||
("alpha/team_admin", Actor.TEAM_ADMIN, "alpha", 403),
|
||||
("alpha/internal_user", Actor.INTERNAL_USER, "alpha", 403),
|
||||
("alpha/owner", Actor.OWNER, "alpha", 403),
|
||||
("alpha/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "alpha", 403),
|
||||
("alpha/cross_org_user", Actor.CROSS_ORG_USER, "alpha", 403),
|
||||
("alpha/service_account", Actor.SERVICE_ACCOUNT, "alpha", 403),
|
||||
("alpha/org_b_admin", Actor.ORG_B_ADMIN, "alpha", 403),
|
||||
("beta/proxy_admin", Actor.PROXY_ADMIN, "beta", 200),
|
||||
("beta/org_admin", Actor.ORG_ADMIN, "beta", 401),
|
||||
("beta/team_admin", Actor.TEAM_ADMIN, "beta", 401),
|
||||
("beta/internal_user", Actor.INTERNAL_USER, "beta", 401),
|
||||
("beta/owner", Actor.OWNER, "beta", 401),
|
||||
("beta/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "beta", 401),
|
||||
("beta/cross_org_user", Actor.CROSS_ORG_USER, "beta", 401),
|
||||
("beta/service_account", Actor.SERVICE_ACCOUNT, "beta", 401),
|
||||
("beta/org_admin", Actor.ORG_ADMIN, "beta", 403),
|
||||
("beta/team_admin", Actor.TEAM_ADMIN, "beta", 403),
|
||||
("beta/internal_user", Actor.INTERNAL_USER, "beta", 403),
|
||||
("beta/owner", Actor.OWNER, "beta", 403),
|
||||
("beta/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "beta", 403),
|
||||
("beta/cross_org_user", Actor.CROSS_ORG_USER, "beta", 403),
|
||||
("beta/service_account", Actor.SERVICE_ACCOUNT, "beta", 403),
|
||||
("beta/org_b_admin", Actor.ORG_B_ADMIN, "beta", 200),
|
||||
]
|
||||
|
||||
|
|
@ -110,8 +110,9 @@ async def test_team_update_org_admin_resolved_from_team_without_org_context(
|
|||
):
|
||||
"""With no organization_id in the body the route gate resolves the target
|
||||
team's org from team_id, so an org admin of the team's own org is allowed
|
||||
(200), same as PROXY_ADMIN. A team admin of that same team stays denied
|
||||
(401): the resolution grants org admins access, not team admins."""
|
||||
(200), same as PROXY_ADMIN. A team admin of that same team reaches the
|
||||
handler but is refused (403) until a proxy admin enables fields for team
|
||||
admins, and the response says so."""
|
||||
await _seed_target(prisma, world, "alpha", scratch.prefix)
|
||||
|
||||
allowed_org_admin = await proxy_client.post(
|
||||
|
|
@ -133,21 +134,25 @@ async def test_team_update_org_admin_resolved_from_team_without_org_context(
|
|||
headers={"Authorization": f"Bearer {world.keys[Actor.TEAM_ADMIN].cleartext}"},
|
||||
json={"team_id": scratch.prefix, "team_alias": MARKER_ALIAS},
|
||||
)
|
||||
assert denied_team_admin.status_code == 401, denied_team_admin.text
|
||||
assert denied_team_admin.status_code == 403, denied_team_admin.text
|
||||
assert "cannot edit team settings" in denied_team_admin.text, denied_team_admin.text
|
||||
assert "Team admin editable fields" in denied_team_admin.text, denied_team_admin.text
|
||||
|
||||
|
||||
# Relocation gate — moving a team to a different org. The scratch team starts
|
||||
# in ORG_A; each scenario relocates it to ORG_B. PROXY_ADMIN bypasses;
|
||||
# ORG_B_ADMIN clears the route gate (dest-org admin) but fails
|
||||
# _verify_team_access on the source team (403); the rest fail the route gate
|
||||
# (401). The relocation-*allowed* branch (caller is org admin of both orgs) is
|
||||
# covered by test_team_update_org_relocation_allowed_for_dual_org_admin below.
|
||||
# ORG_B_ADMIN reaches the handler but holds no role on the source team (403);
|
||||
# ORG_ADMIN holds the source team but not the destination org (403 from the
|
||||
# relocation gate); the team admin is refused by the empty field allow-list and
|
||||
# the internal user holds no role at all (403). The relocation-*allowed* branch
|
||||
# (caller is org admin of both orgs) is covered by
|
||||
# test_team_update_org_relocation_allowed_for_dual_org_admin below.
|
||||
_RELOCATION = [
|
||||
("proxy_admin", Actor.PROXY_ADMIN, 200),
|
||||
("org_b_admin", Actor.ORG_B_ADMIN, 403),
|
||||
("org_admin", Actor.ORG_ADMIN, 401),
|
||||
("team_admin", Actor.TEAM_ADMIN, 401),
|
||||
("internal_user", Actor.INTERNAL_USER, 401),
|
||||
("org_admin", Actor.ORG_ADMIN, 403),
|
||||
("team_admin", Actor.TEAM_ADMIN, 403),
|
||||
("internal_user", Actor.INTERNAL_USER, 403),
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -809,6 +809,7 @@ def test_img_gen(mock_aimage_generation, client_no_auth):
|
|||
n=1,
|
||||
size="1024x1024",
|
||||
imageConfig={"aspectRatio": "9:16", "imageSize": "1K"},
|
||||
litellm_call_id=mock.ANY,
|
||||
metadata=mock.ANY,
|
||||
proxy_server_request=mock.ANY,
|
||||
secret_fields=mock.ANY,
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ class TestSkipPreCallLogic:
|
|||
await processor.base_process_llm_request(
|
||||
request=MagicMock(spec=Request),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
route_type="aresponses",
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
llm_router=MagicMock(),
|
||||
|
|
|
|||
|
|
@ -1934,3 +1934,18 @@ async def test_discovery_auth_fingerprint_tracks_effective_credentials(resolved:
|
|||
assert original != replaced
|
||||
assert len(original) == 64
|
||||
assert "private-original-credential" not in original
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_auth_preview_uses_the_same_effective_headers_as_egress() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
|
||||
|
||||
client: Final = MCPClient(
|
||||
server_url="https://upstream.example/mcp", auth_type=MCPAuth.bearer_token,
|
||||
resolved_auth=StaticHeaderAuth("Bearer resolved"), extra_headers={"X-Trace": "trace"},
|
||||
)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.method == "POST"
|
||||
assert str(request.url) == "https://upstream.example/mcp"
|
||||
assert request.headers["Authorization"] == "Bearer resolved"
|
||||
assert request.headers["X-Trace"] == "trace"
|
||||
|
|
|
|||
|
|
@ -19,12 +19,12 @@ to 0 when the only update we saw was the cursor, allowing the
|
|||
text-based fallback to estimate from the real completion text.
|
||||
"""
|
||||
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.streaming_chunk_builder_utils import ChunkProcessor
|
||||
from litellm.types.utils import (
|
||||
CompletionTokensDetailsWrapper,
|
||||
Delta,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
|
|
@ -35,6 +35,7 @@ from litellm.types.utils import (
|
|||
def _make_chunk(
|
||||
*,
|
||||
content: str = "",
|
||||
reasoning_content: str | None = None,
|
||||
usage: Usage = None,
|
||||
finish_reason: str = None,
|
||||
custom_llm_provider: str = "anthropic",
|
||||
|
|
@ -48,7 +49,7 @@ def _make_chunk(
|
|||
StreamingChoices(
|
||||
finish_reason=finish_reason,
|
||||
index=0,
|
||||
delta=Delta(content=content, role="assistant"),
|
||||
delta=Delta(content=content, role="assistant", reasoning_content=reasoning_content),
|
||||
)
|
||||
],
|
||||
usage=usage,
|
||||
|
|
@ -69,9 +70,7 @@ class TestAnthropicCursorBug:
|
|||
token_counter fallback can estimate from completion text.
|
||||
"""
|
||||
# Anthropic message_start: input_tokens accurate, output_tokens=1 cursor
|
||||
message_start = _make_chunk(
|
||||
usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
|
||||
)
|
||||
message_start = _make_chunk(usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025))
|
||||
# Several content_block_delta chunks (no usage attached)
|
||||
text_chunks = [
|
||||
_make_chunk(content="Hello"),
|
||||
|
|
@ -97,9 +96,7 @@ class TestAnthropicCursorBug:
|
|||
Normal complete stream: message_start cursor=1, then message_delta=3847.
|
||||
Last-wins must give 3847 (the real value).
|
||||
"""
|
||||
message_start = _make_chunk(
|
||||
usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
|
||||
)
|
||||
message_start = _make_chunk(usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025))
|
||||
text_chunks = [_make_chunk(content=t) for t in ["Hello", " world", "!"]]
|
||||
# message_delta with the real cumulative output_tokens
|
||||
message_delta = _make_chunk(
|
||||
|
|
@ -119,19 +116,14 @@ class TestAnthropicCursorBug:
|
|||
End-to-end via calculate_usage(): cursor-only stream + real completion
|
||||
text should produce a token-counter estimate, NOT 1.
|
||||
"""
|
||||
message_start = _make_chunk(
|
||||
usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
|
||||
)
|
||||
message_start = _make_chunk(usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025))
|
||||
# ~50 visible chars ≈ ~12 tokens (anthropic-style tokenizer ballpark)
|
||||
text_chunks = [
|
||||
_make_chunk(content="Based on your question, I think the answer is "),
|
||||
_make_chunk(content="forty-two. Here is my reasoning: "),
|
||||
]
|
||||
chunks = [message_start, *text_chunks]
|
||||
completion_output = (
|
||||
"Based on your question, I think the answer is forty-two. "
|
||||
"Here is my reasoning: "
|
||||
)
|
||||
completion_output = "Based on your question, I think the answer is forty-two. Here is my reasoning: "
|
||||
|
||||
processor = ChunkProcessor(chunks=chunks, messages=[])
|
||||
usage = processor.calculate_usage(
|
||||
|
|
@ -149,9 +141,7 @@ class TestAnthropicCursorBug:
|
|||
|
||||
def test_cache_fields_preserved_from_message_start(self):
|
||||
"""cache_read / cache_creation come from message_start and must survive."""
|
||||
message_start_usage = Usage(
|
||||
prompt_tokens=1024, completion_tokens=1, total_tokens=1025
|
||||
)
|
||||
message_start_usage = Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
|
||||
# Anthropic puts these in message_start
|
||||
message_start_usage.cache_read_input_tokens = 512
|
||||
message_start_usage.cache_creation_input_tokens = 128
|
||||
|
|
@ -193,9 +183,7 @@ class TestAnthropicCursorBug:
|
|||
on a 1-token string also gives ~1, so billing is still approximately
|
||||
correct. This test pins that the result is sane (1 or 0).
|
||||
"""
|
||||
message_start = _make_chunk(
|
||||
usage=Usage(prompt_tokens=20, completion_tokens=1, total_tokens=21)
|
||||
)
|
||||
message_start = _make_chunk(usage=Usage(prompt_tokens=20, completion_tokens=1, total_tokens=21))
|
||||
text_chunk = _make_chunk(content="Yes.")
|
||||
# Anthropic's message_delta also gives output_tokens=1 in this case
|
||||
message_delta = _make_chunk(
|
||||
|
|
@ -231,9 +219,7 @@ class TestAnthropicCursorBug:
|
|||
must fire so token_counter estimates from completion text instead of
|
||||
billing the placeholder.
|
||||
"""
|
||||
message_start_usage = Usage(
|
||||
prompt_tokens=1024, completion_tokens=1, total_tokens=1025
|
||||
)
|
||||
message_start_usage = Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
|
||||
message_start_usage.cache_read_input_tokens = 4096
|
||||
message_start = _make_chunk(usage=message_start_usage)
|
||||
# Subsequent chunks with cache fields but no completion_tokens
|
||||
|
|
@ -253,6 +239,114 @@ class TestAnthropicCursorBug:
|
|||
"Reset to 0 forces token_counter fallback."
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("placeholder", [1, 3, 8])
|
||||
def test_interrupted_reasoning_only_stream_estimates_from_reasoning(self, placeholder: int):
|
||||
message_start = _make_chunk(
|
||||
usage=Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=placeholder,
|
||||
total_tokens=100 + placeholder,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=0, text_tokens=placeholder),
|
||||
)
|
||||
)
|
||||
reasoning_text = "Let me work through the scheduling constraints step by step. " * 40
|
||||
reasoning_chunks = [
|
||||
_make_chunk(reasoning_content=reasoning_text[i : i + 50]) for i in range(0, len(reasoning_text), 50)
|
||||
]
|
||||
|
||||
response = litellm.stream_chunk_builder(
|
||||
chunks=[message_start, *reasoning_chunks],
|
||||
messages=[{"role": "user", "content": "Plan the schedule."}],
|
||||
)
|
||||
|
||||
assert response.choices[0].message.reasoning_content == reasoning_text
|
||||
reasoning_tokens = response.usage.completion_tokens_details.reasoning_tokens
|
||||
assert reasoning_tokens > placeholder
|
||||
assert response.usage.completion_tokens == reasoning_tokens, (
|
||||
f"Expected completion_tokens to be the reasoning estimate, got "
|
||||
f"completion_tokens={response.usage.completion_tokens} reasoning_tokens={reasoning_tokens}"
|
||||
)
|
||||
assert response.usage.total_tokens == response.usage.prompt_tokens + reasoning_tokens
|
||||
details = response.usage.completion_tokens_details
|
||||
assert details.text_tokens + details.reasoning_tokens == response.usage.completion_tokens
|
||||
|
||||
def test_fallback_counts_reasoning_and_text_together(self):
|
||||
reasoning = "First I should check whether the input is sorted. " * 10
|
||||
text = "The list is already sorted, so no work is needed."
|
||||
chunks = [_make_chunk(reasoning_content=reasoning), _make_chunk(content=text)]
|
||||
|
||||
response = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "Sort it."}])
|
||||
|
||||
text_only = litellm.token_counter(model="claude-sonnet-4-6", text=text, count_response_tokens=True)
|
||||
details = response.usage.completion_tokens_details
|
||||
assert details.reasoning_tokens > 0
|
||||
assert response.usage.completion_tokens == text_only + details.reasoning_tokens
|
||||
assert details.text_tokens == text_only
|
||||
|
||||
def test_lone_usage_event_with_finish_reason_is_trusted(self):
|
||||
chunks = [
|
||||
_make_chunk(content="Yes, "),
|
||||
_make_chunk(content="that works."),
|
||||
_make_chunk(
|
||||
usage=Usage(prompt_tokens=20, completion_tokens=5, total_tokens=25),
|
||||
finish_reason="stop",
|
||||
),
|
||||
]
|
||||
processor = ChunkProcessor(chunks=chunks, messages=[])
|
||||
result = processor._calculate_usage_per_chunk(chunks=chunks)
|
||||
assert result["completion_tokens"] == 5
|
||||
|
||||
def test_dict_chunks_with_finish_reason_are_trusted(self):
|
||||
chunks = [
|
||||
{
|
||||
"_hidden_params": {"custom_llm_provider": "anthropic"},
|
||||
"choices": [{"delta": {"content": "Yes, "}, "finish_reason": None}],
|
||||
},
|
||||
{
|
||||
"_hidden_params": {"custom_llm_provider": "anthropic"},
|
||||
"choices": [{"delta": {"content": "that works."}, "finish_reason": "stop"}],
|
||||
"usage": Usage(prompt_tokens=20, completion_tokens=5, total_tokens=25),
|
||||
},
|
||||
]
|
||||
processor = ChunkProcessor(chunks=chunks, messages=[])
|
||||
result = processor._calculate_usage_per_chunk(chunks=chunks)
|
||||
assert result["completion_tokens"] == 5
|
||||
|
||||
def test_dict_chunks_without_finish_reason_reset_placeholder(self):
|
||||
chunks = [
|
||||
{
|
||||
"_hidden_params": {"custom_llm_provider": "anthropic"},
|
||||
"choices": [],
|
||||
"usage": Usage(prompt_tokens=20, completion_tokens=1, total_tokens=21),
|
||||
},
|
||||
{
|
||||
"_hidden_params": {"custom_llm_provider": "anthropic"},
|
||||
"choices": [{"delta": {"content": "partial"}, "finish_reason": None}],
|
||||
},
|
||||
]
|
||||
processor = ChunkProcessor(chunks=chunks, messages=[])
|
||||
result = processor._calculate_usage_per_chunk(chunks=chunks)
|
||||
assert result["completion_tokens"] == 0
|
||||
assert result["completion_tokens_details"] is None
|
||||
|
||||
def test_estimated_reasoning_is_capped_to_trusted_completion_total(self):
|
||||
chunks = [
|
||||
_make_chunk(reasoning_content="Let me reason about this carefully and at length. " * 20),
|
||||
_make_chunk(
|
||||
finish_reason="stop",
|
||||
usage=Usage(prompt_tokens=20, completion_tokens=5, total_tokens=25),
|
||||
),
|
||||
]
|
||||
response = litellm.stream_chunk_builder(
|
||||
chunks=chunks,
|
||||
messages=[{"role": "user", "content": "Go."}],
|
||||
)
|
||||
details = response.usage.completion_tokens_details
|
||||
assert response.usage.completion_tokens == 5
|
||||
assert details.reasoning_tokens <= response.usage.completion_tokens
|
||||
assert details.reasoning_tokens + details.text_tokens == response.usage.completion_tokens
|
||||
assert details.text_tokens >= 0
|
||||
|
||||
|
||||
class TestProviderGuard:
|
||||
"""Class A: the cursor-reset heuristic must NOT silently affect non-Anthropic
|
||||
|
|
@ -297,11 +391,12 @@ class TestNonAnthropicStreamingIntact:
|
|||
"""Make sure providers without cursor pattern still work."""
|
||||
|
||||
def test_completion_tokens_above_one_never_resets(self):
|
||||
"""Any chunk reporting completion_tokens > 1 sets saw_non_cursor
|
||||
and prevents the reset."""
|
||||
"""A non-Anthropic provider reporting completion_tokens > 1 from a
|
||||
single usage event keeps that value."""
|
||||
chunks = [
|
||||
_make_chunk(
|
||||
usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
|
||||
usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
|
||||
custom_llm_provider="openai",
|
||||
),
|
||||
]
|
||||
processor = ChunkProcessor(chunks=chunks, messages=[])
|
||||
|
|
|
|||
|
|
@ -677,3 +677,14 @@ def test_azure_responses_gpt6_astra_rejects_temperature_while_reasoning(local_mo
|
|||
model="gpt-6-astra",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
|
||||
def test_azure_responses_sends_the_deployment_name_when_azure_ai_prefix_survives_provider_remap():
|
||||
request = AzureOpenAIResponsesAPIConfig().transform_responses_api_request(
|
||||
model="azure_ai/gpt-5.4-nano",
|
||||
input="hi",
|
||||
response_api_optional_request_params={},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert request["model"] == "gpt-5.4-nano"
|
||||
|
|
|
|||
|
|
@ -0,0 +1,311 @@
|
|||
import json
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.responses.transformation import AzureAIResponsesAPIConfig
|
||||
from litellm.responses.main import _will_bridge_to_chat_completions
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
FOUNDRY_PROJECT_BASE = "https://res.services.ai.azure.com/api/projects/proj"
|
||||
FOUNDRY_RESPONSES_URL = f"{FOUNDRY_PROJECT_BASE}/openai/v1/responses"
|
||||
SERVERLESS_BASE = "https://endpoint.eastus.models.ai.azure.com"
|
||||
WEATHER_TOOL = {
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_azure_ai_env(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "api_base", None)
|
||||
monkeypatch.setattr(litellm, "api_key", None)
|
||||
monkeypatch.setattr(litellm, "enable_azure_ad_token_refresh", False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
for env_var in (
|
||||
"AZURE_AI_API_BASE",
|
||||
"AZURE_AI_API_KEY",
|
||||
"AZURE_AD_TOKEN",
|
||||
"AZURE_TENANT_ID",
|
||||
"AZURE_CLIENT_ID",
|
||||
"AZURE_CLIENT_SECRET",
|
||||
):
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
|
||||
|
||||
def _responses_payload(model: str) -> dict:
|
||||
return {
|
||||
"id": "resp_123",
|
||||
"object": "response",
|
||||
"created_at": 1741369938,
|
||||
"status": "completed",
|
||||
"model": model,
|
||||
"output": [],
|
||||
"parallel_tool_calls": False,
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
|
||||
"error": None,
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"metadata": None,
|
||||
"temperature": None,
|
||||
"top_p": None,
|
||||
"max_output_tokens": None,
|
||||
"previous_response_id": None,
|
||||
"reasoning": None,
|
||||
"truncation": None,
|
||||
"instructions": None,
|
||||
"incomplete_details": None,
|
||||
"user": None,
|
||||
}
|
||||
|
||||
|
||||
def _chat_completion_payload(model: str) -> dict:
|
||||
return {
|
||||
"id": "chatcmpl-123",
|
||||
"object": "chat.completion",
|
||||
"created": 1741369938,
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["gpt-5.6-luna-20260710154139", "gpt-5.5-20260504143601", "DeepSeek-R1-0528", None])
|
||||
@pytest.mark.parametrize(
|
||||
"api_base", [FOUNDRY_PROJECT_BASE, "https://res.services.ai.azure.com", "https://res.openai.azure.com"]
|
||||
)
|
||||
def test_azure_openai_v1_hosts_resolve_native_config(model, api_base):
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(provider="azure_ai", model=model, api_base=api_base)
|
||||
assert isinstance(config, AzureAIResponsesAPIConfig)
|
||||
|
||||
|
||||
def test_api_base_from_env_resolves_native_config(monkeypatch):
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", FOUNDRY_PROJECT_BASE)
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(provider="azure_ai", model="gpt-5.6-luna", api_base=None)
|
||||
assert isinstance(config, AzureAIResponsesAPIConfig)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["gpt-5.6-luna", None])
|
||||
@pytest.mark.parametrize(
|
||||
"api_base",
|
||||
[SERVERLESS_BASE, "https://endpoint.eastus.inference.ml.azure.com/score", "https://res.cognitiveservices.azure.com"],
|
||||
)
|
||||
def test_other_hosts_keep_chat_bridge(model, api_base):
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(provider="azure_ai", model=model, api_base=api_base)
|
||||
assert config is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["claude-3-5-sonnet", "model_router/gpt-5", "agents/my-agent"])
|
||||
def test_non_openai_surfaces_keep_chat_bridge(model):
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="azure_ai", model=model, api_base=FOUNDRY_PROJECT_BASE
|
||||
)
|
||||
assert config is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("api_base,bridged", [(FOUNDRY_PROJECT_BASE, False), (SERVERLESS_BASE, True)])
|
||||
def test_will_bridge_to_chat_completions_follows_host(api_base, bridged):
|
||||
assert _will_bridge_to_chat_completions("gpt-5.6-luna", "azure_ai", False, None, api_base) is bridged
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base,expected",
|
||||
[
|
||||
(FOUNDRY_PROJECT_BASE, FOUNDRY_RESPONSES_URL),
|
||||
(f"{FOUNDRY_PROJECT_BASE}/", FOUNDRY_RESPONSES_URL),
|
||||
(f"{FOUNDRY_PROJECT_BASE}/openai/v1", FOUNDRY_RESPONSES_URL),
|
||||
(FOUNDRY_RESPONSES_URL, FOUNDRY_RESPONSES_URL),
|
||||
("https://res.services.ai.azure.com", "https://res.services.ai.azure.com/openai/v1/responses"),
|
||||
("https://res.services.ai.azure.com/models", "https://res.services.ai.azure.com/openai/v1/responses"),
|
||||
(
|
||||
"https://res.services.ai.azure.com/models/chat/completions?api-version=2024-05-01-preview",
|
||||
"https://res.services.ai.azure.com/openai/v1/responses",
|
||||
),
|
||||
("https://res.openai.azure.com", "https://res.openai.azure.com/openai/v1/responses"),
|
||||
(
|
||||
"https://res.openai.azure.com/openai/deployments/gpt-5?api-version=2025-04-01-preview",
|
||||
"https://res.openai.azure.com/openai/v1/responses",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_get_complete_url(api_base, expected):
|
||||
assert AzureAIResponsesAPIConfig().get_complete_url(api_base=api_base, litellm_params={}) == expected
|
||||
|
||||
|
||||
def test_get_complete_url_ignores_api_version():
|
||||
url = AzureAIResponsesAPIConfig().get_complete_url(
|
||||
api_base=FOUNDRY_PROJECT_BASE, litellm_params={"api_version": "2025-04-01-preview"}
|
||||
)
|
||||
assert url == FOUNDRY_RESPONSES_URL
|
||||
|
||||
|
||||
def test_get_complete_url_uses_env_api_base(monkeypatch):
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", FOUNDRY_PROJECT_BASE)
|
||||
assert AzureAIResponsesAPIConfig().get_complete_url(api_base=None, litellm_params={}) == FOUNDRY_RESPONSES_URL
|
||||
|
||||
|
||||
def test_get_complete_url_raises_without_api_base():
|
||||
with pytest.raises(ValueError, match="AZURE_AI_API_BASE"):
|
||||
AzureAIResponsesAPIConfig().get_complete_url(api_base=None, litellm_params={})
|
||||
|
||||
|
||||
def test_native_websocket_stays_off():
|
||||
assert AzureAIResponsesAPIConfig().supports_native_websocket() is False
|
||||
|
||||
|
||||
def test_validate_environment_sends_api_key_header():
|
||||
headers = AzureAIResponsesAPIConfig().validate_environment(
|
||||
headers={"x-custom": "1"},
|
||||
model="gpt-5.6-luna",
|
||||
litellm_params=GenericLiteLLMParams(api_key="secret", api_base=FOUNDRY_PROJECT_BASE),
|
||||
)
|
||||
assert headers == {"x-custom": "1", "api-key": "secret", "Content-Type": "application/json"}
|
||||
|
||||
|
||||
def test_validate_environment_reads_api_key_from_env(monkeypatch):
|
||||
monkeypatch.setenv("AZURE_AI_API_KEY", "env-secret")
|
||||
headers = AzureAIResponsesAPIConfig().validate_environment(
|
||||
headers={}, model="gpt-5.6-luna", litellm_params=GenericLiteLLMParams(api_base=FOUNDRY_PROJECT_BASE)
|
||||
)
|
||||
assert headers["api-key"] == "env-secret"
|
||||
|
||||
|
||||
def test_validate_environment_uses_entra_token_without_api_key():
|
||||
headers = AzureAIResponsesAPIConfig().validate_environment(
|
||||
headers={},
|
||||
model="gpt-5.6-luna",
|
||||
litellm_params=GenericLiteLLMParams(azure_ad_token="entra-token", api_base=FOUNDRY_PROJECT_BASE),
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer entra-token"
|
||||
assert "api-key" not in headers
|
||||
|
||||
|
||||
def test_validate_environment_raises_without_credentials():
|
||||
with pytest.raises(ValueError, match="AZURE_AI_API_KEY"):
|
||||
AzureAIResponsesAPIConfig().validate_environment(
|
||||
headers={}, model="gpt-5.6-luna", litellm_params=GenericLiteLLMParams(api_base=FOUNDRY_PROJECT_BASE)
|
||||
)
|
||||
|
||||
|
||||
NATIVE_RESPONSES_CASES = [
|
||||
("azure_ai/gpt-5.6-luna-20260710154139", FOUNDRY_PROJECT_BASE, FOUNDRY_RESPONSES_URL, "gpt-5.6-luna-20260710154139"),
|
||||
(
|
||||
"azure_ai/gpt-5.6-luna",
|
||||
"https://res.services.ai.azure.com/models",
|
||||
"https://res.services.ai.azure.com/openai/v1/responses",
|
||||
"gpt-5.6-luna",
|
||||
),
|
||||
(
|
||||
"azure_ai/gpt-5.6-sol",
|
||||
"https://res.services.ai.azure.com",
|
||||
"https://res.services.ai.azure.com/openai/v1/responses",
|
||||
"gpt-5.6-sol",
|
||||
),
|
||||
(
|
||||
"azure_ai/gpt-5.6-luna-20260710154139",
|
||||
"https://res.openai.azure.com",
|
||||
"https://res.openai.azure.com/openai/v1/responses",
|
||||
"gpt-5.6-luna-20260710154139",
|
||||
),
|
||||
(
|
||||
"azure_ai/gpt-5.6-sol",
|
||||
"https://res.openai.azure.com",
|
||||
"https://res.openai.azure.com/openai/v1/responses",
|
||||
"gpt-5.6-sol",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _assert_native_responses_request(route, expected_url, expected_model):
|
||||
request = route.calls.last.request
|
||||
body = json.loads(request.content)
|
||||
assert f"{request.url.scheme}://{request.url.host}{request.url.path}" == expected_url
|
||||
assert request.headers["api-key"] == "fake-key"
|
||||
assert body["model"] == expected_model
|
||||
assert body["input"] == "What is the weather in SF?"
|
||||
assert "messages" not in body
|
||||
assert body["reasoning"] == {"effort": "high"}
|
||||
assert body["tools"] == [WEATHER_TOOL]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
@pytest.mark.parametrize("model,api_base,expected_url,expected_model", NATIVE_RESPONSES_CASES)
|
||||
async def test_aresponses_sends_reasoning_and_tools_to_native_endpoint(model, api_base, expected_url, expected_model):
|
||||
route = respx.post(url__regex=r".*/openai/v1/responses(\?.*)?$").mock(
|
||||
return_value=httpx.Response(200, json=_responses_payload(expected_model))
|
||||
)
|
||||
|
||||
await litellm.aresponses(
|
||||
model=model,
|
||||
input="What is the weather in SF?",
|
||||
reasoning_effort="high",
|
||||
tools=[WEATHER_TOOL],
|
||||
api_base=api_base,
|
||||
api_key="fake-key",
|
||||
)
|
||||
|
||||
_assert_native_responses_request(route, expected_url, expected_model)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
async def test_aresponses_catalog_name_remapped_to_azure_sends_bare_deployment_name(monkeypatch):
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", "https://res.openai.azure.com")
|
||||
route = respx.post(url__regex=r".*/openai/v1/responses(\?.*)?$").mock(
|
||||
return_value=httpx.Response(200, json=_responses_payload("gpt-5.4-nano"))
|
||||
)
|
||||
|
||||
await litellm.aresponses(
|
||||
model="azure_ai/gpt-5.4-nano",
|
||||
input="What is the weather in SF?",
|
||||
api_base="https://res.openai.azure.com",
|
||||
api_key="fake-key",
|
||||
)
|
||||
|
||||
assert json.loads(route.calls.last.request.content)["model"] == "gpt-5.4-nano"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
@pytest.mark.parametrize("model,api_base,expected_url,expected_model", NATIVE_RESPONSES_CASES)
|
||||
async def test_router_aresponses_sends_bare_deployment_name(model, api_base, expected_url, expected_model):
|
||||
route = respx.post(url__regex=r".*/openai/v1/responses(\?.*)?$").mock(
|
||||
return_value=httpx.Response(200, json=_responses_payload(expected_model))
|
||||
)
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": model, "api_base": api_base, "api_key": "fake-key"}}],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
await router.aresponses(
|
||||
model="gpt-5.6", input="What is the weather in SF?", reasoning={"effort": "high"}, tools=[WEATHER_TOOL]
|
||||
)
|
||||
|
||||
_assert_native_responses_request(route, expected_url, expected_model)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
async def test_aresponses_serverless_host_stays_on_chat_bridge():
|
||||
chat_route = respx.post(url__regex=r".*/chat/completions$").mock(
|
||||
return_value=httpx.Response(200, json=_chat_completion_payload("gpt-5.6-luna"))
|
||||
)
|
||||
responses_route = respx.post(url__regex=r".*/responses$")
|
||||
|
||||
await litellm.aresponses(
|
||||
model="azure_ai/gpt-5.6-luna-20260710154139",
|
||||
input="What is the weather in SF?",
|
||||
tools=[WEATHER_TOOL],
|
||||
api_base=SERVERLESS_BASE,
|
||||
api_key="fake-key",
|
||||
)
|
||||
|
||||
assert chat_route.called
|
||||
assert not responses_route.called
|
||||
assert chat_route.calls.last.request.headers["Authorization"] == "Bearer fake-key"
|
||||
|
|
@ -164,6 +164,21 @@ class TestDashScopeConfig:
|
|||
|
||||
assert transformed_messages[0].get("cache_control") == {"type": "ephemeral"}
|
||||
|
||||
@pytest.mark.parametrize("reasoning_effort", ["none", "minimal", "low", "high"])
|
||||
def test_dashscope_forwards_reasoning_effort(self, reasoning_effort: str):
|
||||
"""DashScope supports reasoning_effort, so it must reach the provider instead of being dropped."""
|
||||
assert "reasoning_effort" in DashScopeChatConfig().get_supported_openai_params(
|
||||
model="qwen3.7-plus"
|
||||
)
|
||||
|
||||
optional_params = litellm.get_optional_params(
|
||||
model="qwen3.7-plus",
|
||||
custom_llm_provider="dashscope",
|
||||
reasoning_effort=reasoning_effort,
|
||||
)
|
||||
|
||||
assert optional_params["reasoning_effort"] == reasoning_effort
|
||||
|
||||
def test_dashscope_preserves_cache_control_in_tools(self):
|
||||
"""DashScope should NOT strip cache_control from tools."""
|
||||
config = DashScopeChatConfig()
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
|
|||
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
|
||||
from litellm.llms.openai.openai import OpenAIConfig
|
||||
from litellm.utils import (
|
||||
_is_explicitly_disabled_factory,
|
||||
is_explicitly_disabled_factory,
|
||||
peek_reasoning_summary_aliases,
|
||||
strip_reasoning_summary_aliases_from_optional_params,
|
||||
)
|
||||
|
|
@ -524,19 +524,19 @@ def test_gpt5_minimal_explicitly_disabled_check(gpt5_config: OpenAIGPT5Config):
|
|||
|
||||
|
||||
def test_is_explicitly_disabled_factory_minimal():
|
||||
"""_is_explicitly_disabled_factory returns True only for explicit False entries.
|
||||
"""is_explicitly_disabled_factory returns True only for explicit False entries.
|
||||
|
||||
Verifies the shared helper used by _is_reasoning_effort_level_explicitly_disabled
|
||||
directly — so future changes to the helper are caught without going through the
|
||||
method wrapper.
|
||||
"""
|
||||
key = "supports_minimal_reasoning_effort"
|
||||
assert _is_explicitly_disabled_factory("gpt-5.4-mini", None, key)
|
||||
assert _is_explicitly_disabled_factory("gpt-5.4-nano", None, key)
|
||||
assert _is_explicitly_disabled_factory("openai/gpt-5.4-mini", None, key)
|
||||
assert _is_explicitly_disabled_factory("gpt-5.4", None, key)
|
||||
assert _is_explicitly_disabled_factory("gpt-5.4-pro", None, key)
|
||||
assert not _is_explicitly_disabled_factory("gpt-5.4-turbo-preview", None, key)
|
||||
assert is_explicitly_disabled_factory("gpt-5.4-mini", None, key)
|
||||
assert is_explicitly_disabled_factory("gpt-5.4-nano", None, key)
|
||||
assert is_explicitly_disabled_factory("openai/gpt-5.4-mini", None, key)
|
||||
assert is_explicitly_disabled_factory("gpt-5.4", None, key)
|
||||
assert is_explicitly_disabled_factory("gpt-5.4-pro", None, key)
|
||||
assert not is_explicitly_disabled_factory("gpt-5.4-turbo-preview", None, key)
|
||||
|
||||
|
||||
def test_gpt5_unknown_model_passes_through_minimal(config: OpenAIConfig):
|
||||
|
|
|
|||
|
|
@ -1108,3 +1108,52 @@ def test_get_optional_params_preserves_max_for_declared_levels_model():
|
|||
)
|
||||
|
||||
assert optional_params["reasoning_effort"] == "max"
|
||||
|
||||
|
||||
def _together_chat_transport() -> tuple[HTTPHandler, list[httpx.Request]]:
|
||||
captured_requests: list[httpx.Request] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
captured_requests.append(request)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "chatcmpl-together",
|
||||
"object": "chat.completion",
|
||||
"created": 1234567890,
|
||||
"model": TOOL_CALLING_MODEL,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Hello!"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
||||
},
|
||||
)
|
||||
|
||||
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond)))
|
||||
return client, captured_requests
|
||||
|
||||
|
||||
def test_custom_role_wrappers_never_reach_the_request():
|
||||
client, captured_requests = _together_chat_transport()
|
||||
messages = [{"role": "user", "content": "Hello!"}]
|
||||
|
||||
litellm.completion(
|
||||
model=f"together_ai/{TOOL_CALLING_MODEL}",
|
||||
messages=messages,
|
||||
roles={
|
||||
"system": {"pre_message": "<|im_start|>system\n", "post_message": "<|im_end|>"},
|
||||
"assistant": {"pre_message": "<|im_start|>assistant\n", "post_message": "<|im_end|>"},
|
||||
"user": {"pre_message": "<|im_start|>user\n", "post_message": "<|im_end|>"},
|
||||
},
|
||||
api_key="fake-key",
|
||||
client=client,
|
||||
)
|
||||
|
||||
request_body = json.loads(captured_requests[0].content)
|
||||
assert request_body["messages"] == messages
|
||||
assert "prompt" not in request_body
|
||||
assert "roles" not in request_body
|
||||
|
|
|
|||
|
|
@ -5,11 +5,14 @@ from copy import deepcopy
|
|||
from typing import Final, List, cast
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm import ModelResponse, completion
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages import handler as anthropic_messages_handler
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.gemini.chat.transformation import GoogleAIStudioGeminiConfig
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAIError
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
|
|
@ -2678,6 +2681,118 @@ def test_reasoning_effort_maps_to_thinking_level_gemini_3():
|
|||
assert result["thinkingConfig"]["includeThoughts"] is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"gemini-3.7-flash",
|
||||
"vertex_ai/gemini-3.8-flash",
|
||||
"gemini/gemini-3.8-flash",
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
("reasoning_effort", "include_thoughts"),
|
||||
[("minimal", True), ("none", False), ("disable", False)],
|
||||
)
|
||||
def test_gemini_37_38_flash_floor_minimal_thinking_level(
|
||||
local_model_cost_map, model, reasoning_effort, include_thoughts
|
||||
):
|
||||
result = VertexGeminiConfig._map_reasoning_effort_to_thinking_level(
|
||||
reasoning_effort, model
|
||||
)
|
||||
|
||||
assert result["thinkingLevel"] == "low"
|
||||
assert result["includeThoughts"] is include_thoughts
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "reasoning_effort", "expected_level", "include_thoughts"),
|
||||
[
|
||||
("gemini-3-flash-preview", "minimal", "minimal", True),
|
||||
("gemini-3-flash-preview", "none", "minimal", False),
|
||||
("gemini-3-flash-preview", "disable", "minimal", False),
|
||||
("gemini-3.6-flash", "minimal", "minimal", True),
|
||||
("gemini-3.6-flash", "none", "minimal", False),
|
||||
("gemini-3.6-flash", "disable", "minimal", False),
|
||||
("gemini-3.5-flash", "minimal", "minimal", True),
|
||||
("gemini-3.5-flash", "none", "minimal", False),
|
||||
("gemini-3.5-flash", "disable", "minimal", False),
|
||||
("gemini-3.8-flash", "medium", "medium", True),
|
||||
],
|
||||
)
|
||||
def test_gemini_flash_minimal_thinking_support(
|
||||
local_model_cost_map, model, reasoning_effort, expected_level, include_thoughts
|
||||
):
|
||||
result = VertexGeminiConfig._map_reasoning_effort_to_thinking_level(
|
||||
reasoning_effort, model
|
||||
)
|
||||
|
||||
assert result["thinkingLevel"] == expected_level
|
||||
assert result["includeThoughts"] is include_thoughts
|
||||
|
||||
|
||||
def test_gemini_38_flash_feature_flag_uses_low_thinking_level(local_model_cost_map, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_gemini_default_thinking_level_low", True)
|
||||
thinking_param = {"type": "enabled", "budget_tokens": 1024}
|
||||
|
||||
result_38 = VertexGeminiConfig._map_thinking_param(
|
||||
thinking_param, model="gemini-3.8-flash"
|
||||
)
|
||||
result_36 = VertexGeminiConfig._map_thinking_param(
|
||||
thinking_param, model="gemini-3.6-flash"
|
||||
)
|
||||
|
||||
assert result_38["thinkingLevel"] == "low"
|
||||
assert result_36["thinkingLevel"] == "minimal"
|
||||
|
||||
|
||||
def test_gemini_38_flash_public_reasoning_effort_none_uses_low(local_model_cost_map):
|
||||
result = VertexGeminiConfig().map_openai_params(
|
||||
non_default_params={"reasoning_effort": "none"},
|
||||
optional_params={},
|
||||
model="gemini-3.8-flash",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["thinkingConfig"] == {
|
||||
"thinkingLevel": "low",
|
||||
"includeThoughts": False,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_38_flash_messages_bridge_thinking_disabled_sends_low_thinking_level(local_model_cost_map):
|
||||
captured: dict[str, dict] = {}
|
||||
|
||||
def upstream(request: httpx.Request) -> httpx.Response:
|
||||
captured["body"] = json.loads(request.content)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"candidates": [{"content": {"parts": [{"text": "hi"}], "role": "model"}, "finishReason": "STOP"}],
|
||||
"usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2},
|
||||
},
|
||||
request=request,
|
||||
)
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=httpx.MockTransport(upstream))
|
||||
|
||||
await anthropic_messages_handler.anthropic_messages(
|
||||
max_tokens=16,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
model="gemini/gemini-3.8-flash",
|
||||
custom_llm_provider="gemini",
|
||||
thinking={"type": "disabled"},
|
||||
api_key="fake-gemini-key",
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert captured["body"]["generationConfig"]["thinkingConfig"] == {
|
||||
"thinkingLevel": "low",
|
||||
"includeThoughts": False,
|
||||
}
|
||||
|
||||
|
||||
def test_reasoning_effort_dict_format_gemini_3():
|
||||
"""
|
||||
Test that reasoning_effort works when passed as dict format from OpenAI Agents SDK.
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ maps each CredError onto its HTTP status. These pin the parity-critical mapping
|
|||
|
||||
import base64
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -20,7 +21,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import
|
|||
raise_user_oauth_challenge,
|
||||
to_server_spec,
|
||||
to_subject,
|
||||
validate_static_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
ApiKeyConfig,
|
||||
AuthorizationCodeConfig,
|
||||
|
|
@ -34,10 +37,28 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
SharedKey,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
@pytest.mark.parametrize("auth_type,header,value", [
|
||||
(MCPAuth.api_key, "Authorization", "Bearer fixture-key"),
|
||||
(MCPAuth.api_key, "Authorization", "ApiKey fixture-key"),
|
||||
(MCPAuth.api_key, "Authorization", "token fixture-key"),
|
||||
(MCPAuth.api_key, "Authorization", "Bearer token"),
|
||||
(MCPAuth.api_key, "Authorization", "opaque-key"),
|
||||
(MCPAuth.api_key, "Authorization", "Custom Custom"),
|
||||
(MCPAuth.api_key, "X-API-Key", "Bearer Bearer"),
|
||||
(MCPAuth.api_key, "X-Custom", "ApiKey ApiKey"),
|
||||
(MCPAuth.authorization, "Authorization", "opaque-secret-value"),
|
||||
])
|
||||
def test_static_credential_preserves_supported_api_key_and_raw_headers(
|
||||
auth_type: MCPAuthType, header: str, value: str,
|
||||
) -> None:
|
||||
result: Final = validate_static_credential(auth_type, {header: value}, upstream_token_header=header)
|
||||
assert isinstance(result, Ok)
|
||||
|
||||
|
||||
def _server(**kwargs) -> MCPServer:
|
||||
return MCPServer(server_id="s", name="n", transport=MCPTransport.http, **kwargs)
|
||||
|
||||
|
|
@ -155,12 +176,6 @@ def test_oauth2_user_token_maps_to_authorization_code(oauth2_flow):
|
|||
_server(auth_type=MCPAuth.api_key), # no token configured
|
||||
_server(auth_type=MCPAuth.bearer_token), # no token configured
|
||||
_server(auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True), # delegated upstream OAuth -> v1
|
||||
_server(auth_type=MCPAuth.oauth2_token_exchange), # no endpoint/client creds -> incomplete -> v1
|
||||
_server(
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
token_exchange_endpoint="https://idp/token",
|
||||
client_id="cid",
|
||||
), # missing client_secret -> incomplete -> v1
|
||||
_server(auth_type=MCPAuth.aws_sigv4),
|
||||
_server(auth_type=None, oauth_passthrough=True, extra_headers=["Authorization"]),
|
||||
],
|
||||
|
|
@ -802,3 +817,14 @@ def test_a_blank_header_name_means_unset_rather_than_an_error(blank):
|
|||
spec = to_server_spec(server)
|
||||
assert spec is not None
|
||||
assert spec.config.header_name == "Authorization"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client_secret", [None, ""])
|
||||
@pytest.mark.parametrize("is_byok", [False, True])
|
||||
def test_incomplete_obo_keeps_exchange_ownership(client_secret: str | None, is_byok: bool) -> None:
|
||||
spec = to_server_spec(_server(auth_type=MCPAuth.oauth2_token_exchange, client_id="client",
|
||||
client_secret=client_secret, is_byok=is_byok))
|
||||
assert spec is not None
|
||||
assert isinstance(spec.config, TokenExchangeConfig)
|
||||
assert spec.config.client_id == "client"
|
||||
assert spec.config.client_secret is None
|
||||
|
|
|
|||
|
|
@ -5,11 +5,13 @@ import logging
|
|||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Final, Literal, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from respx import MockRouter
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPServerListError,
|
||||
|
|
@ -5127,7 +5129,8 @@ class TestMCPServerManager:
|
|||
captured: dict = {}
|
||||
|
||||
def fake_create_tool_function(
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False,
|
||||
auth_type=None, upstream_token_header=None,
|
||||
):
|
||||
captured["headers"] = headers
|
||||
captured["server_label"] = server_label
|
||||
|
|
@ -5212,7 +5215,8 @@ class TestMCPServerManager:
|
|||
captured: dict = {}
|
||||
|
||||
def fake_create_tool_function(
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False,
|
||||
auth_type=None, upstream_token_header=None,
|
||||
):
|
||||
captured["headers"] = headers
|
||||
|
||||
|
|
@ -9401,12 +9405,13 @@ class TestCreateMcpClientV2Graft:
|
|||
assert "misconfigured" in str(exc_info.value.detail)
|
||||
assert "token_url" in str(exc_info.value.detail)
|
||||
|
||||
async def test_static_token_missing_defers_to_v1(self):
|
||||
client = await MCPServerManager()._create_mcp_client(
|
||||
self._http_server(auth_type=MCPAuth.api_key, authentication_token=None)
|
||||
)
|
||||
|
||||
assert client._resolved_auth is None
|
||||
async def test_static_token_missing_rejects_before_connecting(self):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(
|
||||
self._http_server(auth_type=MCPAuth.api_key, authentication_token=None)
|
||||
)
|
||||
assert exc.value.status_code == 500
|
||||
assert "credential" in str(exc.value.detail)
|
||||
|
||||
async def test_stdio_migrated_auth_type_still_defers_to_v1(self):
|
||||
client = await MCPServerManager()._create_mcp_client(
|
||||
|
|
@ -13467,3 +13472,400 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them(
|
|||
result: Final = await cache.get(("server", None), fetch)
|
||||
assert result[0].description == description
|
||||
assert fetch.await_count == 2
|
||||
|
||||
|
||||
class TestProtectedCredentialPreparation:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,credential", [
|
||||
(MCPAuth.bearer_token, None),
|
||||
(MCPAuth.bearer_token, "Bearer"),
|
||||
(MCPAuth.api_key, None),
|
||||
(MCPAuth.basic, "Basic"),
|
||||
])
|
||||
@pytest.mark.parametrize("dispatch", ["managed", "local"])
|
||||
async def test_openapi_dispatch_rejects_unusable_effective_credentials(
|
||||
self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
auth_type: MCPAuthType, credential: str | None, dispatch: str,
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool
|
||||
from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix
|
||||
|
||||
spec_path: Final = tmp_path / "openapi.json"
|
||||
spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"},
|
||||
"paths": {"/echo": {"get": {"operationId": "echo"}}}}))
|
||||
server: Final = MCPServer(
|
||||
server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example",
|
||||
transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential,
|
||||
)
|
||||
manager: Final = MCPServerManager()
|
||||
await manager._register_openapi_tools(str(spec_path), server, server.url)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="unexpected success")
|
||||
result: Final = (
|
||||
await manager._call_openapi_tool_handler(server, "echo", {})
|
||||
if dispatch == "managed"
|
||||
else await _handle_local_mcp_tool(add_server_prefix_to_name("echo", get_server_prefix(server)), {})
|
||||
)
|
||||
assert result.isError is True
|
||||
assert "requires a usable upstream credential" in result.content[0].text
|
||||
assert destination.call_count == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse])
|
||||
@pytest.mark.parametrize("client_secret", [None, ""])
|
||||
@pytest.mark.parametrize("subject", [None, "caller-subject"])
|
||||
async def test_incomplete_obo_rejects_caller_and_static_fallback(
|
||||
self, transport: MCPTransport, client_secret: str | None, subject: str | None
|
||||
) -> None:
|
||||
server = MCPServer(
|
||||
server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp",
|
||||
transport=transport, auth_type=MCPAuth.oauth2_token_exchange,
|
||||
client_id="gateway", client_secret=client_secret,
|
||||
token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header="Bearer override", subject_token=subject,
|
||||
)
|
||||
assert exc.value.status_code == (401 if subject is None else 500)
|
||||
assert "static-fallback" not in str(exc.value.detail)
|
||||
assert "override" not in str(exc.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.api_key, MCPAuth.bearer_token])
|
||||
@pytest.mark.parametrize("credential", [None, "", " ", {"X-Trace": "trace"}])
|
||||
async def test_static_auth_without_usable_credential_rejects(
|
||||
self, auth_type: MCPAuthType, credential: str | dict[str, str] | None
|
||||
) -> None:
|
||||
server = MCPServer(
|
||||
server_id="empty-static", name="empty-static", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential)
|
||||
assert exc.value.status_code == 500
|
||||
assert "credential" in str(exc.value.detail).lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,headers", [
|
||||
(MCPAuth.api_key, {"X-API-Key": "key"}),
|
||||
(MCPAuth.bearer_token, {"Authorization": "Bearer token"}),
|
||||
])
|
||||
async def test_static_auth_accepts_actual_forwarded_credential(
|
||||
self, auth_type: MCPAuthType, headers: dict[str, str]
|
||||
) -> None:
|
||||
server = MCPServer(
|
||||
server_id="header-static", name="header-static", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type,
|
||||
)
|
||||
client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers)
|
||||
assert client._get_auth_headers() == headers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange])
|
||||
async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None:
|
||||
server = MCPServer(
|
||||
server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type,
|
||||
token_exchange_endpoint="https://idp.example/token",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager().resolve_openapi_upstream_auth(
|
||||
mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None,
|
||||
user_api_key_auth=None, forwarded_headers=None,
|
||||
)
|
||||
assert exc.value.status_code in (401, 500)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,slot,value", [
|
||||
(MCPAuth.api_key, "X-API-Key", "token"),
|
||||
(MCPAuth.authorization, "Authorization", "opaque-secret-value"),
|
||||
(MCPAuth.authorization, "Authorization", "Bearer abc"),
|
||||
(MCPAuth.authorization, "Authorization", "Custom abc"),
|
||||
])
|
||||
async def test_raw_static_credentials_are_forwarded_unchanged(
|
||||
self, auth_type: MCPAuthType, slot: str, value: str,
|
||||
) -> None:
|
||||
server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type, authentication_token=value)
|
||||
client = await MCPServerManager()._create_mcp_client(server)
|
||||
assert client._resolved_auth is not None
|
||||
request = httpx.Request("GET", server.url)
|
||||
flow = client._resolved_auth.auth_flow(request)
|
||||
try:
|
||||
assert next(flow).headers[slot] == value
|
||||
finally:
|
||||
flow.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller", "forwarded"])
|
||||
async def test_raw_authorization_rejects_bare_schemes_before_dispatch(
|
||||
self, respx_mock: MockRouter, value: str, source: str,
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.authorization,
|
||||
authentication_token=value if source == "configured" else None,
|
||||
)
|
||||
destination: Final = respx_mock.route().respond(200)
|
||||
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
|
||||
await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header=value if source == "caller" else None,
|
||||
extra_headers={"Authorization": value} if source == "forwarded" else None,
|
||||
)
|
||||
assert exc.value.status_code == 500
|
||||
assert destination.call_count == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None:
|
||||
server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True,
|
||||
token_exchange_endpoint="https://idp.example/token")
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override")
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")])
|
||||
async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None:
|
||||
server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured)
|
||||
client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override)
|
||||
assert client._get_auth_headers()["Authorization"] == override
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("token", [None, "shared"])
|
||||
async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None:
|
||||
server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_slot_uses_its_actual_credential(self) -> None:
|
||||
server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key,
|
||||
upstream_token_header="X-Custom", authentication_token="key")
|
||||
client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"})
|
||||
assert client._credential_slot == "X-Custom"
|
||||
assert await client.discovery_auth_fingerprint()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("static,forwarded,caller", [
|
||||
({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None),
|
||||
({}, {"X-API-Key": "forwarded"}, None),
|
||||
({}, None, "ApiKey caller"),
|
||||
({"X-API-Key": "static"}, {"Authorization": ""}, None),
|
||||
])
|
||||
async def test_openapi_static_credentials_remain_supported(
|
||||
self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header, _request_extra_headers, create_tool_function,
|
||||
)
|
||||
tool: Final = create_tool_function(
|
||||
"/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key,
|
||||
)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
caller_token: Final = _request_auth_header.set(caller)
|
||||
extra_token: Final = _request_extra_headers.set(forwarded)
|
||||
try:
|
||||
assert await tool() == "authenticated"
|
||||
sent: Final = destination.calls.last.request.headers
|
||||
assert sent.get("x-api-key") == static.get("X-API-Key", (forwarded or {}).get("X-API-Key"))
|
||||
if caller:
|
||||
assert sent["authorization"] == caller
|
||||
assert destination.call_count == 1
|
||||
finally:
|
||||
_request_auth_header.reset(caller_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_static_resolution_cancellation_closes_flow(self) -> None:
|
||||
from collections.abc import AsyncGenerator
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import prepare_mcp_client
|
||||
|
||||
class CancelledAuth(httpx.Auth):
|
||||
closed = False
|
||||
|
||||
async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]:
|
||||
try:
|
||||
raise asyncio.CancelledError()
|
||||
yield request
|
||||
finally:
|
||||
self.closed = True
|
||||
|
||||
auth = CancelledAuth()
|
||||
server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key)
|
||||
client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth)
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await prepare_mcp_client(server, client)
|
||||
assert auth.closed
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization])
|
||||
async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None:
|
||||
server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ")
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="])
|
||||
async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None:
|
||||
server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.basic)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller"])
|
||||
async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None:
|
||||
server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.basic,
|
||||
authentication_token=value if source == "configured" else None)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value,default_slot", [
|
||||
(MCPAuth.api_key, "fixture-key", "X-API-Key"),
|
||||
(MCPAuth.bearer_token, "fixture-key", "Authorization"),
|
||||
(MCPAuth.basic, "user:pass", "Authorization"),
|
||||
(MCPAuth.token, "fixture-key", "Authorization"),
|
||||
(MCPAuth.authorization, "fixture-key", "Authorization"),
|
||||
])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller"])
|
||||
async def test_usable_credential_survives_an_empty_alternate_header(
|
||||
self, auth_type: MCPAuthType, value: str, default_slot: str, source: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="alternate", name="alternate", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom",
|
||||
authentication_token=value if source == "configured" else None,
|
||||
)
|
||||
empty_slot: Final = default_slot if source == "configured" else "X-Custom"
|
||||
selected_slot: Final = "X-Custom" if source == "configured" else default_slot
|
||||
client: Final = await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""},
|
||||
)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.headers[selected_slot]
|
||||
assert request.headers[empty_slot] == ""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="both-empty", name="both-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("custom_slot", [None, "X-Custom"])
|
||||
@pytest.mark.parametrize("source", ["caller", "forwarded"])
|
||||
async def test_api_key_preserves_explicit_authorization_credential(
|
||||
self, custom_slot: str | None, source: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot,
|
||||
)
|
||||
headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""}
|
||||
client: Final = await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header=headers if source == "caller" else None,
|
||||
extra_headers=headers if source == "forwarded" else None,
|
||||
)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.headers["Authorization"] == "Bearer caller-credential"
|
||||
assert request.headers["X-API-Key"] == ""
|
||||
assert custom_slot is None or custom_slot not in request.headers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", [
|
||||
"", " ", "Bearer", "Basic", "token", "ApiKey",
|
||||
"Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY",
|
||||
])
|
||||
async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["no-colon", "Basic bm8tY29sb24="])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller"])
|
||||
async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.basic,
|
||||
authentication_token=value if source == "configured" else None,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["user:pass", "user:", ":pass", ":"])
|
||||
async def test_basic_preserves_username_password_pairs(self, value: str) -> None:
|
||||
import base64
|
||||
|
||||
server: Final = MCPServer(
|
||||
server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value,
|
||||
)
|
||||
client: Final = await MCPServerManager()._create_mcp_client(server)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
scheme, encoded = request.headers["Authorization"].split(" ", 1)
|
||||
assert scheme == "Basic"
|
||||
assert base64.b64decode(encoded) == value.encode()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value", [
|
||||
(MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"),
|
||||
(MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"),
|
||||
])
|
||||
@pytest.mark.parametrize("source", ["configured", "caller"])
|
||||
async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix(
|
||||
self, auth_type: MCPAuthType, value: str, source: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type,
|
||||
authentication_token=value if source == "configured" else None,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value,expected", [
|
||||
(MCPAuth.bearer_token, "token", "Bearer token"),
|
||||
(MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"),
|
||||
(MCPAuth.token, "tokenish", "token tokenish"),
|
||||
])
|
||||
async def test_static_credentials_that_resemble_schemes_remain_usable(
|
||||
self, auth_type: MCPAuthType, value: str, expected: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="real-token", name="real-token", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=auth_type, authentication_token=value,
|
||||
)
|
||||
client: Final = await MCPServerManager()._create_mcp_client(server)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.headers["Authorization"] == expected
|
||||
|
|
|
|||
|
|
@ -10,9 +10,14 @@ This test suite ensures that:
|
|||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from respx import MockRouter
|
||||
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header,
|
||||
|
|
@ -35,6 +40,120 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
|
|||
GET_ASYNC_CLIENT_TARGET = "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.get_async_httpx_client"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value,accepted", [
|
||||
(MCPAuth.api_key, "Bearer Bearer", False), (MCPAuth.api_key, "ApiKey ApiKey", False),
|
||||
(MCPAuth.api_key, "token token", False), (MCPAuth.api_key, "bEaReR BEARER", False),
|
||||
(MCPAuth.api_key, "aPiKeY\tAPIKEY", False), (MCPAuth.api_key, "Bearer fixture-key", True),
|
||||
(MCPAuth.api_key, "ApiKey fixture-key", True), (MCPAuth.api_key, "token fixture-key", True),
|
||||
(MCPAuth.authorization, "Bearer", False), (MCPAuth.authorization, "basic", False),
|
||||
(MCPAuth.authorization, "token", False), (MCPAuth.authorization, "ApiKey", False),
|
||||
(MCPAuth.authorization, " bEaReR ", False), (MCPAuth.authorization, "\tTOKEN\t", False),
|
||||
(MCPAuth.authorization, "opaque-secret-value", True), (MCPAuth.authorization, "Bearer abc", True),
|
||||
(MCPAuth.authorization, "Custom abc", True),
|
||||
])
|
||||
async def test_authorization_validates_credentials_before_http(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, auth_type: MCPAuthType, value: str, accepted: bool,
|
||||
) -> None:
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
tool: Final = create_tool_function(
|
||||
"/echo", "get", {}, "https://upstream.example", auth_type=auth_type,
|
||||
)
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
caller_token: Final = _request_auth_header.set(value)
|
||||
try:
|
||||
if accepted:
|
||||
assert await tool() == "authenticated"
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["authorization"] == value
|
||||
else:
|
||||
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
|
||||
await tool()
|
||||
assert exc.value.status_code == 500
|
||||
assert destination.call_count == 0
|
||||
finally:
|
||||
_request_auth_header.reset(caller_token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("static,forwarded,caller,resolved,expected", [
|
||||
({"Authorization": "Bearer configured"}, {"authorization": "Bearer forwarded"}, None, None, "Bearer configured"),
|
||||
({"Authorization": "Bearer configured"}, None, "Bearer caller", None, "Bearer caller"),
|
||||
({"Authorization": "Bearer configured"}, None, "Bearer", None, None),
|
||||
({"Authorization": "Bearer configured"}, None, "Bearer caller", {"authorization": " "}, None),
|
||||
({"Authorization": "Bearer configured"}, None, "Bearer", {"authorization": "Bearer resolved"}, "Bearer resolved"),
|
||||
])
|
||||
async def test_static_auth_validates_headers_after_existing_precedence(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None,
|
||||
resolved: dict[str, str] | None, expected: str | None,
|
||||
) -> None:
|
||||
tool: Final = create_tool_function(
|
||||
"/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.bearer_token,
|
||||
)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
caller_token: Final = _request_auth_header.set(caller)
|
||||
extra_token: Final = _request_extra_headers.set(forwarded)
|
||||
resolved_token: Final = _request_resolved_auth_headers.set(resolved)
|
||||
try:
|
||||
if expected is None:
|
||||
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
|
||||
await tool()
|
||||
assert exc.value.status_code == 500
|
||||
assert destination.call_count == 0
|
||||
else:
|
||||
assert await tool() == "authenticated"
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["authorization"] == expected
|
||||
finally:
|
||||
_request_auth_header.reset(caller_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
_request_resolved_auth_headers.reset(resolved_token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("credential", ["custom-key", ""])
|
||||
async def test_static_auth_uses_configured_custom_header(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, credential: str,
|
||||
) -> None:
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
tool: Final = create_tool_function(
|
||||
"/echo", "get", {}, "https://upstream.example", headers={"x-custom": credential},
|
||||
auth_type=MCPAuth.api_key, upstream_token_header="X-Custom",
|
||||
)
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
if credential:
|
||||
assert await tool() == "authenticated"
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["x-custom"] == credential
|
||||
else:
|
||||
with pytest.raises(HTTPException, match="requires a usable upstream credential"):
|
||||
await tool()
|
||||
assert destination.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,resolved", [
|
||||
(MCPAuth.none, None),
|
||||
(MCPAuth.oauth2, {"Authorization": "Bearer user-oauth"}),
|
||||
])
|
||||
async def test_static_validation_preserves_no_auth_and_resolved_oauth(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
auth_type: MCPAuthType, resolved: dict[str, str] | None,
|
||||
) -> None:
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
tool: Final = create_tool_function("/echo", "get", {}, "https://upstream.example", auth_type=auth_type)
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="echo")
|
||||
token: Final = _request_resolved_auth_headers.set(resolved)
|
||||
try:
|
||||
assert await tool() == "echo"
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers.get("authorization") == (resolved or {}).get("Authorization")
|
||||
finally:
|
||||
_request_resolved_auth_headers.reset(token)
|
||||
|
||||
|
||||
def _create_mock_client(method: str, response_text: str, status_code: int = 200) -> AsyncMock:
|
||||
"""Utility to create a mocked async httpx client for the given method.
|
||||
|
||||
|
|
@ -1458,3 +1577,21 @@ class TestBoundedOpenAPISpecLoading:
|
|||
else:
|
||||
assert await load_openapi_spec_async("https://93.184.216.34/spec.json", max_bytes=100) == {"paths": {}}
|
||||
assert destination.call_count == 1
|
||||
|
||||
|
||||
def test_openapi_generator_import_does_not_require_mcp_sdk() -> None:
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
script = """
|
||||
import builtins
|
||||
original_import = builtins.__import__
|
||||
def without_mcp(name, *args, **kwargs):
|
||||
if name == 'mcp' or name.startswith('mcp.'):
|
||||
raise ModuleNotFoundError('MCP SDK unavailable')
|
||||
return original_import(name, *args, **kwargs)
|
||||
builtins.__import__ = without_mcp
|
||||
import litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator
|
||||
"""
|
||||
result = subprocess.run([sys.executable, "-c", script], capture_output=True, text=True)
|
||||
assert result.returncode == 0, result.stderr
|
||||
|
|
|
|||
|
|
@ -3,12 +3,14 @@ Test for anthropic_endpoints/endpoints.py, focusing on handling dictionary objec
|
|||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import unittest
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
|
||||
|
|
@ -285,6 +287,115 @@ class TestFailureHookRequestData:
|
|||
assert hook_request_data["litellm_logging_obj"] == "logging-obj-sentinel"
|
||||
|
||||
|
||||
class TestErrorLogCarriesCallId:
|
||||
"""LIT-7836: the /v1/messages and /v1/messages/count_tokens error lines must carry
|
||||
the request's litellm_call_id, rendered in the message and as a structured field."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def propagating_proxy_logger(self):
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
@staticmethod
|
||||
def _error_record(caplog: pytest.LogCaptureFixture) -> logging.LogRecord:
|
||||
return next(r for r in caplog.records if "Exception occured" in r.getMessage())
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_messages_failure_log_carries_call_id(self, caplog: pytest.LogCaptureFixture):
|
||||
import litellm.proxy.anthropic_endpoints.endpoints as ep
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
call_id = "messages-call-7836"
|
||||
|
||||
async def fake_process(self, **kwargs):
|
||||
self.data = {**self.data, "litellm_call_id": call_id}
|
||||
raise RuntimeError("provider timeout")
|
||||
|
||||
request = MagicMock()
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam
|
||||
patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), # test-quality-ok: the provider failure happens inside this call; the test targets the endpoint's except block
|
||||
patch.object(proxy_server, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global imported at call time; no injection seam
|
||||
caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"),
|
||||
):
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
response = await ep.anthropic_response(
|
||||
fastapi_response=MagicMock(),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
assert response.status_code == 500
|
||||
record = self._error_record(caplog)
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_messages_already_shaped_failure_answers_with_the_call_id(self):
|
||||
import litellm.proxy.anthropic_endpoints.endpoints as ep
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth
|
||||
|
||||
call_id = "messages-call-7836-shaped"
|
||||
|
||||
async def fake_process(self, **kwargs):
|
||||
self.data = {**self.data, "litellm_call_id": call_id}
|
||||
raise ProxyException(message="budget exceeded", type=ProxyErrorTypes.budget_exceeded, param="key", code=402)
|
||||
|
||||
request = MagicMock()
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam
|
||||
patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), # test-quality-ok: the proxy shaped failure happens inside this call; the test targets the endpoint's except block
|
||||
patch.object(proxy_server, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global imported at call time; no injection seam
|
||||
):
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
response = await ep.anthropic_response(
|
||||
fastapi_response=MagicMock(),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
assert response.status_code == 402
|
||||
assert response.headers["x-litellm-call-id"] == call_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_failure_log_carries_callers_call_id(self, caplog: pytest.LogCaptureFixture):
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm.proxy.anthropic_endpoints.endpoints as ep
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
call_id = "count-tokens-call-7836"
|
||||
request = MagicMock()
|
||||
request.headers = {"x-litellm-call-id": call_id}
|
||||
|
||||
with (
|
||||
patch.object( # test-quality-ok: endpoint reads the body via a module function; no injection seam
|
||||
ep,
|
||||
"_read_request_body",
|
||||
new=AsyncMock(return_value={"model": "claude-sonnet", "messages": [{"role": "user", "content": "hi"}]}),
|
||||
),
|
||||
patch.object(proxy_server, "token_counter", new=AsyncMock(side_effect=RuntimeError("tokenizer down"))), # test-quality-ok: module global imported at call time; the test targets the endpoint's except block
|
||||
caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"),
|
||||
pytest.raises(HTTPException) as raised,
|
||||
):
|
||||
await ep.count_tokens(request=request, user_api_key_dict=UserAPIKeyAuth())
|
||||
|
||||
assert raised.value.status_code == 500
|
||||
record = self._error_record(caplog)
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
|
||||
class TestEventLoggingBatchEndpoint:
|
||||
"""Test the stubbed event logging batch endpoint"""
|
||||
|
||||
|
|
|
|||
|
|
@ -2892,45 +2892,49 @@ def test_team_update_gate_allows_org_admin_with_resolved_org():
|
|||
)
|
||||
|
||||
|
||||
def test_team_update_gate_rejects_without_org_context():
|
||||
"""Without organization_id (i.e. resolution found no org, or a non-org-admin),
|
||||
the gate still rejects /team/update — the fix adds no blanket allow. Guards
|
||||
against re-widening the route (e.g. dropping it into self_managed_routes)."""
|
||||
def test_team_update_gate_admits_internal_user_without_org_context(): # test-quality-ok: the gate's only success signal is not raising; the handler's team-admin 403s are pinned in test_team_endpoints
|
||||
"""/team/update is self-managed (LIT-5722): the coarse gate admits any authenticated
|
||||
caller and update_team resolves proxy, org or team admin itself, then filters team admins
|
||||
through the team_admin_editable_team_fields setting. Before that the gate 401'd every
|
||||
team admin, which left the handler's team-admin branch unreachable."""
|
||||
user_obj = LiteLLM_UserTable(
|
||||
user_id="team-admin-user",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER.value,
|
||||
organization_memberships=None,
|
||||
)
|
||||
valid_token = UserAPIKeyAuth(user_id="team-admin-user", user_role=LitellmUserRoles.INTERNAL_USER.value)
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = "POST"
|
||||
request.query_params = {}
|
||||
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER.value,
|
||||
route="/team/update",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={"team_id": "team-1", "max_budget": 42},
|
||||
)
|
||||
|
||||
|
||||
def test_team_update_gate_defers_cross_org_admin_to_the_handler(): # test-quality-ok: the gate's only success signal is not raising; the handler's 403 it defers to is pinned in test_team_endpoints
|
||||
"""An org admin of a DIFFERENT org clears the coarse gate like any internal user;
|
||||
update_team's _resolve_team_access finds no role on the team and 403s (pinned in
|
||||
test_team_endpoints), so there is still no cross-org escalation."""
|
||||
user_obj = _make_org_admin_user("org-1")
|
||||
valid_token = UserAPIKeyAuth(user_id="org-admin-user", user_role=LitellmUserRoles.INTERNAL_USER.value)
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = "POST"
|
||||
request.query_params = {}
|
||||
|
||||
with pytest.raises(Exception, match="Only proxy admin can be used to generate"):
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER.value,
|
||||
route="/team/update",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={"team_id": "team-1", "max_budget": 42},
|
||||
)
|
||||
|
||||
|
||||
def test_team_update_gate_rejects_cross_org_admin_with_resolved_org():
|
||||
"""Even after the target team's org is resolved, an org admin of a DIFFERENT
|
||||
org is rejected at the gate (no cross-org escalation)."""
|
||||
user_obj = _make_org_admin_user("org-1")
|
||||
valid_token = UserAPIKeyAuth(user_id="org-admin-user", user_role=LitellmUserRoles.INTERNAL_USER.value)
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = "POST"
|
||||
request.query_params = {}
|
||||
|
||||
with pytest.raises(Exception, match="Only proxy admin can be used to generate"):
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER.value,
|
||||
route="/team/update",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={"team_id": "team-1", "organization_id": "org-2"},
|
||||
)
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER.value,
|
||||
route="/team/update",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={"team_id": "team-1", "organization_id": "org-2"},
|
||||
)
|
||||
|
||||
|
||||
# ── PATCH /team/{team_id}: same org-context + role reach as POST /team/update ──
|
||||
|
|
@ -2993,23 +2997,6 @@ async def test_add_team_org_context_noop_for_static_team_route():
|
|||
assert out == body
|
||||
|
||||
|
||||
def test_patch_team_route_has_same_reach_as_team_update():
|
||||
"""/team/{team_id} is reachable by org admins (in org_admin_allowed_routes) but
|
||||
NOT by regular internal users or the role-agnostic self_managed_routes — the
|
||||
latter would open /team/new (the collision footgun) to any authenticated user."""
|
||||
from litellm.proxy._types import LiteLLMRoutes
|
||||
|
||||
assert RouteChecks.check_route_access(
|
||||
route="/team/abc-123", allowed_routes=LiteLLMRoutes.org_admin_allowed_routes.value
|
||||
)
|
||||
assert not RouteChecks.check_route_access(
|
||||
route="/team/abc-123", allowed_routes=LiteLLMRoutes.internal_user_routes.value
|
||||
)
|
||||
assert not RouteChecks.check_route_access(
|
||||
route="/team/abc-123", allowed_routes=LiteLLMRoutes.self_managed_routes.value
|
||||
)
|
||||
|
||||
|
||||
def _patch_team_request() -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = "PATCH"
|
||||
|
|
@ -3897,7 +3884,6 @@ def test_team_disable_logging_stays_proxy_admin_only():
|
|||
"route",
|
||||
[
|
||||
"/team/06bda574-5ca9-43d3-beb8-3b23c2f17112",
|
||||
"/team/update",
|
||||
"/team/06bda574-5ca9-43d3-beb8-3b23c2f17112/model/add",
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ cannot drift without a test failure.
|
|||
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
from contextlib import ExitStack
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
|
@ -1088,6 +1089,28 @@ async def test_create__exception_calls_failure_hook(harness, openai_env_creds):
|
|||
assert harness.logging.post_call_failure_hook.call_args.kwargs["original_exception"].args[0] == "provider boom"
|
||||
|
||||
|
||||
async def test_create__exception_carries_the_litellm_call_id(harness, openai_env_creds, caplog):
|
||||
call_id = "lit7836-batch-call-id"
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": "file-plain",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
"litellm_call_id": call_id,
|
||||
},
|
||||
)
|
||||
harness.litellm_acreate.side_effect = ValueError("provider boom")
|
||||
|
||||
with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"), pytest.raises(ProxyException) as raised:
|
||||
await call_create(harness)
|
||||
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
record = next(r for r in caplog.records if "Exception occured" in r.getMessage())
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# #
|
||||
# GET /v1/batches/{batch_id} - retrieve_batch routing-contract tests #
|
||||
|
|
@ -1953,6 +1976,24 @@ async def test_list__exception_calls_failure_hook(list_harness):
|
|||
assert list_harness.logging.post_call_failure_hook.call_args.kwargs["original_exception"].args[0] == "provider boom"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list__failure_hook_and_response_share_the_request_litellm_call_id(list_harness):
|
||||
call_id = "lit7836-list-batches-call-id"
|
||||
list_harness.pre_call.side_effect = lambda **kw: (
|
||||
{**list_harness.body["body"], "litellm_call_id": call_id},
|
||||
MagicMock(),
|
||||
)
|
||||
list_harness.litellm_alist.side_effect = ValueError("provider boom")
|
||||
|
||||
with pytest.raises(ProxyException) as raised:
|
||||
await call_list(list_harness, after="batch-0", limit=5)
|
||||
|
||||
failure_request_data = list_harness.logging.post_call_failure_hook.call_args.kwargs["request_data"]
|
||||
assert failure_request_data["litellm_call_id"] == call_id
|
||||
assert (failure_request_data["after"], failure_request_data["limit"]) == ("batch-0", 5)
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# #
|
||||
# POST /v1/batches/{batch_id}/cancel - cancel_batch routing-contract tests #
|
||||
|
|
|
|||
|
|
@ -6,8 +6,10 @@ from fastapi import HTTPException
|
|||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
litellm_call_id_headers,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
with_litellm_call_id,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -158,3 +160,32 @@ def test_a_stringified_none_type_or_param_is_treated_as_absent():
|
|||
assert carried.type == "None"
|
||||
assert openai_error_type(carried, 400) == "invalid_request_error"
|
||||
assert openai_error_param(carried) is None
|
||||
|
||||
|
||||
def test_a_failed_request_answers_with_the_call_id_it_was_logged_under():
|
||||
assert litellm_call_id_headers("call-7836") == {"x-litellm-call-id": "call-7836"}
|
||||
assert litellm_call_id_headers(None) is None
|
||||
|
||||
|
||||
def test_an_already_shaped_proxy_error_answers_with_the_call_id_it_was_logged_under():
|
||||
raised_without_id = ProxyException(message="budget exceeded", type="budget_exceeded", param="key", code=402)
|
||||
|
||||
carried = with_litellm_call_id(raised_without_id, "call-7836")
|
||||
|
||||
assert carried is raised_without_id
|
||||
assert carried.headers == {"x-litellm-call-id": "call-7836"}
|
||||
assert (carried.message, carried.type, carried.param, carried.code) == (
|
||||
"budget exceeded",
|
||||
"budget_exceeded",
|
||||
"key",
|
||||
"402",
|
||||
)
|
||||
|
||||
|
||||
def test_a_proxy_error_keeps_the_call_id_it_was_raised_with():
|
||||
raised_with_id = ProxyException(
|
||||
message="nope", type="None", param=None, code=400, headers={"x-litellm-call-id": "first"}
|
||||
)
|
||||
|
||||
assert with_litellm_call_id(raised_with_id, "second").headers == {"x-litellm-call-id": "first"}
|
||||
assert with_litellm_call_id(ProxyException(message="nope", type="None", param=None, code=400), None).headers == {}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import logging
|
||||
from collections.abc import Iterator, Mapping
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Dict
|
||||
|
||||
|
|
@ -10,6 +12,7 @@ from fastapi.testclient import TestClient
|
|||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.image_endpoints import endpoints
|
||||
|
|
@ -211,3 +214,117 @@ async def test_a_model_the_router_cannot_serve_answers_an_openai_typed_error(mon
|
|||
await endpoints.image_generation(request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth())
|
||||
|
||||
assert (raised.value.type, raised.value.param, raised.value.code) == ("invalid_request_error", None, "404")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def propagating_proxy_logger() -> Iterator[None]:
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_log_carries_the_callers_litellm_call_id(
|
||||
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, propagating_proxy_logger: None
|
||||
) -> None:
|
||||
"""LIT-7836: the /v1/images/generations error line must carry the litellm_call_id
|
||||
the client sent, both rendered in the message and as a structured record field."""
|
||||
call_id = "images-call-7836"
|
||||
|
||||
async def fake_add_litellm_data_to_request(**kwargs: object) -> object:
|
||||
return kwargs["data"]
|
||||
|
||||
async def fake_pre_call_hook(*, user_api_key_dict: UserAPIKeyAuth, data: dict[str, object], call_type: str) -> dict[str, object]:
|
||||
return data
|
||||
|
||||
async def fake_post_call_failure_hook(**_: object) -> None:
|
||||
return None
|
||||
|
||||
async def failing_route_request(**_: object) -> None:
|
||||
raise HTTPException(status_code=401, detail={"error": "invalid api key"})
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.add_litellm_data_to_request", fake_add_litellm_data_to_request)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {})
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj",
|
||||
SimpleNamespace(pre_call_hook=fake_pre_call_hook, post_call_failure_hook=fake_post_call_failure_hook),
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.version", "test-version")
|
||||
monkeypatch.setattr("litellm.proxy.image_endpoints.endpoints.route_request", failing_route_request)
|
||||
|
||||
body = orjson.dumps({"model": "dall-e-3", "prompt": "a lighthouse at dusk"})
|
||||
|
||||
async def receive() -> dict[str, object]:
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
request = Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/images/generations",
|
||||
"headers": [(b"x-litellm-call-id", call_id.encode())],
|
||||
},
|
||||
receive,
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"), pytest.raises(ProxyException) as raised:
|
||||
await endpoints.image_generation(request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth())
|
||||
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
record = next(r for r in caplog.records if "Exception occured" in r.getMessage())
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_before_the_provider_call_bills_the_callers_litellm_call_id(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""LIT-7836: when the request is rejected while it is still being prepared, the
|
||||
failure hook must see the same litellm_call_id the response header answers with,
|
||||
otherwise the spend row is stored under a freshly minted id nobody can look up."""
|
||||
call_id = "images-early-7836"
|
||||
hook_request_data: list[Mapping[str, object]] = []
|
||||
|
||||
async def rejecting_add_litellm_data_to_request(**_: object) -> object:
|
||||
raise HTTPException(status_code=400, detail={"error": "tag not allowed"})
|
||||
|
||||
async def fake_post_call_failure_hook(*, request_data: Mapping[str, object], **_: object) -> None:
|
||||
hook_request_data.append(request_data)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.add_litellm_data_to_request", rejecting_add_litellm_data_to_request)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {})
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj",
|
||||
SimpleNamespace(post_call_failure_hook=fake_post_call_failure_hook),
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.version", "test-version")
|
||||
|
||||
body = orjson.dumps({"model": "dall-e-3", "prompt": "a lighthouse at dusk", "litellm_call_id": "from-the-body"})
|
||||
|
||||
async def receive() -> dict[str, object]:
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
request = Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/images/generations",
|
||||
"headers": [(b"x-litellm-call-id", call_id.encode())],
|
||||
},
|
||||
receive,
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as raised:
|
||||
await endpoints.image_generation(request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth())
|
||||
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
assert [data["litellm_call_id"] for data in hook_request_data] == [call_id]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,138 @@
|
|||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import LiteLLM_ModelTable, LiteLLM_TeamTable, UpdateTeamRequest
|
||||
from litellm.proxy.management_endpoints.team_admin_field_permissions import (
|
||||
TeamAdminEditAllowed,
|
||||
TeamAdminEditingDisabled,
|
||||
TeamAdminFieldNotPermitted,
|
||||
changed_team_fields,
|
||||
resolve_team_admin_editable_fields,
|
||||
team_admin_edit_verdict,
|
||||
team_admin_request_or_raise,
|
||||
)
|
||||
|
||||
_SUPPORTED = frozenset({"tpm_limit", "rpm_limit", "team_alias"})
|
||||
|
||||
|
||||
def _team(**overrides):
|
||||
return LiteLLM_TeamTable(team_id="team-1", **overrides)
|
||||
|
||||
|
||||
class TestResolveTeamAdminEditableFields:
|
||||
def test_missing_setting_means_nothing_editable(self):
|
||||
assert resolve_team_admin_editable_fields({}, _SUPPORTED) == frozenset()
|
||||
|
||||
def test_keeps_only_supported_names(self):
|
||||
configured = {"team_admin_editable_team_fields": ["tpm_limit", "blocked", "organization_id"]}
|
||||
assert resolve_team_admin_editable_fields(configured, _SUPPORTED) == frozenset({"tpm_limit"})
|
||||
|
||||
@pytest.mark.parametrize("raw", ["tpm_limit", 7, {"tpm_limit": True}, [1, 2]])
|
||||
def test_malformed_setting_fails_closed(self, raw):
|
||||
assert resolve_team_admin_editable_fields({"team_admin_editable_team_fields": raw}, _SUPPORTED) == frozenset()
|
||||
|
||||
|
||||
class TestChangedTeamFields:
|
||||
def test_team_id_alone_changes_nothing(self):
|
||||
assert changed_team_fields(UpdateTeamRequest(team_id="team-1"), _team()) == frozenset()
|
||||
|
||||
def test_column_echoing_stored_value_is_not_a_change(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", tpm_limit=5, team_alias="alpha", max_budget=None)
|
||||
assert changed_team_fields(data, _team(tpm_limit=5, team_alias="alpha")) == frozenset()
|
||||
|
||||
def test_column_with_different_value_is_a_change(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", tpm_limit=6, team_alias="alpha")
|
||||
assert changed_team_fields(data, _team(tpm_limit=5, team_alias="alpha")) == frozenset({"tpm_limit"})
|
||||
|
||||
def test_explicit_null_clearing_a_stored_column_is_a_change(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", max_budget=None)
|
||||
assert changed_team_fields(data, _team(max_budget=30.0)) == frozenset({"max_budget"})
|
||||
|
||||
def test_folded_field_sent_top_level_is_named_not_metadata(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", guardrails=["b"])
|
||||
assert changed_team_fields(data, _team(metadata={"guardrails": ["a"]})) == frozenset({"guardrails"})
|
||||
|
||||
def test_folded_field_sent_inside_metadata_is_named_not_metadata(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", metadata={"guardrails": ["b"]})
|
||||
assert changed_team_fields(data, _team(metadata={"guardrails": ["a"]})) == frozenset({"guardrails"})
|
||||
|
||||
def test_custom_metadata_key_change_is_attributed_to_metadata(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", metadata={"guardrails": ["a"], "cost_center": "b"})
|
||||
existing = _team(metadata={"guardrails": ["a"], "cost_center": "a"})
|
||||
assert changed_team_fields(data, existing) == frozenset({"metadata"})
|
||||
|
||||
def test_metadata_echo_with_top_level_override_only_names_the_override(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", guardrails=["b"], metadata={"guardrails": ["a"], "cost_center": "a"})
|
||||
existing = _team(metadata={"guardrails": ["a"], "cost_center": "a"})
|
||||
assert changed_team_fields(data, existing) == frozenset({"guardrails"})
|
||||
|
||||
def test_dropping_a_stored_key_from_submitted_metadata_is_a_change(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", metadata={"cost_center": "a"})
|
||||
existing = _team(metadata={"cost_center": "a", "tags": ["x"], "logging": [{"callback": "langfuse"}]})
|
||||
assert changed_team_fields(data, existing) == frozenset({"tags", "logging"})
|
||||
|
||||
def test_server_managed_metadata_key_is_ignored(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", metadata={"cost_center": "a"})
|
||||
existing = _team(metadata={"cost_center": "a", "team_member_budget_id": "budget-1"})
|
||||
assert changed_team_fields(data, existing) == frozenset()
|
||||
|
||||
def test_model_aliases_compare_against_the_model_table(self):
|
||||
table = LiteLLM_ModelTable(model_aliases='{"fast": "gpt-4o-mini"}', created_by="a", updated_by="a")
|
||||
same = UpdateTeamRequest(team_id="team-1", model_aliases={"fast": "gpt-4o-mini"})
|
||||
different = UpdateTeamRequest(team_id="team-1", model_aliases={"fast": "gpt-4o"})
|
||||
assert changed_team_fields(same, _team(litellm_model_table=table)) == frozenset()
|
||||
assert changed_team_fields(different, _team(litellm_model_table=table)) == frozenset({"model_aliases"})
|
||||
|
||||
def test_empty_model_aliases_against_no_model_table_is_not_a_change(self):
|
||||
assert changed_team_fields(UpdateTeamRequest(team_id="team-1", model_aliases={}), _team()) == frozenset()
|
||||
|
||||
def test_field_without_a_stored_counterpart_counts_as_changed_when_sent(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", team_member_budget=10.0)
|
||||
assert changed_team_fields(data, _team()) == frozenset({"team_member_budget"})
|
||||
|
||||
|
||||
class TestTeamAdminEditVerdict:
|
||||
def test_no_permitted_fields_disables_editing_even_for_a_no_op(self):
|
||||
verdict = team_admin_edit_verdict(UpdateTeamRequest(team_id="team-1"), _team(), frozenset())
|
||||
assert verdict == TeamAdminEditingDisabled()
|
||||
|
||||
def test_allowed_request_keeps_only_the_changed_fields(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", tpm_limit=6, team_alias="alpha", budget_duration="30d")
|
||||
existing = _team(team_alias="alpha", budget_duration="30d")
|
||||
verdict = team_admin_edit_verdict(data, existing, frozenset({"tpm_limit"}))
|
||||
assert isinstance(verdict, TeamAdminEditAllowed)
|
||||
assert verdict.request.model_dump(exclude_unset=True) == {"team_id": "team-1", "tpm_limit": 6}
|
||||
|
||||
def test_permitted_field_changed_inside_metadata_keeps_the_metadata(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", metadata={"guardrails": ["b"]}, team_alias="alpha")
|
||||
existing = _team(team_alias="alpha", metadata={"guardrails": ["a"]})
|
||||
verdict = team_admin_edit_verdict(data, existing, frozenset({"guardrails"}))
|
||||
assert isinstance(verdict, TeamAdminEditAllowed)
|
||||
assert verdict.request.model_dump(exclude_unset=True) == {
|
||||
"team_id": "team-1",
|
||||
"metadata": {"guardrails": ["b"]},
|
||||
}
|
||||
|
||||
def test_first_blocked_field_in_sorted_order_is_reported(self):
|
||||
data = UpdateTeamRequest(team_id="team-1", tpm_limit=6, rpm_limit=6, blocked=True)
|
||||
verdict = team_admin_edit_verdict(data, _team(), frozenset({"tpm_limit"}))
|
||||
assert verdict == TeamAdminFieldNotPermitted(field="blocked")
|
||||
|
||||
|
||||
class TestTeamAdminRequestOrRaise:
|
||||
def test_allowed_hands_back_its_request(self):
|
||||
request = UpdateTeamRequest(team_id="team-1", tpm_limit=6)
|
||||
assert team_admin_request_or_raise(TeamAdminEditAllowed(request=request)) is request
|
||||
|
||||
def test_disabled_is_a_403_pointing_at_the_proxy_admin(self):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
team_admin_request_or_raise(TeamAdminEditingDisabled())
|
||||
assert exc.value.status_code == 403
|
||||
assert "cannot edit team settings" in exc.value.detail
|
||||
assert "Settings > UI > Team admin editable fields" in exc.value.detail
|
||||
|
||||
def test_field_not_permitted_is_a_403_naming_the_field(self):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
team_admin_request_or_raise(TeamAdminFieldNotPermitted(field="blocked"))
|
||||
assert exc.value.status_code == 403
|
||||
assert "'blocked'" in exc.value.detail
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
import asyncio
|
||||
import json
|
||||
from contextlib import asynccontextmanager
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Final, Optional, cast
|
||||
|
|
@ -76,6 +76,31 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
client = TestClient(app)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _team_admin_may_edit(*fields: str):
|
||||
"""Let team admins change ``fields`` on /team/update for the duration of the block.
|
||||
|
||||
The registry only lists the fields shipped so far (LIT-5722 adds them one PR at a time), so tests that
|
||||
exercise the gates layered underneath the allow-list widen it here instead of asserting the early 403."""
|
||||
with (
|
||||
patch( # test-quality-ok: the registry is a module constant update_team reads directly; no seam to inject
|
||||
"litellm.proxy.management_endpoints.team_endpoints.SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS",
|
||||
frozenset(fields),
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.general_settings", {"team_admin_editable_team_fields": list(fields)}), # test-quality-ok: update_team reads general_settings as a proxy_server module global
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
def _not_org_admin():
|
||||
"""update_team asks whether the caller administers the team's org before it settles for team admin;
|
||||
a MagicMock prisma cannot answer that lookup, so pin it to False."""
|
||||
return patch( # test-quality-ok: the org-admin lookup needs a real prisma client this file's MagicMock cannot provide
|
||||
"litellm.proxy.management_endpoints.team_endpoints._is_user_org_admin_for_team",
|
||||
AsyncMock(return_value=False),
|
||||
)
|
||||
|
||||
|
||||
def _wire_team_create_tx(prisma_client):
|
||||
"""`/team/new` inserts the team and mirrors it onto the access groups in one transaction,
|
||||
so a mocked client has to hand its team table back out of `db.tx()`.
|
||||
|
|
@ -6393,6 +6418,7 @@ async def test_update_team_standalone_budget_raise_blocked_for_team_admin():
|
|||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("max_budget"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -6549,6 +6575,7 @@ async def test_update_team_standalone_budget_removal_blocked_for_team_admin():
|
|||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("max_budget"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -6618,6 +6645,7 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed(
|
|||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("max_budget"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -6712,6 +6740,7 @@ async def test_update_team_standalone_unchanged_budget_allowed(
|
|||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("max_budget", "tpm_limit"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -6810,6 +6839,7 @@ async def test_update_team_standalone_lower_budget_allowed(
|
|||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("max_budget"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -6912,6 +6942,8 @@ async def test_update_team_org_scoped_budget_exceeds_org_limit():
|
|||
mock_org.litellm_budget_table = mock_budget_table
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("max_budget"),
|
||||
_not_org_admin(),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -6992,6 +7024,7 @@ async def test_update_team_standalone_models_not_gated_by_user_limit(
|
|||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("models"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -7091,6 +7124,8 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit(
|
|||
mock_org.litellm_budget_table = mock_budget_table
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("max_budget"),
|
||||
_not_org_admin(),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -7202,6 +7237,8 @@ async def test_update_team_org_scoped_models_bypasses_user_limit(
|
|||
mock_org.litellm_budget_table = None
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("models"),
|
||||
_not_org_admin(),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -7304,6 +7341,8 @@ async def test_update_team_org_scoped_models_not_in_org_models():
|
|||
mock_org.litellm_budget_table = None
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("models"),
|
||||
_not_org_admin(),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -7393,6 +7432,8 @@ async def test_update_team_org_scoped_models_with_all_proxy_models(
|
|||
mock_org.litellm_budget_table = None
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("models"),
|
||||
_not_org_admin(),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -7502,6 +7543,7 @@ async def test_update_team_tpm_limit_not_gated_by_user_limit(
|
|||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("tpm_limit"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -7584,6 +7626,7 @@ async def test_update_team_rpm_limit_not_gated_by_user_limit(
|
|||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("rpm_limit"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -7981,6 +8024,8 @@ async def test_update_team_org_scoped_tpm_exceeds_org_limit():
|
|||
mock_org.litellm_budget_table = mock_budget_table
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("tpm_limit"),
|
||||
_not_org_admin(),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -8067,6 +8112,8 @@ async def test_update_team_org_scoped_rpm_exceeds_org_limit():
|
|||
mock_org.litellm_budget_table = mock_budget_table
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("rpm_limit"),
|
||||
_not_org_admin(),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -8158,6 +8205,8 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit(
|
|||
mock_org.litellm_budget_table = mock_budget_table
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("tpm_limit", "rpm_limit"),
|
||||
_not_org_admin(),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -8286,6 +8335,7 @@ async def test_update_team_guardrails_with_org_id(
|
|||
}
|
||||
|
||||
with (
|
||||
_team_admin_may_edit("guardrails", "organization_id"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
|
|
@ -11177,8 +11227,8 @@ async def test_update_team_blocks_non_admin_passthrough_routes(mock_db_client):
|
|||
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._verify_team_access",
|
||||
AsyncMock(return_value=None),
|
||||
"litellm.proxy.management_endpoints.team_endpoints._resolve_team_access",
|
||||
AsyncMock(return_value="org_admin"),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await update_team(
|
||||
|
|
@ -13246,6 +13296,7 @@ async def test_update_team_output_token_estimate_lowered_rejected_for_team_admin
|
|||
|
||||
with contextlib.ExitStack() as stack:
|
||||
_wire_update_team(stack, {_TEAM_ESTIMATE: 4000})
|
||||
stack.enter_context(_team_admin_may_edit("default_estimated_output_tokens"))
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", default_estimated_output_tokens=1),
|
||||
|
|
@ -13277,6 +13328,7 @@ async def test_update_team_output_token_estimate_unchanged_allows_team_admin_edi
|
|||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {_TEAM_ESTIMATE: 4000})
|
||||
stack.enter_context(_team_admin_may_edit("team_alias"))
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(
|
||||
team_id="test_team_id",
|
||||
|
|
@ -13336,6 +13388,7 @@ async def test_update_team_batch_enqueued_token_limit_raised_rejected_for_team_a
|
|||
|
||||
with contextlib.ExitStack() as stack:
|
||||
_wire_update_team(stack, {_TEAM_BATCH_LIMIT: 100000})
|
||||
stack.enter_context(_team_admin_may_edit("metadata"))
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", metadata={_TEAM_BATCH_LIMIT: 10**12}),
|
||||
|
|
@ -14894,3 +14947,369 @@ async def test_update_team_model_max_budget_raise_blocked_for_team_admin():
|
|||
assert exc.value.code == "403"
|
||||
assert "proxy admin" in str(exc.value.message).lower()
|
||||
mock_prisma.db.litellm_teamtable.update.assert_not_awaited()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LIT-5722: team admins reach update_team through self_managed_routes and are
|
||||
# filtered by the team_admin_editable_team_fields setting.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_TEAM_ADMIN_CALLER = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-team-admin", user_id="team-admin"
|
||||
)
|
||||
_PROXY_ADMIN_CALLER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin")
|
||||
|
||||
|
||||
def _update_request_stub():
|
||||
from unittest.mock import Mock
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
return Mock(spec=Request)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_team_admin_is_refused_before_any_write_when_no_fields_are_enabled():
|
||||
import contextlib
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
stack.enter_context(_team_admin_may_edit())
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", team_alias="renamed"),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert str(exc.value.code) == "403"
|
||||
assert "cannot edit team settings" in str(exc.value.message)
|
||||
assert not prisma.db.litellm_teamtable.update.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_configured_but_unsupported_field_does_not_open_editing():
|
||||
"""Only fields in SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS count, whatever general_settings says."""
|
||||
import contextlib
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
stack.enter_context(
|
||||
patch("litellm.proxy.proxy_server.general_settings", {"team_admin_editable_team_fields": ["team_alias"]}) # test-quality-ok: update_team reads general_settings as a proxy_server module global
|
||||
)
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", team_alias="renamed"),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert str(exc.value.code) == "403"
|
||||
assert "cannot edit team settings" in str(exc.value.message)
|
||||
assert not prisma.db.litellm_teamtable.update.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_team_admin_changing_an_unpermitted_field_is_refused_by_name():
|
||||
import contextlib
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
stack.enter_context(_team_admin_may_edit("team_alias"))
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", team_alias="renamed", tpm_limit=10),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert str(exc.value.code) == "403"
|
||||
assert "'tpm_limit'" in str(exc.value.message)
|
||||
assert not prisma.db.litellm_teamtable.update.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_team_admin_echoing_unpermitted_fields_unchanged_is_allowed(
|
||||
disable_audit_logging_for_mocked_team,
|
||||
):
|
||||
"""The dashboard resends the whole form, so only a value that differs from what is stored counts."""
|
||||
import contextlib
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
stack.enter_context(_team_admin_may_edit("team_alias"))
|
||||
result = await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", team_alias="renamed", tpm_limit=None, models=[]),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert result["data"].team_id == "test_team_id"
|
||||
assert prisma.db.litellm_teamtable.update.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_team_admin_changes_tpm_limit_once_a_proxy_admin_enables_it(
|
||||
disable_audit_logging_for_mocked_team,
|
||||
):
|
||||
"""tpm_limit is the first field a proxy admin can open to team admins; every other field stays admin-only."""
|
||||
import contextlib
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
stack.enter_context(
|
||||
patch("litellm.proxy.proxy_server.general_settings", {"team_admin_editable_team_fields": ["tpm_limit"]}) # test-quality-ok: update_team reads general_settings as a proxy_server module global
|
||||
)
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", tpm_limit=5000),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
with pytest.raises(ProxyException) as refused:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", tpm_limit=6000, rpm_limit=10),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert prisma.db.litellm_teamtable.update.await_count == 1
|
||||
assert prisma.db.litellm_teamtable.update.call_args.kwargs["data"]["tpm_limit"] == 5000
|
||||
assert str(refused.value.code) == "403"
|
||||
assert "'rpm_limit'" in str(refused.value.message)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_team_admin_resending_budget_settings_does_not_push_back_budget_resets(
|
||||
disable_audit_logging_for_mocked_team,
|
||||
):
|
||||
"""A resent budget_duration or budget_limits would otherwise recompute the reset timestamps from now."""
|
||||
import contextlib
|
||||
|
||||
stored_windows = [{"budget_duration": "7d", "max_budget": 5.0, "reset_at": "2026-09-20T00:00:00Z"}]
|
||||
budgeted_team = MagicMock()
|
||||
budgeted_team.metadata = {}
|
||||
budgeted_team.model_dump.return_value = {
|
||||
"team_id": "test_team_id",
|
||||
"team_alias": "test_team",
|
||||
"metadata": {},
|
||||
"budget_duration": "30d",
|
||||
"budget_limits": stored_windows,
|
||||
"members_with_roles": [{"user_id": "team-admin", "role": "admin"}],
|
||||
}
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=budgeted_team)
|
||||
stack.enter_context(_team_admin_may_edit("tpm_limit"))
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(
|
||||
team_id="test_team_id", tpm_limit=5000, budget_duration="30d", budget_limits=stored_windows
|
||||
),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
written = prisma.db.litellm_teamtable.update.call_args.kwargs["data"]
|
||||
assert written["tpm_limit"] == 5000
|
||||
assert not {"budget_duration", "budget_reset_at", "budget_limits"} & written.keys()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_holds_a_team_admin_to_the_org_tpm_limit(disable_audit_logging_for_mocked_team):
|
||||
"""The org ceiling lives on the org's budget row, so /team/update must load it to enforce the cap."""
|
||||
import contextlib
|
||||
|
||||
capped_org = LiteLLM_OrganizationTable(
|
||||
organization_id="capped-org",
|
||||
budget_id="capped-budget",
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(tpm_limit=10000),
|
||||
)
|
||||
|
||||
async def org_lookup(**kwargs):
|
||||
return capped_org if kwargs.get("include_budget_table") else capped_org.model_copy(
|
||||
update={"litellm_budget_table": None}
|
||||
)
|
||||
|
||||
org_team = MagicMock()
|
||||
org_team.metadata = {}
|
||||
org_team.organization_id = "capped-org"
|
||||
org_team.model_dump.return_value = {
|
||||
"team_id": "test_team_id",
|
||||
"team_alias": "test_team",
|
||||
"organization_id": "capped-org",
|
||||
"metadata": {},
|
||||
"members_with_roles": [{"user_id": "team-admin", "role": "admin"}],
|
||||
}
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=org_team)
|
||||
stack.enter_context(_team_admin_may_edit("tpm_limit"))
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: the org-admin lookup needs a real prisma client this file's MagicMock cannot provide
|
||||
"litellm.proxy.management_endpoints.team_endpoints._is_user_org_admin_for_team",
|
||||
AsyncMock(return_value=False),
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: update_team reads orgs through this module-level import; no seam to inject
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_org_object",
|
||||
AsyncMock(side_effect=org_lookup),
|
||||
)
|
||||
)
|
||||
with pytest.raises(ProxyException) as over_cap:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", tpm_limit=20000),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", tpm_limit=8000),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert str(over_cap.value.code) == "400"
|
||||
assert "exceeds organization's tpm_limit (10000)" in str(over_cap.value.message)
|
||||
assert prisma.db.litellm_teamtable.update.await_count == 1
|
||||
assert prisma.db.litellm_teamtable.update.call_args.kwargs["data"]["tpm_limit"] == 8000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_org_admin_is_not_filtered_by_the_team_admin_field_list(
|
||||
disable_audit_logging_for_mocked_team,
|
||||
):
|
||||
"""A caller who is both org admin and roster admin keeps unrestricted edits."""
|
||||
import contextlib
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
stack.enter_context(_team_admin_may_edit())
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: the org-admin lookup needs a real prisma client this file's MagicMock cannot provide
|
||||
"litellm.proxy.management_endpoints.team_endpoints._is_user_org_admin_for_team",
|
||||
AsyncMock(return_value=True),
|
||||
)
|
||||
)
|
||||
result = await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", team_alias="renamed"),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert result["data"].team_id == "test_team_id"
|
||||
assert prisma.db.litellm_teamtable.update.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_unknown_team_is_403_for_non_proxy_admins_and_404_for_proxy_admins():
|
||||
"""Now that any authenticated caller reaches the handler, 'team not found' must not leak team ids."""
|
||||
import contextlib
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(ProxyException) as denied:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="no-such-team", team_alias="renamed"),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
with pytest.raises(ProxyException) as missing:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="no-such-team", team_alias="renamed"),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_PROXY_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert str(denied.value.code) == "403"
|
||||
assert "do not have access to this team" in str(denied.value.message)
|
||||
assert "no-such-team" not in str(denied.value.message)
|
||||
assert str(missing.value.code) == "404"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_team_access_ranks_proxy_admin_then_org_admin_then_team_admin():
|
||||
from litellm.proxy.management_endpoints.team_endpoints import _resolve_team_access
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team-1",
|
||||
organization_id="org-1",
|
||||
members_with_roles=[Member(user_id="team-admin", role="admin")],
|
||||
)
|
||||
roster_admin = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="team-admin")
|
||||
outsider = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="someone-else")
|
||||
org_lookup = AsyncMock(return_value=False)
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.team_endpoints._is_user_org_admin_for_team", org_lookup): # test-quality-ok: the org-admin lookup needs a real prisma client this file's MagicMock cannot provide
|
||||
assert await _resolve_team_access(team_obj=team, user_api_key_dict=_PROXY_ADMIN_CALLER) == "proxy_admin"
|
||||
assert org_lookup.await_count == 0
|
||||
assert await _resolve_team_access(team_obj=team, user_api_key_dict=roster_admin) == "team_admin"
|
||||
assert await _resolve_team_access(team_obj=team, user_api_key_dict=outsider) is None
|
||||
org_lookup.return_value = True
|
||||
assert await _resolve_team_access(team_obj=team, user_api_key_dict=roster_admin) == "org_admin"
|
||||
|
||||
|
||||
_ROSTER_ADMIN_CALLER = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="admin-1")
|
||||
_MEMBER_CALLER = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="member-1")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"caller, org_admin, enabled_fields, expected",
|
||||
[
|
||||
pytest.param(_PROXY_ADMIN_CALLER, False, (), {"kind": "unrestricted"}, id="proxy-admin"),
|
||||
pytest.param(
|
||||
UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, user_id="viewer"),
|
||||
False,
|
||||
("tpm_limit",),
|
||||
{"kind": "none"},
|
||||
id="proxy-admin-viewer",
|
||||
),
|
||||
pytest.param(_ROSTER_ADMIN_CALLER, True, (), {"kind": "unrestricted"}, id="org-admin-who-is-also-team-admin"),
|
||||
pytest.param(_ROSTER_ADMIN_CALLER, False, (), {"kind": "team_admin_disabled"}, id="team-admin-nothing-enabled"),
|
||||
pytest.param(
|
||||
_ROSTER_ADMIN_CALLER,
|
||||
False,
|
||||
("tpm_limit",),
|
||||
{"kind": "team_admin", "editable_fields": ["tpm_limit"]},
|
||||
id="team-admin-field-enabled",
|
||||
),
|
||||
pytest.param(_MEMBER_CALLER, False, ("tpm_limit",), {"kind": "none"}, id="plain-member"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_info_reports_what_the_caller_may_edit(caller, org_admin, enabled_fields, expected):
|
||||
"""The dashboard gates its edit form on this field instead of guessing the caller's role from the org list,
|
||||
which is premium-gated and can be empty for a dual-role org admin."""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy.management_endpoints import team_endpoints
|
||||
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1",
|
||||
organization_id="org-1",
|
||||
members_with_roles=[Member(user_id="admin-1", role="admin"), Member(user_id="member-1", role="user")],
|
||||
)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
|
||||
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.get_data = AsyncMock(return_value=[])
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: no seam on team_info
|
||||
patch.object(team_endpoints, "get_all_team_memberships", AsyncMock(return_value=[])), # test-quality-ok: no seam on team_info
|
||||
patch.object( # test-quality-ok: the org-admin lookup needs a real prisma client this file's MagicMock cannot provide
|
||||
team_endpoints, "_is_user_org_admin_for_team", AsyncMock(return_value=org_admin)
|
||||
),
|
||||
_team_admin_may_edit(*enabled_fields),
|
||||
):
|
||||
response = await team_endpoints.team_info(
|
||||
http_request=MagicMock(spec=Request),
|
||||
team_id="team-1",
|
||||
user_api_key_dict=caller,
|
||||
)
|
||||
|
||||
assert response["team_info"].caller_edit_access.model_dump(mode="json") == expected
|
||||
|
|
|
|||
|
|
@ -1983,6 +1983,144 @@ class TestBedrockAgentRuntimePassthroughToggle:
|
|||
create_route.assert_called_once()
|
||||
|
||||
|
||||
class TestBedrockAgentRuntimePassthroughVirtualKeyLeak:
|
||||
|
||||
VKEY: Final = "sk-litellm-victim-key"
|
||||
MASTER_KEY: Final = "sk-master-1234"
|
||||
ENDPOINT: Final = "knowledgebases/KB1234567/retrieve"
|
||||
AMBIENT_AWS_ENV: Final = (
|
||||
"AWS_BEARER_TOKEN_BEDROCK",
|
||||
"AWS_SESSION_TOKEN",
|
||||
"AWS_SESSION_NAME",
|
||||
"AWS_PROFILE_NAME",
|
||||
"AWS_ROLE_NAME",
|
||||
"AWS_WEB_IDENTITY_TOKEN",
|
||||
"AWS_STS_ENDPOINT",
|
||||
"AWS_EXTERNAL_ID",
|
||||
)
|
||||
|
||||
async def _upstream_headers(self, monkeypatch, headers: list[tuple[bytes, bytes]]) -> dict:
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import HttpPassThroughEndpointHelpers
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", self.MASTER_KEY)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
for ambient in self.AMBIENT_AWS_ENV:
|
||||
monkeypatch.delenv(ambient, raising=False)
|
||||
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "ak")
|
||||
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "sk")
|
||||
monkeypatch.setenv("AWS_REGION_NAME", "us-east-1")
|
||||
caller: Final = UserAPIKeyAuth(api_key=self.VKEY)
|
||||
|
||||
async def receive():
|
||||
return {"type": "http.request", "body": b'{"retrievalQuery": {"text": "hi"}}', "more_body": False}
|
||||
|
||||
request: Final = Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": f"/bedrock/{self.ENDPOINT}",
|
||||
"headers": headers,
|
||||
"query_string": b"",
|
||||
},
|
||||
receive=receive,
|
||||
)
|
||||
captured: dict = {}
|
||||
|
||||
def fake_create_pass_through_route(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return AsyncMock(return_value={"status": "success"})
|
||||
|
||||
module: Final = "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints"
|
||||
with (
|
||||
patch(f"{module}.create_request_copy", Mock()),
|
||||
patch(f"{module}.create_pass_through_route", side_effect=fake_create_pass_through_route),
|
||||
):
|
||||
await bedrock_proxy_route(
|
||||
endpoint=self.ENDPOINT,
|
||||
request=request,
|
||||
fastapi_response=Response(),
|
||||
user_api_key_dict=caller,
|
||||
)
|
||||
return HttpPassThroughEndpointHelpers.forward_headers_from_request(
|
||||
request_headers=dict(request.headers),
|
||||
headers=dict(captured["custom_headers"] or {}),
|
||||
forward_headers=captured.get("_forward_headers", False),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _blob(upstream: dict) -> str:
|
||||
return " ".join(f"{name}:{value}" for name, value in upstream.items())
|
||||
|
||||
@staticmethod
|
||||
def _names_matching(upstream: dict, lowercase_name: str) -> list[str]:
|
||||
return [name for name in upstream if name.lower() == lowercase_name]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"header_name", ["x-api-key", "x-litellm-api-key", "api-key", "x-goog-api-key", "ocp-apim-subscription-key"]
|
||||
)
|
||||
async def test_virtual_key_in_a_credential_header_never_reaches_aws(self, monkeypatch, header_name: str):
|
||||
upstream: Final = await self._upstream_headers(
|
||||
monkeypatch,
|
||||
[
|
||||
(header_name.encode(), self.VKEY.encode()),
|
||||
(b"content-type", b"application/json"),
|
||||
(b"x-request-id", b"trace-1"),
|
||||
],
|
||||
)
|
||||
|
||||
assert self.VKEY not in self._blob(upstream)
|
||||
assert self._names_matching(upstream, header_name) == []
|
||||
assert upstream["x-request-id"] == "trace-1", "a benign caller header still reaches AWS"
|
||||
assert upstream["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
assert self._names_matching(upstream, "content-type") == ["Content-Type"], "the signed header is the only one"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_credential_headers_are_dropped_by_name_even_when_they_carry_someone_elses_key(self, monkeypatch):
|
||||
other_key: Final = "sk-other-tenant-key"
|
||||
upstream: Final = await self._upstream_headers(
|
||||
monkeypatch,
|
||||
[
|
||||
(b"x-api-key", other_key.encode()),
|
||||
(b"x-litellm-api-key", other_key.encode()),
|
||||
(b"x-request-id", b"trace-3"),
|
||||
],
|
||||
)
|
||||
|
||||
assert other_key not in self._blob(upstream)
|
||||
assert self._names_matching(upstream, "x-api-key") == []
|
||||
assert self._names_matching(upstream, "x-litellm-api-key") == []
|
||||
assert upstream["x-request-id"] == "trace-3"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_virtual_key_in_authorization_bearer_is_replaced_by_the_sigv4_signature(self, monkeypatch):
|
||||
upstream: Final = await self._upstream_headers(
|
||||
monkeypatch,
|
||||
[(b"authorization", f"Bearer {self.VKEY}".encode()), (b"content-type", b"application/json")],
|
||||
)
|
||||
|
||||
assert self.VKEY not in self._blob(upstream)
|
||||
assert self._names_matching(upstream, "authorization") == ["Authorization"]
|
||||
assert upstream["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticated_secrets_in_any_other_header_never_reach_aws(self, monkeypatch):
|
||||
upstream: Final = await self._upstream_headers(
|
||||
monkeypatch,
|
||||
[
|
||||
(b"x-api-key", self.VKEY.encode()),
|
||||
(b"x-forwarded-key", self.VKEY.encode()),
|
||||
(b"x-operator-token", self.MASTER_KEY.encode()),
|
||||
(b"x-request-id", b"trace-2"),
|
||||
],
|
||||
)
|
||||
|
||||
assert self.VKEY not in self._blob(upstream) and self.MASTER_KEY not in self._blob(upstream)
|
||||
assert self._names_matching(upstream, "x-forwarded-key") == []
|
||||
assert self._names_matching(upstream, "x-operator-token") == []
|
||||
assert upstream["x-request-id"] == "trace-2"
|
||||
|
||||
|
||||
class TestLLMPassthroughFactoryProxyRoute:
|
||||
@pytest.mark.asyncio
|
||||
async def test_llm_passthrough_factory_proxy_route_success(self):
|
||||
|
|
|
|||
|
|
@ -6143,3 +6143,42 @@ async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_err
|
|||
)
|
||||
|
||||
assert (raised.value.type, raised.value.param, raised.value.code) == ("invalid_request_error", None, "400")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_pass_through_endpoint_failure_carries_the_callers_litellm_call_id(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
):
|
||||
call_id = "lit7836-pass-through-call-id"
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"])
|
||||
proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
async def fake_add_litellm_data_to_request(**kwargs: object) -> object:
|
||||
return kwargs["data"]
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.add_litellm_data_to_request", fake_add_litellm_data_to_request)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
|
||||
request = MagicMock(spec=Request)
|
||||
request.headers = Headers({"x-litellm-call-id": call_id})
|
||||
request.body = AsyncMock(
|
||||
return_value=json.dumps({"model": "unknown-model", "messages": [{"role": "user", "content": "hi"}]}).encode()
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"), pytest.raises(ProxyException) as raised:
|
||||
await chat_completion_pass_through_endpoint(
|
||||
fastapi_response=Response(),
|
||||
request=request,
|
||||
adapter_id="anthropic",
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
)
|
||||
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
record = next(r for r in caplog.records if "Exception occured" in r.getMessage())
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
|
|
|||
|
|
@ -3938,3 +3938,58 @@ async def test_ProxyConfig__init_guardrails_in_db_skips_only_the_unloadable_row(
|
|||
|
||||
assert sorted(handler.IN_MEMORY_GUARDRAILS) == ["first", "last"]
|
||||
assert handler.reconciled_with == [{"first", "broken", "last"}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# add_deployment: UI settings convergence
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_deployment_re_reads_ui_settings_so_other_pods_converge(monkeypatch):
|
||||
"""The periodic config reload picks up a UI setting written through another pod.
|
||||
|
||||
Startup used to be the only read, so a proxy admin flipping a runtime flag reached the pod
|
||||
that served the PATCH and nowhere else until every other pod restarted.
|
||||
"""
|
||||
general_settings: Dict[str, Any] = {"allow_agents_for_team_admins": False}
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[])
|
||||
prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[])
|
||||
prisma_client.db.litellm_uisettings.find_unique = AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
ui_settings=json.dumps({"allow_agents_for_team_admins": True, "enable_chat_ui": False})
|
||||
)
|
||||
)
|
||||
|
||||
config = ProxyConfig()
|
||||
config._should_load_db_object = MagicMock(return_value=False)
|
||||
config._init_non_llm_objects_in_db = AsyncMock()
|
||||
|
||||
await config.add_deployment(prisma_client=prisma_client, proxy_logging_obj=MagicMock())
|
||||
|
||||
prisma_client.db.litellm_uisettings.find_unique.assert_awaited_once_with(where={"id": "ui_settings"})
|
||||
assert general_settings["allow_agents_for_team_admins"] is True
|
||||
assert "enable_chat_ui" not in general_settings
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_deployment_syncs_ui_settings_even_when_the_model_reconcile_fails(monkeypatch):
|
||||
"""A broken model reconcile must not strand every pod on stale settings."""
|
||||
general_settings: Dict[str, Any] = {"allow_agents_for_team_admins": False}
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_uisettings.find_unique = AsyncMock(
|
||||
return_value=SimpleNamespace(ui_settings={"allow_agents_for_team_admins": True})
|
||||
)
|
||||
|
||||
config = ProxyConfig()
|
||||
config._should_load_db_object = MagicMock(side_effect=RuntimeError("db down"))
|
||||
|
||||
await config.add_deployment(prisma_client=prisma_client, proxy_logging_obj=MagicMock())
|
||||
|
||||
assert general_settings["allow_agents_for_team_admins"] is True
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import copy
|
|||
from collections.abc import Callable
|
||||
from contextlib import AbstractContextManager
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -185,7 +185,6 @@ def test_v1_model_info_no_model_list_error(client, auth_as, null_router, path):
|
|||
assert "LLM Model List not loaded" in response.text
|
||||
|
||||
|
||||
|
||||
def test_get_proxy_model_info_surfaces_supports_parallel_function_calling(local_model_cost_map):
|
||||
"""``GET /v1/model/info`` enriches each deployment through ``_get_proxy_model_info``; a registry
|
||||
entry declaring parallel function calling must land in ``model_info`` instead of null."""
|
||||
|
|
@ -218,9 +217,7 @@ def test_v1_model_info_star_wildcard_filter_keeps_provider_expansion(monkeypatch
|
|||
router.get_model_list = MagicMock(return_value=[deployment])
|
||||
monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models)
|
||||
|
||||
expanded_deployments = proxy_server.expand_wildcard_deployments_for_model_info(
|
||||
[deployment]
|
||||
)
|
||||
expanded_deployments = proxy_server.expand_wildcard_deployments_for_model_info([deployment])
|
||||
allowed_model_names = proxy_server._get_v1_model_info_allowed_model_names(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
|
|
@ -365,6 +362,80 @@ def test_model_group_info_invalid_method(client, auth_as, null_router):
|
|||
assert len(response.content) > 0
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model_group_info_router(monkeypatch):
|
||||
from litellm.types.proxy.management_endpoints.model_management_endpoints import ModelGroupInfoProxy
|
||||
|
||||
model_names = ["gpt-4", "claude-3"]
|
||||
router = MagicMock()
|
||||
router.get_model_names.return_value = model_names
|
||||
router.get_model_access_groups.return_value = {}
|
||||
router.get_model_list.return_value = []
|
||||
|
||||
def model_group_info(*, llm_router, all_models_str, model_group):
|
||||
return [ModelGroupInfoProxy(model_group=name, providers=[]) for name in all_models_str]
|
||||
|
||||
async def append_agents_to_model_group(*, model_groups, user_api_key_dict):
|
||||
return model_groups
|
||||
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(proxy_server, "llm_model_list", [{"model_name": name} for name in model_names])
|
||||
monkeypatch.setattr(proxy_server, "user_model", None)
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", None)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", None)
|
||||
monkeypatch.setattr(proxy_server, "_get_model_group_info", model_group_info)
|
||||
|
||||
from litellm.proxy.agent_endpoints import model_list_helpers
|
||||
|
||||
monkeypatch.setattr(
|
||||
model_list_helpers,
|
||||
"append_agents_to_model_group",
|
||||
AsyncMock(side_effect=append_agents_to_model_group),
|
||||
)
|
||||
return router
|
||||
|
||||
|
||||
@pytest.mark.parametrize("admin_role", ["proxy_admin", "proxy_admin_viewer"])
|
||||
def test_model_group_info_proxy_admin_ignores_key_model_restriction(
|
||||
client, auth_as, model_group_info_router, admin_role
|
||||
):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
with auth_as(LitellmUserRoles(admin_role), models=["no-default-models"]):
|
||||
response = client.get("/model_group/info")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [model["model_group"] for model in response.json()["data"]] == ["gpt-4", "claude-3"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("admin_role", ["proxy_admin", "proxy_admin_viewer"])
|
||||
def test_model_group_info_proxy_admin_expands_wildcard_deployments(client, auth_as, model_group_info_router, admin_role):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.auth.model_checks import get_known_models_from_wildcard
|
||||
|
||||
model_group_info_router.get_model_names.return_value = ["gpt-4", "anthropic/*"]
|
||||
known_anthropic_models = get_known_models_from_wildcard(wildcard_model="anthropic/*")
|
||||
assert known_anthropic_models
|
||||
|
||||
with auth_as(LitellmUserRoles(admin_role), models=["no-default-models"]):
|
||||
response = client.get("/model_group/info")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [model["model_group"] for model in response.json()["data"]] == ["gpt-4", *known_anthropic_models]
|
||||
|
||||
|
||||
def test_model_group_info_internal_user_key_model_restriction_applies(client, auth_as, model_group_info_router):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
with auth_as(LitellmUserRoles.INTERNAL_USER, models=["gpt-4"]):
|
||||
response = client.get("/model_group/info")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [model["model_group"] for model in response.json()["data"]] == ["gpt-4"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /v2/model/info?exclude_auto_routers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -456,14 +527,10 @@ def test_v2_model_info_exclude_auto_routers_shrinks_total_count(client, auth_as,
|
|||
assert len(payload["data"]) == payload["total_count"]
|
||||
|
||||
|
||||
def test_v2_model_info_exclude_auto_routers_paginates_over_the_filtered_set(
|
||||
client, auth_as, mixed_auto_router_router
|
||||
):
|
||||
def test_v2_model_info_exclude_auto_routers_paginates_over_the_filtered_set(client, auth_as, mixed_auto_router_router):
|
||||
"""Page size applies to the filtered list, so no page silently comes back short."""
|
||||
with auth_as():
|
||||
response = client.get(
|
||||
"/v2/model/info", params={"exclude_auto_routers": "true", "page": 1, "size": 1}
|
||||
)
|
||||
response = client.get("/v2/model/info", params={"exclude_auto_routers": "true", "page": 1, "size": 1})
|
||||
payload = response.json()
|
||||
assert payload["total_count"] == 2
|
||||
assert payload["total_pages"] == 2
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ Tests for rerank_endpoints/endpoints.py response headers.
|
|||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Iterator
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -10,6 +12,7 @@ from fastapi import HTTPException, Request, Response
|
|||
|
||||
import litellm.proxy.common_request_processing as common_request_processing_mod
|
||||
import litellm.proxy.proxy_server as proxy_server_mod
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.rerank_endpoints.endpoints import rerank
|
||||
from litellm.types.utils import RerankResponse
|
||||
|
|
@ -28,7 +31,7 @@ HIDDEN_PARAMS = {
|
|||
}
|
||||
|
||||
|
||||
def _build_request() -> Request:
|
||||
def _build_request(headers: tuple[tuple[bytes, bytes], ...] = ()) -> Request:
|
||||
body = json.dumps({"model": "rerank-model", "query": "q", "documents": ["a", "b"]}).encode()
|
||||
|
||||
async def receive():
|
||||
|
|
@ -39,7 +42,7 @@ def _build_request() -> Request:
|
|||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/rerank",
|
||||
"headers": [(b"content-type", b"application/json")],
|
||||
"headers": [(b"content-type", b"application/json"), *headers],
|
||||
"query_string": b"",
|
||||
},
|
||||
receive=receive,
|
||||
|
|
@ -56,7 +59,7 @@ async def _call_rerank(hidden_params: dict = HIDDEN_PARAMS) -> Response:
|
|||
proxy_logging_obj.update_request_status = AsyncMock()
|
||||
|
||||
async def fake_add_litellm_data_to_request(**kwargs):
|
||||
return {**kwargs["data"], "litellm_call_id": "call-123"}
|
||||
return dict(kwargs["data"])
|
||||
|
||||
async def fake_route_request(**kwargs):
|
||||
async def _call():
|
||||
|
|
@ -72,7 +75,7 @@ async def _call_rerank(hidden_params: dict = HIDDEN_PARAMS) -> Response:
|
|||
patch.object(proxy_server_mod, "version", "1.2.3"), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
|
||||
):
|
||||
await rerank(
|
||||
request=_build_request(),
|
||||
request=_build_request(headers=((b"x-litellm-call-id", b"call-123"),)),
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
)
|
||||
|
|
@ -121,7 +124,11 @@ async def test_rerank_omits_detailed_timing_headers_when_disabled():
|
|||
|
||||
|
||||
async def _rerank_failure(
|
||||
failure: Exception, *, raised_before_routing: bool, monkeypatch: pytest.MonkeyPatch
|
||||
failure: Exception,
|
||||
*,
|
||||
raised_before_routing: bool,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
headers: tuple[tuple[bytes, bytes], ...] = (),
|
||||
) -> ProxyException:
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(
|
||||
|
|
@ -143,13 +150,45 @@ async def _rerank_failure(
|
|||
|
||||
with pytest.raises(ProxyException) as raised:
|
||||
await rerank(
|
||||
request=_build_request(),
|
||||
request=_build_request(headers),
|
||||
fastapi_response=Response(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
)
|
||||
return raised.value
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def propagating_proxy_logger() -> Iterator[None]:
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_log_carries_the_callers_litellm_call_id(
|
||||
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, propagating_proxy_logger: None
|
||||
) -> None:
|
||||
"""LIT-7836: the /rerank error line must carry the same litellm_call_id the client
|
||||
sent, both in the rendered message and as a structured log record field."""
|
||||
call_id = "rerank-call-7836"
|
||||
failure = HTTPException(status_code=401, detail={"error": "invalid api key"})
|
||||
|
||||
with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"):
|
||||
raised = await _rerank_failure(
|
||||
failure,
|
||||
raised_before_routing=False,
|
||||
monkeypatch=monkeypatch,
|
||||
headers=((b"x-litellm-call-id", call_id.encode()),),
|
||||
)
|
||||
|
||||
assert raised.headers["x-litellm-call-id"] == call_id
|
||||
record = next(r for r in caplog.records if "Exception occured" in r.getMessage())
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_model_the_router_cannot_serve_answers_an_openai_typed_error(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A bare HTTPException carries no type or param, so the tail used to ship the
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ from litellm.proxy.common_request_processing import (
|
|||
create_response,
|
||||
)
|
||||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy._types import UserAPIKeyAuth as ProxyUserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
|
@ -6365,6 +6366,206 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
|
|||
call_type="acompletion",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _v3_limiter_rig(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
user_api_key_dict: ProxyUserAPIKeyAuth,
|
||||
fallbacks: list[dict[str, list[str]]],
|
||||
) -> tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]]:
|
||||
"""Real v3 limiter (the default ``parallel_request_limiter``) wired in through the
|
||||
``proxy_logging_obj`` seam, so ``common_processing_pre_call_logic`` runs for real:
|
||||
``add_litellm_data_to_request`` with a live OTel span, ``function_setup``, then the limiter."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
|
||||
from litellm.proxy.utils import InternalUsageCache
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
limiter_models: list[str] = []
|
||||
|
||||
async def run_limiter(
|
||||
user_api_key_dict: ProxyUserAPIKeyAuth, data: dict[str, object], call_type: str
|
||||
) -> dict[str, object]:
|
||||
limiter_models.append(str(data["model"]))
|
||||
await limiter.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type=call_type,
|
||||
)
|
||||
return data
|
||||
|
||||
proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=run_limiter)
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": group, "litellm_params": {"model": "openai/gpt-4.1-nano", "api_key": "fake"}}
|
||||
for chain in fallbacks
|
||||
for group in (*chain.keys(), *(m for models in chain.values() for m in models))
|
||||
],
|
||||
fallbacks=fallbacks,
|
||||
)
|
||||
return proxy_logging_obj, router, proxy_server.ProxyConfig(), limiter_models
|
||||
|
||||
@staticmethod
|
||||
def _otel_key(
|
||||
rpm_limit: int | None = None,
|
||||
model_rpm_limit: dict[str, int] | None = None,
|
||||
disable_fallbacks: bool = False,
|
||||
) -> ProxyUserAPIKeyAuth:
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
||||
span = TracerProvider().get_tracer("test").start_span("proxy-request")
|
||||
return ProxyUserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
parent_otel_span=span,
|
||||
rpm_limit=rpm_limit,
|
||||
metadata={
|
||||
**({"model_rpm_limit": model_rpm_limit} if model_rpm_limit else {}),
|
||||
**({"disable_fallbacks": True} if disable_fallbacks else {}),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _chat_request() -> Request:
|
||||
return Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []})
|
||||
|
||||
async def _pre_call(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
user_api_key_dict: ProxyUserAPIKeyAuth,
|
||||
rig: tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]],
|
||||
) -> tuple[ProxyBaseLLMRequestProcessing, tuple[dict[str, object], LiteLLMLoggingObj]]:
|
||||
proxy_logging_obj, router, proxy_config, _ = rig
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
result = await processor._pre_call_with_fallbacks(
|
||||
request=self._chat_request(),
|
||||
general_settings={},
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
version=None,
|
||||
proxy_config=proxy_config,
|
||||
user_model=None,
|
||||
user_temperature=None,
|
||||
user_request_timeout=None,
|
||||
user_max_tokens=None,
|
||||
user_api_base=None,
|
||||
model=None,
|
||||
route_type="acompletion",
|
||||
llm_router=router,
|
||||
)
|
||||
return processor, result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v3_limiter_with_otel_span_falls_back_from_client_request(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Customer path: OTel on, per-key model RPM cap on the primary, a router fallback configured.
|
||||
The first pass enriches ``data["metadata"]`` with the live span, then the limiter raises. The
|
||||
fallback pass must start from the client's request again, so ``add_litellm_data_to_request``
|
||||
never deep-copies the span (the ``cannot pickle '_thread.RLock'`` 500)."""
|
||||
primary_model = "gpt-4.1"
|
||||
fallback_model = "gpt-4.1-mini"
|
||||
key = self._otel_key(model_rpm_limit={primary_model: 1})
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
|
||||
def client_request() -> dict[str, object]:
|
||||
return {
|
||||
"model": primary_model,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": {"tags": ["client-tag"]},
|
||||
}
|
||||
|
||||
_, (first_data, _) = await self._pre_call(client_request(), key, rig)
|
||||
processor, (data, logging_obj) = await self._pre_call(client_request(), key, rig)
|
||||
|
||||
assert first_data["model"] == primary_model
|
||||
assert data["model"] == fallback_model
|
||||
assert processor.data is data
|
||||
assert data["litellm_logging_obj"] is logging_obj
|
||||
assert logging_obj.model == fallback_model
|
||||
requester_metadata = data["metadata"]["requester_metadata"]
|
||||
assert requester_metadata["tags"] == ["client-tag"]
|
||||
assert "litellm_parent_otel_span" not in requester_metadata
|
||||
assert "user_api_key_auth" not in requester_metadata
|
||||
assert data["metadata"]["litellm_parent_otel_span"] is key.parent_otel_span
|
||||
assert rig[3] == [primary_model, primary_model, fallback_model]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v3_limiter_with_otel_span_returns_429_when_fallbacks_exhausted(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
primary_model = "gpt-4.1"
|
||||
fallback_model = "gpt-4.1-mini"
|
||||
key = self._otel_key(rpm_limit=1)
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
request = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
processor = ProxyBaseLLMRequestProcessing(data=dict(request))
|
||||
with pytest.raises(ProxyRateLimitError) as exc_info:
|
||||
await processor._pre_call_with_fallbacks(
|
||||
request=self._chat_request(),
|
||||
general_settings={},
|
||||
proxy_logging_obj=rig[0],
|
||||
user_api_key_dict=key,
|
||||
version=None,
|
||||
proxy_config=rig[2],
|
||||
user_model=None,
|
||||
user_temperature=None,
|
||||
user_request_timeout=None,
|
||||
user_max_tokens=None,
|
||||
user_api_base=None,
|
||||
model=None,
|
||||
route_type="acompletion",
|
||||
llm_router=rig[1],
|
||||
)
|
||||
|
||||
assert rig[3] == [primary_model, primary_model, fallback_model]
|
||||
assert exc_info.value.status_code == 429
|
||||
assert "Rate limit exceeded" in str(exc_info.value.detail)
|
||||
assert exc_info.value.headers["retry-after"]
|
||||
assert processor.data["model"] == primary_model
|
||||
assert processor.data["litellm_logging_obj"].model == primary_model
|
||||
assert processor.data["litellm_call_id"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_lookup_uses_alias_resolved_model_group(self, monkeypatch: pytest.MonkeyPatch):
|
||||
primary_model = "gpt-4.1"
|
||||
fallback_model = "gpt-4.1-mini"
|
||||
monkeypatch.setattr(litellm, "model_alias_map", {"my-alias": primary_model})
|
||||
key = self._otel_key(model_rpm_limit={primary_model: 1})
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
request = {"model": "my-alias", "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
_, (data, _) = await self._pre_call(dict(request), key, rig)
|
||||
|
||||
assert data["model"] == fallback_model
|
||||
assert rig[3] == [primary_model, primary_model, fallback_model]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_metadata_disable_fallbacks_returns_429_instead_of_retrying(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""``disable_fallbacks`` set in key metadata only lands on ``data`` during the first
|
||||
pre-call pass (``add_key_level_controls``), so it must be honored after that pass."""
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
primary_model = "gpt-4.1"
|
||||
fallback_model = "gpt-4.1-mini"
|
||||
key = self._otel_key(model_rpm_limit={primary_model: 1}, disable_fallbacks=True)
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
request = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
with pytest.raises(ProxyRateLimitError) as exc_info:
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert rig[3] == [primary_model, primary_model]
|
||||
|
||||
|
||||
class _RecordingSuccessLogger(CustomLogger):
|
||||
def __init__(self):
|
||||
|
|
@ -8212,7 +8413,7 @@ def test_log_llm_api_exception_traceback_only_for_unexpected_errors(exc, expect_
|
|||
"""Regression for LIT-6043: expected 4xx errors log without formatting a
|
||||
traceback; unexpected errors keep logger.exception behavior."""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_request_processing import _log_llm_api_exception
|
||||
from litellm.proxy.common_request_processing import log_llm_api_exception
|
||||
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
|
|
@ -8220,7 +8421,7 @@ def test_log_llm_api_exception_traceback_only_for_unexpected_errors(exc, expect_
|
|||
try:
|
||||
raise exc
|
||||
except Exception as raised:
|
||||
_log_llm_api_exception(raised, "call-id-for-traceback-test")
|
||||
log_llm_api_exception(raised, "call-id-for-traceback-test")
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
|
|
@ -8778,14 +8979,14 @@ class TestErrorLogCarriesCallId:
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_request_processing import (
|
||||
_CLIENT_DISCONNECT_DETAIL,
|
||||
_log_llm_api_exception,
|
||||
log_llm_api_exception,
|
||||
)
|
||||
|
||||
call_id: Final = str(uuid.uuid4())
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
with caplog.at_level("INFO", logger="LiteLLM Proxy"):
|
||||
_log_llm_api_exception(
|
||||
log_llm_api_exception(
|
||||
HTTPException(status_code=499, detail=_CLIENT_DISCONNECT_DETAIL),
|
||||
call_id,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import contextlib
|
||||
import importlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import socket
|
||||
|
|
@ -19,7 +20,7 @@ import fastapi.routing
|
|||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.encoders import jsonable_encoder
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from fastapi.testclient import TestClient
|
||||
|
|
@ -13003,6 +13004,144 @@ async def test_moderations_response_carries_litellm_call_id_header():
|
|||
assert fastapi_response.headers["x-litellm-model-id"] == "mod-deployment-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_moderations_failure_log_carries_the_callers_litellm_call_id(caplog):
|
||||
"""LIT-7836: the /v1/moderations error line must carry the litellm_call_id the
|
||||
client sent, rendered in the message and as a structured log record field."""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
call_id = "moderations-call-7836"
|
||||
|
||||
async def passthrough_add_litellm_data(data, **kwargs):
|
||||
return data
|
||||
|
||||
request = MagicMock()
|
||||
request.headers = {"x-litellm-call-id": call_id}
|
||||
request.body = AsyncMock(return_value=b'{"input": "hi"}')
|
||||
fake_logging = MagicMock()
|
||||
fake_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
fake_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
with (
|
||||
patch.object(proxy_server_module, "add_litellm_data_to_request", new=passthrough_add_litellm_data), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(proxy_server_module, "route_request", new=AsyncMock(side_effect=Exception("bad key"))), # test-quality-ok: fakes the provider failure so the real route's error log is observable
|
||||
patch.object(proxy_server_module, "proxy_logging_obj", new=fake_logging), # test-quality-ok: module global, no injection point
|
||||
caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"),
|
||||
pytest.raises(ProxyException) as raised,
|
||||
):
|
||||
await proxy_server_module.moderations(
|
||||
request=request,
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", spend=0.0),
|
||||
)
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
record = next(r for r in caplog.records if "Exception occured" in r.getMessage())
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_moderations_unparseable_body_bills_the_callers_litellm_call_id():
|
||||
"""LIT-7836: a body that fails to parse must still hand the failure hook the
|
||||
litellm_call_id the response header answers with, so the spend row is findable."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
call_id = "moderations-early-7836"
|
||||
|
||||
request = MagicMock()
|
||||
request.headers = {"x-litellm-call-id": call_id}
|
||||
request.body = AsyncMock(return_value=b'{"input": ')
|
||||
fake_logging = MagicMock()
|
||||
fake_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(proxy_server_module, "proxy_logging_obj", new=fake_logging), # test-quality-ok: module global, no injection point
|
||||
pytest.raises(ProxyException) as raised,
|
||||
):
|
||||
await proxy_server_module.moderations(
|
||||
request=request,
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", spend=0.0),
|
||||
)
|
||||
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
hook_request_data = fake_logging.post_call_failure_hook.await_args.kwargs["request_data"]
|
||||
assert hook_request_data["litellm_call_id"] == call_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_moderations_already_shaped_failure_answers_with_the_callers_litellm_call_id():
|
||||
"""LIT-7836: a ProxyException raised inside /v1/moderations is re-raised unwrapped but still
|
||||
answers with the caller's x-litellm-call-id so the client can join it to the error log."""
|
||||
call_id = "moderations-call-7836-shaped"
|
||||
exc = ProxyException(message="budget exceeded", type=ProxyErrorTypes.budget_exceeded, param="key", code=402)
|
||||
|
||||
request = MagicMock()
|
||||
request.headers = {"x-litellm-call-id": call_id}
|
||||
request.body = AsyncMock(return_value=b'{"input": "hi"}')
|
||||
fake_logging = MagicMock()
|
||||
fake_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(proxy_server_module, "add_litellm_data_to_request", new=AsyncMock(side_effect=exc)), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(proxy_server_module, "proxy_logging_obj", new=fake_logging), # test-quality-ok: module global, no injection point
|
||||
pytest.raises(ProxyException) as raised,
|
||||
):
|
||||
await proxy_server_module.moderations(
|
||||
request=request,
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", spend=0.0),
|
||||
)
|
||||
|
||||
assert raised.value is exc
|
||||
assert raised.value.code == "402"
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"exc",
|
||||
[
|
||||
HTTPException(status_code=401, detail="bad key"),
|
||||
ProxyException(message="budget exceeded", type=ProxyErrorTypes.budget_exceeded, param="key", code=402),
|
||||
],
|
||||
ids=["http_exception", "proxy_exception"],
|
||||
)
|
||||
async def test_audio_speech_already_shaped_failure_answers_with_the_callers_litellm_call_id(exc: Exception):
|
||||
"""LIT-7836: /v1/audio/speech re-raises HTTP and proxy shaped failures unchanged, and they must
|
||||
still answer with the caller's x-litellm-call-id."""
|
||||
call_id = "speech-call-7836-shaped"
|
||||
|
||||
request = MagicMock()
|
||||
request.headers = {"x-litellm-call-id": call_id}
|
||||
request.body = AsyncMock(return_value=b'{"model": "tts-1", "input": "hi", "voice": "alloy"}')
|
||||
fake_logging = MagicMock()
|
||||
fake_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(proxy_server_module, "add_litellm_data_to_request", new=AsyncMock(side_effect=exc)), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(proxy_server_module, "proxy_logging_obj", new=fake_logging), # test-quality-ok: module global, no injection point
|
||||
pytest.raises(type(exc)) as raised,
|
||||
):
|
||||
await proxy_server_module.audio_speech(
|
||||
request=request,
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", spend=0.0),
|
||||
)
|
||||
|
||||
if isinstance(exc, HTTPException):
|
||||
assert (raised.value.status_code, raised.value.detail) == (401, "bad key")
|
||||
else:
|
||||
assert raised.value is exc
|
||||
assert raised.value.headers["x-litellm-call-id"] == call_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_agents_in_db_rebuilds_registry_under_agent_reconcile_lock(monkeypatch):
|
||||
from litellm.proxy.agent_endpoints.agent_registry import (
|
||||
|
|
|
|||
|
|
@ -160,6 +160,37 @@ async def test_proxy_only_error_log_keeps_litellm_metadata_in_litellm_params():
|
|||
assert "litellm_metadata" not in captured["optional_params"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_only_error_log_keeps_the_request_litellm_call_id(monkeypatch: pytest.MonkeyPatch):
|
||||
"""LIT-7836: a route that already stamped the caller's litellm_call_id must
|
||||
keep it when the failure is a proxy-only error, so the spend-log row and the
|
||||
error line share one id instead of a fresh uuid minted here."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
call_id: Final = "caller-supplied-7836"
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def fake_pre_call(self, *args, **kwargs):
|
||||
captured["litellm_call_id"] = self.litellm_call_id
|
||||
|
||||
async def _noop_async_failure(self, *args, **kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(Logging, "pre_call", fake_pre_call)
|
||||
monkeypatch.setattr(Logging, "async_failure_handler", _noop_async_failure)
|
||||
request_data: Final[dict[str, object]] = {"model": "gpt-4o", "input": "hi", "litellm_call_id": call_id}
|
||||
|
||||
await ProxyLogging(user_api_key_cache=DualCache())._handle_logging_proxy_only_error(
|
||||
request_data=request_data,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-bad", request_route="/v1/moderations"),
|
||||
route="/v1/moderations",
|
||||
original_exception=Exception("bad key"),
|
||||
)
|
||||
|
||||
assert request_data["litellm_call_id"] == call_id
|
||||
assert captured["litellm_call_id"] == call_id
|
||||
|
||||
|
||||
def test_get_model_group_info_order():
|
||||
from litellm import Router
|
||||
from litellm.proxy.proxy_server import _get_model_group_info
|
||||
|
|
|
|||
|
|
@ -3266,3 +3266,174 @@ class TestPtuCostAttributionUISetting:
|
|||
assert response.status_code == 400
|
||||
assert "enable_ptu_cost_attribution" in str(response.json()["detail"])
|
||||
assert not mock_prisma.db.litellm_uisettings.upsert.called
|
||||
|
||||
|
||||
class TestTeamAdminEditableTeamFieldsSetting:
|
||||
"""team_admin_editable_team_fields: the proxy-wide allow-list update_team applies to team admins."""
|
||||
|
||||
def _as_proxy_admin(self, monkeypatch):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="test-user-123",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
|
||||
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
return mock_prisma
|
||||
|
||||
def test_patch_rejects_field_names_the_proxy_does_not_support(self, monkeypatch):
|
||||
mock_prisma = self._as_proxy_admin(monkeypatch)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS",
|
||||
frozenset({"tpm_limit"}),
|
||||
)
|
||||
|
||||
try:
|
||||
response = client.patch(
|
||||
"/update/ui_settings",
|
||||
json={"team_admin_editable_team_fields": ["tpm_limit", "blocked", "organization_id"]},
|
||||
)
|
||||
finally:
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
assert response.status_code == 400
|
||||
detail = response.json()["detail"]["error"]
|
||||
assert "['blocked', 'organization_id']" in detail
|
||||
assert "['tpm_limit']" in detail
|
||||
assert not mock_prisma.db.litellm_uisettings.upsert.called
|
||||
|
||||
def test_patch_rejects_a_non_list_value(self, monkeypatch):
|
||||
self._as_proxy_admin(monkeypatch)
|
||||
|
||||
try:
|
||||
response = client.patch("/update/ui_settings", json={"team_admin_editable_team_fields": "tpm_limit"})
|
||||
finally:
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_patch_persists_and_syncs_the_list_to_general_settings(self, monkeypatch):
|
||||
mock_prisma = self._as_proxy_admin(monkeypatch)
|
||||
general_settings: dict = {"team_admin_editable_team_fields": []}
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
|
||||
try:
|
||||
response = client.patch("/update/ui_settings", json={"team_admin_editable_team_fields": ["tpm_limit"]})
|
||||
finally:
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
assert response.status_code == 200
|
||||
stored = json.loads(mock_prisma.db.litellm_uisettings.upsert.call_args.kwargs["data"]["create"]["ui_settings"])
|
||||
assert stored["team_admin_editable_team_fields"] == ["tpm_limit"]
|
||||
assert general_settings["team_admin_editable_team_fields"] == ["tpm_limit"]
|
||||
|
||||
def test_patch_with_an_empty_list_turns_team_admin_editing_off_again(self, monkeypatch):
|
||||
mock_prisma = self._as_proxy_admin(monkeypatch)
|
||||
general_settings: dict = {"team_admin_editable_team_fields": ["tpm_limit"]}
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
|
||||
try:
|
||||
response = client.patch("/update/ui_settings", json={"team_admin_editable_team_fields": []})
|
||||
finally:
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
assert response.status_code == 200
|
||||
stored = json.loads(mock_prisma.db.litellm_uisettings.upsert.call_args.kwargs["data"]["create"]["ui_settings"])
|
||||
assert stored["team_admin_editable_team_fields"] == []
|
||||
assert general_settings["team_admin_editable_team_fields"] == []
|
||||
|
||||
def test_get_reports_the_stored_list_and_advertises_supported_fields(self, mock_auth, monkeypatch):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_db_record = MagicMock()
|
||||
mock_db_record.ui_settings = {"team_admin_editable_team_fields": ["tpm_limit"]}
|
||||
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=mock_db_record)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
general_settings: dict = {}
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
|
||||
response = client.get("/get/ui_settings")
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["values"]["team_admin_editable_team_fields"] == ["tpm_limit"]
|
||||
assert general_settings["team_admin_editable_team_fields"] == ["tpm_limit"]
|
||||
field_schema = data["field_schema"]["properties"]["team_admin_editable_team_fields"]
|
||||
assert field_schema["type"] == "array"
|
||||
assert field_schema["items"]["type"] == "string"
|
||||
assert "tpm_limit" in field_schema["items"]["enum"]
|
||||
|
||||
|
||||
class TestSyncUiSettingsToGeneralSettings:
|
||||
"""The DB re-read each pod runs on startup and on every config reload."""
|
||||
|
||||
def _sync(self):
|
||||
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
|
||||
sync_ui_settings_to_general_settings,
|
||||
)
|
||||
|
||||
return sync_ui_settings_to_general_settings
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_applies_runtime_flags_and_leaves_other_ui_settings_alone(self, monkeypatch):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
general_settings: dict = {"allow_agents_for_team_admins": False}
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
mock_prisma = MagicMock()
|
||||
record = MagicMock()
|
||||
record.ui_settings = json.dumps(
|
||||
{
|
||||
"allow_agents_for_team_admins": True,
|
||||
"team_admin_editable_team_fields": ["tpm_limit"],
|
||||
"enable_chat_ui": False,
|
||||
}
|
||||
)
|
||||
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=record)
|
||||
|
||||
applied = await self._sync()(mock_prisma)
|
||||
|
||||
assert dict(applied) == {
|
||||
"allow_agents_for_team_admins": True,
|
||||
"team_admin_editable_team_fields": ["tpm_limit"],
|
||||
}
|
||||
assert general_settings["allow_agents_for_team_admins"] is True
|
||||
assert general_settings["team_admin_editable_team_fields"] == ["tpm_limit"]
|
||||
assert "enable_chat_ui" not in general_settings
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reads_a_row_the_prisma_client_already_deserialized(self, monkeypatch):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
general_settings: dict = {}
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
mock_prisma = MagicMock()
|
||||
record = MagicMock()
|
||||
record.ui_settings = {"team_admin_editable_team_fields": ["rpm_limit"]}
|
||||
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=record)
|
||||
|
||||
await self._sync()(mock_prisma)
|
||||
|
||||
assert general_settings["team_admin_editable_team_fields"] == ["rpm_limit"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_without_a_stored_row_general_settings_is_left_untouched(self, monkeypatch):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
general_settings: dict = {"allow_agents_for_team_admins": True}
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
applied = await self._sync()(mock_prisma)
|
||||
|
||||
assert dict(applied) == {}
|
||||
assert general_settings == {"allow_agents_for_team_admins": True}
|
||||
|
|
|
|||
|
|
@ -176,6 +176,25 @@ def test_handle_exception_on_proxy_error_path_none_input_wraps_as_500():
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"exc",
|
||||
[
|
||||
HTTPException(status_code=401, detail="bad key"),
|
||||
ValueError("provider boom"),
|
||||
ProxyException(message="already wrapped", type=ProxyErrorTypes.budget_exceeded.value, param="key", code=402),
|
||||
],
|
||||
ids=["http_exception", "generic_exception", "already_proxy_exception"],
|
||||
)
|
||||
def test_handle_exception_on_proxy_returns_the_litellm_call_id_header(exc: Exception):
|
||||
result = handle_exception_on_proxy(exc, "call-7836")
|
||||
|
||||
assert result.headers == {"x-litellm-call-id": "call-7836"}
|
||||
|
||||
|
||||
def test_handle_exception_on_proxy_sends_no_call_id_header_when_the_request_has_none():
|
||||
assert handle_exception_on_proxy(ValueError("provider boom")).headers == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_exception_on_proxy_read_only_transaction_forces_writer_recreate(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
|
|
|
|||
|
|
@ -115,13 +115,13 @@ def _respx_interceptable_httpx_client(monkeypatch):
|
|||
],
|
||||
)
|
||||
def test_resolver_opt_in_gates_openai_like_config(model_info, expected_type):
|
||||
config = _resolve_responses_api_provider_config("my-model", "custom_openai", model_info)
|
||||
config = _resolve_responses_api_provider_config("my-model", "custom_openai", model_info, None)
|
||||
assert type(config) is expected_type
|
||||
|
||||
|
||||
def test_resolver_keeps_native_provider_config():
|
||||
"""`openai/` already routes /v1/responses natively; the opt-in must not swap its config."""
|
||||
config = _resolve_responses_api_provider_config("gpt-4.1", "openai", OPT_IN)
|
||||
config = _resolve_responses_api_provider_config("gpt-4.1", "openai", OPT_IN, None)
|
||||
assert type(config) is OpenAIResponsesAPIConfig
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import SCIMConfig from "@/components/SCIM";
|
|||
import LoggingSettings from "@/components/Settings/AdminSettings/LoggingSettings/LoggingSettings";
|
||||
import SSOSettings from "@/components/Settings/AdminSettings/SSOSettings/SSOSettings";
|
||||
import UISettings from "@/components/Settings/AdminSettings/UISettings/UISettings";
|
||||
import TeamAdminEditableFieldsSettings from "@/components/Settings/AdminSettings/UISettings/TeamAdminEditableFieldsSettings";
|
||||
import UserBannerSettings from "@/components/Settings/AdminSettings/UserBannerSettings/UserBannerSettings";
|
||||
import CyberArk from "@/components/Settings/AdminSettings/CyberArk/CyberArk";
|
||||
import HashicorpVault from "@/components/Settings/AdminSettings/HashicorpVault/HashicorpVault";
|
||||
|
|
@ -382,6 +383,7 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
|||
children: (
|
||||
<div className="flex flex-col gap-4">
|
||||
<UISettings />
|
||||
<TeamAdminEditableFieldsSettings />
|
||||
<UserBannerSettings />
|
||||
</div>
|
||||
),
|
||||
|
|
|
|||
|
|
@ -1,82 +0,0 @@
|
|||
/* @vitest-environment jsdom */
|
||||
import { renderHook } from "@testing-library/react";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
|
||||
const { mockPush, navState } = vi.hoisted(() => ({
|
||||
mockPush: vi.fn(),
|
||||
navState: { pathname: "/logs" },
|
||||
}));
|
||||
vi.mock("next/navigation", () => ({
|
||||
usePathname: () => navState.pathname,
|
||||
useRouter: () => ({ push: mockPush }),
|
||||
}));
|
||||
|
||||
vi.mock("@/components/networking", () => ({ serverRootPath: "" }));
|
||||
|
||||
import { createTabRoutes } from "@/utils/tabRoutes";
|
||||
import { useTabRouting } from "./useTabRouting";
|
||||
|
||||
const routes = createTabRoutes("logs", ["audit", "deleted-keys", "deleted-teams"] as const);
|
||||
|
||||
const render = (ready = true) => {
|
||||
const config = {
|
||||
routes,
|
||||
baseTabKey: "request-logs",
|
||||
visibleKeys: ["audit", "deleted-keys", "deleted-teams"],
|
||||
ready,
|
||||
};
|
||||
return renderHook(() => useTabRouting(config));
|
||||
};
|
||||
|
||||
describe("useTabRouting", () => {
|
||||
beforeEach(() => {
|
||||
navState.pathname = "/logs";
|
||||
mockPush.mockClear();
|
||||
});
|
||||
|
||||
it("maps the base path to the base tab key", () => {
|
||||
const { result } = render();
|
||||
expect(result.current.activeSlug).toBe("");
|
||||
expect(result.current.activeKey).toBe("request-logs");
|
||||
});
|
||||
|
||||
it("uses the slug itself as the active key for a known nested tab", () => {
|
||||
navState.pathname = "/ui/logs/audit";
|
||||
const { result } = render();
|
||||
expect(result.current.activeKey).toBe("audit");
|
||||
});
|
||||
|
||||
it("falls back to the base tab key for an unknown slug", () => {
|
||||
navState.pathname = "/ui/logs/bogus";
|
||||
const { result } = render();
|
||||
expect(result.current.activeKey).toBe("request-logs");
|
||||
});
|
||||
|
||||
it("redirects an unknown slug to the base href once ready", () => {
|
||||
const replaceMock = vi.fn();
|
||||
const originalLocation = window.location;
|
||||
Object.defineProperty(window, "location", { configurable: true, value: { replace: replaceMock } });
|
||||
navState.pathname = "/ui/logs/bogus";
|
||||
render(true);
|
||||
expect(replaceMock).toHaveBeenCalledWith("/ui/logs/");
|
||||
Object.defineProperty(window, "location", { configurable: true, value: originalLocation });
|
||||
});
|
||||
|
||||
it("does not redirect while not ready (role/creds still loading)", () => {
|
||||
const replaceMock = vi.fn();
|
||||
const originalLocation = window.location;
|
||||
Object.defineProperty(window, "location", { configurable: true, value: { replace: replaceMock } });
|
||||
navState.pathname = "/ui/logs/bogus";
|
||||
render(false);
|
||||
expect(replaceMock).not.toHaveBeenCalled();
|
||||
Object.defineProperty(window, "location", { configurable: true, value: originalLocation });
|
||||
});
|
||||
|
||||
it("pushes the tab href on change, mapping the base key back to the empty slug", () => {
|
||||
const { result } = render();
|
||||
result.current.onTabChange("audit");
|
||||
expect(mockPush).toHaveBeenCalledWith("/ui/logs/audit/");
|
||||
result.current.onTabChange("request-logs");
|
||||
expect(mockPush).toHaveBeenCalledWith("/ui/logs/");
|
||||
});
|
||||
});
|
||||
|
|
@ -1,38 +0,0 @@
|
|||
import { useEffect } from "react";
|
||||
import { usePathname, useRouter } from "next/navigation";
|
||||
import type { TabRoutes } from "@/utils/tabRoutes";
|
||||
|
||||
interface UseTabRoutingArgs {
|
||||
routes: Pick<TabRoutes<string>, "tabHref" | "slugFromPathname">;
|
||||
baseTabKey: string;
|
||||
visibleKeys: readonly string[];
|
||||
ready?: boolean;
|
||||
}
|
||||
|
||||
interface TabRoutingState {
|
||||
activeSlug: string;
|
||||
activeKey: string;
|
||||
onTabChange: (key: string) => void;
|
||||
}
|
||||
|
||||
export function useTabRouting({ routes, baseTabKey, visibleKeys, ready = true }: UseTabRoutingArgs): TabRoutingState {
|
||||
const { tabHref, slugFromPathname } = routes;
|
||||
const pathname = usePathname();
|
||||
const router = useRouter();
|
||||
|
||||
const activeSlug = slugFromPathname(pathname);
|
||||
const isKnownSlug = activeSlug === "" || visibleKeys.includes(activeSlug);
|
||||
const activeKey = isKnownSlug ? activeSlug || baseTabKey : baseTabKey;
|
||||
|
||||
useEffect(() => {
|
||||
if (ready && activeSlug !== "" && !isKnownSlug) {
|
||||
window.location.replace(tabHref(""));
|
||||
}
|
||||
}, [ready, activeSlug, isKnownSlug, tabHref]);
|
||||
|
||||
const onTabChange = (key: string) => {
|
||||
router.push(tabHref(key === baseTabKey ? "" : key));
|
||||
};
|
||||
|
||||
return { activeSlug, activeKey, onTabChange };
|
||||
}
|
||||
|
|
@ -1,5 +1,8 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import { screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import type { OnUrlUpdateFunction } from "nuqs/adapters/testing";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders } from "../../../../tests/test-utils";
|
||||
import PlaygroundPage from "./page";
|
||||
|
||||
const authState = { userRole: "Admin" };
|
||||
|
|
@ -35,14 +38,17 @@ vi.mock("@/app/(dashboard)/playground/components/chat_ui/AgentBuilderView", () =
|
|||
default: () => <div data-testid="agent-builder" />,
|
||||
}));
|
||||
|
||||
describe("PlaygroundPage role guard", () => {
|
||||
beforeEach(() => {
|
||||
authState.userRole = "Admin";
|
||||
});
|
||||
const lastUrlUpdate = (onUrlUpdate: ReturnType<typeof vi.fn<OnUrlUpdateFunction>>) =>
|
||||
onUrlUpdate.mock.calls.at(-1)?.[0];
|
||||
|
||||
beforeEach(() => {
|
||||
authState.userRole = "Admin";
|
||||
});
|
||||
|
||||
describe("PlaygroundPage role guard", () => {
|
||||
it.each(["Internal Viewer", "Admin Viewer"])("blocks the entire playground for %s", (role) => {
|
||||
authState.userRole = role;
|
||||
render(<PlaygroundPage />);
|
||||
renderWithProviders(<PlaygroundPage />);
|
||||
|
||||
expect(screen.getByText("Access Denied")).toBeInTheDocument();
|
||||
expect(screen.queryByRole("tab")).not.toBeInTheDocument();
|
||||
|
|
@ -54,10 +60,43 @@ describe("PlaygroundPage role guard", () => {
|
|||
|
||||
it.each(["Admin", "Internal User", "Org Admin"])("renders the playground for %s", (role) => {
|
||||
authState.userRole = role;
|
||||
render(<PlaygroundPage />);
|
||||
renderWithProviders(<PlaygroundPage />);
|
||||
|
||||
expect(screen.queryByText("Access Denied")).not.toBeInTheDocument();
|
||||
expect(screen.getByRole("tab", { name: "Chat" })).toBeInTheDocument();
|
||||
expect(screen.getByTestId("chat-ui")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("PlaygroundPage ?tab= deep link", () => {
|
||||
it("opens on Chat when the URL has no tab", () => {
|
||||
renderWithProviders(<PlaygroundPage />);
|
||||
|
||||
expect(screen.getByRole("tab", { name: "Chat" })).toHaveAttribute("aria-selected", "true");
|
||||
});
|
||||
|
||||
it("activates the tab named in ?tab=", () => {
|
||||
renderWithProviders(<PlaygroundPage />, { searchParams: { tab: "compare" } });
|
||||
|
||||
expect(screen.getByRole("tab", { name: "Compare" })).toHaveAttribute("aria-selected", "true");
|
||||
expect(screen.getByRole("tab", { name: "Chat" })).toHaveAttribute("aria-selected", "false");
|
||||
});
|
||||
|
||||
it("falls back to Chat when ?tab= is not a playground tab", () => {
|
||||
renderWithProviders(<PlaygroundPage />, { searchParams: { tab: "settings" } });
|
||||
|
||||
expect(screen.getByRole("tab", { name: "Chat" })).toHaveAttribute("aria-selected", "true");
|
||||
});
|
||||
|
||||
it("clicking a tab writes ?tab= with history replace", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onUrlUpdate = vi.fn<OnUrlUpdateFunction>();
|
||||
renderWithProviders(<PlaygroundPage />, { onUrlUpdate });
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: "Compliance" }));
|
||||
|
||||
expect(await screen.findByRole("tab", { name: "Compliance", selected: true })).toBeInTheDocument();
|
||||
await waitFor(() => expect(lastUrlUpdate(onUrlUpdate)?.searchParams.get("tab")).toBe("compliance"));
|
||||
expect(lastUrlUpdate(onUrlUpdate)?.options.history).toBe("replace");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -9,6 +9,9 @@ import { DeprecationBanner } from "@/components/DeprecationBanner";
|
|||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { fetchProxySettings } from "@/utils/proxyUtils";
|
||||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
import { useUrlTab } from "@/hooks/useUrlTab";
|
||||
|
||||
const PLAYGROUND_TABS = ["chat", "compare", "compliance", "agent-builder"] as const;
|
||||
|
||||
interface ProxySettings {
|
||||
PROXY_BASE_URL?: string;
|
||||
|
|
@ -18,6 +21,7 @@ interface ProxySettings {
|
|||
export default function PlaygroundPage() {
|
||||
const { accessToken, userRole, userId, disabledPersonalKeyCreation, token, isViewOnly } = useAuthorized();
|
||||
const [proxySettings, setProxySettings] = useState<ProxySettings | undefined>(undefined);
|
||||
const [activeTab, setActiveTab] = useUrlTab(PLAYGROUND_TABS, "chat");
|
||||
|
||||
useEffect(() => {
|
||||
const initializeProxySettings = async () => {
|
||||
|
|
@ -48,7 +52,11 @@ export default function PlaygroundPage() {
|
|||
|
||||
return (
|
||||
<div className="flex h-full min-h-0 w-full min-w-0 flex-col overflow-hidden">
|
||||
<Tabs defaultValue="chat" className="flex min-h-0 min-w-0 flex-1 flex-col gap-0 overflow-hidden">
|
||||
<Tabs
|
||||
value={activeTab}
|
||||
onValueChange={setActiveTab}
|
||||
className="flex min-h-0 min-w-0 flex-1 flex-col gap-0 overflow-hidden"
|
||||
>
|
||||
<TabsList variant="line" className="w-full shrink-0 justify-start overflow-x-auto pb-1">
|
||||
<TabsTrigger value="chat" className="flex-none">
|
||||
Chat
|
||||
|
|
|
|||
|
|
@ -0,0 +1,175 @@
|
|||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
|
||||
import { fireEvent, renderWithProviders, screen, waitFor } from "@/../tests/test-utils";
|
||||
import { toast } from "@/lib/toast";
|
||||
|
||||
import TeamAdminEditableFieldsSettings from "./TeamAdminEditableFieldsSettings";
|
||||
|
||||
const mockUseUISettings = vi.hoisted(() => vi.fn());
|
||||
const mockUseUpdateUISettings = vi.hoisted(() => vi.fn());
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
||||
default: () => ({ accessToken: "test-token" }),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/uiSettings/useUISettings", () => ({
|
||||
useUISettings: mockUseUISettings,
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/uiSettings/useUpdateUISettings", () => ({
|
||||
useUpdateUISettings: mockUseUpdateUISettings,
|
||||
}));
|
||||
|
||||
const TPM_LABEL = "Tokens per minute Limit (TPM)";
|
||||
|
||||
const mockSettings = (supported: readonly string[], enabled: readonly string[]) =>
|
||||
mockUseUISettings.mockReturnValue({
|
||||
isLoading: false,
|
||||
data: {
|
||||
field_schema: {
|
||||
properties: {
|
||||
team_admin_editable_team_fields: {
|
||||
description: "Fields a team admin may change",
|
||||
items: { type: "string", enum: supported },
|
||||
},
|
||||
},
|
||||
},
|
||||
values: { team_admin_editable_team_fields: enabled },
|
||||
},
|
||||
});
|
||||
|
||||
const mockSave = ({
|
||||
isPending = false,
|
||||
outcome = "success",
|
||||
}: {
|
||||
isPending?: boolean;
|
||||
outcome?: "success" | "error";
|
||||
}) => {
|
||||
const mutate = vi.fn((_settings: unknown, options: { onSuccess: () => void; onError: (error: Error) => void }) =>
|
||||
outcome === "success" ? options.onSuccess() : options.onError(new Error("save failed")),
|
||||
);
|
||||
mockUseUpdateUISettings.mockReturnValue({ mutate, isPending });
|
||||
return mutate;
|
||||
};
|
||||
|
||||
const saveButton = () => screen.getByRole("button", { name: "Save" });
|
||||
|
||||
describe("TeamAdminEditableFieldsSettings", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("explains that nothing can be enabled when the proxy supports no fields", () => {
|
||||
mockSettings([], []);
|
||||
mockSave({});
|
||||
|
||||
renderWithProviders(<TeamAdminEditableFieldsSettings />);
|
||||
|
||||
expect(screen.getByText("Team admins cannot edit team settings")).toBeInTheDocument();
|
||||
expect(screen.getByText(/does not support enabling any team settings fields/)).toBeInTheDocument();
|
||||
expect(screen.queryByRole("checkbox")).not.toBeInTheDocument();
|
||||
expect(screen.queryByRole("button", { name: "Save" })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders one checkbox per supported field, checked for the saved ones, with Save disabled until something changes", () => {
|
||||
mockSettings(["max_budget", "tpm_limit"], ["tpm_limit"]);
|
||||
mockSave({});
|
||||
|
||||
renderWithProviders(<TeamAdminEditableFieldsSettings />);
|
||||
|
||||
expect(screen.getByText("Team admin editable fields")).toBeInTheDocument();
|
||||
expect(screen.getByText("1 field enabled")).toBeInTheDocument();
|
||||
expect(screen.getByText("Fields a team admin may change")).toBeInTheDocument();
|
||||
expect(screen.getByRole("checkbox", { name: "max_budget" })).not.toBeChecked();
|
||||
expect(screen.getByRole("checkbox", { name: TPM_LABEL })).toBeChecked();
|
||||
expect(saveButton()).toBeDisabled();
|
||||
});
|
||||
|
||||
it("only saves a ticked field once Save is clicked", async () => {
|
||||
mockSettings(["max_budget", "tpm_limit"], ["tpm_limit"]);
|
||||
const mutate = mockSave({});
|
||||
|
||||
renderWithProviders(<TeamAdminEditableFieldsSettings />);
|
||||
fireEvent.click(screen.getByRole("checkbox", { name: "max_budget" }));
|
||||
|
||||
expect(screen.getByRole("checkbox", { name: "max_budget" })).toBeChecked();
|
||||
expect(mutate).not.toHaveBeenCalled();
|
||||
|
||||
fireEvent.click(saveButton());
|
||||
|
||||
await waitFor(() => expect(toast.success).toHaveBeenCalledWith("Team admin editable fields updated successfully"));
|
||||
expect(mutate).toHaveBeenCalledWith(
|
||||
{ team_admin_editable_team_fields: ["max_budget", "tpm_limit"] },
|
||||
expect.anything(),
|
||||
);
|
||||
expect(saveButton()).toBeDisabled();
|
||||
});
|
||||
|
||||
it("saves the list without an unticked field", async () => {
|
||||
mockSettings(["max_budget", "tpm_limit"], ["max_budget", "tpm_limit"]);
|
||||
const mutate = mockSave({});
|
||||
|
||||
renderWithProviders(<TeamAdminEditableFieldsSettings />);
|
||||
fireEvent.click(screen.getByRole("checkbox", { name: TPM_LABEL }));
|
||||
fireEvent.click(saveButton());
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalledTimes(1));
|
||||
expect(mutate).toHaveBeenCalledWith({ team_admin_editable_team_fields: ["max_budget"] }, expect.anything());
|
||||
});
|
||||
|
||||
it("disables Save again when the draft is ticked back to the saved list", () => {
|
||||
mockSettings(["tpm_limit"], []);
|
||||
mockSave({});
|
||||
|
||||
renderWithProviders(<TeamAdminEditableFieldsSettings />);
|
||||
fireEvent.click(screen.getByRole("checkbox", { name: TPM_LABEL }));
|
||||
|
||||
expect(saveButton()).toBeEnabled();
|
||||
|
||||
fireEvent.click(screen.getByRole("checkbox", { name: TPM_LABEL }));
|
||||
|
||||
expect(screen.getByRole("checkbox", { name: TPM_LABEL })).not.toBeChecked();
|
||||
expect(saveButton()).toBeDisabled();
|
||||
});
|
||||
|
||||
it("treats a saved list in another order, or with fields this proxy dropped, as the same selection", () => {
|
||||
mockSettings(["max_budget", "tpm_limit"], ["tpm_limit", "retired_field", "max_budget"]);
|
||||
mockSave({});
|
||||
|
||||
renderWithProviders(<TeamAdminEditableFieldsSettings />);
|
||||
|
||||
expect(screen.getByText("2 fields enabled")).toBeInTheDocument();
|
||||
|
||||
fireEvent.click(screen.getByRole("checkbox", { name: TPM_LABEL }));
|
||||
fireEvent.click(screen.getByRole("checkbox", { name: TPM_LABEL }));
|
||||
|
||||
expect(saveButton()).toBeDisabled();
|
||||
});
|
||||
|
||||
it("keeps the draft and shows the error when the save fails", async () => {
|
||||
mockSettings(["tpm_limit"], []);
|
||||
const mutate = mockSave({ outcome: "error" });
|
||||
|
||||
renderWithProviders(<TeamAdminEditableFieldsSettings />);
|
||||
fireEvent.click(screen.getByRole("checkbox", { name: TPM_LABEL }));
|
||||
fireEvent.click(saveButton());
|
||||
|
||||
await waitFor(() => expect(toast.fromError).toHaveBeenCalledTimes(1));
|
||||
expect(mutate).toHaveBeenCalledTimes(1);
|
||||
expect(toast.success).not.toHaveBeenCalled();
|
||||
expect(screen.getByRole("checkbox", { name: TPM_LABEL })).toBeChecked();
|
||||
expect(saveButton()).toBeEnabled();
|
||||
});
|
||||
|
||||
it("blocks ticking and saving while a save is in flight", () => {
|
||||
mockSettings(["tpm_limit"], []);
|
||||
const mutate = mockSave({ isPending: true });
|
||||
|
||||
renderWithProviders(<TeamAdminEditableFieldsSettings />);
|
||||
fireEvent.click(screen.getByRole("checkbox", { name: TPM_LABEL }));
|
||||
|
||||
expect(screen.getByRole("checkbox", { name: TPM_LABEL })).not.toBeChecked();
|
||||
expect(screen.getByRole("button", { name: "Saving..." })).toBeDisabled();
|
||||
expect(mutate).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,138 @@
|
|||
"use client";
|
||||
|
||||
import { Controller } from "react-hook-form";
|
||||
import { z } from "zod/v4";
|
||||
|
||||
import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings";
|
||||
import { useUpdateUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUpdateUISettings";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import {
|
||||
parseSupportedTeamAdminEditableFields,
|
||||
parseTeamAdminEditableFields,
|
||||
teamAdminFieldLabel,
|
||||
} from "@/components/team/teamAdminEditAccess";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card";
|
||||
import { Checkbox } from "@/components/ui/checkbox";
|
||||
import { Skeleton } from "@/components/ui/skeleton";
|
||||
import { useZodForm } from "@/lib/forms/useZodForm";
|
||||
import { toast } from "@/lib/toast";
|
||||
|
||||
const editableFieldsSchema = z.object({ team_admin_editable_team_fields: z.array(z.string()) });
|
||||
|
||||
type SaveEditableFields = ReturnType<typeof useUpdateUISettings>["mutate"];
|
||||
|
||||
export default function TeamAdminEditableFieldsSettings() {
|
||||
const { accessToken } = useAuthorized();
|
||||
const { data, isLoading } = useUISettings();
|
||||
const { mutate: saveSettings, isPending } = useUpdateUISettings(accessToken);
|
||||
const supportedFields = parseSupportedTeamAdminEditableFields(data?.field_schema);
|
||||
const savedFields = parseTeamAdminEditableFields(data?.values);
|
||||
const enabledFields = supportedFields.filter((field) => savedFields.includes(field));
|
||||
|
||||
return (
|
||||
<Card>
|
||||
<CardHeader>
|
||||
<div className="flex items-center gap-2">
|
||||
<CardTitle>Team admin editable fields</CardTitle>
|
||||
<Badge variant={enabledFields.length > 0 ? "secondary" : "outline"}>
|
||||
{enabledFields.length > 0
|
||||
? `${enabledFields.length} field${enabledFields.length !== 1 ? "s" : ""} enabled`
|
||||
: "Team admins cannot edit team settings"}
|
||||
</Badge>
|
||||
</div>
|
||||
<CardDescription>
|
||||
{data?.field_schema?.properties?.team_admin_editable_team_fields?.description ??
|
||||
"Team settings fields a team admin may change on the teams they administer."}
|
||||
</CardDescription>
|
||||
</CardHeader>
|
||||
<CardContent>
|
||||
{isLoading ? (
|
||||
<Skeleton className="h-16 w-full" />
|
||||
) : (
|
||||
<TeamAdminEditableFieldsForm
|
||||
key={enabledFields.join(",")}
|
||||
enabledFields={enabledFields}
|
||||
supportedFields={supportedFields}
|
||||
isPending={isPending}
|
||||
saveSettings={saveSettings}
|
||||
/>
|
||||
)}
|
||||
</CardContent>
|
||||
</Card>
|
||||
);
|
||||
}
|
||||
|
||||
interface TeamAdminEditableFieldsFormProps {
|
||||
enabledFields: readonly string[];
|
||||
supportedFields: readonly string[];
|
||||
isPending: boolean;
|
||||
saveSettings: SaveEditableFields;
|
||||
}
|
||||
|
||||
function TeamAdminEditableFieldsForm({
|
||||
enabledFields,
|
||||
supportedFields,
|
||||
isPending,
|
||||
saveSettings,
|
||||
}: TeamAdminEditableFieldsFormProps) {
|
||||
const form = useZodForm(editableFieldsSchema, {
|
||||
defaultValues: { team_admin_editable_team_fields: [...enabledFields] },
|
||||
});
|
||||
const submit = form.handleSubmit((values) =>
|
||||
saveSettings(values, {
|
||||
onSuccess: () => {
|
||||
form.reset(values);
|
||||
toast.success("Team admin editable fields updated successfully");
|
||||
},
|
||||
onError: (error) => {
|
||||
toast.fromError(error);
|
||||
},
|
||||
}),
|
||||
);
|
||||
|
||||
if (supportedFields.length === 0) {
|
||||
return (
|
||||
<p className="text-sm italic text-muted-foreground">
|
||||
This proxy version does not support enabling any team settings fields for team admins yet.
|
||||
</p>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<form onSubmit={(event) => void submit(event)} className="space-y-4">
|
||||
<Controller
|
||||
control={form.control}
|
||||
name="team_admin_editable_team_fields"
|
||||
render={({ field }) => (
|
||||
<div className="space-y-2">
|
||||
{supportedFields.map((name) => {
|
||||
const checkboxId = `team-admin-editable-${name}`;
|
||||
return (
|
||||
<label key={name} htmlFor={checkboxId} className="flex cursor-pointer items-center gap-2">
|
||||
<Checkbox
|
||||
id={checkboxId}
|
||||
checked={field.value.includes(name)}
|
||||
disabled={isPending}
|
||||
onCheckedChange={(checked) =>
|
||||
field.onChange(
|
||||
supportedFields.filter((item) => (item === name ? checked : field.value.includes(item))),
|
||||
)
|
||||
}
|
||||
/>
|
||||
<span className="text-sm text-foreground">{teamAdminFieldLabel(name)}</span>
|
||||
</label>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
/>
|
||||
<div className="flex justify-end">
|
||||
<Button type="submit" disabled={isPending || !form.formState.isDirty}>
|
||||
{isPending ? "Saving..." : "Save"}
|
||||
</Button>
|
||||
</div>
|
||||
</form>
|
||||
);
|
||||
}
|
||||
|
|
@ -145,7 +145,7 @@ function UserBannerSettingsForm({ persisted, isLoading, isPending, saveBanner }:
|
|||
</div>
|
||||
)}
|
||||
|
||||
<div>
|
||||
<div className="flex justify-end">
|
||||
<Button onClick={handleSave} disabled={isPending || messageMissing}>
|
||||
{isPending ? "Saving..." : "Save banner"}
|
||||
</Button>
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue