1use crate::lex::{Cursor, LexError, Scan};
2use crate::lit::Lit;
3use crate::{Span, Spanner};
4
5#[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 From<LitCStr> for Lit {
86 fn from(value: LitCStr) -> Self {
87 Self::CStr(value)
88 }
89}
90
91#[cfg(feature = "serde")]
92impl From<LitCStr> for String {
93 fn from(value: LitCStr) -> Self {
94 value.repr.into_string()
95 }
96}
97
98fn scan_cooked<'a>(c: Cursor<'a>, open: &str) -> Result<Cursor<'a>, LexError> {
99 if !c.starts_with(open) {
100 return c.error().into();
101 }
102
103 let mut c = c.advance_by(open.len());
104
105 loop {
106 match c.first() {
107 None => return c.error().into(),
108 Some('"') => return Ok(c.advance()),
109 Some('\\') => c = escape(c.advance())?,
110 Some(ch) => c = c.advance_by(ch.len_utf8()),
111 }
112 }
113}
114
115fn scan_raw<'a>(start: Cursor<'a>, open: &str) -> Result<Cursor<'a>, LexError> {
116 if !start.starts_with(open) {
117 return start.error().into();
118 }
119
120 let mut cur = start.advance_by(open.len());
121 let mut hashes = 0u32;
122
123 while cur.starts_with("#") {
124 hashes += 1;
125 cur = cur.advance();
126 }
127
128 if !cur.starts_with("\"") {
129 return start.error().into();
130 }
131
132 cur = cur.advance();
133
134 let closing: String = std::iter::once('"')
135 .chain(std::iter::repeat_n('#', hashes as usize))
136 .collect();
137
138 loop {
139 if cur.is_empty() {
140 return start.error().into();
141 }
142
143 if cur.starts_with(&closing) {
144 return Ok(cur.advance_by(closing.len()));
145 }
146
147 if let Some(ch) = cur.first() {
148 cur = cur.advance_by(ch.len_utf8());
149 } else {
150 return start.error().into();
151 }
152 }
153}
154
155fn escape(c: Cursor<'_>) -> Result<Cursor<'_>, LexError> {
156 match c.first() {
157 None => c.error().into(),
158 Some('n' | 'r' | 't' | '\\' | '\'' | '"' | '0') => Ok(c.advance()),
159 Some('x') => {
160 let c = c.advance();
161 let c = hex_digit(c)?;
162 hex_digit(c)
163 }
164 Some('u') => {
165 let c = c.advance();
166
167 if !c.starts_with("{") {
168 return c.error().into();
169 }
170
171 let mut c = c.advance();
172 let mut count = 0;
173
174 loop {
175 match c.first() {
176 Some('}') if count > 0 => return Ok(c.advance()),
177 Some(ch) if ch.is_ascii_hexdigit() && count < 6 => {
178 count += 1;
179 c = c.advance();
180 }
181 _ => return c.error().into(),
182 }
183 }
184 }
185 _ => c.error().into(),
186 }
187}
188
189fn hex_digit(c: Cursor<'_>) -> Result<Cursor<'_>, LexError> {
190 match c.first() {
191 Some(ch) if ch.is_ascii_hexdigit() => Ok(c.advance()),
192 _ => c.error().into(),
193 }
194}
195
196fn decode_string_body(repr: &str) -> Option<String> {
199 if let Some(rest) = repr.strip_prefix('r') {
200 let hashes = rest.bytes().take_while(|b| *b == b'#').count();
201 let open = 1 + hashes;
202 let inner = &rest[open..rest.len() - open];
203 return Some(inner.to_string());
204 }
205
206 let inner = repr.strip_prefix('"')?.strip_suffix('"')?;
207 let mut out = String::new();
208 let mut chars = inner.chars();
209
210 while let Some(c) = chars.next() {
211 if c == '\\' {
212 out.push(decode_escape(&mut chars)?);
213 } else {
214 out.push(c);
215 }
216 }
217
218 Some(out)
219}
220
221fn decode_escape(chars: &mut std::str::Chars<'_>) -> Option<char> {
222 match chars.next()? {
223 'n' => Some('\n'),
224 'r' => Some('\r'),
225 't' => Some('\t'),
226 '\\' => Some('\\'),
227 '\'' => Some('\''),
228 '"' => Some('"'),
229 '0' => Some('\0'),
230 'x' => {
231 let hi = chars.next()?.to_digit(16)?;
232 let lo = chars.next()?.to_digit(16)?;
233 char::from_u32(hi * 16 + lo)
234 }
235 'u' => {
236 if chars.next()? != '{' {
237 return None;
238 }
239 let mut value = 0u32;
240 let mut count = 0;
241 for c in chars.by_ref() {
242 if c == '}' {
243 break;
244 }
245 value = value.checked_mul(16)?.checked_add(c.to_digit(16)?)?;
246 count += 1;
247 if count > 6 {
248 return None;
249 }
250 }
251 char::from_u32(value)
252 }
253 _ => None,
254 }
255}