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    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    // see blankspace and line breaks: https://www.w3.org/TR/WGSL/#blankspace-and-line-breaks
206    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// following the spec at this date: https://www.w3.org/TR/2024/WD-WGSL-20240731/
253#[derive(Logos, Clone, Debug, PartialEq)]
254#[logos(
255    // see blankspace and line breaks: https://www.w3.org/TR/WGSL/#blankspace-and-line-breaks
256    skip r"[\s\u0085\u200e\u200f\u2028\u2029]+", // blankspace
257    extras = LexerState,
258    error = ParseError)]
259pub enum Token {
260    // HACK: parsing entrypoints.
261    // See: https://github.com/lalrpop/lalrpop/issues/65#issuecomment-516769995
262    // The first token provided by the lexer to the parser must be one of these tokens.
263    // They tell the parser which syntax node to parse. The parser returns a union type
264    // containing the corresponding syntax node type.
265    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    // the parse_ident function can return either Token::Ident or Token::ReservedWord.
280    // Token::Ignored variant is never produced.
281    // It serves as a placeholder for running parse_ident.
282    #[regex(
283        r#"([_\p{XID_Start}][\p{XID_Continue}]+)|([\p{XID_Start}])"#,
284        parse_ident,
285        priority = 1
286    )]
287    Ignored,
288    // syntactic tokens
289    // https://www.w3.org/TR/WGSL/#syntactic-tokens
290    #[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    // keywords
384    // https://www.w3.org/TR/WGSL/#keyword-summary
385    #[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    // Idents and ReservedWord tokens are parsed on the Ignored variant, because of a current
439    // limitation of logos. See logos#295.
440    Ident(String),
441    // variant produced by parse_ident for reserved words.
442    // Reserved words can be used in context-dependent words, e.g. attribute names and module names.
443    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    // hex
454    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    // hex
458    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    // hex
462    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    // hex
479    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    // hex
484    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    // extension: wesl-imports
496    // https://github.com/wgsl-tooling-wg/wesl-spec/blob/imports-update/Imports.md
497    // date: 2025-01-18, hash: 2db8e7f681087db6bdcd4a254963deb5c0159775
498    #[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    /// This display implementation is used for error messages.
628    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(); // safety: lookahead implies lexer looked at a `<` token
780                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 we finished recognition of a template
814        if self.recognizing_template && cur_tok == Token::TemplateArgsEnd {
815            self.opened_templates -= 1;
816            if self.opened_templates == 0 {
817                next = None; // push eof after end of template
818            }
819        }
820
821        self.next_token = next;
822        Some((Ok(cur_tok), cur_span))
823    }
824}
825
826/// Returns `true` if the source starts with a valid template list.
827///
828/// ## Specification
829///
830/// [3.9. Template Lists](https://www.w3.org/TR/WGSL/#template-lists-sec)
831///
832/// Contrary to the specification [template list discovery algorithm], this function also
833/// checks that the template is syntactically valid (syntax: [*template_list*]).
834///
835/// [template list discovery algorigthm]: https://www.w3.org/TR/WGSL/#template-list-discovery
836/// [*template_list*]: https://www.w3.org/TR/WGSL/#syntax-template_list
837pub 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    // cases from the WGSL spec
852    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    // more cases
860    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<'_> {}