Skip to main content

moxy_token/lit/
byte_str.rs

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