//! Model loading and caching infrastructure //! //! Provides ML model loading from S3 with LRU caching and version management. //! Supports historical model retrieval for backtesting scenarios. #![deny(clippy::unwrap_used, clippy::expect_used)] use anyhow::{Context, Result}; use async_trait::async_trait; use lru::LruCache; use parking_lot::Mutex; use semver::Version; use serde::{Deserialize, Serialize}; use std::num::NonZeroUsize; use std::sync::Arc; use std::time::SystemTime; use storage::{ObjectStoreBackend, Storage}; use tracing::{debug, info, warn}; // Re-export canonical ModelType from common crate pub use common::model_types::ModelType; /// Model metadata for version tracking #[derive(Debug, Clone, Serialize, Deserialize)] pub struct LoadedModelInfo { /// Model name pub name: String, /// Semantic version pub version: Version, /// Model type pub model_type: ModelType, /// Timestamp when model was created pub created_at: SystemTime, /// File size in bytes pub size_bytes: usize, /// Checksum for integrity validation pub checksum: String, } /// Cache key for model lookups #[derive(Debug, Clone, PartialEq, Eq, Hash)] struct CacheKey { model_name: String, version: Version, } /// Configuration for model loader #[derive(Debug, Clone)] pub struct ModelLoaderConfig { /// S3 prefix for model files pub prefix: String, /// Maximum number of models to cache pub cache_size: usize, } impl Default for ModelLoaderConfig { fn default() -> Self { Self { prefix: "models/".to_owned(), cache_size: 1000, } } } /// Trait for model loading operations #[async_trait] pub trait ModelLoader: Send + Sync { /// Load a model by name and version async fn load_model(&self, model_name: &str, version: &Version) -> Result>; /// Get model metadata async fn get_metadata(&self, model_name: &str, version: &Version) -> Result; /// List all available versions of a model async fn list_versions(&self, model_name: &str) -> Result>; /// Get the most recent model version for a time period async fn get_model_for_period( &self, model_name: &str, start: SystemTime, end: SystemTime, ) -> Result<(Version, Vec)>; } /// S3-based model loader with LRU caching pub struct S3ModelLoader { storage: Arc, config: ModelLoaderConfig, cache: Arc>>>, } impl S3ModelLoader { /// Create a new S3 model loader pub fn new(storage: ObjectStoreBackend, config: ModelLoaderConfig) -> Self { let cache_size = if config.cache_size == 0 { warn!("Cache size cannot be 0, using default value of 1000"); // SAFETY: This is a compile-time constant that is guaranteed to be non-zero #[allow(clippy::expect_used)] NonZeroUsize::new(1000).expect("1000 is non-zero") } else { // SAFETY: We just checked that cache_size is non-zero #[allow(clippy::expect_used)] NonZeroUsize::new(config.cache_size).expect("cache_size is non-zero after check") }; Self { storage: Arc::new(storage), config, cache: Arc::new(Mutex::new(LruCache::new(cache_size))), } } /// Build S3 key for model file fn build_model_key(&self, model_name: &str, version: &Version) -> String { format!("{}{}/{}/model.bin", self.config.prefix, model_name, version) } /// Build S3 key for metadata file fn build_metadata_key(&self, model_name: &str, version: &Version) -> String { format!( "{}{}/{}/metadata.json", self.config.prefix, model_name, version ) } /// Build S3 prefix for listing versions fn build_version_prefix(&self, model_name: &str) -> String { format!("{}{}/", self.config.prefix, model_name) } /// Get from cache or load from S3 async fn get_cached_or_load(&self, model_name: &str, version: &Version) -> Result> { let cache_key = CacheKey { model_name: model_name.to_owned(), version: version.clone(), }; // Check cache first { let mut cache = self.cache.lock(); if let Some(data) = cache.get(&cache_key) { debug!("Cache hit for model {}/{}", model_name, version); return Ok(data.clone()); } } // Cache miss - load from S3 debug!( "Cache miss for model {}/{}, loading from S3", model_name, version ); let key = self.build_model_key(model_name, version); let data = self.storage.retrieve(&key).await.with_context(|| { format!("Failed to load model {}/{} from S3", model_name, version) })?; // Store in cache { let mut cache = self.cache.lock(); cache.put(cache_key, data.clone()); } info!( "Loaded model {}/{} from S3 ({} bytes)", model_name, version, data.len() ); Ok(data) } } #[async_trait] impl ModelLoader for S3ModelLoader { async fn load_model(&self, model_name: &str, version: &Version) -> Result> { self.get_cached_or_load(model_name, version).await } async fn get_metadata(&self, model_name: &str, version: &Version) -> Result { let key = self.build_metadata_key(model_name, version); let data = self.storage.retrieve(&key).await.with_context(|| { format!("Failed to load metadata for {}/{}", model_name, version) })?; let metadata: LoadedModelInfo = serde_json::from_slice(&data) .with_context(|| "Failed to deserialize model metadata")?; Ok(metadata) } async fn list_versions(&self, model_name: &str) -> Result> { let prefix = self.build_version_prefix(model_name); let objects = self .storage .list(&prefix) .await .with_context(|| format!("Failed to list versions for model {}", model_name))?; let mut versions = Vec::new(); for key in objects { // Extract version from key: models/model_name/VERSION/... if let Some(version_str) = key.strip_prefix(&prefix).and_then(|s| s.split('/').next()) { if let Ok(version) = Version::parse(version_str) { versions.push(version); } } } versions.sort(); versions.reverse(); // Most recent first Ok(versions) } async fn get_model_for_period( &self, model_name: &str, start: SystemTime, _end: SystemTime, ) -> Result<(Version, Vec)> { // List all versions let versions = self.list_versions(model_name).await?; // Find the most recent version that was created before the period start for version in &versions { let metadata = self.get_metadata(model_name, version).await?; if metadata.created_at <= start { let data = self.load_model(model_name, version).await?; info!( "Found model {}/{} for period (created at {:?})", model_name, version, metadata.created_at ); return Ok((version.clone(), data)); } } // If no version found before start, use the oldest version if let Some(oldest) = versions.last() { warn!( "No model version found before period start, using oldest version {}", oldest ); let data = self.load_model(model_name, oldest).await?; return Ok((oldest.clone(), data)); } anyhow::bail!("No versions found for model {}", model_name) } } /// Backtesting model cache module for historical model loading pub mod backtesting_cache { use super::{ Arc, ModelLoader, ModelLoaderConfig, ObjectStoreBackend, Result, S3ModelLoader, SystemTime, Version, }; use anyhow::Context; use std::path::PathBuf; /// Configuration for backtesting model cache #[derive(Debug, Clone)] pub struct BacktestCacheConfig { /// Directory for caching model files pub cache_dir: PathBuf, /// S3 configuration for model storage pub model_loader_config: ModelLoaderConfig, } impl Default for BacktestCacheConfig { fn default() -> Self { Self { cache_dir: PathBuf::from("/tmp/foxhunt/model_cache"), model_loader_config: ModelLoaderConfig::default(), } } } /// Backtesting model cache /// /// Provides access to historical model versions for backtesting scenarios. pub struct BacktestingModelCache { loader: Arc, _config: BacktestCacheConfig, } impl BacktestingModelCache { /// Create a new backtesting model cache /// /// # Errors /// Returns error if the operation fails /// /// # Errors /// Returns error if cache initialization fails pub async fn new(storage: ObjectStoreBackend, config: BacktestCacheConfig) -> Result { let loader = S3ModelLoader::new(storage, config.model_loader_config.clone()); Ok(Self { loader: Arc::new(loader), _config: config, }) } /// Initialize the model cache /// /// # Errors /// Returns error if initialization fails pub async fn initialize(&mut self) -> Result<()> { // Cache is ready immediately with LRU Ok(()) } /// Get a model by name and version /// /// /// # Errors /// Returns error if the operation fails /// # Errors /// Returns error if the version string is invalid or model cannot be loaded pub async fn get_model(&self, model_name: &str, version_str: &str) -> Result> { let version = Version::parse(version_str) .with_context(|| format!("Invalid version string: {}", version_str))?; self.loader.load_model(model_name, &version).await } /// Get a specific model version /// /// # Errors /// Returns error if the model cannot be loaded pub async fn get_model_version( &self, model_name: &str, version: &Version, ) -> Result> { self.loader.load_model(model_name, version).await } /// Get a model for a specific time period /// /// # Errors /// Returns error if the operation fails /// /// This is used for backtesting to load historically accurate model versions /// /// # Errors /// Returns error if no suitable model version is found or model cannot be loaded pub async fn get_model_for_period( &self, model_name: &str, start: SystemTime, end: SystemTime, ) -> Result<(Version, Vec)> { self.loader .get_model_for_period(model_name, start, end) .await } /// List all available versions of a model /// /// # Errors /// Returns error if listing model versions fails pub async fn list_model_versions(&self, model_name: &str) -> Result> { self.loader.list_versions(model_name).await } } } #[cfg(test)] #[allow(clippy::unwrap_used, clippy::expect_used)] mod tests { use super::*; #[test] fn test_model_type_s3_prefix() { assert_eq!(ModelType::TLOB.s3_prefix(), "tlob_transformer"); assert_eq!(ModelType::DQN.s3_prefix(), "dqn"); assert_eq!(ModelType::MAMBA.s3_prefix(), "mamba2"); } #[test] fn test_model_type_as_str() { assert_eq!(ModelType::TLOB.as_str(), "tlob"); assert_eq!(ModelType::DQN.as_str(), "dqn"); assert_eq!(ModelType::MAMBA.as_str(), "mamba"); } #[test] fn test_cache_key_equality() { let key1 = CacheKey { model_name: "test_model".to_owned(), version: Version::new(1, 0, 0), }; let key2 = CacheKey { model_name: "test_model".to_owned(), version: Version::new(1, 0, 0), }; assert_eq!(key1, key2); } #[test] fn test_model_loader_config_default() { let config = ModelLoaderConfig::default(); assert_eq!(config.prefix, "models/"); assert_eq!(config.cache_size, 1000); } }