Skip to main content

typeql/
value.rs

1/*
2 * This Source Code Form is subject to the terms of the Mozilla Public
3 * License, v. 2.0. If a copy of the MPL was not distributed with this
4 * file, You can obtain one at https://mozilla.org/MPL/2.0/.
5 */
6
7use std::fmt::{self, Formatter};
8
9use crate::{
10    Result,
11    common::{Span, Spanned, error::TypeQLError},
12    pretty::Pretty,
13};
14
15#[derive(Debug, Clone, Eq, PartialEq)]
16pub struct BooleanLiteral {
17    pub value: String,
18}
19
20#[derive(Debug, Clone, Eq, PartialEq)]
21pub struct StringLiteral {
22    pub value: String,
23}
24
25#[derive(Debug, Clone, Eq, PartialEq)]
26pub struct IntegerLiteral {
27    pub value: String,
28}
29
30#[derive(Debug, Clone, Eq, PartialEq)]
31pub struct NumericLiteral {
32    pub value: String,
33}
34
35#[derive(Debug, Clone, Copy, Eq, PartialEq)]
36pub enum Sign {
37    Plus,
38    Minus,
39}
40
41#[derive(Debug, Clone, Eq, PartialEq)]
42pub struct SignedIntegerLiteral {
43    pub sign: Option<Sign>,
44    pub integral: String,
45}
46
47#[derive(Debug, Clone, Eq, PartialEq)]
48pub struct SignedDoubleLiteral {
49    pub sign: Option<Sign>,
50    pub double: String,
51}
52
53#[derive(Debug, Clone, Eq, PartialEq)]
54pub struct SignedDecimalLiteral {
55    pub sign: Option<Sign>,
56    pub decimal: String,
57}
58
59#[derive(Debug, Clone, Eq, PartialEq)]
60pub struct DateFragment {
61    pub year: String,
62    pub month: String,
63    pub day: String,
64}
65
66#[derive(Debug, Clone, Eq, PartialEq)]
67pub struct TimeFragment {
68    pub hour: String,
69    pub minute: String,
70    pub second: Option<String>,
71    pub second_fraction: Option<String>,
72}
73
74#[derive(Debug, Clone, Eq, PartialEq)]
75pub struct DateTimeTZLiteral {
76    pub date: DateFragment,
77    pub time: TimeFragment,
78    pub timezone: TimeZone,
79}
80
81#[derive(Debug, Clone, Eq, PartialEq)]
82pub struct DateTimeLiteral {
83    pub date: DateFragment,
84    pub time: TimeFragment,
85}
86
87#[derive(Debug, Clone, Eq, PartialEq)]
88pub struct DateLiteral {
89    pub date: DateFragment,
90}
91
92#[derive(Debug, Clone, Eq, PartialEq)]
93pub enum TimeZone {
94    IANA(String),
95    ISO(String),
96}
97
98#[derive(Debug, Clone, Eq, PartialEq)]
99pub enum DurationLiteral {
100    Weeks(IntegerLiteral),
101    DateAndTime(DurationDate, Option<DurationTime>),
102    Time(DurationTime),
103}
104
105#[derive(Debug, Clone, Eq, PartialEq)]
106pub struct StructLiteral {
107    pub inner: String, // TODO
108}
109
110#[derive(Debug, Clone, Eq, PartialEq)]
111pub struct DurationDate {
112    pub years: Option<IntegerLiteral>,
113    pub months: Option<IntegerLiteral>,
114    pub days: Option<IntegerLiteral>,
115}
116
117#[derive(Debug, Clone, Eq, PartialEq)]
118pub struct DurationTime {
119    pub hours: Option<IntegerLiteral>,
120    pub minutes: Option<IntegerLiteral>,
121    pub seconds: Option<NumericLiteral>,
122}
123
124#[derive(Debug, Clone, Eq, PartialEq)]
125pub enum ValueLiteral {
126    Boolean(BooleanLiteral),
127    Integer(SignedIntegerLiteral),
128    Decimal(SignedDecimalLiteral),
129    Double(SignedDoubleLiteral),
130    Date(DateLiteral),
131    DateTime(DateTimeLiteral),
132    DateTimeTz(DateTimeTZLiteral),
133    Duration(DurationLiteral),
134    String(StringLiteral),
135    Struct(StructLiteral),
136}
137
138#[derive(Debug, Clone, Eq, PartialEq)]
139pub struct Literal {
140    pub span: Option<Span>,
141    pub inner: ValueLiteral,
142}
143
144impl Literal {
145    pub(crate) fn new(span: Option<Span>, inner: ValueLiteral) -> Self {
146        Self { span, inner }
147    }
148}
149
150impl Spanned for Literal {
151    fn span(&self) -> Option<Span> {
152        self.span
153    }
154}
155
156impl Pretty for Literal {}
157
158impl fmt::Display for Literal {
159    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
160        fmt::Display::fmt(&self.inner, f)
161    }
162}
163
164impl fmt::Display for ValueLiteral {
165    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
166        match self {
167            ValueLiteral::Boolean(value) => fmt::Display::fmt(value, f),
168            ValueLiteral::Integer(value) => fmt::Display::fmt(value, f),
169            ValueLiteral::Decimal(value) => fmt::Display::fmt(value, f),
170            ValueLiteral::Double(value) => fmt::Display::fmt(value, f),
171            ValueLiteral::Date(value) => fmt::Display::fmt(value, f),
172            ValueLiteral::DateTime(value) => fmt::Display::fmt(value, f),
173            ValueLiteral::DateTimeTz(value) => fmt::Display::fmt(value, f),
174            ValueLiteral::Duration(value) => fmt::Display::fmt(value, f),
175            ValueLiteral::String(value) => fmt::Display::fmt(value, f),
176            ValueLiteral::Struct(value) => fmt::Display::fmt(value, f),
177        }
178    }
179}
180
181impl fmt::Display for IntegerLiteral {
182    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
183        f.write_str(self.value.as_str())
184    }
185}
186
187impl fmt::Display for NumericLiteral {
188    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
189        f.write_str(self.value.as_str())
190    }
191}
192
193impl fmt::Display for StringLiteral {
194    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
195        f.write_str(self.value.as_str())
196    }
197}
198
199impl fmt::Display for Sign {
200    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
201        match self {
202            Sign::Plus => f.write_str("+"),
203            Sign::Minus => f.write_str("-"),
204        }
205    }
206}
207
208impl fmt::Display for DateFragment {
209    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
210        write!(f, "{}-{}-{}", self.year, self.month, self.day)
211    }
212}
213
214impl fmt::Display for TimeFragment {
215    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
216        let (hour, minute) = (self.hour.as_str(), self.minute.as_str());
217        match &self.second {
218            None => write!(f, "T{hour}:{minute}"),
219            Some(second) => match &self.second_fraction {
220                None => write!(f, "T{hour}:{minute}:{second}"),
221                Some(second_fraction) => write!(f, "T{hour}:{minute}:{second}.{second_fraction}"),
222            },
223        }
224    }
225}
226
227impl fmt::Display for TimeZone {
228    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
229        match self {
230            TimeZone::IANA(value) => f.write_str(value),
231            TimeZone::ISO(value) => f.write_str(value),
232        }
233    }
234}
235
236impl fmt::Display for DurationDate {
237    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
238        if let Some(years) = &self.years {
239            write!(f, "{years}Y")?;
240        }
241        if let Some(months) = &self.months {
242            write!(f, "{months}M")?;
243        }
244        if let Some(days) = &self.days {
245            write!(f, "{days}D")?;
246        }
247        Ok(())
248    }
249}
250
251impl fmt::Display for DurationTime {
252    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
253        if let Some(hours) = &self.hours {
254            write!(f, "{hours}H")?;
255        }
256        if let Some(minutes) = &self.minutes {
257            write!(f, "{minutes}M")?;
258        }
259        if let Some(seconds) = &self.seconds {
260            write!(f, "{seconds}S")?;
261        }
262        Ok(())
263    }
264}
265
266impl fmt::Display for BooleanLiteral {
267    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
268        f.write_str(self.value.as_str())
269    }
270}
271
272impl fmt::Display for SignedIntegerLiteral {
273    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
274        if let Some(sign) = &self.sign {
275            fmt::Display::fmt(sign, f)?;
276        }
277        f.write_str(self.integral.as_str())
278    }
279}
280
281impl fmt::Display for SignedDecimalLiteral {
282    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
283        if let Some(sign) = &self.sign {
284            fmt::Display::fmt(sign, f)?;
285        }
286        write!(f, "{}dec", self.decimal.as_str())
287    }
288}
289
290impl fmt::Display for SignedDoubleLiteral {
291    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
292        if let Some(sign) = &self.sign {
293            fmt::Display::fmt(sign, f)?;
294        }
295        f.write_str(self.double.as_str())
296    }
297}
298
299impl fmt::Display for DateLiteral {
300    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
301        fmt::Display::fmt(&self.date, f)
302    }
303}
304
305impl fmt::Display for DateTimeLiteral {
306    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
307        write!(f, "{}{}", &self.date, &self.time)
308    }
309}
310
311impl fmt::Display for DateTimeTZLiteral {
312    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
313        fmt::Display::fmt(&self.date, f)?;
314        fmt::Display::fmt(&self.time, f)?;
315        fmt::Display::fmt(&self.timezone, f)?;
316        Ok(())
317    }
318}
319
320impl fmt::Display for DurationLiteral {
321    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
322        f.write_str("P")?;
323        match self {
324            DurationLiteral::Weeks(weeks) => write!(f, "{weeks}W")?,
325            DurationLiteral::DateAndTime(date, time) => {
326                fmt::Display::fmt(date, f)?;
327                match time {
328                    None => {}
329                    Some(time) => write!(f, "T{time}")?,
330                }
331            }
332            DurationLiteral::Time(time) => write!(f, "T{time}")?,
333        }
334        Ok(())
335    }
336}
337
338impl fmt::Display for StructLiteral {
339    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
340        f.write_str(self.inner.as_str())
341    }
342}
343
344impl StringLiteral {
345    pub fn unescape(&self) -> Result<String> {
346        self.process_unescape(|bytes| {
347            if bytes.len() < 2 {
348                return Err(1);
349            }
350            match bytes[1] {
351                BSP => Ok(('\x08', 2)),
352                TAB => Ok(('\x09', 2)),
353                LF_ => Ok(('\x0a', 2)),
354                FF_ => Ok(('\x0c', 2)),
355                CR_ => Ok(('\x0d', 2)),
356                c @ (b'"' | b'\'' | b'\\') => Ok((c as char, 2)),
357                b'u' => match decode_unicode_hex_escape(&bytes[2..]) {
358                    Ok((ch, consumed)) => Ok((ch, consumed + 2)),
359                    Err(consumed) => Err(consumed + 2),
360                },
361                _ => Err(2),
362            }
363        })
364    }
365
366    pub fn unescape_regex(&self) -> Result<String> {
367        self.process_unescape(|bytes| match bytes.get(1) {
368            Some(b'"') => Ok(('"', 2)),
369            _ => Ok(('\\', 1)),
370        })
371    }
372
373    fn process_unescape<F>(&self, escape_handler: F) -> Result<String>
374    where
375        F: Fn(&[u8]) -> std::result::Result<(char, usize), usize>,
376    {
377        let bytes = self.value.as_bytes();
378        assert_eq!(bytes[0], bytes[bytes.len() - 1]);
379        assert!(matches!(bytes[0], b'\'' | b'"'));
380
381        let escaped_string = &self.value[1..self.value.len() - 1];
382        let mut buf = Vec::with_capacity(escaped_string.len());
383        let mut rest = escaped_string.as_bytes();
384        while !rest.is_empty() {
385            if rest[0] == b'\\' {
386                match escape_handler(rest) {
387                    Ok((char, escaped_len)) => {
388                        let start = buf.len();
389                        buf.resize(buf.len() + char.len_utf8(), 0);
390                        char.encode_utf8(&mut buf[start..]);
391                        rest = &rest[escaped_len..];
392                    }
393                    Err(considered_byte_length) => {
394                        let offset = escaped_string.len() - rest.len();
395                        let mut end = std::cmp::min(offset + considered_byte_length, escaped_string.len());
396                        while !escaped_string.is_char_boundary(end) {
397                            end += 1;
398                        }
399                        return Err(TypeQLError::InvalidStringEscape {
400                            full_string: escaped_string.to_owned(),
401                            escape: escaped_string[offset..end].to_owned(),
402                        }
403                        .into());
404                    }
405                }
406            } else {
407                buf.push(rest[0]);
408                rest = &rest[1..];
409            }
410        }
411        Ok(String::from_utf8(buf).expect("Expected valid utf8").to_owned())
412    }
413}
414
415const BSP: u8 = b'b';
416const TAB: u8 = b't';
417const LF_: u8 = b'n';
418const FF_: u8 = b'f';
419const CR_: u8 = b'r';
420
421#[allow(arithmetic_overflow)]
422fn decode_unicode_hex_escape(bytes: &[u8]) -> std::result::Result<(char, usize), usize> {
423    if bytes.is_empty() {
424        Err(0)
425    } else if bytes[0] == b'{' {
426        let safe_len = std::cmp::min(bytes.len(), 8);
427        if let Some(i) = bytes[..safe_len].iter().position(|b| *b == b'}') {
428            unicode_char_from_hex(&bytes[1..i]).map(|c| (c, i + 1)).ok_or(i + 1)
429        } else {
430            Err(safe_len)
431        }
432    } else {
433        if bytes.len() >= 4 {
434            unicode_char_from_hex(&bytes[0..4]).map(|c| (c, 4)).ok_or(4)
435        } else {
436            Err(std::cmp::min(bytes.len(), 4))
437        }
438    }
439}
440
441fn unicode_char_from_hex(bytes: &[u8]) -> Option<char> {
442    if bytes.is_empty() || bytes.len() > 6 {
443        return None;
444    }
445    let mut as_u32 = 0u32;
446    // from_ascii_radix is still experimental
447    for b in bytes {
448        as_u32 = (as_u32 << 4) | (*b as char).to_digit(16)?;
449    }
450    char::from_u32(as_u32)
451}
452
453#[cfg(test)]
454pub mod tests {
455    use crate::{
456        ValueLiteral, parse_value,
457        value::{StringLiteral, TypeQLError},
458    };
459
460    fn parse_to_string_literal(escaped: &str) -> StringLiteral {
461        let ValueLiteral::String(parsed) = parse_value(escaped).unwrap() else {
462            panic!("Not parsed as string");
463        };
464        parsed
465    }
466
467    #[test]
468    fn test_unescape_regex() {
469        {
470            let escaped = r#""a\"b\"c""#;
471            let unescaped = parse_to_string_literal(escaped).unescape_regex().unwrap();
472            assert_eq!(unescaped.as_str(), r#"a"b"c"#);
473        }
474        {
475            let escaped = r#""abc\123""#;
476            let unescaped = parse_to_string_literal(escaped).unescape_regex().unwrap();
477            assert_eq!(unescaped.as_str(), r#"abc\123"#);
478        }
479        // Cases that fail at parsing
480        {
481            let escaped = r#""abc\""#;
482            assert!(crate::parse_value(escaped).is_err()); // Parsing fails as incomplete string literal
483            let string_literal = StringLiteral { value: escaped.to_owned() };
484            let unescaped = string_literal.unescape_regex().unwrap();
485            assert_eq!(unescaped.as_str(), r#"abc\"#);
486        }
487    }
488
489    macro_rules! assert_unescapes_to {
490        ($escaped: expr, $expected: expr) => {
491            let unescaped = parse_to_string_literal($escaped).unescape().unwrap();
492            assert_eq!(unescaped, $expected);
493        };
494    }
495
496    macro_rules! assert_unescape_errors {
497        ($escaped: expr, $expected_escape_sequence: expr) => {
498            let error = parse_to_string_literal($escaped).unescape().unwrap_err();
499            let TypeQLError::InvalidStringEscape { escape, .. } = &error.errors()[0] else {
500                panic!("Wrong error type. Was {error:?}")
501            };
502            assert_eq!(escape, $expected_escape_sequence);
503        };
504    }
505
506    #[test]
507    fn test_unescape() {
508        // Succeeds
509        assert_unescapes_to!(r#""a\tb\tc""#, "a\tb\tc"); // works
510        assert_unescapes_to!(r#""a\"b\"c""#, r#"a"b"c"#); // works
511        assert_unescapes_to!(r#""a\'b\'c""#, r#"a'b'c"#); // works
512        assert_unescapes_to!(r#""a\\b\\c""#, r#"a\b\c"#); // works
513        //  - Unicode
514        assert_unescapes_to!(r#""abc \u0ca0\u005f\u0ca0""#, "abc ಠ_ಠ"); // works
515        assert_unescapes_to!(r#""abc \u0CA0\u005F\u0CA0""#, "abc ಠ_ಠ"); // caps
516        assert_unescapes_to!(r#""abc \u0CA01234""#, "abc ಠ1234"); // consumes only 4
517        assert_unescapes_to!(r#""abc \u{0CA0}1234""#, "abc ಠ1234"); // braces with only 4
518        assert_unescapes_to!(r#""abc \u{130ED}\u{13153}1234""#, "abc 𓃭𓅓1234"); // braces with 6
519
520        // Errors
521        assert_unescape_errors!(r#""ab\c""#, r"\c"); // Invalid escape
522
523        //  - Unicode
524        assert_unescape_errors!(r#""abc \u""#, r"\u"); // Not enough bytes
525        assert_unescape_errors!(r#""abc \u012""#, r"\u012"); // Not enough bytes
526        assert_unescape_errors!(r#""abc \uwu/ abc""#, r"\uwu/ "); // Invalid hex
527        assert_unescape_errors!(r#""abc \uΣ12Σ abc""#, r"\uΣ12"); // Invalid hex, 3 chars more than 4 bytes
528        assert_unescape_errors!(r#""abc \u123Σ abc""#, r"\u123Σ"); // Invalid hex, 4 chars more than 4 bytes
529        assert_unescape_errors!(r#""abc \u{""#, r"\u{"); // Not enough bytes
530        assert_unescape_errors!(r#""abc \u{123Σ} abc""#, r"\u{123Σ}"); // Invalid hex with braces
531        assert_unescape_errors!(r#""abc \u{1234567} abc""#, r"\u{1234567"); // Too many characters, stop at 8
532        assert_unescape_errors!(r#""abc \u{213456} abc""#, r"\u{213456}"); // Above valid range
533
534        // Cases that fail at parsing
535        {
536            let escaped = r#""abc\""#;
537            assert!(crate::parse_value(escaped).is_err()); // Parsing fails as incomplete string literal
538            let string_literal = StringLiteral { value: escaped.to_owned() };
539            let error = string_literal.unescape().unwrap_err();
540            let TypeQLError::InvalidStringEscape { escape, .. } = &error.errors()[0] else {
541                panic!("Wrong error type. Was {error:?}")
542            };
543            assert_eq!(escape, r#"\"#);
544        }
545    }
546
547    #[ignore]
548    #[test]
549    fn time_unescape_ascii() {
550        let text = generate_string(TIME_UNESCAPE_TEXT_LEN, |x| 32 + (x % 94));
551        time_unescape(text);
552    }
553
554    #[ignore]
555    #[test]
556    fn time_unescape_unicode() {
557        // assert_eq!(None, (0..0x07ff).filter(|x| char::from_u32(*x).is_none()).next());
558        let text = generate_string(TIME_UNESCAPE_TEXT_LEN, move |x| x & 0x07ff);
559        time_unescape(text);
560    }
561
562    const TIME_UNESCAPE_TEXT_LEN: usize = 100000;
563    fn time_unescape(text: String) {
564        use std::time::Instant;
565        let iters = 10000;
566
567        let string_literal = StringLiteral { value: text };
568        let start = Instant::now();
569        for _ in 0..iters {
570            string_literal.unescape().unwrap();
571        }
572        let end = Instant::now();
573        println!(
574            "{iters} on string of length {} iters in {}",
575            string_literal.value.as_str().len(),
576            (end - start).as_secs_f64()
577        )
578    }
579
580    fn generate_string(length: usize, mapper: fn(u32) -> u32) -> String {
581        use rand::{RngCore, thread_rng};
582        let mut rng = thread_rng();
583        let capacity: i64 = (1.2 * length as f64).ceil() as i64;
584        let mut text = String::with_capacity(capacity as usize);
585        text.push('"');
586
587        for _ in 0..capacity {
588            if text.len() > length {
589                break;
590            }
591            match char::from_u32(mapper(rng.next_u32())) {
592                Some('\\') => text += r"\\",
593                Some('\'') => text += r"\'",
594                Some('\"') => text += r#"\""#,
595                Some('\x08') => text += r"\b",
596                Some('\x09') => text += r"\t",
597                Some('\x0a') => text += r"\n",
598                Some('\x0c') => text += r"\f",
599                Some('\x0d') => text += r"\r",
600                Some(ch) => text.push(ch),
601                None => (),
602            }
603        }
604        text.push('"');
605        assert!(text.len() > length && text.len() < length + 10);
606        text
607    }
608}