- 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>
1147 lines
37 KiB
Rust
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),
|
|
})
|
|
}
|
|
}
|