Files
foxhunt/crates/ml-checkpoint/src/storage.rs
jgrusewski fca2495a73 fix(clippy): add doc backticks and const fn across ML crates
- Add backticks to type names in doc comments (doc_markdown)
- Mark eligible functions as const fn (missing_const_for_fn)
No behavior changes.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-10 11:51:31 +01:00

1147 lines
37 KiB
Rust

//! Storage backends for checkpoint persistence
//!
//! Provides multiple storage options for checkpoint data with consistent interface.
//! Supports local filesystem, in-memory (for testing), and AWS S3 cloud storage.
use std::fs::{self, File};
use std::io::{BufReader, BufWriter, Read, Write};
use std::path::PathBuf;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use tracing::{debug, error, info, warn};
use super::CheckpointMetadata;
use crate::MLError;
// S3 dependencies for AWS SDK
#[cfg(feature = "s3-storage")]
use aws_config::BehaviorVersion;
#[cfg(feature = "s3-storage")]
use aws_credential_types::Credentials;
#[cfg(feature = "s3-storage")]
use aws_sdk_s3::primitives::ByteStream;
#[cfg(feature = "s3-storage")]
use aws_sdk_s3::types::StorageClass;
#[cfg(feature = "s3-storage")]
use aws_sdk_s3::Client as S3Client;
/// Trait for checkpoint storage backends
#[async_trait]
pub trait CheckpointStorage: std::fmt::Debug + Send + Sync {
/// Save a checkpoint to storage
async fn save_checkpoint(
&self,
filename: &str,
data: &[u8],
metadata: &CheckpointMetadata,
) -> Result<(), MLError>;
/// Load a checkpoint from storage
async fn load_checkpoint(&self, filename: &str) -> Result<Vec<u8>, MLError>;
/// Delete a checkpoint from storage
async fn delete_checkpoint(&self, filename: &str) -> Result<(), MLError>;
/// List all checkpoint metadata
async fn list_all_checkpoints(&self) -> Result<Vec<CheckpointMetadata>, MLError>;
/// Check if a checkpoint exists
async fn has_checkpoint(&self, filename: &str) -> bool;
/// Get storage statistics
async fn get_storage_stats(&self) -> Result<StorageStats, MLError>;
}
/// Statistics about storage usage
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StorageStats {
/// Total number of checkpoints
pub total_checkpoints: u64,
/// Total storage used (bytes)
pub total_bytes: u64,
/// Available storage space (bytes)
pub available_bytes: u64,
/// Average checkpoint size (bytes)
pub avg_checkpoint_size: u64,
/// Storage backend type
pub backend_type: String,
}
/// File system storage backend
#[derive(Debug)]
pub struct FileSystemStorage {
/// Base directory for checkpoints
base_dir: PathBuf,
/// Metadata directory
metadata_dir: PathBuf,
}
impl FileSystemStorage {
/// Create a new filesystem storage backend
#[allow(clippy::cognitive_complexity)]
pub fn new(base_dir: PathBuf) -> Self {
let metadata_dir = base_dir.join("metadata");
// Ensure directories exist
debug!(
base_dir = %base_dir.display(),
metadata_dir = %metadata_dir.display(),
"Creating ML checkpoint storage directories"
);
if let Err(e) = fs::create_dir_all(&base_dir) {
error!(
error = %e,
base_dir = %base_dir.display(),
"Failed to create ML checkpoint base directory"
);
} else {
debug!(
base_dir = %base_dir.display(),
"Successfully created ML checkpoint base directory"
);
}
if let Err(e) = fs::create_dir_all(&metadata_dir) {
error!(
error = %e,
metadata_dir = %metadata_dir.display(),
"Failed to create ML checkpoint metadata directory"
);
} else {
debug!(
metadata_dir = %metadata_dir.display(),
"Successfully created ML checkpoint metadata directory"
);
}
Self {
base_dir,
metadata_dir,
}
}
/// Get path for checkpoint data file
fn checkpoint_path(&self, filename: &str) -> PathBuf {
self.base_dir.join(filename)
}
/// Get path for metadata file
fn metadata_path(&self, filename: &str) -> PathBuf {
let metadata_filename = format!("{}.metadata.json", filename);
self.metadata_dir.join(metadata_filename)
}
/// Save metadata to file
fn save_metadata(&self, filename: &str, metadata: &CheckpointMetadata) -> Result<(), MLError> {
use tracing::{debug, error};
let metadata_path = self.metadata_path(filename);
debug!(
filename = filename,
metadata_path = %metadata_path.display(),
model_name = %metadata.model_name,
version = metadata.version,
"Starting metadata save operation"
);
let file = File::create(&metadata_path).map_err(|e| {
error!(
error = %e,
filename = filename,
metadata_path = %metadata_path.display(),
"Failed to create metadata file for ML checkpoint"
);
MLError::ModelError(format!(
"Failed to create metadata file {}: {}",
metadata_path.display(),
e
))
})?;
let mut writer = BufWriter::new(file);
serde_json::to_writer_pretty(&mut writer, metadata).map_err(|e| {
error!(
error = %e,
filename = filename,
metadata_path = %metadata_path.display(),
"Failed to serialize metadata to JSON"
);
MLError::ModelError(format!("Failed to write metadata: {}", e))
})?;
writer.flush().map_err(|e| {
error!(
error = %e,
filename = filename,
metadata_path = %metadata_path.display(),
"Failed to flush metadata buffer to disk"
);
MLError::ModelError(format!("Failed to flush metadata: {}", e))
})?;
debug!(
filename = filename,
metadata_path = %metadata_path.display(),
model_name = %metadata.model_name,
"Successfully saved ML checkpoint metadata"
);
Ok(())
}
/// Load metadata from file
fn load_metadata(&self, filename: &str) -> Result<CheckpointMetadata, MLError> {
use tracing::{debug, error};
let metadata_path = self.metadata_path(filename);
debug!(
filename = filename,
metadata_path = %metadata_path.display(),
"Starting metadata load operation"
);
let file = File::open(&metadata_path).map_err(|e| {
error!(
error = %e,
filename = filename,
metadata_path = %metadata_path.display(),
"Failed to open metadata file for ML checkpoint"
);
MLError::ModelError(format!(
"Failed to open metadata file {}: {}",
metadata_path.display(),
e
))
})?;
let reader = BufReader::new(file);
let metadata: CheckpointMetadata = serde_json::from_reader(reader).map_err(|e| {
error!(
error = %e,
filename = filename,
metadata_path = %metadata_path.display(),
"Failed to parse JSON metadata from file"
);
MLError::ModelError(format!("Failed to parse metadata: {}", e))
})?;
debug!(
filename = filename,
metadata_path = %metadata_path.display(),
model_name = %metadata.model_name,
version = metadata.version,
"Successfully loaded ML checkpoint metadata"
);
Ok(metadata)
}
}
#[async_trait]
impl CheckpointStorage for FileSystemStorage {
async fn save_checkpoint(
&self,
filename: &str,
data: &[u8],
metadata: &CheckpointMetadata,
) -> Result<(), MLError> {
let checkpoint_path = self.checkpoint_path(filename);
// Save checkpoint data
let file = File::create(&checkpoint_path).map_err(|e| {
MLError::ModelError(format!(
"Failed to create checkpoint file {}: {}",
checkpoint_path.display(),
e
))
})?;
let mut writer = BufWriter::new(file);
writer
.write_all(data)
.map_err(|e| MLError::ModelError(format!("Failed to write checkpoint data: {}", e)))?;
writer
.flush()
.map_err(|e| MLError::ModelError(format!("Failed to flush checkpoint data: {}", e)))?;
// Save metadata
self.save_metadata(filename, metadata)?;
info!("Saved checkpoint {} ({} bytes)", filename, data.len());
Ok(())
}
async fn load_checkpoint(&self, filename: &str) -> Result<Vec<u8>, MLError> {
let checkpoint_path = self.checkpoint_path(filename);
let file = File::open(&checkpoint_path).map_err(|e| {
MLError::ModelError(format!(
"Failed to open checkpoint file {}: {}",
checkpoint_path.display(),
e
))
})?;
let mut reader = BufReader::new(file);
let mut data = Vec::new();
reader
.read_to_end(&mut data)
.map_err(|e| MLError::ModelError(format!("Failed to read checkpoint data: {}", e)))?;
debug!("Loaded checkpoint {} ({} bytes)", filename, data.len());
Ok(data)
}
async fn delete_checkpoint(&self, filename: &str) -> Result<(), MLError> {
let checkpoint_path = self.checkpoint_path(filename);
let metadata_path = self.metadata_path(filename);
// Delete checkpoint file
if checkpoint_path.exists() {
fs::remove_file(&checkpoint_path).map_err(|e| {
MLError::ModelError(format!(
"Failed to delete checkpoint file {}: {}",
checkpoint_path.display(),
e
))
})?;
}
// Delete metadata file
if metadata_path.exists() {
fs::remove_file(&metadata_path).map_err(|e| {
MLError::ModelError(format!(
"Failed to delete metadata file {}: {}",
metadata_path.display(),
e
))
})?;
}
info!("Deleted checkpoint {}", filename);
Ok(())
}
async fn list_all_checkpoints(&self) -> Result<Vec<CheckpointMetadata>, MLError> {
let mut checkpoints = Vec::new();
// Read metadata directory
let entries = fs::read_dir(&self.metadata_dir).map_err(|e| {
MLError::ModelError(format!("Failed to read metadata directory: {}", e))
})?;
for entry in entries {
let entry = entry.map_err(|e| {
MLError::ModelError(format!("Failed to read metadata directory entry: {}", e))
})?;
let path = entry.path();
if path.extension().and_then(|s| s.to_str()) == Some("json") {
if let Some(filename) = path.file_stem().and_then(|s| s.to_str()) {
// Remove .metadata suffix
if let Some(base_filename) = filename.strip_suffix(".metadata") {
match self.load_metadata(base_filename) {
Ok(metadata) => checkpoints.push(metadata),
Err(e) => {
warn!("Failed to load metadata for {}: {}", base_filename, e);
},
}
}
}
}
}
debug!("Listed {} checkpoints", checkpoints.len());
Ok(checkpoints)
}
async fn has_checkpoint(&self, filename: &str) -> bool {
let checkpoint_path = self.checkpoint_path(filename);
let metadata_path = self.metadata_path(filename);
checkpoint_path.exists() && metadata_path.exists()
}
async fn get_storage_stats(&self) -> Result<StorageStats, MLError> {
let checkpoints = self.list_all_checkpoints().await?;
let total_checkpoints = checkpoints.len() as u64;
let mut total_bytes = 0_u64;
// Calculate total storage used
for metadata in &checkpoints {
total_bytes += metadata.compressed_size.unwrap_or(metadata.file_size);
}
// Get available space
let available_bytes = match fs2::available_space(&self.base_dir) {
Ok(space) => space,
Err(_) => u64::MAX, // Fallback if we can't determine available space
};
let avg_checkpoint_size = if total_checkpoints > 0 {
total_bytes / total_checkpoints
} else {
0
};
Ok(StorageStats {
total_checkpoints,
total_bytes,
available_bytes,
avg_checkpoint_size,
backend_type: "filesystem".to_owned(),
})
}
}
/// In-memory storage backend for testing
#[derive(Debug, Default)]
pub struct MemoryStorage {
/// Checkpoint data
checkpoints: std::sync::RwLock<std::collections::HashMap<String, Vec<u8>>>,
/// Metadata
metadata: std::sync::RwLock<std::collections::HashMap<String, CheckpointMetadata>>,
}
impl MemoryStorage {
/// Create a new memory storage backend
pub fn new() -> Self {
Self::default()
}
}
#[async_trait]
impl CheckpointStorage for MemoryStorage {
async fn save_checkpoint(
&self,
filename: &str,
data: &[u8],
metadata: &CheckpointMetadata,
) -> Result<(), MLError> {
{
let mut checkpoints =
self.checkpoints
.write()
.map_err(|e| MLError::ConcurrencyError {
operation: format!("write lock checkpoints: {}", e),
})?;
checkpoints.insert(filename.to_owned(), data.to_vec());
}
{
let mut metadata_map =
self.metadata
.write()
.map_err(|e| MLError::ConcurrencyError {
operation: format!("write lock metadata: {}", e),
})?;
metadata_map.insert(filename.to_owned(), metadata.clone());
}
debug!(
"Saved checkpoint {} to memory ({} bytes)",
filename,
data.len()
);
Ok(())
}
async fn load_checkpoint(&self, filename: &str) -> Result<Vec<u8>, MLError> {
let checkpoints = self
.checkpoints
.read()
.map_err(|e| MLError::ConcurrencyError {
operation: format!("read lock checkpoints: {}", e),
})?;
let data = checkpoints
.get(filename)
.cloned()
.ok_or_else(|| MLError::ModelError(format!("Checkpoint not found: {}", filename)))?;
debug!(
"Loaded checkpoint {} from memory ({} bytes)",
filename,
data.len()
);
Ok(data)
}
async fn delete_checkpoint(&self, filename: &str) -> Result<(), MLError> {
{
let mut checkpoints =
self.checkpoints
.write()
.map_err(|e| MLError::ConcurrencyError {
operation: format!("write lock checkpoints for delete: {}", e),
})?;
checkpoints.remove(filename);
}
{
let mut metadata_map =
self.metadata
.write()
.map_err(|e| MLError::ConcurrencyError {
operation: format!("write lock metadata: {}", e),
})?;
metadata_map.remove(filename);
}
debug!("Deleted checkpoint {} from memory", filename);
Ok(())
}
async fn list_all_checkpoints(&self) -> Result<Vec<CheckpointMetadata>, MLError> {
let metadata_map = self
.metadata
.read()
.map_err(|e| MLError::ConcurrencyError {
operation: format!("read lock metadata: {}", e),
})?;
let checkpoints: Vec<_> = metadata_map.values().cloned().collect();
debug!("Listed {} checkpoints from memory", checkpoints.len());
Ok(checkpoints)
}
async fn has_checkpoint(&self, filename: &str) -> bool {
match self.checkpoints.read() {
Ok(checkpoints) => checkpoints.contains_key(filename),
Err(_) => false, // If we can't read, assume it doesn't exist
}
}
async fn get_storage_stats(&self) -> Result<StorageStats, MLError> {
let checkpoints = self
.checkpoints
.read()
.map_err(|e| MLError::ConcurrencyError {
operation: format!("read lock checkpoints: {}", e),
})?;
let _metadata_map = self
.metadata
.read()
.map_err(|e| MLError::ConcurrencyError {
operation: format!("read lock metadata: {}", e),
})?;
let total_checkpoints = checkpoints.len() as u64;
let total_bytes: u64 = checkpoints.values().map(|data| data.len() as u64).sum();
let avg_checkpoint_size = if total_checkpoints > 0 {
total_bytes / total_checkpoints
} else {
0
};
Ok(StorageStats {
total_checkpoints,
total_bytes,
available_bytes: u64::MAX, // Unlimited for memory storage
avg_checkpoint_size,
backend_type: "memory".to_owned(),
})
}
}
/// S3 storage backend for cloud checkpoint storage
#[cfg(feature = "s3-storage")]
#[derive(Debug, Clone)]
pub struct S3CheckpointStorage {
/// S3 client
client: S3Client,
/// S3 bucket name
bucket_name: String,
/// Key prefix for checkpoints
key_prefix: String,
/// Storage class for checkpoints
storage_class: StorageClass,
/// Enable server-side encryption
encryption_enabled: bool,
}
#[cfg(feature = "s3-storage")]
impl S3CheckpointStorage {
/// Create a new S3 checkpoint storage backend from environment variables
pub async fn from_env() -> Result<Self, MLError> {
let bucket_name = std::env::var("S3_CHECKPOINT_BUCKET")
.unwrap_or_else(|_| "foxhunt-checkpoints".to_owned());
let key_prefix = std::env::var("S3_CHECKPOINT_PREFIX")
.ok()
.unwrap_or_else(|| "ml-checkpoints".to_owned());
let region = std::env::var("AWS_REGION").unwrap_or_else(|_| "us-east-1".to_owned());
info!(
"Initializing S3 checkpoint storage from environment: bucket={}, prefix={}, region={}",
bucket_name, key_prefix, region
);
let client = Self::create_s3_client_from_env().await?;
// Test bucket access
client
.head_bucket()
.bucket(&bucket_name)
.send()
.await
.map_err(|e| {
MLError::ModelError(format!(
"Failed to access S3 bucket '{}': {}. Check credentials and bucket permissions.",
bucket_name, e
))
})?;
info!(
"Successfully connected to S3 bucket '{}' for checkpoint storage",
bucket_name
);
Ok(Self {
client,
bucket_name,
key_prefix,
storage_class: StorageClass::StandardIa, // Infrequent Access for cost optimization
encryption_enabled: std::env::var("S3_ENABLE_ENCRYPTION")
.map(|v| v.to_lowercase() == "true")
.unwrap_or(true),
})
}
/// Create a new S3 checkpoint storage backend with explicit credentials
pub async fn new(
bucket_name: String,
key_prefix: Option<String>,
region: Option<String>,
access_key_id: String,
secret_access_key: String,
) -> Result<Self, MLError> {
info!(
"Initializing S3 checkpoint storage: bucket={}, prefix={:?}, region={:?}",
bucket_name, key_prefix, region
);
// Configure AWS SDK
let region_name = region.unwrap_or_else(|| "us-east-1".to_owned());
let aws_region = aws_types::region::Region::new(region_name);
let aws_config = aws_config::defaults(BehaviorVersion::latest())
.region(aws_region)
.credentials_provider(Credentials::new(
access_key_id,
secret_access_key,
None, // session_token
None, // expiration
"ml_checkpoints", // provider_name
))
.load()
.await;
let s3_client = S3Client::new(&aws_config);
// Test bucket access
s3_client
.head_bucket()
.bucket(&bucket_name)
.send()
.await
.map_err(|e| {
MLError::ModelError(format!(
"Failed to access S3 bucket '{}': {}. Check credentials and bucket permissions.",
bucket_name, e
))
})?;
info!(
"Successfully connected to S3 bucket '{}' for checkpoint storage",
bucket_name
);
Ok(Self {
client: s3_client,
bucket_name,
key_prefix: key_prefix.unwrap_or_else(|| "ml-checkpoints".to_owned()),
storage_class: StorageClass::StandardIa, // Infrequent Access for cost optimization
encryption_enabled: true,
})
}
/// Create AWS S3 client using environment variables or AWS credential chain
async fn create_s3_client_from_env() -> Result<S3Client, MLError> {
// Try to use explicit credentials first, then fall back to AWS credential chain
let region_name = std::env::var("AWS_REGION").unwrap_or_else(|_| "us-east-1".to_owned());
let aws_region = aws_types::region::Region::new(region_name);
let config = if let (Ok(access_key), Ok(secret_key)) = (
std::env::var("AWS_ACCESS_KEY_ID"),
std::env::var("AWS_SECRET_ACCESS_KEY"),
) {
info!("Using explicit AWS credentials from environment variables");
let creds = Credentials::new(
access_key,
secret_key,
std::env::var("AWS_SESSION_TOKEN").ok(),
None,
"environment",
);
aws_config::defaults(BehaviorVersion::latest())
.region(aws_region.clone())
.credentials_provider(creds)
.load()
.await
} else {
info!("Using AWS default credential chain (IAM roles, profiles, etc.)");
aws_config::defaults(BehaviorVersion::latest())
.region(aws_region)
.load()
.await
};
Ok(S3Client::new(&config))
}
/// Generate S3 key for a checkpoint
fn generate_s3_key(&self, filename: &str) -> String {
format!("{}/{}", self.key_prefix, filename)
}
/// Generate S3 key for metadata
fn generate_metadata_key(&self, filename: &str) -> String {
format!("{}/metadata/{}.json", self.key_prefix, filename)
}
/// Create object metadata for S3
fn create_object_metadata(
&self,
checkpoint_metadata: &CheckpointMetadata,
) -> std::collections::HashMap<String, String> {
let mut metadata = std::collections::HashMap::new();
metadata.insert(
"model_type".to_owned(),
format!("{:?}", checkpoint_metadata.model_type),
);
metadata.insert(
"model_name".to_owned(),
checkpoint_metadata.model_name.clone(),
);
metadata.insert("version".to_owned(), checkpoint_metadata.version.clone());
metadata.insert(
"checkpoint_id".to_owned(),
checkpoint_metadata.checkpoint_id.clone(),
);
metadata.insert(
"created_at".to_owned(),
checkpoint_metadata.created_at.to_rfc3339(),
);
if let Some(epoch) = checkpoint_metadata.epoch {
metadata.insert("epoch".to_owned(), epoch.to_string());
}
if let Some(step) = checkpoint_metadata.step {
metadata.insert("step".to_owned(), step.to_string());
}
if let Some(loss) = checkpoint_metadata.loss {
metadata.insert("loss".to_owned(), loss.to_string());
}
if let Some(accuracy) = checkpoint_metadata.accuracy {
metadata.insert("accuracy".to_owned(), accuracy.to_string());
}
metadata.insert("service".to_owned(), "ml-training-service".to_owned());
metadata.insert("purpose".to_owned(), "model-checkpoint".to_owned());
metadata
}
/// Create S3 tags for organizing checkpoints
fn create_object_tags(
&self,
checkpoint_metadata: &CheckpointMetadata,
) -> Result<Vec<aws_sdk_s3::types::Tag>, MLError> {
let mut tags = vec![
aws_sdk_s3::types::Tag::builder()
.key("model_type")
.value(format!("{:?}", checkpoint_metadata.model_type))
.build()
.map_err(|e| {
MLError::CheckpointError(format!("Failed to build model_type tag: {:?}", e))
})?,
aws_sdk_s3::types::Tag::builder()
.key("model_name")
.value(&checkpoint_metadata.model_name)
.build()
.map_err(|e| {
MLError::CheckpointError(format!("Failed to build model_name tag: {:?}", e))
})?,
aws_sdk_s3::types::Tag::builder()
.key("version")
.value(&checkpoint_metadata.version)
.build()
.map_err(|e| {
MLError::CheckpointError(format!("Failed to build version tag: {:?}", e))
})?,
aws_sdk_s3::types::Tag::builder()
.key("service")
.value("ml-training")
.build()
.map_err(|e| {
MLError::CheckpointError(format!("Failed to build service tag: {:?}", e))
})?,
];
// Add custom tags from metadata
for tag in &checkpoint_metadata.tags {
tags.push(
aws_sdk_s3::types::Tag::builder()
.key("custom_tag")
.value(tag)
.build()
.map_err(|e| {
MLError::CheckpointError(format!("Failed to build custom tag: {:?}", e))
})?,
);
}
Ok(tags)
}
/// Save metadata to S3 as a separate JSON object
async fn save_metadata_to_s3(
&self,
filename: &str,
metadata: &CheckpointMetadata,
) -> Result<(), MLError> {
let metadata_key = self.generate_metadata_key(filename);
let metadata_json = serde_json::to_string_pretty(metadata)
.map_err(|e| MLError::ModelError(format!("Failed to serialize metadata: {}", e)))?;
let body = ByteStream::from(metadata_json.into_bytes());
let mut request = self
.client
.put_object()
.bucket(&self.bucket_name)
.key(&metadata_key)
.body(body)
.content_type("application/json");
if self.encryption_enabled {
request =
request.server_side_encryption(aws_sdk_s3::types::ServerSideEncryption::Aes256);
}
request
.send()
.await
.map_err(|e| MLError::ModelError(format!("Failed to save metadata to S3: {}", e)))?;
debug!("Saved checkpoint metadata to S3: {}", metadata_key);
Ok(())
}
/// Load metadata from S3
async fn load_metadata_from_s3(&self, filename: &str) -> Result<CheckpointMetadata, MLError> {
let metadata_key = self.generate_metadata_key(filename);
let response = self
.client
.get_object()
.bucket(&self.bucket_name)
.key(&metadata_key)
.send()
.await
.map_err(|e| MLError::ModelError(format!("Failed to load metadata from S3: {}", e)))?;
let body = response
.body
.collect()
.await
.map_err(|e| MLError::ModelError(format!("Failed to read metadata body: {}", e)))?;
let metadata_json = String::from_utf8(body.into_bytes().to_vec())
.map_err(|e| MLError::ModelError(format!("Invalid UTF-8 in metadata: {}", e)))?;
let metadata: CheckpointMetadata = serde_json::from_str(&metadata_json)
.map_err(|e| MLError::ModelError(format!("Failed to parse metadata JSON: {}", e)))?;
debug!("Loaded checkpoint metadata from S3: {}", metadata_key);
Ok(metadata)
}
}
#[cfg(feature = "s3-storage")]
#[async_trait]
impl CheckpointStorage for S3CheckpointStorage {
async fn save_checkpoint(
&self,
filename: &str,
data: &[u8],
metadata: &CheckpointMetadata,
) -> Result<(), MLError> {
let s3_key = self.generate_s3_key(filename);
debug!("Saving checkpoint to S3: {} ({} bytes)", s3_key, data.len());
// Create object metadata and tags
let object_metadata = self.create_object_metadata(metadata);
let tags = self.create_object_tags(metadata)?;
// Convert tags to URL-encoded string format (key1=value1&key2=value2)
let tagging_str = tags
.iter()
.map(|tag| {
let key = tag.key();
let value = tag.value();
format!(
"{}={}",
urlencoding::encode(key),
urlencoding::encode(value)
)
})
.collect::<Vec<_>>()
.join("&");
// Upload checkpoint data
let body = ByteStream::from(data.to_vec());
let mut request = self
.client
.put_object()
.bucket(&self.bucket_name)
.key(&s3_key)
.body(body)
.set_metadata(Some(object_metadata))
.storage_class(self.storage_class.clone())
.tagging(tagging_str)
.content_type("application/octet-stream");
if self.encryption_enabled {
request =
request.server_side_encryption(aws_sdk_s3::types::ServerSideEncryption::Aes256);
}
request
.send()
.await
.map_err(|e| MLError::ModelError(format!("Failed to save checkpoint to S3: {}", e)))?;
// Save metadata as separate object for easier querying
self.save_metadata_to_s3(filename, metadata).await?;
info!(
"Successfully saved checkpoint to S3: {} ({} bytes)",
s3_key,
data.len()
);
Ok(())
}
async fn load_checkpoint(&self, filename: &str) -> Result<Vec<u8>, MLError> {
let s3_key = self.generate_s3_key(filename);
debug!("Loading checkpoint from S3: {}", s3_key);
let response = self
.client
.get_object()
.bucket(&self.bucket_name)
.key(&s3_key)
.send()
.await
.map_err(|e| {
MLError::ModelError(format!("Failed to load checkpoint from S3: {}", e))
})?;
let body =
response.body.collect().await.map_err(|e| {
MLError::ModelError(format!("Failed to read checkpoint body: {}", e))
})?;
let data = body.into_bytes().to_vec();
debug!(
"Successfully loaded checkpoint from S3: {} ({} bytes)",
s3_key,
data.len()
);
Ok(data)
}
async fn delete_checkpoint(&self, filename: &str) -> Result<(), MLError> {
let s3_key = self.generate_s3_key(filename);
let metadata_key = self.generate_metadata_key(filename);
debug!("Deleting checkpoint from S3: {}", s3_key);
// Delete checkpoint data
self.client
.delete_object()
.bucket(&self.bucket_name)
.key(&s3_key)
.send()
.await
.map_err(|e| {
MLError::ModelError(format!("Failed to delete checkpoint from S3: {}", e))
})?;
// Delete metadata
self.client
.delete_object()
.bucket(&self.bucket_name)
.key(&metadata_key)
.send()
.await
.map_err(|e| {
MLError::ModelError(format!("Failed to delete metadata from S3: {}", e))
})?;
info!("Successfully deleted checkpoint from S3: {}", s3_key);
Ok(())
}
async fn list_all_checkpoints(&self) -> Result<Vec<CheckpointMetadata>, MLError> {
debug!(
"Listing all checkpoints from S3 bucket: {}",
self.bucket_name
);
let metadata_prefix = format!("{}/metadata/", self.key_prefix);
let mut checkpoints = Vec::new();
let mut continuation_token: Option<String> = None;
// List all metadata files
loop {
let mut request = self
.client
.list_objects_v2()
.bucket(&self.bucket_name)
.prefix(&metadata_prefix);
if let Some(token) = &continuation_token {
request = request.continuation_token(token);
}
let response = request
.send()
.await
.map_err(|e| MLError::ModelError(format!("Failed to list objects in S3: {}", e)))?;
if let Some(contents) = response.contents {
for object in contents {
if let Some(key) = object.key {
if key.ends_with(".json") {
// Extract filename from metadata key
if let Some(filename) = key
.strip_prefix(&metadata_prefix)
.and_then(|s| s.strip_suffix(".json"))
{
match self.load_metadata_from_s3(filename).await {
Ok(metadata) => checkpoints.push(metadata),
Err(e) => {
warn!("Failed to load metadata for {}: {}", filename, e);
},
}
}
}
}
}
}
// Check for more results
if response.is_truncated.unwrap_or(false) {
continuation_token = response.next_continuation_token;
} else {
break;
}
}
debug!("Listed {} checkpoints from S3", checkpoints.len());
Ok(checkpoints)
}
async fn has_checkpoint(&self, filename: &str) -> bool {
let s3_key = self.generate_s3_key(filename);
match self
.client
.head_object()
.bucket(&self.bucket_name)
.key(&s3_key)
.send()
.await
{
Ok(_) => true,
Err(_) => false,
}
}
async fn get_storage_stats(&self) -> Result<StorageStats, MLError> {
debug!(
"Calculating S3 storage statistics for bucket: {}",
self.bucket_name
);
let mut total_checkpoints = 0_u64;
let mut total_bytes = 0_u64;
let mut continuation_token: Option<String> = None;
// List all checkpoint objects (not metadata)
loop {
let mut request = self
.client
.list_objects_v2()
.bucket(&self.bucket_name)
.prefix(&format!("{}/", self.key_prefix));
if let Some(token) = &continuation_token {
request = request.continuation_token(token);
}
let response = request.send().await.map_err(|e| {
MLError::ModelError(format!("Failed to list objects for stats: {}", e))
})?;
if let Some(contents) = response.contents {
for object in contents {
if let Some(key) = &object.key {
// Skip metadata files
if !key.contains("/metadata/") {
total_checkpoints += 1;
total_bytes += object.size.unwrap_or(0) as u64;
}
}
}
}
if response.is_truncated.unwrap_or(false) {
continuation_token = response.next_continuation_token;
} else {
break;
}
}
let avg_checkpoint_size = if total_checkpoints > 0 {
total_bytes / total_checkpoints
} else {
0
};
Ok(StorageStats {
total_checkpoints,
total_bytes,
available_bytes: u64::MAX, // S3 has virtually unlimited storage
avg_checkpoint_size,
backend_type: format!("s3:{}", self.bucket_name),
})
}
}