Skip to main content

f_ck/engine/
joiner.rs

1use crate::dsl::{QueryPlan, Mapping, MergePolicy};
2use crate::engine::DataReader;
3use anyhow::Result;
4use polars::lazy::frame::LazyFrame;
5use polars::prelude::*;
6use std::collections::HashMap;
7
8pub struct JoinEngine;
9
10impl JoinEngine {
11    pub fn execute_query(query: &QueryPlan) -> Result<LazyFrame> {
12        query.validate()?;
13
14        // Load all source data
15        let mut dataframes: HashMap<String, LazyFrame> = HashMap::new();
16        for source in &query.sources {
17            let df = DataReader::read_source(source)?;
18            dataframes.insert(source.id.clone(), df);
19        }
20
21        // For now, implement a simple join strategy
22        // This is a basic implementation - will be enhanced with proper transitive closure later
23        Self::simple_join(query, dataframes)
24    }
25
26    fn simple_join(query: &QueryPlan, mut dataframes: HashMap<String, LazyFrame>) -> Result<LazyFrame> {
27        if dataframes.is_empty() {
28            return Err(anyhow::anyhow!("No dataframes to join"));
29        }
30
31        // Start with the first dataframe
32        let first_source_id = query.sources[0].id.clone();
33        let mut result = dataframes.remove(&first_source_id)
34            .ok_or_else(|| anyhow::anyhow!("First source not found"))?;
35
36        // For each additional source, perform joins based on mappings
37        for source in query.sources.iter().skip(1) {
38            if let Some(right_df) = dataframes.remove(&source.id) {
39                result = Self::join_dataframes(result, right_df, query)?;
40            }
41        }
42
43        // Apply field mappings and create destination schema with aggregation
44        Self::apply_mappings_with_aggregation(result, query)
45    }
46
47    fn join_dataframes(left: LazyFrame, right: LazyFrame, query: &QueryPlan) -> Result<LazyFrame> {
48        // Find the primary key mapping to determine join columns
49        let primary_key = &query.primary_keys.keys[0]; // Using first primary key
50        
51        // Find the mapping for this primary key
52        let pk_mapping = query.mappings.iter()
53            .find(|m| m.destination_field == *primary_key)
54            .ok_or_else(|| anyhow::anyhow!("No mapping found for primary key: {}", primary_key))?;
55        
56        // Extract the source columns for the join
57        // Assume first source field is from left table, second from right table
58        if pk_mapping.source_fields.len() < 2 {
59            return Err(anyhow::anyhow!("Primary key mapping needs at least 2 source fields for join"));
60        }
61        
62        let left_col = &pk_mapping.source_fields[0].column_name;
63        let right_col = &pk_mapping.source_fields[1].column_name;
64
65        let result = left.join(
66            right,
67            [col(left_col)],
68            [col(right_col)],
69            JoinArgs::new(JoinType::Left),
70        );
71
72        Ok(result)
73    }
74
75    fn apply_mappings_with_aggregation(df: LazyFrame, query: &QueryPlan) -> Result<LazyFrame> {
76        // First, determine if we need aggregation by checking if any mapping uses Sum, Count, Average
77        let needs_aggregation = query.mappings.iter().any(|mapping| {
78            matches!(mapping.policy, MergePolicy::Sum | MergePolicy::Count | MergePolicy::Average)
79        });
80
81        if needs_aggregation {
82            // Find the primary key columns for grouping
83            let primary_key = &query.primary_keys.keys[0];
84            let pk_mapping = query.mappings.iter()
85                .find(|m| m.destination_field == *primary_key)
86                .ok_or_else(|| anyhow::anyhow!("No mapping found for primary key: {}", primary_key))?;
87            
88            // Use the left table's column for grouping (from the join)
89            let group_col = &pk_mapping.source_fields[0].column_name;
90            
91            // Create aggregation expressions
92            let mut agg_exprs = Vec::new();
93            
94            for dest_field in &query.destination_schema {
95                if let Some(mapping) = query.mappings.iter()
96                    .find(|m| m.destination_field == dest_field.name) {
97                    
98                    let expr = Self::create_aggregation_expression(mapping)?;
99                    agg_exprs.push(expr.alias(&dest_field.name));
100                } else {
101                    // For non-mapped fields, use first() to get one value per group
102                    agg_exprs.push(col(&dest_field.name).first().alias(&dest_field.name));
103                }
104            }
105            
106            Ok(df.group_by([col(group_col)]).agg(agg_exprs))
107        } else {
108            // No aggregation needed, just apply regular mappings
109            Self::apply_simple_mappings(df, query)
110        }
111    }
112
113    fn apply_simple_mappings(df: LazyFrame, query: &QueryPlan) -> Result<LazyFrame> {
114        let mut exprs = Vec::new();
115
116        // Create expressions for each destination field based on mappings
117        for dest_field in &query.destination_schema {
118            if let Some(mapping) = query.mappings.iter()
119                .find(|m| m.destination_field == dest_field.name) {
120                
121                let expr = Self::create_mapping_expression(mapping)?;
122                exprs.push(expr.alias(&dest_field.name));
123            } else {
124                // If no mapping found, try to find a column with the same name
125                exprs.push(col(&dest_field.name));
126            }
127        }
128
129        Ok(df.select(exprs))
130    }
131
132    fn create_aggregation_expression(mapping: &Mapping) -> Result<Expr> {
133        if mapping.source_fields.is_empty() {
134            return Err(anyhow::anyhow!("Mapping has no source fields"));
135        }
136
137        match &mapping.policy {
138            MergePolicy::FirstMatch { priority: _ } => {
139                // For aggregation, use first() to get one value per group
140                let first_col = &mapping.source_fields[0].column_name;
141                Ok(col(first_col).first())
142            },
143            MergePolicy::Sum => {
144                // Sum the values across the group
145                let first_col = &mapping.source_fields[0].column_name;
146                Ok(col(first_col).sum())
147            },
148            MergePolicy::Count => {
149                // Count non-null values in the group
150                let first_col = &mapping.source_fields[0].column_name;
151                Ok(col(first_col).count())
152            },
153            MergePolicy::Average => {
154                // Average the values across the group
155                let first_col = &mapping.source_fields[0].column_name;
156                Ok(col(first_col).mean())
157            },
158            MergePolicy::Min => {
159                // Minimum value in the group
160                let first_col = &mapping.source_fields[0].column_name;
161                Ok(col(first_col).min())
162            },
163            MergePolicy::Max => {
164                // Maximum value in the group
165                let first_col = &mapping.source_fields[0].column_name;
166                Ok(col(first_col).max())
167            },
168        }
169    }
170
171    fn create_mapping_expression(mapping: &Mapping) -> Result<Expr> {
172        if mapping.source_fields.is_empty() {
173            return Err(anyhow::anyhow!("Mapping has no source fields"));
174        }
175
176        match &mapping.policy {
177            MergePolicy::FirstMatch { priority: _ } => {
178                // For FirstMatch, use the first available column
179                // This handles the case where join keys might not be available after join
180                let first_col = &mapping.source_fields[0].column_name;
181                Ok(col(first_col))
182            },
183            MergePolicy::Sum => {
184                // For sum, we'll fold over the columns
185                let mut expr = lit(0);
186                for sf in &mapping.source_fields {
187                    expr = expr + col(&sf.column_name);
188                }
189                Ok(expr)
190            },
191            MergePolicy::Count => {
192                // Count non-null values across source fields
193                let mut expr = lit(0);
194                for sf in &mapping.source_fields {
195                    expr = expr + col(&sf.column_name).is_not_null().cast(DataType::Int32);
196                }
197                Ok(expr)
198            },
199            MergePolicy::Average => {
200                // Calculate average manually
201                let mut sum_expr = lit(0.0);
202                let count = mapping.source_fields.len() as f64;
203                for sf in &mapping.source_fields {
204                    sum_expr = sum_expr + col(&sf.column_name).cast(DataType::Float64);
205                }
206                Ok(sum_expr / lit(count))
207            },
208            MergePolicy::Min => {
209                // Use coalesce-like approach for minimum across columns
210                let cols: Vec<Expr> = mapping.source_fields.iter()
211                    .map(|sf| col(&sf.column_name))
212                    .collect();
213                
214                if cols.len() == 1 {
215                    Ok(cols[0].clone())
216                } else {
217                    // For now, use the first non-null value as a placeholder
218                    // TODO: Implement proper minimum across columns
219                    Ok(coalesce(&cols))
220                }
221            },
222            MergePolicy::Max => {
223                // Use coalesce-like approach for maximum across columns
224                let cols: Vec<Expr> = mapping.source_fields.iter()
225                    .map(|sf| col(&sf.column_name))
226                    .collect();
227                
228                if cols.len() == 1 {
229                    Ok(cols[0].clone())
230                } else {
231                    // For now, use the first non-null value as a placeholder
232                    // TODO: Implement proper maximum across columns
233                    Ok(coalesce(&cols))
234                }
235            },
236        }
237    }
238}