1use crate::error::ParseError;
4use itertools::Itertools;
5use logos::{Logos, SpannedIter};
6use std::{fmt::Display, num::NonZeroU8, sync::LazyLock};
7use wgsl_types::idents::RESERVED_WORDS;
8
9type Span = std::ops::Range<usize>;
10
11fn maybe_template_end(
12 lex: &mut logos::Lexer<Token>,
13 current: Token,
14 lookahead: Option<Token>,
15) -> Token {
16 if let Some(depth) = lex.extras.template_depths.last() {
17 if lex.extras.depth == *depth {
19 lex.extras.template_depths.pop();
20 if let Some(depth) = lex.extras.template_depths.last() {
23 if lex.extras.depth == *depth && lookahead == Some(Token::SymGreaterThan) {
24 lex.extras.template_depths.pop();
25 lex.extras.lookahead = Some(Token::TemplateArgsEnd);
26 } else {
27 lex.extras.lookahead = lookahead;
28 }
29 } else {
30 lex.extras.lookahead = lookahead;
31 }
32 return Token::TemplateArgsEnd;
33 }
34 }
35
36 current
37}
38
39fn maybe_fail_template(lex: &mut logos::Lexer<Token>) -> bool {
42 if let Some(depth) = lex.extras.template_depths.last()
43 && lex.extras.depth == *depth
44 {
45 return false;
46 }
47 true
48}
49
50fn incr_depth(lex: &mut logos::Lexer<Token>) {
51 lex.extras.depth += 1;
52}
53
54fn decr_depth(lex: &mut logos::Lexer<Token>) {
55 lex.extras.depth -= 1;
56}
57
58const DEC_FORMAT: u128 = lexical::NumberFormatBuilder::new().build_unchecked();
62
63const HEX_FORMAT: u128 = lexical::NumberFormatBuilder::new()
65 .mantissa_radix(16)
66 .base_prefix(NonZeroU8::new(b'x'))
67 .exponent_base(NonZeroU8::new(16))
68 .exponent_radix(NonZeroU8::new(10))
69 .build_unchecked();
70
71static FLOAT_HEX_OPTIONS: LazyLock<lexical::parse_float_options::Options> = LazyLock::new(|| {
72 lexical::parse_float_options::OptionsBuilder::new()
73 .exponent(b'p')
74 .decimal_point(b'.')
75 .build()
76 .unwrap()
77});
78
79fn parse_dec_abstract_int(lex: &mut logos::Lexer<Token>) -> Option<i64> {
80 let options = &lexical::parse_integer_options::STANDARD;
81 let str = lex.slice();
82 lexical::parse_with_options::<i64, _, DEC_FORMAT>(str, options).ok()
83}
84
85fn parse_hex_abstract_int(lex: &mut logos::Lexer<Token>) -> Option<i64> {
86 let options = &lexical::parse_integer_options::STANDARD;
87 let str = lex.slice();
88 lexical::parse_with_options::<i64, _, HEX_FORMAT>(str, options).ok()
89}
90
91fn parse_dec_i32(lex: &mut logos::Lexer<Token>) -> Option<i32> {
92 let options = &lexical::parse_integer_options::STANDARD;
93 let str = lex.slice();
94 let str = &str[..str.len() - 1];
95 lexical::parse_with_options::<i32, _, DEC_FORMAT>(str, options).ok()
96}
97
98fn parse_hex_i32(lex: &mut logos::Lexer<Token>) -> Option<i32> {
99 let options = &lexical::parse_integer_options::STANDARD;
100 let str = lex.slice();
101 let str = &str[..str.len() - 1];
102 lexical::parse_with_options::<i32, _, HEX_FORMAT>(str, options).ok()
103}
104
105fn parse_dec_u32(lex: &mut logos::Lexer<Token>) -> Option<u32> {
106 let options = &lexical::parse_integer_options::STANDARD;
107 let str = lex.slice();
108 let str = &str[..str.len() - 1];
109 lexical::parse_with_options::<u32, _, DEC_FORMAT>(str, options).ok()
110}
111
112fn parse_hex_u32(lex: &mut logos::Lexer<Token>) -> Option<u32> {
113 let options = &lexical::parse_integer_options::STANDARD;
114 let str = lex.slice();
115 let str = &str[..str.len() - 1];
116 lexical::parse_with_options::<u32, _, HEX_FORMAT>(str, options).ok()
117}
118
119fn parse_dec_abs_float(lex: &mut logos::Lexer<Token>) -> Option<f64> {
120 let options = &lexical::parse_float_options::STANDARD;
121 let str = lex.slice();
122 lexical::parse_with_options::<f64, _, DEC_FORMAT>(str, options).ok()
123}
124
125fn parse_hex_abs_float(lex: &mut logos::Lexer<Token>) -> Option<f64> {
126 let str = lex.slice();
127 lexical::parse_with_options::<f64, _, HEX_FORMAT>(str, &FLOAT_HEX_OPTIONS).ok()
128}
129
130fn parse_dec_f32(lex: &mut logos::Lexer<Token>) -> Option<f32> {
131 let options = &lexical::parse_float_options::STANDARD;
132 let str = lex.slice();
133 let str = &str[..str.len() - 1];
134 lexical::parse_with_options::<f32, _, DEC_FORMAT>(str, options).ok()
135}
136
137fn parse_hex_f32(lex: &mut logos::Lexer<Token>) -> Option<f32> {
138 let str = lex.slice();
139 let options = &lexical::parse_float_options::STANDARD;
141 let str = &str[..str.len() - 1];
142 lexical::parse_with_options::<f32, _, HEX_FORMAT>(str, options).ok()
143}
144
145fn parse_dec_f16(lex: &mut logos::Lexer<Token>) -> Option<f32> {
146 let options = &lexical::parse_float_options::STANDARD;
147 let str = lex.slice();
148 let str = &str[..str.len() - 1];
149 lexical::parse_with_options::<f32, _, DEC_FORMAT>(str, options).ok()
150}
151
152fn parse_hex_f16(lex: &mut logos::Lexer<Token>) -> Option<f32> {
153 let str = lex.slice();
154 let str = &str[..str.len() - 1];
155 lexical::parse_with_options::<f32, _, HEX_FORMAT>(str, &FLOAT_HEX_OPTIONS).ok()
156}
157
158#[cfg(feature = "naga-ext")]
159fn parse_dec_i64(lex: &mut logos::Lexer<Token>) -> Option<i64> {
160 let options = &lexical::parse_integer_options::STANDARD;
161 let str = lex.slice();
162 let str = &str[..str.len() - 2];
163 lexical::parse_with_options::<i64, _, DEC_FORMAT>(str, options).ok()
164}
165
166#[cfg(feature = "naga-ext")]
167fn parse_hex_i64(lex: &mut logos::Lexer<Token>) -> Option<i64> {
168 let options = &lexical::parse_integer_options::STANDARD;
169 let str = lex.slice();
170 let str = &str[..str.len() - 2];
171 lexical::parse_with_options::<i64, _, HEX_FORMAT>(str, options).ok()
172}
173
174#[cfg(feature = "naga-ext")]
175fn parse_dec_u64(lex: &mut logos::Lexer<Token>) -> Option<u64> {
176 let options = &lexical::parse_integer_options::STANDARD;
177 let str = lex.slice();
178 let str = &str[..str.len() - 2];
179 lexical::parse_with_options::<u64, _, DEC_FORMAT>(str, options).ok()
180}
181
182#[cfg(feature = "naga-ext")]
183fn parse_hex_u64(lex: &mut logos::Lexer<Token>) -> Option<u64> {
184 let options = &lexical::parse_integer_options::STANDARD;
185 let str = lex.slice();
186 let str = &str[..str.len() - 2];
187 lexical::parse_with_options::<u64, _, HEX_FORMAT>(str, options).ok()
188}
189
190#[cfg(feature = "naga-ext")]
191fn parse_dec_f64(lex: &mut logos::Lexer<Token>) -> Option<f64> {
192 let options = &lexical::parse_float_options::STANDARD;
193 let str = lex.slice();
194 let str = &str[..str.len() - 2];
195 lexical::parse_with_options::<f64, _, DEC_FORMAT>(str, options).ok()
196}
197
198#[cfg(feature = "naga-ext")]
199fn parse_hex_f64(lex: &mut logos::Lexer<Token>) -> Option<f64> {
200 let str = lex.slice();
201 let options = &lexical::parse_float_options::STANDARD;
203 let str = &str[..str.len() - 2];
204 lexical::parse_with_options::<f64, _, HEX_FORMAT>(str, options).ok()
205}
206
207fn parse_line_comment(lex: &mut logos::Lexer<Token>) {
208 let rem = lex.remainder();
209 let line_end = rem
211 .char_indices()
212 .find(|(_, c)| "\n\u{000B}\u{000C}\r\u{0085}\u{2028}\u{2029}".contains(*c))
213 .map(|(i, _)| i)
214 .unwrap_or(rem.len());
215 lex.bump(line_end);
216}
217
218fn parse_block_comment(lex: &mut logos::Lexer<Token>) {
219 let mut depth = 1;
220 while depth > 0 {
221 let rem = lex.remainder();
222 if rem.is_empty() {
223 break;
224 } else if rem.starts_with("/*") {
225 lex.bump(2);
226 depth += 1;
227 } else if rem.starts_with("*/") {
228 lex.bump(2);
229 depth -= 1;
230 } else {
231 let mut next_char = 1;
232 while !rem.is_char_boundary(next_char) {
233 next_char += 1;
234 }
235 lex.bump(next_char);
236 }
237 }
238}
239
240fn parse_ident(lex: &mut logos::Lexer<Token>) -> Token {
241 let ident = lex.slice().to_string();
242 if RESERVED_WORDS.iter().contains(&ident.as_str()) {
243 Token::ReservedWord(ident)
244 } else {
245 Token::Ident(ident)
246 }
247}
248
249#[derive(Default, Clone, Debug, PartialEq)]
250pub struct LexerState {
251 depth: i32,
252 template_depths: Vec<i32>,
253 lookahead: Option<Token>,
254}
255
256#[derive(Logos, Clone, Debug, PartialEq)]
258#[logos(
259 skip r"[\s\u0085\u200e\u200f\u2028\u2029]+", extras = LexerState,
262 error = ParseError)]
263pub enum Token {
264 EntryPointTryTemplateList,
270 EntryPointTranslationUnit,
271 EntryPointGlobalDecl,
272 EntryPointLiteral,
273 EntryPointGlobalDirective,
274 EntryPointExpression,
275 EntryPointStatement,
276 #[cfg(feature = "imports")]
277 EntryPointImportStatement,
278
279 #[token("//", parse_line_comment)]
280 LineComment,
281 #[token("/*", parse_block_comment, priority = 2)]
282 BlockComment,
283 #[regex(
287 r#"([_\p{XID_Start}][\p{XID_Continue}]+)|([\p{XID_Start}])"#,
288 parse_ident,
289 priority = 1
290 )]
291 Ignored,
292 #[token("&")]
295 SymAnd,
296 #[token("&&", maybe_fail_template)]
297 SymAndAnd,
298 #[token("->")]
299 SymArrow,
300 #[token("@")]
301 SymAttr,
302 #[token("/")]
303 SymForwardSlash,
304 #[token("!")]
305 SymBang,
306 #[token("[", incr_depth)]
307 SymBracketLeft,
308 #[token("]", decr_depth)]
309 SymBracketRight,
310 #[token("{")]
311 SymBraceLeft,
312 #[token("}")]
313 SymBraceRight,
314 #[token(":")]
315 SymColon,
316 #[token(",")]
317 SymComma,
318 #[token("=")]
319 SymEqual,
320 #[token("==")]
321 SymEqualEqual,
322 #[token("!=")]
323 SymNotEqual,
324 #[token(">", |lex| maybe_template_end(lex, Token::SymGreaterThan, None))]
325 SymGreaterThan,
326 #[token(">=", |lex| maybe_template_end(lex, Token::SymGreaterThanEqual, Some(Token::SymEqual)))]
327 SymGreaterThanEqual,
328 #[token(">>", |lex| maybe_template_end(lex, Token::SymShiftRight, Some(Token::SymGreaterThan)))]
329 SymShiftRight,
330 #[token("<")]
331 SymLessThan,
332 #[token("<=")]
333 SymLessThanEqual,
334 #[token("<<")]
335 SymShiftLeft,
336 #[token("%")]
337 SymModulo,
338 #[token("-")]
339 SymMinus,
340 #[token("--")]
341 SymMinusMinus,
342 #[token(".")]
343 SymPeriod,
344 #[token("+")]
345 SymPlus,
346 #[token("++")]
347 SymPlusPlus,
348 #[token("|")]
349 SymOr,
350 #[token("||", maybe_fail_template)]
351 SymOrOr,
352 #[token("(", incr_depth)]
353 SymParenLeft,
354 #[token(")", decr_depth)]
355 SymParenRight,
356 #[token(";")]
357 SymSemicolon,
358 #[token("*")]
359 SymStar,
360 #[token("~")]
361 SymTilde,
362 #[token("_")]
363 SymUnderscore,
364 #[token("^")]
365 SymXor,
366 #[token("+=")]
367 SymPlusEqual,
368 #[token("-=")]
369 SymMinusEqual,
370 #[token("*=")]
371 SymTimesEqual,
372 #[token("/=")]
373 SymDivisionEqual,
374 #[token("%=")]
375 SymModuloEqual,
376 #[token("&=")]
377 SymAndEqual,
378 #[token("|=")]
379 SymOrEqual,
380 #[token("^=")]
381 SymXorEqual,
382 #[token(">>=", |lex| maybe_template_end(lex, Token::SymShiftRightAssign, Some(Token::SymGreaterThanEqual)))]
383 SymShiftRightAssign,
384 #[token("<<=")]
385 SymShiftLeftAssign,
386
387 #[token("alias")]
390 KwAlias,
391 #[token("break")]
392 KwBreak,
393 #[token("case")]
394 KwCase,
395 #[token("const", priority = 2)]
396 KwConst,
397 #[token("const_assert")]
398 KwConstAssert,
399 #[token("continue")]
400 KwContinue,
401 #[token("continuing")]
402 KwContinuing,
403 #[token("default")]
404 KwDefault,
405 #[token("diagnostic")]
406 KwDiagnostic,
407 #[token("discard")]
408 KwDiscard,
409 #[token("else")]
410 KwElse,
411 #[token("enable")]
412 KwEnable,
413 #[token("false")]
414 KwFalse,
415 #[token("fn")]
416 KwFn,
417 #[token("for")]
418 KwFor,
419 #[token("if")]
420 KwIf,
421 #[token("let")]
422 KwLet,
423 #[token("loop")]
424 KwLoop,
425 #[token("override")]
426 KwOverride,
427 #[token("requires")]
428 KwRequires,
429 #[token("return")]
430 KwReturn,
431 #[token("struct")]
432 KwStruct,
433 #[token("switch")]
434 KwSwitch,
435 #[token("true")]
436 KwTrue,
437 #[token("var")]
438 KwVar,
439 #[token("while")]
440 KwWhile,
441
442 Ident(String),
445 ReservedWord(String),
448
449 #[regex(r#"0|[1-9]\d*"#, parse_dec_abstract_int)]
450 #[regex(r#"0[xX][\da-fA-F]+"#, parse_hex_abstract_int)]
451 AbstractInt(i64),
452 #[regex(r#"(\d+\.\d*|\.\d+)([eE][+-]?\d+)?"#, parse_dec_abs_float)]
453 #[regex(r#"\d+[eE][+-]?\d+"#, parse_dec_abs_float)]
454 #[regex(r#"0[xX][\da-fA-F]+\.[\da-fA-F]*([pP][+-]?\d+)?"#, parse_hex_abs_float)]
455 #[regex(r#"0[xX]\.[\da-fA-F]+([pP][+-]?\d+)?"#, parse_hex_abs_float)]
456 #[regex(r#"0[xX][\da-fA-F]+[pP][+-]?\d+"#, parse_hex_abs_float)]
457 AbstractFloat(f64),
459 #[regex(r#"(0|[1-9]\d*)i"#, parse_dec_i32)]
460 #[regex(r#"0[xX][\da-fA-F]+i"#, parse_hex_i32)]
461 I32(i32),
463 #[regex(r#"(0|[1-9]\d*)u"#, parse_dec_u32)]
464 #[regex(r#"0[xX][\da-fA-F]+u"#, parse_hex_u32)]
465 U32(u32),
467 #[regex(r#"(\d+\.\d*|\.\d+)([eE][+-]?\d+)?f"#, parse_dec_f32)]
468 #[regex(r#"\d+([eE][+-]?\d+)?f"#, parse_dec_f32)]
469 #[regex(r#"0[xX][\da-fA-F]+\.[\da-fA-F]*[pP][+-]?\d+f"#, parse_hex_f32)]
470 #[regex(r#"0[xX]\.[\da-fA-F]+[pP][+-]?\d+f"#, parse_hex_f32)]
471 #[regex(r#"0[xX][\da-fA-F]+[pP][+-]?\d+f"#, parse_hex_f32)]
472 F32(f32),
473 #[regex(r#"(\d+\.\d*|\.\d+)([eE][+-]?\d+)?h"#, parse_dec_f16)]
474 #[regex(r#"\d+([eE][+-]?\d+)?h"#, parse_dec_f16)]
475 #[regex(r#"0[xX][\da-fA-F]+\.[\da-fA-F]*[pP][+-]?\d+h"#, parse_hex_f16)]
476 #[regex(r#"0[xX]\.[\da-fA-F]+[pP][+-]?\d+h"#, parse_hex_f16)]
477 #[regex(r#"0[xX][\da-fA-F]+[pP][+-]?\d+h"#, parse_hex_f16)]
478 F16(f32),
479 #[cfg(feature = "naga-ext")]
480 #[regex(r#"(0|[1-9]\d*)li"#, parse_dec_i64)]
481 #[regex(r#"0[xX][\da-fA-F]+li"#, parse_hex_i64)]
482 I64(i64),
484 #[cfg(feature = "naga-ext")]
485 #[regex(r#"(0|[1-9]\d*)lu"#, parse_dec_u64)]
486 #[regex(r#"0[xX][\da-fA-F]+lu"#, parse_hex_u64)]
487 U64(u64),
489 #[cfg(feature = "naga-ext")]
490 #[regex(r#"(\d+\.\d*|\.\d+)([eE][+-]?\d+)?lf"#, parse_dec_f64)]
491 #[regex(r#"\d+([eE][+-]?\d+)?lf"#, parse_dec_f64)]
492 #[regex(r#"0[xX][\da-fA-F]+\.[\da-fA-F]*[pP][+-]?\d+lf"#, parse_hex_f64)]
493 #[regex(r#"0[xX]\.[\da-fA-F]+[pP][+-]?\d+lf"#, parse_hex_f64)]
494 #[regex(r#"0[xX][\da-fA-F]+[pP][+-]?\d+lf"#, parse_hex_f64)]
495 F64(f64),
496 TemplateArgsStart,
497 TemplateArgsEnd,
498
499 #[cfg(feature = "imports")]
503 #[token("::")]
504 SymColonColon,
505 #[cfg(feature = "imports")]
506 #[token("self")]
507 KwSelf,
508 #[cfg(feature = "imports")]
509 #[token("super")]
510 KwSuper,
511 #[cfg(feature = "imports")]
512 #[token("package")]
513 KwPackage,
514 #[cfg(feature = "imports")]
515 #[token("as")]
516 KwAs,
517 #[cfg(feature = "imports")]
518 #[token("import")]
519 KwImport,
520}
521
522impl Token {
523 pub fn is_trivia(&self) -> bool {
524 matches!(
525 self,
526 Token::LineComment | Token::BlockComment | Token::Ignored
527 )
528 }
529
530 #[allow(unused)]
531 pub fn is_symbol(&self) -> bool {
532 matches!(
533 self,
534 Token::SymAnd
535 | Token::SymAndAnd
536 | Token::SymArrow
537 | Token::SymAttr
538 | Token::SymForwardSlash
539 | Token::SymBang
540 | Token::SymBracketLeft
541 | Token::SymBracketRight
542 | Token::SymBraceLeft
543 | Token::SymBraceRight
544 | Token::SymColon
545 | Token::SymComma
546 | Token::SymEqual
547 | Token::SymEqualEqual
548 | Token::SymNotEqual
549 | Token::SymGreaterThan
550 | Token::SymGreaterThanEqual
551 | Token::SymShiftRight
552 | Token::SymLessThan
553 | Token::SymLessThanEqual
554 | Token::SymShiftLeft
555 | Token::SymModulo
556 | Token::SymMinus
557 | Token::SymMinusMinus
558 | Token::SymPeriod
559 | Token::SymPlus
560 | Token::SymPlusPlus
561 | Token::SymOr
562 | Token::SymOrOr
563 | Token::SymParenLeft
564 | Token::SymParenRight
565 | Token::SymSemicolon
566 | Token::SymStar
567 | Token::SymTilde
568 | Token::SymUnderscore
569 | Token::SymXor
570 | Token::SymPlusEqual
571 | Token::SymMinusEqual
572 | Token::SymTimesEqual
573 | Token::SymDivisionEqual
574 | Token::SymModuloEqual
575 | Token::SymAndEqual
576 | Token::SymOrEqual
577 | Token::SymXorEqual
578 | Token::SymShiftRightAssign
579 | Token::SymShiftLeftAssign
580 )
581 }
582
583 #[allow(unused)]
584 pub fn is_keyword(&self) -> bool {
585 matches!(
586 self,
587 Token::KwAlias
588 | Token::KwBreak
589 | Token::KwCase
590 | Token::KwConst
591 | Token::KwConstAssert
592 | Token::KwContinue
593 | Token::KwContinuing
594 | Token::KwDefault
595 | Token::KwDiagnostic
596 | Token::KwDiscard
597 | Token::KwElse
598 | Token::KwEnable
599 | Token::KwFalse
600 | Token::KwFn
601 | Token::KwFor
602 | Token::KwIf
603 | Token::KwLet
604 | Token::KwLoop
605 | Token::KwOverride
606 | Token::KwRequires
607 | Token::KwReturn
608 | Token::KwStruct
609 | Token::KwSwitch
610 | Token::KwTrue
611 | Token::KwVar
612 | Token::KwWhile
613 )
614 }
615
616 #[allow(unused)]
617 pub fn is_numeric_literal(&self) -> bool {
618 matches!(
619 self,
620 Token::AbstractInt(_)
621 | Token::AbstractFloat(_)
622 | Token::I32(_)
623 | Token::U32(_)
624 | Token::F32(_)
625 | Token::F16(_)
626 )
627 }
628}
629
630impl Display for Token {
631 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
633 match self {
634 Token::EntryPointTryTemplateList => f.write_str("EntryPointTryTemplateList"),
635 Token::EntryPointTranslationUnit => f.write_str("EntryPointTranslationUnit"),
636 Token::EntryPointGlobalDecl => f.write_str("EntryPointGlobalDecl"),
637 Token::EntryPointLiteral => f.write_str("EntryPointLiteral"),
638 Token::EntryPointGlobalDirective => f.write_str("EntryPointGlobalDirective"),
639 Token::EntryPointExpression => f.write_str("EntryPointExpression"),
640 Token::EntryPointStatement => f.write_str("EntryPointStatement"),
641 #[cfg(feature = "imports")]
642 Token::EntryPointImportStatement => f.write_str("EntryPointImportStatement"),
643 Token::LineComment => f.write_str("// line comment"),
644 Token::BlockComment => f.write_str("/* block comment */"),
645 Token::Ignored => unreachable!(),
646 Token::SymAnd => f.write_str("&"),
647 Token::SymAndAnd => f.write_str("&&"),
648 Token::SymArrow => f.write_str("->"),
649 Token::SymAttr => f.write_str("@"),
650 Token::SymForwardSlash => f.write_str("/"),
651 Token::SymBang => f.write_str("!"),
652 Token::SymBracketLeft => f.write_str("["),
653 Token::SymBracketRight => f.write_str("]"),
654 Token::SymBraceLeft => f.write_str("{"),
655 Token::SymBraceRight => f.write_str("}"),
656 Token::SymColon => f.write_str(":"),
657 Token::SymComma => f.write_str(","),
658 Token::SymEqual => f.write_str("="),
659 Token::SymEqualEqual => f.write_str("=="),
660 Token::SymNotEqual => f.write_str("!="),
661 Token::SymGreaterThan => f.write_str(">"),
662 Token::SymGreaterThanEqual => f.write_str(">="),
663 Token::SymShiftRight => f.write_str(">>"),
664 Token::SymLessThan => f.write_str("<"),
665 Token::SymLessThanEqual => f.write_str("<="),
666 Token::SymShiftLeft => f.write_str("<<"),
667 Token::SymModulo => f.write_str("%"),
668 Token::SymMinus => f.write_str("-"),
669 Token::SymMinusMinus => f.write_str("--"),
670 Token::SymPeriod => f.write_str("."),
671 Token::SymPlus => f.write_str("+"),
672 Token::SymPlusPlus => f.write_str("++"),
673 Token::SymOr => f.write_str("|"),
674 Token::SymOrOr => f.write_str("||"),
675 Token::SymParenLeft => f.write_str("("),
676 Token::SymParenRight => f.write_str(")"),
677 Token::SymSemicolon => f.write_str(";"),
678 Token::SymStar => f.write_str("*"),
679 Token::SymTilde => f.write_str("~"),
680 Token::SymUnderscore => f.write_str("_"),
681 Token::SymXor => f.write_str("^"),
682 Token::SymPlusEqual => f.write_str("+="),
683 Token::SymMinusEqual => f.write_str("-="),
684 Token::SymTimesEqual => f.write_str("*="),
685 Token::SymDivisionEqual => f.write_str("/="),
686 Token::SymModuloEqual => f.write_str("%="),
687 Token::SymAndEqual => f.write_str("&="),
688 Token::SymOrEqual => f.write_str("|="),
689 Token::SymXorEqual => f.write_str("^="),
690 Token::SymShiftRightAssign => f.write_str(">>="),
691 Token::SymShiftLeftAssign => f.write_str("<<="),
692 Token::KwAlias => f.write_str("alias"),
693 Token::KwBreak => f.write_str("break"),
694 Token::KwCase => f.write_str("case"),
695 Token::KwConst => f.write_str("const"),
696 Token::KwConstAssert => f.write_str("const_assert"),
697 Token::KwContinue => f.write_str("continue"),
698 Token::KwContinuing => f.write_str("continuing"),
699 Token::KwDefault => f.write_str("default"),
700 Token::KwDiagnostic => f.write_str("diagnostic"),
701 Token::KwDiscard => f.write_str("discard"),
702 Token::KwElse => f.write_str("else"),
703 Token::KwEnable => f.write_str("enable"),
704 Token::KwFalse => f.write_str("false"),
705 Token::KwFn => f.write_str("fn"),
706 Token::KwFor => f.write_str("for"),
707 Token::KwIf => f.write_str("if"),
708 Token::KwLet => f.write_str("let"),
709 Token::KwLoop => f.write_str("loop"),
710 Token::KwOverride => f.write_str("override"),
711 Token::KwRequires => f.write_str("requires"),
712 Token::KwReturn => f.write_str("return"),
713 Token::KwStruct => f.write_str("struct"),
714 Token::KwSwitch => f.write_str("switch"),
715 Token::KwTrue => f.write_str("true"),
716 Token::KwVar => f.write_str("var"),
717 Token::KwWhile => f.write_str("while"),
718 Token::Ident(s) => write!(f, "identifier `{s}`"),
719 Token::ReservedWord(s) => write!(f, "reserved word `{s}`"),
720 Token::AbstractInt(n) => write!(f, "{n}"),
721 Token::AbstractFloat(n) => write!(f, "{n}"),
722 Token::I32(n) => write!(f, "{n}i"),
723 Token::U32(n) => write!(f, "{n}u"),
724 Token::F32(n) => write!(f, "{n}f"),
725 Token::F16(n) => write!(f, "{n}h"),
726 #[cfg(feature = "naga-ext")]
727 Token::I64(n) => write!(f, "{n}li"),
728 #[cfg(feature = "naga-ext")]
729 Token::U64(n) => write!(f, "{n}lu"),
730 #[cfg(feature = "naga-ext")]
731 Token::F64(n) => write!(f, "{n}lf"),
732 Token::TemplateArgsStart => f.write_str("start of template"),
733 Token::TemplateArgsEnd => f.write_str("end of template"),
734 #[cfg(feature = "imports")]
735 Token::SymColonColon => write!(f, "::"),
736 #[cfg(feature = "imports")]
737 Token::KwSelf => write!(f, "self"),
738 #[cfg(feature = "imports")]
739 Token::KwSuper => write!(f, "super"),
740 #[cfg(feature = "imports")]
741 Token::KwPackage => write!(f, "package"),
742 #[cfg(feature = "imports")]
743 Token::KwAs => write!(f, "as"),
744 #[cfg(feature = "imports")]
745 Token::KwImport => write!(f, "import"),
746 }
747 }
748}
749
750type Spanned<Tok, Loc, ParseError> = Result<(Loc, Tok, Loc), (Loc, ParseError, Loc)>;
751type NextToken = Option<(Result<Token, ParseError>, Span)>;
752
753#[derive(Clone)]
754pub struct Lexer<'s> {
755 source: &'s str,
756 token_stream: SpannedIter<'s, Token>,
757 next_token: NextToken,
758 recognizing_template: bool,
759 opened_templates: u32,
760}
761
762impl<'s> Lexer<'s> {
763 pub fn new(source: &'s str) -> Self {
764 let mut token_stream = Token::lexer_with_extras(source, LexerState::default()).spanned();
765 let next_token =
766 token_stream.find(|(tok, _)| tok.as_ref().is_ok_and(|tok| !tok.is_trivia()));
767
768 Self {
769 source,
770 token_stream,
771 next_token,
772 recognizing_template: false,
773 opened_templates: 0,
774 }
775 }
776
777 fn take_two_tokens(&mut self) -> (NextToken, NextToken) {
778 let mut tok1 = self.next_token.take();
779
780 let lookahead = self.token_stream.extras.lookahead.take();
781 let tok2 = match lookahead {
782 Some(tok) => {
783 let (_, span1) = tok1.as_mut().unwrap(); let span2 = span1.start + 1..span1.end;
785 Some((Ok(tok), span2))
786 }
787 None => self
788 .token_stream
789 .find(|(tok, _)| tok.as_ref().is_ok_and(|tok| !tok.is_trivia())),
790 };
791
792 (tok1, tok2)
793 }
794
795 fn next_token(&mut self) -> NextToken {
796 let (cur, mut next) = self.take_two_tokens();
797
798 let (cur_tok, cur_span) = match cur {
799 Some((Ok(tok), span)) => (tok, span),
800 Some((Err(e), span)) => return Some((Err(e), span)),
801 None => return None,
802 };
803
804 if let Some((Ok(next_tok), next_span)) = &mut next
805 && (matches!(cur_tok, Token::Ident(_)) || cur_tok.is_keyword())
806 && *next_tok == Token::SymLessThan
807 {
808 let source = &self.source[next_span.start..];
809 if recognize_template_list(source) {
810 *next_tok = Token::TemplateArgsStart;
811 let cur_depth = self.token_stream.extras.depth;
812 self.token_stream.extras.template_depths.push(cur_depth);
813 self.opened_templates += 1;
814 }
815 }
816
817 if self.recognizing_template && cur_tok == Token::TemplateArgsEnd {
819 self.opened_templates -= 1;
820 if self.opened_templates == 0 {
821 next = None; }
823 }
824
825 self.next_token = next;
826 Some((Ok(cur_tok), cur_span))
827 }
828}
829
830pub fn recognize_template_list(source: &str) -> bool {
842 let mut lexer = Lexer::new(source);
843 match lexer.next_token {
844 Some((Ok(ref mut t), _)) if *t == Token::SymLessThan => *t = Token::TemplateArgsStart,
845 _ => return false,
846 };
847 lexer.recognizing_template = true;
848 lexer.opened_templates = 1;
849 lexer.token_stream.extras.template_depths.push(0);
850 crate::parser::recognize_template_list(lexer).is_ok()
851}
852
853#[test]
854fn test_recognize_template() {
855 assert!(recognize_template_list("<i32,select(2,3,a>b)>"));
857 assert!(!recognize_template_list("<d]>"));
858 assert!(recognize_template_list("<B<<C>"));
859 assert!(recognize_template_list("<B<=C>"));
860 assert!(recognize_template_list("<(B>=C)>"));
861 assert!(recognize_template_list("<(B!=C)>"));
862 assert!(recognize_template_list("<(B==C)>"));
863 assert!(recognize_template_list("<X>"));
865 assert!(recognize_template_list("<X<Y>>"));
866 assert!(recognize_template_list("<X<Y<Z>>>"));
867 assert!(!recognize_template_list(""));
868 assert!(!recognize_template_list(""));
869 assert!(!recognize_template_list("<>"));
870 assert!(!recognize_template_list("<b || c>d"));
871}
872
873pub trait TokenIterator: IntoIterator<Item = Spanned<Token, usize, ParseError>> {}
874
875impl Iterator for Lexer<'_> {
876 type Item = Spanned<Token, usize, ParseError>;
877
878 fn next(&mut self) -> Option<Self::Item> {
879 let tok = self.next_token();
880 tok.map(|(tok, span)| match tok {
881 Ok(tok) => Ok((span.start, tok, span.end)),
882 Err(err) => Err((span.start, err, span.end)),
883 })
884 }
885}
886
887impl TokenIterator for Lexer<'_> {}