Fix(llm): complete response cache — set_model invalidation, stats logging, sync mutex (#290)

* fix(llm): complete response cache — set_model invalidation, stats logging, sync mutex

# Conflicts:
#	src/llm/response_cache.rs

* fix(llm): address response cache review comments

- Add total_hit_count AtomicU64 that is never decremented on eviction;
  maybe_log_stats now uses this counter so hit_rate_pct stays accurate
  under high eviction pressure
- Log cache stats before returning on provider error so milestone
  intervals (every 100 requests) are never silently skipped
- Add tracing-test dev-dep and three new tests: total_hits_survives_eviction,
  stats_logged_at_request_100, stats_logged_on_provider_error_at_interval
- Update PR description to reflect actual set_model() behavior (key
  isolation, not cache clear)

Co-Authored-By: Claude Sonnet 4.6 <[email protected]>

---------

Co-authored-by: Claude Sonnet 4.6 <[email protected]>
This commit is contained in:
Nitanshu Lokhande
2026-03-06 08:36:42 +00:00
committed by GitHub
co-authored by Claude Sonnet 4.6
parent 26d274ac79
commit 7806273aa6
3 changed files with 356 additions and 33 deletions
Generated
+22
View File
@@ -2902,6 +2902,7 @@ dependencies = [
"tower-http 0.6.8", "tower-http 0.6.8",
"tracing", "tracing",
"tracing-subscriber", "tracing-subscriber",
"tracing-test",
"url", "url",
"urlencoding", "urlencoding",
"uuid", "uuid",
@@ -6229,6 +6230,27 @@ dependencies = [
"tracing-serde", "tracing-serde",
] ]
[[package]]
name = "tracing-test"
version = "0.2.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "19a4c448db514d4f24c5ddb9f73f2ee71bfb24c526cf0c570ba142d1119e0051"
dependencies = [
"tracing-core",
"tracing-subscriber",
"tracing-test-macro",
]
[[package]]
name = "tracing-test-macro"
version = "0.2.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ad06847b7afb65c7866a36664b75c40b895e318cea4f71299f013fb22965329d"
dependencies = [
"quote",
"syn 2.0.114",
]
[[package]] [[package]]
name = "try-lock" name = "try-lock"
version = "0.2.5" version = "0.2.5"
+1
View File
@@ -174,6 +174,7 @@ zbus = "4"
[dev-dependencies] [dev-dependencies]
tokio-test = "0.4" tokio-test = "0.4"
tracing-test = "0.2"
tokio-tungstenite = "0.26" tokio-tungstenite = "0.26"
testcontainers-modules = { version = "0.11", features = ["postgres"] } testcontainers-modules = { version = "0.11", features = ["postgres"] }
pretty_assertions = "1" pretty_assertions = "1"
+333 -33
View File
@@ -16,13 +16,14 @@
//! ``` //! ```
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use async_trait::async_trait; use async_trait::async_trait;
use rust_decimal::Decimal; use rust_decimal::Decimal;
use sha2::{Digest, Sha256}; use sha2::{Digest, Sha256};
use tokio::sync::Mutex;
use crate::error::LlmError; use crate::error::LlmError;
use crate::llm::provider::{ use crate::llm::provider::{
@@ -30,6 +31,9 @@ use crate::llm::provider::{
ToolCompletionResponse, ToolCompletionResponse,
}; };
/// How often (in requests) to emit a cache statistics log line.
const STATS_LOG_EVERY_N: u64 = 100;
/// Configuration for the response cache. /// Configuration for the response cache.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct ResponseCacheConfig { pub struct ResponseCacheConfig {
@@ -61,8 +65,16 @@ struct CacheEntry {
/// tool calls can have side effects that should not be replayed. /// tool calls can have side effects that should not be replayed.
pub struct CachedProvider { pub struct CachedProvider {
inner: Arc<dyn LlmProvider>, inner: Arc<dyn LlmProvider>,
/// `std::sync::Mutex` (not tokio) — never held across an `.await` point,
/// so blocking acquisition is safe and keeps `set_model()` synchronous.
cache: Mutex<HashMap<String, CacheEntry>>, cache: Mutex<HashMap<String, CacheEntry>>,
config: ResponseCacheConfig, config: ResponseCacheConfig,
/// Total `complete()` calls (hits + misses) for periodic stats logging.
request_count: AtomicU64,
/// Running total of cache hits, independent of entry lifecycle.
/// Never decremented on eviction, so `hit_rate_pct` in stats doesn't
/// drift down as entries expire or are LRU-evicted.
total_hit_count: AtomicU64,
} }
impl CachedProvider { impl CachedProvider {
@@ -72,27 +84,53 @@ impl CachedProvider {
inner, inner,
cache: Mutex::new(HashMap::new()), cache: Mutex::new(HashMap::new()),
config, config,
request_count: AtomicU64::new(0),
total_hit_count: AtomicU64::new(0),
} }
} }
/// Number of entries currently in the cache. /// Number of entries currently in the cache.
pub async fn len(&self) -> usize { pub fn len(&self) -> usize {
self.cache.lock().await.len() self.cache.lock().unwrap_or_else(|e| e.into_inner()).len()
} }
/// Whether the cache is empty. /// Whether the cache is empty.
pub async fn is_empty(&self) -> bool { pub fn is_empty(&self) -> bool {
self.cache.lock().await.is_empty() self.cache
.lock()
.unwrap_or_else(|e| e.into_inner())
.is_empty()
} }
/// Total cache hits across all entries. /// Total cache hits since this provider was created.
pub async fn total_hits(&self) -> u64 { ///
self.cache.lock().await.values().map(|e| e.hit_count).sum() /// Backed by an atomic counter that is never decremented on eviction,
/// so the value is accurate even under high eviction pressure.
pub fn total_hits(&self) -> u64 {
self.total_hit_count.load(Ordering::Relaxed)
} }
/// Clear all cached entries. /// Clear all cached entries.
pub async fn clear(&self) { pub fn clear(&self) {
self.cache.lock().await.clear(); self.cache.lock().unwrap_or_else(|e| e.into_inner()).clear();
}
/// Emit a cache statistics log line if `req_no` is a multiple of
/// [`STATS_LOG_EVERY_N`]. `total_hits` must come from the `total_hit_count`
/// atomic so it accurately reflects hits that occurred on since-evicted
/// entries. Must be called while holding the cache lock so that
/// `entry_count` is consistent with the snapshot.
fn maybe_log_stats(guard: &HashMap<String, CacheEntry>, req_no: u64, total_hits: u64) {
if req_no.is_multiple_of(STATS_LOG_EVERY_N) {
let hit_rate = total_hits as f64 / req_no as f64 * 100.0;
tracing::info!(
total_requests = req_no,
total_hits,
hit_rate_pct = format!("{hit_rate:.1}"),
entry_count = guard.len(),
"LLM response cache statistics"
);
}
} }
} }
@@ -147,28 +185,47 @@ impl LlmProvider for CachedProvider {
let effective_model = self.inner.effective_model_name(request.model.as_deref()); let effective_model = self.inner.effective_model_name(request.model.as_deref());
let key = cache_key(&effective_model, &request); let key = cache_key(&effective_model, &request);
let now = Instant::now(); let now = Instant::now();
let req_no = self.request_count.fetch_add(1, Ordering::Relaxed) + 1;
// Check cache // Check cache — lock not held across the .await below.
{ {
let mut guard = self.cache.lock().await; let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner());
if let Some(entry) = guard.get_mut(&key) { if let Some(entry) = guard.get_mut(&key) {
if now.duration_since(entry.created_at) < self.config.ttl { if now.duration_since(entry.created_at) < self.config.ttl {
entry.last_accessed = now; entry.last_accessed = now;
entry.hit_count += 1; entry.hit_count += 1;
tracing::debug!(hits = entry.hit_count, "response cache hit"); let hit_count = entry.hit_count;
return Ok(entry.response.clone()); // Clone now so we can release the mutable borrow before stats.
let cached_response = entry.response.clone();
tracing::debug!(hits = hit_count, "response cache hit");
// Drop the mutable borrow of `entry` before reading `guard` immutably.
let _ = entry;
let total_hits = self.total_hit_count.fetch_add(1, Ordering::Relaxed) + 1;
Self::maybe_log_stats(&guard, req_no, total_hits);
return Ok(cached_response);
} }
// Expired, remove it // Expired, remove it
guard.remove(&key); guard.remove(&key);
} }
} }
// Cache miss, call the real provider // Cache miss call the real provider.
let response = self.inner.complete(request).await?; let result = self.inner.complete(request).await;
// Store in cache // Store result and maybe log stats, all within one lock acquisition.
// Stats are logged even on provider error so milestone intervals are
// not silently skipped.
{ {
let mut guard = self.cache.lock().await; let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner());
let total_hits = self.total_hit_count.load(Ordering::Relaxed);
let response = match result {
Err(e) => {
Self::maybe_log_stats(&guard, req_no, total_hits);
return Err(e);
}
Ok(r) => r,
};
// Evict expired entries // Evict expired entries
guard.retain(|_, entry| now.duration_since(entry.created_at) < self.config.ttl); guard.retain(|_, entry| now.duration_since(entry.created_at) < self.config.ttl);
@@ -196,9 +253,10 @@ impl LlmProvider for CachedProvider {
hit_count: 0, hit_count: 0,
}, },
); );
}
Ok(response) Self::maybe_log_stats(&guard, req_no, total_hits);
Ok(response)
}
} }
async fn complete_with_tools( async fn complete_with_tools(
@@ -226,16 +284,91 @@ impl LlmProvider for CachedProvider {
} }
fn set_model(&self, model: &str) -> Result<(), LlmError> { fn set_model(&self, model: &str) -> Result<(), LlmError> {
// Cache keys embed the active model name via `effective_model_name()`, so
// requests to the new model automatically land in a separate cache slot.
// Entries for the old model remain valid: if we switch back, they will be
// hit again rather than wasted. Natural TTL / LRU eviction cleans them up.
self.inner.set_model(model) self.inner.set_model(model)
} }
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use crate::llm::provider::ChatMessage; use std::sync::atomic::{AtomicU32, Ordering};
use rust_decimal::Decimal;
use tracing_test::traced_test;
use crate::error::LlmError;
use crate::llm::provider::{
ChatMessage, CompletionResponse, FinishReason, ToolCompletionRequest,
ToolCompletionResponse,
};
use crate::llm::response_cache::*; use crate::llm::response_cache::*;
use crate::testing::StubLlm; use crate::testing::StubLlm;
/// Minimal provider stub that supports `set_model()` — used to test
/// per-model cache key isolation.
struct SwitchableStub {
call_count: AtomicU32,
active_model: std::sync::RwLock<String>,
}
impl SwitchableStub {
fn new() -> Self {
Self {
call_count: AtomicU32::new(0),
active_model: std::sync::RwLock::new("stub-model".to_string()),
}
}
}
#[async_trait]
impl LlmProvider for SwitchableStub {
fn model_name(&self) -> &str {
"stub-model"
}
fn active_model_name(&self) -> String {
self.active_model.read().unwrap().clone()
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
(Decimal::ZERO, Decimal::ZERO)
}
fn set_model(&self, model: &str) -> Result<(), LlmError> {
*self.active_model.write().unwrap() = model.to_string();
Ok(())
}
async fn complete(
&self,
_request: CompletionRequest,
) -> Result<CompletionResponse, LlmError> {
self.call_count.fetch_add(1, Ordering::Relaxed);
Ok(CompletionResponse {
content: "ok".into(),
input_tokens: 1,
output_tokens: 1,
finish_reason: FinishReason::Stop,
})
}
async fn complete_with_tools(
&self,
_request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
Ok(ToolCompletionResponse {
content: Some("ok".into()),
tool_calls: vec![],
input_tokens: 1,
output_tokens: 1,
finish_reason: FinishReason::Stop,
})
}
}
fn simple_request() -> CompletionRequest { fn simple_request() -> CompletionRequest {
CompletionRequest { CompletionRequest {
messages: vec![ChatMessage::user("hello")], messages: vec![ChatMessage::user("hello")],
@@ -321,7 +454,7 @@ mod tests {
assert_eq!(stub.calls(), 1); // still 1 assert_eq!(stub.calls(), 1); // still 1
assert_eq!(r2.content, "cached response"); assert_eq!(r2.content, "cached response");
assert_eq!(cached.total_hits().await, 1); assert_eq!(cached.total_hits(), 1);
} }
#[tokio::test] #[tokio::test]
@@ -333,7 +466,7 @@ mod tests {
cached.complete(different_request()).await.unwrap(); cached.complete(different_request()).await.unwrap();
assert_eq!(stub.calls(), 2); assert_eq!(stub.calls(), 2);
assert_eq!(cached.len().await, 2); assert_eq!(cached.len(), 2);
} }
#[tokio::test] #[tokio::test]
@@ -372,7 +505,7 @@ mod tests {
// Fill cache with 2 entries // Fill cache with 2 entries
cached.complete(simple_request()).await.unwrap(); cached.complete(simple_request()).await.unwrap();
cached.complete(different_request()).await.unwrap(); cached.complete(different_request()).await.unwrap();
assert_eq!(cached.len().await, 2); assert_eq!(cached.len(), 2);
// Add a third: should evict the oldest // Add a third: should evict the oldest
let third = CompletionRequest { let third = CompletionRequest {
@@ -384,7 +517,7 @@ mod tests {
metadata: Default::default(), metadata: Default::default(),
}; };
cached.complete(third).await.unwrap(); cached.complete(third).await.unwrap();
assert_eq!(cached.len().await, 2); assert_eq!(cached.len(), 2);
assert_eq!(stub.calls(), 3); assert_eq!(stub.calls(), 3);
} }
@@ -408,7 +541,7 @@ mod tests {
// Both should have called through // Both should have called through
assert_eq!(stub.calls(), 2); assert_eq!(stub.calls(), 2);
assert!(cached.is_empty().await); assert!(cached.is_empty());
} }
#[tokio::test] #[tokio::test]
@@ -425,12 +558,12 @@ mod tests {
stub.set_failing(true); stub.set_failing(true);
let result = cached.complete(simple_request()).await; let result = cached.complete(simple_request()).await;
assert!(result.is_err()); assert!(result.is_err());
assert!(cached.is_empty().await); assert!(cached.is_empty());
// After fixing the provider, should succeed and cache // After fixing the provider, should succeed and cache
stub.set_failing(false); stub.set_failing(false);
cached.complete(simple_request()).await.unwrap(); cached.complete(simple_request()).await.unwrap();
assert_eq!(cached.len().await, 1); assert_eq!(cached.len(), 1);
} }
#[tokio::test] #[tokio::test]
@@ -439,10 +572,10 @@ mod tests {
let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default()); let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default());
cached.complete(simple_request()).await.unwrap(); cached.complete(simple_request()).await.unwrap();
assert_eq!(cached.len().await, 1); assert_eq!(cached.len(), 1);
cached.clear().await; cached.clear();
assert!(cached.is_empty().await); assert!(cached.is_empty());
} }
#[tokio::test] #[tokio::test]
@@ -459,7 +592,7 @@ mod tests {
cached.complete(req_b).await.unwrap(); cached.complete(req_b).await.unwrap();
assert_eq!(stub.calls(), 2); assert_eq!(stub.calls(), 2);
assert_eq!(cached.len().await, 2); assert_eq!(cached.len(), 2);
} }
#[test] #[test]
@@ -475,4 +608,171 @@ mod tests {
let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default()); let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default());
assert_eq!(cached.model_name(), "stub-model"); assert_eq!(cached.model_name(), "stub-model");
} }
/// Switching models preserves existing cached entries and routes subsequent
/// requests to a separate cache slot. Switching back replays the old slot.
#[tokio::test]
async fn set_model_isolates_per_model_via_key() {
let stub = Arc::new(SwitchableStub::new());
let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default());
// Populate cache under the initial model ("stub-model").
cached.complete(simple_request()).await.unwrap();
assert_eq!(stub.call_count.load(Ordering::Relaxed), 1);
assert_eq!(cached.len(), 1, "one entry cached for stub-model");
// Switch to a different model — old entries must survive.
cached.set_model("model-b").unwrap();
assert_eq!(cached.len(), 1, "old entries preserved after model switch");
// Same request under model-b is a cache miss (different key).
cached.complete(simple_request()).await.unwrap();
assert_eq!(
stub.call_count.load(Ordering::Relaxed),
2,
"cache miss for model-b"
);
assert_eq!(cached.len(), 2, "separate slots for stub-model and model-b");
// Switch back — original slot is still valid (cache hit, no extra call).
cached.set_model("stub-model").unwrap();
cached.complete(simple_request()).await.unwrap();
assert_eq!(
stub.call_count.load(Ordering::Relaxed),
2,
"cache hit when switching back to stub-model"
);
}
/// When `set_model()` fails the error is propagated and the cache is unaffected.
#[tokio::test]
async fn set_model_error_leaves_cache_intact() {
// StubLlm does not override set_model() — returns an error by default.
let stub = Arc::new(StubLlm::default());
let cached = CachedProvider::new(stub, ResponseCacheConfig::default());
cached.complete(simple_request()).await.unwrap();
assert_eq!(cached.len(), 1);
let result = cached.set_model("new-model");
assert!(result.is_err());
assert_eq!(cached.len(), 1, "cache unaffected by failed set_model");
}
/// `hit_rate_pct` stays accurate even after entries are evicted.
/// The `total_hit_count` atomic is never decremented on eviction.
#[tokio::test]
async fn total_hits_survives_eviction() {
let stub = Arc::new(StubLlm::new("response"));
// max_entries = 1 so the first entry is LRU-evicted when a second arrives.
let cached = CachedProvider::new(
stub.clone(),
ResponseCacheConfig {
ttl: Duration::from_secs(60),
max_entries: 1,
},
);
// Populate the cache and score a hit.
cached.complete(simple_request()).await.unwrap();
cached.complete(simple_request()).await.unwrap();
assert_eq!(cached.total_hits(), 1);
// Add a different request — LRU evicts the first entry.
cached.complete(different_request()).await.unwrap();
assert_eq!(cached.len(), 1, "first entry was evicted");
// The hit from the evicted entry must still be counted.
assert_eq!(cached.total_hits(), 1, "hit count survives eviction");
}
/// A stats line is emitted exactly at the 100th request.
#[tokio::test]
#[traced_test]
async fn stats_logged_at_request_100() {
let stub = Arc::new(StubLlm::new("response"));
let cached = CachedProvider::new(
stub.clone(),
ResponseCacheConfig {
ttl: Duration::from_secs(60),
max_entries: 2000,
},
);
// 99 distinct requests — no stats line yet.
for i in 0..99u32 {
let req = CompletionRequest {
messages: vec![ChatMessage::user(format!("request {i}"))],
model: None,
max_tokens: None,
temperature: None,
stop_sequences: None,
metadata: Default::default(),
};
cached.complete(req).await.unwrap();
}
assert!(
!logs_contain("LLM response cache statistics"),
"no stats before request 100"
);
// 100th request triggers the first stats line.
let req = CompletionRequest {
messages: vec![ChatMessage::user("request 99")],
model: None,
max_tokens: None,
temperature: None,
stop_sequences: None,
metadata: Default::default(),
};
cached.complete(req).await.unwrap();
assert!(
logs_contain("LLM response cache statistics"),
"stats emitted at request 100"
);
}
/// Stats are emitted even when the inner provider returns an error.
#[tokio::test]
#[traced_test]
async fn stats_logged_on_provider_error_at_interval() {
let stub = Arc::new(StubLlm::new("response"));
let cached = CachedProvider::new(
stub.clone(),
ResponseCacheConfig {
ttl: Duration::from_secs(60),
max_entries: 2000,
},
);
// 99 successful requests.
for i in 0..99u32 {
let req = CompletionRequest {
messages: vec![ChatMessage::user(format!("req {i}"))],
model: None,
max_tokens: None,
temperature: None,
stop_sequences: None,
metadata: Default::default(),
};
cached.complete(req).await.unwrap();
}
// 100th request fails — stats must still be logged.
stub.set_failing(true);
let req = CompletionRequest {
messages: vec![ChatMessage::user("req 99")],
model: None,
max_tokens: None,
temperature: None,
stop_sequences: None,
metadata: Default::default(),
};
let result = cached.complete(req).await;
assert!(result.is_err());
assert!(
logs_contain("LLM response cache statistics"),
"stats emitted even when provider errors on request 100"
);
}
} }