@@ -629,30 +629,8 @@ fn create_field_getters<'a>(
629629 field : & ' a node_types:: Field ,
630630 nodes : & ' a node_types:: NodeTypeMap ,
631631) -> ( Vec < ql:: Predicate < ' a > > , Option < ql:: Expression < ' a > > ) {
632- let return_type = match & field. type_info {
633- node_types:: FieldTypeInfo :: Single ( t) => {
634- Some ( ql:: Type :: Facade ( & nodes. get ( t) . unwrap ( ) . ql_class_name ) )
635- }
636- node_types:: FieldTypeInfo :: Multiple {
637- types : _,
638- dbscheme_union : _,
639- ql_class,
640- } => Some ( ql:: Type :: Facade ( ql_class) ) ,
641- node_types:: FieldTypeInfo :: ReservedWordInt ( _) => Some ( ql:: Type :: String ) ,
642- } ;
643- let formal_parameters = match & field. storage {
644- node_types:: Storage :: Column { .. } => vec ! [ ] ,
645- node_types:: Storage :: Table { has_index, .. } => {
646- if * has_index {
647- vec ! [ ql:: FormalParameter {
648- name: "i" ,
649- param_type: ql:: Type :: Int ,
650- } ]
651- } else {
652- vec ! [ ]
653- }
654- }
655- } ;
632+ let return_type = field_getter_return_type ( field, nodes) ;
633+ let formal_parameters = field_getter_formal_parameters ( field) ;
656634
657635 // For the expression to get a value, what variable name should the result
658636 // be bound to?
@@ -742,16 +720,7 @@ fn create_field_getters<'a>(
742720 ( get_value, Some ( get_value_any_index) )
743721 }
744722 } ;
745- let qldoc = match & field. name {
746- Some ( name) => format ! ( "Gets the node corresponding to the field `{name}`." ) ,
747- None => {
748- if formal_parameters. is_empty ( ) {
749- "Gets the child of this node." . to_owned ( )
750- } else {
751- "Gets the `i`th child of this node." . to_owned ( )
752- }
753- }
754- } ;
723+ let qldoc = field_getter_qldoc ( field, !formal_parameters. is_empty ( ) ) ;
755724 let mut predicates = vec ! [ ql:: Predicate {
756725 qldoc: Some ( qldoc. clone( ) ) ,
757726 name: & field. getter_name,
@@ -766,7 +735,7 @@ fn create_field_getters<'a>(
766735
767736 if let Some ( any_getter_name) = & field. any_getter_name {
768737 predicates. push ( ql:: Predicate {
769- qldoc : Some ( qldoc. clone ( ) ) ,
738+ qldoc : Some ( qldoc) ,
770739 name : any_getter_name,
771740 overridden : false ,
772741 is_private : false ,
@@ -788,6 +757,86 @@ fn create_field_getters<'a>(
788757 ( predicates, optional_expr)
789758}
790759
760+ fn field_getter_return_type < ' a > (
761+ field : & ' a node_types:: Field ,
762+ nodes : & ' a node_types:: NodeTypeMap ,
763+ ) -> Option < ql:: Type < ' a > > {
764+ match & field. type_info {
765+ node_types:: FieldTypeInfo :: Single ( t) => {
766+ Some ( ql:: Type :: Facade ( & nodes. get ( t) . unwrap ( ) . ql_class_name ) )
767+ }
768+ node_types:: FieldTypeInfo :: Multiple {
769+ types : _,
770+ dbscheme_union : _,
771+ ql_class,
772+ } => Some ( ql:: Type :: Facade ( ql_class) ) ,
773+ node_types:: FieldTypeInfo :: ReservedWordInt ( _) => Some ( ql:: Type :: String ) ,
774+ }
775+ }
776+
777+ fn field_getter_formal_parameters ( field : & node_types:: Field ) -> Vec < ql:: FormalParameter < ' _ > > {
778+ match & field. storage {
779+ node_types:: Storage :: Column { .. } => vec ! [ ] ,
780+ node_types:: Storage :: Table { has_index, .. } => {
781+ if * has_index {
782+ vec ! [ ql:: FormalParameter {
783+ name: "i" ,
784+ param_type: ql:: Type :: Int ,
785+ } ]
786+ } else {
787+ vec ! [ ]
788+ }
789+ }
790+ }
791+ }
792+
793+ fn field_getter_qldoc ( field : & node_types:: Field , has_index : bool ) -> String {
794+ match & field. name {
795+ Some ( name) => format ! ( "Gets the node corresponding to the field `{name}`." ) ,
796+ None => {
797+ if has_index {
798+ "Gets the `i`th child of this node." . to_owned ( )
799+ } else {
800+ "Gets the child of this node." . to_owned ( )
801+ }
802+ }
803+ }
804+ }
805+
806+ fn create_supertype_field_getters < ' a > (
807+ field : & ' a node_types:: Field ,
808+ nodes : & ' a node_types:: NodeTypeMap ,
809+ ) -> Vec < ql:: Predicate < ' a > > {
810+ let return_type = field_getter_return_type ( field, nodes) ;
811+ let formal_parameters = field_getter_formal_parameters ( field) ;
812+ let qldoc = field_getter_qldoc ( field, !formal_parameters. is_empty ( ) ) ;
813+ let mut predicates = vec ! [ ql:: Predicate {
814+ qldoc: Some ( qldoc. clone( ) ) ,
815+ name: & field. getter_name,
816+ overridden: false ,
817+ is_private: false ,
818+ is_final: false ,
819+ return_type: return_type. clone( ) ,
820+ formal_parameters,
821+ body: ql:: Expression :: Pred ( "none" , vec![ ] ) ,
822+ overlay: None ,
823+ } ] ;
824+ if let Some ( any_getter_name) = & field. any_getter_name {
825+ predicates. push ( ql:: Predicate {
826+ qldoc : Some ( qldoc) ,
827+ name : any_getter_name,
828+ overridden : false ,
829+ is_private : false ,
830+ is_final : false ,
831+ return_type,
832+ formal_parameters : vec ! [ ] ,
833+ body : ql:: Expression :: Pred ( "none" , vec ! [ ] ) ,
834+ overlay : None ,
835+ } ) ;
836+ }
837+ predicates
838+ }
839+
791840fn compute_direct_supertypes (
792841 nodes : & node_types:: NodeTypeMap ,
793842) -> std:: collections:: BTreeMap < node_types:: TypeName , BTreeSet < & str > > {
@@ -894,12 +943,34 @@ fn compute_exposed_predicate_signatures(
894943 exposed
895944}
896945
946+ fn is_predicate_inherited (
947+ type_name : & node_types:: TypeName ,
948+ predicate : & ql:: Predicate ,
949+ nodes : & node_types:: NodeTypeMap ,
950+ exposed : & std:: collections:: BTreeMap < node_types:: TypeName , BTreeSet < PredicateSignature < ' _ > > > ,
951+ ) -> bool {
952+ let signature = PredicateSignature {
953+ name : predicate. name ,
954+ arity : predicate. formal_parameters . len ( ) ,
955+ } ;
956+ nodes. iter ( ) . any ( |( supertype, node) | {
957+ matches ! (
958+ & node. kind,
959+ node_types:: EntryKind :: Union { members, .. }
960+ if members. contains( type_name)
961+ && exposed
962+ . get( supertype)
963+ . is_some_and( |signatures| signatures. contains( & signature) )
964+ )
965+ } )
966+ }
967+
897968/// Converts the given node types into CodeQL classes wrapping the dbscheme.
898969pub fn convert_nodes ( nodes : & node_types:: NodeTypeMap ) -> Vec < ql:: TopLevel < ' _ > > {
899970 let mut classes = Vec :: new ( ) ;
900971 let mut token_kinds = BTreeSet :: new ( ) ;
901972 let direct_supertypes = compute_direct_supertypes ( nodes) ;
902- let _exposed_predicate_signatures = compute_exposed_predicate_signatures ( nodes) ;
973+ let exposed_predicate_signatures = compute_exposed_predicate_signatures ( nodes) ;
903974 for ( type_name, node) in nodes {
904975 if let node_types:: EntryKind :: Token { .. } = & node. kind
905976 && type_name. named
@@ -930,9 +1001,22 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec<ql::TopLevel<'_>> {
9301001 } ) ) ;
9311002 }
9321003 }
933- node_types:: EntryKind :: Union { .. } => {
1004+ node_types:: EntryKind :: Union { fields , .. } => {
9341005 // It's a tree-sitter supertype node, so we're wrapping a dbscheme
9351006 // union type.
1007+ let predicates = fields
1008+ . iter ( )
1009+ . flat_map ( |field| create_supertype_field_getters ( field, nodes) )
1010+ . map ( |mut predicate| {
1011+ predicate. overridden = is_predicate_inherited (
1012+ type_name,
1013+ & predicate,
1014+ nodes,
1015+ & exposed_predicate_signatures,
1016+ ) ;
1017+ predicate
1018+ } )
1019+ . collect ( ) ;
9361020 classes. push ( ql:: TopLevel :: Class ( ql:: Class {
9371021 qldoc : None ,
9381022 name : & node. ql_class_name ,
@@ -946,7 +1030,7 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec<ql::TopLevel<'_>> {
9461030 & direct_supertypes,
9471031 ) ,
9481032 characteristic_predicate : None ,
949- predicates : vec ! [ ] ,
1033+ predicates,
9501034 } ) ) ;
9511035 }
9521036 node_types:: EntryKind :: Table {
@@ -990,13 +1074,21 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec<ql::TopLevel<'_>> {
9901074 // - predicates to access the fields,
9911075 // - the QL expressions to access the fields that will be part of getAFieldOrChild.
9921076 for field in fields {
993- let ( get_preds, get_child_expr) = create_field_getters (
1077+ let ( mut get_preds, get_child_expr) = create_field_getters (
9941078 main_table_name,
9951079 main_table_arity,
9961080 & mut main_table_column_index,
9971081 field,
9981082 nodes,
9991083 ) ;
1084+ for predicate in & mut get_preds {
1085+ predicate. overridden = is_predicate_inherited (
1086+ type_name,
1087+ predicate,
1088+ nodes,
1089+ & exposed_predicate_signatures,
1090+ ) ;
1091+ }
10001092 main_class. predicates . extend ( get_preds) ;
10011093 if let Some ( get_child_expr) = get_child_expr {
10021094 get_child_exprs. push ( get_child_expr)
@@ -1172,4 +1264,45 @@ named:
11721264 arity: 0 ,
11731265 } ) ) ;
11741266 }
1267+
1268+ #[ test]
1269+ fn generates_supertype_getters_and_concrete_overrides ( ) {
1270+ let yaml = r#"
1271+ supertypes:
1272+ callable:
1273+ subtypes: [function]
1274+ fields:
1275+ parameter*: parameter
1276+ body?: block
1277+ named:
1278+ function:
1279+ parameter*: parameter
1280+ body?: block
1281+ parameter:
1282+ block:
1283+ "# ;
1284+ let json = yeast:: node_types_yaml:: convert ( yaml) . unwrap ( ) ;
1285+ let nodes = node_types:: read_node_types_str ( "test" , & json) . unwrap ( ) ;
1286+ let generated = convert_nodes ( & nodes)
1287+ . into_iter ( )
1288+ . map ( |element| element. to_string ( ) )
1289+ . collect :: < Vec < _ > > ( )
1290+ . join ( "\n " ) ;
1291+
1292+ assert ! ( generated. contains( "F::Parameter getParameter(int i) { none() }" ) ) ;
1293+ assert ! ( generated. contains( "F::Parameter getAParameter() { none() }" ) ) ;
1294+ assert ! ( generated. contains( "F::Block getBody() { none() }" ) ) ;
1295+ assert ! (
1296+ generated. contains(
1297+ "final override F::Parameter getParameter(int i) { test_function_parameter"
1298+ )
1299+ ) ;
1300+ assert ! ( generated. contains(
1301+ "final override F::Parameter getAParameter() { result = this.getParameter(_) }"
1302+ ) ) ;
1303+ assert ! (
1304+ generated
1305+ . contains( "final override F::Block getBody() { test_function_body(this, result) }" )
1306+ ) ;
1307+ }
11751308}
0 commit comments