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