use alloc::format;
use alloc::string::String;
use crate::query::adj_ast;
use crate::syntax::builders::utils::{
get_event_subpath, receiver_base_symbol, register_usage_event,
};
use crate::syntax::{
SanitizationEvent, SyntaxGraph, SyntaxGraphArgs, SyntaxGraphError, SyntaxNode, UsageRole,
};
use crate::NodeId;
#[derive(Clone, Copy, Default)]
#[allow(
clippy::struct_field_names,
reason = "mirror the invocation node's child attribute names"
)]
pub struct DirectChildren {
pub expression_id: Option<NodeId>,
pub arguments_id: Option<NodeId>,
pub object_id: Option<NodeId>,
pub block_id: Option<NodeId>,
}
pub fn build_method_invocation_node(
args: &mut SyntaxGraphArgs<'_>,
n_id: NodeId,
expr: String,
direct_children: DirectChildren,
object_name: Option<String>,
) -> Result<NodeId, SyntaxGraphError> {
let symbol_scope = args.metadata.scope_stack.last().copied();
args.syntax_graph.add_node(
n_id,
SyntaxNode::MethodInvocation {
expression: expr,
object: object_name.filter(|value| !value.is_empty()),
symbol_scope,
expression_id: direct_children.expression_id,
arguments_id: direct_children.arguments_id,
object_id: direct_children.object_id,
block_id: direct_children.block_id,
receiver_type_fqn: None,
},
);
let condition_parent = args
.metadata
.sanitization_stack
.last()
.filter(|(kind, _)| kind == "condition_pending")
.map(|(_, parent_id)| *parent_id);
if let Some(parent_id) = condition_parent {
args.metadata
.sanitization_stack
.push((String::from("condition"), parent_id));
}
process_invocation_children(args, n_id, direct_children)?;
if condition_parent.is_some() {
args.metadata.sanitization_stack.pop();
}
index_direct_arguments(args, n_id, direct_children.arguments_id);
index_usage_facts(args, n_id, direct_children.arguments_id);
Ok(n_id)
}
fn process_invocation_children(
args: &mut SyntaxGraphArgs<'_>,
n_id: NodeId,
children: DirectChildren,
) -> Result<(), SyntaxGraphError> {
for child_id in [
children.expression_id,
children.arguments_id,
children.object_id,
children.block_id,
]
.into_iter()
.flatten()
{
let built = args.generic(child_id)?;
args.syntax_graph.add_ast_edge(n_id, built);
}
Ok(())
}
fn index_direct_arguments(
args: &mut SyntaxGraphArgs<'_>,
n_id: NodeId,
arguments_id: Option<NodeId>,
) {
let Some(arguments_id) = arguments_id else {
return;
};
let Some(scope) = args.metadata.scope_stack.last().copied() else {
return;
};
let subpath = get_event_subpath(args.metadata);
for child_id in adj_ast(args.syntax_graph, arguments_id, Some(1), &[]) {
if let Some(symbol) = method_arg_symbol(args.syntax_graph, child_id) {
args.syntax_graph.register_sanitization(
&symbol,
scope,
SanitizationEvent {
node_id: n_id,
kind: String::from("method_arg"),
subpath,
},
);
}
}
}
fn method_arg_symbol(graph: &SyntaxGraph, n_id: NodeId) -> Option<String> {
match graph.nodes.get(&n_id)? {
SyntaxNode::SymbolLookup { symbol, .. } => (!symbol.is_empty()).then(|| symbol.clone()),
SyntaxNode::MemberAccess {
expression, member, ..
} => {
(!expression.is_empty() && !member.is_empty()).then(|| format!("{expression}.{member}"))
}
SyntaxNode::NamedArgument { value_id, .. } => method_arg_symbol(graph, *value_id),
_ => None,
}
}
fn index_usage_facts(args: &mut SyntaxGraphArgs<'_>, n_id: NodeId, arguments_id: Option<NodeId>) {
let subpath = get_event_subpath(args.metadata);
index_receiver_usage(args, n_id, subpath);
index_argument_usage(args, n_id, arguments_id, subpath);
}
fn index_receiver_usage(args: &mut SyntaxGraphArgs<'_>, n_id: NodeId, subpath: Option<NodeId>) {
let (base, dotted) = resolve_receiver(args.syntax_graph, n_id);
for name in base.into_iter().chain(dotted) {
register_usage_event(
args.syntax_graph,
args.metadata,
n_id,
&name,
UsageRole::Receiver,
None,
subpath,
);
}
}
fn index_argument_usage(
args: &mut SyntaxGraphArgs<'_>,
n_id: NodeId,
arguments_id: Option<NodeId>,
subpath: Option<NodeId>,
) {
let Some(arguments_id) = arguments_id else {
return;
};
for (arg_index, child_id) in adj_ast(args.syntax_graph, arguments_id, Some(1), &[])
.into_iter()
.enumerate()
{
let index = i64::try_from(arg_index).ok();
let (base, dotted) = argument_symbol_names(args.syntax_graph, child_id);
for name in base.into_iter().chain(dotted) {
register_usage_event(
args.syntax_graph,
args.metadata,
n_id,
&name,
UsageRole::Argument,
index,
subpath,
);
}
}
}
fn resolve_receiver(graph: &SyntaxGraph, n_id: NodeId) -> (Option<String>, Option<String>) {
let Some(SyntaxNode::MethodInvocation {
object_id,
expression,
expression_id,
..
}) = graph.nodes.get(&n_id)
else {
return (None, None);
};
if let Some(object_id) = *object_id {
if is_callee_wrapper(graph, object_id, expression) {
if let Some(inner_id) = member_access_expression_id(graph, object_id) {
return receiver_node_names(graph, inner_id);
}
}
let (base, dotted) = receiver_node_names(graph, object_id);
if base.is_some() || dotted.is_some() {
return (base, dotted);
}
}
if let Some(expression_id) = *expression_id {
if let Some(inner_id) = member_access_expression_id(graph, expression_id) {
return receiver_node_names(graph, inner_id);
}
}
(None, None)
}
fn member_access_expression_id(graph: &SyntaxGraph, n_id: NodeId) -> Option<NodeId> {
match graph.nodes.get(&n_id) {
Some(SyntaxNode::MemberAccess { expression_id, .. }) => Some(*expression_id),
_ => None,
}
}
fn is_callee_wrapper(
graph: &SyntaxGraph,
candidate_id: NodeId,
invocation_expression: &str,
) -> bool {
let Some(SyntaxNode::MemberAccess {
expression, member, ..
}) = graph.nodes.get(&candidate_id)
else {
return false;
};
!expression.is_empty()
&& !member.is_empty()
&& format!("{expression}.{member}") == invocation_expression
}
fn receiver_node_names(
graph: &SyntaxGraph,
receiver_id: NodeId,
) -> (Option<String>, Option<String>) {
match graph.nodes.get(&receiver_id) {
Some(SyntaxNode::SymbolLookup { symbol, .. }) => {
((!symbol.is_empty()).then(|| symbol.clone()), None)
}
Some(SyntaxNode::MemberAccess {
expression, member, ..
}) => {
let base = receiver_base_symbol(graph, receiver_id);
let dotted =
(!expression.is_empty() && !member.is_empty() && !expression.contains('('))
.then(|| format!("{expression}.{member}"));
let dotted = if dotted == base { None } else { dotted };
(base, dotted)
}
Some(
SyntaxNode::MethodInvocation { .. }
| SyntaxNode::ElementAccess { .. }
| SyntaxNode::ParenthesizedExpression,
) => (receiver_base_symbol(graph, receiver_id), None),
_ => (None, None),
}
}
fn argument_symbol_names(
graph: &SyntaxGraph,
child_id: NodeId,
) -> (Option<String>, Option<String>) {
match graph.nodes.get(&child_id) {
Some(SyntaxNode::NamedArgument { value_id, .. }) => argument_symbol_names(graph, *value_id),
Some(SyntaxNode::SymbolLookup { symbol, .. }) => {
((!symbol.is_empty()).then(|| symbol.clone()), None)
}
Some(SyntaxNode::MemberAccess {
expression, member, ..
}) => {
let base = receiver_base_symbol(graph, child_id);
let dotted =
(!expression.is_empty() && !member.is_empty() && !expression.contains('('))
.then(|| format!("{expression}.{member}"));
let dotted = if dotted == base { None } else { dotted };
(base, dotted)
}
_ => (None, None),
}
}
#[cfg(test)]
mod tests {
use super::{index_direct_arguments, index_usage_facts, method_arg_symbol};
use crate::ast::AstGraph;
use crate::syntax::{
SanitizationEvent, SyntaxGraph, SyntaxGraphArgs, SyntaxMetadata, SyntaxNode, SyntaxReader,
UsageEvent, UsageRole,
};
use crate::{Language, NodeId};
use alloc::borrow::ToOwned;
use alloc::vec;
fn no_dispatch(_: &str) -> Option<SyntaxReader> {
None
}
fn symbol(name: &str) -> SyntaxNode {
SyntaxNode::SymbolLookup {
symbol: name.to_owned(),
symbol_scope: None,
value: None,
}
}
#[test]
fn method_arg_symbol_reads_symbol_member_and_named_argument() {
let mut graph = SyntaxGraph::new();
graph.add_node(NodeId(1), symbol("user"));
graph.add_node(
NodeId(2),
SyntaxNode::MemberAccess {
member: "name".to_owned(),
expression: "user".to_owned(),
expression_id: NodeId(1),
symbol_scope: None,
},
);
graph.add_node(
NodeId(3),
SyntaxNode::NamedArgument {
value_id: NodeId(1),
argument_name: Some("arg".to_owned()),
},
);
graph.add_node(NodeId(4), SyntaxNode::ParenthesizedExpression);
assert_eq!(
method_arg_symbol(&graph, NodeId(1)),
Some("user".to_owned())
);
assert_eq!(
method_arg_symbol(&graph, NodeId(2)),
Some("user.name".to_owned())
);
assert_eq!(
method_arg_symbol(&graph, NodeId(3)),
Some("user".to_owned())
);
assert_eq!(method_arg_symbol(&graph, NodeId(4)), None);
}
#[test]
fn index_direct_arguments_registers_method_arg_sanitization() {
let ast = AstGraph::new();
let mut graph = SyntaxGraph::new();
graph.add_node(NodeId(3), symbol("payload"));
graph.add_node(NodeId(2), SyntaxNode::ArrayInitializer);
graph.add_ast_edge(NodeId(2), NodeId(3));
let mut meta = SyntaxMetadata::seeded(NodeId(3));
{
let mut args =
SyntaxGraphArgs::new(Language::Java, &ast, &mut graph, &mut meta, no_dispatch);
index_direct_arguments(&mut args, NodeId(1), Some(NodeId(2)));
}
assert_eq!(
graph
.sanitization_index
.get("payload")
.and_then(|by_scope| by_scope.get(&NodeId(1))),
Some(&vec![SanitizationEvent {
node_id: NodeId(1),
kind: "method_arg".to_owned(),
subpath: None,
}])
);
}
#[test]
fn index_usage_facts_registers_receiver_and_argument_events() {
let ast = AstGraph::new();
let mut graph = SyntaxGraph::new();
graph.add_node(NodeId(2), symbol("db"));
graph.add_node(NodeId(4), symbol("q"));
graph.add_node(NodeId(3), SyntaxNode::ArrayInitializer);
graph.add_ast_edge(NodeId(3), NodeId(4));
graph.add_node(
NodeId(1),
SyntaxNode::MethodInvocation {
expression: "db.query".to_owned(),
object: None,
symbol_scope: None,
expression_id: None,
arguments_id: Some(NodeId(3)),
object_id: Some(NodeId(2)),
block_id: None,
receiver_type_fqn: None,
},
);
let mut meta = SyntaxMetadata::seeded(NodeId(4));
{
let mut args =
SyntaxGraphArgs::new(Language::Java, &ast, &mut graph, &mut meta, no_dispatch);
index_usage_facts(&mut args, NodeId(1), Some(NodeId(3)));
}
assert_eq!(
graph
.usage_index
.get("db")
.and_then(|by_scope| by_scope.get(&NodeId(1))),
Some(&vec![UsageEvent {
node_id: NodeId(1),
role: UsageRole::Receiver,
arg_index: None,
subpath: None,
}])
);
assert_eq!(
graph
.usage_index
.get("q")
.and_then(|by_scope| by_scope.get(&NodeId(1))),
Some(&vec![UsageEvent {
node_id: NodeId(1),
role: UsageRole::Argument,
arg_index: Some(0),
subpath: None,
}])
);
}
}