mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
Merge pull request #42317 from BerriAI/litellm_redis_cluster_topology
feat(rust-cache): serve RedisClusterCache natively as a Redis topology
This commit is contained in:
commit
4b9aa222a6
12 changed files with 1436 additions and 225 deletions
6
.github/scripts/verify_linux_native_wheel.py
vendored
6
.github/scripts/verify_linux_native_wheel.py
vendored
|
|
@ -205,7 +205,7 @@ def main(
|
|||
native_module: Final = load_native_module(native_path)
|
||||
native_module_loads: Final = native_module is not None
|
||||
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
|
||||
native_size_limit: Final = 25_000_000
|
||||
native_size_limit: Final = 30_000_000
|
||||
native_size_within_limit: Final = native_member.file_size <= native_size_limit
|
||||
validations: Final = (
|
||||
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),
|
||||
|
|
@ -222,7 +222,7 @@ def main(
|
|||
("Python extension entry point is present", extension_entry_point_present),
|
||||
("Native module loads", native_module_loads),
|
||||
("Production module omits the panic test hook", panic_test_hook_absent),
|
||||
("Native extension does not exceed 25 MB", native_size_within_limit),
|
||||
("Native extension does not exceed 30 MB", native_size_within_limit),
|
||||
("Wheel contents are valid", not unexpected_members),
|
||||
)
|
||||
|
||||
|
|
@ -267,7 +267,7 @@ def main(
|
|||
),
|
||||
(
|
||||
not native_size_within_limit,
|
||||
f"native extension exceeds 20 MB: {native_member.file_size / 1_000_000:.2f} MB",
|
||||
f"native extension exceeds 30 MB: {native_member.file_size / 1_000_000:.2f} MB",
|
||||
),
|
||||
(bool(unexpected_members), f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}"),
|
||||
)
|
||||
|
|
|
|||
8
litellm-rust/Cargo.lock
generated
8
litellm-rust/Cargo.lock
generated
|
|
@ -927,6 +927,12 @@ dependencies = [
|
|||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "crc16"
|
||||
version = "0.4.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "338089f42c427b86394a5ee60ff321da23a5c89c9d89514c829687b26359fcff"
|
||||
|
||||
[[package]]
|
||||
name = "crc32fast"
|
||||
version = "1.5.1"
|
||||
|
|
@ -3709,9 +3715,11 @@ checksum = "2acbc41a996f7652b2ddd9dfd98cc4ff602cfd742ae35382f07f608405ab50ed"
|
|||
dependencies = [
|
||||
"arcstr",
|
||||
"combine",
|
||||
"crc16",
|
||||
"itoa",
|
||||
"num-bigint 0.5.1",
|
||||
"percent-encoding",
|
||||
"rand 0.10.2",
|
||||
"rustls 0.23.42",
|
||||
"rustls-native-certs",
|
||||
"ryu",
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ repository.workspace = true
|
|||
|
||||
[dependencies]
|
||||
litellm-cache.workspace = true
|
||||
redis = { version = "1.7.0", features = ["tls-rustls"] }
|
||||
redis = { version = "1.7.0", features = ["cluster", "tls-rustls"] }
|
||||
r2d2 = "0.8.10"
|
||||
tokio.workspace = true
|
||||
|
||||
|
|
|
|||
|
|
@ -9,8 +9,14 @@ use litellm_cache::{
|
|||
};
|
||||
use redis::Commands;
|
||||
|
||||
use crate::topology::RedisTopology;
|
||||
|
||||
mod connection;
|
||||
mod operations;
|
||||
|
||||
pub(crate) use connection::ConnectionRef;
|
||||
use connection::{ClusterConnectionManager, ConnectionManager};
|
||||
|
||||
pub use operations::{
|
||||
RedisArg, RedisLpopOperation, RedisLpopResult, RedisRpushOperation, RedisScript,
|
||||
};
|
||||
|
|
@ -19,40 +25,6 @@ const DEFAULT_TTL: Duration = Duration::from_secs(600);
|
|||
const REDIS_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
const REDIS_POOL_SIZE: u32 = 16;
|
||||
|
||||
struct PooledConnection {
|
||||
connection: redis::Connection,
|
||||
failed: bool,
|
||||
}
|
||||
|
||||
/// Pools connections without a checkout PING, which would double every operation's round trips.
|
||||
/// A timed-out command leaves its reply on the socket while redis still reports the connection
|
||||
/// open, so any connection whose operation failed is discarded instead of being reused.
|
||||
struct ConnectionManager(redis::Client);
|
||||
|
||||
impl r2d2::ManageConnection for ConnectionManager {
|
||||
type Connection = PooledConnection;
|
||||
type Error = redis::RedisError;
|
||||
|
||||
fn connect(&self) -> Result<PooledConnection, redis::RedisError> {
|
||||
let connection = self.0.get_connection()?;
|
||||
connection.set_read_timeout(Some(REDIS_TIMEOUT))?;
|
||||
connection.set_write_timeout(Some(REDIS_TIMEOUT))?;
|
||||
Ok(PooledConnection {
|
||||
connection,
|
||||
failed: false,
|
||||
})
|
||||
}
|
||||
|
||||
fn is_valid(&self, connection: &mut PooledConnection) -> Result<(), redis::RedisError> {
|
||||
redis::cmd("PING").query::<String>(&mut connection.connection)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn has_broken(&self, connection: &mut PooledConnection) -> bool {
|
||||
connection.failed || !redis::ConnectionLike::is_open(&connection.connection)
|
||||
}
|
||||
}
|
||||
|
||||
const INCREMENT_SCRIPT: &str = concat!(
|
||||
"local value = redis.call('INCRBYFLOAT', KEYS[1], ARGV[1]); ",
|
||||
"if redis.call('TTL', KEYS[1]) == -1 then ",
|
||||
|
|
@ -70,42 +42,10 @@ const CLAIM_ATTEMPTS: usize = 8;
|
|||
|
||||
enum Connections<C> {
|
||||
Pool(r2d2::Pool<ConnectionManager>),
|
||||
Cluster(r2d2::Pool<ClusterConnectionManager>),
|
||||
Fixed(Mutex<C>),
|
||||
}
|
||||
|
||||
struct ConnectionRef<'a>(&'a mut dyn redis::ConnectionLike);
|
||||
|
||||
impl redis::ConnectionLike for ConnectionRef<'_> {
|
||||
fn req_packed_command(&mut self, cmd: &[u8]) -> redis::RedisResult<redis::Value> {
|
||||
self.0.req_packed_command(cmd)
|
||||
}
|
||||
|
||||
fn req_packed_commands(
|
||||
&mut self,
|
||||
cmd: &[u8],
|
||||
offset: usize,
|
||||
count: usize,
|
||||
) -> redis::RedisResult<Vec<redis::Value>> {
|
||||
self.0.req_packed_commands(cmd, offset, count)
|
||||
}
|
||||
|
||||
fn get_db(&self) -> i64 {
|
||||
self.0.get_db()
|
||||
}
|
||||
|
||||
fn supports_pipelining(&self) -> bool {
|
||||
self.0.supports_pipelining()
|
||||
}
|
||||
|
||||
fn check_connection(&mut self) -> bool {
|
||||
self.0.check_connection()
|
||||
}
|
||||
|
||||
fn is_open(&self) -> bool {
|
||||
self.0.is_open()
|
||||
}
|
||||
}
|
||||
|
||||
impl<C> Connections<C>
|
||||
where
|
||||
C: redis::ConnectionLike + Send + 'static,
|
||||
|
|
@ -117,13 +57,19 @@ where
|
|||
match self {
|
||||
Self::Pool(pool) => {
|
||||
let mut pooled = pool.get().map_err(|_| Error::Unavailable)?;
|
||||
let result = operation(&mut ConnectionRef(&mut pooled.connection));
|
||||
let result = operation(&mut ConnectionRef::Node(&mut pooled.connection));
|
||||
pooled.failed = matches!(result, Err(Error::Unavailable));
|
||||
result
|
||||
}
|
||||
Self::Cluster(pool) => {
|
||||
let mut pooled = pool.get().map_err(|_| Error::Unavailable)?;
|
||||
let result = operation(&mut ConnectionRef::Cluster(&mut pooled.connection));
|
||||
pooled.failed = matches!(result, Err(Error::Unavailable));
|
||||
result
|
||||
}
|
||||
Self::Fixed(connection) => {
|
||||
let mut connection = connection.lock().map_err(|_| Error::Unavailable)?;
|
||||
operation(&mut ConnectionRef(&mut *connection))
|
||||
operation(&mut ConnectionRef::Node(&mut *connection))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -134,27 +80,46 @@ pub struct RedisCache<S, C = redis::Connection> {
|
|||
default_ttl: Duration,
|
||||
codec: S,
|
||||
namespace: Option<String>,
|
||||
topology: RedisTopology,
|
||||
}
|
||||
|
||||
impl<S: CacheCodec> RedisCache<S> {
|
||||
pub fn new(url: &str, default_ttl: Option<Duration>, codec: S) -> Result<Self, Error> {
|
||||
let client = redis::Client::open(url).map_err(|_| Error::Unavailable)?;
|
||||
let pool = r2d2::Pool::builder()
|
||||
.max_size(REDIS_POOL_SIZE)
|
||||
.min_idle(Some(0))
|
||||
.connection_timeout(REDIS_TIMEOUT)
|
||||
.test_on_check_out(false)
|
||||
.build(ConnectionManager(client))
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
Self::connect(url, &RedisTopology::Standalone, default_ttl, codec)
|
||||
}
|
||||
|
||||
pub fn connect(
|
||||
url: &str,
|
||||
topology: &RedisTopology,
|
||||
default_ttl: Option<Duration>,
|
||||
codec: S,
|
||||
) -> Result<Self, Error> {
|
||||
let connections = match topology {
|
||||
RedisTopology::Standalone => Connections::Pool(pool(ConnectionManager::open(url)?)?),
|
||||
RedisTopology::Cluster { startup_nodes } => {
|
||||
Connections::Cluster(pool(ClusterConnectionManager::open(url, startup_nodes)?)?)
|
||||
}
|
||||
};
|
||||
Ok(Self {
|
||||
connections: Arc::new(Connections::Pool(pool)),
|
||||
connections: Arc::new(connections),
|
||||
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
|
||||
codec,
|
||||
namespace: None,
|
||||
topology: topology.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn pool<M: r2d2::ManageConnection>(manager: M) -> Result<r2d2::Pool<M>, Error> {
|
||||
r2d2::Pool::builder()
|
||||
.max_size(REDIS_POOL_SIZE)
|
||||
.min_idle(Some(0))
|
||||
.connection_timeout(REDIS_TIMEOUT)
|
||||
.test_on_check_out(false)
|
||||
.build(manager)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
}
|
||||
|
||||
impl<S, C> RedisCache<S, C>
|
||||
where
|
||||
S: CacheCodec,
|
||||
|
|
@ -166,6 +131,7 @@ where
|
|||
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
|
||||
codec,
|
||||
namespace: None,
|
||||
topology: RedisTopology::Standalone,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -180,6 +146,10 @@ where
|
|||
self.namespace.as_deref()
|
||||
}
|
||||
|
||||
pub fn topology(&self) -> &RedisTopology {
|
||||
&self.topology
|
||||
}
|
||||
|
||||
fn namespaced_key(&self, key: &str) -> String {
|
||||
namespaced_key(self.namespace.as_deref(), key)
|
||||
}
|
||||
|
|
@ -200,26 +170,14 @@ where
|
|||
}
|
||||
|
||||
fn flush_matching(connection: &mut ConnectionRef<'_>, pattern: &str) -> Result<(), Error> {
|
||||
let mut cursor = 0u64;
|
||||
loop {
|
||||
let (next_cursor, keys): (u64, Vec<String>) = redis::cmd("SCAN")
|
||||
.cursor_arg(cursor)
|
||||
.arg("MATCH")
|
||||
.arg(pattern)
|
||||
.arg("COUNT")
|
||||
.arg(1000)
|
||||
.query(connection)
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
connection.scan(pattern, 1000, |connection, keys| {
|
||||
if !keys.is_empty() {
|
||||
connection
|
||||
.del::<_, usize>(keys)
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
}
|
||||
if next_cursor == 0 {
|
||||
return Ok(());
|
||||
}
|
||||
cursor = next_cursor;
|
||||
}
|
||||
Ok(true)
|
||||
})
|
||||
}
|
||||
|
||||
fn decode_response(&self, value: redis::Value) -> Result<Option<S::Value>, Error> {
|
||||
|
|
@ -350,19 +308,19 @@ where
|
|||
})
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
|
||||
if entries.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
|
||||
let mut pipeline = redis::pipe();
|
||||
for (key, payload) in entries {
|
||||
pipeline
|
||||
.cmd("SETEX")
|
||||
.arg(key)
|
||||
.arg(ttl)
|
||||
.arg(payload)
|
||||
.ignore();
|
||||
}
|
||||
pipeline
|
||||
.query::<()>(connection)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
let commands = entries
|
||||
.into_iter()
|
||||
.map(|(key, payload)| {
|
||||
let mut command = redis::cmd("SETEX");
|
||||
command.arg(key).arg(ttl).arg(payload);
|
||||
command
|
||||
})
|
||||
.collect();
|
||||
connection.pipeline(commands).map(drop)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
|
@ -373,7 +331,7 @@ where
|
|||
|
||||
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
|
||||
match Self::run_blocking(Arc::clone(&self.connections), |connection| {
|
||||
Ok(match redis::cmd("PING").query::<String>(connection) {
|
||||
Ok(match connection.ping() {
|
||||
Ok(_) => CacheConnectionResult {
|
||||
status: CacheConnectionStatus::Success,
|
||||
message: "Redis cache connection test successful".into(),
|
||||
|
|
|
|||
392
litellm-rust/crates/cache-redis/src/cache/connection.rs
vendored
Normal file
392
litellm-rust/crates/cache-redis/src/cache/connection.rs
vendored
Normal file
|
|
@ -0,0 +1,392 @@
|
|||
use std::collections::HashMap;
|
||||
|
||||
use litellm_cache::Error;
|
||||
use redis::{
|
||||
ConnectionAddr, ConnectionInfo, ConnectionLike, IntoConnectionInfo,
|
||||
cluster::{ClusterClient, ClusterClientBuilder, ClusterConnection, NodeAddress},
|
||||
cluster_routing::{
|
||||
MultipleNodeRoutingInfo, ResponsePolicy, RoutingInfo, SingleNodeRoutingInfo, Slot,
|
||||
},
|
||||
};
|
||||
|
||||
use super::REDIS_TIMEOUT;
|
||||
use crate::topology::RedisNode;
|
||||
|
||||
pub(super) struct PooledConnection<C> {
|
||||
pub(super) connection: C,
|
||||
pub(super) failed: bool,
|
||||
}
|
||||
|
||||
/// Pools connections without a checkout PING, which would double every operation's round trips.
|
||||
/// A timed-out command leaves its reply on the socket while redis still reports the connection
|
||||
/// open, so any connection whose operation failed is discarded instead of being reused.
|
||||
pub(super) struct ConnectionManager(redis::Client);
|
||||
|
||||
impl ConnectionManager {
|
||||
pub(super) fn open(url: &str) -> Result<Self, Error> {
|
||||
redis::Client::open(url)
|
||||
.map(Self)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
}
|
||||
}
|
||||
|
||||
impl r2d2::ManageConnection for ConnectionManager {
|
||||
type Connection = PooledConnection<redis::Connection>;
|
||||
type Error = redis::RedisError;
|
||||
|
||||
fn connect(&self) -> Result<Self::Connection, redis::RedisError> {
|
||||
let connection = self.0.get_connection()?;
|
||||
connection.set_read_timeout(Some(REDIS_TIMEOUT))?;
|
||||
connection.set_write_timeout(Some(REDIS_TIMEOUT))?;
|
||||
Ok(PooledConnection {
|
||||
connection,
|
||||
failed: false,
|
||||
})
|
||||
}
|
||||
|
||||
fn is_valid(&self, connection: &mut Self::Connection) -> Result<(), redis::RedisError> {
|
||||
redis::cmd("PING").query::<String>(&mut connection.connection)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn has_broken(&self, connection: &mut Self::Connection) -> bool {
|
||||
connection.failed || !redis::ConnectionLike::is_open(&connection.connection)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct ClusterConnectionManager(ClusterClient);
|
||||
|
||||
impl ClusterConnectionManager {
|
||||
pub(super) fn open(url: &str, startup_nodes: &[RedisNode]) -> Result<Self, Error> {
|
||||
if startup_nodes.is_empty() {
|
||||
return Err(Error::Unavailable);
|
||||
}
|
||||
let info = url.into_connection_info().map_err(|_| Error::Unavailable)?;
|
||||
let nodes = startup_nodes
|
||||
.iter()
|
||||
.map(|node| node_info(&info, node))
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
ClusterClientBuilder::new(nodes)
|
||||
.connection_timeout(REDIS_TIMEOUT)
|
||||
.response_timeout(REDIS_TIMEOUT)
|
||||
.build()
|
||||
.map(Self)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
}
|
||||
}
|
||||
|
||||
fn node_info(info: &ConnectionInfo, node: &RedisNode) -> Result<ConnectionInfo, Error> {
|
||||
let addr = match info.addr() {
|
||||
ConnectionAddr::Tcp(..) => ConnectionAddr::Tcp(node.host.clone(), node.port),
|
||||
ConnectionAddr::TcpTls {
|
||||
insecure,
|
||||
tls_params,
|
||||
..
|
||||
} => ConnectionAddr::TcpTls {
|
||||
host: node.host.clone(),
|
||||
port: node.port,
|
||||
insecure: *insecure,
|
||||
tls_params: tls_params.clone(),
|
||||
},
|
||||
_ => return Err(Error::Unavailable),
|
||||
};
|
||||
Ok(info.clone().set_addr(addr))
|
||||
}
|
||||
|
||||
impl r2d2::ManageConnection for ClusterConnectionManager {
|
||||
type Connection = PooledConnection<ClusterConnection>;
|
||||
type Error = redis::RedisError;
|
||||
|
||||
fn connect(&self) -> Result<Self::Connection, redis::RedisError> {
|
||||
let connection = self.0.get_connection()?;
|
||||
connection.set_read_timeout(Some(REDIS_TIMEOUT))?;
|
||||
connection.set_write_timeout(Some(REDIS_TIMEOUT))?;
|
||||
Ok(PooledConnection {
|
||||
connection,
|
||||
failed: false,
|
||||
})
|
||||
}
|
||||
|
||||
fn is_valid(&self, connection: &mut Self::Connection) -> Result<(), redis::RedisError> {
|
||||
redis::cmd("PING").query::<String>(&mut connection.connection)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn has_broken(&self, connection: &mut Self::Connection) -> bool {
|
||||
connection.failed || !redis::ConnectionLike::is_open(&connection.connection)
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) enum ConnectionRef<'a> {
|
||||
Node(&'a mut dyn redis::ConnectionLike),
|
||||
Cluster(&'a mut ClusterConnection),
|
||||
}
|
||||
|
||||
impl redis::ConnectionLike for ConnectionRef<'_> {
|
||||
fn req_packed_command(&mut self, cmd: &[u8]) -> redis::RedisResult<redis::Value> {
|
||||
match self {
|
||||
Self::Node(connection) => connection.req_packed_command(cmd),
|
||||
Self::Cluster(connection) => connection.req_packed_command(cmd),
|
||||
}
|
||||
}
|
||||
|
||||
fn req_packed_commands(
|
||||
&mut self,
|
||||
cmd: &[u8],
|
||||
offset: usize,
|
||||
count: usize,
|
||||
) -> redis::RedisResult<Vec<redis::Value>> {
|
||||
match self {
|
||||
Self::Node(connection) => connection.req_packed_commands(cmd, offset, count),
|
||||
Self::Cluster(connection) => connection.req_packed_commands(cmd, offset, count),
|
||||
}
|
||||
}
|
||||
|
||||
fn get_db(&self) -> i64 {
|
||||
match self {
|
||||
Self::Node(connection) => connection.get_db(),
|
||||
Self::Cluster(connection) => redis::ConnectionLike::get_db(*connection),
|
||||
}
|
||||
}
|
||||
|
||||
fn supports_pipelining(&self) -> bool {
|
||||
match self {
|
||||
Self::Node(connection) => connection.supports_pipelining(),
|
||||
Self::Cluster(connection) => redis::ConnectionLike::supports_pipelining(*connection),
|
||||
}
|
||||
}
|
||||
|
||||
fn check_connection(&mut self) -> bool {
|
||||
match self {
|
||||
Self::Node(connection) => connection.check_connection(),
|
||||
Self::Cluster(connection) => connection.check_connection(),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_open(&self) -> bool {
|
||||
match self {
|
||||
Self::Node(connection) => connection.is_open(),
|
||||
Self::Cluster(connection) => redis::ConnectionLike::is_open(*connection),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl ConnectionRef<'_> {
|
||||
pub(crate) fn pipeline(
|
||||
&mut self,
|
||||
commands: Vec<redis::Cmd>,
|
||||
) -> Result<Vec<redis::Value>, Error> {
|
||||
match self {
|
||||
Self::Node(connection) => {
|
||||
let mut pipeline = redis::pipe();
|
||||
for command in &commands {
|
||||
pipeline.add_command(command.clone());
|
||||
}
|
||||
pipeline
|
||||
.query::<Vec<redis::Value>>(*connection)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
}
|
||||
Self::Cluster(connection) => {
|
||||
let mut replies: Vec<Option<redis::Value>> = vec![None; commands.len()];
|
||||
for indices in slot_groups(&commands).into_values() {
|
||||
let mut pipeline = redis::pipe();
|
||||
for index in &indices {
|
||||
pipeline.add_command(commands[*index].clone());
|
||||
}
|
||||
let values = connection
|
||||
.req_packed_commands(&pipeline.get_packed_pipeline(), 0, indices.len())
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
if values.len() != indices.len() {
|
||||
return Err(Error::Unavailable);
|
||||
}
|
||||
for (index, value) in indices.into_iter().zip(values) {
|
||||
replies[index] = Some(value);
|
||||
}
|
||||
}
|
||||
replies
|
||||
.into_iter()
|
||||
.collect::<Option<Vec<_>>>()
|
||||
.ok_or(Error::Unavailable)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn scan(
|
||||
&mut self,
|
||||
pattern: &str,
|
||||
count: usize,
|
||||
mut visit: impl FnMut(&mut Self, Vec<String>) -> Result<bool, Error>,
|
||||
) -> Result<(), Error> {
|
||||
let pages = match self {
|
||||
Self::Node(connection) => {
|
||||
let page = scan_command(0, pattern, count)
|
||||
.query::<ScanPage>(*connection)
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
vec![(None, page)]
|
||||
}
|
||||
Self::Cluster(connection) => connection
|
||||
.route_command(
|
||||
&scan_command(0, pattern, count),
|
||||
RoutingInfo::MultiNode((
|
||||
MultipleNodeRoutingInfo::AllMasters,
|
||||
Some(ResponsePolicy::Special),
|
||||
)),
|
||||
)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
.and_then(primary_pages)?
|
||||
.into_iter()
|
||||
.map(|(node, page)| (Some(node), page))
|
||||
.collect(),
|
||||
};
|
||||
for (node, (mut cursor, mut keys)) in pages {
|
||||
loop {
|
||||
if !visit(self, keys)? {
|
||||
return Ok(());
|
||||
}
|
||||
if cursor == 0 {
|
||||
break;
|
||||
}
|
||||
(cursor, keys) = self.scan_page(node.as_ref(), cursor, pattern, count)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn ping(&mut self) -> Result<bool, redis::RedisError> {
|
||||
let command = redis::cmd("PING");
|
||||
match self {
|
||||
Self::Node(connection) => command
|
||||
.query::<String>(*connection)
|
||||
.map(|response| response == "PONG"),
|
||||
Self::Cluster(connection) => connection
|
||||
.route_command(
|
||||
&command,
|
||||
RoutingInfo::MultiNode((
|
||||
MultipleNodeRoutingInfo::AllNodes,
|
||||
Some(ResponsePolicy::AllSucceeded),
|
||||
)),
|
||||
)
|
||||
.map(|_| true),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn node_text(&mut self, command: &redis::Cmd) -> Result<String, Error> {
|
||||
match self {
|
||||
Self::Node(connection) => command.query(*connection).map_err(|_| Error::Unavailable),
|
||||
Self::Cluster(connection) => {
|
||||
let value = connection
|
||||
.route_command(
|
||||
command,
|
||||
RoutingInfo::MultiNode((
|
||||
MultipleNodeRoutingInfo::AllNodes,
|
||||
Some(ResponsePolicy::Special),
|
||||
)),
|
||||
)
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
let redis::Value::Map(entries) = value else {
|
||||
return Err(Error::Unavailable);
|
||||
};
|
||||
let mut replies = entries
|
||||
.into_iter()
|
||||
.map(|(node, reply)| {
|
||||
Ok((
|
||||
redis::from_redis_value::<String>(node)
|
||||
.map_err(|_| Error::Unavailable)?,
|
||||
redis::from_redis_value::<String>(reply)
|
||||
.map_err(|_| Error::Unavailable)?,
|
||||
))
|
||||
})
|
||||
.collect::<Result<Vec<(String, String)>, Error>>()?;
|
||||
replies.sort();
|
||||
Ok(replies
|
||||
.into_iter()
|
||||
.map(|(_, reply)| reply)
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn flushall(&mut self) -> Result<(), Error> {
|
||||
let command = redis::cmd("FLUSHALL");
|
||||
match self {
|
||||
Self::Node(connection) => command.query(*connection).map_err(|_| Error::Unavailable),
|
||||
Self::Cluster(connection) => connection
|
||||
.route_command(
|
||||
&command,
|
||||
RoutingInfo::MultiNode((
|
||||
MultipleNodeRoutingInfo::AllMasters,
|
||||
Some(ResponsePolicy::AllSucceeded),
|
||||
)),
|
||||
)
|
||||
.map(|_| ())
|
||||
.map_err(|_| Error::Unavailable),
|
||||
}
|
||||
}
|
||||
|
||||
fn scan_page(
|
||||
&mut self,
|
||||
node: Option<&NodeAddress>,
|
||||
cursor: u64,
|
||||
pattern: &str,
|
||||
count: usize,
|
||||
) -> Result<ScanPage, Error> {
|
||||
let command = scan_command(cursor, pattern, count);
|
||||
match (self, node) {
|
||||
(Self::Node(connection), None) => {
|
||||
command.query(*connection).map_err(|_| Error::Unavailable)
|
||||
}
|
||||
(Self::Cluster(connection), Some(node)) => connection
|
||||
.route_command(
|
||||
&command,
|
||||
RoutingInfo::SingleNode(SingleNodeRoutingInfo::ByAddress {
|
||||
host: node.host().to_string(),
|
||||
port: node.port(),
|
||||
}),
|
||||
)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
.and_then(|value| redis::from_redis_value(value).map_err(|_| Error::Unavailable)),
|
||||
_ => Err(Error::Unavailable),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type ScanPage = (u64, Vec<String>);
|
||||
|
||||
fn primary_pages(value: redis::Value) -> Result<Vec<(NodeAddress, ScanPage)>, Error> {
|
||||
let redis::Value::Map(entries) = value else {
|
||||
return Err(Error::Unavailable);
|
||||
};
|
||||
entries
|
||||
.into_iter()
|
||||
.map(|(node, page)| {
|
||||
let node = redis::from_redis_value::<String>(node).map_err(|_| Error::Unavailable)?;
|
||||
let node = NodeAddress::try_from(node.as_str()).map_err(|_| Error::Unavailable)?;
|
||||
let page = redis::from_redis_value::<ScanPage>(page).map_err(|_| Error::Unavailable)?;
|
||||
Ok((node, page))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn scan_command(cursor: u64, pattern: &str, count: usize) -> redis::Cmd {
|
||||
let mut command = redis::cmd("SCAN");
|
||||
command
|
||||
.cursor_arg(cursor)
|
||||
.arg("MATCH")
|
||||
.arg(pattern)
|
||||
.arg("COUNT")
|
||||
.arg(count);
|
||||
command
|
||||
}
|
||||
|
||||
fn slot_groups(commands: &[redis::Cmd]) -> HashMap<Slot, Vec<usize>> {
|
||||
let mut groups: HashMap<Slot, Vec<usize>> = HashMap::new();
|
||||
for (index, command) in commands.iter().enumerate() {
|
||||
let key = match command.args_iter().nth(1) {
|
||||
Some(redis::Arg::Simple(key)) => key,
|
||||
_ => b"",
|
||||
};
|
||||
groups.entry(Slot::for_key(key)).or_default().push(index);
|
||||
}
|
||||
groups
|
||||
}
|
||||
|
|
@ -183,20 +183,13 @@ where
|
|||
}
|
||||
|
||||
pub fn sync_ping(&self) -> Result<bool, Error> {
|
||||
self.connections.execute(|connection| {
|
||||
redis::cmd("PING")
|
||||
.query::<String>(connection)
|
||||
.map(|response| response == "PONG")
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
self.connections
|
||||
.execute(|connection| connection.ping().map_err(|_| Error::Unavailable))
|
||||
}
|
||||
|
||||
pub async fn ping(&self) -> Result<bool, Error> {
|
||||
Self::run_blocking(Arc::clone(&self.connections), |connection| {
|
||||
redis::cmd("PING")
|
||||
.query::<String>(connection)
|
||||
.map(|response| response == "PONG")
|
||||
.map_err(|_| Error::Unavailable)
|
||||
connection.ping().map_err(|_| Error::Unavailable)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
|
@ -216,24 +209,13 @@ where
|
|||
pub async fn async_scan_iter(&self, pattern: &str, count: usize) -> Result<Vec<String>, Error> {
|
||||
let pattern = format!("{}*", self.namespaced_key(pattern));
|
||||
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
|
||||
let mut cursor = 0u64;
|
||||
let mut matches = Vec::new();
|
||||
loop {
|
||||
let (next_cursor, keys): (u64, Vec<String>) = redis::cmd("SCAN")
|
||||
.cursor_arg(cursor)
|
||||
.arg("MATCH")
|
||||
.arg(&pattern)
|
||||
.arg("COUNT")
|
||||
.arg(count)
|
||||
.query(connection)
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
connection.scan(&pattern, count, |_, keys| {
|
||||
matches.extend(keys);
|
||||
if matches.len() >= count || next_cursor == 0 {
|
||||
matches.truncate(count);
|
||||
return Ok(matches);
|
||||
}
|
||||
cursor = next_cursor;
|
||||
}
|
||||
Ok(matches.len() < count)
|
||||
})?;
|
||||
matches.truncate(count);
|
||||
Ok(matches)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
|
@ -250,13 +232,18 @@ where
|
|||
let key = self.namespaced_key(key);
|
||||
let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl));
|
||||
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
|
||||
let mut pipeline = redis::pipe();
|
||||
pipeline.cmd("SADD").arg(&key).arg(values);
|
||||
pipeline.cmd("EXPIRE").arg(&key).arg(ttl).ignore();
|
||||
pipeline
|
||||
.query::<(usize,)>(connection)
|
||||
.map(|(added,)| added)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
let mut sadd = redis::cmd("SADD");
|
||||
sadd.arg(&key).arg(values);
|
||||
let mut expire = redis::cmd("EXPIRE");
|
||||
expire.arg(&key).arg(ttl);
|
||||
let replies = connection.pipeline(vec![sadd, expire])?;
|
||||
replies
|
||||
.into_iter()
|
||||
.next()
|
||||
.map(redis::from_redis_value::<usize>)
|
||||
.transpose()
|
||||
.map_err(|_| Error::Unavailable)?
|
||||
.ok_or(Error::Unavailable)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
|
@ -293,11 +280,19 @@ where
|
|||
return Ok(Vec::new());
|
||||
}
|
||||
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
|
||||
let mut pipeline = redis::pipe();
|
||||
for (key, values) in operations {
|
||||
pipeline.cmd("RPUSH").arg(key).arg(values);
|
||||
}
|
||||
pipeline.query(connection).map_err(|_| Error::Unavailable)
|
||||
let commands = operations
|
||||
.into_iter()
|
||||
.map(|(key, values)| {
|
||||
let mut command = redis::cmd("RPUSH");
|
||||
command.arg(key).arg(values);
|
||||
command
|
||||
})
|
||||
.collect();
|
||||
connection
|
||||
.pipeline(commands)?
|
||||
.into_iter()
|
||||
.map(|value| redis::from_redis_value(value).map_err(|_| Error::Unavailable))
|
||||
.collect()
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
|
@ -339,16 +334,18 @@ where
|
|||
.map(|(_, count)| count.is_some())
|
||||
.collect::<Vec<_>>();
|
||||
let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| {
|
||||
let mut pipeline = redis::pipe();
|
||||
for (key, count) in operations {
|
||||
let command = pipeline.cmd("LPOP").arg(key);
|
||||
if let Some(count) = count {
|
||||
command.arg(count);
|
||||
}
|
||||
}
|
||||
pipeline
|
||||
.query::<Vec<redis::Value>>(connection)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
let commands = operations
|
||||
.into_iter()
|
||||
.map(|(key, count)| {
|
||||
let mut command = redis::cmd("LPOP");
|
||||
command.arg(key);
|
||||
if let Some(count) = count {
|
||||
command.arg(count);
|
||||
}
|
||||
command
|
||||
})
|
||||
.collect();
|
||||
connection.pipeline(commands)
|
||||
})
|
||||
.await?;
|
||||
values
|
||||
|
|
@ -381,28 +378,17 @@ where
|
|||
}
|
||||
|
||||
pub fn client_list(&self) -> Result<String, Error> {
|
||||
self.connections.execute(|connection| {
|
||||
redis::cmd("CLIENT")
|
||||
.arg("LIST")
|
||||
.query(connection)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
self.connections
|
||||
.execute(|connection| connection.node_text(redis::cmd("CLIENT").arg("LIST")))
|
||||
}
|
||||
|
||||
pub fn info(&self) -> Result<String, Error> {
|
||||
self.connections.execute(|connection| {
|
||||
redis::cmd("INFO")
|
||||
.query(connection)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
self.connections
|
||||
.execute(|connection| connection.node_text(&redis::cmd("INFO")))
|
||||
}
|
||||
|
||||
pub fn flushall(&self) -> Result<(), Error> {
|
||||
self.connections.execute(|connection| {
|
||||
redis::cmd("FLUSHALL")
|
||||
.query(connection)
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
self.connections.execute(|connection| connection.flushall())
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -441,14 +427,27 @@ where
|
|||
return Ok(Vec::new());
|
||||
}
|
||||
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
|
||||
let mut pipeline = redis::pipe();
|
||||
let mut commands = Vec::with_capacity(operations.len() * 2);
|
||||
let mut increments = Vec::with_capacity(operations.len());
|
||||
for (key, amount, ttl) in operations {
|
||||
pipeline.cmd("INCRBYFLOAT").arg(&key).arg(amount);
|
||||
let mut increment = redis::cmd("INCRBYFLOAT");
|
||||
increment.arg(&key).arg(amount);
|
||||
increments.push(commands.len());
|
||||
commands.push(increment);
|
||||
if let Some(ttl) = ttl {
|
||||
pipeline.cmd("EXPIRE").arg(key).arg(ttl).ignore();
|
||||
let mut expire = redis::cmd("EXPIRE");
|
||||
expire.arg(key).arg(ttl);
|
||||
commands.push(expire);
|
||||
}
|
||||
}
|
||||
pipeline.query(connection).map_err(|_| Error::Unavailable)
|
||||
let mut replies = connection.pipeline(commands)?;
|
||||
increments
|
||||
.into_iter()
|
||||
.map(|index| {
|
||||
redis::from_redis_value(std::mem::take(&mut replies[index]))
|
||||
.map_err(|_| Error::Unavailable)
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
|
|
|||
492
litellm-rust/crates/cache-redis/tests/cluster.rs
Normal file
492
litellm-rust/crates/cache-redis/tests/cluster.rs
Normal file
|
|
@ -0,0 +1,492 @@
|
|||
//! Contract tests against a real Redis Cluster. Set `LITELLM_TEST_REDIS_CLUSTER_NODES` to a
|
||||
//! comma separated `host:port` list (for example `127.0.0.1:7000,127.0.0.1:7001`) to run them.
|
||||
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use litellm_cache::{
|
||||
BaseCache, BatchCache, BatchEntry, CacheConnectionStatus, CacheScript, ClaimCache,
|
||||
CounterCache, DeleteCache, Error, ExactCacheContext, FlushCache, IncrementOperation, JsonCodec,
|
||||
ScriptCache,
|
||||
};
|
||||
use litellm_cache_redis::{
|
||||
RedisArg, RedisCache, RedisLpopOperation, RedisLpopResult, RedisNode, RedisRpushOperation,
|
||||
RedisTopology,
|
||||
};
|
||||
use redis::cluster_routing::Slot;
|
||||
|
||||
type Cache = RedisCache<JsonCodec<serde_json::Value>>;
|
||||
|
||||
fn topology() -> Option<RedisTopology> {
|
||||
let nodes = std::env::var("LITELLM_TEST_REDIS_CLUSTER_NODES").ok()?;
|
||||
let startup_nodes = nodes
|
||||
.split(',')
|
||||
.map(|node| {
|
||||
let (host, port) = node.trim().rsplit_once(':').expect("host:port");
|
||||
RedisNode {
|
||||
host: host.to_string(),
|
||||
port: port.parse().expect("port"),
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
Some(RedisTopology::Cluster { startup_nodes })
|
||||
}
|
||||
|
||||
fn namespace(label: &str) -> String {
|
||||
let nanos = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_nanos();
|
||||
format!("cluster-test:{label}:{nanos}")
|
||||
}
|
||||
|
||||
fn cluster_url() -> String {
|
||||
std::env::var("LITELLM_TEST_REDIS_CLUSTER_URL")
|
||||
.unwrap_or_else(|_| "redis://127.0.0.1:7000".into())
|
||||
}
|
||||
|
||||
fn cluster_cache(label: &str) -> Option<Cache> {
|
||||
let topology = topology()?;
|
||||
Some(
|
||||
Cache::connect(
|
||||
&cluster_url(),
|
||||
&topology,
|
||||
Some(Duration::from_secs(120)),
|
||||
JsonCodec::new(),
|
||||
)
|
||||
.expect("cluster connection")
|
||||
.with_namespace(Some(namespace(label))),
|
||||
)
|
||||
}
|
||||
|
||||
fn counter_cache(label: &str) -> Option<RedisCache<JsonCodec<f64>>> {
|
||||
let topology = topology()?;
|
||||
Some(
|
||||
RedisCache::connect(
|
||||
&cluster_url(),
|
||||
&topology,
|
||||
Some(Duration::from_secs(60)),
|
||||
JsonCodec::new(),
|
||||
)
|
||||
.expect("cluster connection")
|
||||
.with_namespace(Some(namespace(label))),
|
||||
)
|
||||
}
|
||||
|
||||
fn multi_slot_keys(count: usize) -> Vec<String> {
|
||||
let keys: Vec<String> = (0..count).map(|index| format!("key-{index}")).collect();
|
||||
let slots: std::collections::HashSet<Slot> = keys.iter().map(Slot::for_key).collect();
|
||||
assert!(slots.len() > 1, "keys must span multiple slots");
|
||||
keys
|
||||
}
|
||||
|
||||
macro_rules! cluster_or_skip {
|
||||
($label:expr) => {
|
||||
match cluster_cache($label) {
|
||||
Some(cache) => cache,
|
||||
None => return,
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn constructor_rejects_clusters_without_startup_nodes() {
|
||||
let error = Cache::connect(
|
||||
"redis://127.0.0.1:7000",
|
||||
&RedisTopology::Cluster {
|
||||
startup_nodes: Vec::new(),
|
||||
},
|
||||
None,
|
||||
JsonCodec::new(),
|
||||
)
|
||||
.err();
|
||||
assert!(matches!(error, Some(Error::Unavailable)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn constructor_rejects_unix_socket_urls_for_clusters() {
|
||||
let error = Cache::connect(
|
||||
"redis+unix:///tmp/redis.sock",
|
||||
&RedisTopology::Cluster {
|
||||
startup_nodes: vec![RedisNode {
|
||||
host: "127.0.0.1".into(),
|
||||
port: 7000,
|
||||
}],
|
||||
},
|
||||
None,
|
||||
JsonCodec::new(),
|
||||
)
|
||||
.err();
|
||||
assert!(matches!(error, Some(Error::Unavailable)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_key_operations_round_trip_with_ttl_rounding() {
|
||||
let cache = cluster_or_skip!("single");
|
||||
let context = ExactCacheContext {
|
||||
ttl: Some(Duration::from_millis(1500)),
|
||||
};
|
||||
let keys = multi_slot_keys(12);
|
||||
for (index, key) in keys.iter().enumerate() {
|
||||
cache
|
||||
.set_cache(key, serde_json::json!({ "index": index }), &context)
|
||||
.unwrap();
|
||||
}
|
||||
for (index, key) in keys.iter().enumerate() {
|
||||
assert_eq!(
|
||||
cache.get_cache(key, &context).unwrap(),
|
||||
Some(serde_json::json!({ "index": index }))
|
||||
);
|
||||
}
|
||||
let runtime = tokio::runtime::Runtime::new().unwrap();
|
||||
let ttl = runtime.block_on(cache.async_get_ttl(&keys[0])).unwrap();
|
||||
assert_eq!(ttl, Some(2));
|
||||
cache.delete_cache(&keys[0]).unwrap();
|
||||
assert_eq!(cache.get_cache(&keys[0], &context).unwrap(), None);
|
||||
assert!(cache.sync_ping().unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn batch_reads_span_slots_and_preserve_order_with_malformed_entries() {
|
||||
let cache = cluster_or_skip!("batch");
|
||||
let context = ExactCacheContext::default();
|
||||
let keys = multi_slot_keys(40);
|
||||
for (index, key) in keys.iter().enumerate() {
|
||||
if index % 5 == 0 {
|
||||
continue;
|
||||
}
|
||||
cache
|
||||
.async_set_cache(key, serde_json::json!(index), context.clone())
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
let mut raw = redis::cluster::ClusterClient::new(vec![cluster_url()])
|
||||
.unwrap()
|
||||
.get_connection()
|
||||
.unwrap();
|
||||
let malformed = format!("{}:{}", cache.namespace().unwrap(), keys[1]);
|
||||
redis::cmd("SET")
|
||||
.arg(&malformed)
|
||||
.arg("not json")
|
||||
.exec(&mut raw)
|
||||
.unwrap();
|
||||
|
||||
let entries = cache
|
||||
.async_batch_get_cache(keys.clone(), context.clone())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(entries.len(), keys.len());
|
||||
for (index, entry) in entries.iter().enumerate() {
|
||||
let expected = if index == 1 {
|
||||
BatchEntry::Invalid
|
||||
} else if index % 5 == 0 {
|
||||
BatchEntry::Miss
|
||||
} else {
|
||||
BatchEntry::Hit(serde_json::json!(index))
|
||||
};
|
||||
assert_eq!(*entry, expected, "entry {index}");
|
||||
}
|
||||
let sync_entries = cache.batch_get_cache(&keys, &context).unwrap();
|
||||
assert_eq!(sync_entries, entries);
|
||||
|
||||
cache.delete_cache_keys(keys.clone()).await.unwrap();
|
||||
let entries = cache.async_batch_get_cache(keys, context).await.unwrap();
|
||||
assert!(entries.iter().all(|entry| *entry == BatchEntry::Miss));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pipelines_group_by_slot_and_return_results_in_submission_order() {
|
||||
let cache = cluster_or_skip!("pipeline");
|
||||
let keys = multi_slot_keys(30);
|
||||
let entries = keys
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, key)| (key.clone(), serde_json::json!(index)))
|
||||
.collect();
|
||||
cache
|
||||
.async_set_cache_pipeline(entries, ExactCacheContext::default())
|
||||
.await
|
||||
.unwrap();
|
||||
let hits = cache
|
||||
.async_batch_get_cache(keys.clone(), ExactCacheContext::default())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(
|
||||
hits.iter()
|
||||
.enumerate()
|
||||
.all(|(index, entry)| *entry == BatchEntry::Hit(serde_json::json!(index)))
|
||||
);
|
||||
|
||||
let queues: Vec<String> = keys.iter().map(|key| format!("queue:{key}")).collect();
|
||||
let pushed = cache
|
||||
.async_rpush_pipeline(
|
||||
queues
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, key)| RedisRpushOperation {
|
||||
key: key.clone(),
|
||||
values: (0..=index)
|
||||
.map(|value| RedisArg::Integer(value as i64))
|
||||
.collect(),
|
||||
})
|
||||
.collect(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(pushed, (1..=keys.len()).collect::<Vec<_>>());
|
||||
let popped = cache
|
||||
.async_lpop_pipeline(
|
||||
queues
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, key)| RedisLpopOperation {
|
||||
key: key.clone(),
|
||||
count: (index % 2 == 0).then_some(2),
|
||||
})
|
||||
.collect(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
for (index, result) in popped.into_iter().enumerate() {
|
||||
match result {
|
||||
RedisLpopResult::Value(value) => {
|
||||
assert_eq!(index % 2, 1, "queue {index}");
|
||||
assert_eq!(value, b"0");
|
||||
}
|
||||
RedisLpopResult::Values(values) => {
|
||||
assert_eq!(index % 2, 0, "queue {index}");
|
||||
let expected: Vec<Vec<u8>> = (0..=index)
|
||||
.take(2)
|
||||
.map(|value| value.to_string().into_bytes())
|
||||
.collect();
|
||||
assert_eq!(values, expected);
|
||||
}
|
||||
other => panic!("queue {index}: {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
let counters: Vec<String> = keys.iter().map(|key| format!("counter:{key}")).collect();
|
||||
let Some(counter) = counter_cache("counter") else {
|
||||
return;
|
||||
};
|
||||
let totals = counter
|
||||
.async_increment_pipeline(
|
||||
counters
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, key)| IncrementOperation {
|
||||
key: key.clone(),
|
||||
amount: index as f64 + 0.5,
|
||||
ttl: (index % 3 == 0).then_some(Duration::from_secs(30)),
|
||||
})
|
||||
.collect(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let expected: Vec<f64> = (0..keys.len()).map(|index| index as f64 + 0.5).collect();
|
||||
assert_eq!(totals, expected);
|
||||
assert_eq!(counter.async_get_ttl(&counters[0]).await.unwrap(), Some(30));
|
||||
assert_eq!(counter.async_get_ttl(&counters[1]).await.unwrap(), None);
|
||||
counter.async_flush_cache().await.unwrap();
|
||||
cache.async_flush_cache().await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn scan_and_scoped_flush_cover_every_primary() {
|
||||
let cache = cluster_or_skip!("flush");
|
||||
let other = cluster_or_skip!("other");
|
||||
let context = ExactCacheContext::default();
|
||||
let keys = multi_slot_keys(60);
|
||||
for key in &keys {
|
||||
cache
|
||||
.async_set_cache(key, serde_json::json!(true), context.clone())
|
||||
.await
|
||||
.unwrap();
|
||||
other
|
||||
.async_set_cache(key, serde_json::json!(true), context.clone())
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
let mut scanned = cache.async_scan_iter("key-", 1000).await.unwrap();
|
||||
scanned.sort();
|
||||
let mut expected: Vec<String> = keys
|
||||
.iter()
|
||||
.map(|key| format!("{}:{key}", cache.namespace().unwrap()))
|
||||
.collect();
|
||||
expected.sort();
|
||||
assert_eq!(scanned, expected);
|
||||
assert_eq!(cache.async_scan_iter("key-", 7).await.unwrap().len(), 7);
|
||||
|
||||
cache.flush_cache().unwrap();
|
||||
let flushed = cache
|
||||
.async_batch_get_cache(keys.clone(), context.clone())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(flushed.iter().all(|entry| *entry == BatchEntry::Miss));
|
||||
let kept = other.async_batch_get_cache(keys, context).await.unwrap();
|
||||
assert!(
|
||||
kept.iter()
|
||||
.all(|entry| *entry == BatchEntry::Hit(serde_json::json!(true)))
|
||||
);
|
||||
other.async_flush_cache().await.unwrap();
|
||||
}
|
||||
|
||||
fn ping_calls_per_node(startup: &redis::Client) -> Vec<(String, u64)> {
|
||||
let mut connection = startup.get_connection().unwrap();
|
||||
let nodes: String = redis::cmd("CLUSTER")
|
||||
.arg("NODES")
|
||||
.query(&mut connection)
|
||||
.unwrap();
|
||||
let mut counts: Vec<(String, u64)> = nodes
|
||||
.lines()
|
||||
.map(|line| {
|
||||
let address = line.split_whitespace().nth(1).unwrap();
|
||||
let address = address.split('@').next().unwrap();
|
||||
let mut node = redis::Client::open(format!("redis://{address}"))
|
||||
.unwrap()
|
||||
.get_connection()
|
||||
.unwrap();
|
||||
let stats: String = redis::cmd("INFO")
|
||||
.arg("commandstats")
|
||||
.query(&mut node)
|
||||
.unwrap();
|
||||
let calls = stats
|
||||
.lines()
|
||||
.find_map(|stat| stat.strip_prefix("cmdstat_ping:calls="))
|
||||
.and_then(|rest| rest.split(',').next())
|
||||
.map_or(0, |calls| calls.parse().unwrap());
|
||||
(address.to_string(), calls)
|
||||
})
|
||||
.collect();
|
||||
counts.sort();
|
||||
counts
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ping_reaches_every_node() {
|
||||
let cache = cluster_or_skip!("ping");
|
||||
let startup = redis::Client::open(cluster_url()).unwrap();
|
||||
let before = ping_calls_per_node(&startup);
|
||||
assert!(before.len() >= 2, "{before:?}");
|
||||
assert!(cache.ping().await.unwrap());
|
||||
let after = ping_calls_per_node(&startup);
|
||||
for ((node, calls_before), (_, calls_after)) in before.iter().zip(&after) {
|
||||
assert!(calls_after > calls_before, "{node} was not pinged");
|
||||
}
|
||||
assert!(cache.sync_ping().unwrap());
|
||||
let result = cache.test_connection().await.unwrap();
|
||||
assert_eq!(result.status, CacheConnectionStatus::Success);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn counters_claims_scripts_and_sets_work_on_the_cluster() {
|
||||
let Some(counter) = counter_cache("counter") else {
|
||||
return;
|
||||
};
|
||||
let context = ExactCacheContext::default();
|
||||
assert_eq!(
|
||||
counter
|
||||
.increment_cache("spend", 1.5, context.clone())
|
||||
.unwrap(),
|
||||
1.5
|
||||
);
|
||||
assert_eq!(
|
||||
counter
|
||||
.async_increment("spend", 2.0, context.clone())
|
||||
.await
|
||||
.unwrap(),
|
||||
3.5
|
||||
);
|
||||
assert_eq!(
|
||||
counter
|
||||
.increment_with_floor("budget", -3, Duration::from_secs(30))
|
||||
.unwrap(),
|
||||
0
|
||||
);
|
||||
assert_eq!(
|
||||
counter
|
||||
.async_increment_with_floor("budget", 7, Duration::from_secs(30))
|
||||
.await
|
||||
.unwrap(),
|
||||
7
|
||||
);
|
||||
assert_eq!(counter.async_set_max("peak", 4.0, None).await.unwrap(), 4.0);
|
||||
assert_eq!(counter.async_set_max("peak", 2.0, None).await.unwrap(), 4.0);
|
||||
counter.flush_cache().unwrap();
|
||||
|
||||
let cache = cluster_or_skip!("claim");
|
||||
let owner = serde_json::json!("owner-a");
|
||||
let rival = serde_json::json!("owner-b");
|
||||
assert_eq!(
|
||||
cache
|
||||
.claim_cache("lock", owner.clone(), &[], context.clone())
|
||||
.unwrap(),
|
||||
owner
|
||||
);
|
||||
assert_eq!(
|
||||
cache
|
||||
.async_claim_cache("lock", rival.clone(), vec![owner.clone()], context.clone())
|
||||
.await
|
||||
.unwrap(),
|
||||
owner
|
||||
);
|
||||
assert_eq!(
|
||||
cache
|
||||
.claim_cache("lock", rival.clone(), &[], context.clone())
|
||||
.unwrap(),
|
||||
owner
|
||||
);
|
||||
assert_eq!(
|
||||
cache
|
||||
.async_claim_cache("lock", rival.clone(), vec![rival.clone()], context.clone())
|
||||
.await
|
||||
.unwrap(),
|
||||
rival
|
||||
);
|
||||
|
||||
let script = cache
|
||||
.async_register_script("return redis.call('SET', KEYS[1], ARGV[1], 'EX', ARGV[2])".into());
|
||||
let reply = script
|
||||
.invoke(
|
||||
vec!["scripted".into()],
|
||||
vec![RedisArg::Bytes(b"payload".to_vec()), RedisArg::Integer(5)],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(reply, redis::Value::Okay);
|
||||
assert_eq!(cache.async_get_ttl("scripted").await.unwrap(), Some(5));
|
||||
let evaluated: redis::Value = cache
|
||||
.async_eval(
|
||||
"return redis.call('GET', KEYS[1])".into(),
|
||||
vec!["scripted".into()],
|
||||
Vec::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(evaluated, redis::Value::BulkString(b"payload".to_vec()));
|
||||
|
||||
assert_eq!(
|
||||
cache
|
||||
.async_set_cache_sadd(
|
||||
"members",
|
||||
vec![
|
||||
RedisArg::Bytes(b"a".to_vec()),
|
||||
RedisArg::Bytes(b"b".to_vec())
|
||||
],
|
||||
Some(Duration::from_secs(9)),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
2
|
||||
);
|
||||
assert_eq!(cache.async_get_ttl("members").await.unwrap(), Some(9));
|
||||
|
||||
let result = cache.test_connection().await.unwrap();
|
||||
assert_eq!(result.status, CacheConnectionStatus::Success);
|
||||
assert!(cache.ping().await.unwrap());
|
||||
let info = cache.info().unwrap();
|
||||
assert!(info.matches("redis_version").count() > 1, "{info}");
|
||||
assert!(cache.client_list().unwrap().contains("id="));
|
||||
cache.async_flush_cache().await.unwrap();
|
||||
assert_eq!(cache.async_get_ttl("members").await.unwrap(), None);
|
||||
assert_eq!(cache.get_cache("lock", &context).unwrap(), None);
|
||||
}
|
||||
|
|
@ -1,10 +1,11 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_cache::CacheType;
|
||||
use litellm_cache_redis::{RedisNode, RedisTopology};
|
||||
use pyo3::{
|
||||
exceptions::{PyTypeError, PyValueError},
|
||||
prelude::*,
|
||||
types::{PyAny, PyDict, PyString},
|
||||
types::{PyAny, PyDict, PyList, PyString},
|
||||
};
|
||||
|
||||
use super::{native::NativeResponseCache, request::duration};
|
||||
|
|
@ -70,9 +71,21 @@ pub(super) struct RedisCacheConfig {
|
|||
pub(super) default_ttl: Duration,
|
||||
pub(super) namespace: Option<String>,
|
||||
pub(super) flush_size: usize,
|
||||
pub(super) topology: RedisTopology,
|
||||
pub(super) connection: RedisConnectionConfig,
|
||||
}
|
||||
|
||||
struct RedisClientProjection<'py> {
|
||||
topology: RedisTopology,
|
||||
host: String,
|
||||
port: u16,
|
||||
pool_size: usize,
|
||||
resolved: Bound<'py, PyDict>,
|
||||
tls: Option<RedisTlsConfig>,
|
||||
}
|
||||
|
||||
const REDIS_PY_DEFAULT_MAX_CONNECTIONS: usize = 1 << 31;
|
||||
|
||||
pub(super) enum CacheBackendConfig {
|
||||
Memory(MemoryCacheConfig),
|
||||
Redis(Box<RedisCacheConfig>),
|
||||
|
|
@ -182,6 +195,9 @@ impl NativeCacheConfig {
|
|||
CacheBackendConfig::Redis(_) if service.kind() != "redis" => {
|
||||
Some("facade and native backend types must match")
|
||||
}
|
||||
CacheBackendConfig::Redis(config) if service.topology() != Some(&config.topology) => {
|
||||
Some("facade and native backend topologies must match")
|
||||
}
|
||||
CacheBackendConfig::Redis(config) => (service.namespace()
|
||||
!= config.namespace.as_deref())
|
||||
.then_some("facade and native backend namespaces must match"),
|
||||
|
|
@ -206,9 +222,6 @@ fn project_redis(
|
|||
backend: &Bound<'_, PyAny>,
|
||||
) -> PyResult<Result<RedisCacheConfig, UnsupportedCacheConfig>> {
|
||||
let source = backend.getattr("redis_kwargs")?.cast_into::<PyDict>()?;
|
||||
if has_value(&source, "startup_nodes")? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisTopology));
|
||||
}
|
||||
if has_value(&source, "sentinel_nodes")? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisTopology));
|
||||
}
|
||||
|
|
@ -248,26 +261,25 @@ fn project_redis(
|
|||
}
|
||||
|
||||
let client = backend.getattr("redis_client")?;
|
||||
let pool = client.getattr("connection_pool")?;
|
||||
if !instance_class_is(&pool, "redis.connection", "ConnectionPool")? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
|
||||
}
|
||||
let resolved = pool.getattr("connection_kwargs")?.cast_into::<PyDict>()?;
|
||||
for key in ["credential_provider", "redis_connect_func"] {
|
||||
if has_value(&resolved, key)? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
|
||||
}
|
||||
}
|
||||
let connection_class = resolved
|
||||
.get_item("connection_class")?
|
||||
.unwrap_or(pool.getattr("connection_class")?);
|
||||
let tls = if class_is(&connection_class, "redis.connection", "Connection")? {
|
||||
None
|
||||
} else if class_is(&connection_class, "redis.connection", "SSLConnection")? {
|
||||
Some(project_tls(&resolved)?)
|
||||
let projection = if has_value(&source, "startup_nodes")? {
|
||||
project_cluster_client(&source, &client)?
|
||||
} else {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
|
||||
project_standalone_client(&client)?
|
||||
};
|
||||
let RedisClientProjection {
|
||||
topology,
|
||||
host,
|
||||
port,
|
||||
pool_size,
|
||||
resolved,
|
||||
tls,
|
||||
} = match projection {
|
||||
Ok(projection) => projection,
|
||||
Err(reason) => return Ok(Err(reason)),
|
||||
};
|
||||
if has_value(&resolved, "credential_provider")? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
|
||||
}
|
||||
|
||||
let protocol = match optional_i64(&resolved, "protocol")?.unwrap_or(2) {
|
||||
2 => RedisProtocol::Resp2,
|
||||
|
|
@ -280,15 +292,15 @@ fn project_redis(
|
|||
default_ttl: duration(backend.getattr("default_ttl")?.extract::<f64>()?)?,
|
||||
namespace: optional_attribute_string(backend, "namespace")?,
|
||||
flush_size: backend.getattr("redis_flush_size")?.extract::<usize>()?,
|
||||
topology,
|
||||
connection: RedisConnectionConfig {
|
||||
host: required_string(&resolved, "host")?,
|
||||
port: u16::try_from(required_i64(&resolved, "port")?)
|
||||
.map_err(|_| PyValueError::new_err("invalid Redis port"))?,
|
||||
host,
|
||||
port,
|
||||
database: optional_i64(&resolved, "db")?.unwrap_or(0),
|
||||
username: optional_dict_string(&resolved, "username")?,
|
||||
password: optional_dict_string(&resolved, "password")?,
|
||||
protocol,
|
||||
pool_size: pool.getattr("max_connections")?.extract::<usize>()?,
|
||||
pool_size,
|
||||
read_timeout: optional_dict_duration(&resolved, "socket_timeout")?,
|
||||
connect_timeout: optional_dict_duration(&resolved, "socket_connect_timeout")?,
|
||||
socket_keepalive: optional_bool(&resolved, "socket_keepalive")?,
|
||||
|
|
@ -299,6 +311,128 @@ fn project_redis(
|
|||
}))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
fn project_standalone_client<'py>(
|
||||
client: &Bound<'py, PyAny>,
|
||||
) -> PyResult<Result<RedisClientProjection<'py>, UnsupportedCacheConfig>> {
|
||||
let pool = client.getattr("connection_pool")?;
|
||||
if !instance_class_is(&pool, "redis.connection", "ConnectionPool")? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
|
||||
}
|
||||
let resolved = pool.getattr("connection_kwargs")?.cast_into::<PyDict>()?;
|
||||
if has_value(&resolved, "redis_connect_func")? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
|
||||
}
|
||||
let connection_class = resolved
|
||||
.get_item("connection_class")?
|
||||
.unwrap_or(pool.getattr("connection_class")?);
|
||||
let tls = if class_is(&connection_class, "redis.connection", "Connection")? {
|
||||
None
|
||||
} else if class_is(&connection_class, "redis.connection", "SSLConnection")? {
|
||||
Some(project_tls(&resolved)?)
|
||||
} else {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
|
||||
};
|
||||
Ok(Ok(RedisClientProjection {
|
||||
topology: RedisTopology::Standalone,
|
||||
host: required_string(&resolved, "host")?,
|
||||
port: port(required_i64(&resolved, "port")?)?,
|
||||
pool_size: pool.getattr("max_connections")?.extract::<usize>()?,
|
||||
resolved,
|
||||
tls,
|
||||
}))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
fn project_cluster_client<'py>(
|
||||
source: &Bound<'py, PyDict>,
|
||||
client: &Bound<'py, PyAny>,
|
||||
) -> PyResult<Result<RedisClientProjection<'py>, UnsupportedCacheConfig>> {
|
||||
let Some(startup_nodes) = startup_nodes(source)? else {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisTopology));
|
||||
};
|
||||
if !instance_class_is(client, "redis.cluster", "RedisCluster")? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
|
||||
}
|
||||
let nodes = client.getattr("nodes_manager")?;
|
||||
if !class_is(
|
||||
&nodes.getattr("connection_pool_class")?,
|
||||
"redis.connection",
|
||||
"ConnectionPool",
|
||||
)? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
|
||||
}
|
||||
let resolved = nodes.getattr("connection_kwargs")?.cast_into::<PyDict>()?;
|
||||
if let Some(connect) = resolved.get_item("redis_connect_func")?
|
||||
&& !connect.is_none()
|
||||
{
|
||||
let own_hook = connect
|
||||
.getattr("__self__")
|
||||
.is_ok_and(|owner| owner.is(client))
|
||||
&& connect
|
||||
.getattr("__func__")
|
||||
.and_then(|function| Ok(function.is(&client.get_type().getattr("on_connect")?)))
|
||||
.unwrap_or(false);
|
||||
if !own_hook {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
|
||||
}
|
||||
}
|
||||
let tls = if optional_bool(&resolved, "ssl")?.unwrap_or(false) {
|
||||
Some(project_tls(&resolved)?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let first = &startup_nodes[0];
|
||||
Ok(Ok(RedisClientProjection {
|
||||
host: first.host.clone(),
|
||||
port: first.port,
|
||||
pool_size: optional_i64(&resolved, "max_connections")?
|
||||
.map(|value| {
|
||||
usize::try_from(value).map_err(|_| PyValueError::new_err("invalid Redis pool size"))
|
||||
})
|
||||
.transpose()?
|
||||
.unwrap_or(REDIS_PY_DEFAULT_MAX_CONNECTIONS),
|
||||
topology: RedisTopology::Cluster { startup_nodes },
|
||||
resolved,
|
||||
tls,
|
||||
}))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
fn startup_nodes(source: &Bound<'_, PyDict>) -> PyResult<Option<Vec<RedisNode>>> {
|
||||
let Some(nodes) = source.get_item("startup_nodes")? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Ok(nodes) = nodes.cast_into::<PyList>() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if nodes.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
let mut parsed = Vec::with_capacity(nodes.len());
|
||||
for node in nodes.iter() {
|
||||
let Ok(node) = node.cast_into::<PyDict>() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if node.len() != 2 || !has_value(&node, "host")? || !has_value(&node, "port")? {
|
||||
return Ok(None);
|
||||
}
|
||||
let (Ok(host), Ok(port)) = (
|
||||
required_string(&node, "host"),
|
||||
required_i64(&node, "port").and_then(port),
|
||||
) else {
|
||||
return Ok(None);
|
||||
};
|
||||
parsed.push(RedisNode { host, port });
|
||||
}
|
||||
Ok(Some(parsed))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
fn port(value: i64) -> PyResult<u16> {
|
||||
u16::try_from(value).map_err(|_| PyValueError::new_err("invalid Redis port"))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
fn project_tls(values: &Bound<'_, PyDict>) -> PyResult<RedisTlsConfig> {
|
||||
Ok(RedisTlsConfig {
|
||||
|
|
@ -467,12 +601,27 @@ mod tests {
|
|||
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
|
||||
use litellm_cache_redis::{RedisNode, RedisTopology};
|
||||
|
||||
use super::{
|
||||
CacheBackendConfig, CacheConfigProjection, CertificateRequirement, NativeCacheConfig,
|
||||
RedisProtocol,
|
||||
};
|
||||
use crate::cache::native::NativeResponseCache;
|
||||
|
||||
fn cluster_facade<'py>(py: Python<'py>, startup_nodes: &str, hook: &str) -> Bound<'py, PyAny> {
|
||||
facade(
|
||||
py,
|
||||
&format!(
|
||||
"RedisCluster = type('RedisCluster', (), {{'__module__': 'redis.cluster', 'on_connect': lambda self, connection: None}})\n\
|
||||
client = RedisCluster()\n\
|
||||
client.nodes_manager = SimpleNamespace(connection_pool_class=ConnectionPool, connection_kwargs={{'password': 'secret', 'redis_connect_func': {hook}, 'protocol': 3, 'ssl': True, 'ssl_cert_reqs': 'none'}})\n\
|
||||
backend = SimpleNamespace(default_ttl=120, namespace='team', redis_flush_size=100, redis_kwargs={{'startup_nodes': {startup_nodes}, 'password': 'secret'}}, redis_client=client)\n\
|
||||
facade = SimpleNamespace(type='redis', mode='default-on', ttl=None, namespace='team', supported_call_types=None, redis_flush_size=100, semantic_cache_scope='key', cache=backend)"
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
fn facade<'py>(py: Python<'py>, body: &str) -> Bound<'py, PyAny> {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
|
|
@ -591,4 +740,87 @@ mod tests {
|
|||
assert_eq!(reason.message(), "native Redis credentials require Python");
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn projects_cluster_startup_nodes_as_redis_topology() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let facade = cluster_facade(
|
||||
py,
|
||||
"[{'host': 'node-a', 'port': 7000}, {'host': 'node-b', 'port': 7001}]",
|
||||
"client.on_connect",
|
||||
);
|
||||
let CacheConfigProjection::Native(config) =
|
||||
NativeCacheConfig::project(&facade).unwrap()
|
||||
else {
|
||||
panic!("cluster startup nodes should project natively");
|
||||
};
|
||||
let CacheBackendConfig::Redis(redis) = &config.backend else {
|
||||
panic!("expected Redis configuration");
|
||||
};
|
||||
let expected = RedisTopology::Cluster {
|
||||
startup_nodes: vec![
|
||||
RedisNode {
|
||||
host: "node-a".into(),
|
||||
port: 7000,
|
||||
},
|
||||
RedisNode {
|
||||
host: "node-b".into(),
|
||||
port: 7001,
|
||||
},
|
||||
],
|
||||
};
|
||||
assert_eq!(redis.topology, expected);
|
||||
assert_eq!(redis.connection.host, "node-a");
|
||||
assert_eq!(redis.connection.port, 7000);
|
||||
assert_eq!(redis.connection.password.as_deref(), Some("secret"));
|
||||
assert_eq!(redis.connection.protocol, RedisProtocol::Resp3);
|
||||
assert_eq!(
|
||||
redis
|
||||
.connection
|
||||
.tls
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.certificate_requirement,
|
||||
CertificateRequirement::None
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_startup_nodes_and_foreign_connect_hooks_stay_on_python() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
for (startup_nodes, hook, message) in [
|
||||
(
|
||||
"[{'host': 'node-a', 'port': 7000, 'server_type': 'primary'}]",
|
||||
"client.on_connect",
|
||||
"native Redis topology is not implemented",
|
||||
),
|
||||
(
|
||||
"[{'host': 'node-a', 'port': 'seven'}]",
|
||||
"client.on_connect",
|
||||
"native Redis topology is not implemented",
|
||||
),
|
||||
(
|
||||
"[]",
|
||||
"client.on_connect",
|
||||
"native Redis topology is not implemented",
|
||||
),
|
||||
(
|
||||
"[{'host': 'node-a', 'port': 7000}]",
|
||||
"lambda connection: None",
|
||||
"native Redis credentials require Python",
|
||||
),
|
||||
] {
|
||||
let facade = cluster_facade(py, startup_nodes, hook);
|
||||
let CacheConfigProjection::Unsupported(reason) =
|
||||
NativeCacheConfig::project(&facade).unwrap()
|
||||
else {
|
||||
panic!("{startup_nodes} with {hook} must stay on Python");
|
||||
};
|
||||
assert_eq!(reason.message(), message, "{startup_nodes} with {hook}");
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use litellm_cache_redis::RedisTopology;
|
||||
use litellm_host_python::from_py;
|
||||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
|
|
@ -29,9 +30,28 @@ struct RedisPoolGuard {
|
|||
reference: Py<PyAny>,
|
||||
connection_class: Py<PyAny>,
|
||||
connection_kwargs: Py<PyAny>,
|
||||
max_connections: usize,
|
||||
max_connections: Option<usize>,
|
||||
attributes: RedisPoolAttributes,
|
||||
}
|
||||
|
||||
struct RedisPoolAttributes {
|
||||
pool: &'static str,
|
||||
connection_class: &'static str,
|
||||
max_connections: Option<&'static str>,
|
||||
}
|
||||
|
||||
const STANDALONE_POOL: RedisPoolAttributes = RedisPoolAttributes {
|
||||
pool: "connection_pool",
|
||||
connection_class: "connection_class",
|
||||
max_connections: Some("max_connections"),
|
||||
};
|
||||
|
||||
const CLUSTER_POOL: RedisPoolAttributes = RedisPoolAttributes {
|
||||
pool: "nodes_manager",
|
||||
connection_class: "connection_pool_class",
|
||||
max_connections: None,
|
||||
};
|
||||
|
||||
pub(super) struct FacadeGuard {
|
||||
outer: ObjectGuard,
|
||||
backend: ObjectGuard,
|
||||
|
|
@ -138,31 +158,40 @@ impl ObjectGuard {
|
|||
}
|
||||
|
||||
impl RedisPoolGuard {
|
||||
fn capture(backend: &Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
let pool = backend
|
||||
.getattr("redis_client")?
|
||||
.getattr("connection_pool")?;
|
||||
fn capture(backend: &Bound<'_, PyAny>, attributes: RedisPoolAttributes) -> PyResult<Self> {
|
||||
let pool = backend.getattr("redis_client")?.getattr(attributes.pool)?;
|
||||
Ok(Self {
|
||||
reference: pool.clone().unbind(),
|
||||
connection_class: pool.getattr("connection_class")?.unbind(),
|
||||
connection_class: pool.getattr(attributes.connection_class)?.unbind(),
|
||||
connection_kwargs: pool
|
||||
.getattr("connection_kwargs")?
|
||||
.call_method0("copy")?
|
||||
.unbind(),
|
||||
max_connections: pool.getattr("max_connections")?.extract::<usize>()?,
|
||||
max_connections: Self::max_connections(&pool, &attributes)?,
|
||||
attributes,
|
||||
})
|
||||
}
|
||||
|
||||
fn max_connections(
|
||||
pool: &Bound<'_, PyAny>,
|
||||
attributes: &RedisPoolAttributes,
|
||||
) -> PyResult<Option<usize>> {
|
||||
attributes
|
||||
.max_connections
|
||||
.map(|name| pool.getattr(name)?.extract::<usize>())
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
let pool = backend
|
||||
.getattr("redis_client")?
|
||||
.getattr("connection_pool")?;
|
||||
.getattr(self.attributes.pool)?;
|
||||
Ok(self.reference.bind(py).is(&pool)
|
||||
&& self
|
||||
.connection_class
|
||||
.bind(py)
|
||||
.is(&pool.getattr("connection_class")?)
|
||||
&& self.max_connections == pool.getattr("max_connections")?.extract::<usize>()?
|
||||
.is(&pool.getattr(self.attributes.connection_class)?)
|
||||
&& self.max_connections == Self::max_connections(&pool, &self.attributes)?
|
||||
&& self
|
||||
.connection_kwargs
|
||||
.bind(py)
|
||||
|
|
@ -189,9 +218,15 @@ impl FacadeGuard {
|
|||
"only exact built-in Cache facades can be registered",
|
||||
));
|
||||
}
|
||||
let (module, name, cache_kind) = match kind {
|
||||
"memory" => ("litellm.caching.in_memory_cache", "InMemoryCache", "local"),
|
||||
"redis" => ("litellm.caching.redis_cache", "RedisCache", "redis"),
|
||||
let cluster = matches!(service.topology(), Some(RedisTopology::Cluster { .. }));
|
||||
let (module, name, cache_kind) = match (kind, cluster) {
|
||||
("memory", _) => ("litellm.caching.in_memory_cache", "InMemoryCache", "local"),
|
||||
("redis", false) => ("litellm.caching.redis_cache", "RedisCache", "redis"),
|
||||
("redis", true) => (
|
||||
"litellm.caching.redis_cluster_cache",
|
||||
"RedisClusterCache",
|
||||
"redis",
|
||||
),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
let backend = facade.getattr("cache")?;
|
||||
|
|
@ -237,9 +272,11 @@ impl FacadeGuard {
|
|||
"redis_flush_size",
|
||||
],
|
||||
)?,
|
||||
redis_pool: (kind == "redis")
|
||||
.then(|| RedisPoolGuard::capture(&backend))
|
||||
.transpose()?,
|
||||
redis_pool: match (kind, cluster) {
|
||||
("redis", false) => Some(RedisPoolGuard::capture(&backend, STANDALONE_POOL)?),
|
||||
("redis", true) => Some(RedisPoolGuard::capture(&backend, CLUSTER_POOL)?),
|
||||
_ => None,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use litellm_cache_redis::{RedisNode, RedisTopology};
|
||||
use litellm_host_python::release_gil;
|
||||
use pyo3::{PyTraverseError, PyVisit, exceptions::PyRuntimeError, prelude::*};
|
||||
|
||||
|
|
@ -34,16 +35,28 @@ impl CacheTestHandle {
|
|||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (url, *, ttl_seconds=60.0, namespace=None))]
|
||||
#[pyo3(signature = (url, *, ttl_seconds=60.0, namespace=None, startup_nodes=None))]
|
||||
fn redis(
|
||||
py: Python<'_>,
|
||||
url: String,
|
||||
ttl_seconds: f64,
|
||||
namespace: Option<String>,
|
||||
startup_nodes: Option<Vec<(String, u16)>>,
|
||||
) -> PyResult<Self> {
|
||||
let ttl = Some(duration(ttl_seconds)?);
|
||||
let service = release_gil(py, move || NativeResponseCache::redis(&url, ttl, namespace))
|
||||
.map_err(cache_error)?;
|
||||
let topology = match startup_nodes {
|
||||
None => RedisTopology::Standalone,
|
||||
Some(nodes) => RedisTopology::Cluster {
|
||||
startup_nodes: nodes
|
||||
.into_iter()
|
||||
.map(|(host, port)| RedisNode { host, port })
|
||||
.collect(),
|
||||
},
|
||||
};
|
||||
let service = release_gil(py, move || {
|
||||
NativeResponseCache::redis(&url, &topology, ttl, namespace)
|
||||
})
|
||||
.map_err(cache_error)?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ use std::{sync::Arc, time::Duration};
|
|||
|
||||
use litellm_cache::{CacheCodec, CacheConnectionResult, Error};
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_redis::RedisCache;
|
||||
use litellm_cache_redis::{RedisCache, RedisTopology};
|
||||
use litellm_cache_response::{
|
||||
CacheEntry, PartialHits, ResponseCache, ResponseCacheCodec, ResponseCacheRequest, WriteBuffer,
|
||||
};
|
||||
|
|
@ -34,10 +34,12 @@ impl NativeResponseCache {
|
|||
|
||||
pub fn redis(
|
||||
url: &str,
|
||||
topology: &RedisTopology,
|
||||
ttl: Option<Duration>,
|
||||
namespace: Option<String>,
|
||||
) -> Result<Self, Error> {
|
||||
let backend = RedisCache::new(url, ttl, ResponseCacheCodec)?.with_namespace(namespace);
|
||||
let backend =
|
||||
RedisCache::connect(url, topology, ttl, ResponseCacheCodec)?.with_namespace(namespace);
|
||||
Ok(Self::Redis {
|
||||
cache: Arc::new(ResponseCache::new(Arc::new(backend))),
|
||||
buffer: None,
|
||||
|
|
@ -67,6 +69,13 @@ impl NativeResponseCache {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn topology(&self) -> Option<&RedisTopology> {
|
||||
match self {
|
||||
Self::Memory(_) => None,
|
||||
Self::Redis { cache, .. } => Some(cache.backend().topology()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn capacity(&self) -> Option<usize> {
|
||||
match self {
|
||||
Self::Memory(cache) => Some(cache.backend().max_size_in_memory()),
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import contextvars
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
|
|
@ -17,6 +18,7 @@ import redis
|
|||
import litellm
|
||||
from litellm.caching.caching import Cache, disable_cache, enable_cache, update_cache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||||
from litellm.rust_bridge import _native
|
||||
from litellm.types.caching import LiteLLMCacheType
|
||||
from tests.test_litellm_rust.support.isolation import rebound
|
||||
|
|
@ -45,6 +47,14 @@ def redis_url() -> Generator[str]:
|
|||
worker.join(timeout=5)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def cluster_nodes() -> tuple[tuple[str, int], ...]:
|
||||
configured: Final = os.environ.get("LITELLM_TEST_REDIS_CLUSTER_NODES")
|
||||
if not configured:
|
||||
pytest.skip("LITELLM_TEST_REDIS_CLUSTER_NODES is not set")
|
||||
return tuple((host, int(port)) for host, _, port in (node.partition(":") for node in configured.split(",")))
|
||||
|
||||
|
||||
def test_existing_constructor_and_global_are_unchanged() -> None:
|
||||
facade: Final = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
assert type(facade.cache) is InMemoryCache
|
||||
|
|
@ -393,3 +403,64 @@ async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None:
|
|||
assert client.get("second") is not None
|
||||
await facade.cache.disconnect()
|
||||
client.close()
|
||||
|
||||
|
||||
async def test_redis_cluster_facade_serves_multi_slot_batches_and_scoped_flush_natively(
|
||||
cluster_nodes: tuple[tuple[str, int], ...],
|
||||
) -> None:
|
||||
startup_nodes: Final = [{"host": host, "port": port} for host, port in cluster_nodes]
|
||||
url: Final = f"redis://{cluster_nodes[0][0]}:{cluster_nodes[0][1]}"
|
||||
with rebound(litellm, "default_redis_ttl", 60):
|
||||
facade: Final = Cache(type=LiteLLMCacheType.REDIS, redis_startup_nodes=startup_nodes, namespace="parity")
|
||||
assert type(facade.cache) is RedisClusterCache
|
||||
with pytest.raises(TypeError, match="types must match"):
|
||||
_native._CacheTestHandle.redis(url, namespace="parity")._bind_facade(facade)
|
||||
_native._CacheTestHandle.redis(url, namespace="parity", startup_nodes=list(cluster_nodes))._bind_facade(facade)
|
||||
resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade))
|
||||
assert resolver.resolve().kind == "native"
|
||||
|
||||
manager: Final = facade.cache.redis_client.nodes_manager
|
||||
with rebound(manager, "connection_kwargs", {**manager.connection_kwargs, "db": 1}):
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "startup_nodes": startup_nodes[:1]}):
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
binding: Final = resolver.resolve()
|
||||
assert binding.kind == "native"
|
||||
|
||||
client: Final = redis.RedisCluster(startup_nodes=[redis.cluster.ClusterNode(*node) for node in cluster_nodes])
|
||||
keys: Final = tuple(f"slot-{index}" for index in range(12))
|
||||
slots: Final = {client.keyslot(f"parity:{key}") for key in keys}
|
||||
assert len(slots) > 1, slots
|
||||
requests: Final = [request(key) for key in keys]
|
||||
values: Final = [{"index": index} for index in range(len(keys))]
|
||||
await binding.async_store_batch(requests, values)
|
||||
client.set("parity:slot-3", "not a cache entry")
|
||||
client.set("parity:slot-7", json.dumps({"timestamp": time.time(), "response": {"index": 7, "python": True}}))
|
||||
|
||||
batch: Final = await binding.async_lookup_batch(requests)
|
||||
assert batch == {
|
||||
"values": [
|
||||
None if index == 3 else {"index": 7, "python": True} if index == 7 else value
|
||||
for index, value in enumerate(values)
|
||||
],
|
||||
"missing_indices": [3],
|
||||
}
|
||||
assert facade.cache.get_cache("parity:slot-0")["response"] == {"index": 0}
|
||||
assert (await facade.cache.async_get_cache("parity:slot-11"))["response"] == {"index": 11}
|
||||
assert facade.cache.redis_client.mget_nonatomic([f"parity:{key}" for key in keys[:2]]) == [
|
||||
client.get("parity:slot-0"),
|
||||
client.get("parity:slot-1"),
|
||||
]
|
||||
|
||||
await binding.async_store({**request("pinned"), "ttl_seconds": 12.0}, {"pinned": True})
|
||||
assert 0 < client.ttl("parity:pinned") <= 12
|
||||
client.set("unscoped", "stays")
|
||||
|
||||
await binding.async_flush()
|
||||
|
||||
remaining: Final = tuple(sorted(key for node in client.get_primaries() for key in client.keys("parity:*", target_nodes=node)))
|
||||
assert remaining == (), remaining
|
||||
assert client.get("unscoped") == b"stays"
|
||||
client.delete("unscoped")
|
||||
client.close()
|
||||
facade.cache.redis_client.close()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue