Skip to main content

mz_repr/adt/
range.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10use std::any::type_name;
11use std::cmp::Ordering;
12use std::error::Error;
13use std::fmt::{self, Debug, Display};
14use std::hash::{Hash, Hasher};
15
16use bitflags::bitflags;
17use chrono::{DateTime, NaiveDateTime, Utc};
18use dec::OrderedDecimal;
19use mz_proto::{RustType, TryFromProtoError};
20use postgres_protocol::types;
21#[cfg(any(test, feature = "proptest"))]
22use proptest_derive::Arbitrary;
23use serde::{Deserialize, Serialize};
24use tokio_postgres::types::{FromSql, Type as PgType};
25
26use crate::Datum;
27use crate::adt::date::Date;
28use crate::adt::numeric::Numeric;
29use crate::adt::timestamp::CheckedTimestamp;
30use crate::scalar::{DatumKind, SqlScalarType};
31
32include!(concat!(env!("OUT_DIR"), "/mz_repr.adt.range.rs"));
33
34bitflags! {
35    #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
36    pub(crate) struct InternalFlags: u8 {
37        const EMPTY = 1;
38        const LB_INCLUSIVE = 1 << 1;
39        const LB_INFINITE = 1 << 2;
40        const UB_INCLUSIVE = 1 << 3;
41        const UB_INFINITE = 1 << 4;
42    }
43}
44
45bitflags! {
46    #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
47    pub(crate) struct PgFlags: u8 {
48        const EMPTY = 0b0000_0001;
49        const LB_INCLUSIVE = 0b0000_0010;
50        const UB_INCLUSIVE = 0b0000_0100;
51        const LB_INFINITE = 0b0000_1000;
52        const UB_INFINITE = 0b0001_0000;
53    }
54}
55
56/// A range of values along the domain `D`.
57///
58/// `D` is generic to facilitate interoperating over multiple representation,
59/// e.g. `Datum` and `mz_pgrepr::Value`. Because of the latter, we have to
60/// "manually derive" traits over `Range`.
61///
62/// Also notable, is that `Datum`s themselves store ranges as
63/// `Range<DatumNested<'a>>`, which lets us avoid unnecessary boxing of the
64/// range's finite bounds, which are most often expressed as `Datum`.
65pub struct Range<D> {
66    /// None value represents empty range
67    pub inner: Option<RangeInner<D>>,
68}
69
70impl crate::scalar::SqlContainerType for Range<Datum<'_>> {
71    fn unwrap_element_type(container: &SqlScalarType) -> &SqlScalarType {
72        container.unwrap_range_element_type()
73    }
74    fn wrap_element_type(element: SqlScalarType) -> SqlScalarType {
75        SqlScalarType::Range {
76            element_type: Box::new(element),
77        }
78    }
79}
80
81impl<D: Display> Display for Range<D> {
82    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
83        match &self.inner {
84            None => f.write_str("empty"),
85            Some(i) => i.fmt(f),
86        }
87    }
88}
89
90impl<D: Debug> Debug for Range<D> {
91    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
92        f.debug_struct("Range").field("inner", &self.inner).finish()
93    }
94}
95
96impl<D: Clone> Clone for Range<D> {
97    fn clone(&self) -> Self {
98        Self {
99            inner: self.inner.clone(),
100        }
101    }
102}
103
104impl<D: Copy> Copy for Range<D> {}
105
106impl<D: PartialEq> PartialEq for Range<D> {
107    fn eq(&self, other: &Self) -> bool {
108        self.inner == other.inner
109    }
110}
111
112impl<D: Eq> Eq for Range<D> {}
113
114impl<D: Ord + PartialOrd> PartialOrd for Range<D> {
115    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
116        Some(self.cmp(other))
117    }
118}
119
120impl<D: Ord> Ord for Range<D> {
121    fn cmp(&self, other: &Self) -> std::cmp::Ordering {
122        self.inner.cmp(&other.inner)
123    }
124}
125
126impl<D: Hash> Hash for Range<D> {
127    fn hash<H: Hasher>(&self, hasher: &mut H) {
128        self.inner.hash(hasher)
129    }
130}
131
132/// Trait alias for traits required for generic range function implementations.
133pub trait RangeOps<'a>:
134    Debug + Ord + PartialOrd + Eq + PartialEq + TryFrom<Datum<'a>> + Into<Datum<'a>>
135where
136    <Self as TryFrom<Datum<'a>>>::Error: std::fmt::Debug,
137{
138    /// Increment `self` one step forward, if applicable. Return `None` if
139    /// overflows.
140    fn step(self) -> Option<Self> {
141        Some(self)
142    }
143
144    fn unwrap_datum(d: Datum<'a>) -> Self {
145        <Self>::try_from(d)
146            .unwrap_or_else(|_| panic!("cannot take {} to {}", d, type_name::<Self>()))
147    }
148
149    fn err_type_name() -> &'static str;
150}
151
152impl<'a> RangeOps<'a> for i32 {
153    fn step(self) -> Option<i32> {
154        self.checked_add(1)
155    }
156
157    fn err_type_name() -> &'static str {
158        "integer"
159    }
160}
161
162impl<'a> RangeOps<'a> for i64 {
163    fn step(self) -> Option<i64> {
164        self.checked_add(1)
165    }
166
167    fn err_type_name() -> &'static str {
168        "bigint"
169    }
170}
171
172impl<'a> RangeOps<'a> for Date {
173    fn step(self) -> Option<Date> {
174        self.checked_add(1).ok()
175    }
176
177    fn err_type_name() -> &'static str {
178        "date"
179    }
180}
181
182impl<'a> RangeOps<'a> for OrderedDecimal<Numeric> {
183    fn err_type_name() -> &'static str {
184        "numeric"
185    }
186}
187
188impl<'a> RangeOps<'a> for CheckedTimestamp<NaiveDateTime> {
189    fn err_type_name() -> &'static str {
190        "timestamp"
191    }
192}
193
194impl<'a> RangeOps<'a> for CheckedTimestamp<DateTime<Utc>> {
195    fn err_type_name() -> &'static str {
196        "timestamptz"
197    }
198}
199
200// Totally generic range implementations.
201impl<D> Range<D> {
202    /// Create a new range.
203    ///
204    /// Note that when constructing `Range<Datum<'a>>`, the range must still be
205    /// canonicalized. If this becomes a common operation, we should consider
206    /// addinga `new_canonical` function that performs both steps.
207    pub fn new(inner: Option<(RangeLowerBound<D>, RangeUpperBound<D>)>) -> Range<D> {
208        Range {
209            inner: inner.map(|(lower, upper)| RangeInner { lower, upper }),
210        }
211    }
212
213    /// Get the flag bits appropriate to use in our internal (i.e. row) encoding
214    /// of range values.
215    ///
216    /// Note that this differs from the flags appropriate to encode with
217    /// Postgres, which has `UB_INFINITE` and `LB_INCLUSIVE` in the alternate
218    /// position.
219    pub fn internal_flag_bits(&self) -> u8 {
220        let mut flags = InternalFlags::empty();
221
222        match &self.inner {
223            None => {
224                flags.set(InternalFlags::EMPTY, true);
225            }
226            Some(RangeInner { lower, upper }) => {
227                flags.set(InternalFlags::EMPTY, false);
228                flags.set(InternalFlags::LB_INFINITE, lower.bound.is_none());
229                flags.set(InternalFlags::UB_INFINITE, upper.bound.is_none());
230                flags.set(InternalFlags::LB_INCLUSIVE, lower.inclusive);
231                flags.set(InternalFlags::UB_INCLUSIVE, upper.inclusive);
232            }
233        }
234
235        flags.bits()
236    }
237
238    /// Get the flag bits appropriate to use in PG-compatible encodings of range
239    /// values.
240    ///
241    /// Note that this differs from the flags appropriate for our internal
242    /// encoding, which has `UB_INFINITE` and `LB_INCLUSIVE` in the alternate
243    /// position.
244    pub fn pg_flag_bits(&self) -> u8 {
245        let mut flags = PgFlags::empty();
246
247        match &self.inner {
248            None => {
249                flags.set(PgFlags::EMPTY, true);
250            }
251            Some(RangeInner { lower, upper }) => {
252                flags.set(PgFlags::EMPTY, false);
253                flags.set(PgFlags::LB_INFINITE, lower.bound.is_none());
254                flags.set(PgFlags::UB_INFINITE, upper.bound.is_none());
255                flags.set(PgFlags::LB_INCLUSIVE, lower.inclusive);
256                flags.set(PgFlags::UB_INCLUSIVE, upper.inclusive);
257            }
258        }
259
260        flags.bits()
261    }
262
263    /// Converts `self` from having bounds of type `D` to type `O`, converting
264    /// the current bounds using `conv`.
265    pub fn into_bounds<F, O>(self, conv: F) -> Range<O>
266    where
267        F: Fn(D) -> O,
268    {
269        Range {
270            inner: self
271                .inner
272                .map(|RangeInner::<D> { lower, upper }| RangeInner::<O> {
273                    lower: RangeLowerBound {
274                        inclusive: lower.inclusive,
275                        bound: lower.bound.map(&conv),
276                    },
277                    upper: RangeUpperBound {
278                        inclusive: upper.inclusive,
279                        bound: upper.bound.map(&conv),
280                    },
281                }),
282        }
283    }
284
285    /// Like [`into_bounds`](Self::into_bounds), but the conversion may fail.
286    ///
287    /// Use this when converting each bound with a fallible function (e.g.
288    /// `into_datum`). Callers need not reach into `Range`'s internals
289    /// (`RangeInner`, `RangeLowerBound`, `RangeUpperBound`).
290    pub fn try_into_bounds<F, O, E>(self, conv: F) -> Result<Range<O>, E>
291    where
292        F: Fn(D) -> Result<O, E>,
293    {
294        let inner = match self.inner {
295            None => None,
296            Some(RangeInner { lower, upper }) => Some(RangeInner {
297                lower: RangeLowerBound {
298                    inclusive: lower.inclusive,
299                    bound: lower.bound.map(&conv).transpose()?,
300                },
301                upper: RangeUpperBound {
302                    inclusive: upper.inclusive,
303                    bound: upper.bound.map(&conv).transpose()?,
304                },
305            }),
306        };
307        Ok(Range { inner })
308    }
309}
310
311/// Range operations on `Range<Datum>` and `Range<DatumNested>`.
312impl<'a, B: Copy + Ord> Range<B> {
313    pub fn contains_elem<T: RangeOps<'a>>(&self, elem: &T) -> bool
314    where
315        Datum<'a>: From<B>,
316        <T as TryFrom<Datum<'a>>>::Error: std::fmt::Debug,
317    {
318        match self.inner {
319            None => false,
320            Some(inner) => inner.lower.satisfied_by(elem) && inner.upper.satisfied_by(elem),
321        }
322    }
323
324    pub fn contains_range(&self, other: &Range<B>) -> bool {
325        match (self.inner, other.inner) {
326            (None, None) | (Some(_), None) => true,
327            (None, Some(_)) => false,
328            (Some(i), Some(j)) => i.lower <= j.lower && j.upper <= i.upper,
329        }
330    }
331
332    pub fn overlaps(&self, other: &Range<B>) -> bool {
333        match (self.inner, other.inner) {
334            (Some(s), Some(o)) => {
335                let r = match s.cmp(&o) {
336                    Ordering::Equal => Ordering::Equal,
337                    Ordering::Less => s.upper.range_bound_cmp(&o.lower),
338                    Ordering::Greater => o.upper.range_bound_cmp(&s.lower),
339                };
340
341                // If smaller upper is >= larger lower, elements overlap.
342                matches!(r, Ordering::Greater | Ordering::Equal)
343            }
344            _ => false,
345        }
346    }
347
348    pub fn before(&self, other: &Range<B>) -> bool {
349        match (self.inner, other.inner) {
350            (Some(s), Some(o)) => {
351                matches!(s.upper.range_bound_cmp(&o.lower), Ordering::Less)
352            }
353            _ => false,
354        }
355    }
356
357    pub fn after(&self, other: &Range<B>) -> bool {
358        match (self.inner, other.inner) {
359            (Some(s), Some(o)) => {
360                matches!(s.lower.range_bound_cmp(&o.upper), Ordering::Greater)
361            }
362            _ => false,
363        }
364    }
365
366    pub fn overleft(&self, other: &Range<B>) -> bool {
367        match (self.inner, other.inner) {
368            (Some(s), Some(o)) => {
369                matches!(
370                    s.upper.range_bound_cmp(&o.upper),
371                    Ordering::Less | Ordering::Equal
372                )
373            }
374            _ => false,
375        }
376    }
377
378    pub fn overright(&self, other: &Range<B>) -> bool {
379        match (self.inner, other.inner) {
380            (Some(s), Some(o)) => {
381                matches!(
382                    s.lower.range_bound_cmp(&o.lower),
383                    Ordering::Greater | Ordering::Equal
384                )
385            }
386            _ => false,
387        }
388    }
389
390    pub fn adjacent(&self, other: &Range<B>) -> bool {
391        match (self.inner, other.inner) {
392            (Some(s), Some(o)) => {
393                // Look at each (lower, upper) pair.
394                for (lower, upper) in [(s.lower, o.upper), (o.lower, s.upper)] {
395                    if let (Some(l), Some(u)) = (lower.bound, upper.bound) {
396                        // If ..x](x.. or ..x)[x.., adjacent
397                        if lower.inclusive ^ upper.inclusive && l == u {
398                            return true;
399                        }
400                    }
401                }
402                false
403            }
404            _ => false,
405        }
406    }
407
408    pub fn union(&self, other: &Range<B>) -> Result<Range<B>, InvalidRangeError> {
409        // Handle self or other being empty
410        let (s, o) = match (self.inner, other.inner) {
411            (None, None) => return Ok(Range { inner: None }),
412            (inner @ Some(_), None) | (None, inner @ Some(_)) => return Ok(Range { inner }),
413            (Some(s), Some(o)) => {
414                // if not overlapping or adjacent, then result would not present continuity, so error.
415                if !(self.overlaps(other) || self.adjacent(other)) {
416                    return Err(InvalidRangeError::DiscontiguousUnion);
417                }
418                (s, o)
419            }
420        };
421
422        let lower = std::cmp::min(s.lower, o.lower);
423        let upper = std::cmp::max(s.upper, o.upper);
424
425        Ok(Range {
426            inner: Some(RangeInner { lower, upper }),
427        })
428    }
429
430    pub fn intersection(&self, other: &Range<B>) -> Range<B> {
431        // Handle self or other being empty
432        let (s, o) = match (self.inner, other.inner) {
433            (Some(s), Some(o)) => {
434                if !self.overlaps(other) {
435                    return Range { inner: None };
436                }
437
438                (s, o)
439            }
440            _ => return Range { inner: None },
441        };
442
443        let lower = std::cmp::max(s.lower, o.lower);
444        let upper = std::cmp::min(s.upper, o.upper);
445
446        Range {
447            inner: Some(RangeInner { lower, upper }),
448        }
449    }
450
451    // Function requires canonicalization so must be taken into `Range<Datum>`,
452    // which can be taken back into `Range<DatumNested>` by the caller if need
453    // be.
454    pub fn difference(&self, other: &Range<B>) -> Result<Range<Datum<'a>>, InvalidRangeError>
455    where
456        Datum<'a>: From<B>,
457    {
458        use std::cmp::Ordering::*;
459
460        // Difference op does nothing if no overlap.
461        if !self.overlaps(other) {
462            return Ok(self.into_bounds(Datum::from));
463        }
464
465        let (s, o) = match (self.inner, other.inner) {
466            (None, _) | (_, None) => unreachable!("already returned from overlap check"),
467            (Some(s), Some(o)) => (s, o),
468        };
469
470        let ll = s.lower.cmp(&o.lower);
471        let uu = s.upper.cmp(&o.upper);
472
473        let r = match (ll, uu) {
474            // `self` totally contains `other`
475            (Less, Greater) => return Err(InvalidRangeError::DiscontiguousDifference),
476            // `other` totally contains `self`
477            (Greater | Equal, Less | Equal) => Range { inner: None },
478            (Greater | Equal, Greater) => {
479                let lower = RangeBound {
480                    inclusive: !o.upper.inclusive,
481                    bound: o.upper.bound,
482                };
483                Range {
484                    inner: Some(RangeInner {
485                        lower,
486                        upper: s.upper,
487                    }),
488                }
489            }
490            (Less, Less | Equal) => {
491                let upper = RangeBound {
492                    inclusive: !o.lower.inclusive,
493                    bound: o.lower.bound,
494                };
495                Range {
496                    inner: Some(RangeInner {
497                        lower: s.lower,
498                        upper,
499                    }),
500                }
501            }
502        };
503
504        let mut r = r.into_bounds(Datum::from);
505
506        r.canonicalize()?;
507
508        Ok(r)
509    }
510}
511
512impl<'a> Range<Datum<'a>> {
513    /// Canonicalize the range by PG's heuristics, which are:
514    /// - Infinite bounds are always exclusive
515    /// - If type has step:
516    ///  - Exclusive lower bounds are rewritten as inclusive += step
517    ///  - Inclusive lower bounds are rewritten as exclusive += step
518    /// - Ranges are empty if lower >= upper after prev. step unless range type
519    ///   does not have step and both bounds are inclusive
520    ///
521    /// # Panics
522    /// - If the upper and lower bounds are finite and of different types.
523    pub fn canonicalize(&mut self) -> Result<(), InvalidRangeError> {
524        let (lower, upper) = match &mut self.inner {
525            Some(inner) => (&mut inner.lower, &mut inner.upper),
526            None => return Ok(()),
527        };
528
529        match (lower.bound, upper.bound) {
530            (Some(l), Some(u)) => {
531                assert_eq!(
532                    DatumKind::from(l),
533                    DatumKind::from(u),
534                    "finite bounds must be of same type"
535                );
536                if l > u {
537                    return Err(InvalidRangeError::MisorderedRangeBounds);
538                }
539            }
540            _ => {}
541        };
542
543        lower.canonicalize()?;
544        upper.canonicalize()?;
545
546        // The only way that you have two inclusive bounds with equal value are
547        // if type does not have step.
548        if !(lower.inclusive && upper.inclusive)
549            && lower.bound >= upper.bound
550            // None is less than any Some, so only need to check this condition.
551            && upper.bound.is_some()
552        {
553            // emtpy range
554            self.inner = None
555        }
556
557        Ok(())
558    }
559}
560
561/// Holds the upper and lower bounds for non-empty ranges.
562#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
563pub struct RangeInner<B> {
564    pub lower: RangeLowerBound<B>,
565    pub upper: RangeUpperBound<B>,
566}
567
568impl<B: Display> Display for RangeInner<B> {
569    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
570        f.write_str(if self.lower.inclusive { "[" } else { "(" })?;
571        self.lower.fmt(f)?;
572        f.write_str(",")?;
573        Display::fmt(&self.upper, f)?;
574        f.write_str(if self.upper.inclusive { "]" } else { ")" })
575    }
576}
577
578impl<B: Ord> Ord for RangeInner<B> {
579    fn cmp(&self, other: &Self) -> std::cmp::Ordering {
580        self.lower
581            .cmp(&other.lower)
582            .then(self.upper.cmp(&other.upper))
583    }
584}
585
586impl<B: PartialOrd + Ord> PartialOrd for RangeInner<B> {
587    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
588        Some(self.cmp(other))
589    }
590}
591
592/// Represents a terminal point of a range.
593#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
594pub struct RangeBound<B, const UPPER: bool = false> {
595    pub inclusive: bool,
596    /// None value represents an infinite bound.
597    pub bound: Option<B>,
598}
599
600impl<const UPPER: bool, D: Display> Display for RangeBound<D, UPPER> {
601    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
602        match &self.bound {
603            None => Ok(()),
604            Some(bound) => bound.fmt(f),
605        }
606    }
607}
608
609impl<const UPPER: bool, D: Ord> Ord for RangeBound<D, UPPER> {
610    fn cmp(&self, other: &Self) -> Ordering {
611        // 1. Sort by bounds
612        let mut cmp = self.bound.cmp(&other.bound);
613        // 2. Infinite bounds vs. finite bounds are reversed for uppers.
614        if UPPER && other.bound.is_none() ^ self.bound.is_none() {
615            cmp = cmp.reverse();
616        }
617        // 3. Tie break by sorting by inclusivity, which is inverted between
618        //    lowers and uppers.
619        cmp.then(if self.inclusive == other.inclusive {
620            Ordering::Equal
621        } else if self.inclusive {
622            if UPPER {
623                Ordering::Greater
624            } else {
625                Ordering::Less
626            }
627        } else if UPPER {
628            Ordering::Less
629        } else {
630            Ordering::Greater
631        })
632    }
633}
634
635impl<const UPPER: bool, D: PartialOrd + Ord> PartialOrd for RangeBound<D, UPPER> {
636    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
637        Some(self.cmp(other))
638    }
639}
640
641/// A `RangeBound` that sorts correctly for use as a lower bound.
642pub type RangeLowerBound<B> = RangeBound<B, false>;
643
644/// A `RangeBound` that sorts correctly for use as an upper bound.
645pub type RangeUpperBound<B> = RangeBound<B, true>;
646
647// Generic RangeBound implementations meant to work over `RangeBound<Datum,..>`
648// and `RangeBound<DatumNested,..>`.
649impl<'a, const UPPER: bool, B: Copy + Ord> RangeBound<B, UPPER> {
650    /// Determines where `elem` lies in relation to the range bound.
651    ///
652    /// # Panics
653    /// - If `self.bound.datum()` is not convertible to `T`.
654    fn elem_cmp<T: RangeOps<'a>>(&self, elem: &T) -> Ordering
655    where
656        Datum<'a>: From<B>,
657        <T as TryFrom<Datum<'a>>>::Error: std::fmt::Debug,
658    {
659        match self.bound.map(|bound| <T>::unwrap_datum(bound.into())) {
660            None if UPPER => Ordering::Greater,
661            None => Ordering::Less,
662            Some(bound) => bound.cmp(elem),
663        }
664    }
665
666    /// Does `elem` satisfy this bound?
667    fn satisfied_by<T: RangeOps<'a>>(&self, elem: &T) -> bool
668    where
669        Datum<'a>: From<B>,
670        <T as TryFrom<Datum<'a>>>::Error: std::fmt::Debug,
671    {
672        match self.elem_cmp(elem) {
673            // Inclusive always satisfied with equality, regardless of upper or
674            // lower.
675            Ordering::Equal => self.inclusive,
676            // Upper satisfied with values less than itself
677            Ordering::Greater => UPPER,
678            // Lower satisfied with values greater than itself
679            Ordering::Less => !UPPER,
680        }
681    }
682
683    // Compares two `RangeBound`, which do not need to both be of the same
684    // `UPPER`.
685    fn range_bound_cmp<const OTHER_UPPER: bool>(
686        &self,
687        other: &RangeBound<B, OTHER_UPPER>,
688    ) -> Ordering {
689        if UPPER == OTHER_UPPER {
690            return self.cmp(&RangeBound {
691                inclusive: other.inclusive,
692                bound: other.bound,
693            });
694        }
695
696        // Handle cases where either are infinite bounds, which have special
697        // semantics.
698        if self.bound.is_none() || other.bound.is_none() {
699            return if UPPER {
700                Ordering::Greater
701            } else {
702                Ordering::Less
703            };
704        }
705        // 1. Sort by bounds
706        let cmp = self.bound.cmp(&other.bound);
707        // 2. Tie break by sorting by inclusivity, which is inverted between
708        //    lowers and uppers.
709        cmp.then(if self.inclusive && other.inclusive {
710            Ordering::Equal
711        } else if UPPER {
712            Ordering::Less
713        } else {
714            Ordering::Greater
715        })
716    }
717}
718
719impl<'a, const UPPER: bool> RangeBound<Datum<'a>, UPPER> {
720    /// Create a new `RangeBound` whose value is "infinite" (i.e. None) if `d ==
721    /// Datum::Null`, otherwise finite (i.e. Some).
722    ///
723    /// There is not a corresponding generic implementation of this because
724    /// genericizing how to express infinite bounds is less clear.
725    pub fn new(d: Datum<'a>, inclusive: bool) -> RangeBound<Datum<'a>, UPPER> {
726        RangeBound {
727            inclusive,
728            bound: match d {
729                Datum::Null => None,
730                o => Some(o),
731            },
732        }
733    }
734
735    /// Rewrite the bounds to the consistent format. This is absolutely
736    /// necessary to perform the correct equality/comparison operations on
737    /// types.
738    fn canonicalize(&mut self) -> Result<(), InvalidRangeError> {
739        Ok(match self.bound {
740            None => {
741                self.inclusive = false;
742            }
743            // Valid range types are defined in typeconv.rs:validate_range_element_type
744            Some(value) => match value {
745                d @ Datum::Int32(_) => self.canonicalize_inner::<i32>(d)?,
746                d @ Datum::Int64(_) => self.canonicalize_inner::<i64>(d)?,
747                d @ Datum::Date(_) => self.canonicalize_inner::<Date>(d)?,
748                Datum::Numeric(..) | Datum::Timestamp(..) | Datum::TimestampTz(..) => {}
749                d => unreachable!("{d:?} not yet supported in ranges"),
750            },
751        })
752    }
753
754    /// Canonicalize `self`'s representation for types that have discrete steps
755    /// between values.
756    ///
757    /// Continuous values (e.g. timestamps, numeric) must not be
758    /// canonicalized.
759    fn canonicalize_inner<T: RangeOps<'a>>(&mut self, d: Datum<'a>) -> Result<(), InvalidRangeError>
760    where
761        <T as TryFrom<Datum<'a>>>::Error: std::fmt::Debug,
762    {
763        // Upper bounds must be exclusive, lower bounds inclusive
764        if UPPER == self.inclusive {
765            let cur = <T>::unwrap_datum(d);
766            self.bound = Some(
767                cur.step()
768                    .ok_or_else(|| {
769                        InvalidRangeError::CanonicalizationOverflow(T::err_type_name().into())
770                    })?
771                    .into(),
772            );
773            self.inclusive = !UPPER;
774        }
775
776        Ok(())
777    }
778}
779
780#[derive(
781    Ord,
782    PartialOrd,
783    Clone,
784    Debug,
785    Eq,
786    PartialEq,
787    Serialize,
788    Deserialize,
789    Hash
790)]
791#[cfg_attr(any(test, feature = "proptest"), derive(Arbitrary))]
792pub enum InvalidRangeError {
793    MisorderedRangeBounds,
794    CanonicalizationOverflow(Box<str>),
795    InvalidRangeBoundFlags,
796    DiscontiguousUnion,
797    DiscontiguousDifference,
798    NullRangeBoundFlags,
799    /// The encoded range data is structurally invalid (e.g. a null bound,
800    /// bounds of inconsistent types, or the wrong number of bounds). Only
801    /// reachable by decoding untrusted/corrupted bytes.
802    InvalidRangeData,
803}
804
805impl Display for InvalidRangeError {
806    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
807        match self {
808            InvalidRangeError::MisorderedRangeBounds => {
809                f.write_str("range lower bound must be less than or equal to range upper bound")
810            }
811            InvalidRangeError::CanonicalizationOverflow(t) => {
812                write!(f, "{} out of range", t)
813            }
814            InvalidRangeError::InvalidRangeBoundFlags => f.write_str("invalid range bound flags"),
815            InvalidRangeError::DiscontiguousUnion => {
816                f.write_str("result of range union would not be contiguous")
817            }
818            InvalidRangeError::DiscontiguousDifference => {
819                f.write_str("result of range difference would not be contiguous")
820            }
821            InvalidRangeError::NullRangeBoundFlags => {
822                f.write_str("range constructor flags argument must not be null")
823            }
824            InvalidRangeError::InvalidRangeData => f.write_str("invalid range data"),
825        }
826    }
827}
828
829impl Error for InvalidRangeError {
830    fn source(&self) -> Option<&(dyn Error + 'static)> {
831        None
832    }
833}
834
835// Required due to Proto decoding using string as its error type
836impl From<InvalidRangeError> for String {
837    fn from(e: InvalidRangeError) -> Self {
838        e.to_string()
839    }
840}
841
842impl RustType<ProtoInvalidRangeError> for InvalidRangeError {
843    fn into_proto(&self) -> ProtoInvalidRangeError {
844        use Kind::*;
845        use proto_invalid_range_error::*;
846        let kind = match self {
847            InvalidRangeError::MisorderedRangeBounds => MisorderedRangeBounds(()),
848            InvalidRangeError::CanonicalizationOverflow(s) => {
849                CanonicalizationOverflow(s.into_proto())
850            }
851            InvalidRangeError::InvalidRangeBoundFlags => InvalidRangeBoundFlags(()),
852            InvalidRangeError::DiscontiguousUnion => DiscontiguousUnion(()),
853            InvalidRangeError::DiscontiguousDifference => DiscontiguousDifference(()),
854            InvalidRangeError::NullRangeBoundFlags => NullRangeBoundFlags(()),
855            InvalidRangeError::InvalidRangeData => InvalidRangeData(()),
856        };
857        ProtoInvalidRangeError { kind: Some(kind) }
858    }
859
860    fn from_proto(proto: ProtoInvalidRangeError) -> Result<Self, TryFromProtoError> {
861        use proto_invalid_range_error::Kind::*;
862        match proto.kind {
863            Some(kind) => Ok(match kind {
864                MisorderedRangeBounds(()) => InvalidRangeError::MisorderedRangeBounds,
865                CanonicalizationOverflow(s) => {
866                    InvalidRangeError::CanonicalizationOverflow(s.into())
867                }
868                InvalidRangeBoundFlags(()) => InvalidRangeError::InvalidRangeBoundFlags,
869                DiscontiguousUnion(()) => InvalidRangeError::DiscontiguousUnion,
870                DiscontiguousDifference(()) => InvalidRangeError::DiscontiguousDifference,
871                NullRangeBoundFlags(()) => InvalidRangeError::NullRangeBoundFlags,
872                InvalidRangeData(()) => InvalidRangeError::InvalidRangeData,
873            }),
874            None => Err(TryFromProtoError::missing_field(
875                "`ProtoInvalidRangeError::kind`",
876            )),
877        }
878    }
879}
880
881pub fn parse_range_bound_flags<'a>(flags: &'a str) -> Result<(bool, bool), InvalidRangeError> {
882    let mut flags = flags.chars();
883
884    let lower = match flags.next() {
885        Some('(') => false,
886        Some('[') => true,
887        _ => return Err(InvalidRangeError::InvalidRangeBoundFlags),
888    };
889
890    let upper = match flags.next() {
891        Some(')') => false,
892        Some(']') => true,
893        _ => return Err(InvalidRangeError::InvalidRangeBoundFlags),
894    };
895
896    match flags.next() {
897        Some(_) => Err(InvalidRangeError::InvalidRangeBoundFlags),
898        None => Ok((lower, upper)),
899    }
900}
901
902impl<'a, T: FromSql<'a>> FromSql<'a> for Range<T> {
903    fn from_sql(ty: &PgType, raw: &'a [u8]) -> Result<Range<T>, Box<dyn Error + Sync + Send>> {
904        let inner_typ = match ty {
905            &PgType::INT4_RANGE => PgType::INT4,
906            &PgType::INT8_RANGE => PgType::INT8,
907            &PgType::DATE_RANGE => PgType::DATE,
908            &PgType::NUM_RANGE => PgType::NUMERIC,
909            &PgType::TS_RANGE => PgType::TIMESTAMP,
910            &PgType::TSTZ_RANGE => PgType::TIMESTAMPTZ,
911            _ => unreachable!(),
912        };
913
914        let inner = match types::range_from_sql(raw)? {
915            types::Range::Empty => None,
916            types::Range::Nonempty(lower, upper) => {
917                let mut bounds = Vec::with_capacity(2);
918
919                for bound_outer in [lower, upper].into_iter() {
920                    let bound = match bound_outer {
921                        types::RangeBound::Exclusive(bound)
922                        | types::RangeBound::Inclusive(bound) => bound
923                            .map(|bound| T::from_sql(&inner_typ, bound))
924                            .transpose()?,
925                        types::RangeBound::Unbounded => None,
926                    };
927                    let inclusive = matches!(bound_outer, types::RangeBound::Inclusive(_));
928                    bounds.push(RangeBound { bound, inclusive });
929                }
930
931                let lower = bounds.remove(0);
932                let upper = bounds.remove(0);
933                assert!(bounds.is_empty());
934
935                Some(RangeInner {
936                    lower,
937                    // Rewrite bound in terms of appropriate `UPPER`
938                    upper: RangeBound {
939                        bound: upper.bound,
940                        inclusive: upper.inclusive,
941                    },
942                })
943            }
944        };
945
946        Ok(Range { inner })
947    }
948
949    fn accepts(ty: &PgType) -> bool {
950        matches!(
951            ty,
952            &PgType::INT4_RANGE
953                | &PgType::INT8_RANGE
954                | &PgType::DATE_RANGE
955                | &PgType::NUM_RANGE
956                | &PgType::TS_RANGE
957                | &PgType::TSTZ_RANGE
958        )
959    }
960}