1mod boolean;
2mod data_set;
3mod date_time;
4mod decimal;
5mod fixed_char;
6mod float;
7mod integer;
8mod interval_day_to_second;
9mod interval_year_to_month;
10mod ip_address;
11pub mod json;
12mod map;
13mod option;
14mod row;
15mod seq;
16mod string;
17mod util;
18pub mod uuid;
19mod var_binary;
20
21pub use self::uuid::*;
22pub use boolean::*;
23pub use data_set::*;
24pub use date_time::*;
25pub use decimal::*;
26pub use fixed_char::*;
27pub use float::*;
28pub use integer::*;
29pub use interval_day_to_second::*;
30pub use interval_year_to_month::*;
31pub use ip_address::*;
32pub use map::*;
33pub use option::*;
34pub use row::*;
35pub use seq::*;
36pub use string::*;
37pub use var_binary::*;
38
39use std::borrow::Cow;
43use std::collections::HashMap;
44use std::iter::FromIterator;
45use std::sync::Arc;
46
47use crate::{
48 ClientTypeSignatureParameter, Column, NamedTypeSignature, RawTrinoTy, RowFieldName,
49 TypeSignature,
50};
51use derive_more::Display;
52use iterable::*;
53use serde::de::DeserializeSeed;
54use serde::Serialize;
55
56#[derive(Display, Debug)]
58pub enum Error {
59 InvalidTrinoType,
60 InvalidColumn,
61 InvalidTypeSignature,
62 #[display("unsupported Trino type: {_0}")]
63 UnsupportedType(String),
64 ParseDecimalFailed(String),
65 ParseIntervalMonthFailed,
66 ParseIntervalDayFailed,
67 EmptyInTrinoRow,
68 NoneTrinoRow,
69}
70
71pub trait Trino {
72 type ValueType<'a>: Serialize
73 where
74 Self: 'a;
75 type Seed<'a, 'de>: DeserializeSeed<'de, Value = Self>;
76
77 fn value(&self) -> Self::ValueType<'_>;
78
79 fn ty() -> TrinoTy;
80
81 fn seed<'a, 'de>(ctx: &'a Context<'a>) -> Self::Seed<'a, 'de>;
83
84 fn empty() -> Self;
85}
86
87pub trait TrinoMapKey: Trino {}
88
89#[derive(Debug)]
90pub struct Context<'a> {
91 ty: &'a TrinoTy,
92 map: Arc<HashMap<usize, Vec<usize>>>,
93}
94
95impl<'a> Context<'a> {
96 pub fn new<T: Trino>(provided: &'a TrinoTy) -> Result<Self, Error> {
97 let target = T::ty();
98 let ret = extract(&target, provided)?;
99 let map = HashMap::from_iter(ret);
100 Ok(Context {
101 ty: provided,
102 map: Arc::new(map),
103 })
104 }
105
106 pub fn with_ty(&'a self, ty: &'a TrinoTy) -> Context<'a> {
107 Context {
108 ty,
109 map: self.map.clone(),
110 }
111 }
112
113 pub fn ty(&self) -> &TrinoTy {
114 self.ty
115 }
116
117 pub fn row_map(&self) -> Option<&[usize]> {
118 let key = self.ty as *const TrinoTy as usize;
119 self.map.get(&key).map(|r| &**r)
120 }
121}
122
123fn extract(target: &TrinoTy, provided: &TrinoTy) -> Result<Vec<(usize, Vec<usize>)>, Error> {
124 use TrinoTy::*;
125
126 match (target, provided) {
127 (Unknown, _) => Ok(vec![]),
128 (Decimal(p1, s1), Decimal(p2, s2)) if p1 == p2 && s1 == s2 => Ok(vec![]),
129 (Option(ty), provided) => extract(ty, provided),
130 (Boolean, Boolean) => Ok(vec![]),
131 (Date, Date) => Ok(vec![]),
132 (Time, Time) => Ok(vec![]),
133 (TimeWithTimeZone, TimeWithTimeZone) => Ok(vec![]),
134 (Timestamp, Timestamp) => Ok(vec![]),
135 (TimestampWithTimeZone, TimestampWithTimeZone) => Ok(vec![]),
136 (IntervalYearToMonth, IntervalYearToMonth) => Ok(vec![]),
137 (IntervalDayToSecond, IntervalDayToSecond) => Ok(vec![]),
138 (TrinoInt(_), TrinoInt(_)) => Ok(vec![]),
139 (TrinoFloat(_), TrinoFloat(_)) => Ok(vec![]),
140 (Varchar, Varchar) => Ok(vec![]),
141 (Char(a), Char(b)) if a == b => Ok(vec![]),
142 (Tuple(t1), Tuple(t2)) => {
143 if t1.len() != t2.len() {
144 Err(Error::InvalidTrinoType)
145 } else {
146 t1.lazy_zip(t2).try_flat_map(|(l, r)| extract(l, r))
147 }
148 }
149 (Row(t1), Row(t2)) => {
150 if t1.len() != t2.len() {
151 Err(Error::InvalidTrinoType)
152 } else {
153 let t1k = t1.sorted_by(|t1, t2| Ord::cmp(&t1.0, &t2.0));
155 let t2k = t2.sorted_by(|t1, t2| Ord::cmp(&t1.0, &t2.0));
156
157 let ret = t1k.lazy_zip(t2k).try_flat_map(|(l, r)| {
158 if l.0 == r.0 {
159 extract(&l.1, &r.1)
160 } else {
161 Err(Error::InvalidTrinoType)
162 }
163 })?;
164
165 let map = t2.map(|provided| t1.position(|target| provided.0 == target.0).unwrap());
166 let key = provided as *const TrinoTy as usize;
167 Ok(ret.add_one((key, map)))
168 }
169 }
170 (Array(t1), Array(t2)) => extract(t1, t2),
171 (Map(t1k, t1v), Map(t2k, t2v)) => Ok(extract(t1k, t2k)?.chain(extract(t1v, t2v)?)),
172 (IpAddress, IpAddress) => Ok(vec![]),
173 (Uuid, Uuid) => Ok(vec![]),
174 (Json, Json) => Ok(vec![]),
175 (VarBinary, VarBinary) => Ok(vec![]),
176 _ => Err(Error::InvalidTrinoType),
177 }
178}
179
180#[derive(Clone, Debug, Eq, PartialEq)]
186pub enum TrinoTy {
187 Date,
188 Time,
189 TimeWithTimeZone,
190 Timestamp,
191 TimestampWithTimeZone,
192 Uuid,
193 IntervalYearToMonth,
194 IntervalDayToSecond,
195 Option(Box<TrinoTy>),
196 Boolean,
197 TrinoInt(TrinoInt),
198 TrinoFloat(TrinoFloat),
199 Varchar,
200 Char(usize),
201 Tuple(Vec<TrinoTy>),
202 Row(Vec<(String, TrinoTy)>),
203 Array(Box<TrinoTy>),
204 Map(Box<TrinoTy>, Box<TrinoTy>),
205 Decimal(usize, usize),
206 IpAddress,
207 Json,
208 VarBinary,
209 Unknown,
210}
211
212#[derive(Clone, Debug, Eq, PartialEq)]
213pub enum TrinoInt {
214 I8,
215 I16,
216 I32,
217 I64,
218}
219
220#[derive(Clone, Debug, Eq, PartialEq)]
221pub enum TrinoFloat {
222 F32,
223 F64,
224}
225
226impl TrinoTy {
227 pub fn from_type_signature(mut sig: TypeSignature) -> Result<Self, Error> {
228 use TrinoFloat::*;
229 use TrinoInt::*;
230
231 let ty = match sig.raw_type {
232 RawTrinoTy::Date => TrinoTy::Date,
233 RawTrinoTy::Time => TrinoTy::Time,
234 RawTrinoTy::TimeWithTimeZone => TrinoTy::TimeWithTimeZone,
235 RawTrinoTy::Timestamp => TrinoTy::Timestamp,
236 RawTrinoTy::TimestampWithTimeZone => TrinoTy::TimestampWithTimeZone,
237 RawTrinoTy::IntervalYearToMonth => TrinoTy::IntervalYearToMonth,
238 RawTrinoTy::IntervalDayToSecond => TrinoTy::IntervalDayToSecond,
239 RawTrinoTy::Unknown => TrinoTy::Unknown,
240 RawTrinoTy::Decimal if sig.arguments.len() == 2 => {
241 let s_sig = sig.arguments.pop().unwrap();
242 let p_sig = sig.arguments.pop().unwrap();
243 if let (
244 ClientTypeSignatureParameter::LongLiteral(p),
245 ClientTypeSignatureParameter::LongLiteral(s),
246 ) = (p_sig, s_sig)
247 {
248 TrinoTy::Decimal(p as usize, s as usize)
249 } else {
250 return Err(Error::InvalidTypeSignature);
251 }
252 }
253 RawTrinoTy::Boolean => TrinoTy::Boolean,
254 RawTrinoTy::TinyInt => TrinoTy::TrinoInt(I8),
255 RawTrinoTy::SmallInt => TrinoTy::TrinoInt(I16),
256 RawTrinoTy::Integer => TrinoTy::TrinoInt(I32),
257 RawTrinoTy::BigInt => TrinoTy::TrinoInt(I64),
258 RawTrinoTy::Real => TrinoTy::TrinoFloat(F32),
259 RawTrinoTy::Double => TrinoTy::TrinoFloat(F64),
260 RawTrinoTy::VarChar => TrinoTy::Varchar,
261 RawTrinoTy::Char if sig.arguments.len() == 1 => {
262 if let ClientTypeSignatureParameter::LongLiteral(p) = sig.arguments.pop().unwrap() {
263 TrinoTy::Char(p as usize)
264 } else {
265 return Err(Error::InvalidTypeSignature);
266 }
267 }
268 RawTrinoTy::Array if sig.arguments.len() == 1 => {
269 let sig = sig.arguments.pop().unwrap();
270 if let ClientTypeSignatureParameter::TypeSignature(sig) = sig {
271 let inner = Self::from_type_signature(sig)?;
272 TrinoTy::Array(Box::new(inner))
273 } else {
274 return Err(Error::InvalidTypeSignature);
275 }
276 }
277 RawTrinoTy::Map if sig.arguments.len() == 2 => {
278 let v_sig = sig.arguments.pop().unwrap();
279 let k_sig = sig.arguments.pop().unwrap();
280 if let (
281 ClientTypeSignatureParameter::TypeSignature(k_sig),
282 ClientTypeSignatureParameter::TypeSignature(v_sig),
283 ) = (k_sig, v_sig)
284 {
285 let k_inner = Self::from_type_signature(k_sig)?;
286 let v_inner = Self::from_type_signature(v_sig)?;
287 TrinoTy::Map(Box::new(k_inner), Box::new(v_inner))
288 } else {
289 return Err(Error::InvalidTypeSignature);
290 }
291 }
292 RawTrinoTy::Row if !sig.arguments.is_empty() => {
293 let ir = sig.arguments.try_map(|arg| match arg {
294 ClientTypeSignatureParameter::NamedTypeSignature(sig) => {
295 let name = sig.field_name.map(|n| n.name);
296 let ty = Self::from_type_signature(sig.type_signature)?;
297 Ok((name, ty))
298 }
299 _ => Err(Error::InvalidTypeSignature),
300 })?;
301
302 let is_named = ir[0].0.is_some();
303
304 if is_named {
305 let row = ir.try_map(|(name, ty)| match name {
306 Some(n) => Ok((n, ty)),
307 None => Err(Error::InvalidTypeSignature),
308 })?;
309 TrinoTy::Row(row)
310 } else {
311 let tuple = ir.try_map(|(name, ty)| match name {
312 Some(_) => Err(Error::InvalidTypeSignature),
313 None => Ok(ty),
314 })?;
315 TrinoTy::Tuple(tuple)
316 }
317 }
318 RawTrinoTy::IpAddress => TrinoTy::IpAddress,
319 RawTrinoTy::Uuid => TrinoTy::Uuid,
320 RawTrinoTy::Json => TrinoTy::Json,
321 RawTrinoTy::VarBinary => TrinoTy::VarBinary,
322 other => return Err(Error::UnsupportedType(other.to_str().to_string())),
323 };
324
325 Ok(ty)
326 }
327
328 pub fn from_column(column: Column) -> Result<(String, Self), Error> {
329 let name = column.name;
330 if let Some(sig) = column.type_signature {
331 let ty = Self::from_type_signature(sig)?;
332 Ok((name, ty))
333 } else {
334 Err(Error::InvalidColumn)
335 }
336 }
337
338 pub fn from_columns(columns: Vec<Column>) -> Result<Self, Error> {
339 let row = columns.try_map(Self::from_column)?;
340 Ok(TrinoTy::Row(row))
341 }
342
343 pub fn into_type_signature(self) -> TypeSignature {
344 use TrinoTy::*;
345
346 let raw_ty = self.raw_type();
347
348 let params = match self {
349 Unknown => vec![],
350 Decimal(p, s) => vec![
351 ClientTypeSignatureParameter::LongLiteral(p as u64),
352 ClientTypeSignatureParameter::LongLiteral(s as u64),
353 ],
354 Date => vec![],
355 Time => vec![],
356 TimeWithTimeZone => vec![],
357 Timestamp => vec![],
358 TimestampWithTimeZone => vec![],
359 IntervalYearToMonth => vec![],
360 IntervalDayToSecond => vec![],
361 Option(t) => return t.into_type_signature(),
362 Boolean => vec![],
363 TrinoInt(_) => vec![],
364 TrinoFloat(_) => vec![],
365 Varchar => vec![ClientTypeSignatureParameter::LongLiteral(2147483647)],
366 Char(a) => vec![ClientTypeSignatureParameter::LongLiteral(a as u64)],
367 Tuple(ts) => ts.map(|ty| {
368 ClientTypeSignatureParameter::NamedTypeSignature(NamedTypeSignature {
369 field_name: None,
370 type_signature: ty.into_type_signature(),
371 })
372 }),
373 Row(ts) => ts.map(|(name, ty)| {
374 ClientTypeSignatureParameter::NamedTypeSignature(NamedTypeSignature {
375 field_name: Some(RowFieldName::new(name)),
376 type_signature: ty.into_type_signature(),
377 })
378 }),
379 Array(t) => vec![ClientTypeSignatureParameter::TypeSignature(
380 t.into_type_signature(),
381 )],
382 Map(t1, t2) => vec![
383 ClientTypeSignatureParameter::TypeSignature(t1.into_type_signature()),
384 ClientTypeSignatureParameter::TypeSignature(t2.into_type_signature()),
385 ],
386 IpAddress => vec![],
387 Uuid => vec![],
388 Json => vec![],
389 VarBinary => vec![],
390 };
391
392 TypeSignature::new(raw_ty, params)
393 }
394
395 pub fn full_type(&self) -> Cow<'static, str> {
396 use TrinoTy::*;
397
398 match self {
399 Unknown => RawTrinoTy::Unknown.to_str().into(),
400 Decimal(p, s) => format!("{}({},{})", RawTrinoTy::Decimal.to_str(), p, s).into(),
401 Option(t) => t.full_type(),
402 Date => RawTrinoTy::Date.to_str().into(),
403 Time => RawTrinoTy::Time.to_str().into(),
404 TimeWithTimeZone => RawTrinoTy::TimeWithTimeZone.to_str().into(),
405 Timestamp => RawTrinoTy::Timestamp.to_str().into(),
406 TimestampWithTimeZone => RawTrinoTy::TimestampWithTimeZone.to_str().into(),
407 IntervalYearToMonth => RawTrinoTy::IntervalYearToMonth.to_str().into(),
408 IntervalDayToSecond => RawTrinoTy::IntervalDayToSecond.to_str().into(),
409 Boolean => RawTrinoTy::Boolean.to_str().into(),
410 TrinoInt(ty) => ty.raw_type().to_str().into(),
411 TrinoFloat(ty) => ty.raw_type().to_str().into(),
412 Varchar => RawTrinoTy::VarChar.to_str().into(),
413 Char(a) => format!("{}({})", RawTrinoTy::Char.to_str(), a).into(),
414 Tuple(ts) => format!(
415 "{}({})",
416 RawTrinoTy::Row.to_str(),
417 ts.lazy_map(|ty| ty.full_type()).join(",")
418 )
419 .into(),
420 Row(ts) => format!(
421 "{}({})",
422 RawTrinoTy::Row.to_str(),
423 ts.lazy_map(|(name, ty)| format!("{} {}", name, ty.full_type()))
424 .join(",")
425 )
426 .into(),
427 Array(t) => format!("{}({})", RawTrinoTy::Array.to_str(), t.full_type()).into(),
428 Map(t1, t2) => format!(
429 "{}({},{})",
430 RawTrinoTy::Map.to_str(),
431 t1.full_type(),
432 t2.full_type()
433 )
434 .into(),
435 IpAddress => RawTrinoTy::IpAddress.to_str().into(),
436 Uuid => RawTrinoTy::Uuid.to_str().into(),
437 Json => RawTrinoTy::Json.to_str().into(),
438 VarBinary => RawTrinoTy::VarBinary.to_str().into(),
439 }
440 }
441
442 pub fn raw_type(&self) -> RawTrinoTy {
443 use TrinoTy::*;
444
445 match self {
446 Unknown => RawTrinoTy::Unknown,
447 Date => RawTrinoTy::Date,
448 Time => RawTrinoTy::Time,
449 TimeWithTimeZone => RawTrinoTy::TimeWithTimeZone,
450 Timestamp => RawTrinoTy::Timestamp,
451 TimestampWithTimeZone => RawTrinoTy::TimestampWithTimeZone,
452 IntervalYearToMonth => RawTrinoTy::IntervalYearToMonth,
453 IntervalDayToSecond => RawTrinoTy::IntervalDayToSecond,
454 Decimal(_, _) => RawTrinoTy::Decimal,
455 Option(ty) => ty.raw_type(),
456 Boolean => RawTrinoTy::Boolean,
457 TrinoInt(ty) => ty.raw_type(),
458 TrinoFloat(ty) => ty.raw_type(),
459 Varchar => RawTrinoTy::VarChar,
460 Char(_) => RawTrinoTy::Char,
461 Tuple(_) => RawTrinoTy::Row,
462 Row(_) => RawTrinoTy::Row,
463 Array(_) => RawTrinoTy::Array,
464 Map(_, _) => RawTrinoTy::Map,
465 IpAddress => RawTrinoTy::IpAddress,
466 Uuid => RawTrinoTy::Uuid,
467 Json => RawTrinoTy::Json,
468 VarBinary => RawTrinoTy::VarBinary,
469 }
470 }
471}
472
473impl TrinoInt {
474 pub fn raw_type(&self) -> RawTrinoTy {
475 use TrinoInt::*;
476 match self {
477 I8 => RawTrinoTy::TinyInt,
478 I16 => RawTrinoTy::SmallInt,
479 I32 => RawTrinoTy::Integer,
480 I64 => RawTrinoTy::BigInt,
481 }
482 }
483}
484
485impl TrinoFloat {
486 pub fn raw_type(&self) -> RawTrinoTy {
487 use TrinoFloat::*;
488 match self {
489 F32 => RawTrinoTy::Real,
490 F64 => RawTrinoTy::Double,
491 }
492 }
493}