1use super::{
10 array_scalar_type, column_type_name, validate_vector_dimensions, ColumnType, SQLError,
11};
12use crate::expr::{
13 cast_value_from_with_control, parse_pg_array_literal_with_control,
14 value_to_tensor_with_control, value_to_text_with_control, value_to_vector_with_control,
15};
16use uqa_core::{
17 memory::{MemoryReservation, Produced, ProductionControl, ProductionString, ProductionVec},
18 ArrayValue, DecimalValue, Value,
19};
20
21type Result<T> = std::result::Result<T, SQLError>;
22
23#[expect(
24 clippy::too_many_lines,
25 reason = "assignment matrix preserves conversion and validation order"
26)]
27pub fn convert_value_to_column_type_with_control(
28 value: Produced<Value>,
29 ty: &ColumnType,
30 control: &ProductionControl<'_>,
31) -> Result<Produced<Value>> {
32 let (value, memory) = value.into_parts();
33 let value = control.finish(value, memory)?;
34 if matches!(&*value, Value::Null) {
35 return Ok(value);
36 }
37 match ty {
38 ColumnType::Named(name) => Err(SQLError::Routine {
39 sqlstate: "42704".into(),
40 message: format!("type \"{name}\" does not exist"),
41 }),
42 ColumnType::Enum(reference) => match &*value {
44 Value::Enum(label) if label.type_oid() == reference.oid => Ok(value),
45 _ => Err(SQLError::Internal(format!(
46 "enum input for type OID {} requires catalog-aware conversion",
47 reference.oid
48 ))),
49 },
50 ColumnType::Composite(reference) => match &*value {
52 Value::Record(_) => Ok(value),
53 _ => Err(SQLError::Internal(format!(
54 "composite input for type OID {} requires catalog-aware conversion",
55 reference.oid
56 ))),
57 },
58 ColumnType::SmallInteger => cast_value_from_with_control(&value, "smallint", None, control),
59 ColumnType::Integer => cast_value_from_with_control(&value, "integer", None, control),
60 ColumnType::BigInteger => cast_value_from_with_control(&value, "bigint", None, control),
61 ColumnType::Oid => cast_value_from_with_control(&value, "oid", None, control),
63 ColumnType::Xid => {
64 let converted = cast_value_from_with_control(&value, "bigint", None, control)?;
65 let Value::Int(number) = &*converted else {
66 unreachable!("bigint cast returned a non-integer value");
67 };
68 u32::try_from(*number).map_err(|_| {
69 SQLError::TypeMismatch(format!(
70 "value {number} is out of range for type {}",
71 column_type_name(ty)
72 ))
73 })?;
74 Ok(converted)
75 }
76 ColumnType::Boolean => match &*value {
77 Value::Bool(_) => Ok(value),
78 Value::Str(text) => {
79 let boolean = parse_boolean_text(text)
80 .ok_or_else(|| crate::expr::invalid_boolean_input(text))?;
81 Ok(control.finish(Value::Bool(boolean), control.empty_reservation())?)
82 }
83 other => Err(SQLError::TypeMismatch(format!(
84 "cannot cast {other:?} to boolean"
85 ))),
86 },
87 ColumnType::Void => Ok(control.finish(Value::Void, control.empty_reservation())?),
88 ColumnType::Text | ColumnType::RefCursor | ColumnType::Varchar(None) => {
89 text_value(value_to_text_with_control(&value, control)?, false, control)
90 }
91 ColumnType::Name => cast_value_from_with_control(&value, "name", None, control),
92 ColumnType::Uuid => cast_value_from_with_control(&value, "uuid", None, control),
93 ColumnType::Varchar(Some(length)) => character_value(&value, *length, false, control),
94 ColumnType::Bpchar => {
95 text_value(value_to_text_with_control(&value, control)?, true, control)
96 }
97 ColumnType::Character(length) => character_value(&value, *length, true, control),
98 ColumnType::Real => cast_value_from_with_control(&value, "real", None, control),
99 ColumnType::DoublePrecision => {
100 cast_value_from_with_control(&value, "double precision", None, control)
101 }
102 ColumnType::Numeric { precision, scale } => {
103 numeric_value(value, *precision, *scale, control)
104 }
105 ColumnType::Json => cast_value_from_with_control(&value, "json", None, control),
106 ColumnType::JsonB => cast_value_from_with_control(&value, "jsonb", None, control),
107 ColumnType::Bytea => match &*value {
109 Value::Bytes(_) => Ok(value),
110 Value::Str(_) | Value::FixedChar(_) => {
111 cast_value_from_with_control(&value, "bytea", Some("text"), control)
112 }
113 other => Err(SQLError::TypeMismatch(format!(
114 "cannot cast {other:?} to bytea"
115 ))),
116 },
117 ColumnType::InternalChar => {
118 let text = value_to_text_with_control(&value, control)?;
119 if text.len() != 1 {
120 return Err(SQLError::TypeMismatch(format!(
121 "value `{}` must be exactly one byte for type \"char\"",
122 text.as_str()
123 )));
124 }
125 text_value(text, false, control)
126 }
127 ColumnType::Regproc
128 | ColumnType::Regprocedure
129 | ColumnType::Regclass
130 | ColumnType::Regnamespace
131 | ColumnType::Regtype
132 | ColumnType::PgNodeTree
133 | ColumnType::AclItem => {
134 if matches!(&*value, Value::Int(_) | Value::Str(_)) {
135 Ok(value)
136 } else {
137 text_value(value_to_text_with_control(&value, control)?, false, control)
138 }
139 }
140 ColumnType::Regrole => match &*value {
141 Value::Int(number) => {
142 u32::try_from(*number).map_err(|_| {
143 SQLError::TypeMismatch(format!(
144 "value {number} is out of range for type regrole"
145 ))
146 })?;
147 Ok(value)
148 }
149 Value::Str(_) | Value::FixedChar(_) => Err(SQLError::Internal(
150 "regrole name conversion requires catalog resolution".into(),
151 )),
152 other => Err(SQLError::TypeMismatch(format!(
153 "cannot cast {other:?} to regrole"
154 ))),
155 },
156 ColumnType::Int2Vector | ColumnType::OidVector => {
157 let name = column_type_name(ty);
158 if matches!(&*value, Value::LegacyVector(vector) if vector.kind().type_name() == name) {
159 Ok(value)
160 } else {
161 cast_value_from_with_control(&value, name, None, control)
162 }
163 }
164 ColumnType::AnyArray => {
165 if value.array_view().is_some() {
166 Ok(value)
167 } else {
168 Err(SQLError::TypeMismatch(format!(
169 "cannot cast {:?} to anyarray",
170 &*value
171 )))
172 }
173 }
174 ColumnType::Record => match &*value {
175 Value::Record(_) => Ok(value),
176 Value::Row(_) => record_value(value, control),
177 other => Err(SQLError::TypeMismatch(format!(
178 "cannot cast {other:?} to record"
179 ))),
180 },
181 ColumnType::Array(element) => array_value(value, element, control),
182 ColumnType::Date
183 | ColumnType::Time
184 | ColumnType::TimePrecision(_)
185 | ColumnType::TimeTz
186 | ColumnType::TimeTzPrecision(_)
187 | ColumnType::Timestamp
188 | ColumnType::TimestampPrecision(_)
189 | ColumnType::TimestampTz
190 | ColumnType::TimestampTzPrecision(_)
191 | ColumnType::Interval
192 | ColumnType::IntervalWithFields { .. } => {
193 let name = ty.sql_name_with_control(control)?;
194 cast_value_from_with_control(&value, &name, None, control)
195 }
196 ColumnType::Range(subtype) => {
197 cast_value_from_with_control(&value, subtype.range_name(), None, control)
198 }
199 ColumnType::Multirange(subtype) => {
200 cast_value_from_with_control(&value, subtype.multirange_name(), None, control)
201 }
202 ColumnType::Vector(dimensions) => {
203 let vector = value_to_vector_with_control(&value, control)?;
204 validate_vector_dimensions(*dimensions, vector.len())?;
205 vector_value(&vector, control)
206 }
207 ColumnType::Tensor(dimensions) => {
208 let tensor = value_to_tensor_with_control(&value, control)?;
209 for vector in &*tensor {
210 control.check()?;
211 validate_vector_dimensions(*dimensions, vector.len())?;
212 }
213 let mut output = ProductionVec::new(*control);
214 output.reserve(tensor.len())?;
215 for vector in &*tensor {
216 output.push_produced(vector_value(vector, control)?)?;
217 }
218 let (values, memory) = output.finish()?.into_parts();
219 Ok(control.finish(Value::List(values), memory)?)
220 }
221 ColumnType::Domain { base, .. } => {
222 convert_value_to_column_type_with_control(value, base, control)
223 }
224 }
225}
226
227fn numeric_value(
228 value: Produced<Value>,
229 precision: Option<u32>,
230 scale: Option<i32>,
231 control: &ProductionControl<'_>,
232) -> Result<Produced<Value>> {
233 let decimal = match &*value {
234 Value::Decimal(_) => {
235 let (Value::Decimal(value), memory) = value.into_parts() else {
236 unreachable!();
237 };
238 control.finish(value, memory)?
239 }
240 Value::Int(number) => DecimalValue::from_i64_with_control(*number, control)?,
241 Value::Bool(boolean) => DecimalValue::from_i64_with_control(i64::from(*boolean), control)?,
242 Value::Float(number) => DecimalValue::from_f64_lossy_with_control(*number, control)?
243 .ok_or_else(|| SQLError::TypeMismatch(format!("cannot cast {number:?} to numeric")))?,
244 Value::Str(text) => DecimalValue::parse_with_control(text, control)?
245 .ok_or_else(|| crate::expr::invalid_numeric_input(text))?,
246 other => {
247 return Err(SQLError::TypeMismatch(format!(
248 "cannot cast {other:?} to numeric"
249 )))
250 }
251 };
252 let rounded = match scale {
253 Some(scale) => decimal
254 .round_to_scale_with_control(scale, control)?
255 .ok_or_else(|| {
256 SQLError::TypeMismatch(format!("cannot round numeric to scale {scale}"))
257 })?,
258 None => decimal,
259 };
260 if let Some(precision) = precision {
261 let scale = scale.unwrap_or(0);
262 if !rounded.fits_precision_with_control(precision, scale, control)? {
263 return Err(numeric_field_overflow(precision, scale));
264 }
265 }
266 let (decimal, memory) = rounded.into_parts();
267 Ok(control.finish(Value::Decimal(decimal), memory)?)
268}
269
270fn text_value(
271 text: Produced<String>,
272 fixed: bool,
273 control: &ProductionControl<'_>,
274) -> Result<Produced<Value>> {
275 let (text, memory) = text.into_parts();
276 Ok(control.finish(
277 if fixed {
278 Value::FixedChar(text)
279 } else {
280 Value::Str(text)
281 },
282 memory,
283 )?)
284}
285
286fn character_value(
287 value: &Value,
288 length: u32,
289 fixed: bool,
290 control: &ProductionControl<'_>,
291) -> Result<Produced<Value>> {
292 let name = if fixed {
293 "character"
294 } else {
295 "character varying"
296 };
297 let length = usize::try_from(length).map_err(|_| {
298 SQLError::TypeMismatch(format!(
299 "{name} length {length} exceeds the platform addressable range"
300 ))
301 })?;
302 let text = value_to_text_with_control(value, control)?;
303 let mut count = 0;
304 let mut end = 0;
305 for (index, character) in text.char_indices() {
306 control.check()?;
307 if count < length {
308 count += 1;
309 end = index + character.len_utf8();
310 } else if character != ' ' {
311 return Err(SQLError::Routine {
312 sqlstate: "22001".into(),
313 message: format!("value too long for type {name}({length})"),
314 });
315 }
316 }
317 let mut output = ProductionString::from_produced(text, *control)?;
318 output.truncate(end)?;
319 if fixed {
320 for _ in count..length {
321 output.push(' ')?;
322 }
323 }
324 text_value(output.finish()?, fixed, control)
325}
326
327fn array_value(
328 value: Produced<Value>,
329 element: &ColumnType,
330 control: &ProductionControl<'_>,
331) -> Result<Produced<Value>> {
332 let array = match &*value {
333 Value::Array(_) => {
334 let (Value::Array(array), memory) = value.into_parts() else {
335 unreachable!();
336 };
337 control.finish(array, memory)?
338 }
339 Value::LegacyVector(_) => {
340 let (Value::LegacyVector(vector), memory) = value.into_parts() else {
341 unreachable!();
342 };
343 control.finish(vector.into_array(), memory)?
344 }
345 Value::List(_) => {
346 let (Value::List(elements), memory) = value.into_parts() else {
347 unreachable!();
348 };
349 ArrayValue::try_new_with_control(control.finish(elements, memory)?, control)?
350 .ok_or_else(array_shape_error)?
351 }
352 Value::Str(text) => parse_pg_array_literal_with_control(text, control)?,
353 other => {
354 return Err(SQLError::TypeMismatch(format!(
355 "cannot cast {other:?} to {}[]",
356 column_type_name(element)
357 )))
358 }
359 };
360 let converted = array_elements(array.elements(), array_scalar_type(element), control)?;
361 let mut bounds = ProductionVec::new(*control);
362 bounds.reserve(array.lower_bounds().len())?;
363 for bound in array.lower_bounds() {
364 bounds.push_copy(*bound)?;
365 }
366 let array = ArrayValue::with_lower_bounds_with_control(converted, bounds.finish()?, control)?
367 .ok_or_else(array_shape_error)?;
368 let (array, memory) = array.into_parts();
369 Ok(control.finish(Value::Array(array), memory)?)
370}
371
372fn array_elements(
373 values: &[Value],
374 element: &ColumnType,
375 control: &ProductionControl<'_>,
376) -> Result<Produced<Vec<Value>>> {
377 let mut output = ProductionVec::new(*control);
378 output.reserve(values.len())?;
379 for value in values {
380 let value = match value {
381 Value::List(values) => {
382 let (values, memory) = array_elements(values, element, control)?.into_parts();
383 control.finish(Value::List(values), memory)?
384 }
385 scalar => convert_value_to_column_type_with_control(
386 control.copy_value(scalar)?,
387 element,
388 control,
389 )?,
390 };
391 output.push_produced(value)?;
392 }
393 Ok(output.finish()?)
394}
395
396fn array_shape_error() -> SQLError {
397 SQLError::TypeMismatch("multidimensional arrays must have matching dimensions".into())
398}
399
400fn record_value(
401 value: Produced<Value>,
402 control: &ProductionControl<'_>,
403) -> Result<Produced<Value>> {
404 let Value::Row(source) = &*value else {
405 unreachable!();
406 };
407 if source.field_types().is_some() {
409 return Ok(value);
410 }
411 let mut records = ProductionVec::new(*control);
412 records.reserve(source.len())?;
413 for index in 0..source.len() {
414 let (name, memory) = control.format(format_args!("f{}", index + 1))?.into_parts();
415 records.push_produced(control.finish((name, Value::Null), memory)?)?;
416 }
417 let (records, records_memory) = records.finish()?.into_parts();
418 let (Value::Row(source), source_memory) = value.into_parts() else {
419 unreachable!();
420 };
421 let old_buffer_bytes =
422 source.capacity() * size_of::<Value>() + uqa_core::RowValue::retained_header_bytes();
423 let mut parts = RecordParts {
424 source: source.into_values(),
425 records,
426 memory: control.combine(source_memory, records_memory),
427 };
428 for ((_, destination), source) in parts.records.iter_mut().zip(parts.source) {
429 control.check()?;
430 *destination = source;
431 }
432 if let Some(memory) = &mut parts.memory {
433 drop(memory.split(old_buffer_bytes));
434 }
435 Ok(control.finish(Value::Record(parts.records), parts.memory)?)
436}
437
438struct RecordParts {
439 source: Vec<Value>,
440 records: Vec<(String, Value)>,
441 memory: Option<MemoryReservation>,
442}
443
444fn vector_value(vector: &[f32], control: &ProductionControl<'_>) -> Result<Produced<Value>> {
445 let mut output = ProductionVec::new(*control);
446 output.reserve(vector.len())?;
447 for value in vector {
448 output.push_produced(
449 control.finish(Value::Float(f64::from(*value)), control.empty_reservation())?,
450 )?;
451 }
452 let (values, memory) = output.finish()?.into_parts();
453 Ok(control.finish(Value::List(values), memory)?)
454}
455
456fn parse_boolean_text(text: &str) -> Option<bool> {
457 let text = text.trim();
458 if ["true", "t", "yes", "y", "on", "1"]
459 .iter()
460 .any(|candidate| text.eq_ignore_ascii_case(candidate))
461 {
462 Some(true)
463 } else if ["false", "f", "no", "n", "off", "0"]
464 .iter()
465 .any(|candidate| text.eq_ignore_ascii_case(candidate))
466 {
467 Some(false)
468 } else {
469 None
470 }
471}
472
473#[cfg(test)]
474mod tests;
475
476#[must_use]
478pub fn numeric_field_overflow(precision: u32, scale: i32) -> SQLError {
479 let maxdigits = i64::from(precision) - i64::from(scale);
480 SQLError::Diagnostic {
481 sqlstate: "22003".into(),
482 message: "numeric field overflow".into(),
483 detail: Some(format!(
484 "A field with precision {precision}, scale {scale} must round to an absolute value less than {}.",
485 if maxdigits == 0 {
486 "1".to_string()
487 } else {
488 format!("10^{maxdigits}")
489 }
490 )),
491 hint: None,
492 }
493}