mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
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:
co-authored by
Claude Sonnet 4.6
parent
26d274ac79
commit
7806273aa6
Generated
+22
@@ -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"
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
+332
-32
@@ -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,10 +253,11 @@ impl LlmProvider for CachedProvider {
|
|||||||
hit_count: 0,
|
hit_count: 0,
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
}
|
|
||||||
|
|
||||||
|
Self::maybe_log_stats(&guard, req_no, total_hits);
|
||||||
Ok(response)
|
Ok(response)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
async fn complete_with_tools(
|
async fn complete_with_tools(
|
||||||
&self,
|
&self,
|
||||||
@@ -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"
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user