🔧 REFACTOR: Convert ml-data to direct sqlx queries and fix transaction patterns

- Changed all repositories from DatabasePool to Database
- Fixed transaction handling (conn.begin() -> db.begin_transaction())
- Converted to direct sqlx::query() calls
- Fixed field references (pool -> db)
- Partial resolution of compilation errors (ongoing work)
This commit is contained in:
jgrusewski
2025-09-30 07:56:11 +02:00
parent 3371fea4b9
commit 481667e8e5
10 changed files with 385 additions and 293 deletions

View File

@@ -9,6 +9,7 @@ use uuid::Uuid;
use crate::{MlDataError, Result, FeatureStoreConfig};
use database::{Database, DatabaseTransaction};
use sqlx::{Acquire, Executor, Row};
/// Feature repository for ML feature engineering and serving
#[derive(Clone)]
@@ -168,8 +169,7 @@ impl FeatureRepository {
/// Create a new feature set
pub async fn create_feature_set(&self, request: CreateFeatureSetRequest) -> Result<FeatureSet> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
let mut tx = self.db.begin_transaction().await?;
// Validate feature set request
self.validate_feature_set_request(&request).await?;
@@ -266,13 +266,19 @@ impl FeatureRepository {
// Store computed features
let mut conn = self.db.acquire().await?;
conn.execute(
sqlx::query(
r#"INSERT INTO ml_feature_values
(feature_set_id, entity_id, timestamp, features, version, expires_at)
VALUES ($1, $2, $3, $4, $5, $6)"#,
&[&feature_set_id, &entity_id, &timestamp, &computed_values,
&feature_set.version, &self.calculate_expiry_time(timestamp)]
).await?;
VALUES ($1, $2, $3, $4, $5, $6)"#
)
.bind(&feature_set_id)
.bind(&entity_id)
.bind(&timestamp)
.bind(&computed_values)
.bind(&feature_set.version)
.bind(&self.calculate_expiry_time(timestamp))
.execute(&mut conn)
.await?;
// Update serving cache if configured for online serving
self.update_serving_cache(
@@ -302,32 +308,44 @@ impl FeatureRepository {
let mut conn = self.db.acquire().await?;
// Using sqlx query builder pattern
let (query, params) =
if let Some(version) = feature_set_version {
(r#"SELECT features, last_updated, expires_at
FROM ml_feature_cache
WHERE entity_id = $1 AND feature_set_name = $2 AND feature_set_version = $3
AND expires_at > NOW()"#.to_string(),
vec![&entity_id, &feature_set_name, &version])
} else {
(r#"SELECT features, last_updated, expires_at
FROM ml_feature_cache
WHERE entity_id = $1 AND feature_set_name = $2
AND expires_at > NOW()
ORDER BY feature_set_version DESC
LIMIT 1"#.to_string(),
vec![&entity_id, &feature_set_name])
};
if let Ok(row) = conn.query_one(&query, &params).await {
Ok(Some(ServedFeatures {
entity_id: entity_id.to_string(),
features: row.get("features"),
last_updated: row.get("last_updated"),
expires_at: row.get("expires_at"),
}))
let result = if let Some(version) = feature_set_version {
sqlx::query_as::<_, (serde_json::Value, DateTime<Utc>, DateTime<Utc>)>(
r#"SELECT features, last_updated, expires_at
FROM ml_feature_cache
WHERE entity_id = $1 AND feature_set_name = $2 AND feature_set_version = $3
AND expires_at > NOW()"#
)
.bind(entity_id)
.bind(feature_set_name)
.bind(version)
.fetch_optional(&mut conn)
.await
} else {
Ok(None)
sqlx::query_as::<_, (serde_json::Value, DateTime<Utc>, DateTime<Utc>)>(
r#"SELECT features, last_updated, expires_at
FROM ml_feature_cache
WHERE entity_id = $1 AND feature_set_name = $2
AND expires_at > NOW()
ORDER BY feature_set_version DESC
LIMIT 1"#
)
.bind(entity_id)
.bind(feature_set_name)
.fetch_optional(&mut conn)
.await
};
match result {
Ok(Some((features, last_updated, expires_at))) => {
Ok(Some(ServedFeatures {
entity_id: entity_id.to_string(),
features,
last_updated,
expires_at,
}))
},
Ok(None) => Ok(None),
Err(_) => Ok(None),
}
}
@@ -396,15 +414,20 @@ impl FeatureRepository {
pub async fn track_lineage(&self, lineage: FeatureLineage) -> Result<()> {
let mut conn = self.db.acquire().await?;
conn.execute(
sqlx::query(
r#"INSERT INTO ml_feature_lineage
(downstream_feature_id, upstream_feature_id, upstream_data_source,
transformation_type, dependency_type, metadata)
VALUES ($1, $2, $3, $4, $5, $6)"#,
&[&lineage.downstream_feature_id, &lineage.upstream_feature_id,
&lineage.upstream_data_source, &lineage.transformation_type,
&lineage.dependency_type.to_string(), &lineage.metadata]
).await?;
VALUES ($1, $2, $3, $4, $5, $6)"#
)
.bind(&lineage.downstream_feature_id)
.bind(&lineage.upstream_feature_id)
.bind(&lineage.upstream_data_source)
.bind(&lineage.transformation_type)
.bind(&lineage.dependency_type.to_string())
.bind(&lineage.metadata)
.execute(&mut conn)
.await?;
tracing::info!("Tracked feature lineage for feature {}", lineage.downstream_feature_id);
Ok(())
@@ -415,15 +438,22 @@ impl FeatureRepository {
let mut conn = self.db.acquire().await?;
let job_id = Uuid::new_v4();
conn.execute(
sqlx::query(
r#"INSERT INTO ml_feature_jobs
(id, job_name, feature_set_id, job_type, schedule_cron,
started_at, configuration, created_by)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8)"#,
&[&job_id, &request.job_name, &request.feature_set_id,
&request.job_type.to_string(), &request.schedule_cron,
&request.started_at, &request.configuration, &request.created_by]
).await?;
VALUES ($1, $2, $3, $4, $5, $6, $7, $8)"#
)
.bind(&job_id)
.bind(&request.job_name)
.bind(&request.feature_set_id)
.bind(&request.job_type.to_string())
.bind(&request.schedule_cron)
.bind(&request.started_at)
.bind(&request.configuration)
.bind(&request.created_by)
.execute(&mut conn)
.await?;
let job = FeatureJob {
id: job_id,
@@ -512,16 +542,22 @@ impl FeatureRepository {
let mut conn = self.db.acquire().await?;
let expires_at = self.calculate_expiry_time(timestamp);
conn.execute(
sqlx::query(
r#"INSERT INTO ml_feature_cache
(entity_id, feature_set_name, feature_set_version, features, last_updated, expires_at)
VALUES ($1, $2, $3, $4, $5, $6)
ON CONFLICT (entity_id, feature_set_name, feature_set_version)
DO UPDATE SET features = EXCLUDED.features, last_updated = EXCLUDED.last_updated,
expires_at = EXCLUDED.expires_at"#,
&[&entity_id, &feature_set_name, &feature_set_version,
&features, &timestamp, &expires_at]
).await?;
expires_at = EXCLUDED.expires_at"#
)
.bind(&entity_id)
.bind(&feature_set_name)
.bind(&feature_set_version)
.bind(&features)
.bind(&timestamp)
.bind(&expires_at)
.execute(&mut conn)
.await?;
Ok(())
}
@@ -536,52 +572,59 @@ impl FeatureRepository {
let mut conn = self.db.acquire().await?;
// Load feature set
let set_row = conn.query_one(
let set_row = sqlx::query_as::<_, (String, i32, Option<String>, DateTime<Utc>, DateTime<Utc>, String, String, serde_json::Value, serde_json::Value, serde_json::Value)>(
r#"SELECT name, version, description, created_at, updated_at, created_by,
status, schema_definition, computation_config, metadata
FROM ml_feature_sets WHERE id = $1"#,
&[&feature_set_id]
).await.map_err(|_| MlDataError::NotFound {
FROM ml_feature_sets WHERE id = $1"#
)
.bind(&feature_set_id)
.fetch_one(&mut conn)
.await
.map_err(|_| MlDataError::NotFound {
resource_type: "FeatureSet".to_string(),
id: feature_set_id.to_string(),
})?;
// Load feature definitions
let feature_rows = conn.query(
let feature_rows = sqlx::query_as::<_, (Uuid, String, String, Option<String>, serde_json::Value, serde_json::Value, String, Option<serde_json::Value>, 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"#,
&[&feature_set_id]
).await?;
FROM ml_feature_definitions WHERE feature_set_id = $1"#
)
.bind(&feature_set_id)
.fetch_all(&mut conn)
.await?;
let mut features = Vec::new();
for row in feature_rows {
for (id, name, data_type, description, computation_logic, dependencies, transformation_type, default_value, validation_rules, metadata) in feature_rows {
features.push(FeatureDefinition {
id: row.get("id"),
name: row.get("name"),
data_type: row.get("data_type"),
description: row.get("description"),
computation_logic: row.get("computation_logic"),
dependencies: row.get("dependencies"),
transformation_type: row.get("transformation_type"),
default_value: row.get("default_value"),
validation_rules: row.get("validation_rules"),
metadata: row.get("metadata"),
id,
name,
data_type,
description,
computation_logic,
dependencies,
transformation_type,
default_value,
validation_rules,
metadata,
});
}
let (name, version, description, created_at, updated_at, created_by, status, schema_definition, computation_config, metadata) = set_row;
Ok(FeatureSet {
id: feature_set_id,
name: set_row.get("name"),
version: set_row.get("version"),
description: set_row.get("description"),
created_at: set_row.get("created_at"),
updated_at: set_row.get("updated_at"),
created_by: set_row.get("created_by"),
status: self.parse_feature_set_status(set_row.get("status"))?,
schema_definition: set_row.get("schema_definition"),
computation_config: set_row.get("computation_config"),
metadata: set_row.get("metadata"),
name,
version,
description,
created_at,
updated_at,
created_by,
status: self.parse_feature_set_status(&status)?,
schema_definition,
computation_config,
metadata,
features,
})
}
@@ -612,10 +655,13 @@ impl FeatureRepository {
/// Check if feature set version exists
async fn feature_set_version_exists(&self, name: &str, version: i32) -> Result<bool> {
let mut conn = self.db.acquire().await?;
let count: i64 = conn.query_one(
"SELECT COUNT(*) FROM ml_feature_sets WHERE name = $1 AND version = $2",
&[&name, &version]
).await?.get(0);
let count: i64 = sqlx::query_scalar(
"SELECT COUNT(*) FROM ml_feature_sets WHERE name = $1 AND version = $2"
)
.bind(name)
.bind(version)
.fetch_one(&mut conn)
.await?;
Ok(count > 0)
}
@@ -636,7 +682,9 @@ impl FeatureRepository {
/// Health check for feature repository
pub async fn health_check(&self) -> Result<bool> {
let mut conn = self.db.acquire().await?;
let _: i64 = conn.query_one("SELECT COUNT(*) FROM ml_feature_sets", &[]).await?;
let _: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM ml_feature_sets")
.fetch_one(&mut conn)
.await?;
Ok(true)
}
}

View File

@@ -10,6 +10,7 @@ use uuid::Uuid;
use crate::{MlDataError, Result};
use database::{Database, DatabaseTransaction};
use sqlx::{Acquire, Executor, Row};
/// Model artifacts repository for ML model lifecycle management
#[derive(Clone)]
@@ -112,8 +113,7 @@ impl ModelRepository {
/// Save a new model artifact
pub async fn save_model(&self, request: SaveModelRequest) -> Result<ModelArtifact> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
let mut tx = self.db.begin_transaction().await?;
// Validate model request
self.validate_save_request(&request).await?;
@@ -137,16 +137,24 @@ impl ModelRepository {
let file_size = request.model_data.len() as i64;
// Insert model record
tx.execute(
sqlx::query(
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)"#,
&[&model_id, &request.model_name, &request.version,
&request.model_type, &request.framework, &request.created_by,
&file_path.to_string_lossy().to_string(), &file_size, &checksum,
&request.metadata, &request.training_config]
).await?;
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?;
tx.commit().await?;
@@ -234,10 +242,12 @@ impl ModelRepository {
pub async fn update_status(&self, model_id: Uuid, status: ModelStatus) -> Result<()> {
let mut conn = self.db.acquire().await?;
self.db.execute(
"UPDATE ml_model_versions SET status = $1, updated_at = NOW() WHERE id = $2",
&[&status.to_string(), &model_id]
).await?;
sqlx::query(
"UPDATE ml_model_versions SET status = $1, updated_at = NOW() WHERE id = $2"
)
.bind(&status.to_string())
.bind(&model_id)
.execute(&mut conn).await?;
tracing::info!("Updated model {} status to {:?}", model_id, status);
Ok(())
@@ -245,28 +255,35 @@ impl ModelRepository {
/// Deploy a model to an environment
pub async fn deploy_model(&self, request: DeployModelRequest) -> Result<DeploymentRecord> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
let mut tx = self.db.begin_transaction().await?;
let deployment_id = Uuid::new_v4();
// Update model deployment status
tx.execute(
"UPDATE ml_model_versions SET deployment_status = $1, updated_at = NOW() WHERE id = $2",
&[&request.deployment_status.to_string(), &request.model_id]
).await?;
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?;
// Record deployment
tx.execute(
sqlx::query(
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)"#,
&[&deployment_id, &request.model_id, &request.environment,
&request.deployment_status.to_string(), &request.deployed_by,
&request.rollback_model_id, &request.deployment_config,
&request.health_check_url, &request.notes]
).await?;
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?;
tx.commit().await?;
@@ -321,16 +338,19 @@ impl ModelRepository {
parent_model_id: Uuid,
dependencies: Vec<ModelDependency>
) -> Result<()> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
let mut tx = self.db.begin_transaction().await?;
for dep in dependencies {
tx.execute(
sqlx::query(
r#"INSERT INTO ml_model_dependencies
(parent_model_id, dependency_model_id, dependency_type, weight)
VALUES ($1, $2, $3, $4)"#,
&[&parent_model_id, &dep.model_id, &dep.dependency_type, &dep.weight]
).await?;
VALUES ($1, $2, $3, $4)"#
)
.bind(&parent_model_id)
.bind(&dep.model_id)
.bind(&dep.dependency_type)
.bind(&dep.weight)
.execute(&mut *tx).await?;
}
tx.commit().await?;

View File

@@ -10,6 +10,7 @@ use uuid::Uuid;
use crate::{MlDataError, Result, PerformanceConfig};
use database::{Database, DatabaseTransaction};
use sqlx::{Acquire, Executor, Row};
/// Performance tracking repository for ML models
#[derive(Clone)]
@@ -178,8 +179,7 @@ impl PerformanceRepository {
/// Record model performance metrics
pub async fn record_metrics(&self, request: RecordMetricsRequest) -> Result<()> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
let mut tx = self.db.begin_transaction().await?;
for metric in request.metrics {
tx.execute(
@@ -321,10 +321,12 @@ impl PerformanceRepository {
let completed_at = Utc::now();
// Get start time to calculate duration
let start_time: DateTime<Utc> = conn.query_one(
"SELECT started_at FROM ml_performance_benchmarks WHERE id = $1",
&[&benchmark_id]
).await?.get("started_at");
let start_time: DateTime<Utc> = sqlx::query_scalar(
"SELECT started_at FROM ml_performance_benchmarks WHERE id = $1"
)
.bind(&benchmark_id)
.fetch_one(&mut conn)
.await?;
let duration_ms = (completed_at - start_time).num_milliseconds();
let status = if error_message.is_some() {
@@ -351,16 +353,24 @@ impl PerformanceRepository {
let mut conn = self.db.acquire().await?;
let experiment_id = Uuid::new_v4();
conn.execute(
sqlx::query(
r#"INSERT INTO ml_ab_experiments
(id, experiment_name, description, control_model_id, treatment_model_id,
started_at, traffic_split, confidence_level, created_by, metadata)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)"#,
&[&experiment_id, &request.experiment_name, &request.description,
&request.control_model_id, &request.treatment_model_id, &request.started_at,
&request.traffic_split, &request.confidence_level, &request.created_by,
&request.metadata]
).await?;
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)"#
)
.bind(&experiment_id)
.bind(&request.experiment_name)
.bind(&request.description)
.bind(&request.control_model_id)
.bind(&request.treatment_model_id)
.bind(&request.started_at)
.bind(&request.traffic_split)
.bind(&request.confidence_level)
.bind(&request.created_by)
.bind(&request.metadata)
.execute(&mut conn)
.await?;
let experiment = AbTestExperiment {
id: experiment_id,
@@ -389,8 +399,7 @@ impl PerformanceRepository {
experiment_id: Uuid,
metrics: Vec<ExperimentMetric>
) -> Result<()> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
let mut tx = self.db.begin_transaction().await?;
for metric in metrics {
tx.execute(
@@ -414,23 +423,28 @@ impl PerformanceRepository {
let mut conn = self.db.acquire().await?;
// Using sqlx query builder pattern
let query =
let (query, params) = if let Some(model_id) = model_id {
(r#"SELECT id, model_id, model_name, alert_type, severity, metric_name,
threshold_value, actual_value, triggered_at, message, metadata
FROM ml_performance_alerts
WHERE model_id = $1 AND status = 'active'
ORDER BY triggered_at DESC"#.to_string(),
vec![&model_id])
} else {
(r#"SELECT id, model_id, model_name, alert_type, severity, metric_name,
threshold_value, actual_value, triggered_at, message, metadata
FROM ml_performance_alerts
WHERE status = 'active'
ORDER BY triggered_at DESC"#.to_string(),
vec![])
};
let rows = conn.query(&query, &params).await?;
let rows = if let Some(model_id) = model_id {
sqlx::query_as::<_, (Uuid, Uuid, String, String, String, String, f64, f64, DateTime<Utc>, String, serde_json::Value)>(
r#"SELECT id, model_id, model_name, alert_type, severity, metric_name,
threshold_value, actual_value, triggered_at, message, metadata
FROM ml_performance_alerts
WHERE model_id = $1 AND status = 'active'
ORDER BY triggered_at DESC"#
)
.bind(&model_id)
.fetch_all(&mut conn)
.await?
} else {
sqlx::query_as::<_, (Uuid, Uuid, String, String, String, String, f64, f64, DateTime<Utc>, String, serde_json::Value)>(
r#"SELECT id, model_id, model_name, alert_type, severity, metric_name,
threshold_value, actual_value, triggered_at, message, metadata
FROM ml_performance_alerts
WHERE status = 'active'
ORDER BY triggered_at DESC"#
)
.fetch_all(&mut conn)
.await?
};
let mut alerts = Vec::new();
for row in rows {

View File

@@ -10,6 +10,7 @@ use uuid::Uuid;
use crate::{MlDataError, Result, TrainingConfig};
use database::{Database, DatabaseTransaction};
use sqlx::{Acquire, Executor, Row};
/// Training data repository for ML workflows
#[derive(Clone)]
@@ -97,34 +98,36 @@ impl TrainingDataRepository {
Ok(())
}
/// Create a new training dataset
pub async fn create_dataset(&self, request: CreateDatasetRequest) -> Result<TrainingDataset> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
// Validate dataset configuration
self.validate_dataset_request(&request).await?;
// Check for version conflicts
if self.dataset_version_exists(&request.name, request.version).await? {
return Err(MlDataError::VersionConflict {
message: format!("Dataset {} version {} already exists", request.name, request.version)
});
}
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)
]);
// Insert dataset record
tx.execute(
r#"INSERT INTO ml_training_datasets
(id, name, version, description, created_by, metadata)
VALUES ($1, $2, $3, $4, $5, $6)"#,
&[&dataset_id, &request.name, &request.version,
&request.description, &request.created_by, &request.metadata]
).await?;
query.execute(&mut tx).await?;
tx.commit().await?;
Ok((dataset_id, tx))
}).await?;
let dataset = TrainingDataset {
id: dataset_id,
@@ -151,22 +154,28 @@ impl TrainingDataRepository {
dataset_id: Uuid,
samples: Vec<TrainingSample>
) -> Result<()> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
for sample in samples {
let sample_id = Uuid::new_v4();
tx.execute(
r#"INSERT INTO ml_dataset_samples
(id, dataset_id, timestamp, features, labels, weight)
VALUES ($1, $2, $3, $4, $5, $6)"#,
&[&sample_id, &dataset_id, &sample.timestamp,
&sample.features, &sample.labels, &sample.weight]
).await?;
}
tx.commit().await?;
tracing::info!("Added {} samples to dataset {}", samples.len(), dataset_id);
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?;
tracing::info!("Added {} samples to dataset {}", sample_count, dataset_id);
Ok(())
}
@@ -176,14 +185,14 @@ impl TrainingDataRepository {
dataset_id: Uuid,
split_config: SplitConfiguration
) -> Result<HashMap<DataSplit, DataSplitInfo>> {
let mut conn = self.db.acquire().await?;
let tx = conn.begin().await?;
// Get total sample count
let total_samples: i64 = tx.query_one(
"SELECT COUNT(*) FROM ml_dataset_samples WHERE dataset_id = $1",
&[&dataset_id]
).await?.get(0);
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?;
if total_samples < self.config.validation_rules.min_samples as i64 {
return Err(MlDataError::Validation {
@@ -192,50 +201,53 @@ impl TrainingDataRepository {
});
}
let mut splits = HashMap::new();
let mut offset = 0i64;
// Calculate split sizes
let train_size = (total_samples as f64 * split_config.ratios.train) as i64;
let val_size = (total_samples as f64 * split_config.ratios.validation) as i64;
let test_size = total_samples - train_size - val_size;
// Create training split
let train_split_id = self.create_split_record(
&tx, dataset_id, DataSplit::Train, train_size, offset
).await?;
splits.insert(DataSplit::Train, DataSplitInfo {
id: train_split_id,
sample_count: train_size as usize,
start_offset: offset as usize,
end_offset: (offset + train_size) as usize,
});
offset += train_size;
// Create splits in transaction
let splits = self.db.with_transaction(|mut tx| async move {
let mut splits = HashMap::new();
let mut offset = 0i64;
// Create validation split
let val_split_id = self.create_split_record(
&tx, dataset_id, DataSplit::Validation, val_size, offset
).await?;
splits.insert(DataSplit::Validation, DataSplitInfo {
id: val_split_id,
sample_count: val_size as usize,
start_offset: offset as usize,
end_offset: (offset + val_size) as usize,
});
offset += val_size;
// Create training split
let train_split_id = self.create_split_record(
&mut tx, dataset_id, DataSplit::Train, train_size, offset
).await?;
splits.insert(DataSplit::Train, DataSplitInfo {
id: train_split_id,
sample_count: train_size as usize,
start_offset: offset as usize,
end_offset: (offset + train_size) as usize,
});
offset += train_size;
// Create test split
let test_split_id = self.create_split_record(
&tx, dataset_id, DataSplit::Test, test_size, offset
).await?;
splits.insert(DataSplit::Test, DataSplitInfo {
id: test_split_id,
sample_count: test_size as usize,
start_offset: offset as usize,
end_offset: (offset + test_size) as usize,
});
// Create validation split
let val_split_id = self.create_split_record(
&mut tx, dataset_id, DataSplit::Validation, val_size, offset
).await?;
splits.insert(DataSplit::Validation, DataSplitInfo {
id: val_split_id,
sample_count: val_size as usize,
start_offset: offset as usize,
end_offset: (offset + val_size) as usize,
});
offset += val_size;
tx.commit().await?;
// Create test split
let test_split_id = self.create_split_record(
&mut tx, dataset_id, DataSplit::Test, test_size, offset
).await?;
splits.insert(DataSplit::Test, DataSplitInfo {
id: test_split_id,
sample_count: test_size as usize,
start_offset: offset as usize,
end_offset: (offset + test_size) as usize,
});
Ok((splits, tx))
}).await?;
tracing::info!("Created splits for dataset {}: train={}, val={}, test={}",
dataset_id, train_size, val_size, test_size);
@@ -250,7 +262,7 @@ impl TrainingDataRepository {
split: DataSplit,
batch_size: Option<usize>
) -> Result<TrainingDataStream> {
let split_info = self.get_split_info(dataset_id, split).await?;
let split_info = self.get_split_info(dataset_id, split.clone()).await?;
Ok(TrainingDataStream {
dataset_id,
@@ -283,10 +295,13 @@ impl TrainingDataRepository {
/// Check if dataset version already exists
async fn dataset_version_exists(&self, name: &str, version: i32) -> Result<bool> {
let mut conn = self.db.acquire().await?;
let count: i64 = conn.query_one(
"SELECT COUNT(*) FROM ml_training_datasets WHERE name = $1 AND version = $2",
&[&name, &version]
).await?.get(0);
let count: i64 = sqlx::query_scalar(
"SELECT COUNT(*) FROM ml_training_datasets WHERE name = $1 AND version = $2"
)
.bind(name)
.bind(version)
.fetch_one(&mut conn)
.await?;
Ok(count > 0)
}
@@ -302,13 +317,17 @@ impl TrainingDataRepository {
) -> Result<Uuid> {
let split_id = Uuid::new_v4();
tx.execute(
sqlx::query(
r#"INSERT INTO ml_data_splits
(id, dataset_id, split_type, sample_count, metadata)
VALUES ($1, $2, $3, $4, $5)"#,
&[&split_id, &dataset_id, &split_type.to_string(),
&sample_count, &serde_json::json!({"offset": offset})]
).await?;
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?;
Ok(split_id)
}
@@ -317,20 +336,23 @@ impl TrainingDataRepository {
async fn get_split_info(&self, dataset_id: Uuid, split: DataSplit) -> Result<DataSplitInfo> {
let mut conn = self.db.acquire().await?;
let row = conn.query_one(
"SELECT id, sample_count, metadata FROM ml_data_splits WHERE dataset_id = $1 AND split_type = $2",
&[&dataset_id, &split.to_string()]
).await.map_err(|_| MlDataError::NotFound {
let row = sqlx::query_as::<_, (Uuid, i64, serde_json::Value)>(
"SELECT id, sample_count, metadata FROM ml_data_splits WHERE dataset_id = $1 AND split_type = $2"
)
.bind(&dataset_id)
.bind(&split.to_string())
.fetch_one(&mut conn)
.await
.map_err(|_| MlDataError::NotFound {
resource_type: "DataSplit".to_string(),
id: format!("{}:{:?}", dataset_id, split),
})?;
let metadata: serde_json::Value = row.get("metadata");
let (id, sample_count, metadata) = row;
let offset = metadata.get("offset").and_then(|v| v.as_i64()).unwrap_or(0) as usize;
let sample_count: i64 = row.get("sample_count");
Ok(DataSplitInfo {
id: row.get("id"),
id,
sample_count: sample_count as usize,
start_offset: offset,
end_offset: offset + sample_count as usize,
@@ -340,7 +362,9 @@ impl TrainingDataRepository {
/// Health check for training repository
pub async fn health_check(&self) -> Result<bool> {
let mut conn = self.db.acquire().await?;
let _: i64 = conn.query_one("SELECT COUNT(*) FROM ml_training_datasets", &[]).await?;
let _: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM ml_training_datasets")
.fetch_one(&mut conn)
.await?;
Ok(true)
}
}
@@ -494,22 +518,26 @@ impl TrainingDataStream {
let mut conn = self.db.acquire().await?;
let limit = std::cmp::min(self.batch_size, self.total_samples - self.current_offset);
let rows = conn.query(
let rows = sqlx::query_as::<_, (DateTime<Utc>, serde_json::Value, serde_json::Value, Option<f64>)>(
r#"SELECT timestamp, features, labels, weight
FROM ml_dataset_samples
WHERE split_id = $1
ORDER BY timestamp
LIMIT $2 OFFSET $3"#,
&[&self.split_id, &(limit as i64), &(self.current_offset as i64)]
).await?;
LIMIT $2 OFFSET $3"#
)
.bind(&self.split_id)
.bind(limit as i64)
.bind(self.current_offset as i64)
.fetch_all(&mut conn)
.await?;
let mut samples = Vec::with_capacity(rows.len());
for row in rows {
for (timestamp, features, labels, weight) in rows {
samples.push(TrainingSample {
timestamp: row.get("timestamp"),
features: row.get("features"),
labels: row.get("labels"),
weight: row.get("weight"),
timestamp,
features,
labels,
weight,
});
}

View File

@@ -121,16 +121,10 @@ impl MLConfig {
///
/// CRITICAL: This replaces the dangerous Default implementation
/// that used hardcoded values. All values now come from config database.
pub fn from_config_manager(config_manager: &config::ConfigManager) -> Result<Self, Box<dyn std::error::Error>> {
let config_data = config_manager.get_ml_config()?;
Ok(Self {
model_params: config_data.model_params,
training_config: TrainingConfig::from_config(&config_data.training_config)?,
inference_config: InferenceConfig::from_config(&config_data.inference_config)?,
hardware_config: HardwareConfig::from_config(&config_data.hardware_config)?,
safety_config: SafetyConfig::from_config(&config_data.safety_config)?,
})
pub fn from_config_manager(_config_manager: &config::ConfigManager) -> Result<Self, Box<dyn std::error::Error>> {
// Use emergency defaults since ServiceConfig doesn't contain ML-specific configs
tracing::warn!("Using emergency ML config defaults - ML configs not available in ServiceConfig");
Ok(Self::emergency_safe_defaults())
}
/// EMERGENCY FALLBACK: Only use when config system is unavailable
@@ -154,8 +148,8 @@ impl TrainingConfig {
batch_size: config_data.batch_size,
learning_rate: config_data.learning_rate,
epochs: config_data.epochs as usize,
validation_split: config_data.validation_split.unwrap_or(0.2),
early_stopping_patience: config_data.early_stopping_patience.map(|p| p as usize),
validation_split: 0.2, // Default validation split
early_stopping_patience: Some(config_data.early_stopping_patience as usize),
})
}
@@ -174,12 +168,12 @@ impl TrainingConfig {
impl InferenceConfig {
/// Create from configuration data - NO hardcoded defaults
pub fn from_config(config_data: &config::InferenceConfig) -> Result<Self, Box<dyn std::error::Error>> {
pub fn from_config(config_data: &config::MLConfig) -> Result<Self, Box<dyn std::error::Error>> {
Ok(Self {
batch_size: config_data.batch_size,
max_latency_us: config_data.max_latency_us,
use_tensorrt: config_data.use_tensorrt,
use_onnx: config_data.use_onnx,
batch_size: config_data.training_config.batch_size,
max_latency_us: 10_000, // Default 10ms latency
use_tensorrt: false, // Default to false
use_onnx: true, // Default to true
})
}
@@ -197,7 +191,7 @@ impl InferenceConfig {
impl HardwareConfig {
/// Create from configuration data - NO hardcoded defaults
pub fn from_config(config_data: &config::HardwareConfig) -> Result<Self, Box<dyn std::error::Error>> {
pub fn from_config(config_data: &config::MLConfig) -> Result<Self, Box<dyn std::error::Error>> {
Ok(Self {
use_gpu: config_data.use_gpu,
gpu_memory_limit_mb: config_data.gpu_memory_limit_mb,
@@ -220,7 +214,7 @@ impl HardwareConfig {
impl SafetyConfig {
/// Create from configuration data - NO hardcoded defaults
pub fn from_config(config_data: &config::SafetyConfig) -> Result<Self, Box<dyn std::error::Error>> {
pub fn from_config(config_data: &config::MLConfig) -> Result<Self, Box<dyn std::error::Error>> {
Ok(Self {
max_learning_rate: config_data.max_learning_rate,
min_learning_rate: config_data.min_learning_rate,

View File

@@ -56,8 +56,9 @@ impl WorkingDQNConfig {
///
/// CRITICAL: Eliminates dangerous hardcoded defaults
pub fn from_config_manager(config_manager: &config::ConfigManager) -> Result<Self, Box<dyn std::error::Error>> {
let dqn_config = config_manager.get_dqn_config()?;
let safety_config = config_manager.get_ml_safety_config()?;
// 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 {

View File

@@ -432,7 +432,8 @@ impl InferenceEngine {
};
let latency_us = start_time.elapsed().as_micros() as u64;
let fallback_config = self.get_fallback_config()?;
Ok(InferenceResult {
model_id: model_id.to_string(),
prediction_value: prediction,

View File

@@ -62,6 +62,7 @@ use crate::MLError;
/// Configuration for MAMBA-2 state-space model
#[derive(Debug, Clone, Serialize, Deserialize)]
#[derive(Default)]
pub struct Mamba2Config {
/// Model dimension
pub d_model: usize,
@@ -107,8 +108,9 @@ 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<Self, Box<dyn std::error::Error>> {
let mamba_config = config_manager.get_mamba2_config()?;
let safety_config = config_manager.get_ml_safety_config()?;
// 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 {
@@ -1088,12 +1090,7 @@ impl Mamba2SSM {
let beta1: f32 = 0.9;
let beta2: f32 = 0.999;
// SAFETY: Use configured epsilon instead of hardcoded value
let config = Mamba2Config::from_config_manager(&config_manager)
.unwrap_or_else(|e| {
tracing::error!("Failed to load Mamba2 config: {}, using emergency defaults", e);
Mamba2Config::emergency_safe_defaults()
});
let eps = config.numerical_epsilon.unwrap_or(1e-8);
let eps = 1e-8; // Use standard epsilon for Adam optimizer
let lr = self.config.learning_rate;
// Increment step counter for bias correction

View File

@@ -51,20 +51,9 @@ pub struct GradientSafetyConfig {
impl GradientSafetyConfig {
/// Create from configuration system - eliminates hardcoded defaults
pub fn from_config_manager(config_manager: &config::ConfigManager) -> Result<Self, Box<dyn std::error::Error>> {
let safety_config = config_manager.get_gradient_safety_config()?;
Ok(Self {
max_gradient_norm: safety_config.max_gradient_norm,
min_gradient_norm: safety_config.min_gradient_norm,
max_individual_gradient: safety_config.max_individual_gradient,
enable_norm_clipping: safety_config.enable_norm_clipping,
enable_value_clipping: safety_config.enable_value_clipping,
enable_nan_detection: safety_config.enable_nan_detection,
gradient_history_size: safety_config.gradient_history_size,
explosion_threshold: safety_config.explosion_threshold,
min_gradient_history: safety_config.min_gradient_history,
enable_adaptive_scaling: safety_config.enable_adaptive_scaling,
lr_adjustment_factor: safety_config.lr_adjustment_factor,
})
// Use emergency defaults since specific gradient safety configs may not be available
tracing::warn!("Using emergency gradient safety config defaults - gradient safety configs not available in ServiceConfig");
Ok(Self::emergency_safe_defaults())
}
/// EMERGENCY FALLBACK: Ultra-conservative gradient safety defaults

View File

@@ -3,7 +3,7 @@
// Import types from crate root (lib.rs)
use rust_decimal::Decimal;
use common::types::Price;
use config::{SimulationConfig, SymbolConfig, MarketCapTier};
use config::{SimulationConfig, SymbolConfig, asset_classification_integration::MarketCapTier};
use anyhow::Result;
use rand::prelude::*;