1use std::collections::{BTreeMap, BTreeSet};
10
11use crate::ast::{BinaryOp, ColumnType, JoinKind, JoinUsing};
12use crate::schema::join_output::JoinOutputSource;
13use crate::SQLError;
14use crate::{ColumnIdentity, RowSchema, ScalarExpr};
15
16#[derive(Debug, Clone)]
17pub struct ResolvedJoinUsing {
18 columns: Vec<ResolvedJoinColumn>,
19 alias: Option<String>,
20}
21
22pub type JoinUsingLayout = (
23 Vec<(String, ColumnIdentity, JoinOutputSource)>,
24 Vec<(ColumnIdentity, JoinOutputSource)>,
25);
26
27#[derive(Debug, Clone)]
28struct ResolvedJoinColumn {
29 name: String,
30 left: usize,
31 right: usize,
32 comparison_type: Option<ColumnType>,
33 output_type: Option<ColumnType>,
34}
35
36fn visible_columns(schema: &RowSchema) -> Vec<(&str, usize)> {
37 schema
38 .columns()
39 .iter()
40 .enumerate()
41 .filter_map(|(position, _)| Some((schema.public_name(position)?, position)))
42 .collect()
43}
44
45fn positions_by_name<'a>(columns: &'a [(&'a str, usize)]) -> BTreeMap<&'a str, Vec<usize>> {
46 let mut positions = BTreeMap::<&str, Vec<usize>>::new();
47 for (name, position) in columns {
48 positions.entry(name).or_default().push(*position);
49 }
50 positions
51}
52
53fn missing_using_column(column: &str, side: &str) -> SQLError {
54 SQLError::Routine {
55 sqlstate: "42703".into(),
56 message: format!(
57 "column \"{column}\" specified in USING clause does not exist in {side} table"
58 ),
59 }
60}
61
62fn ambiguous_using_column(column: &str, side: &str) -> SQLError {
63 SQLError::Routine {
64 sqlstate: "42702".into(),
65 message: format!("common column name \"{column}\" appears more than once in {side} table"),
66 }
67}
68
69fn unique_position(
70 positions: &BTreeMap<&str, Vec<usize>>,
71 column: &str,
72 side: &str,
73) -> Result<usize, SQLError> {
74 match positions.get(column).map(Vec::as_slice) {
75 None | Some([]) => Err(missing_using_column(column, side)),
76 Some([position]) => Ok(*position),
77 Some(_) => Err(ambiguous_using_column(column, side)),
78 }
79}
80
81pub fn resolve_join_using(
82 explicit: Option<&JoinUsing>,
83 natural: bool,
84 left: &RowSchema,
85 right: &RowSchema,
86) -> Result<Option<ResolvedJoinUsing>, SQLError> {
87 if explicit.is_none() && !natural {
88 return Ok(None);
89 }
90 if explicit.is_some() && natural {
91 return Err(SQLError::Internal(
92 "JOIN cannot be both USING-qualified and NATURAL".into(),
93 ));
94 }
95
96 let left_columns = visible_columns(left);
97 let right_columns = visible_columns(right);
98 let left_positions = positions_by_name(&left_columns);
99 let right_positions = positions_by_name(&right_columns);
100 let (names, alias) = if let Some(using) = explicit {
101 let mut seen = BTreeSet::new();
102 for name in &using.columns {
103 if !seen.insert(name.as_str()) {
104 return Err(SQLError::Routine {
105 sqlstate: "42701".into(),
106 message: format!(
107 "column name \"{name}\" appears more than once in USING clause"
108 ),
109 });
110 }
111 }
112 (using.columns.clone(), using.alias.clone())
113 } else {
114 let mut names = Vec::new();
115 let mut seen = BTreeSet::new();
116 for (name, _) in &left_columns {
117 if right_positions.contains_key(name) && seen.insert(*name) {
118 names.push((*name).to_string());
119 }
120 }
121 (names, None)
122 };
123
124 if let Some(alias) = alias.as_deref() {
125 if left.has_qualifier(alias) || right.has_qualifier(alias) {
126 return Err(SQLError::Routine {
127 sqlstate: "42712".into(),
128 message: format!("table name \"{alias}\" specified more than once"),
129 });
130 }
131 }
132
133 let mut columns = Vec::with_capacity(names.len());
134 for name in names {
135 let left_position = unique_position(&left_positions, &name, "left")?;
136 let right_position = unique_position(&right_positions, &name, "right")?;
137 let (comparison_type, output_type) = match (
138 left.column_type(left_position),
139 right.column_type(right_position),
140 ) {
141 (Some(left_type), Some(right_type)) => (
142 Some(crate::equality_operand_type(left_type, right_type)?),
143 Some(crate::common_type(left_type, right_type)?),
144 ),
145 _ => (None, None),
146 };
147 columns.push(ResolvedJoinColumn {
148 left: left_position,
149 right: right_position,
150 name,
151 comparison_type,
152 output_type,
153 });
154 }
155 Ok(Some(ResolvedJoinUsing { columns, alias }))
156}
157
158pub fn join_using_predicate(
159 using: &ResolvedJoinUsing,
160 left: &RowSchema,
161 right: &RowSchema,
162) -> Option<ScalarExpr> {
163 let mut predicates = using
164 .columns
165 .iter()
166 .map(|column| {
167 let lhs = coerced_join_column(
168 scalar_column(left, column.left),
169 left.column_type(column.left),
170 column.comparison_type.as_ref(),
171 );
172 let rhs = coerced_join_column(
173 scalar_column(right, column.right),
174 right.column_type(column.right),
175 column.comparison_type.as_ref(),
176 );
177 ScalarExpr::Binary {
178 op: BinaryOp::Equal,
179 lhs: Box::new(lhs),
180 rhs: Box::new(rhs),
181 }
182 })
183 .collect::<Vec<_>>();
184 match predicates.len() {
185 0 => None,
186 1 => predicates.pop(),
187 _ => Some(ScalarExpr::And(predicates)),
188 }
189}
190
191fn scalar_column(schema: &RowSchema, position: usize) -> ScalarExpr {
192 let identity = &schema.identities()[position];
193 identity.qualifier().map_or_else(
194 || ScalarExpr::Column(identity.column().to_string()),
195 |qualifier| ScalarExpr::qualified_column(qualifier, identity.column()),
196 )
197}
198
199fn coerced_join_column(
200 expression: ScalarExpr,
201 source: Option<&ColumnType>,
202 target: Option<&ColumnType>,
203) -> ScalarExpr {
204 match (source, target) {
205 (Some(source), Some(target)) if source != target => ScalarExpr::Cast {
206 expr: Box::new(expression),
207 ty: target.sql_name(),
208 },
209 _ => expression,
210 }
211}
212
213pub fn join_using_output_schema(
218 kind: JoinKind,
219 left: &RowSchema,
220 right: &RowSchema,
221 using: &ResolvedJoinUsing,
222) -> Result<RowSchema, SQLError> {
223 let input = RowSchema::join(left, right, std::iter::empty());
224 let (columns, aliases) = join_using_layout(kind, left, right, using)?;
225 crate::schema::join_output::compile_layout(&input, &columns, &aliases)
226 .map(|(schema, _)| schema)
227 .map_err(|error| SQLError::Internal(error.0))
228}
229
230#[expect(
232 clippy::too_many_lines,
233 reason = "preserves source schema and row identity"
234)]
235pub fn join_alias_input_schemas(
236 kind: JoinKind,
237 left: &RowSchema,
238 right: &RowSchema,
239 using: Option<&ResolvedJoinUsing>,
240 alias: &str,
241 column_aliases: &[String],
242) -> Result<(RowSchema, RowSchema), SQLError> {
243 #[derive(Clone, Copy)]
244 enum Owner {
245 Left(usize),
246 Right(usize),
247 Neither,
248 }
249
250 let mut output = Vec::<(String, Owner)>::new();
251 if let Some(using) = using {
252 let mut left_used = vec![false; left.len()];
253 let mut right_used = vec![false; right.len()];
254 for column in &using.columns {
255 left_used[column.left] = true;
256 right_used[column.right] = true;
257 let owner = match kind {
258 JoinKind::Inner | JoinKind::Left => Owner::Left(column.left),
259 JoinKind::Right => Owner::Right(column.right),
260 JoinKind::Full => Owner::Neither,
261 JoinKind::Cross => {
262 return Err(SQLError::Internal(
263 "CROSS JOIN cannot carry a USING qualification".into(),
264 ));
265 }
266 };
267 output.push((column.name.clone(), owner));
268 }
269 output.extend(
270 left.columns()
271 .iter()
272 .enumerate()
273 .filter(|(position, _)| !left_used[*position])
274 .map(|(position, column)| {
275 (
276 left.public_name(position).unwrap_or(column).to_string(),
277 Owner::Left(position),
278 )
279 }),
280 );
281 output.extend(
282 right
283 .columns()
284 .iter()
285 .enumerate()
286 .filter(|(position, _)| !right_used[*position])
287 .map(|(position, column)| {
288 (
289 right.public_name(position).unwrap_or(column).to_string(),
290 Owner::Right(position),
291 )
292 }),
293 );
294 } else {
295 output.extend(left.columns().iter().enumerate().map(|(position, column)| {
296 (
297 left.public_name(position).unwrap_or(column).to_string(),
298 Owner::Left(position),
299 )
300 }));
301 output.extend(
302 right
303 .columns()
304 .iter()
305 .enumerate()
306 .map(|(position, column)| {
307 (
308 right.public_name(position).unwrap_or(column).to_string(),
309 Owner::Right(position),
310 )
311 }),
312 );
313 }
314
315 if column_aliases.len() > output.len() {
316 return Err(SQLError::Routine {
317 sqlstate: "42P10".into(),
318 message: format!(
319 "join expression \"{alias}\" has {} columns available but {} columns specified",
320 output.len(),
321 column_aliases.len()
322 ),
323 });
324 }
325 for (position, column_alias) in column_aliases.iter().enumerate() {
326 output[position].0.clone_from(column_alias);
327 }
328
329 let mut counts = BTreeMap::<String, usize>::new();
330 for (column, _) in &output {
331 *counts.entry(column.clone()).or_default() += 1;
332 }
333 let mut left_aliases = Vec::new();
334 let mut right_aliases = Vec::new();
335 for (column, owner) in output {
336 if counts.get(&column) != Some(&1) {
337 continue;
338 }
339 let identity = ColumnIdentity::qualified(alias, column);
340 match owner {
341 Owner::Left(position) => left_aliases.push((identity, position)),
342 Owner::Right(position) => right_aliases.push((identity, position)),
343 Owner::Neither => {}
344 }
345 }
346 Ok((
347 RowSchema::with_identity_aliases(left, &left_aliases),
348 RowSchema::with_identity_aliases(right, &right_aliases),
349 ))
350}
351
352pub fn join_using_layout(
353 kind: JoinKind,
354 left: &RowSchema,
355 right: &RowSchema,
356 using: &ResolvedJoinUsing,
357) -> Result<JoinUsingLayout, SQLError> {
358 let right_offset = left.len();
359 let mut left_used = vec![false; left.len()];
360 let mut right_used = vec![false; right.len()];
361 let mut columns = Vec::with_capacity(left.len() + right.len() - using.columns.len());
362 let mut merged_sources = Vec::with_capacity(using.columns.len());
363
364 for column in &using.columns {
365 left_used[column.left] = true;
366 right_used[column.right] = true;
367 let source = match kind {
368 JoinKind::Inner | JoinKind::Left => coerced_output_source(
369 column.left,
370 left.column_type(column.left),
371 column.output_type.as_ref(),
372 ),
373 JoinKind::Right => coerced_output_source(
374 right_offset + column.right,
375 right.column_type(column.right),
376 column.output_type.as_ref(),
377 ),
378 JoinKind::Full => {
379 column
380 .output_type
381 .as_ref()
382 .map_or(JoinOutputSource::Input(column.left), |ty| {
383 JoinOutputSource::Coalesce {
384 left: column.left,
385 right: right_offset + column.right,
386 ty: ty.clone(),
387 }
388 })
389 }
390 JoinKind::Cross => {
391 return Err(SQLError::Internal(
392 "CROSS JOIN cannot carry a USING qualification".into(),
393 ));
394 }
395 };
396 columns.push((
397 column.name.clone(),
398 ColumnIdentity::unqualified(&column.name),
399 source.clone(),
400 ));
401 merged_sources.push((column.name.clone(), source));
402 }
403 for (position, column) in left.columns().iter().enumerate() {
404 if !left_used[position] {
405 columns.push((
406 left.public_name(position).unwrap_or(column).to_string(),
407 left.identities()[position].clone(),
408 JoinOutputSource::Input(position),
409 ));
410 }
411 }
412 for (position, column) in right.columns().iter().enumerate() {
413 if !right_used[position] {
414 columns.push((
415 right.public_name(position).unwrap_or(column).to_string(),
416 right.identities()[position].clone(),
417 JoinOutputSource::Input(right_offset + position),
418 ));
419 }
420 }
421
422 let mut aliases = Vec::new();
423 for (position, identity) in left.identities().iter().enumerate() {
424 if identity.qualifier().is_some() {
425 aliases.push((identity.clone(), JoinOutputSource::Input(position)));
426 }
427 }
428 for (position, identity) in right.identities().iter().enumerate() {
429 if identity.qualifier().is_some() {
430 aliases.push((
431 identity.clone(),
432 JoinOutputSource::Input(right_offset + position),
433 ));
434 }
435 }
436 if let Some(alias) = using.alias.as_ref() {
437 aliases.extend(
438 merged_sources
439 .into_iter()
440 .map(|(column, source)| (ColumnIdentity::qualified(alias, column), source)),
441 );
442 }
443
444 Ok((columns, aliases))
445}
446
447fn coerced_output_source(
448 input: usize,
449 source: Option<&ColumnType>,
450 target: Option<&ColumnType>,
451) -> JoinOutputSource {
452 match (source, target) {
453 (Some(source), Some(target)) if source != target => JoinOutputSource::Cast {
454 input,
455 ty: target.clone(),
456 },
457 _ => JoinOutputSource::Input(input),
458 }
459}
460
461#[cfg(test)]
462mod tests {
463 use super::*;
464
465 #[test]
466 fn natural_columns_follow_left_input_order() {
467 let left = RowSchema::with_qualified_types(
468 "l",
469 vec!["b".into(), "a".into(), "only".into()],
470 vec![None; 3],
471 );
472 let right = RowSchema::with_qualified_types(
473 "r",
474 vec!["a".into(), "b".into(), "other".into()],
475 vec![None; 3],
476 );
477 let resolved = resolve_join_using(None, true, &left, &right)
478 .unwrap()
479 .unwrap();
480 assert_eq!(
481 resolved
482 .columns
483 .iter()
484 .map(|column| column.name.as_str())
485 .collect::<Vec<_>>(),
486 ["b", "a"]
487 );
488 }
489
490 #[test]
491 fn using_requires_one_column_on_each_side() {
492 let left = RowSchema::with_identities(
493 vec!["id".into(), "id".into()],
494 vec![
495 ColumnIdentity::qualified("l", "id"),
496 ColumnIdentity::qualified("other", "id"),
497 ],
498 vec![None; 2],
499 );
500 let right = RowSchema::with_qualified_types("r", vec!["id".into()], vec![None]);
501 let error = resolve_join_using(
502 Some(&JoinUsing {
503 columns: vec!["id".into()],
504 alias: None,
505 }),
506 false,
507 &left,
508 &right,
509 )
510 .unwrap_err();
511 assert_eq!(error.sqlstate(), Some("42702"));
512 }
513}