Skip to main content

wgsl_parse/
lexer.rs

1//! Prefer using [`crate::parse_str`]. You shouldn't need to manipulate the lexer.
2
3use 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 found a ">" on the same nesting level as the opening "<", it is a template end.
18        if lex.extras.depth == *depth {
19            lex.extras.template_depths.pop();
20            // if lookahead is GreaterThan, we may have a second closing template.
21            // note that >>= can never be (TemplateEnd, TemplateEnd, Equal).
22            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
39// operators && and || have lower precedence than < and >.
40// therefore, this is not a template: a < b || c > d
41fn 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
58// TODO: get rid of crate `lexical`
59
60// don't have to be super strict, the lexer regex already did the heavy lifting
61const DEC_FORMAT: u128 = lexical::NumberFormatBuilder::new().build_unchecked();
62
63// don't have to be super strict, the lexer regex already did the heavy lifting
64const 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    // TODO
140    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    // TODO
202    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    // see blankspace and line breaks: https://www.w3.org/TR/WGSL/#blankspace-and-line-breaks
210    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// following the spec at this date: https://www.w3.org/TR/2024/WD-WGSL-20240731/
257#[derive(Logos, Clone, Debug, PartialEq)]
258#[logos(
259    // see blankspace and line breaks: https://www.w3.org/TR/WGSL/#blankspace-and-line-breaks
260    skip r"[\s\u0085\u200e\u200f\u2028\u2029]+", // blankspace
261    extras = LexerState,
262    error = ParseError)]
263pub enum Token {
264    // HACK: parsing entrypoints.
265    // See: https://github.com/lalrpop/lalrpop/issues/65#issuecomment-516769995
266    // The first token provided by the lexer to the parser must be one of these tokens.
267    // They tell the parser which syntax node to parse. The parser returns a union type
268    // containing the corresponding syntax node type.
269    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    // the parse_ident function can return either Token::Ident or Token::ReservedWord.
284    // Token::Ignored variant is never produced.
285    // It serves as a placeholder for running parse_ident.
286    #[regex(
287        r#"([_\p{XID_Start}][\p{XID_Continue}]+)|([\p{XID_Start}])"#,
288        parse_ident,
289        priority = 1
290    )]
291    Ignored,
292    // syntactic tokens
293    // https://www.w3.org/TR/WGSL/#syntactic-tokens
294    #[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    // keywords
388    // https://www.w3.org/TR/WGSL/#keyword-summary
389    #[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    // Idents and ReservedWord tokens are parsed on the Ignored variant, because of a current
443    // limitation of logos. See logos#295.
444    Ident(String),
445    // variant produced by parse_ident for reserved words.
446    // Reserved words can be used in context-dependent words, e.g. attribute names and module names.
447    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    // hex
458    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    // hex
462    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    // hex
466    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    // hex
483    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    // hex
488    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    // extension: wesl-imports
500    // https://github.com/wgsl-tooling-wg/wesl-spec/blob/imports-update/Imports.md
501    // date: 2025-01-18, hash: 2db8e7f681087db6bdcd4a254963deb5c0159775
502    #[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    /// This display implementation is used for error messages.
632    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(); // safety: lookahead implies lexer looked at a `<` token
784                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 we finished recognition of a template
818        if self.recognizing_template && cur_tok == Token::TemplateArgsEnd {
819            self.opened_templates -= 1;
820            if self.opened_templates == 0 {
821                next = None; // push eof after end of template
822            }
823        }
824
825        self.next_token = next;
826        Some((Ok(cur_tok), cur_span))
827    }
828}
829
830/// Returns `true` if the source starts with a valid template list.
831///
832/// ## Specification
833///
834/// [3.9. Template Lists](https://www.w3.org/TR/WGSL/#template-lists-sec)
835///
836/// Contrary to the specification [template list discovery algorithm], this function also
837/// checks that the template is syntactically valid (syntax: [*template_list*]).
838///
839/// [template list discovery algorigthm]: https://www.w3.org/TR/WGSL/#template-list-discovery
840/// [*template_list*]: https://www.w3.org/TR/WGSL/#syntax-template_list
841pub 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    // cases from the WGSL spec
856    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    // more cases
864    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<'_> {}