Skip to main content

rapx/analysis/api_dependency/graph/
resolve.rs

1use super::Config;
2use super::dep_edge::DepEdge;
3use super::dep_node::DepNode;
4use super::transform::TransformKind;
5use super::ty_wrapper::TyWrapper;
6use crate::analysis::api_dependency::ApiDependencyGraph;
7use crate::analysis::api_dependency::graph::std_tys;
8use crate::analysis::api_dependency::mono::{Mono, get_mono_complexity};
9use crate::analysis::api_dependency::utils::{
10    fn_requires_monomorphization, is_fuzzable_ty, ty_complexity,
11};
12use crate::analysis::api_dependency::visit::FnVisitor;
13use crate::analysis::api_dependency::{mono, utils};
14use crate::helpers::def_path::path_str_def_id;
15use crate::limit::MAX_TY_COMPLX;
16use crate::utils::fs::rap_create_file;
17use crate::{rap_debug, rap_info, rap_trace};
18use petgraph::Direction::{self, Incoming};
19use petgraph::Graph;
20use petgraph::dot;
21use petgraph::graph::NodeIndex;
22use petgraph::visit::{EdgeRef, NodeIndexable, Visitable};
23use rand::Rng;
24use rustc_hir::def_id::DefId;
25use rustc_middle::ty::{self, GenericArgsRef, TraitRef, Ty, TyCtxt};
26use rustc_span::sym::{self, require};
27use std::collections::HashMap;
28use std::collections::HashSet;
29use std::collections::VecDeque;
30use std::hash::Hash;
31use std::io::Write;
32use std::path::Path;
33use std::time;
34
35const RESOLVE_DEBUG: bool = false;
36
37fn add_return_type_if_reachable<'tcx>(
38    fn_did: DefId,
39    args: &[ty::GenericArg<'tcx>],
40    reachable_tys: &HashSet<TyWrapper<'tcx>>,
41    new_tys: &mut HashSet<Ty<'tcx>>,
42    tcx: TyCtxt<'tcx>,
43) -> bool {
44    let fn_sig = utils::fn_sig_with_generic_args(fn_did, args, tcx);
45    let inputs = fn_sig.inputs();
46    for input_ty in inputs {
47        if !is_fuzzable_ty(*input_ty, tcx) && !reachable_tys.contains(&TyWrapper::from(*input_ty)) {
48            return false;
49        }
50    }
51    let output_ty = fn_sig.output();
52    if !output_ty.is_unit() {
53        new_tys.insert(output_ty);
54    }
55    true
56}
57
58#[derive(Clone)]
59struct TypeCandidates<'tcx> {
60    tcx: TyCtxt<'tcx>,
61    candidates: HashSet<TyWrapper<'tcx>>,
62    max_complexity: usize,
63}
64
65impl<'tcx> TypeCandidates<'tcx> {
66    pub fn new(tcx: TyCtxt<'tcx>, max_complexity: usize) -> Self {
67        TypeCandidates {
68            tcx,
69            candidates: HashSet::new(),
70            max_complexity,
71        }
72    }
73
74    pub fn insert(&mut self, ty: Ty<'tcx>) -> bool {
75        if ty_complexity(ty) <= self.max_complexity {
76            self.candidates.insert(ty.into())
77        } else {
78            false
79        }
80    }
81
82    pub fn insert_all(&mut self, ty: Ty<'tcx>) -> bool {
83        let complexity = ty_complexity(ty);
84        self.insert_all_with_complexity(ty, complexity)
85    }
86
87    pub fn insert_all_with_complexity(&mut self, ty: Ty<'tcx>, current_cmplx: usize) -> bool {
88        if current_cmplx > self.max_complexity {
89            return false;
90        }
91
92        // add T
93        let mut changed = self.candidates.insert(ty.into());
94
95        // add &T
96        changed |= self.insert_all_with_complexity(
97            Ty::new_ref(
98                self.tcx,
99                self.tcx.lifetimes.re_erased,
100                ty,
101                ty::Mutability::Not,
102            ),
103            current_cmplx + 1,
104        );
105
106        // add &mut T
107        changed |= self.insert_all_with_complexity(
108            Ty::new_ref(
109                self.tcx,
110                self.tcx.lifetimes.re_erased,
111                ty,
112                ty::Mutability::Mut,
113            ),
114            current_cmplx + 1,
115        );
116
117        // add &[T]
118        changed |= self.insert_all_with_complexity(
119            Ty::new_ref(
120                self.tcx,
121                self.tcx.lifetimes.re_erased,
122                Ty::new_slice(self.tcx, ty),
123                ty::Mutability::Not,
124            ),
125            current_cmplx + 2,
126        );
127
128        // add &mut [T]
129        changed |= self.insert_all_with_complexity(
130            Ty::new_ref(
131                self.tcx,
132                self.tcx.lifetimes.re_erased,
133                Ty::new_slice(self.tcx, ty),
134                ty::Mutability::Mut,
135            ),
136            current_cmplx + 2,
137        );
138
139        changed
140    }
141
142    pub fn add_prelude_tys(&mut self) {
143        let tcx = self.tcx;
144
145        let primitive_tys = [
146            tcx.types.bool,
147            tcx.types.char,
148            tcx.types.f32,
149            tcx.types.i8,
150            tcx.types.u8,
151            tcx.types.i32,
152            tcx.types.u32,
153            tcx.types.i64,
154            tcx.types.u64,
155            tcx.types.isize,
156            tcx.types.usize,
157        ];
158
159        let mut prelude_tys = Vec::new();
160
161        prelude_tys.extend_from_slice(&primitive_tys);
162        // &str
163        prelude_tys.push(Ty::new_imm_ref(
164            tcx,
165            tcx.lifetimes.re_erased,
166            tcx.types.str_,
167        ));
168
169        if let Some(string_ty) = std_tys::std_string(tcx) {
170            prelude_tys.push(string_ty);
171        }
172
173        for element_ty in &primitive_tys {
174            if let Some(vec_ty) = std_tys::std_vec(*element_ty, tcx) {
175                prelude_tys.push(vec_ty);
176            }
177        }
178
179        prelude_tys.into_iter().for_each(|ty| {
180            self.insert_all(ty);
181        });
182    }
183
184    pub fn candidates(&self) -> &HashSet<TyWrapper<'tcx>> {
185        &self.candidates
186    }
187}
188
189pub fn partion_generic_api<'tcx>(
190    all_apis: &HashSet<DefId>,
191    tcx: TyCtxt<'tcx>,
192) -> (HashSet<DefId>, HashSet<DefId>) {
193    let mut generic_api = HashSet::new();
194    let mut non_generic_api = HashSet::new();
195    for api_id in all_apis.iter() {
196        if tcx.generics_of(*api_id).requires_monomorphization(tcx) {
197            generic_api.insert(*api_id);
198        } else {
199            non_generic_api.insert(*api_id);
200        }
201    }
202    (non_generic_api, generic_api)
203}
204
205impl<'tcx> ApiDependencyGraph<'tcx> {
206    pub fn resolve_generic_api(
207        &mut self,
208        non_generic_apis: &[DefId],
209        generic_apis: &[DefId],
210        max_iteration: usize,
211    ) {
212        rap_info!("start resolving generic APIs");
213
214        // 1. Reachable generic API search
215        let generic_map = self.search_reachable_apis(non_generic_apis, generic_apis, max_iteration);
216
217        self.add_mono_apis_from_map(&generic_map);
218        self.update_transform_edges();
219
220        rap_info!("finish resolving generic APIs");
221        self.statistics().info();
222
223        if RESOLVE_DEBUG {
224            self.dump_to_file(Path::new("api_graph_unpruned.dot"));
225        }
226
227        let reserved = self.prune_by_similarity(generic_map);
228
229        let count = self.reserve_nodes(&reserved);
230        rap_info!("remove {} nodes by pruning", count);
231    }
232
233    pub fn search_reachable_apis(
234        &mut self,
235        non_generic_apis: &[DefId],
236        generic_apis: &[DefId],
237        max_iteration: usize,
238    ) -> HashMap<DefId, HashSet<Mono<'tcx>>> {
239        let tcx = self.tcx;
240        let mut type_candidates = TypeCandidates::new(self.tcx, MAX_TY_COMPLX);
241
242        type_candidates.add_prelude_tys();
243
244        let mut generic_map: HashMap<DefId, HashSet<Mono>> = HashMap::new();
245        let mut unreachable_non_generic_api = Vec::from(non_generic_apis);
246
247        rap_debug!("[resolve_generic] non_generic_api = {unreachable_non_generic_api:?}");
248        rap_debug!("[resolve_generic] generic_api = {generic_apis:?}");
249
250        let mut num_iter = 0;
251
252        loop {
253            num_iter += 1;
254            let all_reachable_tys = type_candidates.candidates();
255            rap_info!(
256                "start iter #{num_iter}, # of reachble types = {}",
257                all_reachable_tys.len()
258            );
259
260            // dump all reachable types to files, each line output a type
261            if RESOLVE_DEBUG {
262                let mut file =
263                    rap_create_file(Path::new("reachable_types.txt"), "create file fail");
264                for ty in all_reachable_tys.iter() {
265                    writeln!(file, "{}", ty.ty()).unwrap();
266                }
267            }
268
269            let mut current_tys = HashSet::new();
270
271            // check whether there is any non-generic reachable in this iteration.
272            // if the api is reachable, add output type to reachble_tys,
273            // and remove it from the set.
274            unreachable_non_generic_api.retain(|fn_did| {
275                !add_return_type_if_reachable(
276                    *fn_did,
277                    ty::GenericArgs::identity_for_item(tcx, *fn_did),
278                    all_reachable_tys,
279                    &mut current_tys,
280                    tcx,
281                )
282            });
283
284            // check each generic API for new monomorphic API
285            for fn_did in generic_apis.iter() {
286                let mono_set = mono::resolve_mono_apis(*fn_did, all_reachable_tys, tcx);
287                rap_debug!(
288                    "[search_reachable_apis] {} -> {:?}",
289                    tcx.def_path_str(*fn_did),
290                    mono_set
291                );
292                for mono in mono_set.monos {
293                    let fn_sig = utils::fn_sig_with_generic_args(*fn_did, &mono.value, tcx);
294                    let output_ty = fn_sig.output();
295                    if generic_map.entry(*fn_did).or_default().insert(mono) {
296                        if !output_ty.is_unit() && ty_complexity(output_ty) <= MAX_TY_COMPLX {
297                            current_tys.insert(output_ty);
298                        }
299                    }
300                }
301            }
302
303            let mut changed = false;
304            for ty in current_tys {
305                changed = changed | type_candidates.insert_all(ty);
306            }
307
308            if !changed {
309                rap_info!("Terminate. Reachable types unchange in this iteration.");
310                break;
311            }
312            if num_iter >= max_iteration {
313                rap_info!("Terminate. Max iteration reached.");
314                break;
315            }
316        }
317
318        let mono_cnt = generic_map.values().fold(0, |acc, monos| acc + monos.len());
319
320        rap_debug!("# reachable types: {}", type_candidates.candidates().len());
321        rap_debug!("# mono APIs: {}", mono_cnt);
322
323        generic_map
324    }
325
326    pub fn add_mono_apis_from_map(&mut self, generic_map: &HashMap<DefId, HashSet<Mono<'tcx>>>) {
327        for (fn_did, mono_set) in generic_map {
328            for mono in mono_set {
329                let args = self.tcx.mk_args(&mono.value);
330                self.add_api(*fn_did, args);
331            }
332        }
333    }
334
335    /// heuristic strategy: prioritize to reserve APIs that first arg of which is reachable.
336    /// This is based on that we want to reserve APIs that have the same Self type ASAP.
337    pub fn heuristic_select(&mut self, reserved: &mut [bool]) {
338        let mut worklist = VecDeque::new();
339        let mut visited = vec![false; self.graph.node_count()];
340        let mut impl_map: HashMap<DefId, HashSet<DefId>> = HashMap::new();
341        let mut count_map: HashMap<DefId, usize> = HashMap::new();
342
343        // traverse from start node, if a node can achieve a reserved node,
344        // this node should be reserved
345        for node in self.graph.node_indices() {
346            if self.is_start_node_index(node) {
347                rap_trace!("initial node {:?}", self.graph[node]);
348                worklist.push_back(node);
349            }
350        }
351
352        while let Some(node) = worklist.pop_front() {
353            if visited[node.index()] {
354                continue;
355            }
356            visited[node.index()] = true;
357
358            match self.graph[node] {
359                DepNode::Api(fn_did, args) => {
360                    if fn_requires_monomorphization(fn_did, self.tcx) {
361                        let impl_entry = impl_map.entry(fn_did).or_default();
362                        let count_entry = count_map.entry(fn_did).or_default();
363                        let impls = mono::get_impls(self.tcx, fn_did, args);
364                        let size = impls
365                            .iter()
366                            .fold(0, |cnt, did| cnt + (!impl_entry.contains(did)) as usize);
367                        if *count_entry == 0 || size > 0 {
368                            *count_entry += 1;
369                            impls.iter().for_each(|did| {
370                                impl_entry.insert(*did);
371                            });
372                            reserved[node.index()] = true;
373                        }
374                    }
375                    for neighbor in self.graph.neighbors(node) {
376                        worklist.push_back(neighbor);
377                    }
378                }
379                DepNode::Ty(..) => {
380                    for edge in self.graph.edges_directed(node, Direction::Outgoing) {
381                        let weight = self.graph.edge_weight(edge.id()).unwrap();
382                        if let DepEdge::Transform { .. } | DepEdge::Arg { .. } = weight {
383                            worklist.push_back(edge.target());
384                        }
385                    }
386                }
387            }
388
389            if reserved[node.index()] {
390                rap_debug!(
391                    "[propagate_reserved] reserve: {:?}",
392                    self.graph.node_weight(node).unwrap()
393                );
394            }
395        }
396    }
397
398    pub fn minimal_select(
399        &mut self,
400        reserved: &mut [bool],
401        generic_map: &HashMap<DefId, HashSet<Mono<'tcx>>>,
402    ) {
403        let mut rng = rand::rng();
404        let mut reserved_map: HashMap<DefId, Vec<(GenericArgsRef<'tcx>, bool)>> = HashMap::new();
405
406        // transform into reserved map
407        for (fn_did, mono_set) in generic_map {
408            let entry = reserved_map.entry(*fn_did).or_default();
409            mono_set.into_iter().for_each(|mono| {
410                let args = self.tcx.mk_args(&mono.value);
411                entry.push((args, false));
412            });
413        }
414        // add all monomorphic APIs to API Graph, but select minimal set cover to be reserved
415        for (fn_did, monos) in &mut reserved_map {
416            select_minimal_set_cover(self.tcx, *fn_did, monos, &mut rng);
417            for (args, r) in monos {
418                if *r {
419                    let idx = self.get_index(DepNode::Api(*fn_did, args)).unwrap();
420                    reserved[idx.index()] = true;
421                }
422            }
423        }
424    }
425
426    pub fn prune_by_similarity(
427        &mut self,
428        generic_map: HashMap<DefId, HashSet<Mono<'tcx>>>,
429    ) -> Vec<bool> {
430        let (estimate, total) = self.estimate_coverage_distinct();
431        rap_info!(
432            "estimate API coverage before pruning: {:.2} ({}/{})",
433            estimate as f64 / total as f64,
434            estimate,
435            total
436        );
437
438        let mut visited = vec![false; self.graph.node_count()];
439        let mut reserved = vec![false; self.graph.node_count()];
440
441        // initialize reserved
442        // all non-generic API should be reserved
443        for idx in self.graph.node_indices() {
444            if let DepNode::Api(fn_did, _) = self.graph[idx] {
445                if !utils::fn_requires_monomorphization(fn_did, self.tcx) {
446                    reserved[idx.index()] = true;
447                }
448            }
449        }
450
451        // minimal set cover strategy
452        // self.minimal_select(&mut reserved, &generic_map);
453
454        // heuristic strategy
455        self.heuristic_select(&mut reserved);
456
457        // traverse from start node, if a node can achieve a reserved node,
458        // this node should be reserved as well
459        for node in self.graph.node_indices() {
460            if !visited[node.index()] && self.is_start_node_index(node) {
461                rap_trace!("start propagate from {:?}", self.graph[node]);
462                self.propagate_reserved(node, &mut visited, &mut reserved);
463            }
464        }
465
466        for node in self.graph.node_indices() {
467            if !visited[node.index()] {
468                rap_trace!("{:?} is unvisited", self.graph[node]);
469                self.propagate_reserved(node, &mut visited, &mut reserved);
470            }
471        }
472
473        reserved
474    }
475
476    pub fn reserve_nodes(&mut self, reserved: &[bool]) -> usize {
477        let mut count = 0;
478        for idx in (0..self.graph.node_count()).rev() {
479            if !reserved[idx] {
480                self.graph
481                    .remove_node(NodeIndex::new(idx))
482                    .expect("remove should not fail");
483                count += 1;
484            }
485        }
486        self.recache();
487        count
488    }
489
490    pub fn propagate_reserved(
491        &self,
492        node: NodeIndex,
493        visited: &mut [bool],
494        reserved: &mut [bool],
495    ) -> bool {
496        visited[node.index()] = true;
497
498        match self.graph[node] {
499            // Api should be reserved if must_reserve is true,
500            // or at least one its neighbor is reserved
501            DepNode::Api(fn_did, args) => {
502                for neighbor in self.graph.neighbors(node) {
503                    if !visited[neighbor.index()] {
504                        reserved[node.index()] |=
505                            self.propagate_reserved(neighbor, visited, reserved);
506                    }
507                }
508            }
509
510            // Ty should be reserved if at least one its neighbor is reserved
511            DepNode::Ty(..) => {
512                // self.graph.edges_directed(node, dir)
513                for neighbor in self.graph.neighbors(node) {
514                    if !visited[neighbor.index()] {
515                        self.propagate_reserved(neighbor, visited, reserved);
516                    }
517                    reserved[node.index()] |= reserved[neighbor.index()]
518                }
519            }
520        }
521
522        if reserved[node.index()] {
523            rap_trace!(
524                "[propagate_reserved] reserve: {:?}",
525                self.graph.node_weight(node).unwrap()
526            );
527        }
528        reserved[node.index()]
529    }
530}
531
532fn select_minimal_set_cover<'tcx>(
533    tcx: TyCtxt<'tcx>,
534    fn_did: DefId,
535    monos: &mut Vec<(ty::GenericArgsRef<'tcx>, bool)>,
536    rng: &mut impl Rng,
537) {
538    rap_debug!("select minimal set for: {}", tcx.def_path_str(fn_did));
539    let mut impl_vec = Vec::new();
540    let mut cmplx_vec = Vec::new();
541    for (args, _) in monos.iter() {
542        impl_vec.push(mono::get_impls(tcx, fn_did, args));
543        cmplx_vec.push(get_mono_complexity(args));
544    }
545
546    let mut selected_cnt = 0;
547    let mut complete = HashSet::new();
548    loop {
549        let mut current_max = 0;
550        let mut current_cmplx = usize::MAX;
551        let mut idx = 0;
552        for i in 0..impl_vec.len() {
553            let size = impl_vec[i]
554                .iter()
555                .fold(0, |cnt, did| cnt + (!complete.contains(did)) as usize);
556
557            if size > current_max || (size == current_max && cmplx_vec[i] < current_cmplx) {
558                current_max = size;
559                current_cmplx = cmplx_vec[i];
560                idx = i;
561            }
562        }
563        // though maybe all impls is empty, we have to select at least one
564        if current_max == 0 && selected_cnt > 0 {
565            break;
566        }
567        selected_cnt += 1;
568        monos[idx].1 = true;
569        rap_debug!("select: {:?}", monos[idx].0);
570        impl_vec[idx].iter().for_each(|did| {
571            complete.insert(*did);
572        });
573    }
574
575    // if selected_cnt == 0 {
576    //     let idx = rng.random_range(0..impl_vec.len());
577    //     rap_debug!("random select: {:?}", monos[idx].0);
578    //     monos[idx].1 = true;
579    // }
580}