Skip to main content

moxy_token/lit/
str.rs

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