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 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 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 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 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 Self::apply_mappings_with_aggregation(result, query)
45 }
46
47 fn join_dataframes(left: LazyFrame, right: LazyFrame, query: &QueryPlan) -> Result<LazyFrame> {
48 let primary_key = &query.primary_keys.keys[0]; 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 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 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 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 let group_col = &pk_mapping.source_fields[0].column_name;
90
91 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 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 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 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 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 let first_col = &mapping.source_fields[0].column_name;
141 Ok(col(first_col).first())
142 },
143 MergePolicy::Sum => {
144 let first_col = &mapping.source_fields[0].column_name;
146 Ok(col(first_col).sum())
147 },
148 MergePolicy::Count => {
149 let first_col = &mapping.source_fields[0].column_name;
151 Ok(col(first_col).count())
152 },
153 MergePolicy::Average => {
154 let first_col = &mapping.source_fields[0].column_name;
156 Ok(col(first_col).mean())
157 },
158 MergePolicy::Min => {
159 let first_col = &mapping.source_fields[0].column_name;
161 Ok(col(first_col).min())
162 },
163 MergePolicy::Max => {
164 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 let first_col = &mapping.source_fields[0].column_name;
181 Ok(col(first_col))
182 },
183 MergePolicy::Sum => {
184 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 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 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 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 Ok(coalesce(&cols))
220 }
221 },
222 MergePolicy::Max => {
223 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 Ok(coalesce(&cols))
234 }
235 },
236 }
237 }
238}