1use super::declaration::RoutineTypeCatalog;
10use crate::{
11 ast::{ColumnType, CreateFunction, FunctionReturns},
12 binding::statements::AnalyzedResult,
13 expr::composites,
14 type_resolution::canonical_routine_type_name,
15 SQLError,
16};
17
18mod anonymous;
19pub use anonymous::validate_anonymous_record_result;
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum SQLFunctionResultKind {
24 Void,
25 Value,
26 Tuple,
27}
28
29#[derive(Debug, Clone, PartialEq, Eq)]
31pub struct SQLFunctionResultColumn {
32 pub name: String,
33 pub ty: ColumnType,
34}
35
36#[derive(Debug, Clone, PartialEq, Eq)]
38pub struct SQLFunctionResultLayout {
39 pub kind: SQLFunctionResultKind,
40 pub declared_type: ColumnType,
41 pub source_record: Option<Vec<Option<ColumnType>>>,
43 pub columns: Option<Vec<SQLFunctionResultColumn>>,
45}
46
47pub fn check_sql_function_result(
49 types: &dyn RoutineTypeCatalog,
50 def: &CreateFunction,
51 last: Option<&AnalyzedResult>,
52) -> Result<(), SQLError> {
53 sql_function_result_layout(types, def, last).map(|_| ())
54}
55
56pub fn sql_function_result_layout(
58 types: &dyn RoutineTypeCatalog,
59 def: &CreateFunction,
60 last: Option<&AnalyzedResult>,
61) -> Result<SQLFunctionResultLayout, SQLError> {
62 let mut layout = declared_sql_function_result(types, def)?;
63 if layout.kind == SQLFunctionResultKind::Void {
64 return Ok(layout);
65 }
66 let mismatch = |detail: String| -> Result<SQLError, SQLError> {
67 Ok(SQLError::Diagnostic {
68 sqlstate: "42P13".into(),
69 message: format!(
70 "return type mismatch in function declared to return {}",
71 types.format_type(&layout.declared_type)?
72 ),
73 detail: Some(detail),
74 hint: None,
75 })
76 };
77 let Some(columns) = last.and_then(AnalyzedResult::column_types) else {
78 return Err(mismatch(
79 "Function's final statement must be SELECT or INSERT/UPDATE/DELETE/MERGE RETURNING."
80 .into(),
81 )?);
82 };
83 if let [actual] = columns {
84 layout.source_record = last
85 .and_then(|result| result.record_fields(0))
86 .map(|fields| fields.to_vec())
87 .or(sql_function_composite_columns(types, actual.as_ref())?);
88 }
89 if layout.kind == SQLFunctionResultKind::Value {
90 let [actual] = columns else {
91 return Err(mismatch(
92 "Final statement must return exactly one column.".into(),
93 )?);
94 };
95 if let Some(actual) = actual {
96 if !crate::assignment_type_compatible(actual, &layout.declared_type) {
97 return Err(mismatch(format!(
98 "Actual return type is {}.",
99 types.format_type(actual)?
100 ))?);
101 }
102 }
103 return Ok(layout);
104 }
105 if !def.is_procedure {
106 if let [Some(actual)] = columns {
107 if crate::assignment_type_compatible(actual, &layout.declared_type) {
108 layout.kind = SQLFunctionResultKind::Value;
109 if matches!(actual, ColumnType::Record)
110 && matches!(layout.declared_type, ColumnType::Composite(_))
111 {
112 if let Some(source) = &layout.source_record {
113 check_composite_assignment(types, source, &layout)?;
114 }
115 }
116 return Ok(layout);
117 }
118 }
119 }
120 if let Some(expected) = &layout.columns {
121 for (position, actual) in columns.iter().enumerate() {
122 let Some(expected) = expected.get(position) else {
123 return Err(mismatch(
124 "Final statement returns too many columns.".into(),
125 )?);
126 };
127 if let Some(actual) = actual {
128 if !crate::assignment_type_compatible(actual, &expected.ty) {
129 return Err(mismatch(format!(
130 "Final statement returns {} instead of {} at column {}.",
131 types.format_type(actual)?,
132 types.format_type(&expected.ty)?,
133 position + 1
134 ))?);
135 }
136 }
137 }
138 if columns.len() < expected.len() {
139 return Err(mismatch("Final statement returns too few columns.".into())?);
140 }
141 }
142 Ok(layout)
143}
144
145pub fn declared_sql_function_result(
147 types: &dyn RoutineTypeCatalog,
148 def: &CreateFunction,
149) -> Result<SQLFunctionResultLayout, SQLError> {
150 let outputs = def.output_params();
151 if def.is_procedure || outputs.len() > 1 {
152 let columns = outputs
153 .iter()
154 .map(|parameter| {
155 Ok(SQLFunctionResultColumn {
156 name: parameter.name.clone(),
157 ty: types.resolve_catalog_column_type_name(¶meter.type_name)?,
158 })
159 })
160 .collect::<Result<Vec<_>, SQLError>>()?;
161 return Ok(SQLFunctionResultLayout {
162 kind: if columns.is_empty() {
163 SQLFunctionResultKind::Void
164 } else {
165 SQLFunctionResultKind::Tuple
166 },
167 declared_type: ColumnType::Record,
168 source_record: None,
169 columns: Some(columns),
170 });
171 }
172 let type_name = match (outputs.first(), &def.returns) {
173 (Some(parameter), _) => parameter.type_name.as_str(),
174 (None, FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name }) => {
175 type_name.as_str()
176 }
177 (None, FunctionReturns::None | FunctionReturns::Table) => "void",
178 };
179 let declared_type = match canonical_routine_type_name(type_name).as_str() {
180 "void" => ColumnType::Void,
181 "record" => ColumnType::Record,
182 _ => types.resolve_catalog_column_type_name(type_name)?,
183 };
184 let (kind, columns) = match &declared_type {
185 ColumnType::Void => (SQLFunctionResultKind::Void, None),
186 ColumnType::Record => (SQLFunctionResultKind::Tuple, None),
187 ColumnType::Composite(reference) => {
188 let descriptor = composites::descriptor(types.composite_types(), reference.oid)?;
189 let columns = descriptor
190 .attributes
191 .iter()
192 .map(|attribute| SQLFunctionResultColumn {
193 name: attribute.name.clone(),
194 ty: attribute.ty.clone(),
195 })
196 .collect();
197 (SQLFunctionResultKind::Tuple, Some(columns))
198 }
199 _ => (SQLFunctionResultKind::Value, None),
200 };
201 Ok(SQLFunctionResultLayout {
202 kind,
203 declared_type,
204 source_record: None,
205 columns,
206 })
207}
208
209fn check_composite_assignment(
210 types: &dyn RoutineTypeCatalog,
211 source: &[Option<ColumnType>],
212 layout: &SQLFunctionResultLayout,
213) -> Result<(), SQLError> {
214 let target = layout.columns.as_deref().unwrap_or_default();
215 let detail = match source.len().cmp(&target.len()) {
216 std::cmp::Ordering::Less => Some("Input has too few columns.".into()),
217 std::cmp::Ordering::Greater => Some("Input has too many columns.".into()),
218 std::cmp::Ordering::Equal => source
219 .iter()
220 .zip(target)
221 .enumerate()
222 .find_map(|(index, (source, target))| {
223 source
224 .as_ref()
225 .filter(|source| !crate::assignment_type_compatible(source, &target.ty))
226 .map(|source| {
227 Ok::<_, SQLError>(format!(
228 "Cannot cast type {} to {} in column {}.",
229 types.format_type(source)?,
230 types.format_type(&target.ty)?,
231 index + 1
232 ))
233 })
234 })
235 .transpose()?,
236 };
237 if let Some(detail) = detail {
238 return Err(SQLError::Diagnostic {
239 sqlstate: "42846".into(),
240 message: format!(
241 "cannot cast type record to {}",
242 types.format_type(&layout.declared_type)?
243 ),
244 detail: Some(detail),
245 hint: None,
246 });
247 }
248 Ok(())
249}
250
251pub fn sql_function_composite_columns(
253 types: &dyn RoutineTypeCatalog,
254 source: Option<&ColumnType>,
255) -> Result<Option<Vec<Option<ColumnType>>>, SQLError> {
256 match source {
257 Some(ColumnType::Composite(reference)) => {
258 let descriptor = composites::descriptor(types.composite_types(), reference.oid)?;
259 Ok(Some(
260 descriptor
261 .attributes
262 .iter()
263 .map(|attribute| Some(attribute.ty.clone()))
264 .collect(),
265 ))
266 }
267 Some(ColumnType::Domain { base, .. }) => sql_function_composite_columns(types, Some(base)),
268 _ => Ok(None),
269 }
270}
271
272pub fn validate_sql_function_record(
274 types: &dyn RoutineTypeCatalog,
275 source: &[Option<ColumnType>],
276 target: &[ColumnType],
277) -> Result<(), SQLError> {
278 if source.len() != target.len() {
279 return Err(record_mismatch(format!(
280 "Returned row contains {} attribute{}, but query expects {}.",
281 source.len(),
282 if source.len() == 1 { "" } else { "s" },
283 target.len()
284 )));
285 }
286 for (index, (source, target)) in source.iter().zip(target).enumerate() {
287 let Some(source) = source else {
288 return Err(record_mismatch(format!(
289 "Returned type unknown at ordinal position {}, but query expects {}.",
290 index + 1,
291 types.format_type(target)?
292 )));
293 };
294 let source_oid = crate::catalog::type_metadata::pg_type_oid(source);
295 let target_oid = crate::catalog::type_metadata::pg_type_oid(target);
296 let source_modifier = crate::catalog::type_metadata::pg_type_modifier(source);
297 let target_modifier = crate::catalog::type_metadata::pg_type_modifier(target);
298 if source_oid != target_oid || (target_modifier >= 0 && source_modifier != target_modifier)
299 {
300 return Err(record_mismatch(format!(
301 "Returned type {} at ordinal position {}, but query expects {}.",
302 types.format_type(source)?,
303 index + 1,
304 types.format_type(target)?
305 )));
306 }
307 }
308 Ok(())
309}
310
311fn record_mismatch(detail: String) -> SQLError {
312 SQLError::Diagnostic {
313 sqlstate: "42804".into(),
314 message: "function return row and query-specified return row do not match".into(),
315 detail: Some(detail),
316 hint: None,
317 }
318}
319
320pub fn validate_sql_function_record_rows(result: &crate::SQLResult) -> Result<(), SQLError> {
322 let mut descriptor = None;
323 for index in 0..result.rows.len() {
324 if let Some(uqa_core::Value::Row(row)) = result.value_at(index, 0) {
325 if let Some(fields) = row.field_types() {
326 if descriptor.is_some_and(|previous| previous != fields) {
327 return Err(SQLError::Routine {
328 sqlstate: "42804".into(),
329 message: "rows returned by function are not all of the same row type"
330 .into(),
331 });
332 }
333 descriptor = Some(fields);
334 }
335 }
336 }
337 Ok(())
338}
339
340pub fn validate_sql_function_record_identity(
342 types: &dyn RoutineTypeCatalog,
343 source: &[uqa_core::RecordFieldType],
344 target: &[ColumnType],
345) -> Result<(), SQLError> {
346 if source.len() != target.len() {
347 return Err(record_mismatch(format!(
348 "Returned row contains {} attribute{}, but query expects {}.",
349 source.len(),
350 if source.len() == 1 { "" } else { "s" },
351 target.len(),
352 )));
353 }
354 for (index, (source, target)) in source.iter().zip(target).enumerate() {
355 let target_modifier = crate::catalog::type_metadata::pg_type_modifier(target);
356 if i64::from(source.oid) == crate::catalog::type_metadata::pg_type_oid(target)
357 && (target_modifier < 0 || i64::from(source.type_modifier) == target_modifier)
358 {
359 continue;
360 }
361 let source_name = if source.oid == 705 {
362 "unknown".into()
363 } else {
364 types.format_type_oid(source.oid)?
365 };
366 return Err(record_mismatch(format!(
367 "Returned type {source_name} at ordinal position {}, but query expects {}.",
368 index + 1,
369 types.format_type(target)?,
370 )));
371 }
372 Ok(())
373}
374
375#[cfg(test)]
376mod tests;