Skip to main content

rapx/check/opt/checking/encoding_checking/
array_encoding.rs

1use 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        // must be Index edge
37        let dst_node = &graph.nodes[edge.dst];
38        if dst_node.in_edges.len() > 2 {
39            // must be the left value
40            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                // The index must be an lvalue and the rvalue must come from a const
62                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}