use candle_core::{Device, Tensor};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::net::SocketAddr;
use std::path::Path;
use std::sync::Arc;
use tokio::sync::RwLock;
use crate::cached_loader::CachedModelLoader;
use crate::LoadedModel;
use anyhow::{anyhow, Result};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DistributedConfig {
pub nodes: Vec<NodeConfig>,
pub sharding_strategy: ShardingStrategy,
pub device_placement: DevicePlacementConfig,
pub communication: CommunicationConfig,
pub load_balancing: LoadBalancingConfig,
pub fault_tolerance: FaultToleranceConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NodeConfig {
pub node_id: String,
pub address: SocketAddr,
pub devices: Vec<DeviceInfo>,
pub capabilities: NodeCapabilities,
pub role: NodeRole,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DeviceInfo {
pub device_id: String,
pub device_type: DeviceType,
pub memory_bytes: u64,
pub compute_score: f32,
pub utilization: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub enum DeviceType {
Cpu,
Cuda,
Metal,
Custom(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NodeCapabilities {
pub total_memory: u64,
pub network_bandwidth: u64,
pub storage_capacity: u64,
pub supported_dtypes: Vec<String>,
pub special_capabilities: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub enum NodeRole {
Coordinator,
Worker,
Storage,
Hybrid(Vec<NodeRole>),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ShardingStrategy {
NoSharding,
LayerSharding {
layers_per_shard: usize,
},
DimensionSharding {
dimension: ShardingDimension,
shard_size: usize,
},
PipelineSharding {
num_stages: usize,
},
TensorSharding {
degree: usize,
},
Custom {
placement_fn: String, },
ModalitySpecific {
modality_assignments: HashMap<crate::multimodal::Modality, Vec<String>>,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ShardingDimension {
Rows,
Columns,
Batch,
Sequence,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DevicePlacementConfig {
pub strategy: PlacementStrategy,
pub constraints: PlacementConstraints,
pub memory_allocation: MemoryAllocationConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum PlacementStrategy {
RoundRobin,
MemoryBased,
ComputeBased,
LoadBalanced,
LocalityAware,
Custom(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PlacementConstraints {
pub min_memory_per_device: u64,
pub max_devices: Option<usize>,
pub preferred_device_types: Vec<DeviceType>,
pub colocation_groups: Vec<Vec<String>>,
pub anti_affinity_groups: Vec<Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MemoryAllocationConfig {
pub strategy: MemoryAllocationStrategy,
pub system_reserve_percent: f32,
pub enable_memory_pooling: bool,
pub fragmentation_threshold: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum MemoryAllocationStrategy {
Eager,
Lazy,
Pooled,
Streaming,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CommunicationConfig {
pub protocol: CommunicationProtocol,
pub timeouts: TimeoutConfig,
pub compression: CompressionConfig,
pub security: SecurityConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum CommunicationProtocol {
Http,
Grpc,
Tcp,
Udp,
Nccl,
Mpi,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TimeoutConfig {
pub connection_timeout_ms: u64,
pub request_timeout_ms: u64,
pub heartbeat_interval_ms: u64,
pub failure_detection_timeout_ms: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompressionConfig {
pub enable_compression: bool,
pub algorithm: CompressionAlgorithm,
pub level: u32,
pub min_compress_size: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum CompressionAlgorithm {
None,
Lz4,
Zstd,
Gzip,
Custom(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SecurityConfig {
pub enable_tls: bool,
pub cert_path: Option<String>,
pub key_path: Option<String>,
pub auth_method: AuthMethod,
pub api_key: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum AuthMethod {
None,
ApiKey,
Token,
MutualTls,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LoadBalancingConfig {
pub strategy: LoadBalancingStrategy,
pub health_check: HealthCheckConfig,
pub routing: RoutingConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum LoadBalancingStrategy {
RoundRobin,
LeastConnections,
WeightedRoundRobin,
ResponseTime,
ResourceBased,
Custom(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HealthCheckConfig {
pub enabled: bool,
pub interval_ms: u64,
pub timeout_ms: u64,
pub failure_threshold: u32,
pub success_threshold: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RoutingConfig {
pub sticky_sessions: bool,
pub affinity_method: AffinityMethod,
pub retry: RetryConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum AffinityMethod {
None,
ClientIp,
SessionToken,
ModelBased,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RetryConfig {
pub max_retries: u32,
pub initial_delay_ms: u64,
pub backoff_multiplier: f32,
pub max_delay_ms: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FaultToleranceConfig {
pub replication: ReplicationConfig,
pub failure_handling: FailureHandlingConfig,
pub recovery: RecoveryConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReplicationConfig {
pub enabled: bool,
pub factor: u32,
pub strategy: ReplicationStrategy,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ReplicationStrategy {
Synchronous,
Asynchronous,
Quorum { min_replicas: u32 },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FailureHandlingConfig {
pub auto_failover: bool,
pub failover_timeout_ms: u64,
pub circuit_breaker: CircuitBreakerConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CircuitBreakerConfig {
pub enabled: bool,
pub failure_threshold: u32,
pub recovery_timeout_ms: u64,
pub half_open_requests: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RecoveryConfig {
pub auto_recovery: bool,
pub strategy: RecoveryStrategy,
pub recovery_timeout_ms: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum RecoveryStrategy {
RestartNodes,
RedistributeShards,
ScaleOut,
Manual,
}
#[derive(Debug, Clone)]
pub struct DistributedModel {
pub metadata: DistributedModelMetadata,
pub shards: Vec<ModelShard>,
pub active_nodes: HashMap<String, NodeConfig>,
pub status: DeploymentStatus,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DistributedModelMetadata {
pub model_id: String,
pub total_size: u64,
pub num_shards: usize,
pub sharding_strategy: ShardingStrategy,
pub created_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug, Clone)]
pub struct ModelShard {
pub shard_id: String,
pub node_id: String,
pub device_id: String,
pub size_bytes: u64,
pub loaded: bool,
pub status: ShardStatus,
}
#[derive(Debug, Clone, PartialEq)]
pub enum ShardStatus {
NotLoaded,
Loading,
Ready,
Failed(String),
Migrating,
}
#[derive(Debug, Clone, PartialEq)]
pub enum DeploymentStatus {
Planning,
Deploying,
Ready,
Failed(String),
Updating,
ShuttingDown,
}
impl Default for DistributedConfig {
fn default() -> Self {
Self {
nodes: Vec::new(),
sharding_strategy: ShardingStrategy::NoSharding,
device_placement: DevicePlacementConfig::default(),
communication: CommunicationConfig::default(),
load_balancing: LoadBalancingConfig::default(),
fault_tolerance: FaultToleranceConfig::default(),
}
}
}
impl Default for DevicePlacementConfig {
fn default() -> Self {
Self {
strategy: PlacementStrategy::LoadBalanced,
constraints: PlacementConstraints::default(),
memory_allocation: MemoryAllocationConfig::default(),
}
}
}
impl Default for PlacementConstraints {
fn default() -> Self {
Self {
min_memory_per_device: 1024 * 1024 * 1024, max_devices: None,
preferred_device_types: vec![DeviceType::Cuda, DeviceType::Cpu],
colocation_groups: Vec::new(),
anti_affinity_groups: Vec::new(),
}
}
}
impl Default for MemoryAllocationConfig {
fn default() -> Self {
Self {
strategy: MemoryAllocationStrategy::Lazy,
system_reserve_percent: 0.1, enable_memory_pooling: true,
fragmentation_threshold: 0.2,
}
}
}
impl Default for CommunicationConfig {
fn default() -> Self {
Self {
protocol: CommunicationProtocol::Http,
timeouts: TimeoutConfig::default(),
compression: CompressionConfig::default(),
security: SecurityConfig::default(),
}
}
}
impl Default for TimeoutConfig {
fn default() -> Self {
Self {
connection_timeout_ms: 30_000,
request_timeout_ms: 60_000,
heartbeat_interval_ms: 10_000,
failure_detection_timeout_ms: 30_000,
}
}
}
impl Default for CompressionConfig {
fn default() -> Self {
Self {
enable_compression: true,
algorithm: CompressionAlgorithm::Lz4,
level: 1,
min_compress_size: 1024, }
}
}
impl Default for SecurityConfig {
fn default() -> Self {
Self {
enable_tls: false,
cert_path: None,
key_path: None,
auth_method: AuthMethod::None,
api_key: None,
}
}
}
impl Default for LoadBalancingConfig {
fn default() -> Self {
Self {
strategy: LoadBalancingStrategy::RoundRobin,
health_check: HealthCheckConfig::default(),
routing: RoutingConfig::default(),
}
}
}
impl Default for HealthCheckConfig {
fn default() -> Self {
Self {
enabled: true,
interval_ms: 30_000,
timeout_ms: 5_000,
failure_threshold: 3,
success_threshold: 2,
}
}
}
impl Default for RoutingConfig {
fn default() -> Self {
Self {
sticky_sessions: false,
affinity_method: AffinityMethod::None,
retry: RetryConfig::default(),
}
}
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
max_retries: 3,
initial_delay_ms: 1000,
backoff_multiplier: 2.0,
max_delay_ms: 30_000,
}
}
}
impl Default for FaultToleranceConfig {
fn default() -> Self {
Self {
replication: ReplicationConfig::default(),
failure_handling: FailureHandlingConfig::default(),
recovery: RecoveryConfig::default(),
}
}
}
impl Default for ReplicationConfig {
fn default() -> Self {
Self {
enabled: false,
factor: 1,
strategy: ReplicationStrategy::Asynchronous,
}
}
}
impl Default for FailureHandlingConfig {
fn default() -> Self {
Self {
auto_failover: true,
failover_timeout_ms: 60_000,
circuit_breaker: CircuitBreakerConfig::default(),
}
}
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
enabled: true,
failure_threshold: 5,
recovery_timeout_ms: 60_000,
half_open_requests: 3,
}
}
}
impl Default for RecoveryConfig {
fn default() -> Self {
Self {
auto_recovery: true,
strategy: RecoveryStrategy::RedistributeShards,
recovery_timeout_ms: 300_000, }
}
}
pub struct DistributedConfigBuilder {
config: DistributedConfig,
}
impl DistributedConfigBuilder {
pub fn new() -> Self {
Self {
config: DistributedConfig::default(),
}
}
pub fn add_node(mut self, node: NodeConfig) -> Self {
self.config.nodes.push(node);
self
}
pub fn sharding_strategy(mut self, strategy: ShardingStrategy) -> Self {
self.config.sharding_strategy = strategy;
self
}
pub fn device_placement(mut self, placement: DevicePlacementConfig) -> Self {
self.config.device_placement = placement;
self
}
pub fn communication(mut self, communication: CommunicationConfig) -> Self {
self.config.communication = communication;
self
}
pub fn load_balancing(mut self, load_balancing: LoadBalancingConfig) -> Self {
self.config.load_balancing = load_balancing;
self
}
pub fn fault_tolerance(mut self, fault_tolerance: FaultToleranceConfig) -> Self {
self.config.fault_tolerance = fault_tolerance;
self
}
pub fn build(self) -> DistributedConfig {
self.config
}
}
impl Default for DistributedConfigBuilder {
fn default() -> Self {
Self::new()
}
}