@@ -12,7 +12,7 @@ use crate::{
1212 extensions:: visit:: Scope , Def , EnumDef , FieldDef , OptionDef , Schema , StructDef , TypeDef ,
1313 TypeId , VecDef ,
1414 } ,
15- utils:: create_ident_tokens,
15+ utils:: { create_ident , create_ident_tokens} ,
1616 Codegen , Generator , Result , AST_CRATE_PATH ,
1717} ;
1818
@@ -27,12 +27,12 @@ define_generator!(VisitGenerator);
2727
2828impl Generator for VisitGenerator {
2929 /// Register that accept:
30- /// * `#[visit]` attr on struct fields or enum variants.
30+ /// * `#[visit]` attr on structs, struct fields, or enum variants.
3131 /// * `#[ast(visit)]` on structs or enums.
3232 /// * `#[scope]` on structs or struct fields.
3333 fn attrs ( & self ) -> & [ ( & ' static str , AttrPositions ) ] {
3434 & [
35- ( "visit" , attr_positions ! ( AstAttr | StructField | EnumVariant ) ) ,
35+ ( "visit" , attr_positions ! ( AstAttr | Struct | StructField | EnumVariant ) ) ,
3636 ( "scope" , attr_positions ! ( Struct | StructField ) ) ,
3737 ]
3838 }
@@ -86,6 +86,9 @@ fn parse_visit_attr(location: AttrLocation, part: AttrPart) -> Result<()> {
8686 }
8787
8888 match location {
89+ AttrLocation :: Struct ( struct_def) => {
90+ struct_def. visit . visit_args = Some ( args) ;
91+ }
8992 AttrLocation :: StructField ( struct_def, field_index) => {
9093 struct_def. fields [ field_index] . visit . visit_args = Some ( args) ;
9194 }
@@ -349,9 +352,16 @@ impl<'s> VisitBuilder<'s> {
349352 let visit_fn_ident = format_ident ! ( "visit_{type_snake_name}" ) ;
350353 let walk_fn_ident = format_ident ! ( "walk_{type_snake_name}" ) ;
351354
352- // TODO: Don't hard-code this. Represent it in an attr in AST type definition instead.
353- let ( extra_params, extra_args) = if struct_def. name ( ) == "Function" {
354- ( quote ! ( , flags: ScopeFlags ) , quote ! ( , flags ) )
355+ // Get additional params
356+ let ( extra_params, extra_args) = if let Some ( visit_args) = & struct_def. visit . visit_args {
357+ visit_args
358+ . iter ( )
359+ . map ( |( arg_name, arg_type_name) | {
360+ let param_ident = create_ident ( arg_name) ;
361+ let arg_type_ident = create_ident ( arg_type_name) ;
362+ ( quote ! ( , #param_ident: #arg_type_ident ) , quote ! ( , #param_ident ) )
363+ } )
364+ . unzip ( )
355365 } else {
356366 ( quote ! ( ) , quote ! ( ) )
357367 } ;
0 commit comments