moxy_token/lit/
byte_str.rs1use 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 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 Scan for LitByteStr {
72 fn scan(cursor: Cursor<'_>) -> Result<(Cursor<'_>, Self), LexError> {
73 let end = scan_cooked(cursor, "b\"").or_else(|_| scan_raw(cursor, "br"))?;
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('b').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 LitByteStr {
86 fn parse(stream: &mut ParseStream) -> Result<Self, ParseError> {
87 let at = stream.span();
88
89 match stream.parse::<Lit>()? {
90 Lit::ByteStr(v) => Ok(v),
91 _ => Err(LexError::new(at).message("expected byte string literal").into()),
92 }
93 }
94}
95
96impl From<LitByteStr> for Lit {
97 fn from(value: LitByteStr) -> Self {
98 Self::ByteStr(value)
99 }
100}
101
102#[cfg(feature = "serde")]
103impl From<LitByteStr> for String {
104 fn from(value: LitByteStr) -> 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
207fn 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}