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::type_resolution::common_type_in(
144 crate::type_resolution::CommonTypeContext::JoinUsing,
145 left_type,
146 right_type,
147 )?),
148 ),
149 _ => (None, None),
150 };
151 columns.push(ResolvedJoinColumn {
152 left: left_position,
153 right: right_position,
154 name,
155 comparison_type,
156 output_type,
157 });
158 }
159 Ok(Some(ResolvedJoinUsing { columns, alias }))
160}
161
162pub fn join_using_predicate(
163 using: &ResolvedJoinUsing,
164 left: &RowSchema,
165 right: &RowSchema,
166) -> Option<ScalarExpr> {
167 let mut predicates = using
168 .columns
169 .iter()
170 .map(|column| {
171 let lhs = coerced_join_column(
172 scalar_column(left, column.left),
173 left.column_type(column.left),
174 column.comparison_type.as_ref(),
175 );
176 let rhs = coerced_join_column(
177 scalar_column(right, column.right),
178 right.column_type(column.right),
179 column.comparison_type.as_ref(),
180 );
181 ScalarExpr::Binary {
182 op: BinaryOp::Equal,
183 lhs: Box::new(lhs),
184 rhs: Box::new(rhs),
185 }
186 })
187 .collect::<Vec<_>>();
188 match predicates.len() {
189 0 => None,
190 1 => predicates.pop(),
191 _ => Some(ScalarExpr::And(predicates)),
192 }
193}
194
195fn scalar_column(schema: &RowSchema, position: usize) -> ScalarExpr {
196 let identity = &schema.identities()[position];
197 identity.qualifier().map_or_else(
198 || ScalarExpr::Column(identity.column().to_string()),
199 |qualifier| ScalarExpr::qualified_column(qualifier, identity.column()),
200 )
201}
202
203fn coerced_join_column(
204 expression: ScalarExpr,
205 source: Option<&ColumnType>,
206 target: Option<&ColumnType>,
207) -> ScalarExpr {
208 match (source, target) {
209 (Some(source), Some(target)) if source != target => ScalarExpr::Cast {
210 implicit: true,
211 expr: Box::new(expression),
212 ty: target.sql_name(),
213 },
214 _ => expression,
215 }
216}
217
218pub fn join_using_output_schema(
223 kind: JoinKind,
224 left: &RowSchema,
225 right: &RowSchema,
226 using: &ResolvedJoinUsing,
227) -> Result<RowSchema, SQLError> {
228 let input = RowSchema::join(left, right, std::iter::empty());
229 let (columns, aliases) = join_using_layout(kind, left, right, using)?;
230 crate::schema::join_output::compile_layout(&input, &columns, &aliases)
231 .map(|(schema, _)| schema)
232 .map_err(|error| SQLError::Internal(error.0))
233}
234
235#[expect(
237 clippy::too_many_lines,
238 reason = "preserves source schema and row identity"
239)]
240pub fn join_alias_input_schemas(
241 kind: JoinKind,
242 left: &RowSchema,
243 right: &RowSchema,
244 using: Option<&ResolvedJoinUsing>,
245 alias: &str,
246 column_aliases: &[String],
247) -> Result<(RowSchema, RowSchema), SQLError> {
248 #[derive(Clone, Copy)]
249 enum Owner {
250 Left(usize),
251 Right(usize),
252 Neither,
253 }
254
255 let mut output = Vec::<(String, Owner)>::new();
256 if let Some(using) = using {
257 let mut left_used = vec![false; left.len()];
258 let mut right_used = vec![false; right.len()];
259 for column in &using.columns {
260 left_used[column.left] = true;
261 right_used[column.right] = true;
262 let owner = match kind {
263 JoinKind::Inner | JoinKind::Left => Owner::Left(column.left),
264 JoinKind::Right => Owner::Right(column.right),
265 JoinKind::Full => Owner::Neither,
266 JoinKind::Cross => {
267 return Err(SQLError::Internal(
268 "CROSS JOIN cannot carry a USING qualification".into(),
269 ));
270 }
271 };
272 output.push((column.name.clone(), owner));
273 }
274 output.extend(
275 left.columns()
276 .iter()
277 .enumerate()
278 .filter(|(position, _)| !left_used[*position])
279 .map(|(position, column)| {
280 (
281 left.public_name(position).unwrap_or(column).to_string(),
282 Owner::Left(position),
283 )
284 }),
285 );
286 output.extend(
287 right
288 .columns()
289 .iter()
290 .enumerate()
291 .filter(|(position, _)| !right_used[*position])
292 .map(|(position, column)| {
293 (
294 right.public_name(position).unwrap_or(column).to_string(),
295 Owner::Right(position),
296 )
297 }),
298 );
299 } else {
300 output.extend(left.columns().iter().enumerate().map(|(position, column)| {
301 (
302 left.public_name(position).unwrap_or(column).to_string(),
303 Owner::Left(position),
304 )
305 }));
306 output.extend(
307 right
308 .columns()
309 .iter()
310 .enumerate()
311 .map(|(position, column)| {
312 (
313 right.public_name(position).unwrap_or(column).to_string(),
314 Owner::Right(position),
315 )
316 }),
317 );
318 }
319
320 if column_aliases.len() > output.len() {
321 return Err(SQLError::Routine {
322 sqlstate: "42P10".into(),
323 message: format!(
324 "join expression \"{alias}\" has {} columns available but {} columns specified",
325 output.len(),
326 column_aliases.len()
327 ),
328 });
329 }
330 for (position, column_alias) in column_aliases.iter().enumerate() {
331 output[position].0.clone_from(column_alias);
332 }
333
334 let mut counts = BTreeMap::<String, usize>::new();
335 for (column, _) in &output {
336 *counts.entry(column.clone()).or_default() += 1;
337 }
338 let mut left_aliases = Vec::new();
339 let mut right_aliases = Vec::new();
340 for (column, owner) in output {
341 if counts.get(&column) != Some(&1) {
342 continue;
343 }
344 let identity = ColumnIdentity::qualified(alias, column);
345 match owner {
346 Owner::Left(position) => left_aliases.push((identity, position)),
347 Owner::Right(position) => right_aliases.push((identity, position)),
348 Owner::Neither => {}
349 }
350 }
351 Ok((
352 RowSchema::with_identity_aliases(left, &left_aliases),
353 RowSchema::with_identity_aliases(right, &right_aliases),
354 ))
355}
356
357pub fn join_using_layout(
358 kind: JoinKind,
359 left: &RowSchema,
360 right: &RowSchema,
361 using: &ResolvedJoinUsing,
362) -> Result<JoinUsingLayout, SQLError> {
363 let right_offset = left.len();
364 let mut left_used = vec![false; left.len()];
365 let mut right_used = vec![false; right.len()];
366 let mut columns = Vec::with_capacity(left.len() + right.len() - using.columns.len());
367 let mut merged_sources = Vec::with_capacity(using.columns.len());
368
369 for column in &using.columns {
370 left_used[column.left] = true;
371 right_used[column.right] = true;
372 let source = match kind {
373 JoinKind::Inner | JoinKind::Left => coerced_output_source(
374 column.left,
375 left.column_type(column.left),
376 column.output_type.as_ref(),
377 ),
378 JoinKind::Right => coerced_output_source(
379 right_offset + column.right,
380 right.column_type(column.right),
381 column.output_type.as_ref(),
382 ),
383 JoinKind::Full => {
384 column
385 .output_type
386 .as_ref()
387 .map_or(JoinOutputSource::Input(column.left), |ty| {
388 JoinOutputSource::Coalesce {
389 left: column.left,
390 right: right_offset + column.right,
391 ty: ty.clone(),
392 }
393 })
394 }
395 JoinKind::Cross => {
396 return Err(SQLError::Internal(
397 "CROSS JOIN cannot carry a USING qualification".into(),
398 ));
399 }
400 };
401 columns.push((
402 column.name.clone(),
403 ColumnIdentity::unqualified(&column.name),
404 source.clone(),
405 ));
406 merged_sources.push((column.name.clone(), source));
407 }
408 for (position, column) in left.columns().iter().enumerate() {
409 if !left_used[position] {
410 columns.push((
411 left.public_name(position).unwrap_or(column).to_string(),
412 left.identities()[position].clone(),
413 JoinOutputSource::Input(position),
414 ));
415 }
416 }
417 for (position, column) in right.columns().iter().enumerate() {
418 if !right_used[position] {
419 columns.push((
420 right.public_name(position).unwrap_or(column).to_string(),
421 right.identities()[position].clone(),
422 JoinOutputSource::Input(right_offset + position),
423 ));
424 }
425 }
426
427 let mut aliases = Vec::new();
428 for (position, identity) in left.identities().iter().enumerate() {
429 if identity.qualifier().is_some() {
430 aliases.push((identity.clone(), JoinOutputSource::Input(position)));
431 }
432 }
433 for (position, identity) in right.identities().iter().enumerate() {
434 if identity.qualifier().is_some() {
435 aliases.push((
436 identity.clone(),
437 JoinOutputSource::Input(right_offset + position),
438 ));
439 }
440 }
441 if let Some(alias) = using.alias.as_ref() {
442 aliases.extend(
443 merged_sources
444 .into_iter()
445 .map(|(column, source)| (ColumnIdentity::qualified(alias, column), source)),
446 );
447 }
448
449 Ok((columns, aliases))
450}
451
452fn coerced_output_source(
453 input: usize,
454 source: Option<&ColumnType>,
455 target: Option<&ColumnType>,
456) -> JoinOutputSource {
457 match (source, target) {
458 (Some(source), Some(target)) if source != target => JoinOutputSource::Cast {
459 input,
460 ty: target.clone(),
461 },
462 _ => JoinOutputSource::Input(input),
463 }
464}
465
466#[cfg(test)]
467mod tests {
468 use super::*;
469
470 #[test]
471 fn natural_columns_follow_left_input_order() {
472 let left = RowSchema::with_qualified_types(
473 "l",
474 vec!["b".into(), "a".into(), "only".into()],
475 vec![None; 3],
476 );
477 let right = RowSchema::with_qualified_types(
478 "r",
479 vec!["a".into(), "b".into(), "other".into()],
480 vec![None; 3],
481 );
482 let resolved = resolve_join_using(None, true, &left, &right)
483 .unwrap()
484 .unwrap();
485 assert_eq!(
486 resolved
487 .columns
488 .iter()
489 .map(|column| column.name.as_str())
490 .collect::<Vec<_>>(),
491 ["b", "a"]
492 );
493 }
494
495 #[test]
496 fn using_requires_one_column_on_each_side() {
497 let left = RowSchema::with_identities(
498 vec!["id".into(), "id".into()],
499 vec![
500 ColumnIdentity::qualified("l", "id"),
501 ColumnIdentity::qualified("other", "id"),
502 ],
503 vec![None; 2],
504 );
505 let right = RowSchema::with_qualified_types("r", vec!["id".into()], vec![None]);
506 let error = resolve_join_using(
507 Some(&JoinUsing {
508 columns: vec!["id".into()],
509 alias: None,
510 }),
511 false,
512 &left,
513 &right,
514 )
515 .unwrap_err();
516 assert_eq!(error.sqlstate(), Some("42702"));
517 }
518}