Skip to main content

moxy_token/lit/
str.rs

1use crate::lex::{Cursor, LexError, Scan};
2use crate::lit::Lit;
3use crate::{Span, Spanner};
4
5#[derive(Debug, Clone)]
6#[cfg_attr(feature = "serde", derive(serde::Serialize), serde(into = "String"))]
7pub struct LitStr {
8    value: String,
9    repr: Box<str>,
10    span: Span,
11}
12
13impl LitStr {
14    #[inline]
15    pub fn new(value: &str, span: Span) -> Self {
16        Self {
17            value: value.to_string(),
18            repr: format!("{value:?}").into_boxed_str(),
19            span,
20        }
21    }
22
23    #[inline]
24    pub(crate) fn from_parts(value: String, repr: &str, span: Span) -> Self {
25        Self {
26            value,
27            repr: repr.into(),
28            span,
29        }
30    }
31
32    #[inline]
33    pub fn value(&self) -> &str {
34        &self.value
35    }
36
37    #[inline]
38    pub fn repr(&self) -> &str {
39        &self.repr
40    }
41
42    #[inline]
43    pub fn span(&self) -> Span {
44        self.span
45    }
46
47    #[inline]
48    pub fn set_span(&mut self, span: Span) {
49        self.span = span;
50    }
51}
52
53impl PartialEq for LitStr {
54    fn eq(&self, other: &Self) -> bool {
55        self.value == other.value
56    }
57}
58
59impl Eq for LitStr {}
60
61impl std::hash::Hash for LitStr {
62    fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
63        self.value.hash(state);
64    }
65}
66
67impl std::fmt::Display for LitStr {
68    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
69        f.write_str(&self.repr)
70    }
71}
72
73impl Spanner for LitStr {
74    fn span(&self) -> Span {
75        self.span
76    }
77}
78
79impl Scan for LitStr {
80    fn scan(cursor: Cursor<'_>) -> Result<(Cursor<'_>, Self), LexError> {
81        let end = scan_cooked(cursor, "\"").or_else(|_| scan_raw(cursor, "r"))?;
82        let len = end.offset() as usize - cursor.offset() as usize;
83        let repr = &cursor.rest()[..len];
84        let span = cursor.span_to(&end);
85
86        match decode_string_body(repr) {
87            Some(value) => Ok((end, Self::from_parts(value, repr, span))),
88            None => cursor.error().into(),
89        }
90    }
91}
92
93impl From<LitStr> for Lit {
94    fn from(value: LitStr) -> Self {
95        Self::Str(value)
96    }
97}
98
99#[cfg(feature = "serde")]
100impl From<LitStr> for String {
101    fn from(value: LitStr) -> Self {
102        value.repr.into_string()
103    }
104}
105
106fn scan_cooked<'a>(c: Cursor<'a>, open: &str) -> Result<Cursor<'a>, LexError> {
107    if !c.starts_with(open) {
108        return c.error().into();
109    }
110
111    let mut c = c.advance_by(open.len());
112
113    loop {
114        match c.first() {
115            None => return c.error().into(),
116            Some('"') => return Ok(c.advance()),
117            Some('\\') => c = escape(c.advance())?,
118            Some(ch) => c = c.advance_by(ch.len_utf8()),
119        }
120    }
121}
122
123fn scan_raw<'a>(start: Cursor<'a>, open: &str) -> Result<Cursor<'a>, LexError> {
124    if !start.starts_with(open) {
125        return start.error().into();
126    }
127
128    let mut cur = start.advance_by(open.len());
129    let mut hashes = 0u32;
130
131    while cur.starts_with("#") {
132        hashes += 1;
133        cur = cur.advance();
134    }
135
136    if !cur.starts_with("\"") {
137        return start.error().into();
138    }
139
140    cur = cur.advance();
141
142    let closing: String = std::iter::once('"')
143        .chain(std::iter::repeat_n('#', hashes as usize))
144        .collect();
145
146    loop {
147        if cur.is_empty() {
148            return start.error().into();
149        }
150
151        if cur.starts_with(&closing) {
152            return Ok(cur.advance_by(closing.len()));
153        }
154
155        if let Some(ch) = cur.first() {
156            cur = cur.advance_by(ch.len_utf8());
157        } else {
158            return start.error().into();
159        }
160    }
161}
162
163fn escape(c: Cursor<'_>) -> Result<Cursor<'_>, LexError> {
164    match c.first() {
165        None => c.error().into(),
166        Some('n' | 'r' | 't' | '\\' | '\'' | '"' | '0') => Ok(c.advance()),
167        Some('x') => {
168            let c = c.advance();
169            let c = hex_digit(c)?;
170            hex_digit(c)
171        }
172        Some('u') => {
173            let c = c.advance();
174
175            if !c.starts_with("{") {
176                return c.error().into();
177            }
178
179            let mut c = c.advance();
180            let mut count = 0;
181
182            loop {
183                match c.first() {
184                    Some('}') if count > 0 => return Ok(c.advance()),
185                    Some(ch) if ch.is_ascii_hexdigit() && count < 6 => {
186                        count += 1;
187                        c = c.advance();
188                    }
189                    _ => return c.error().into(),
190                }
191            }
192        }
193        _ => c.error().into(),
194    }
195}
196
197fn hex_digit(c: Cursor<'_>) -> Result<Cursor<'_>, LexError> {
198    match c.first() {
199        Some(ch) if ch.is_ascii_hexdigit() => Ok(c.advance()),
200        _ => c.error().into(),
201    }
202}
203
204/// Decode the body of a string repr (after any `b`/`c` prefix the caller strips).
205/// Handles cooked (`"…"`) and raw (`r"…"` / `r#"…"#`) forms.
206fn decode_string_body(repr: &str) -> Option<String> {
207    if let Some(rest) = repr.strip_prefix('r') {
208        let hashes = rest.bytes().take_while(|b| *b == b'#').count();
209        let open = 1 + hashes;
210        let inner = &rest[open..rest.len() - open];
211        return Some(inner.to_string());
212    }
213
214    let inner = repr.strip_prefix('"')?.strip_suffix('"')?;
215    let mut out = String::new();
216    let mut chars = inner.chars();
217
218    while let Some(c) = chars.next() {
219        if c == '\\' {
220            out.push(decode_escape(&mut chars)?);
221        } else {
222            out.push(c);
223        }
224    }
225
226    Some(out)
227}
228
229fn decode_escape(chars: &mut std::str::Chars<'_>) -> Option<char> {
230    match chars.next()? {
231        'n' => Some('\n'),
232        'r' => Some('\r'),
233        't' => Some('\t'),
234        '\\' => Some('\\'),
235        '\'' => Some('\''),
236        '"' => Some('"'),
237        '0' => Some('\0'),
238        'x' => {
239            let hi = chars.next()?.to_digit(16)?;
240            let lo = chars.next()?.to_digit(16)?;
241            char::from_u32(hi * 16 + lo)
242        }
243        'u' => {
244            if chars.next()? != '{' {
245                return None;
246            }
247            let mut value = 0u32;
248            let mut count = 0;
249            for c in chars.by_ref() {
250                if c == '}' {
251                    break;
252                }
253                value = value.checked_mul(16)?.checked_add(c.to_digit(16)?)?;
254                count += 1;
255                if count > 6 {
256                    return None;
257                }
258            }
259            char::from_u32(value)
260        }
261        _ => None,
262    }
263}