Skip to main content

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