🔧 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:
@@ -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, ×tamp, &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(×tamp)
|
||||
.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, ¶ms).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, ×tamp, &expires_at]
|
||||
).await?;
|
||||
expires_at = EXCLUDED.expires_at"#
|
||||
)
|
||||
.bind(&entity_id)
|
||||
.bind(&feature_set_name)
|
||||
.bind(&feature_set_version)
|
||||
.bind(&features)
|
||||
.bind(×tamp)
|
||||
.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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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?;
|
||||
|
||||
@@ -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, ¶ms).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 {
|
||||
|
||||
@@ -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,
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user