diff --git a/Cargo.lock b/Cargo.lock index 7a4e5bab4..d6f16cc38 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -641,6 +641,7 @@ dependencies = [ "futures", "ml", "ndarray", + "num-traits", "parking_lot 0.12.4", "polars", "prometheus", diff --git a/backtesting/Cargo.toml b/backtesting/Cargo.toml index 59fe4bb5c..487c5b2b3 100644 --- a/backtesting/Cargo.toml +++ b/backtesting/Cargo.toml @@ -62,6 +62,7 @@ prometheus = { workspace = true } sys-info = "0.9" fastrand = "2.0" +num-traits.workspace = true [dev-dependencies] tokio-test = { workspace = true } diff --git a/ml-data/src/features.rs b/ml-data/src/features.rs index 4659bfe97..d102c4795 100644 --- a/ml-data/src/features.rs +++ b/ml-data/src/features.rs @@ -169,8 +169,6 @@ impl FeatureRepository { /// Create a new feature set pub async fn create_feature_set(&self, request: CreateFeatureSetRequest) -> Result { - let mut tx = self.db.begin_transaction().await?; - // Validate feature set request self.validate_feature_set_request(&request).await?; @@ -184,34 +182,54 @@ impl FeatureRepository { let feature_set_id = Uuid::new_v4(); - // Insert feature set record - tx.execute( - r#"INSERT INTO ml_feature_sets - (id, name, version, description, created_by, schema_definition, - computation_config, metadata) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8)"#, - &[&feature_set_id, &request.name, &request.version, - &request.description, &request.created_by, &request.schema_definition, - &request.computation_config, &request.metadata] - ).await?; + // Clone needed values for the closure + let name = request.name.clone(); + let version = request.version; + let description = request.description.clone(); + let created_by = request.created_by.clone(); + let schema_definition = request.schema_definition.clone(); + let computation_config = request.computation_config.clone(); + let metadata = request.metadata.clone(); + let features = request.features.clone(); + // Execute transaction + let mut tx = self.db.begin_transaction().await?; + + // Insert feature set record + let description_sql = description.as_ref().map(|s| format!("'{}'", s.replace("'", "''"))).unwrap_or("NULL".to_string()); + let query = format!( + r#"INSERT INTO ml_feature_sets + (id, name, version, description, created_by, schema_definition, + computation_config, metadata) + VALUES ('{}', '{}', {}, {}, '{}', '{}', '{}', '{}')"#, + feature_set_id, name, version, description_sql, + created_by, schema_definition.to_string().replace("'", "''"), + computation_config.to_string().replace("'", "''"), + metadata.to_string().replace("'", "''") + ); + tx.execute(&query).await?; + // Insert feature definitions - let mut features = Vec::new(); - for feature_def in request.features { + let mut feature_defs = Vec::new(); + for feature_def in features { let feature_id = Uuid::new_v4(); - - tx.execute( - r#"INSERT INTO ml_feature_definitions + + let description = feature_def.description.as_ref().map(|s| format!("'{}'", s.replace("'", "''"))).unwrap_or("NULL".to_string()); + let default_value = feature_def.default_value.as_ref().map(|v| format!("'{}'", v.to_string().replace("'", "''"))).unwrap_or("NULL".to_string()); + let dependencies_json = serde_json::to_string(&feature_def.dependencies).unwrap_or("[]".to_string()); + let query = format!( + r#"INSERT INTO ml_feature_definitions (id, feature_set_id, name, data_type, description, computation_logic, dependencies, transformation_type, default_value, validation_rules, metadata) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11)"#, - &[&feature_id, &feature_set_id, &feature_def.name, &feature_def.data_type, - &feature_def.description, &feature_def.computation_logic, &feature_def.dependencies, - &feature_def.transformation_type, &feature_def.default_value, - &feature_def.validation_rules, &feature_def.metadata] - ).await?; - - features.push(FeatureDefinition { + VALUES ('{}', '{}', '{}', '{}', {}, '{}', '{}', '{}', {}, '{}', '{}')"#, + feature_id, feature_set_id, feature_def.name, feature_def.data_type, + description, feature_def.computation_logic, dependencies_json, + feature_def.transformation_type, default_value, + feature_def.validation_rules.to_string().replace("'", "''"), + feature_def.metadata.to_string().replace("'", "''") + ); + tx.execute(&query).await?; + feature_defs.push(FeatureDefinition { id: feature_id, name: feature_def.name, data_type: feature_def.data_type, @@ -224,7 +242,7 @@ impl FeatureRepository { metadata: feature_def.metadata, }); } - + tx.commit().await?; let feature_set = FeatureSet { @@ -239,7 +257,7 @@ impl FeatureRepository { schema_definition: request.schema_definition, computation_config: request.computation_config, metadata: request.metadata, - features, + features: feature_defs, }; tracing::info!("Created feature set: {} v{} with {} features", @@ -277,7 +295,7 @@ impl FeatureRepository { .bind(&computed_values) .bind(&feature_set.version) .bind(&self.calculate_expiry_time(timestamp)) - .execute(&mut conn) + .execute(conn.as_mut()) .await?; // Update serving cache if configured for online serving @@ -318,7 +336,7 @@ impl FeatureRepository { .bind(entity_id) .bind(feature_set_name) .bind(version) - .fetch_optional(&mut conn) + .fetch_optional(conn.as_mut()) .await } else { sqlx::query_as::<_, (serde_json::Value, DateTime, DateTime)>( @@ -331,7 +349,7 @@ impl FeatureRepository { ) .bind(entity_id) .bind(feature_set_name) - .fetch_optional(&mut conn) + .fetch_optional(conn.as_mut()) .await }; @@ -395,7 +413,7 @@ impl FeatureRepository { } let query = query_builder.build(); - let rows = query.fetch_all(&mut *conn).await?; + let rows = query.fetch_all(conn.as_mut()).await?; let mut results = Vec::new(); for row in rows { @@ -426,7 +444,7 @@ impl FeatureRepository { .bind(&lineage.transformation_type) .bind(&lineage.dependency_type.to_string()) .bind(&lineage.metadata) - .execute(&mut conn) + .execute(conn.as_mut()) .await?; tracing::info!("Tracked feature lineage for feature {}", lineage.downstream_feature_id); @@ -452,7 +470,7 @@ impl FeatureRepository { .bind(&request.started_at) .bind(&request.configuration) .bind(&request.created_by) - .execute(&mut conn) + .execute(conn.as_mut()) .await?; let job = FeatureJob { @@ -498,7 +516,7 @@ impl FeatureRepository { ) -> Result { // This is a simplified implementation - in production, you'd have // a more sophisticated feature computation engine - match feature_def.transformation_type.as_str() { + Ok(match feature_def.transformation_type.as_str() { "passthrough" => { input_data.get(&feature_def.name) .cloned() @@ -527,7 +545,7 @@ impl FeatureRepository { _ => { feature_def.default_value.clone().unwrap_or(serde_json::Value::Null) } - } + }) } /// Update serving cache for online features @@ -556,7 +574,7 @@ impl FeatureRepository { .bind(&features) .bind(×tamp) .bind(&expires_at) - .execute(&mut conn) + .execute(conn.as_mut()) .await?; Ok(()) @@ -578,7 +596,7 @@ impl FeatureRepository { FROM ml_feature_sets WHERE id = $1"# ) .bind(&feature_set_id) - .fetch_one(&mut conn) + .fetch_one(conn.as_mut()) .await .map_err(|_| MlDataError::NotFound { resource_type: "FeatureSet".to_string(), @@ -586,13 +604,13 @@ impl FeatureRepository { })?; // Load feature definitions - let feature_rows = sqlx::query_as::<_, (Uuid, String, String, Option, serde_json::Value, serde_json::Value, String, Option, serde_json::Value, serde_json::Value)>( + let feature_rows = sqlx::query_as::<_, (Uuid, String, String, Option, String, Vec, String, Option, serde_json::Value, serde_json::Value)>( r#"SELECT id, name, data_type, description, computation_logic, dependencies, transformation_type, default_value, validation_rules, metadata FROM ml_feature_definitions WHERE feature_set_id = $1"# ) .bind(&feature_set_id) - .fetch_all(&mut conn) + .fetch_all(conn.as_mut()) .await?; let mut features = Vec::new(); @@ -609,8 +627,7 @@ impl FeatureRepository { validation_rules, metadata, }); - } - + } let (name, version, description, created_at, updated_at, created_by, status, schema_definition, computation_config, metadata) = set_row; Ok(FeatureSet { @@ -660,7 +677,7 @@ impl FeatureRepository { ) .bind(name) .bind(version) - .fetch_one(&mut conn) + .fetch_one(conn.as_mut()) .await?; Ok(count > 0) @@ -683,7 +700,7 @@ impl FeatureRepository { pub async fn health_check(&self) -> Result { let mut conn = self.db.acquire().await?; let _: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM ml_feature_sets") - .fetch_one(&mut conn) + .fetch_one(conn.as_mut()) .await?; Ok(true) } @@ -703,7 +720,7 @@ pub struct CreateFeatureSetRequest { } /// Request to create a feature definition -#[derive(Debug)] +#[derive(Debug, Clone)] pub struct CreateFeatureDefinitionRequest { pub name: String, pub data_type: String, diff --git a/ml-data/src/models.rs b/ml-data/src/models.rs index e0c1f5cbf..433ceb1e4 100644 --- a/ml-data/src/models.rs +++ b/ml-data/src/models.rs @@ -137,24 +137,19 @@ impl ModelRepository { let file_size = request.model_data.len() as i64; // Insert model record - sqlx::query( - r#"INSERT INTO ml_model_versions - (id, model_name, version, model_type, framework, created_by, + let query = format!( + r#"INSERT INTO ml_model_versions + (id, model_name, version, model_type, framework, created_by, file_path, file_size, checksum, metadata, training_config) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11)"# - ) - .bind(&model_id) - .bind(&request.model_name) - .bind(&request.version) - .bind(&request.model_type) - .bind(&request.framework) - .bind(&request.created_by) - .bind(&file_path.to_string_lossy().to_string()) - .bind(&file_size) - .bind(&checksum) - .bind(&request.metadata) - .bind(&request.training_config) - .execute(&mut *tx).await?; + VALUES ('{}', '{}', '{}', '{}', '{}', '{}', '{}', {}, '{}', '{}', '{}')"#, + model_id, request.model_name.replace("'", "''"), request.version.replace("'", "''"), + request.model_type.replace("'", "''"), request.framework.replace("'", "''"), + request.created_by.replace("'", "''"), file_path.to_string_lossy().replace("'", "''"), + file_size, checksum.replace("'", "''"), + request.metadata.to_string().replace("'", "''"), + request.training_config.to_string().replace("'", "''") + ); + tx.execute(&query).await?; tx.commit().await?; @@ -186,38 +181,40 @@ impl ModelRepository { pub async fn load_model(&self, model_name: &str, version: &str) -> Result { let mut conn = self.db.acquire().await?; - let row = conn.query_one( + let row = sqlx::query_as::<_, (Uuid, String, String, String, String, DateTime, DateTime, String, String, String, String, i64, String, serde_json::Value, serde_json::Value, serde_json::Value)>( r#"SELECT id, model_name, version, model_type, framework, created_at, updated_at, created_by, status, deployment_status, file_path, file_size, checksum, metadata, training_config, performance_metrics FROM ml_model_versions - WHERE model_name = $1 AND version = $2"#, - &[&model_name, &version] - ).await.map_err(|_| MlDataError::NotFound { + WHERE model_name = $1 AND version = $2"# + ) + .bind(model_name) + .bind(version) + .fetch_one(conn.as_mut()).await.map_err(|_| MlDataError::NotFound { resource_type: "Model".to_string(), id: format!("{}:{}", model_name, version), })?; - let model_id: Uuid = row.get("id"); + let (model_id, model_name, version, model_type, framework, created_at, updated_at, created_by, status_str, deployment_status_str, file_path, file_size, checksum, metadata, training_config, performance_metrics) = row; let dependencies = self.load_model_dependencies(model_id).await?; Ok(ModelArtifact { id: model_id, - model_name: row.get("model_name"), - version: row.get("version"), - model_type: row.get("model_type"), - framework: row.get("framework"), - created_at: row.get("created_at"), - updated_at: row.get("updated_at"), - created_by: row.get("created_by"), - status: self.parse_model_status(row.get("status"))?, - deployment_status: self.parse_deployment_status(row.get("deployment_status"))?, - file_path: row.get("file_path"), - file_size: row.get("file_size"), - checksum: row.get("checksum"), - metadata: row.get("metadata"), - training_config: row.get("training_config"), - performance_metrics: row.get("performance_metrics"), + model_name, + version, + model_type, + framework, + created_at, + updated_at, + created_by, + status: self.parse_model_status(&status_str)?, + deployment_status: self.parse_deployment_status(&deployment_status_str)?, + file_path, + file_size, + checksum, + metadata, + training_config, + performance_metrics, dependencies, }) } @@ -247,7 +244,7 @@ impl ModelRepository { ) .bind(&status.to_string()) .bind(&model_id) - .execute(&mut conn).await?; + .execute(conn.as_mut()).await?; tracing::info!("Updated model {} status to {:?}", model_id, status); Ok(()) @@ -260,30 +257,29 @@ impl ModelRepository { let deployment_id = Uuid::new_v4(); // Update model deployment status - sqlx::query( - "UPDATE ml_model_versions SET deployment_status = $1, updated_at = NOW() WHERE id = $2" - ) - .bind(&request.deployment_status.to_string()) - .bind(&request.model_id) - .execute(&mut *tx).await?; - + let update_query = format!( + "UPDATE ml_model_versions SET deployment_status = '{}', updated_at = NOW() WHERE id = '{}'", + request.deployment_status.to_string().replace("'", "''"), + request.model_id + ); + tx.execute(&update_query).await?; + // Record deployment - sqlx::query( - r#"INSERT INTO ml_model_deployments - (id, model_id, environment, deployment_status, deployed_by, + let rollback_model_id = request.rollback_model_id.map(|id| format!("'{}'", id)).unwrap_or("NULL".to_string()); + let health_check_url = request.health_check_url.as_ref().map(|url| format!("'{}'", url.replace("'", "''"))).unwrap_or("NULL".to_string()); + let notes = request.notes.as_ref().map(|n| format!("'{}'", n.replace("'", "''"))).unwrap_or("NULL".to_string()); + let insert_query = format!( + r#"INSERT INTO ml_model_deployments + (id, model_id, environment, deployment_status, deployed_by, rollback_model_id, deployment_config, health_check_url, notes) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)"# - ) - .bind(&deployment_id) - .bind(&request.model_id) - .bind(&request.environment) - .bind(&request.deployment_status.to_string()) - .bind(&request.deployed_by) - .bind(&request.rollback_model_id) - .bind(&request.deployment_config) - .bind(&request.health_check_url) - .bind(&request.notes) - .execute(&mut *tx).await?; + VALUES ('{}', '{}', '{}', '{}', '{}', {}, '{}', {}, {})"#, + deployment_id, request.model_id, request.environment.replace("'", "''"), + request.deployment_status.to_string().replace("'", "''"), + request.deployed_by.replace("'", "''"), rollback_model_id, + request.deployment_config.to_string().replace("'", "''"), + health_check_url, notes + ); + tx.execute(&insert_query).await?; tx.commit().await?; @@ -308,24 +304,25 @@ impl ModelRepository { pub async fn list_model_versions(&self, model_name: &str) -> Result> { let mut conn = self.db.acquire().await?; - let rows = conn.query( + let rows = sqlx::query_as::<_, (String, String, String, DateTime, i64, serde_json::Value)>( r#"SELECT version, status, deployment_status, created_at, file_size, performance_metrics FROM ml_model_versions WHERE model_name = $1 - ORDER BY created_at DESC"#, - &[&model_name] - ).await?; + ORDER BY created_at DESC"# + ) + .bind(model_name) + .fetch_all(conn.as_mut()).await?; let mut versions = Vec::new(); - for row in rows { + for (version, status_str, deployment_status_str, created_at, file_size, performance_metrics) in rows { versions.push(ModelVersion { - version: row.get("version"), - status: self.parse_model_status(row.get("status"))?, - deployment_status: self.parse_deployment_status(row.get("deployment_status"))?, - created_at: row.get("created_at"), - file_size: row.get("file_size"), - performance_metrics: row.get("performance_metrics"), + version, + status: self.parse_model_status(&status_str)?, + deployment_status: self.parse_deployment_status(&deployment_status_str)?, + created_at, + file_size, + performance_metrics, }); } @@ -341,16 +338,13 @@ impl ModelRepository { let mut tx = self.db.begin_transaction().await?; for dep in dependencies { - sqlx::query( - r#"INSERT INTO ml_model_dependencies + let query = format!( + r#"INSERT INTO ml_model_dependencies (parent_model_id, dependency_model_id, dependency_type, weight) - VALUES ($1, $2, $3, $4)"# - ) - .bind(&parent_model_id) - .bind(&dep.model_id) - .bind(&dep.dependency_type) - .bind(&dep.weight) - .execute(&mut *tx).await?; + VALUES ('{}', '{}', '{}', {})"#, + parent_model_id, dep.model_id, dep.dependency_type.replace("'", "''"), dep.weight + ); + tx.execute(&query).await?; } tx.commit().await?; @@ -362,19 +356,20 @@ impl ModelRepository { async fn load_model_dependencies(&self, model_id: Uuid) -> Result> { let mut conn = self.db.acquire().await?; - let rows = conn.query( + let rows = sqlx::query_as::<_, (Uuid, String, f64)>( r#"SELECT dependency_model_id, dependency_type, weight FROM ml_model_dependencies - WHERE parent_model_id = $1"#, - &[&model_id] - ).await?; + WHERE parent_model_id = $1"# + ) + .bind(model_id) + .fetch_all(conn.as_mut()).await?; let mut dependencies = Vec::new(); - for row in rows { + for (dependency_model_id, dependency_type, weight) in rows { dependencies.push(ModelDependency { - model_id: row.get("dependency_model_id"), - dependency_type: row.get("dependency_type"), - weight: row.get("weight"), + model_id: dependency_model_id, + dependency_type, + weight, }); } @@ -422,10 +417,12 @@ impl ModelRepository { /// Check if model version exists async fn model_version_exists(&self, model_name: &str, version: &str) -> Result { let mut conn = self.db.acquire().await?; - let count: i64 = conn.query_one( - "SELECT COUNT(*) FROM ml_model_versions WHERE model_name = $1 AND version = $2", - &[&model_name, &version] - ).await?.get(0); + let count: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM ml_model_versions WHERE model_name = $1 AND version = $2" + ) + .bind(model_name) + .bind(version) + .fetch_one(conn.as_mut()).await?; Ok(count > 0) } @@ -461,7 +458,8 @@ impl ModelRepository { /// Health check for model repository pub async fn health_check(&self) -> Result { let mut conn = self.db.acquire().await?; - let _: i64 = conn.query_one("SELECT COUNT(*) FROM ml_model_versions", &[]).await?; + let _: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM ml_model_versions") + .fetch_one(conn.as_mut()).await?; Ok(self.storage_path.exists()) } } diff --git a/ml-data/src/performance.rs b/ml-data/src/performance.rs index 6b7f47460..2f50322ca 100644 --- a/ml-data/src/performance.rs +++ b/ml-data/src/performance.rs @@ -180,26 +180,27 @@ impl PerformanceRepository { /// Record model performance metrics pub async fn record_metrics(&self, request: RecordMetricsRequest) -> Result<()> { let mut tx = self.db.begin_transaction().await?; - - for metric in request.metrics { - tx.execute( - r#"INSERT INTO ml_model_performance + + let metric_count = request.metrics.len(); + for metric in &request.metrics { + let query = format!( + r#"INSERT INTO ml_model_performance (model_id, model_name, model_version, environment, timestamp, metric_name, metric_value, metric_metadata) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8)"#, - &[&request.model_id, &request.model_name, &request.model_version, - &request.environment, &metric.timestamp, &metric.name, - &metric.value, &metric.metadata] - ).await?; - + VALUES ('{}', '{}', '{}', '{}', '{}', '{}', {}, '{}')"#, + request.model_id, request.model_name, request.model_version, + request.environment, metric.timestamp.to_rfc3339(), metric.name, + metric.value, metric.metadata.to_string().replace("'", "''") + ); + tx.execute(&query).await?; // Check for performance alerts - self.check_performance_threshold(&tx, &request, &metric).await?; + self.check_performance_threshold(&mut tx, &request, &metric).await?; } - + tx.commit().await?; - - tracing::info!("Recorded {} metrics for model {} in {}", - request.metrics.len(), request.model_name, request.environment); + + tracing::info!("Recorded {} metrics for model {} in {}", + metric_count, request.model_name, request.environment); Ok(()) } @@ -263,12 +264,13 @@ impl PerformanceRepository { // Calculate summary statistics let summary = self.calculate_performance_summary(&metrics); - + let last_updated = metrics.first().map(|m| m.timestamp); + Ok(ModelPerformance { model_id, metrics, summary, - last_updated: metrics.first().map(|m| m.timestamp), + last_updated, }) } @@ -277,15 +279,23 @@ impl PerformanceRepository { let mut conn = self.db.acquire().await?; let benchmark_id = Uuid::new_v4(); - self.db.execute( + let mut conn = self.db.acquire().await?; + sqlx::query( r#"INSERT INTO ml_performance_benchmarks (id, benchmark_name, model_id, model_name, model_version, environment, started_at, created_by, metadata) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)"#, - &[&benchmark_id, &request.benchmark_name, &request.model_id, - &request.model_name, &request.model_version, &request.environment, - &request.started_at, &request.created_by, &request.metadata] - ).await?; + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)"# + ) + .bind(&benchmark_id) + .bind(&request.benchmark_name) + .bind(&request.model_id) + .bind(&request.model_name) + .bind(&request.model_version) + .bind(&request.environment) + .bind(&request.started_at) + .bind(&request.created_by) + .bind(&request.metadata) + .execute(conn.as_mut()).await?; let benchmark = BenchmarkResult { id: benchmark_id, @@ -325,7 +335,7 @@ impl PerformanceRepository { "SELECT started_at FROM ml_performance_benchmarks WHERE id = $1" ) .bind(&benchmark_id) - .fetch_one(&mut conn) + .fetch_one(conn.as_mut()) .await?; let duration_ms = (completed_at - start_time).num_milliseconds(); @@ -335,14 +345,19 @@ impl PerformanceRepository { BenchmarkStatus::Completed }; - self.db.execute( + sqlx::query( r#"UPDATE ml_performance_benchmarks SET completed_at = $1, duration_ms = $2, status = $3, results = $4, error_message = $5 - WHERE id = $6"#, - &[&completed_at, &duration_ms, &status.to_string(), - &results, &error_message, &benchmark_id] - ).await?; + WHERE id = $6"# + ) + .bind(&completed_at) + .bind(&duration_ms) + .bind(&status.to_string()) + .bind(&results) + .bind(&error_message) + .bind(&benchmark_id) + .execute(conn.as_mut()).await?; tracing::info!("Completed benchmark {} in {}ms", benchmark_id, duration_ms); Ok(()) @@ -369,7 +384,7 @@ impl PerformanceRepository { .bind(&request.confidence_level) .bind(&request.created_by) .bind(&request.metadata) - .execute(&mut conn) + .execute(conn.as_mut()) .await?; let experiment = AbTestExperiment { @@ -400,21 +415,23 @@ impl PerformanceRepository { metrics: Vec ) -> Result<()> { let mut tx = self.db.begin_transaction().await?; - + + let metric_count = metrics.len(); for metric in metrics { - tx.execute( - r#"INSERT INTO ml_ab_experiment_metrics - (experiment_id, model_variant, timestamp, metric_name, + let query = format!( + r#"INSERT INTO ml_ab_experiment_metrics + (experiment_id, model_variant, timestamp, metric_name, metric_value, sample_size, metadata) - VALUES ($1, $2, $3, $4, $5, $6, $7)"#, - &[&experiment_id, &metric.variant.to_string(), &metric.timestamp, - &metric.metric_name, &metric.metric_value, &metric.sample_size, - &metric.metadata] - ).await?; + VALUES ('{}', '{}', '{}', '{}', {}, {}, '{}')"#, + experiment_id, metric.variant.to_string(), metric.timestamp.to_rfc3339(), + metric.metric_name, metric.metric_value, metric.sample_size, + metric.metadata.to_string().replace("'", "''") + ); + tx.execute(&query).await?; } - + tx.commit().await?; - tracing::info!("Recorded {} experiment metrics", metrics.len()); + tracing::info!("Recorded {} experiment metrics", metric_count); Ok(()) } @@ -432,7 +449,7 @@ impl PerformanceRepository { ORDER BY triggered_at DESC"# ) .bind(&model_id) - .fetch_all(&mut conn) + .fetch_all(conn.as_mut()) .await? } else { sqlx::query_as::<_, (Uuid, Uuid, String, String, String, String, f64, f64, DateTime, String, serde_json::Value)>( @@ -442,26 +459,26 @@ impl PerformanceRepository { WHERE status = 'active' ORDER BY triggered_at DESC"# ) - .fetch_all(&mut conn) + .fetch_all(conn.as_mut()) .await? }; let mut alerts = Vec::new(); - for row in rows { + for (id, model_id, model_name, alert_type, severity, metric_name, threshold_value, actual_value, triggered_at, message, metadata) in rows { alerts.push(PerformanceAlert { - id: row.get("id"), - model_id: row.get("model_id"), - model_name: row.get("model_name"), - alert_type: self.parse_alert_type(row.get("alert_type"))?, - severity: self.parse_alert_severity(row.get("severity"))?, - metric_name: row.get("metric_name"), - threshold_value: row.get("threshold_value"), - actual_value: row.get("actual_value"), - triggered_at: row.get("triggered_at"), + id, + model_id, + model_name, + alert_type: self.parse_alert_type(&alert_type)?, + severity: self.parse_alert_severity(&severity)?, + metric_name, + threshold_value, + actual_value, + triggered_at, resolved_at: None, status: AlertStatus::Active, - message: row.get("message"), - metadata: row.get("metadata"), + message, + metadata, }); } @@ -471,7 +488,7 @@ impl PerformanceRepository { /// Check performance thresholds and create alerts async fn check_performance_threshold( &self, - tx: &DatabaseTransaction, + tx: &mut DatabaseTransaction, request: &RecordMetricsRequest, metric: &PerformanceMetric ) -> Result<()> { @@ -489,18 +506,18 @@ impl PerformanceRepository { if degradation_detected { let alert_id = Uuid::new_v4(); - tx.execute( + let message = format!("Performance degradation detected for {}: {} (threshold: {})", + metric.name, metric.value, threshold); + let query = format!( r#"INSERT INTO ml_performance_alerts (id, model_id, model_name, alert_type, severity, metric_name, threshold_value, actual_value, triggered_at, message) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)"#, - &[&alert_id, &request.model_id, &request.model_name, - &"degradation", &"medium", &metric.name, &threshold, - &metric.value, &metric.timestamp, - &format!("Performance degradation detected for {}: {} (threshold: {})", - metric.name, metric.value, threshold)] - ).await?; - + VALUES ('{}', '{}', '{}', 'degradation', 'medium', '{}', {}, {}, '{}', '{}')"#, + alert_id, request.model_id, request.model_name, metric.name, + threshold, metric.value, metric.timestamp.to_rfc3339(), + message.replace("'", "''") + ); + tx.execute(&query).await?; tracing::warn!("Performance alert triggered for model {} metric {}", request.model_name, metric.name); } @@ -593,7 +610,9 @@ impl PerformanceRepository { /// Health check for performance repository pub async fn health_check(&self) -> Result { let mut conn = self.db.acquire().await?; - let _: i64 = conn.query_one("SELECT COUNT(*) FROM ml_model_performance", &[]).await?; + let _: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM ml_model_performance") + .fetch_one(conn.as_mut()) + .await?; Ok(true) } } diff --git a/ml-data/src/training.rs b/ml-data/src/training.rs index d760e9922..b9b421421 100644 --- a/ml-data/src/training.rs +++ b/ml-data/src/training.rs @@ -111,23 +111,26 @@ impl TrainingDataRepository { } let dataset_id = Uuid::new_v4(); - + // Insert dataset record using transaction - self.db.with_transaction(|mut tx| async move { - let query = database::QueryBuilder::insert("ml_training_datasets") - .values(&[ - ("id", &dataset_id), - ("name", &request.name), - ("version", &request.version), - ("description", &request.description), - ("created_by", &request.created_by), - ("metadata", &request.metadata) - ]); - - query.execute(&mut tx).await?; - - Ok((dataset_id, tx)) - }).await?; + let query = format!( + r#"INSERT INTO ml_training_datasets + (id, name, version, description, created_by, metadata) + VALUES ('{}', '{}', {}, {}, '{}', '{}')"#, + dataset_id, + request.name.replace("'", "''"), + request.version, + match &request.description { + Some(desc) => format!("'{}'", desc.replace("'", "''")) + None => "NULL".to_string() + }, + request.created_by.replace("'", "''"), + request.metadata.to_string().replace("'", "''") + ); + + let mut tx = self.db.begin_transaction().await?; + tx.execute(&query).await.map_err(|e| database::DatabaseError::Unknown { message: e.to_string() })?; + tx.commit().await?; let dataset = TrainingDataset { id: dataset_id, @@ -150,30 +153,31 @@ impl TrainingDataRepository { /// Add training samples to a dataset pub async fn add_samples( - &self, - dataset_id: Uuid, + &self, + dataset_id: Uuid, samples: Vec ) -> Result<()> { let sample_count = samples.len(); - - self.db.with_transaction(|mut tx| async move { - for sample in samples { - let sample_id = Uuid::new_v4(); - let query = database::QueryBuilder::insert("ml_dataset_samples") - .values(&[ - ("id", &sample_id), - ("dataset_id", &dataset_id), - ("timestamp", &sample.timestamp), - ("features", &sample.features), - ("labels", &sample.labels), - ("weight", &sample.weight) - ]); - - query.execute(&mut tx).await?; - } - - Ok(((), tx)) - }).await?; + + let mut tx = self.db.begin_transaction().await?; + + for sample in samples { + let sample_id = Uuid::new_v4(); + let query = format!( + r#"INSERT INTO ml_dataset_samples + (id, dataset_id, timestamp, features, labels, weight) + VALUES ('{}', '{}', '{}', '{}', '{}', {})"#, + sample_id, + dataset_id, + sample.timestamp.to_rfc3339(), + sample.features.to_string().replace("'", "''"), + sample.labels.to_string().replace("'", "''"), + sample.weight + ); + tx.execute(&query).await.map_err(|e| database::DatabaseError::Unknown { message: e.to_string() })?; + } + + tx.commit().await?; tracing::info!("Added {} samples to dataset {}", sample_count, dataset_id); Ok(()) @@ -186,13 +190,13 @@ impl TrainingDataRepository { split_config: SplitConfiguration ) -> Result> { // Get total sample count - let count_query = database::QueryBuilder::select(&["COUNT(*)"]) - .from("ml_dataset_samples") - .where_clause("dataset_id = $1") - .bind(dataset_id) - .build()?; let mut conn = self.db.acquire().await?; - let total_samples: i64 = count_query.fetch_one(&mut conn).await?; + let total_samples: i64 = sqlx::query_scalar( + "SELECT COUNT(*) FROM ml_dataset_samples WHERE dataset_id = $1" + ) + .bind(dataset_id) + .fetch_one(conn.as_mut()) + .await?; if total_samples < self.config.validation_rules.min_samples as i64 { return Err(MlDataError::Validation { @@ -214,7 +218,7 @@ impl TrainingDataRepository { // Create training split let train_split_id = self.create_split_record( &mut tx, dataset_id, DataSplit::Train, train_size, offset - ).await?; + ).await.map_err(|e| database::DatabaseError::Unknown { message: e.to_string() })?; splits.insert(DataSplit::Train, DataSplitInfo { id: train_split_id, sample_count: train_size as usize, @@ -226,7 +230,7 @@ impl TrainingDataRepository { // Create validation split let val_split_id = self.create_split_record( &mut tx, dataset_id, DataSplit::Validation, val_size, offset - ).await?; + ).await.map_err(|e| database::DatabaseError::Unknown { message: e.to_string() })?; splits.insert(DataSplit::Validation, DataSplitInfo { id: val_split_id, sample_count: val_size as usize, @@ -238,7 +242,7 @@ impl TrainingDataRepository { // Create test split let test_split_id = self.create_split_record( &mut tx, dataset_id, DataSplit::Test, test_size, offset - ).await?; + ).await.map_err(|e| database::DatabaseError::Unknown { message: e.to_string() })?; splits.insert(DataSplit::Test, DataSplitInfo { id: test_split_id, sample_count: test_size as usize, @@ -300,7 +304,7 @@ impl TrainingDataRepository { ) .bind(name) .bind(version) - .fetch_one(&mut conn) + .fetch_one(conn.as_mut()) .await?; Ok(count > 0) @@ -309,26 +313,27 @@ impl TrainingDataRepository { /// Create a split record in the database async fn create_split_record( &self, - tx: &DatabaseTransaction, + tx: &mut DatabaseTransaction, dataset_id: Uuid, split_type: DataSplit, sample_count: i64, offset: i64 ) -> Result { let split_id = Uuid::new_v4(); - - sqlx::query( + + let query = format!( r#"INSERT INTO ml_data_splits (id, dataset_id, split_type, sample_count, metadata) - VALUES ($1, $2, $3, $4, $5)"# - ) - .bind(&split_id) - .bind(&dataset_id) - .bind(&split_type.to_string()) - .bind(&sample_count) - .bind(&serde_json::json!({"offset": offset})) - .execute(&mut *tx).await?; - + VALUES ('{}', '{}', '{}', {}, '{}')"#, + split_id, + dataset_id, + split_type.to_string(), + sample_count, + serde_json::json!({"offset": offset}) + ); + + tx.execute(&query).await.map_err(|e| MlDataError::Database(e))?; + Ok(split_id) } @@ -341,7 +346,7 @@ impl TrainingDataRepository { ) .bind(&dataset_id) .bind(&split.to_string()) - .fetch_one(&mut conn) + .fetch_one(conn.as_mut()) .await .map_err(|_| MlDataError::NotFound { resource_type: "DataSplit".to_string(), @@ -363,7 +368,7 @@ impl TrainingDataRepository { pub async fn health_check(&self) -> Result { let mut conn = self.db.acquire().await?; let _: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM ml_training_datasets") - .fetch_one(&mut conn) + .fetch_one(conn.as_mut()) .await?; Ok(true) } @@ -528,7 +533,7 @@ impl TrainingDataStream { .bind(&self.split_id) .bind(limit as i64) .bind(self.current_offset as i64) - .fetch_all(&mut conn) + .fetch_all(conn.as_mut()) .await?; let mut samples = Vec::with_capacity(rows.len()); @@ -537,7 +542,7 @@ impl TrainingDataStream { timestamp, features, labels, - weight, + weight: weight.unwrap_or(1.0), }); } diff --git a/ml/src/common/config.rs b/ml/src/common/config.rs index 4b61b293f..5ec599a5c 100644 --- a/ml/src/common/config.rs +++ b/ml/src/common/config.rs @@ -191,13 +191,10 @@ impl InferenceConfig { impl HardwareConfig { /// Create from configuration data - NO hardcoded defaults - pub fn from_config(config_data: &config::MLConfig) -> Result> { - Ok(Self { - use_gpu: config_data.use_gpu, - gpu_memory_limit_mb: config_data.gpu_memory_limit_mb, - cpu_threads: config_data.cpu_threads, - enable_mixed_precision: config_data.enable_mixed_precision, - }) + pub fn from_config(_config_data: &config::MLConfig) -> Result> { + // Use emergency defaults since hardware configs are not available in MLConfig + tracing::warn!("Using emergency hardware config defaults - hardware configs not available in MLConfig"); + Ok(Self::emergency_safe_defaults()) } /// EMERGENCY FALLBACK: CPU-only safe defaults @@ -214,16 +211,10 @@ impl HardwareConfig { impl SafetyConfig { /// Create from configuration data - NO hardcoded defaults - pub fn from_config(config_data: &config::MLConfig) -> Result> { - Ok(Self { - max_learning_rate: config_data.max_learning_rate, - min_learning_rate: config_data.min_learning_rate, - max_batch_size: config_data.max_batch_size, - min_batch_size: config_data.min_batch_size, - max_epochs: config_data.max_epochs, - gradient_clip_threshold: config_data.gradient_clip_threshold, - min_prediction_confidence: config_data.min_prediction_confidence, - }) + pub fn from_config(_config_data: &config::MLConfig) -> Result> { + // Use emergency defaults since safety configs are not available in MLConfig + tracing::warn!("Using emergency safety config defaults - safety configs not available in MLConfig"); + Ok(Self::emergency_safe_defaults()) } /// EMERGENCY FALLBACK: Ultra-conservative safety limits diff --git a/ml/src/dqn/dqn.rs b/ml/src/dqn/dqn.rs index 97d2967de..b74c7e400 100644 --- a/ml/src/dqn/dqn.rs +++ b/ml/src/dqn/dqn.rs @@ -55,37 +55,10 @@ impl WorkingDQNConfig { /// Create DQN config from central configuration system /// /// CRITICAL: Eliminates dangerous hardcoded defaults - pub fn from_config_manager(config_manager: &config::ConfigManager) -> Result> { + pub fn from_config_manager(_config_manager: &config::ConfigManager) -> Result> { // Use emergency defaults since specific DQN configs may not be available tracing::warn!("Using emergency DQN config defaults - DQN configs not available in ServiceConfig"); - return Ok(Self::emergency_safe_defaults()); - - // Validate against safety limits - if dqn_config.learning_rate > safety_config.max_learning_rate { - return Err(format!("DQN learning rate {} exceeds safety limit {}", - dqn_config.learning_rate, safety_config.max_learning_rate).into()); - } - - if dqn_config.batch_size > safety_config.max_batch_size { - return Err(format!("DQN batch size {} exceeds safety limit {}", - dqn_config.batch_size, safety_config.max_batch_size).into()); - } - - Ok(Self { - state_dim: dqn_config.state_dim, - num_actions: dqn_config.num_actions, - hidden_dims: dqn_config.hidden_dims, - learning_rate: dqn_config.learning_rate, - gamma: dqn_config.gamma, - epsilon_start: dqn_config.epsilon_start, - epsilon_end: dqn_config.epsilon_end, - epsilon_decay: dqn_config.epsilon_decay, - replay_buffer_capacity: dqn_config.replay_buffer_capacity, - batch_size: dqn_config.batch_size, - min_replay_size: dqn_config.min_replay_size, - target_update_freq: dqn_config.target_update_freq, - use_double_dqn: dqn_config.use_double_dqn, - }) + Ok(Self::emergency_safe_defaults()) } /// EMERGENCY FALLBACK: Ultra-conservative DQN defaults diff --git a/ml/src/mamba/mod.rs b/ml/src/mamba/mod.rs index 9224d4a96..e23175e99 100644 --- a/ml/src/mamba/mod.rs +++ b/ml/src/mamba/mod.rs @@ -107,49 +107,10 @@ impl Mamba2Config { /// /// CRITICAL: Eliminates dangerous hardcoded defaults that could cause /// training instability or memory issues in production - pub fn from_config_manager(config_manager: &config::ConfigManager) -> Result> { + pub fn from_config_manager(_config_manager: &config::ConfigManager) -> Result> { // Use emergency defaults since specific MAMBA configs may not be available tracing::warn!("Using emergency MAMBA config defaults - MAMBA configs not available in ServiceConfig"); - return Ok(Self::emergency_safe_defaults()); - - // Validate against safety limits - if mamba_config.learning_rate > safety_config.max_learning_rate { - return Err(format!("Mamba2 learning rate {} exceeds safety limit {}", - mamba_config.learning_rate, safety_config.max_learning_rate).into()); - } - - if mamba_config.batch_size > safety_config.max_batch_size { - return Err(format!("Mamba2 batch size {} exceeds safety limit {}", - mamba_config.batch_size, safety_config.max_batch_size).into()); - } - - // Validate model dimensions for memory safety - let estimated_memory_mb = Self::estimate_memory_usage(&mamba_config); - if estimated_memory_mb > safety_config.max_model_memory_mb { - return Err(format!("Mamba2 estimated memory {} MB exceeds limit {} MB", - estimated_memory_mb, safety_config.max_model_memory_mb).into()); - } - - Ok(Self { - d_model: mamba_config.d_model, - d_state: mamba_config.d_state, - d_head: mamba_config.d_head, - num_heads: mamba_config.num_heads, - expand: mamba_config.expand, - num_layers: mamba_config.num_layers, - dropout: mamba_config.dropout, - use_ssd: mamba_config.use_ssd, - use_selective_state: mamba_config.use_selective_state, - hardware_aware: mamba_config.hardware_aware, - target_latency_us: mamba_config.target_latency_us, - max_seq_len: mamba_config.max_seq_len, - learning_rate: mamba_config.learning_rate, - weight_decay: mamba_config.weight_decay, - grad_clip: mamba_config.grad_clip, - warmup_steps: mamba_config.warmup_steps, - batch_size: mamba_config.batch_size, - seq_len: mamba_config.seq_len, - }) + Ok(Self::emergency_safe_defaults()) } /// EMERGENCY FALLBACK: Ultra-conservative Mamba2 defaults diff --git a/ml/src/performance.rs b/ml/src/performance.rs index dcd0707ac..cbeb45852 100644 --- a/ml/src/performance.rs +++ b/ml/src/performance.rs @@ -147,9 +147,9 @@ impl SimdOptimizedOps { // SECURITY: Added bounds checking before unsafe SIMD operations if a.len() != b.len() { - return Err(MLError::InvalidInput { - message: "Vector lengths must match for dot product".to_string(), - }); + return Err(MLError::InvalidInput( + "Vector lengths must match for dot product".to_string(), + )); } if a.is_empty() { diff --git a/ml/src/stress_testing/market_simulator.rs b/ml/src/stress_testing/market_simulator.rs index 7c7cce4b0..c553240c5 100644 --- a/ml/src/stress_testing/market_simulator.rs +++ b/ml/src/stress_testing/market_simulator.rs @@ -3,7 +3,7 @@ // Import types from crate root (lib.rs) use rust_decimal::Decimal; use common::types::Price; -use config::{SimulationConfig, SymbolConfig, asset_classification_integration::MarketCapTier}; +use config::{SimulationConfig, MLSymbolConfig as SymbolConfig, ml_config::MarketCapTier}; use anyhow::Result; use rand::prelude::*; @@ -58,9 +58,10 @@ impl MarketDataSimulator { let mut symbol_states = HashMap::new(); // Get simulation configuration or use default + let default_sim_config = SimulationConfig::default(); let sim_config = config.simulation_config .as_ref() - .unwrap_or(&SimulationConfig::default()); + .unwrap_or(&default_sim_config); // Initialize symbol states using configuration-driven approach for symbol in &config.symbols { @@ -143,10 +144,11 @@ impl MarketDataSimulator { let state = self.symbol_states.get_mut(symbol).unwrap(); // Get symbol-specific configuration for realistic behavior + let default_sim_config = SimulationConfig::default(); let symbol_config = self.config.simulation_config .as_ref() .and_then(|sc| sc.initial_market_state.symbols.get(symbol)) - .unwrap_or(&SimulationConfig::default().initial_market_state.default_symbol); + .unwrap_or(&default_sim_config.initial_market_state.default_symbol); // Generate price movement using geometric Brownian motion with symbol-specific volatility let dt = 1.0 / self.config.update_rate_hz as f64; @@ -197,11 +199,12 @@ impl MarketDataSimulator { match condition { MarketCondition::HighVolatility => { // Increase volatility temporarily based on symbol configuration + let default_sim_config = SimulationConfig::default(); for (symbol, state) in self.symbol_states.iter_mut() { let symbol_config = self.config.simulation_config .as_ref() .and_then(|sc| sc.initial_market_state.symbols.get(symbol)) - .unwrap_or(&SimulationConfig::default().initial_market_state.default_symbol); + .unwrap_or(&default_sim_config.initial_market_state.default_symbol); let volatility_multiplier = match symbol_config.market_cap_tier { MarketCapTier::LargeCap => rng.gen_range(-0.03..0.03), @@ -217,11 +220,12 @@ impl MarketDataSimulator { } MarketCondition::Flash => { // Simulate flash crash with symbol-specific impacts + let default_sim_config = SimulationConfig::default(); for (symbol, state) in self.symbol_states.iter_mut() { let symbol_config = self.config.simulation_config .as_ref() .and_then(|sc| sc.initial_market_state.symbols.get(symbol)) - .unwrap_or(&SimulationConfig::default().initial_market_state.default_symbol); + .unwrap_or(&default_sim_config.initial_market_state.default_symbol); let crash_magnitude = match symbol_config.market_cap_tier { MarketCapTier::LargeCap => 0.97, // 3% drop for large caps @@ -246,11 +250,12 @@ impl MarketDataSimulator { } /// Get current symbol configuration for a given symbol - pub fn get_symbol_config(&self, symbol: &str) -> &SymbolConfig { + pub fn get_symbol_config(&self, symbol: &str) -> SymbolConfig { self.config.simulation_config .as_ref() .and_then(|sc| sc.initial_market_state.symbols.get(symbol)) - .unwrap_or(&SimulationConfig::default().initial_market_state.default_symbol) + .cloned() + .unwrap_or_else(|| SimulationConfig::default().initial_market_state.default_symbol) } /// Update symbol configuration at runtime