use std::time::Duration; use litellm_cache::CacheType; use litellm_cache_redis::{RedisNode, RedisTopology}; use pyo3::{ exceptions::{PyTypeError, PyValueError}, prelude::*, types::{PyAny, PyDict, PyList, PyString}, }; use super::{native::NativeResponseCache, request::duration}; #[allow(dead_code, reason = "consumed by the cache activation follow-up")] pub(super) struct CachePolicy { pub(super) mode: String, pub(super) ttl: Option, pub(super) namespace: Option, pub(super) supported_call_types: Option>, pub(super) redis_flush_size: Option, pub(super) semantic_cache_scope: String, } pub(super) struct MemoryCacheConfig { pub(super) default_ttl: Duration, pub(super) capacity: usize, pub(super) max_entry_bytes: usize, } #[derive(Debug, PartialEq)] pub(super) enum RedisProtocol { Resp2, Resp3, } #[derive(Debug, PartialEq)] pub(super) enum CertificateRequirement { None, Optional, Required, } #[allow(dead_code, reason = "consumed by the cache activation follow-up")] pub(super) struct RedisTlsConfig { pub(super) certificate_requirement: CertificateRequirement, pub(super) check_hostname: bool, pub(super) ca_certificate: Option, pub(super) ca_data: Option, pub(super) client_certificate: Option, pub(super) client_key: Option, } #[allow(dead_code, reason = "consumed by the cache activation follow-up")] pub(super) struct RedisConnectionConfig { pub(super) host: String, pub(super) port: u16, pub(super) database: i64, pub(super) username: Option, pub(super) password: Option, pub(super) protocol: RedisProtocol, pub(super) pool_size: usize, pub(super) read_timeout: Option, pub(super) connect_timeout: Option, pub(super) socket_keepalive: Option, pub(super) health_check_interval: Duration, pub(super) client_name: Option, pub(super) tls: Option, } #[allow(dead_code, reason = "consumed by the cache activation follow-up")] pub(super) struct RedisCacheConfig { pub(super) default_ttl: Duration, pub(super) namespace: Option, pub(super) flush_size: usize, pub(super) topology: RedisTopology, pub(super) connection: RedisConnectionConfig, } #[derive(Debug, PartialEq)] pub(super) struct GcsCacheConfig { pub(super) bucket_name: String, pub(super) key_prefix: String, pub(super) path_service_account: Option, } pub(super) struct AzureBlobCacheConfig { pub(super) account_url: String, pub(super) container: String, } struct RedisClientProjection<'py> { topology: RedisTopology, host: String, port: u16, pool_size: usize, resolved: Bound<'py, PyDict>, tls: Option, } const REDIS_PY_DEFAULT_MAX_CONNECTIONS: usize = 1 << 31; pub(super) enum CacheBackendConfig { Memory(MemoryCacheConfig), Redis(Box), Gcs(GcsCacheConfig), AzureBlob(AzureBlobCacheConfig), } #[allow(dead_code, reason = "consumed by the cache activation follow-up")] pub(super) struct NativeCacheConfig { pub(super) policy: CachePolicy, pub(super) backend: CacheBackendConfig, } pub(super) enum UnsupportedCacheConfig { Backend, RedisTopology, RedisCredentials, RedisConnection, RedisOption, GcsBucket, } impl UnsupportedCacheConfig { pub(super) fn message(&self) -> &'static str { match self { Self::Backend => "native cache backend is not implemented", Self::RedisTopology => "native Redis topology is not implemented", Self::RedisCredentials => "native Redis credentials require Python", Self::RedisConnection => "native Redis connection type is not implemented", Self::RedisOption => "native Redis configuration requires Python", Self::GcsBucket => "native GCS cache requires a configured bucket name", } } } pub(super) enum CacheConfigProjection { Native(Box), Unsupported(UnsupportedCacheConfig), } impl NativeCacheConfig { #[inline(never)] pub(super) fn project(facade: &Bound<'_, PyAny>) -> PyResult { let backend_name = facade.getattr("type")?.extract::()?; let policy = CachePolicy { mode: facade.getattr("mode")?.extract::()?, ttl: optional_duration(facade.getattr("ttl")?)?, namespace: optional_string(facade.getattr("namespace")?)?, supported_call_types: facade .getattr("supported_call_types")? .extract::>>()?, redis_flush_size: facade .getattr("redis_flush_size")? .extract::>()?, semantic_cache_scope: facade .getattr("semantic_cache_scope")? .extract::()?, }; let backend = facade.getattr("cache")?; match CacheType::from_python_name(&backend_name) { Some(CacheType::Local) => project_memory(&backend).map(|backend| { CacheConfigProjection::Native(Box::new(Self { policy, backend: CacheBackendConfig::Memory(backend), })) }), Some(CacheType::Redis) => match project_redis(&backend)? { Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { policy, backend: CacheBackendConfig::Redis(Box::new(backend)), }))), Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), }, Some(CacheType::Gcs) => match project_gcs(&backend)? { Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { policy, backend: CacheBackendConfig::Gcs(backend), }))), Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), }, Some(CacheType::AzureBlob) => project_azure_blob(&backend).map(|backend| { CacheConfigProjection::Native(Box::new(Self { policy, backend: CacheBackendConfig::AzureBlob(backend), })) }), Some( CacheType::RedisSemantic | CacheType::ValkeySemantic | CacheType::S3 | CacheType::Disk | CacheType::QdrantSemantic, ) | None => Ok(CacheConfigProjection::Unsupported( UnsupportedCacheConfig::Backend, )), } } pub(super) fn service_mismatch(&self, service: &NativeResponseCache) -> Option<&'static str> { let default_ttl = match &self.backend { CacheBackendConfig::Memory(config) => Some(config.default_ttl), CacheBackendConfig::Redis(config) => Some(config.default_ttl), CacheBackendConfig::AzureBlob(_) | CacheBackendConfig::Gcs(_) => None, }; if service.default_ttl() != default_ttl { return Some("facade and native backend default TTLs must match"); } match &self.backend { CacheBackendConfig::Memory(config) if service.kind() != "memory" => { Some("facade and native backend types must match") } CacheBackendConfig::Memory(config) if service.capacity() != Some(config.capacity) => { Some("facade and native backend capacities must match") } CacheBackendConfig::Memory(config) if service.max_entry_bytes() != Some(config.max_entry_bytes) => { Some("facade and native backend item limits must match") } CacheBackendConfig::Memory(_) => None, 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"), CacheBackendConfig::Gcs(_) if service.kind() != "gcs" => { Some("facade and native backend types must match") } CacheBackendConfig::Gcs(config) if service .gcs_backend() .is_none_or(|backend| backend.bucket_name() != config.bucket_name) => { Some("facade and native backend buckets must match") } CacheBackendConfig::Gcs(config) if service .gcs_backend() .is_none_or(|backend| backend.key_prefix() != config.key_prefix) => { Some("facade and native backend key prefixes must match") } CacheBackendConfig::Gcs(config) if service.gcs_backend().is_none_or(|backend| { backend.path_service_account() != config.path_service_account.as_deref() }) => { Some("facade and native backend credentials must match") } CacheBackendConfig::Gcs(_) => None, CacheBackendConfig::AzureBlob(config) => match service.azure_blob_identity() { None => Some("facade and native backend types must match"), Some((account_url, container)) if account_url != config.account_url || container != config.container => { Some("facade and native backend containers must match") } Some(_) => None, }, } } } #[inline(never)] fn project_azure_blob(backend: &Bound<'_, PyAny>) -> PyResult { let client = backend.getattr("container_client")?; let container = client.getattr("container_name")?.extract::()?; let url = client.getattr("url")?.extract::()?; let account_url = url .strip_suffix(container.as_str()) .and_then(|url| url.strip_suffix('/')) .ok_or_else(|| PyValueError::new_err("Azure Blob container URL is malformed"))?; Ok(AzureBlobCacheConfig { account_url: account_url.to_string(), container, }) } #[inline(never)] fn project_memory(backend: &Bound<'_, PyAny>) -> PyResult { let max_size_kib = backend.getattr("max_size_per_item")?.extract::()?; Ok(MemoryCacheConfig { default_ttl: duration(backend.getattr("default_ttl")?.extract::()?)?, capacity: backend.getattr("max_size_in_memory")?.extract::()?, max_entry_bytes: max_size_kib .checked_mul(1024) .ok_or_else(|| PyValueError::new_err("memory cache item limit is too large"))?, }) } #[inline(never)] fn project_gcs( backend: &Bound<'_, PyAny>, ) -> PyResult> { let bucket_name = match backend.getattr("bucket_name")?.extract::>() { Ok(Some(bucket_name)) if !bucket_name.is_empty() => bucket_name, _ => return Ok(Err(UnsupportedCacheConfig::GcsBucket)), }; Ok(Ok(GcsCacheConfig { bucket_name, key_prefix: backend.getattr("key_prefix")?.extract::()?, path_service_account: backend .getattr("path_service_account")? .extract::>()?, })) } #[inline(never)] fn project_redis( backend: &Bound<'_, PyAny>, ) -> PyResult> { let source = backend.getattr("redis_kwargs")?.cast_into::()?; if has_value(&source, "sentinel_nodes")? { return Ok(Err(UnsupportedCacheConfig::RedisTopology)); } for key in ["credential_provider", "redis_connect_func"] { if has_value(&source, key)? { return Ok(Err(UnsupportedCacheConfig::RedisCredentials)); } } if has_value(&source, "connection_pool")? { return Ok(Err(UnsupportedCacheConfig::RedisConnection)); } for key in [ "retry", "retry_on_error", "socket_keepalive_options", "unix_socket_path", "cache", "cache_config", "event_dispatcher", "ssl_ca_path", "ssl_password", "ssl_min_version", "ssl_ciphers", "ssl_validate_ocsp", "ssl_validate_ocsp_stapled", "ssl_ocsp_context", "ssl_ocsp_expected_cert", ] { if has_value(&source, key)? { return Ok(Err(UnsupportedCacheConfig::RedisOption)); } } for key in ["retry_on_timeout", "single_connection_client"] { if optional_coerced_bool(&source, key)?.unwrap_or(false) { return Ok(Err(UnsupportedCacheConfig::RedisOption)); } } let client = backend.getattr("redis_client")?; let projection = if has_value(&source, "startup_nodes")? { project_cluster_client(&source, &client)? } else { 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, 3 => RedisProtocol::Resp3, _ => return Err(PyValueError::new_err("unsupported Redis protocol version")), }; let health_check_interval = duration(optional_f64(&resolved, "health_check_interval")?.unwrap_or(0.0))?; Ok(Ok(RedisCacheConfig { default_ttl: duration(backend.getattr("default_ttl")?.extract::()?)?, namespace: optional_attribute_string(backend, "namespace")?, flush_size: backend.getattr("redis_flush_size")?.extract::()?, topology, connection: RedisConnectionConfig { 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, read_timeout: optional_dict_duration(&resolved, "socket_timeout")?, connect_timeout: optional_dict_duration(&resolved, "socket_connect_timeout")?, socket_keepalive: optional_bool(&resolved, "socket_keepalive")?, health_check_interval, client_name: optional_dict_string(&resolved, "client_name")?, tls, }, })) } #[inline(never)] fn project_standalone_client<'py>( client: &Bound<'py, PyAny>, ) -> PyResult, 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::()?; 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::()?, resolved, tls, })) } #[inline(never)] fn project_cluster_client<'py>( source: &Bound<'py, PyDict>, client: &Bound<'py, PyAny>, ) -> PyResult, 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::()?; 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>> { let Some(nodes) = source.get_item("startup_nodes")? else { return Ok(None); }; let Ok(nodes) = nodes.cast_into::() 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::() 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::try_from(value).map_err(|_| PyValueError::new_err("invalid Redis port")) } #[inline(never)] fn project_tls(values: &Bound<'_, PyDict>) -> PyResult { Ok(RedisTlsConfig { certificate_requirement: certificate_requirement(values)?, check_hostname: optional_bool(values, "ssl_check_hostname")?.unwrap_or(false), ca_certificate: optional_dict_string(values, "ssl_ca_certs")?, ca_data: optional_dict_string(values, "ssl_ca_data")?, client_certificate: optional_dict_string(values, "ssl_certfile")?, client_key: optional_dict_string(values, "ssl_keyfile")?, }) } #[inline(never)] fn certificate_requirement(values: &Bound<'_, PyDict>) -> PyResult { let Some(value) = values.get_item("ssl_cert_reqs")? else { return Ok(CertificateRequirement::Required); }; if value.is_none() { return Ok(CertificateRequirement::Required); } if let Ok(number) = value.extract::() { return match number { 0 => Ok(CertificateRequirement::None), 1 => Ok(CertificateRequirement::Optional), 2 => Ok(CertificateRequirement::Required), _ => Err(PyValueError::new_err( "invalid Redis TLS certificate requirement", )), }; } let text = value.str()?; let text = text.to_str()?; if text.eq_ignore_ascii_case("none") || text.eq_ignore_ascii_case("cert_none") { return Ok(CertificateRequirement::None); } if text.eq_ignore_ascii_case("optional") || text.eq_ignore_ascii_case("cert_optional") { return Ok(CertificateRequirement::Optional); } if text.eq_ignore_ascii_case("required") || text.eq_ignore_ascii_case("cert_required") { return Ok(CertificateRequirement::Required); } Err(PyValueError::new_err( "invalid Redis TLS certificate requirement", )) } #[inline(never)] fn instance_class_is(value: &Bound<'_, PyAny>, module: &str, name: &str) -> PyResult { class_is(value.get_type().as_any(), module, name) } #[inline(never)] fn class_is(value: &Bound<'_, PyAny>, module: &str, name: &str) -> PyResult { Ok(value .getattr("__module__")? .cast_into::()? .to_str()? == module && value .getattr("__qualname__")? .cast_into::()? .to_str()? == name) } #[inline(never)] fn optional_duration(value: Bound<'_, PyAny>) -> PyResult> { value.extract::>()?.map(duration).transpose() } #[inline(never)] fn optional_attribute_string(value: &Bound<'_, PyAny>, name: &str) -> PyResult> { match value.getattr(name) { Ok(value) => optional_string(value), Err(error) if error.is_instance_of::(value.py()) => { Ok(None) } Err(error) => Err(error), } } #[inline(never)] fn optional_string(value: Bound<'_, PyAny>) -> PyResult> { Ok(value .extract::>()? .filter(|value| !value.is_empty())) } #[inline(never)] fn has_value(values: &Bound<'_, PyDict>, key: &str) -> PyResult { Ok(values.get_item(key)?.is_some_and(|value| !value.is_none())) } #[inline(never)] fn required_string(values: &Bound<'_, PyDict>, key: &str) -> PyResult { values .get_item(key)? .ok_or_else(|| PyTypeError::new_err("Redis connection is incomplete"))? .extract::() } #[inline(never)] fn required_i64(values: &Bound<'_, PyDict>, key: &str) -> PyResult { values .get_item(key)? .ok_or_else(|| PyTypeError::new_err("Redis connection is incomplete"))? .extract::() } #[inline(never)] fn optional_dict_string(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { match values.get_item(key)? { Some(value) if !value.is_none() => optional_string(value), _ => Ok(None), } } #[inline(never)] fn optional_f64(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { match values.get_item(key)? { Some(value) => value.extract::>(), None => Ok(None), } } #[inline(never)] fn optional_i64(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { match values.get_item(key)? { Some(value) => value.extract::>(), None => Ok(None), } } #[inline(never)] fn optional_bool(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { match values.get_item(key)? { Some(value) => value.extract::>(), None => Ok(None), } } #[inline(never)] fn optional_coerced_bool(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { let Some(value) = values.get_item(key)? else { return Ok(None); }; if value.is_none() { return Ok(None); } if let Ok(text) = value.extract::() { return Ok(Some( text == "1" || text.eq_ignore_ascii_case("true") || text.eq_ignore_ascii_case("yes"), )); } value.extract::().map(Some) } #[inline(never)] fn optional_dict_duration(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { optional_f64(values, key)?.map(duration).transpose() } #[cfg(test)] mod tests { use std::ffi::CString; use pyo3::{prelude::*, types::PyDict}; use litellm_cache_redis::{RedisNode, RedisTopology}; use super::{ CacheBackendConfig, CacheConfigProjection, CertificateRequirement, GcsCacheConfig, NativeCacheConfig, RedisProtocol, UnsupportedCacheConfig, }; 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( &CString::new(format!( "from types import SimpleNamespace\n\ ConnectionPool = type('ConnectionPool', (), {{'__module__': 'redis.connection'}})\n\ Connection = type('Connection', (), {{'__module__': 'redis.connection'}})\n\ SSLConnection = type('SSLConnection', (), {{'__module__': 'redis.connection'}})\n\ {body}" )) .unwrap(), None, Some(&locals), ) .unwrap(); locals.get_item("facade").unwrap().unwrap() } #[test] fn projects_effective_memory_configuration() { Python::initialize(); Python::attach(|py| { let facade = facade( py, "backend = SimpleNamespace(default_ttl=913, max_size_in_memory=37, max_size_per_item=8)\n\ facade = SimpleNamespace(type='local', mode='default-on', ttl=11.5, namespace=None, supported_call_types=['completion'], redis_flush_size=None, semantic_cache_scope='key', cache=backend)", ); let CacheConfigProjection::Native(config) = NativeCacheConfig::project(&facade).unwrap() else { panic!("memory cache should be supported"); }; assert_eq!( config.policy.ttl.unwrap(), std::time::Duration::from_secs_f64(11.5) ); let CacheBackendConfig::Memory(memory) = config.backend else { panic!("expected memory configuration"); }; assert_eq!(memory.default_ttl, std::time::Duration::from_secs(913)); assert_eq!(memory.capacity, 37); assert_eq!(memory.max_entry_bytes, 8192); let matching = NativeResponseCache::memory(37, std::time::Duration::from_secs(913), 8192); let mismatched = NativeResponseCache::memory(37, std::time::Duration::from_secs(913), 8191); let matching_config = NativeCacheConfig { policy: config.policy, backend: CacheBackendConfig::Memory(memory), }; assert_eq!(matching_config.service_mismatch(&matching), None); assert_eq!( matching_config.service_mismatch(&mismatched), Some("facade and native backend item limits must match") ); }); } #[test] fn projects_gcs_configuration() { Python::initialize(); Python::attach(|py| { let facade = facade( py, "backend = SimpleNamespace(bucket_name='bucket', key_prefix='cache/', path_service_account='credentials.json')\n\ facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", ); let CacheConfigProjection::Native(config) = NativeCacheConfig::project(&facade).unwrap() else { panic!("GCS cache should be supported"); }; let CacheBackendConfig::Gcs(gcs) = config.backend else { panic!("expected GCS configuration"); }; assert_eq!( gcs, GcsCacheConfig { bucket_name: "bucket".into(), key_prefix: "cache/".into(), path_service_account: Some("credentials.json".into()), } ); let matching = NativeResponseCache::gcs( litellm_cache_gcs::GcsConfig { bucket_name: "bucket".into(), gcs_path: Some("cache/".into()), path_service_account: Some("credentials.json".into()), endpoint: litellm_cache_gcs::DEFAULT_ENDPOINT.into(), }, Some("token".into()), ) .unwrap(); let matching_config = NativeCacheConfig { policy: config.policy, backend: CacheBackendConfig::Gcs(gcs), }; assert_eq!(matching_config.service_mismatch(&matching), None); }); } #[test] fn rejects_gcs_without_a_bucket_name() { Python::initialize(); Python::attach(|py| { let facade = facade( py, "backend = SimpleNamespace(bucket_name=None, key_prefix='', path_service_account=None)\n\ facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", ); let CacheConfigProjection::Unsupported(reason) = NativeCacheConfig::project(&facade).unwrap() else { panic!("GCS cache without a bucket should be unsupported"); }; assert!(matches!(&reason, UnsupportedCacheConfig::GcsBucket)); assert_eq!( reason.message(), "native GCS cache requires a configured bucket name" ); }); } #[test] fn projects_resolved_redis_tls_configuration() { Python::initialize(); Python::attach(|py| { let facade = facade( py, "pool = ConnectionPool()\n\ pool.connection_class = SSLConnection\n\ pool.max_connections = 29\n\ pool.connection_kwargs = {'host': 'cache.internal', 'port': 6380, 'db': 4, 'username': 'user', 'password': 'secret', 'protocol': 3, 'socket_timeout': 7.5, 'socket_connect_timeout': 2, 'socket_keepalive': True, 'health_check_interval': 15, 'client_name': 'litellm', 'ssl_cert_reqs': 'optional', 'ssl_check_hostname': True, 'ssl_ca_certs': '/ca.pem', 'ssl_ca_data': 'CA DATA', 'ssl_certfile': '/client.pem', 'ssl_keyfile': '/client.key'}\n\ client = SimpleNamespace(connection_pool=pool)\n\ backend = SimpleNamespace(default_ttl=777, namespace='team', redis_flush_size=31, redis_kwargs={}, redis_client=client)\n\ facade = SimpleNamespace(type='redis', mode='default-off', ttl=None, namespace='team', supported_call_types=None, redis_flush_size=31, semantic_cache_scope='key', cache=backend)", ); let CacheConfigProjection::Native(config) = NativeCacheConfig::project(&facade).unwrap() else { panic!("Redis cache should be supported"); }; let CacheBackendConfig::Redis(redis) = config.backend else { panic!("expected Redis configuration"); }; assert_eq!(redis.default_ttl, std::time::Duration::from_secs(777)); assert_eq!(redis.namespace.as_deref(), Some("team")); assert_eq!(redis.flush_size, 31); assert_eq!(redis.connection.host, "cache.internal"); assert_eq!(redis.connection.port, 6380); assert_eq!(redis.connection.database, 4); assert_eq!(redis.connection.protocol, RedisProtocol::Resp3); assert_eq!(redis.connection.pool_size, 29); let tls = redis.connection.tls.unwrap(); assert_eq!( tls.certificate_requirement, CertificateRequirement::Optional ); assert!(tls.check_hostname); assert_eq!(tls.ca_certificate.as_deref(), Some("/ca.pem")); assert_eq!(tls.ca_data.as_deref(), Some("CA DATA")); assert_eq!(tls.client_certificate.as_deref(), Some("/client.pem")); assert_eq!(tls.client_key.as_deref(), Some("/client.key")); }); } #[test] fn dynamic_redis_auth_stays_on_python() { Python::initialize(); Python::attach(|py| { let facade = facade( py, "backend = SimpleNamespace(redis_kwargs={'credential_provider': object()})\n\ facade = SimpleNamespace(type='redis', mode='default-on', ttl=None, namespace=None, supported_call_types=[], redis_flush_size=None, semantic_cache_scope='key', cache=backend)", ); let CacheConfigProjection::Unsupported(reason) = NativeCacheConfig::project(&facade).unwrap() else { panic!("dynamic authentication must stay on Python"); }; 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}"); } }); } }