diff --git a/shared/tree-sitter-extractor/src/extractor/mod.rs b/shared/tree-sitter-extractor/src/extractor/mod.rs index 13ee3264133b..06505f9b1fc1 100644 --- a/shared/tree-sitter-extractor/src/extractor/mod.rs +++ b/shared/tree-sitter-extractor/src/extractor/mod.rs @@ -904,7 +904,8 @@ impl<'a> Visitor<'a> { if tp == single_type { return true; } - if let EntryKind::Union { members } = &self.schema.get(single_type).unwrap().kind + if let EntryKind::Union { members, .. } = + &self.schema.get(single_type).unwrap().kind && self.type_matches_set(tp, members) { return true; @@ -926,7 +927,7 @@ impl<'a> Visitor<'a> { return true; } for other in types.iter() { - if let EntryKind::Union { members } = &self.schema.get(other).unwrap().kind + if let EntryKind::Union { members, .. } = &self.schema.get(other).unwrap().kind && self.type_matches_set(tp, members) { return true; diff --git a/shared/tree-sitter-extractor/src/generator/mod.rs b/shared/tree-sitter-extractor/src/generator/mod.rs index ea90a8482894..81bf0955cd6e 100644 --- a/shared/tree-sitter-extractor/src/generator/mod.rs +++ b/shared/tree-sitter-extractor/src/generator/mod.rs @@ -399,7 +399,9 @@ fn convert_nodes( .collect(); for node in nodes.values() { match &node.kind { - node_types::EntryKind::Union { members: n_members } => { + node_types::EntryKind::Union { + members: n_members, .. + } => { // It's a tree-sitter supertype node, for which we create a union // type. let members: Set<&str> = n_members diff --git a/shared/tree-sitter-extractor/src/generator/ql.rs b/shared/tree-sitter-extractor/src/generator/ql.rs index 19419a8465d0..c780aec681cd 100644 --- a/shared/tree-sitter-extractor/src/generator/ql.rs +++ b/shared/tree-sitter-extractor/src/generator/ql.rs @@ -109,7 +109,7 @@ impl fmt::Display for Class<'_> { is_final: false, return_type: None, formal_parameters: vec![], - body: Some(charpred.clone()), + body: charpred.clone(), overlay: None, } )?; @@ -307,9 +307,7 @@ pub struct Predicate<'a> { pub is_final: bool, pub return_type: Option>, pub formal_parameters: Vec>, - /// The body of the predicate, or `None` if this is an `abstract` - /// predicate declaration with no body. - pub body: Option>, + pub body: Expression<'a>, pub overlay: Option, } @@ -332,9 +330,6 @@ impl fmt::Display for Predicate<'_> { if self.is_final { write!(f, "final ")?; } - if self.body.is_none() { - write!(f, "abstract ")?; - } if self.overridden { write!(f, "override ")?; } @@ -349,10 +344,7 @@ impl fmt::Display for Predicate<'_> { } write!(f, "{param}")?; } - match &self.body { - Some(body) => write!(f, ") {{ {body} }}")?, - None => write!(f, ");")?, - } + write!(f, ") {{ {} }}", self.body)?; Ok(()) } diff --git a/shared/tree-sitter-extractor/src/generator/ql_gen.rs b/shared/tree-sitter-extractor/src/generator/ql_gen.rs index b6f3d45f4b12..bd05de8316d8 100644 --- a/shared/tree-sitter-extractor/src/generator/ql_gen.rs +++ b/shared/tree-sitter-extractor/src/generator/ql_gen.rs @@ -1,4 +1,3 @@ -use std::collections::BTreeMap; use std::collections::BTreeSet; use crate::{generator::ql, node_types}; @@ -21,14 +20,14 @@ pub fn create_ast_node_class<'a>( is_final: false, return_type: Some(ql::Type::String), formal_parameters: vec![], - body: Some(ql::Expression::Equals( + body: ql::Expression::Equals( Box::new(ql::Expression::Var("result")), Box::new(ql::Expression::Dot( Box::new(ql::Expression::Var("this")), "getAPrimaryQlClass", vec![], )), - )), + ), overlay: None, }; let get_location = ql::Predicate { @@ -39,10 +38,10 @@ pub fn create_ast_node_class<'a>( is_final: true, return_type: Some(ql::Type::Normal("L::Location")), formal_parameters: vec![], - body: Some(ql::Expression::Pred( + body: ql::Expression::Pred( node_location_table, vec![ql::Expression::Var("this"), ql::Expression::Var("result")], - )), + ), overlay: None, }; let get_a_field_or_child = create_none_predicate( @@ -59,14 +58,14 @@ pub fn create_ast_node_class<'a>( is_final: true, return_type: Some(ql::Type::Facade("AstNode")), formal_parameters: vec![], - body: Some(ql::Expression::Pred( + body: ql::Expression::Pred( node_parent_table, vec![ ql::Expression::Var("this"), ql::Expression::Var("result"), ql::Expression::Var("_"), ], - )), + ), overlay: None, }; let get_parent_index = ql::Predicate { @@ -79,14 +78,14 @@ pub fn create_ast_node_class<'a>( is_final: true, return_type: Some(ql::Type::Int), formal_parameters: vec![], - body: Some(ql::Expression::Pred( + body: ql::Expression::Pred( node_parent_table, vec![ ql::Expression::Var("this"), ql::Expression::Var("_"), ql::Expression::Var("result"), ], - )), + ), overlay: None, }; let get_a_primary_ql_class = ql::Predicate { @@ -99,10 +98,10 @@ pub fn create_ast_node_class<'a>( is_final: false, return_type: Some(ql::Type::String), formal_parameters: vec![], - body: Some(ql::Expression::Equals( + body: ql::Expression::Equals( Box::new(ql::Expression::Var("result")), Box::new(ql::Expression::String("???")), - )), + ), overlay: None, }; let get_primary_ql_classes = ql::Predicate { @@ -117,7 +116,7 @@ pub fn create_ast_node_class<'a>( is_final: false, return_type: Some(ql::Type::String), formal_parameters: vec![], - body: Some(ql::Expression::Equals( + body: ql::Expression::Equals( Box::new(ql::Expression::Var("result")), Box::new(ql::Expression::Aggregate { name: "concat", @@ -130,7 +129,7 @@ pub fn create_ast_node_class<'a>( )), second_expr: Some(Box::new(ql::Expression::String(","))), }), - )), + ), overlay: None, }; ql::Class { @@ -164,12 +163,7 @@ pub fn create_token_class<'a>(token_type: &'a str, tokeninfo: &'a str) -> ql::Cl is_final: true, return_type: Some(ql::Type::String), formal_parameters: vec![], - body: Some(create_get_field_expr_for_column_storage( - "result", - tokeninfo, - 1, - tokeninfo_arity, - )), + body: create_get_field_expr_for_column_storage("result", tokeninfo, 1, tokeninfo_arity), overlay: None, }; let to_string = ql::Predicate { @@ -182,14 +176,14 @@ pub fn create_token_class<'a>(token_type: &'a str, tokeninfo: &'a str) -> ql::Cl is_final: true, return_type: Some(ql::Type::String), formal_parameters: vec![], - body: Some(ql::Expression::Equals( + body: ql::Expression::Equals( Box::new(ql::Expression::Var("result")), Box::new(ql::Expression::Dot( Box::new(ql::Expression::Var("this")), "getValue", vec![], )), - )), + ), overlay: None, }; ql::Class { @@ -229,12 +223,12 @@ pub fn create_trivia_token_class<'a>( is_final: true, return_type: Some(ql::Type::String), formal_parameters: vec![], - body: Some(create_get_field_expr_for_column_storage( + body: create_get_field_expr_for_column_storage( "result", trivia_tokeninfo, 1, trivia_tokeninfo_arity, - )), + ), overlay: None, }; let to_string = ql::Predicate { @@ -247,14 +241,14 @@ pub fn create_trivia_token_class<'a>( is_final: true, return_type: Some(ql::Type::String), formal_parameters: vec![], - body: Some(ql::Expression::Equals( + body: ql::Expression::Equals( Box::new(ql::Expression::Var("result")), Box::new(ql::Expression::Dot( Box::new(ql::Expression::Var("this")), "getValue", vec![], )), - )), + ), overlay: None, }; ql::Class { @@ -312,7 +306,7 @@ fn create_none_predicate<'a>( is_final: false, return_type, formal_parameters: Vec::new(), - body: Some(ql::Expression::Pred("none", vec![])), + body: ql::Expression::Pred("none", vec![]), overlay: None, } } @@ -330,10 +324,10 @@ fn create_get_a_primary_ql_class(class_name: &str, is_final: bool) -> ql::Predic is_final, return_type: Some(ql::Type::String), formal_parameters: vec![], - body: Some(ql::Expression::Equals( + body: ql::Expression::Equals( Box::new(ql::Expression::Var("result")), Box::new(ql::Expression::String(class_name)), - )), + ), overlay: None, } } @@ -348,13 +342,13 @@ pub fn create_is_overlay_predicate() -> ql::Predicate<'static> { return_type: None, overlay: Some(ql::OverlayAnnotation::Local), formal_parameters: vec![], - body: Some(ql::Expression::Pred( + body: ql::Expression::Pred( "databaseMetadata", vec![ ql::Expression::String("isOverlay"), ql::Expression::String("true"), ], - )), + ), } } @@ -374,7 +368,7 @@ pub fn create_get_node_file_predicate<'a>( name: "node", param_type: ql::Type::At(ast_node_name), }], - body: Some(ql::Expression::Aggregate { + body: ql::Expression::Aggregate { name: "exists", vars: vec![ql::FormalParameter { name: "loc", @@ -396,7 +390,7 @@ pub fn create_get_node_file_predicate<'a>( ], )), second_expr: None, - }), + }, } } @@ -421,7 +415,7 @@ pub fn create_discardable_ast_node_predicate(ast_node_name: &str) -> ql::Predica param_type: ql::Type::At(ast_node_name), }, ], - body: Some(ql::Expression::And(vec![ + body: ql::Expression::And(vec![ ql::Expression::Negation(Box::new(ql::Expression::Pred("isOverlay", vec![]))), ql::Expression::Equals( Box::new(ql::Expression::Var("file")), @@ -430,7 +424,7 @@ pub fn create_discardable_ast_node_predicate(ast_node_name: &str) -> ql::Predica vec![ql::Expression::Var("node")], )), ), - ])), + ]), } } @@ -450,7 +444,7 @@ pub fn create_discard_ast_node_predicate(ast_node_name: &str) -> ql::Predicate<' name: "node", param_type: ql::Type::At(ast_node_name), }], - body: Some(ql::Expression::Aggregate { + body: ql::Expression::Aggregate { name: "exists", vars: vec![ ql::FormalParameter { @@ -474,7 +468,7 @@ pub fn create_discard_ast_node_predicate(ast_node_name: &str) -> ql::Predicate<' ql::Expression::Pred("overlayChangedFiles", vec![ql::Expression::Var("path")]), ])), second_expr: None, - }), + }, } } @@ -499,7 +493,7 @@ pub fn create_discardable_location_predicate() -> ql::Predicate<'static> { param_type: ql::Type::At("location_default"), }, ], - body: Some(ql::Expression::And(vec![ + body: ql::Expression::And(vec![ ql::Expression::Negation(Box::new(ql::Expression::Pred("isOverlay", vec![]))), ql::Expression::Pred( "locations_default", @@ -512,7 +506,7 @@ pub fn create_discardable_location_predicate() -> ql::Predicate<'static> { ql::Expression::Var("_"), ], ), - ])), + ]), } } @@ -535,7 +529,7 @@ pub fn create_discard_location_predicate() -> ql::Predicate<'static> { name: "loc", param_type: ql::Type::At("location_default"), }], - body: Some(ql::Expression::Aggregate { + body: ql::Expression::Aggregate { name: "exists", vars: vec![ ql::FormalParameter { @@ -559,7 +553,7 @@ pub fn create_discard_location_predicate() -> ql::Predicate<'static> { ql::Expression::Pred("overlayChangedFiles", vec![ql::Expression::Var("path")]), ])), second_expr: None, - }), + }, } } @@ -635,30 +629,8 @@ fn create_field_getters<'a>( field: &'a node_types::Field, nodes: &'a node_types::NodeTypeMap, ) -> (Vec>, Option>) { - let return_type = match &field.type_info { - node_types::FieldTypeInfo::Single(t) => { - Some(ql::Type::Facade(&nodes.get(t).unwrap().ql_class_name)) - } - node_types::FieldTypeInfo::Multiple { - types: _, - dbscheme_union: _, - ql_class, - } => Some(ql::Type::Facade(ql_class)), - node_types::FieldTypeInfo::ReservedWordInt(_) => Some(ql::Type::String), - }; - let formal_parameters = match &field.storage { - node_types::Storage::Column { .. } => vec![], - node_types::Storage::Table { has_index, .. } => { - if *has_index { - vec![ql::FormalParameter { - name: "i", - param_type: ql::Type::Int, - }] - } else { - vec![] - } - } - }; + let return_type = field_getter_return_type(field, nodes); + let formal_parameters = field_getter_formal_parameters(field); // For the expression to get a value, what variable name should the result // be bound to? @@ -748,16 +720,7 @@ fn create_field_getters<'a>( (get_value, Some(get_value_any_index)) } }; - let qldoc = match &field.name { - Some(name) => format!("Gets the node corresponding to the field `{name}`."), - None => { - if formal_parameters.is_empty() { - "Gets the child of this node.".to_owned() - } else { - "Gets the `i`th child of this node.".to_owned() - } - } - }; + let qldoc = field_getter_qldoc(field, !formal_parameters.is_empty()); let mut predicates = vec![ql::Predicate { qldoc: Some(qldoc.clone()), name: &field.getter_name, @@ -766,27 +729,27 @@ fn create_field_getters<'a>( is_final: true, return_type: return_type.clone(), formal_parameters, - body: Some(body), + body, overlay: None, }]; if let Some(any_getter_name) = &field.any_getter_name { predicates.push(ql::Predicate { - qldoc: Some(qldoc.clone()), + qldoc: Some(qldoc), name: any_getter_name, overridden: false, is_private: false, is_final: true, return_type, formal_parameters: vec![], - body: Some(ql::Expression::Equals( + body: ql::Expression::Equals( Box::new(ql::Expression::Var("result")), Box::new(ql::Expression::Dot( Box::new(ql::Expression::Var("this")), &field.getter_name, vec![ql::Expression::Var("_")], )), - )), + ), overlay: None, }); } @@ -794,12 +757,92 @@ fn create_field_getters<'a>( (predicates, optional_expr) } +fn field_getter_return_type<'a>( + field: &'a node_types::Field, + nodes: &'a node_types::NodeTypeMap, +) -> Option> { + match &field.type_info { + node_types::FieldTypeInfo::Single(t) => { + Some(ql::Type::Facade(&nodes.get(t).unwrap().ql_class_name)) + } + node_types::FieldTypeInfo::Multiple { + types: _, + dbscheme_union: _, + ql_class, + } => Some(ql::Type::Facade(ql_class)), + node_types::FieldTypeInfo::ReservedWordInt(_) => Some(ql::Type::String), + } +} + +fn field_getter_formal_parameters(field: &node_types::Field) -> Vec> { + match &field.storage { + node_types::Storage::Column { .. } => vec![], + node_types::Storage::Table { has_index, .. } => { + if *has_index { + vec![ql::FormalParameter { + name: "i", + param_type: ql::Type::Int, + }] + } else { + vec![] + } + } + } +} + +fn field_getter_qldoc(field: &node_types::Field, has_index: bool) -> String { + match &field.name { + Some(name) => format!("Gets the node corresponding to the field `{name}`."), + None => { + if has_index { + "Gets the `i`th child of this node.".to_owned() + } else { + "Gets the child of this node.".to_owned() + } + } + } +} + +fn create_supertype_field_getters<'a>( + field: &'a node_types::Field, + nodes: &'a node_types::NodeTypeMap, +) -> Vec> { + let return_type = field_getter_return_type(field, nodes); + let formal_parameters = field_getter_formal_parameters(field); + let qldoc = field_getter_qldoc(field, !formal_parameters.is_empty()); + let mut predicates = vec![ql::Predicate { + qldoc: Some(qldoc.clone()), + name: &field.getter_name, + overridden: false, + is_private: false, + is_final: false, + return_type: return_type.clone(), + formal_parameters, + body: ql::Expression::Pred("none", vec![]), + overlay: None, + }]; + if let Some(any_getter_name) = &field.any_getter_name { + predicates.push(ql::Predicate { + qldoc: Some(qldoc), + name: any_getter_name, + overridden: false, + is_private: false, + is_final: false, + return_type, + formal_parameters: vec![], + body: ql::Expression::Pred("none", vec![]), + overlay: None, + }); + } + predicates +} + fn compute_direct_supertypes( nodes: &node_types::NodeTypeMap, ) -> std::collections::BTreeMap> { let mut supertypes = std::collections::BTreeMap::new(); for node in nodes.values() { - if let node_types::EntryKind::Union { members } = &node.kind { + if let node_types::EntryKind::Union { members, .. } = &node.kind { for member in members { supertypes .entry(member.clone()) @@ -834,86 +877,91 @@ fn class_supertypes<'a>( supertypes } -/// Returns whether `a` and `b` have the same signature, i.e. the same name, -/// return type, and formal parameters. Predicates with the same signature can -/// override one another. -fn same_predicate_signature(a: &ql::Predicate, b: &ql::Predicate) -> bool { - a.name == b.name && a.return_type == b.return_type && a.formal_parameters == b.formal_parameters +#[derive(Clone, Copy, Debug, Eq, Ord, PartialEq, PartialOrd)] +struct PredicateSignature<'a> { + name: &'a str, + arity: usize, } -/// Computes, for each tree-sitter supertype (union) node, the list of -/// predicates that are guaranteed to be defined identically (in terms of -/// name, return type, and formal parameters, though not necessarily body) by -/// every one of its members. These are the predicates that can be hoisted to -/// an `abstract` predicate on the union's class, with the corresponding -/// predicates on its members becoming `override`s. -/// -/// The result for a given node is memoized in `cache` (keyed by its QL class -/// name), and also used to answer the query for any other node that -/// (directly, or transitively through further supertypes) has that node as a -/// member. The same cache also serves as the answer to "what does the class -/// named X expose?", used by `is_predicate_inherited`. -fn compute_exposed_predicates<'a, 'b>( - type_name: &'a node_types::TypeName, - nodes: &'a node_types::NodeTypeMap, - field_predicates: &BTreeMap<&node_types::TypeName, Vec>>, - cache: &'b mut BTreeMap<&'a str, Vec>>, -) -> &'b Vec> { - let node = nodes.get(type_name); - let class_name = node.map_or(type_name.kind.as_str(), |node| node.ql_class_name.as_str()); - if !cache.contains_key(class_name) { - // Supertype declarations that recursively refer to themselves are a mistake, but we don't - // want to cause infinite recursion, so we insert a temporary sentinel. - cache.insert(class_name, Vec::new()); - let exposed = match node.map(|node| &node.kind) { - Some(node_types::EntryKind::Table { .. }) => { - field_predicates.get(type_name).cloned().unwrap_or_default() - } - Some(node_types::EntryKind::Union { members }) => { - let mut members = members.iter(); - let mut common = match members.next() { - Some(first) => { - compute_exposed_predicates(first, nodes, field_predicates, cache).clone() - } - None => Vec::new(), - }; - for member in members { - let member_predicates = - compute_exposed_predicates(member, nodes, field_predicates, cache); - common.retain(|predicate| { - member_predicates - .iter() - .any(|other| same_predicate_signature(predicate, other)) - }); - } - common +fn field_predicate_signatures(field: &node_types::Field) -> Vec> { + let getter_arity = match field.storage { + node_types::Storage::Table { + has_index: true, .. + } => 1, + _ => 0, + }; + let mut signatures = vec![PredicateSignature { + name: &field.getter_name, + arity: getter_arity, + }]; + if let Some(any_getter_name) = &field.any_getter_name { + signatures.push(PredicateSignature { + name: any_getter_name, + arity: 0, + }); + } + signatures +} + +/// Builds an index of the field getter signatures exposed by each generated +/// class, including getters inherited from supertypes. +fn compute_exposed_predicate_signatures( + nodes: &node_types::NodeTypeMap, +) -> std::collections::BTreeMap>> { + let mut exposed = nodes + .iter() + .map(|(type_name, node)| { + let fields = match &node.kind { + node_types::EntryKind::Union { fields, .. } + | node_types::EntryKind::Table { fields, .. } => fields.as_slice(), + node_types::EntryKind::Token { .. } => &[], + }; + let signatures = fields.iter().flat_map(field_predicate_signatures).collect(); + (type_name.clone(), signatures) + }) + .collect::>>(); + + loop { + let mut changed = false; + for (supertype, node) in nodes { + let node_types::EntryKind::Union { members, .. } = &node.kind else { + continue; + }; + let inherited = exposed.get(supertype).cloned().unwrap_or_default(); + for member in members { + let member_signatures = exposed.entry(member.clone()).or_default(); + let previous_len = member_signatures.len(); + member_signatures.extend(inherited.iter().copied()); + changed |= member_signatures.len() != previous_len; } - Some(node_types::EntryKind::Token { .. }) | None => Vec::new(), - }; - cache.insert(class_name, exposed); + } + if !changed { + break; + } } - cache.get(class_name).unwrap() + + exposed } -/// Returns whether `predicate` (declared, or about to be declared, on the -/// class for `type_name`) is already exposed by one of `type_name`'s direct -/// supertypes, and therefore must be marked as an `override` (for a concrete -/// predicate) or can be omitted entirely (for an `abstract` one, since it's -/// already inherited). fn is_predicate_inherited( - predicate: &ql::Predicate, type_name: &node_types::TypeName, - direct_supertypes: &BTreeMap>, - exposed_predicates: &BTreeMap<&str, Vec>, + predicate: &ql::Predicate, + nodes: &node_types::NodeTypeMap, + exposed: &std::collections::BTreeMap>>, ) -> bool { - direct_supertypes.get(type_name).is_some_and(|supertypes| { - supertypes.iter().any(|supertype| { - exposed_predicates.get(supertype).is_some_and(|predicates| { - predicates - .iter() - .any(|other| same_predicate_signature(predicate, other)) - }) - }) + let signature = PredicateSignature { + name: predicate.name, + arity: predicate.formal_parameters.len(), + }; + nodes.iter().any(|(supertype, node)| { + matches!( + &node.kind, + node_types::EntryKind::Union { members, .. } + if members.contains(type_name) + && exposed + .get(supertype) + .is_some_and(|signatures| signatures.contains(&signature)) + ) }) } @@ -922,6 +970,7 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec> { let mut classes = Vec::new(); let mut token_kinds = BTreeSet::new(); let direct_supertypes = compute_direct_supertypes(nodes); + let exposed_predicate_signatures = compute_exposed_predicate_signatures(nodes); for (type_name, node) in nodes { if let node_types::EntryKind::Token { .. } = &node.kind && type_name.named @@ -930,71 +979,6 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec> { } } - // First, compute the field-getter predicates (and the expressions used by - // `getAFieldOrChild`) for every table node, without yet knowing whether - // any of them will need to be marked `override`. These are needed both - // to build the final classes below, and to figure out which fields are - // shared identically by all the members of a supertype. - let mut field_predicates: BTreeMap<&node_types::TypeName, Vec>> = - BTreeMap::new(); - let mut get_child_exprs: BTreeMap<&node_types::TypeName, Vec>> = - BTreeMap::new(); - for (type_name, node) in nodes { - if let node_types::EntryKind::Table { - name: main_table_name, - fields, - } = &node.kind - { - if fields.is_empty() { - panic!("Encountered node '{}' with no fields", type_name.kind); - } - - // Count how many columns there will be in the main table. There - // will be one for the id, plus one for each field that's stored - // as a column. - let main_table_arity = 1 + fields - .iter() - .filter(|&f| matches!(f.storage, node_types::Storage::Column { .. })) - .count(); - - let mut main_table_column_index: usize = 0; - let mut predicates = Vec::new(); - let mut exprs = Vec::new(); - for field in fields { - let (get_preds, get_child_expr) = create_field_getters( - main_table_name, - main_table_arity, - &mut main_table_column_index, - field, - nodes, - ); - predicates.extend(get_preds); - if let Some(get_child_expr) = get_child_expr { - exprs.push(get_child_expr) - } - } - field_predicates.insert(type_name, predicates); - get_child_exprs.insert(type_name, exprs); - } - } - - // Next, for every supertype (union) node, compute the predicates that are - // guaranteed to be defined identically (in name, return type, and formal - // parameters) by every one of its members. Such predicates can be hoisted - // to an `abstract` predicate on the supertype's class, with the - // corresponding predicates on its members becoming `override`s. - let mut exposed_predicates: BTreeMap<&str, Vec>> = BTreeMap::new(); - for (type_name, node) in nodes { - if let node_types::EntryKind::Union { .. } = &node.kind { - compute_exposed_predicates( - type_name, - nodes, - &field_predicates, - &mut exposed_predicates, - ); - } - } - for (type_name, node) in nodes { match &node.kind { node_types::EntryKind::Token { kind_id: _ } => { @@ -1017,26 +1001,20 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec> { })); } } - node_types::EntryKind::Union { members: _ } => { + node_types::EntryKind::Union { fields, .. } => { // It's a tree-sitter supertype node, so we're wrapping a dbscheme - // union type. Any predicate that's identically defined by every - // member becomes an `abstract` predicate here. - let predicates = exposed_predicates - .get(node.ql_class_name.as_str()) - .cloned() - .unwrap_or_default() - .into_iter() - .map(|predicate| ql::Predicate { - overridden: is_predicate_inherited( - &predicate, + // union type. + let predicates = fields + .iter() + .flat_map(|field| create_supertype_field_getters(field, nodes)) + .map(|mut predicate| { + predicate.overridden = is_predicate_inherited( type_name, - &direct_supertypes, - &exposed_predicates, - ), - is_private: false, - is_final: false, - body: None, - ..predicate + &predicate, + nodes, + &exposed_predicate_signatures, + ); + predicate }) .collect(); classes.push(ql::TopLevel::Class(ql::Class { @@ -1055,7 +1033,22 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec> { predicates, })); } - node_types::EntryKind::Table { .. } => { + node_types::EntryKind::Table { + name: main_table_name, + fields, + } => { + if fields.is_empty() { + panic!("Encountered node '{}' with no fields", type_name.kind); + } + + // Count how many columns there will be in the main table. There + // will be one for the id, plus one for each field that's stored + // as a column. + let main_table_arity = 1 + fields + .iter() + .filter(|&f| matches!(f.storage, node_types::Storage::Column { .. })) + .count(); + let main_class_name = &node.ql_class_name; let mut main_class = ql::Class { qldoc: Some(format!("A class representing `{}` nodes.", type_name.kind)), @@ -1073,30 +1066,34 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec> { predicates: vec![create_get_a_primary_ql_class(main_class_name, true)], }; - // A field getter that's identically defined (in signature) by - // every member of one of this node's direct supertypes is an - // override of the corresponding `abstract` predicate declared - // there. - main_class.predicates.extend( - field_predicates - .get(type_name) - .cloned() - .unwrap_or_default() - .into_iter() - .map(|predicate| { - let overridden = predicate.overridden - || is_predicate_inherited( - &predicate, - type_name, - &direct_supertypes, - &exposed_predicates, - ); - ql::Predicate { - overridden, - ..predicate - } - }), - ); + let mut main_table_column_index: usize = 0; + let mut get_child_exprs: Vec = Vec::new(); + + // Iterate through the fields, creating: + // - classes to wrap union types if fields need them, + // - predicates to access the fields, + // - the QL expressions to access the fields that will be part of getAFieldOrChild. + for field in fields { + let (mut get_preds, get_child_expr) = create_field_getters( + main_table_name, + main_table_arity, + &mut main_table_column_index, + field, + nodes, + ); + for predicate in &mut get_preds { + predicate.overridden = is_predicate_inherited( + type_name, + predicate, + nodes, + &exposed_predicate_signatures, + ); + } + main_class.predicates.extend(get_preds); + if let Some(get_child_expr) = get_child_expr { + get_child_exprs.push(get_child_expr) + } + } main_class.predicates.push(ql::Predicate { qldoc: Some(String::from("Gets a field or child node of this node.")), @@ -1106,9 +1103,7 @@ pub fn convert_nodes(nodes: &node_types::NodeTypeMap) -> Vec> { is_final: true, return_type: Some(ql::Type::Facade("AstNode")), formal_parameters: vec![], - body: Some(ql::Expression::Or( - get_child_exprs.get(type_name).cloned().unwrap_or_default(), - )), + body: ql::Expression::Or(get_child_exprs), overlay: None, }); @@ -1202,7 +1197,7 @@ pub fn create_print_ast_module(nodes: &node_types::NodeTypeMap) -> ql::TopLevel< param_type: ql::Type::Int, }, ], - body: Some(ql::Expression::Or(disjuncts)), + body: ql::Expression::Or(disjuncts), overlay: None, }; @@ -1215,3 +1210,99 @@ pub fn create_print_ast_module(nodes: &node_types::NodeTypeMap) -> ql::TopLevel< overlay: None, }) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn indexes_predicate_signatures_exposed_by_classes() { + let yaml = r#" +supertypes: + callable: + subtypes: [function_like] + fields: + parameter*: parameter + body?: block + function_like: + subtypes: [function] + fields: + name: identifier +named: + function: + parameter*: parameter + body?: block + name: identifier + parameter: + block: + identifier: +"#; + let json = yeast::node_types_yaml::convert(yaml).unwrap(); + let nodes = node_types::read_node_types_str("test", &json).unwrap(); + let signatures = compute_exposed_predicate_signatures(&nodes); + let function = signatures + .get(&node_types::TypeName { + kind: "function".to_owned(), + named: true, + }) + .unwrap(); + + assert!(function.contains(&PredicateSignature { + name: "getParameter", + arity: 1, + })); + assert!(function.contains(&PredicateSignature { + name: "getAParameter", + arity: 0, + })); + assert!(function.contains(&PredicateSignature { + name: "getBody", + arity: 0, + })); + assert!(function.contains(&PredicateSignature { + name: "getName", + arity: 0, + })); + } + + #[test] + fn generates_supertype_getters_and_concrete_overrides() { + let yaml = r#" +supertypes: + callable: + subtypes: [function] + fields: + parameter*: parameter + body?: block +named: + function: + parameter*: parameter + body?: block + parameter: + block: +"#; + let json = yeast::node_types_yaml::convert(yaml).unwrap(); + let nodes = node_types::read_node_types_str("test", &json).unwrap(); + let generated = convert_nodes(&nodes) + .into_iter() + .map(|element| element.to_string()) + .collect::>() + .join("\n"); + + assert!(generated.contains("F::Parameter getParameter(int i) { none() }")); + assert!(generated.contains("F::Parameter getAParameter() { none() }")); + assert!(generated.contains("F::Block getBody() { none() }")); + assert!( + generated.contains( + "final override F::Parameter getParameter(int i) { test_function_parameter" + ) + ); + assert!(generated.contains( + "final override F::Parameter getAParameter() { result = this.getParameter(_) }" + )); + assert!( + generated + .contains("final override F::Block getBody() { test_function_body(this, result) }") + ); + } +} diff --git a/shared/tree-sitter-extractor/src/node_types.rs b/shared/tree-sitter-extractor/src/node_types.rs index 2967bc845802..50cd03ba00d0 100644 --- a/shared/tree-sitter-extractor/src/node_types.rs +++ b/shared/tree-sitter-extractor/src/node_types.rs @@ -17,9 +17,17 @@ pub struct Entry { #[derive(Debug)] pub enum EntryKind { - Union { members: Set }, - Table { name: String, fields: Vec }, - Token { kind_id: usize }, + Union { + members: Set, + fields: Vec, + }, + Table { + name: String, + fields: Vec, + }, + Token { + kind_id: usize, + }, } #[derive(Clone, Debug, Ord, PartialOrd, Eq, PartialEq)] @@ -135,16 +143,39 @@ pub fn convert_nodes(prefix: &str, nodes: &[NodeInfo]) -> NodeTypeMap { if !subtypes.is_empty() { // It's a tree-sitter supertype node, for which we create a union // type. + let type_name = TypeName { + kind: node.kind.clone(), + named: node.named, + }; + let mut fields = Vec::new(); + for (field_name, field_info) in &node.fields { + add_field( + prefix, + &type_name, + Some(field_name.to_string()), + field_info, + &mut fields, + &token_kinds, + ); + } + if let Some(children) = &node.children { + add_field( + prefix, + &type_name, + None, + children, + &mut fields, + &token_kinds, + ); + } entries.insert( - TypeName { - kind: node.kind.clone(), - named: node.named, - }, + type_name, Entry { dbscheme_name, ql_class_name, kind: EntryKind::Union { members: convert_types(subtypes), + fields, }, }, ); @@ -454,3 +485,49 @@ fn to_snake_case_test() { assert_eq!("erb", to_snake_case("ERB")); assert_eq!("embedded_template", to_snake_case("EmbeddedTemplate")); } + +#[test] +fn supertype_fields_are_preserved() { + let yaml = r#" +supertypes: + callable: + subtypes: [function] + fields: + parameter*: parameter + body?: block +named: + function: + parameter: + block: +"#; + let json = yeast::node_types_yaml::convert(yaml).unwrap(); + let nodes = read_node_types_str("test", &json).unwrap(); + let callable = nodes + .get(&TypeName { + kind: "callable".to_owned(), + named: true, + }) + .unwrap(); + let EntryKind::Union { fields, .. } = &callable.kind else { + panic!("callable should be a union"); + }; + + assert_eq!(fields.len(), 2); + assert_eq!(fields[0].getter_name, "getBody"); + assert!(matches!( + fields[0].storage, + Storage::Table { + has_index: false, + .. + } + )); + assert_eq!(fields[1].getter_name, "getParameter"); + assert_eq!(fields[1].any_getter_name.as_deref(), Some("getAParameter")); + assert!(matches!( + fields[1].storage, + Storage::Table { + has_index: true, + .. + } + )); +} diff --git a/shared/yeast-schema/src/node_types_yaml.rs b/shared/yeast-schema/src/node_types_yaml.rs index c2d27793a18c..2bdcc02945e9 100644 --- a/shared/yeast-schema/src/node_types_yaml.rs +++ b/shared/yeast-schema/src/node_types_yaml.rs @@ -5,8 +5,11 @@ /// ```yaml /// supertypes: /// _expression: -/// - assignment -/// - binary +/// subtypes: +/// - assignment +/// - binary +/// fields: +/// operand*: _expression /// /// named: /// assignment: @@ -31,13 +34,39 @@ use serde_json::json; #[derive(Deserialize, Default)] struct YamlNodeTypes { #[serde(default)] - supertypes: BTreeMap>, + supertypes: BTreeMap, #[serde(default)] named: BTreeMap>>, #[serde(default)] unnamed: Vec, } +#[derive(Deserialize)] +#[serde(untagged)] +enum YamlSupertype { + Subtypes(Vec), + Detailed { + subtypes: Vec, + #[serde(default)] + fields: BTreeMap, + }, +} + +impl YamlSupertype { + fn subtypes(&self) -> &[TypeRef] { + match self { + Self::Subtypes(subtypes) | Self::Detailed { subtypes, .. } => subtypes, + } + } + + fn fields(&self) -> Option<&BTreeMap> { + match self { + Self::Subtypes(_) => None, + Self::Detailed { fields, .. } => Some(fields), + } + } +} + /// A reference to a node type. Can be: /// - a plain string (resolved by looking up named vs unnamed) /// - a map `{unnamed: "name"}` to force unnamed interpretation @@ -151,16 +180,19 @@ pub fn convert(yaml_input: &str) -> Result { let mut output = Vec::new(); // 1. Supertypes - for (name, members) in &yaml.supertypes { - let subtypes: Vec<_> = members + for (name, supertype) in &yaml.supertypes { + let subtypes: Vec<_> = supertype + .subtypes() .iter() .map(|m| resolve_type_ref(m, &named_types, &unnamed_types)) .collect(); - output.push(json!({ + let mut entry = json!({ "type": name, "named": true, "subtypes": subtypes, - })); + }); + add_fields_to_json_entry(&mut entry, supertype.fields(), &named_types, &unnamed_types); + output.push(entry); } // 2. Named nodes @@ -186,19 +218,33 @@ pub fn convert(yaml_input: &str) -> Result { Some(m) => m, }; + let mut entry = json!({ + "type": name, + "named": true, + "fields": {}, + }); + add_fields_to_json_entry(&mut entry, Some(fields_map), &named_types, &unnamed_types); + + output.push(entry); + } + + fn add_fields_to_json_entry( + entry: &mut serde_json::Value, + fields: Option<&BTreeMap>, + named_types: &BTreeSet, + unnamed_types: &BTreeSet, + ) { let mut json_fields = serde_json::Map::new(); - let mut json_children: Option = None; + let mut json_children = None; - for (raw_field_name, type_refs) in fields_map { + for (raw_field_name, type_refs) in fields.into_iter().flatten() { let spec = parse_field_name(raw_field_name); let types: Vec<_> = type_refs .clone() .into_vec() .iter() - .map(|t| resolve_type_ref(t, &named_types, &unnamed_types)) + .map(|t| resolve_type_ref(t, named_types, unnamed_types)) .collect(); - - // Cloning to make the borrow checker happy let field_info = json!({ "multiple": spec.multiple, "required": spec.required, @@ -208,25 +254,20 @@ pub fn convert(yaml_input: &str) -> Result { if let Some(name) = spec.name { json_fields.insert(name, field_info); } else { - // $children json_children = Some(field_info); } } - let mut entry = json!({ - "type": name, - "named": true, - "fields": json_fields, - }); - + entry + .as_object_mut() + .unwrap() + .insert("fields".to_string(), json_fields.into()); if let Some(children) = json_children { entry .as_object_mut() .unwrap() .insert("children".to_string(), children); } - - output.push(entry); } // 3. Unnamed tokens @@ -265,35 +306,53 @@ pub fn extend_schema_from_yaml( fn record_field_order(schema: &mut crate::schema::Schema, yaml_input: &str) -> Result<(), String> { let value: serde_yaml::Value = serde_yaml::from_str(yaml_input) .map_err(|e| format!("Failed to parse YAML for field order: {e}"))?; - let Some(named) = value.get("named").and_then(|v| v.as_mapping()) else { - return Ok(()); - }; - for (node_name, fields) in named { - let Some(node_name) = node_name.as_str() else { - continue; - }; - let Some(fields) = fields.as_mapping() else { - continue; // node with no fields (null) - }; - let mut order = Vec::new(); - for (raw_field_name, _) in fields { - let Some(raw) = raw_field_name.as_str() else { - continue; - }; - // Skip the unnamed/`child` slot; the dump handles it separately. - if let Some(name) = parse_field_name(raw).name { - order.push(schema.register_field(&name)); - } + record_field_order_for_nodes(schema, value.get("named").and_then(|v| v.as_mapping())); + let supertypes = value.get("supertypes").and_then(|v| v.as_mapping()); + if let Some(supertypes) = supertypes { + for (node_name, declaration) in supertypes { + let fields = declaration + .as_mapping() + .and_then(|mapping| mapping.get("fields")); + record_field_order_for_node(schema, node_name, fields); } - schema.set_field_order(node_name, order); } Ok(()) } -fn apply_yaml_to_schema( - yaml: &YamlNodeTypes, +fn record_field_order_for_nodes( + schema: &mut crate::schema::Schema, + nodes: Option<&serde_yaml::Mapping>, +) { + for (node_name, fields) in nodes.into_iter().flatten() { + record_field_order_for_node(schema, node_name, Some(fields)); + } +} + +fn record_field_order_for_node( schema: &mut crate::schema::Schema, + node_name: &serde_yaml::Value, + fields: Option<&serde_yaml::Value>, ) { + let Some(node_name) = node_name.as_str() else { + return; + }; + let Some(fields) = fields.and_then(|fields| fields.as_mapping()) else { + return; + }; + let mut order = Vec::new(); + for (raw_field_name, _) in fields { + let Some(raw) = raw_field_name.as_str() else { + continue; + }; + // Skip the unnamed/`child` slot; the dump handles it separately. + if let Some(name) = parse_field_name(raw).name { + order.push(schema.register_field(&name)); + } + } + schema.set_field_order(node_name, order); +} + +fn apply_yaml_to_schema(yaml: &YamlNodeTypes, schema: &mut crate::schema::Schema) { // Register all supertypes as node kinds for name in yaml.supertypes.keys() { schema.register_kind(name); @@ -302,14 +361,10 @@ fn apply_yaml_to_schema( // Register named node kinds and their fields for (name, fields_opt) in &yaml.named { schema.register_kind(name); - if let Some(fields) = fields_opt { - for raw_field_name in fields.keys() { - let spec = parse_field_name(raw_field_name); - if let Some(field_name) = &spec.name { - schema.register_field(field_name); - } - } - } + register_fields(schema, fields_opt.as_ref()); + } + for supertype in yaml.supertypes.values() { + register_fields(schema, supertype.fields()); } // Register unnamed tokens as node kinds @@ -326,8 +381,9 @@ fn apply_yaml_to_schema( } let unnamed_types: BTreeSet = yaml.unnamed.iter().cloned().collect(); - for (supertype, members) in &yaml.supertypes { - let node_types = members + for (supertype, declaration) in &yaml.supertypes { + let node_types = declaration + .subtypes() .iter() .map(|m| { let (kind, named) = resolve_type_ref_pair(m, &named_types, &unnamed_types); @@ -339,41 +395,74 @@ fn apply_yaml_to_schema( // Register allowed field child types for type checking. for (parent_kind, fields_opt) in &yaml.named { - let Some(fields) = fields_opt else { - continue; - }; - - for (raw_field_name, type_refs) in fields { - let spec = parse_field_name(raw_field_name); - let field_id = match &spec.name { - Some(name) => schema.register_field(name), - None => CHILD_FIELD, - }; + apply_fields_to_schema( + schema, + parent_kind, + fields_opt.as_ref(), + &named_types, + &unnamed_types, + ); + } + for (parent_kind, supertype) in &yaml.supertypes { + apply_fields_to_schema( + schema, + parent_kind, + supertype.fields(), + &named_types, + &unnamed_types, + ); + } +} - let mut node_types = type_refs - .clone() - .into_vec() - .into_iter() - .map(|type_ref| { - let (kind, named) = resolve_type_ref_pair(&type_ref, &named_types, &unnamed_types); - crate::schema::NodeType { kind, named } - }) - .collect::>(); - node_types.sort_by(|a, b| a.kind.cmp(&b.kind).then(a.named.cmp(&b.named))); - node_types.dedup_by(|a, b| a.kind == b.kind && a.named == b.named); - schema.set_field_types(parent_kind, field_id, node_types); - schema.set_field_cardinality( - parent_kind, - field_id, - crate::schema::FieldCardinality { - multiple: spec.multiple, - required: spec.required, - }, - ); +fn register_fields( + schema: &mut crate::schema::Schema, + fields: Option<&BTreeMap>, +) { + if let Some(fields) = fields { + for raw_field_name in fields.keys() { + if let Some(field_name) = parse_field_name(raw_field_name).name { + schema.register_field(&field_name); + } } } } +fn apply_fields_to_schema( + schema: &mut crate::schema::Schema, + parent_kind: &str, + fields: Option<&BTreeMap>, + named_types: &BTreeSet, + unnamed_types: &BTreeSet, +) { + for (raw_field_name, type_refs) in fields.into_iter().flatten() { + let spec = parse_field_name(raw_field_name); + let field_id = match &spec.name { + Some(name) => schema.register_field(name), + None => CHILD_FIELD, + }; + let mut node_types = type_refs + .clone() + .into_vec() + .into_iter() + .map(|type_ref| { + let (kind, named) = resolve_type_ref_pair(&type_ref, named_types, unnamed_types); + crate::schema::NodeType { kind, named } + }) + .collect::>(); + node_types.sort_by(|a, b| a.kind.cmp(&b.kind).then(a.named.cmp(&b.named))); + node_types.dedup_by(|a, b| a.kind == b.kind && a.named == b.named); + schema.set_field_types(parent_kind, field_id, node_types); + schema.set_field_cardinality( + parent_kind, + field_id, + crate::schema::FieldCardinality { + multiple: spec.multiple, + required: spec.required, + }, + ); + } +} + pub fn schema_from_yaml(yaml_input: &str) -> Result { let mut schema = crate::schema::Schema::new(); extend_schema_from_yaml(&mut schema, yaml_input)?; @@ -411,6 +500,12 @@ struct JsonFieldInfo { types: Vec, } +struct JsonSupertype { + subtypes: Vec, + fields: BTreeMap, + children: Option, +} + /// Convert a tree-sitter node-types.json string to the YAML format. pub fn convert_from_json(json_input: &str) -> Result { let nodes: Vec = @@ -427,7 +522,7 @@ pub fn convert_from_json(json_input: &str) -> Result { } } - let mut supertypes: BTreeMap> = BTreeMap::new(); + let mut supertypes: BTreeMap = BTreeMap::new(); let mut named: BTreeMap>> = BTreeMap::new(); let mut unnamed: Vec = Vec::new(); @@ -438,7 +533,14 @@ pub fn convert_from_json(json_input: &str) -> Result { } if !node.subtypes.is_empty() { - supertypes.insert(node.kind, node.subtypes); + supertypes.insert( + node.kind, + JsonSupertype { + subtypes: node.subtypes, + fields: node.fields, + children: node.children, + }, + ); continue; } @@ -463,13 +565,35 @@ pub fn convert_from_json(json_input: &str) -> Result { // Supertypes if !supertypes.is_empty() { writeln!(out, "supertypes:").unwrap(); - for (name, members) in &supertypes { - writeln!(out, " {name}:").unwrap(); - for member in members { + for (name, supertype) in &supertypes { + if supertype.fields.is_empty() && supertype.children.is_none() { + writeln!(out, " {name}:").unwrap(); + } else { + writeln!(out, " {name}:").unwrap(); + writeln!(out, " subtypes:").unwrap(); + } + let indent = if supertype.fields.is_empty() && supertype.children.is_none() { + " " + } else { + " " + }; + for member in &supertype.subtypes { let ref_str = format_type_ref(&member.kind, member.named, &all_named, &all_unnamed); - writeln!(out, " - {ref_str}").unwrap(); + writeln!(out, "{indent}- {ref_str}").unwrap(); + } + if !supertype.fields.is_empty() || supertype.children.is_some() { + writeln!(out, " fields:").unwrap(); + write_json_fields( + &mut out, + &supertype.fields, + supertype.children.as_ref(), + " ", + &all_named, + &all_unnamed, + ); } } + writeln!(out).unwrap(); } @@ -483,31 +607,7 @@ pub fn convert_from_json(json_input: &str) -> Result { } Some(fields) => { writeln!(out, " {name}:").unwrap(); - for (field_name, info) in fields { - let suffix = field_suffix(info.multiple, info.required); - let yaml_name = if field_name == "$children" { - format!("$children{suffix}") - } else { - format!("{field_name}{suffix}") - }; - - let type_refs: Vec = info - .types - .iter() - .map(|t| format_type_ref(&t.kind, t.named, &all_named, &all_unnamed)) - .collect(); - - if type_refs.len() == 1 { - writeln!(out, " {yaml_name}: {}", type_refs[0]).unwrap(); - } else { - let list = type_refs - .iter() - .map(|s| s.as_str()) - .collect::>() - .join(", "); - writeln!(out, " {yaml_name}: [{list}]").unwrap(); - } - } + write_json_fields(&mut out, fields, None, " ", &all_named, &all_unnamed); } } } @@ -525,6 +625,38 @@ pub fn convert_from_json(json_input: &str) -> Result { Ok(out) } +fn write_json_fields( + out: &mut String, + fields: &BTreeMap, + children: Option<&JsonFieldInfo>, + indent: &str, + all_named: &BTreeSet, + all_unnamed: &BTreeSet, +) { + for (field_name, info) in fields + .iter() + .map(|(name, info)| (Some(name.as_str()), info)) + .chain(children.map(|info| (None, info))) + { + let suffix = field_suffix(info.multiple, info.required); + let yaml_name = match field_name { + Some(name) => format!("{name}{suffix}"), + None => format!("$children{suffix}"), + }; + let type_refs: Vec = info + .types + .iter() + .map(|t| format_type_ref(&t.kind, t.named, all_named, all_unnamed)) + .collect(); + + if type_refs.len() == 1 { + writeln!(out, "{indent}{yaml_name}: {}", type_refs[0]).unwrap(); + } else { + writeln!(out, "{indent}{yaml_name}: [{}]", type_refs.join(", ")).unwrap(); + } + } +} + fn field_suffix(multiple: bool, required: bool) -> &'static str { match (multiple, required) { (false, true) => "", @@ -724,6 +856,50 @@ named: assert_eq!(fields["optional_multiple"]["multiple"], true); } + #[test] + fn test_supertype_fields() { + let yaml = r#" +supertypes: + callable: + subtypes: + - function + - closure + fields: + parameter*: parameter + body?: block + +named: + function: + closure: + parameter: + block: +"#; + + let json = convert(yaml).unwrap(); + let result: Vec = serde_json::from_str(&json).unwrap(); + let callable = result.iter().find(|n| n["type"] == "callable").unwrap(); + assert_eq!(callable["subtypes"].as_array().unwrap().len(), 2); + assert_eq!(callable["fields"]["parameter"]["multiple"], true); + assert_eq!(callable["fields"]["parameter"]["required"], false); + assert_eq!(callable["fields"]["body"]["multiple"], false); + assert_eq!(callable["fields"]["body"]["required"], false); + + let schema = schema_from_yaml(yaml).unwrap(); + let parameter = schema.field_id_for_name("parameter").unwrap(); + assert!(schema.field_types("callable", parameter).is_some()); + assert_eq!( + schema.field_cardinality("callable", parameter), + Some(crate::schema::FieldCardinality { + multiple: true, + required: false, + }) + ); + assert_eq!( + schema.field_order("callable"), + Some(&vec![parameter, schema.field_id_for_name("body").unwrap()]) + ); + } + #[test] fn test_json_to_yaml() { let json = r#"[ @@ -794,4 +970,30 @@ unnamed: let v2: serde_json::Value = serde_json::from_str(&json2).unwrap(); assert_eq!(v1, v2); } + + #[test] + fn test_supertype_fields_round_trip() { + let yaml = r#" +supertypes: + callable: + subtypes: [function, closure] + fields: + parameter*: parameter + body?: block + +named: + function: + closure: + parameter: + block: +"#; + + let json = convert(yaml).unwrap(); + let yaml_output = convert_from_json(&json).unwrap(); + let json_output = convert(&yaml_output).unwrap(); + assert_eq!( + serde_json::from_str::(&json).unwrap(), + serde_json::from_str::(&json_output).unwrap() + ); + } } diff --git a/unified/extractor/ast_types.yml b/unified/extractor/ast_types.yml index 5c97b12b2ab0..2b59c721b5dc 100644 --- a/unified/extractor/ast_types.yml +++ b/unified/extractor/ast_types.yml @@ -66,13 +66,17 @@ supertypes: - do_while_stmt - labeled_stmt callable: - - top_level - - function_expr - - function_declaration - - constructor_declaration - - destructor_declaration - - accessor_declaration - - initializer_declaration + subtypes: + - top_level + - function_expr + - function_declaration + - constructor_declaration + - destructor_declaration + - accessor_declaration + - initializer_declaration + fields: + parameter*: parameter + body?: block # A member is anything that can appear in the body of a class-like declaration member: - constructor_declaration diff --git a/unified/ql/lib/codeql/unified/internal/Ast.qll b/unified/ql/lib/codeql/unified/internal/Ast.qll index cfc917b85580..a376131d5ba3 100644 --- a/unified/ql/lib/codeql/unified/internal/Ast.qll +++ b/unified/ql/lib/codeql/unified/internal/Ast.qll @@ -116,12 +116,12 @@ module Unified { final F::Identifier getNameNode() { unified_accessor_declaration_def(this, _, result) } /** Gets the node corresponding to the field `parameter`. */ - final F::Parameter getParameter(int i) { + final override F::Parameter getParameter(int i) { unified_accessor_declaration_parameter(this, i, result) } /** Gets the node corresponding to the field `parameter`. */ - final F::Parameter getAParameter() { result = this.getParameter(_) } + final override F::Parameter getAParameter() { result = this.getParameter(_) } /** Gets the node corresponding to the field `type`. */ final F::Expr getType() { unified_accessor_declaration_type(this, result) } @@ -360,7 +360,13 @@ module Unified { class Callable extends @unified_callable, F::AstNode { /** Gets the node corresponding to the field `body`. */ - abstract F::Block getBody(); + F::Block getBody() { none() } + + /** Gets the node corresponding to the field `parameter`. */ + F::Parameter getParameter(int i) { none() } + + /** Gets the node corresponding to the field `parameter`. */ + F::Parameter getAParameter() { none() } } /** A class representing `catch_clause` nodes. */ @@ -498,12 +504,12 @@ module Unified { final F::Identifier getNameNode() { unified_constructor_declaration_name_node(this, result) } /** Gets the node corresponding to the field `parameter`. */ - final F::Parameter getParameter(int i) { + final override F::Parameter getParameter(int i) { unified_constructor_declaration_parameter(this, i, result) } /** Gets the node corresponding to the field `parameter`. */ - final F::Parameter getAParameter() { result = this.getParameter(_) } + final override F::Parameter getAParameter() { result = this.getParameter(_) } /** Gets a field or child node of this node. */ final override F::AstNode getAFieldOrChild() { @@ -689,12 +695,12 @@ module Unified { final F::Identifier getNameNode() { unified_function_declaration_def(this, result) } /** Gets the node corresponding to the field `parameter`. */ - final F::Parameter getParameter(int i) { + final override F::Parameter getParameter(int i) { unified_function_declaration_parameter(this, i, result) } /** Gets the node corresponding to the field `parameter`. */ - final F::Parameter getAParameter() { result = this.getParameter(_) } + final override F::Parameter getAParameter() { result = this.getParameter(_) } /** Gets the node corresponding to the field `return_type`. */ final F::Expr getReturnType() { unified_function_declaration_return_type(this, result) } @@ -750,10 +756,12 @@ module Unified { final F::Modifier getAModifier() { result = this.getModifier(_) } /** Gets the node corresponding to the field `parameter`. */ - final F::Parameter getParameter(int i) { unified_function_expr_parameter(this, i, result) } + final override F::Parameter getParameter(int i) { + unified_function_expr_parameter(this, i, result) + } /** Gets the node corresponding to the field `parameter`. */ - final F::Parameter getAParameter() { result = this.getParameter(_) } + final override F::Parameter getAParameter() { result = this.getParameter(_) } /** Gets the node corresponding to the field `return_type`. */ final F::Expr getReturnType() { unified_function_expr_return_type(this, result) }