//! Enhanced inference endpoints with comprehensive context caching use axum::Json; use std::sync::Arc; use std::time::Instant; use tokio::sync::OnceCell; use crate::cache::{ BeamSearchConfig, CacheConfig, CacheManager, CacheValue, WindowConfig, config::ContextCacheConfig, manager::PersistenceConfig, metrics::{AlertThresholds, CacheMetrics, CacheOperationType, OperationMetrics}, speculative::TokenLogit, }; use crate::{ApiError, ApiResult}; /// Enhanced inference request with caching parameters #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct CachedInferenceRequest { /// Model ID to use for inference pub model: String, /// Input prompt pub prompt: String, /// Maximum tokens to generate pub max_tokens: Option, /// Temperature for sampling pub temperature: Option, /// Whether to stream the response pub stream: Option, /// Enable context caching pub enable_caching: Option, /// Enable speculative decoding pub enable_speculation: Option, /// Enable sliding window attention pub enable_sliding_window: Option, /// Cache key override pub cache_key: Option, /// Request ID for tracking pub request_id: Option, } /// Enhanced inference response with caching metadata #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct CachedInferenceResponse { /// Generated text pub text: String, /// Finish reason pub finish_reason: String, /// Usage statistics pub usage: TokenUsage, /// Cache statistics for this request pub cache_stats: CacheRequestStats, /// Performance metrics pub performance: PerformanceMetrics, /// Request ID pub request_id: String, } /// Token usage statistics #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct TokenUsage { /// Prompt tokens pub prompt_tokens: u32, /// Completion tokens pub completion_tokens: u32, /// Total tokens pub total_tokens: u32, /// Cached tokens (reused from cache) pub cached_tokens: u32, } /// Cache-related statistics for a request #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct CacheRequestStats { /// Whether cache was used pub cache_used: bool, /// Cache hit rate for this request pub hit_rate: f64, /// Number of cache hits pub cache_hits: u32, /// Number of cache misses pub cache_misses: u32, /// Speculative decoding hits pub speculation_hits: u32, /// Radix tree prefix matches pub prefix_matches: u32, /// Sliding window cache reuse pub window_reuse: u32, } /// Performance metrics for a request #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct PerformanceMetrics { /// Total inference time in milliseconds pub total_time_ms: f64, /// Time spent on caching operations in milliseconds pub cache_time_ms: f64, /// Time spent on actual inference in milliseconds pub inference_time_ms: f64, /// Tokens per second pub tokens_per_second: f64, /// Memory usage in bytes pub memory_usage_bytes: usize, } /// Global cache manager instance static CACHE_MANAGER: OnceCell> = OnceCell::const_new(); static CACHE_METRICS: OnceCell> = OnceCell::const_new(); /// Initialize the cache system pub async fn initialize_cache_system(config: ContextCacheConfig) -> ApiResult<()> { let cache_config = CacheConfig { l1_capacity: config.kv_cache.l1_capacity, l2_capacity: config.kv_cache.l2_capacity, l3_capacity_bytes: config.kv_cache.l3_capacity_bytes, max_memory_bytes: config.kv_cache.max_memory_bytes, eviction_policy: config.kv_cache.eviction_policy, ttl_seconds: config.kv_cache.ttl_seconds, enable_compression: config.kv_cache.enable_compression, cleanup_interval_seconds: config.kv_cache.cleanup_interval_seconds, }; let beam_config = BeamSearchConfig { beam_width: config.speculative.beam_width, max_length: config.speculative.max_length, early_stopping_threshold: config.speculative.early_stopping_threshold, length_penalty: config.speculative.length_penalty, temperature: config.speculative.temperature, top_k: config.speculative.top_k, top_p: config.speculative.top_p, repetition_penalty: config.speculative.repetition_penalty, }; let window_config = WindowConfig { window_size: config.sliding_window.window_size, window_overlap: config.sliding_window.window_overlap, max_cached_windows: config.sliding_window.max_cached_windows, use_flash_attention: config.sliding_window.use_flash_attention, attention_block_size: config.sliding_window.attention_block_size, gradient_checkpointing: config.sliding_window.gradient_checkpointing, }; let persistence_config = PersistenceConfig { cache_dir: config.persistence.cache_dir, auto_persist: config.persistence.enabled, persist_interval_seconds: config.persistence.persist_interval_seconds, enable_compression: config.persistence.compression.enabled, max_file_size_bytes: config.persistence.max_file_size_bytes, backup_count: config.persistence.backup_count, }; let cache_manager = CacheManager::new(cache_config, persistence_config, beam_config, window_config).await?; let alert_thresholds = AlertThresholds { min_hit_rate: config.monitoring.alerts.hit_rate_threshold, max_latency_us: config.monitoring.alerts.latency_threshold_us as f64, max_memory_bytes: config.monitoring.alerts.memory_threshold_percent as usize * 1024 * 1024, // Convert to bytes max_error_rate: config.monitoring.alerts.error_rate_threshold, }; let cache_metrics = CacheMetrics::new(10000, alert_thresholds) .map_err(|e| ApiError::internal(format!("Failed to initialize metrics: {e}")))?; if config.monitoring.enabled { cache_metrics .start_background_collection(config.monitoring.collection_interval_seconds) .await; } CACHE_MANAGER .set(Arc::new(cache_manager)) .map_err(|_| ApiError::internal("Cache manager already initialized"))?; CACHE_METRICS .set(Arc::new(cache_metrics)) .map_err(|_| ApiError::internal("Cache metrics already initialized"))?; Ok(()) } /// Get the cache manager instance fn get_cache_manager() -> ApiResult<&'static Arc> { CACHE_MANAGER .get() .ok_or_else(|| ApiError::internal("Cache system not initialized")) } /// Get the cache metrics instance fn get_cache_metrics() -> ApiResult<&'static Arc> { CACHE_METRICS .get() .ok_or_else(|| ApiError::internal("Cache metrics not initialized")) } /// Enhanced inference handler with context caching pub async fn cached_inference( Json(request): Json, ) -> ApiResult> { let start_time = Instant::now(); let request_id = request .request_id .clone() .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); let cache_manager = get_cache_manager()?; let cache_metrics = get_cache_metrics()?; // Initialize request statistics let mut cache_stats = CacheRequestStats { cache_used: request.enable_caching.unwrap_or(true), hit_rate: 0.0, cache_hits: 0, cache_misses: 0, speculation_hits: 0, prefix_matches: 0, window_reuse: 0, }; let cache_key = request .cache_key .clone() .unwrap_or_else(|| generate_cache_key(&request.model, &request.prompt)); // Try cache lookup first if caching is enabled let mut cached_response = None; if cache_stats.cache_used { let cache_start = Instant::now(); if let Some(cached_value) = cache_manager.get(&cache_key).await { cache_stats.cache_hits += 1; cached_response = Some(reconstruct_response_from_cache(&cached_value, &request_id)); // Record cache hit metrics cache_metrics.record_operation(OperationMetrics { operation: CacheOperationType::Get, duration: cache_start.elapsed(), success: true, hit: Some(true), data_size: Some(cached_value.size_bytes), timestamp: cache_start, labels: std::collections::HashMap::new(), }); } else { cache_stats.cache_misses += 1; // Record cache miss metrics cache_metrics.record_operation(OperationMetrics { operation: CacheOperationType::Get, duration: cache_start.elapsed(), success: true, hit: Some(false), data_size: None, timestamp: cache_start, labels: std::collections::HashMap::new(), }); } } let inference_start = Instant::now(); let mut response_text = String::new(); let mut finish_reason = "stop".to_string(); // If we have a cached response, use it let is_cached = cached_response.is_some(); if let Some(cached_resp) = cached_response { response_text = cached_resp.text; finish_reason = cached_resp.finish_reason; } else { // Perform inference with optional enhancements let tokens = tokenize_prompt(&request.prompt)?; if request.enable_speculation.unwrap_or(false) { // Use speculative decoding let generated_tokens = cache_manager .speculative_decode( tokens.clone(), request.max_tokens.unwrap_or(100) as usize, create_mock_logits_function(request.temperature.unwrap_or(1.0)), ) .await; cache_stats.speculation_hits += (generated_tokens.len().saturating_sub(tokens.len())) as u32; response_text = detokenize(&generated_tokens); } else if request.enable_sliding_window.unwrap_or(false) { // Use sliding window attention let sequence = tokens_to_sequence(&tokens); let processed = cache_manager .sliding_window_attention(&sequence, 64, 8) .await; cache_stats.window_reuse += estimate_window_reuse(&processed); response_text = sequence_to_text(&processed); } else { // Standard inference with potential prefix matching response_text = perform_standard_inference(&request, &mut cache_stats).await?; } // Cache the result if caching is enabled if cache_stats.cache_used { let cache_value = create_cache_value(&response_text, &tokens); if let Err(e) = cache_manager.put(cache_key, cache_value).await { tracing::warn!("Failed to cache inference result: {}", e); } } } let inference_time = inference_start.elapsed(); let total_time = start_time.elapsed(); // Calculate token counts let prompt_tokens = count_tokens(&request.prompt); let completion_tokens = count_tokens(&response_text); let cached_tokens = if is_cached { completion_tokens } else { 0 }; let usage = TokenUsage { prompt_tokens, completion_tokens, total_tokens: prompt_tokens + completion_tokens, cached_tokens, }; // Calculate hit rate cache_stats.hit_rate = if cache_stats.cache_hits + cache_stats.cache_misses > 0 { f64::from(cache_stats.cache_hits) / f64::from(cache_stats.cache_hits + cache_stats.cache_misses) * 100.0 } else { 0.0 }; // Performance metrics let cache_time = total_time.saturating_sub(inference_time); let performance = PerformanceMetrics { total_time_ms: total_time.as_secs_f64() * 1000.0, cache_time_ms: cache_time.as_secs_f64() * 1000.0, inference_time_ms: inference_time.as_secs_f64() * 1000.0, tokens_per_second: if total_time.as_secs_f64() > 0.0 { f64::from(completion_tokens) / total_time.as_secs_f64() } else { 0.0 }, memory_usage_bytes: estimate_memory_usage(&response_text), }; // Update cache manager statistics let manager_stats = cache_manager.stats().await; cache_metrics.update_cache_stats(&manager_stats); let response = CachedInferenceResponse { text: response_text, finish_reason, usage, cache_stats, performance, request_id, }; Ok(Json(response)) } /// Get cache statistics endpoint pub async fn cache_stats() -> ApiResult> { let cache_manager = get_cache_manager()?; let cache_metrics = get_cache_metrics()?; let manager_stats = cache_manager.stats().await; let efficiency_score = cache_metrics.calculate_efficiency_score(&manager_stats); let alerts = cache_metrics.check_alerts(&manager_stats); let trends = cache_metrics.get_performance_trends(); let stats = serde_json::json!({ "cache_manager": { "kv_cache": { "hits": manager_stats.kv_cache.hits, "misses": manager_stats.kv_cache.misses, "hit_rate": manager_stats.kv_cache.hit_rate(), "memory_bytes": manager_stats.kv_cache.memory_bytes, "evictions": manager_stats.kv_cache.evictions, "l1_hits": manager_stats.kv_cache.l1_hits, "l2_hits": manager_stats.kv_cache.l2_hits, "l3_hits": manager_stats.kv_cache.l3_hits, }, "radix_tree": { "prefix_matches": manager_stats.radix_tree.prefix_matches, "sharing_hits": manager_stats.radix_tree.sharing_hits, "memory_saved": manager_stats.radix_tree.memory_saved, "active_prefixes": manager_stats.radix_tree.active_prefixes, "avg_prefix_length": manager_stats.radix_tree.avg_prefix_length, }, "speculation": { "hits": manager_stats.speculation.speculation_hits, "misses": manager_stats.speculation.speculation_misses, "hit_rate": manager_stats.speculation.hit_rate(), "tokens_generated": manager_stats.speculation.speculative_tokens, "avg_depth": manager_stats.speculation.avg_speculation_depth, "early_stops": manager_stats.speculation.early_stops, }, "sliding_window": { "hits": manager_stats.sliding_window.window_hits, "misses": manager_stats.sliding_window.window_misses, "hit_rate": manager_stats.sliding_window.hit_rate(), "memory_bytes": manager_stats.sliding_window.total_memory_bytes, "evictions": manager_stats.sliding_window.windows_evicted, "utilization": manager_stats.sliding_window.avg_window_utilization, } }, "overall": { "hit_rate": manager_stats.overall_hit_rate, "total_memory_bytes": manager_stats.total_memory_bytes, "operations_per_second": manager_stats.operations_per_second, "avg_response_time_us": manager_stats.avg_response_time_us, "efficiency_score": efficiency_score, "background_tasks_active": manager_stats.background_tasks_active, }, "alerts": alerts.iter().map(|alert| { serde_json::json!({ "severity": format!("{:?}", alert.severity), "message": alert.message, "metric": alert.metric, "current_value": alert.current_value, "threshold_value": alert.threshold_value, }) }).collect::>(), "trends": { "window_seconds": trends.window.as_secs(), "data_points": trends.hit_rate_trend.len(), } }); Ok(Json(stats)) } /// Clear all caches endpoint pub async fn clear_cache() -> ApiResult> { let cache_manager = get_cache_manager()?; cache_manager.clear_all().await?; Ok(Json(serde_json::json!({ "status": "success", "message": "All caches cleared successfully" }))) } /// Warm cache endpoint pub async fn warm_cache( Json(_request): Json, ) -> ApiResult> { let cache_manager = get_cache_manager()?; let strategy = crate::cache::manager::WarmingStrategy::FrequentPatterns; // Default strategy let warmed_count = cache_manager.warm_cache(strategy).await?; Ok(Json(serde_json::json!({ "status": "success", "warmed_entries": warmed_count, "message": format!("Cache warmed with {} entries", warmed_count) }))) } // Helper functions fn generate_cache_key(model: &str, prompt: &str) -> String { use std::collections::hash_map::DefaultHasher; use std::hash::{Hash, Hasher}; let mut hasher = DefaultHasher::new(); model.hash(&mut hasher); prompt.hash(&mut hasher); format!("inference:{}:{:016x}", model, hasher.finish()) } fn tokenize_prompt(prompt: &str) -> ApiResult> { // Simplified tokenization - in production would use actual tokenizer Ok(prompt .split_whitespace() .enumerate() .map(|(i, _)| i as u32) .collect()) } fn detokenize(tokens: &[u32]) -> String { // Simplified detokenization format!("Generated text from {} tokens", tokens.len()) } fn tokens_to_sequence(tokens: &[u32]) -> Vec { // Convert tokens to float sequence for attention computation tokens.iter().map(|&t| t as f32 * 0.1).collect() } fn sequence_to_text(sequence: &[f32]) -> String { format!( "Attention processed sequence with {} elements", sequence.len() ) } fn count_tokens(text: &str) -> u32 { text.split_whitespace().count() as u32 } fn create_cache_value(text: &str, tokens: &[u32]) -> CacheValue { let keys: Vec = tokens.iter().map(|&t| t as f32).collect(); let values: Vec = text.chars().map(|c| c as u32 as f32).collect(); CacheValue::new(keys, values, vec![tokens.len(), 1]) } fn reconstruct_response_from_cache( cache_value: &CacheValue, request_id: &str, ) -> CachedInferenceResponse { let text = format!( "Cached response with {} keys and {} values", cache_value.keys.len(), cache_value.values.len() ); CachedInferenceResponse { text, finish_reason: "cached".to_string(), usage: TokenUsage { prompt_tokens: 0, completion_tokens: cache_value.keys.len() as u32, total_tokens: cache_value.keys.len() as u32, cached_tokens: cache_value.keys.len() as u32, }, cache_stats: CacheRequestStats { cache_used: true, hit_rate: 100.0, cache_hits: 1, cache_misses: 0, speculation_hits: 0, prefix_matches: 0, window_reuse: 0, }, performance: PerformanceMetrics { total_time_ms: 0.1, // Very fast cache hit cache_time_ms: 0.1, inference_time_ms: 0.0, tokens_per_second: 10000.0, // Very fast memory_usage_bytes: cache_value.size_bytes, }, request_id: request_id.to_string(), } } fn create_mock_logits_function(temperature: f32) -> impl Fn(&[u32]) -> Vec { move |tokens: &[u32]| { let base = tokens.last().unwrap_or(&0); vec![ TokenLogit { token: base + 1, logit: 2.0 / temperature, probability: 0.6, }, TokenLogit { token: base + 2, logit: 1.0 / temperature, probability: 0.3, }, TokenLogit { token: base + 3, logit: 0.5 / temperature, probability: 0.1, }, ] } } fn estimate_window_reuse(sequence: &[f32]) -> u32 { // Estimate how much of the sequence was reused from cache (sequence.len() / 4) as u32 // Assume 25% reuse } fn estimate_memory_usage(text: &str) -> usize { text.len() * 4 + std::mem::size_of::() // Rough estimate } async fn perform_standard_inference( request: &CachedInferenceRequest, cache_stats: &mut CacheRequestStats, ) -> ApiResult { // Mock inference - in production this would call actual model let response_text = format!( "Mock response to: '{}' with model '{}'", request.prompt, request.model ); // Simulate some prefix matching if request.prompt.len() > 10 { cache_stats.prefix_matches += 1; } Ok(response_text) } #[cfg(test)] mod tests { use super::*; use crate::cache::config::ContextCacheConfig; #[tokio::test] async fn test_cache_key_generation() { let key1 = generate_cache_key("model1", "hello world"); let key2 = generate_cache_key("model1", "hello world"); let key3 = generate_cache_key("model2", "hello world"); assert_eq!(key1, key2); // Same inputs should generate same key assert_ne!(key1, key3); // Different inputs should generate different keys assert!(key1.starts_with("inference:model1:")); } #[test] fn test_tokenization() { let tokens = tokenize_prompt("hello world test").unwrap(); assert_eq!(tokens.len(), 3); assert_eq!(tokens[0], 0); assert_eq!(tokens[1], 1); assert_eq!(tokens[2], 2); } #[test] fn test_token_counting() { assert_eq!(count_tokens("hello world"), 2); assert_eq!(count_tokens("single"), 1); assert_eq!(count_tokens(""), 0); assert_eq!(count_tokens(" multiple spaces "), 2); } #[test] fn test_cache_value_creation() { let text = "test response"; let tokens = vec![1, 2, 3]; let cache_value = create_cache_value(text, &tokens); assert_eq!(cache_value.keys.len(), 3); assert_eq!(cache_value.values.len(), text.len()); assert_eq!(cache_value.shape, vec![3, 1]); } #[test] fn test_response_reconstruction() { let cache_value = CacheValue::new(vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0], vec![3, 1]); let response = reconstruct_response_from_cache(&cache_value, "test-request-123"); assert_eq!(response.request_id, "test-request-123"); assert_eq!(response.finish_reason, "cached"); assert!(response.cache_stats.cache_used); assert_eq!(response.cache_stats.hit_rate, 100.0); assert_eq!(response.usage.cached_tokens, 3); } #[tokio::test] async fn test_cache_initialization() { let config = ContextCacheConfig::development(); // Should not fail with development config // Note: This would require proper async test setup in production assert!(config.validate().is_ok()); } #[test] fn test_memory_usage_estimation() { let text = "hello"; let usage = estimate_memory_usage(text); // Should be greater than just the text length assert!(usage > text.len()); assert!(usage > 0); } #[test] fn test_mock_logits_function() { let logits_fn = create_mock_logits_function(1.0); let tokens = vec![5, 10, 15]; let logits = logits_fn(&tokens); assert_eq!(logits.len(), 3); assert_eq!(logits[0].token, 16); // last token (15) + 1 assert!(logits[0].probability > logits[1].probability); assert!(logits[1].probability > logits[2].probability); // Test with temperature scaling let logits_fn_cold = create_mock_logits_function(0.5); let logits_cold = logits_fn_cold(&tokens); // Lower temperature should increase logits (making distribution sharper) assert!(logits_cold[0].logit > logits[0].logit); } #[test] fn test_window_reuse_estimation() { let sequence = vec![1.0; 100]; let reuse = estimate_window_reuse(&sequence); assert_eq!(reuse, 25); // Should be 25% of sequence length let small_sequence = vec![1.0; 4]; let small_reuse = estimate_window_reuse(&small_sequence); assert_eq!(small_reuse, 1); } }