Skip to main content

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