rapx/check/opt/checking/encoding_checking/
array_encoding.rs1use std::collections::HashSet;
2
3use rustc_middle::{mir::Local, ty::TyCtxt};
4use rustc_span::Span;
5
6use super::{report_encoding_bug, value_is_from_const};
7use crate::analysis::dataflow::*;
8use crate::check::opt::OptCheck;
9
10crate::def_paths! {
11 str_from_utf8: "std::str::from_utf8",
12}
13
14pub struct ArrayEncodingCheck {
15 record: Vec<Span>,
16}
17
18fn extract_ancestor_set_if_is_str_from(
19 graph: &Graph,
20 node_idx: Local,
21 node: &GraphNode,
22) -> Option<HashSet<Local>> {
23 let def_paths = DEFPATHS.get().unwrap();
24 for op in node.ops.iter() {
25 if let NodeOp::Call(def_id) = op {
26 if *def_id == def_paths.str_from_utf8.last_def_id() {
27 return Some(graph.collect_ancestor_locals(node_idx, false));
28 }
29 }
30 }
31 None
32}
33
34fn is_valid_index_edge(graph: &Graph, edge: &GraphEdge) -> bool {
35 if let EdgeOp::Index = edge.op {
36 let dst_node = &graph.nodes[edge.dst];
38 if dst_node.in_edges.len() > 2 {
39 let rvalue_edge_idx = dst_node.in_edges[2];
41 let rvalue_idx = graph.edges[rvalue_edge_idx].src;
42 if value_is_from_const(graph, rvalue_idx) {
43 return true;
44 }
45 }
46 }
47 false
48}
49
50impl OptCheck for ArrayEncodingCheck {
51 fn new() -> Self {
52 Self { record: Vec::new() }
53 }
54
55 fn check(&mut self, graph: &Graph, tcx: &TyCtxt) {
56 DEFPATHS.get_or_init(|| DefPaths::new(tcx));
57 let common_ancestor = graph
58 .edges
59 .iter()
60 .filter_map(|edge| {
61 if is_valid_index_edge(graph, edge) {
63 Some(graph.collect_ancestor_locals(edge.src, true))
64 } else {
65 None
66 }
67 })
68 .reduce(|set1, set2| set1.into_iter().filter(|k| set2.contains(k)).collect());
69
70 if let Some(common_ancestor) = common_ancestor {
71 for (node_idx, node) in graph.nodes.iter_enumerated() {
72 if let Some(str_from_ancestor_set) =
73 extract_ancestor_set_if_is_str_from(graph, node_idx, node)
74 {
75 if !common_ancestor
76 .intersection(&str_from_ancestor_set)
77 .next()
78 .is_some()
79 {
80 self.record.clear();
81 return;
82 }
83 self.record.push(node.span);
84 }
85 }
86 }
87 }
88
89 fn report(&self, graph: &Graph) {
90 for span in self.record.iter() {
91 report_encoding_bug(graph, *span);
92 }
93 }
94
95 fn cnt(&self) -> usize {
96 self.record.len()
97 }
98}