@@ -770,10 +770,51 @@ fn create_field_getters<'a>(
770770 )
771771}
772772
773+ fn compute_direct_supertypes < ' a > (
774+ nodes : & ' a node_types:: NodeTypeMap ,
775+ ) -> std:: collections:: BTreeMap < node_types:: TypeName , BTreeSet < & ' a str > > {
776+ let mut supertypes = std:: collections:: BTreeMap :: new ( ) ;
777+ for node in nodes. values ( ) {
778+ if let node_types:: EntryKind :: Union { members } = & node. kind {
779+ for member in members {
780+ supertypes
781+ . entry ( member. clone ( ) )
782+ . or_insert_with ( BTreeSet :: new)
783+ . insert ( node. ql_class_name . as_str ( ) ) ;
784+ }
785+ }
786+ }
787+ supertypes
788+ }
789+
790+ fn ast_base_types < ' a > (
791+ type_name : & node_types:: TypeName ,
792+ direct_supertypes : & std:: collections:: BTreeMap < node_types:: TypeName , BTreeSet < & ' a str > > ,
793+ ) -> BTreeSet < ql:: Type < ' a > > {
794+ match direct_supertypes. get ( type_name) {
795+ Some ( supertypes) if !supertypes. is_empty ( ) => supertypes
796+ . iter ( )
797+ . map ( |name| ql:: Type :: Normal ( name) )
798+ . collect ( ) ,
799+ _ => vec ! [ ql:: Type :: Normal ( "AstNode" ) ] . into_iter ( ) . collect ( ) ,
800+ }
801+ }
802+
803+ fn class_supertypes < ' a > (
804+ type_name : & node_types:: TypeName ,
805+ dbscheme_name : & ' a str ,
806+ direct_supertypes : & std:: collections:: BTreeMap < node_types:: TypeName , BTreeSet < & ' a str > > ,
807+ ) -> BTreeSet < ql:: Type < ' a > > {
808+ let mut supertypes = ast_base_types ( type_name, direct_supertypes) ;
809+ supertypes. insert ( ql:: Type :: At ( dbscheme_name) ) ;
810+ supertypes
811+ }
812+
773813/// Converts the given node types into CodeQL classes wrapping the dbscheme.
774814pub fn convert_nodes ( nodes : & node_types:: NodeTypeMap ) -> Vec < ql:: TopLevel < ' _ > > {
775815 let mut classes = Vec :: new ( ) ;
776816 let mut token_kinds = BTreeSet :: new ( ) ;
817+ let direct_supertypes = compute_direct_supertypes ( nodes) ;
777818 for ( type_name, node) in nodes {
778819 if let node_types:: EntryKind :: Token { .. } = & node. kind {
779820 if type_name. named {
@@ -788,8 +829,8 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec<ql::TopLevel<'_>> {
788829 if type_name. named {
789830 let get_a_primary_ql_class =
790831 create_get_a_primary_ql_class ( & node. ql_class_name , true ) ;
791- let mut supertypes: BTreeSet < ql :: Type > = BTreeSet :: new ( ) ;
792- supertypes . insert ( ql :: Type :: At ( & node. dbscheme_name ) ) ;
832+ let mut supertypes =
833+ class_supertypes ( type_name , & node. dbscheme_name , & direct_supertypes ) ;
793834 supertypes. insert ( ql:: Type :: Normal ( "Token" ) ) ;
794835 classes. push ( ql:: TopLevel :: Class ( ql:: Class {
795836 qldoc : Some ( format ! ( "A class representing `{}` tokens." , type_name. kind) ) ,
@@ -814,12 +855,11 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec<ql::TopLevel<'_>> {
814855 is_final : false ,
815856 is_private : false ,
816857 alias : None ,
817- supertypes : vec ! [
818- ql:: Type :: At ( & node. dbscheme_name) ,
819- ql:: Type :: Normal ( "AstNode" ) ,
820- ]
821- . into_iter ( )
822- . collect ( ) ,
858+ supertypes : class_supertypes (
859+ type_name,
860+ & node. dbscheme_name ,
861+ & direct_supertypes,
862+ ) ,
823863 characteristic_predicate : None ,
824864 predicates : vec ! [ ] ,
825865 } ) ) ;
@@ -848,12 +888,11 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec<ql::TopLevel<'_>> {
848888 is_final : false ,
849889 is_private : false ,
850890 alias : None ,
851- supertypes : vec ! [
852- ql:: Type :: At ( & node. dbscheme_name) ,
853- ql:: Type :: Normal ( "AstNode" ) ,
854- ]
855- . into_iter ( )
856- . collect ( ) ,
891+ supertypes : class_supertypes (
892+ type_name,
893+ & node. dbscheme_name ,
894+ & direct_supertypes,
895+ ) ,
857896 characteristic_predicate : None ,
858897 predicates : vec ! [ create_get_a_primary_ql_class( main_class_name, true ) ] ,
859898 } ;
0 commit comments