Skip to main content

dataprof_db/
lib.rs

1//! Database connectivity module for dataprof.
2//!
3//! This crate owns the database profiling surface, including connection
4//! handling, secure configuration helpers, sampling, streaming, and the
5//! feature-gated sqlx-based connectors.
6
7pub(crate) use dataprof_core::DataProfilerError;
8use dataprof_core::{DataSource, ExecutionMetadata, QualityDimension, QueryEngine};
9use dataprof_metrics::analyze_column;
10use dataprof_runtime::{ProfileReport, ReportAssembler};
11use std::collections::HashMap;
12
13pub mod connection;
14pub mod connectors;
15pub mod retry;
16pub mod sampling;
17pub mod security;
18pub mod streaming;
19
20pub use connection::*;
21pub use connectors::*;
22pub use retry::*;
23pub use sampling::*;
24pub use security::*;
25
26/// Database configuration for connection strings and settings
27#[derive(Debug, Clone)]
28pub struct DatabaseConfig {
29    pub connection_string: String,
30    pub batch_size: usize,
31    pub max_connections: Option<u32>,
32    pub connection_timeout: Option<std::time::Duration>,
33    pub retry_config: Option<RetryConfig>,
34    pub sampling_config: Option<SamplingConfig>,
35    pub ssl_config: Option<SslConfig>,
36    pub load_credentials_from_env: bool,
37}
38
39impl Default for DatabaseConfig {
40    fn default() -> Self {
41        Self {
42            connection_string: String::new(),
43            batch_size: 10000,
44            max_connections: Some(10),
45            connection_timeout: Some(std::time::Duration::from_secs(30)),
46            retry_config: Some(RetryConfig::default()),
47            sampling_config: None,
48            ssl_config: Some(SslConfig::default()),
49            load_credentials_from_env: true,
50        }
51    }
52}
53
54/// Trait that all database connectors must implement
55#[async_trait::async_trait]
56pub trait DatabaseConnector: Send + Sync {
57    /// Connect to the database
58    async fn connect(&mut self) -> Result<(), DataProfilerError>;
59
60    /// Disconnect from the database
61    async fn disconnect(&mut self) -> Result<(), DataProfilerError>;
62
63    /// Execute a query and get column data for profiling
64    async fn profile_query(
65        &mut self,
66        query: &str,
67    ) -> Result<HashMap<String, Vec<String>>, DataProfilerError>;
68
69    /// Execute a query with streaming for large result sets
70    async fn profile_query_streaming(
71        &mut self,
72        query: &str,
73        batch_size: usize,
74    ) -> Result<HashMap<String, Vec<String>>, DataProfilerError>;
75
76    /// Get table schema information
77    async fn get_table_schema(
78        &mut self,
79        table_name: &str,
80    ) -> Result<Vec<String>, DataProfilerError>;
81
82    /// Count total rows in table (for progress tracking)
83    async fn count_table_rows(&mut self, table_name: &str) -> Result<u64, DataProfilerError>;
84
85    /// Test connection
86    async fn test_connection(&mut self) -> Result<bool, DataProfilerError>;
87}
88
89/// Factory function to create appropriate database connector
90pub fn create_connector(
91    mut config: DatabaseConfig,
92) -> Result<Box<dyn DatabaseConnector>, DataProfilerError> {
93    if config.load_credentials_from_env || config.connection_string.is_empty() {
94        config = apply_environment_configuration(config)?;
95    }
96
97    let connection_str = config.connection_string.as_str();
98
99    if connection_str.starts_with("postgresql://") || connection_str.starts_with("postgres://") {
100        Ok(Box::new(connectors::postgres::PostgresConnector::new(
101            config,
102        )?))
103    } else if connection_str.starts_with("mysql://") {
104        Ok(Box::new(connectors::mysql::MySqlConnector::new(config)?))
105    } else if connection_str.starts_with("sqlite://")
106        || connection_str.ends_with(".db")
107        || connection_str.ends_with(".sqlite")
108        || connection_str == ":memory:"
109    {
110        Ok(Box::new(connectors::sqlite::SqliteConnector::new(config)?))
111    } else {
112        Err(DataProfilerError::DatabaseConfigError {
113            message: format!(
114                "Unsupported database connection string: {}. Supported: postgresql://, mysql://, sqlite://",
115                connection_str
116            ),
117        })
118    }
119}
120
121/// Apply environment configuration to database config
122fn apply_environment_configuration(
123    mut config: DatabaseConfig,
124) -> Result<DatabaseConfig, DataProfilerError> {
125    let database_type = if config.connection_string.is_empty() {
126        if std::env::var("POSTGRES_URL").is_ok()
127            || std::env::var("DATABASE_URL")
128                .map(|url| url.starts_with("postgres"))
129                .unwrap_or(false)
130        {
131            "postgresql".to_string()
132        } else if std::env::var("MYSQL_URL").is_ok() {
133            "mysql".to_string()
134        } else {
135            "postgresql".to_string()
136        }
137    } else {
138        let conn_info = ConnectionInfo::parse(&config.connection_string)?;
139        conn_info.database_type().to_string()
140    };
141    let database_type = database_type.as_str();
142
143    if config.connection_string.is_empty() {
144        let (secure_connection_string, ssl_config) = load_secure_database_config(database_type)?;
145        config.connection_string = secure_connection_string;
146        config.ssl_config = Some(ssl_config);
147    } else {
148        if let Some(ssl_config) = &config.ssl_config {
149            config.connection_string = ssl_config
150                .apply_to_connection_string(config.connection_string.clone(), database_type);
151        }
152
153        if config.load_credentials_from_env {
154            let credentials = DatabaseCredentials::from_environment(database_type);
155            config.connection_string =
156                credentials.apply_to_connection_string(&config.connection_string);
157        }
158    }
159
160    Ok(config)
161}
162
163/// High-level function to analyze a database table or query.
164pub async fn analyze_database(
165    config: DatabaseConfig,
166    query: &str,
167    calculate_quality: bool,
168    quality_dimensions: Option<Vec<QualityDimension>>,
169) -> Result<ProfileReport, DataProfilerError> {
170    let mut connector = create_connector(config.clone())?;
171
172    connector.connect().await?;
173
174    let start = std::time::Instant::now();
175
176    let (actual_query, is_table) = if query.trim().to_uppercase().starts_with("SELECT") {
177        let validated_query = security::validate_base_query(query)?;
178        (validated_query, false)
179    } else {
180        security::validate_sql_identifier(query)?;
181        (format!("SELECT * FROM {}", query), true)
182    };
183
184    let total_rows = if is_table {
185        // decode-audit: unknown — a failed COUNT means "row count unknown", not
186        // "zero rows". 0 disables sampling below, so log the failure instead of
187        // silently pretending the table is empty.
188        match connector.count_table_rows(query).await {
189            Ok(count) => count,
190            Err(e) => {
191                log::warn!(
192                    "count_table_rows failed for '{}': {}; row count unknown, sampling disabled",
193                    query,
194                    e
195                );
196                0
197            }
198        }
199    } else {
200        0
201    };
202
203    let (final_query, sample_info) = if let Some(sampling_config) = &config.sampling_config {
204        if total_rows > sampling_config.sample_size as u64 {
205            let sampled_query = sampling_config.generate_sample_query(&actual_query, total_rows)?;
206            let info = SampleInfo::new(
207                total_rows,
208                sampling_config.sample_size.min(total_rows as usize) as u64,
209                sampling_config.strategy.clone(),
210            );
211            (sampled_query, Some(info))
212        } else {
213            (actual_query, None)
214        }
215    } else {
216        (actual_query, None)
217    };
218
219    let columns = connector
220        .profile_query_streaming(&final_query, config.batch_size)
221        .await?;
222
223    connector.disconnect().await?;
224
225    let query_engine = detect_query_engine(&config.connection_string);
226
227    if columns.is_empty() {
228        let mut exec = ExecutionMetadata::new(0, 0, start.elapsed().as_millis())
229            .with_engine(query_engine.to_string());
230        if let Some(ref info) = sample_info
231            && info.sampling_ratio < 1.0
232        {
233            exec = exec
234                .with_sampling(info.sampling_ratio)
235                .with_source_exhausted(false);
236        }
237        return Ok(ReportAssembler::new(
238            DataSource::Query {
239                engine: query_engine.clone(),
240                statement: query.to_string(),
241                database: extract_database_name(&config.connection_string),
242                execution_id: None,
243            },
244            exec,
245        )
246        .skip_quality()
247        .build());
248    }
249
250    let mut column_profiles = Vec::new();
251    // decode-audit: no-data — the empty-columns case returned early above, so
252    // this default is only a guard; every column vec has the same length.
253    let actual_rows_processed = columns.values().next().map(|v| v.len()).unwrap_or(0);
254
255    for (name, data) in &columns {
256        let profile = analyze_column(name, data);
257        column_profiles.push(profile);
258    }
259
260    let scan_time_ms = start.elapsed().as_millis();
261    let sampling_ratio = sample_info.map(|s| s.sampling_ratio).unwrap_or(1.0);
262    let num_columns = column_profiles.len();
263
264    let mut execution = ExecutionMetadata::new(actual_rows_processed, num_columns, scan_time_ms)
265        .with_engine(query_engine.to_string());
266    if sampling_ratio < 1.0 {
267        execution = execution
268            .with_sampling(sampling_ratio)
269            .with_source_exhausted(false);
270    }
271
272    let mut assembler = ReportAssembler::new(
273        DataSource::Query {
274            engine: query_engine,
275            statement: query.to_string(),
276            database: extract_database_name(&config.connection_string),
277            execution_id: None,
278        },
279        execution,
280    )
281    .columns(column_profiles);
282
283    if calculate_quality {
284        assembler = assembler.with_quality_data(columns);
285        if let Some(dims) = quality_dimensions {
286            assembler = assembler.with_requested_dimensions(dims);
287        }
288    } else {
289        assembler = assembler.skip_quality();
290    }
291
292    Ok(assembler.build())
293}
294
295/// Detect query engine from connection string
296fn detect_query_engine(connection_string: &str) -> QueryEngine {
297    let conn = connection_string.to_lowercase();
298    if conn.starts_with("postgres") || conn.starts_with("postgresql") {
299        QueryEngine::Postgres
300    } else if conn.starts_with("mysql") || conn.starts_with("mariadb") {
301        QueryEngine::MySql
302    } else if conn.starts_with("sqlite") {
303        QueryEngine::Sqlite
304    } else {
305        QueryEngine::Custom("unknown".to_string())
306    }
307}
308
309/// Extract database name from connection string
310fn extract_database_name(connection_string: &str) -> Option<String> {
311    if let Some(pos) = connection_string.rfind('/') {
312        let db_part = &connection_string[pos + 1..];
313        let db_name = db_part.split('?').next().unwrap_or(db_part);
314        if !db_name.is_empty() {
315            return Some(db_name.to_string());
316        }
317    }
318    None
319}