Skip to main content

trino_rust_client/types/
mod.rs

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
39//mod str;
40//pub use self::str::*;
41
42use 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//TODO: refine it
57#[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    /// caller must provide a valid context
82    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                // create a vector of the original element's reference
154                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// Not yet decodable into a Rust type (these produce `Error::UnsupportedType`):
181// * TimeWithTimeZone
182// * HyperLogLog / P4HyperLogLog
183// * QDigest
184// * Geometry (Trino 482 added support alongside with Iceberg v3 tables)
185#[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}