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