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