1use crate::lex::{Cursor, LexError, Scan};
2use crate::lit::Lit;
3use crate::{Span, Spanner, ToTokens, TokenStream, TokenTree};
4
5#[derive(Debug, Clone)]
7#[cfg_attr(feature = "serde", derive(serde::Serialize), serde(into = "String"))]
8pub struct LitByte {
9 value: u8,
10 repr: Box<str>,
11 span: Span,
12}
13
14impl LitByte {
15 #[inline]
16 pub(crate) fn from_parts(value: 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 LitByte {
46 fn eq(&self, other: &Self) -> bool {
47 self.value == other.value
48 }
49}
50
51impl Eq for LitByte {}
52
53impl std::hash::Hash for LitByte {
54 fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
55 self.value.hash(state);
56 }
57}
58
59impl std::fmt::Display for LitByte {
60 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
61 f.write_str(&self.repr)
62 }
63}
64
65impl Spanner for LitByte {
66 fn span(&self) -> Span {
67 self.span
68 }
69}
70
71impl ToTokens for LitByte {
72 fn to_tokens(&self, tokens: &mut TokenStream) {
73 tokens.extend_one(TokenTree::Literal(self.clone().into()));
74 }
75}
76
77impl Scan for LitByte {
78 fn scan(cursor: Cursor<'_>) -> Result<(Cursor<'_>, Self), LexError> {
79 let end = scan(cursor, "b'")?;
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 let inner = repr.strip_prefix("b'").and_then(|r| r.strip_suffix('\''));
84
85 match inner.and_then(decode_one_char) {
86 Some(value) => Ok((end, Self::from_parts(value as u8, repr, span))),
87 _ => cursor.error().into(),
88 }
89 }
90}
91
92impl From<LitByte> for Lit {
93 fn from(value: LitByte) -> Self {
94 Self::Byte(value)
95 }
96}
97
98#[cfg(feature = "serde")]
99impl From<LitByte> for String {
100 fn from(value: LitByte) -> Self {
101 value.repr.into_string()
102 }
103}
104
105fn scan<'a>(c: Cursor<'a>, open: &str) -> Result<Cursor<'a>, LexError> {
106 if !c.starts_with(open) {
107 return c.error().into();
108 }
109
110 let start = c;
111 let c = c.advance_by(open.len());
112 let c = match c.first() {
113 None | Some('\'') => return start.error().into(),
114 Some('\\') => escape(c.advance())?,
115 Some(ch) => c.advance_by(ch.len_utf8()),
116 };
117
118 if !c.starts_with("'") {
119 return start.error().into();
120 }
121
122 Ok(c.advance())
123}
124
125fn escape(c: Cursor<'_>) -> Result<Cursor<'_>, LexError> {
126 match c.first() {
127 None => c.error().into(),
128 Some('n' | 'r' | 't' | '\\' | '\'' | '"' | '0') => Ok(c.advance()),
129 Some('x') => {
130 let c = c.advance();
131 let c = hex_digit(c)?;
132 hex_digit(c)
133 }
134 Some('u') => {
135 let c = c.advance();
136
137 if !c.starts_with("{") {
138 return c.error().into();
139 }
140
141 let mut c = c.advance();
142 let mut count = 0;
143
144 loop {
145 match c.first() {
146 Some('}') if count > 0 => return Ok(c.advance()),
147 Some(ch) if ch.is_ascii_hexdigit() && count < 6 => {
148 count += 1;
149 c = c.advance();
150 }
151 _ => return c.error().into(),
152 }
153 }
154 }
155 _ => c.error().into(),
156 }
157}
158
159fn hex_digit(c: Cursor<'_>) -> Result<Cursor<'_>, LexError> {
160 match c.first() {
161 Some(ch) if ch.is_ascii_hexdigit() => Ok(c.advance()),
162 _ => c.error().into(),
163 }
164}
165
166fn decode_one_char(inner: &str) -> Option<char> {
168 let mut chars = inner.chars();
169 let value = match chars.next()? {
170 '\\' => decode_escape(&mut chars)?,
171 c => c,
172 };
173
174 if chars.next().is_some() {
175 return None;
176 }
177
178 Some(value)
179}
180
181fn decode_escape(chars: &mut std::str::Chars<'_>) -> Option<char> {
182 match chars.next()? {
183 'n' => Some('\n'),
184 'r' => Some('\r'),
185 't' => Some('\t'),
186 '\\' => Some('\\'),
187 '\'' => Some('\''),
188 '"' => Some('"'),
189 '0' => Some('\0'),
190 'x' => {
191 let hi = chars.next()?.to_digit(16)?;
192 let lo = chars.next()?.to_digit(16)?;
193 char::from_u32(hi * 16 + lo)
194 }
195 'u' => {
196 if chars.next()? != '{' {
197 return None;
198 }
199
200 let mut value = 0u32;
201 let mut count = 0;
202
203 for c in chars.by_ref() {
204 if c == '}' {
205 break;
206 }
207
208 value = value.checked_mul(16)?.checked_add(c.to_digit(16)?)?;
209 count += 1;
210
211 if count > 6 {
212 return None;
213 }
214 }
215
216 char::from_u32(value)
217 }
218 _ => None,
219 }
220}