Skip to content

Commit e36aa9f

Browse files
asgerfCopilot
andcommitted
Generate getters for supertype fields
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent 0af0b60 commit e36aa9f

2 files changed

Lines changed: 199 additions & 55 deletions

File tree

  • shared/tree-sitter-extractor/src/generator
  • unified/ql/lib/codeql/unified/internal

‎shared/tree-sitter-extractor/src/generator/ql_gen.rs‎

Lines changed: 172 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
791840
fn 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.
898969
pub 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

Comments
 (0)