🔧 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,
});
}