Skip to main content

rapx/analysis/api_dependency/
mono.rs

1#![allow(warnings, unused)]
2
3use super::graph::TyWrapper;
4use super::utils::{self, fn_sig_with_generic_args};
5use crate::compat;
6use crate::limit::MAX_STEP_SET_SIZE;
7#[cfg(not(rapx_has_skip_norm_wip))]
8use crate::compat::SkipNormWip;
9use crate::helpers::def_path::path_str_def_id;
10use crate::{rap_debug, rap_trace};
11use rand::Rng;
12use rand::seq::SliceRandom;
13
14#[cfg(rapx_has_attr_ir)]
15use rustc_attr_ir::LangItem;
16#[cfg(all(not(rapx_has_attr_ir), not(rapx_ge_100)))]
17use rustc_hir::LangItem;
18#[cfg(all(not(rapx_has_attr_ir), rapx_ge_100))]
19use rustc_hir::attrs::lang_items::LangItem;
20use rustc_hir::def_id::DefId;
21use rustc_infer::infer::DefineOpaqueTypes;
22use rustc_infer::infer::{InferCtxt, TyCtxtInferExt};
23use rustc_infer::traits::{ImplSource, Obligation, ObligationCause};
24use rustc_middle::ty::{
25    self, GenericArgKind, GenericArgsRef, Ty, TyCtxt, TypeVisitableExt, TypingEnv,
26};
27use rustc_span::DUMMY_SP;
28use rustc_trait_selection::traits::query::evaluate_obligation::InferCtxtExt as _;
29use rustc_type_ir::InferCtxtLike;
30use std::collections::HashSet;
31
32#[derive(Clone, Debug, Hash, PartialEq, Eq)]
33pub struct Mono<'tcx> {
34    pub value: Vec<ty::GenericArg<'tcx>>,
35}
36
37impl<'tcx> FromIterator<ty::GenericArg<'tcx>> for Mono<'tcx> {
38    fn from_iter<T>(iter: T) -> Self
39    where
40        T: IntoIterator<Item = ty::GenericArg<'tcx>>,
41    {
42        Mono {
43            value: iter.into_iter().collect(),
44        }
45    }
46}
47
48impl<'tcx> Mono<'tcx> {
49    pub fn new(identity: &[ty::GenericArg<'tcx>]) -> Self {
50        Mono {
51            value: Vec::from(identity),
52        }
53    }
54
55    fn has_infer_types(&self) -> bool {
56        self.value.iter().any(|arg| match arg.kind() {
57            ty::GenericArgKind::Type(ty) => ty.has_infer_types(),
58            _ => false,
59        })
60    }
61
62    fn mut_arg_at(&mut self, idx: usize) -> &mut ty::GenericArg<'tcx> {
63        &mut self.value[idx]
64    }
65
66    fn merge(&self, other: &Mono<'tcx>, tcx: TyCtxt<'tcx>) -> Option<Mono<'tcx>> {
67        assert!(self.value.len() == other.value.len());
68        let mut res = Vec::new();
69        for i in 0..self.value.len() {
70            let arg = self.value[i];
71            let other_arg = other.value[i];
72            let new_arg = if let GenericArgKind::Type(ty) = arg.kind() {
73                let other_ty = other_arg.expect_ty();
74                if ty.is_ty_var() && other_ty.is_ty_var() {
75                    arg
76                } else if ty.is_ty_var() {
77                    other_arg
78                } else if other_ty.is_ty_var() {
79                    arg
80                } else if utils::is_ty_eq(ty, other_ty, tcx) {
81                    arg
82                } else {
83                    return None;
84                }
85            } else {
86                arg
87            };
88            res.push(new_arg);
89        }
90        Some(Mono { value: res })
91    }
92
93    fn fill_unbound_var(&self, tcx: TyCtxt<'tcx>) -> Vec<Mono<'tcx>> {
94        let candidates = get_unbound_generic_candidates(tcx);
95        let mut res = vec![self.clone()];
96        rap_trace!("fill unbound: {:?}", self);
97
98        for (i, arg) in self.value.iter().enumerate() {
99            if let GenericArgKind::Type(ty) = arg.kind() {
100                if ty.is_ty_var() {
101                    let mut last = Vec::new();
102                    std::mem::swap(&mut res, &mut last);
103                    last.into_iter().for_each(|mono| {
104                        for candidate in &candidates {
105                            let mut new_mono = mono.clone();
106                            *new_mono.mut_arg_at(i) = (*candidate).into();
107                            res.push(new_mono);
108                        }
109                    });
110                }
111            }
112        }
113        res
114    }
115}
116
117#[derive(Clone, Debug, Default)]
118pub struct MonoSet<'tcx> {
119    pub monos: Vec<Mono<'tcx>>,
120}
121
122impl<'tcx> MonoSet<'tcx> {
123    pub fn all(identity: &[ty::GenericArg<'tcx>]) -> MonoSet<'tcx> {
124        MonoSet {
125            monos: vec![Mono::new(identity)],
126        }
127    }
128
129    pub fn empty() -> MonoSet<'tcx> {
130        MonoSet { monos: Vec::new() }
131    }
132
133    pub fn count(&self) -> usize {
134        self.monos.len()
135    }
136
137    pub fn at(&self, no: usize) -> &Mono<'tcx> {
138        &self.monos[no]
139    }
140
141    pub fn is_empty(&self) -> bool {
142        self.monos.is_empty()
143    }
144
145    pub fn new() -> MonoSet<'tcx> {
146        MonoSet { monos: Vec::new() }
147    }
148
149    pub fn insert(&mut self, mono: Mono<'tcx>) {
150        self.monos.push(mono);
151    }
152
153    pub fn merge(&mut self, other: &MonoSet<'tcx>, tcx: TyCtxt<'tcx>) -> MonoSet<'tcx> {
154        let mut res = MonoSet::new();
155
156        for args in self.monos.iter() {
157            for other_args in other.monos.iter() {
158                let merged = args.merge(other_args, tcx);
159                if let Some(mono) = merged {
160                    res.insert(mono);
161                }
162            }
163        }
164        res
165    }
166
167    // if the unbound generic type is still exist (this could happen
168    // if `T` has no trait bounds at all)
169    // we substitute the unbound generic type with predefined type candidates
170    fn instantiate_unbound(&self, tcx: TyCtxt<'tcx>) -> Self {
171        let mut res = MonoSet::new();
172        for mono in &self.monos {
173            let filled = mono.fill_unbound_var(tcx);
174            res.monos.extend(filled);
175        }
176        res
177    }
178
179    fn erase_region_var(&mut self, tcx: TyCtxt<'tcx>) {
180        for mono in &mut self.monos {
181            mono.value
182                .iter_mut()
183                .for_each(|arg| *arg = tcx.erase_and_anonymize_regions(*arg))
184        }
185    }
186
187    pub fn filter(mut self, f: impl Fn(&Mono<'tcx>) -> bool) -> Self {
188        self.monos.retain(|args| f(args));
189        self
190    }
191
192    pub fn random_sample<R: Rng>(&mut self, rng: &mut R) {
193        if self.monos.len() <= MAX_STEP_SET_SIZE {
194            return;
195        }
196        self.monos.shuffle(rng);
197        self.monos.truncate(MAX_STEP_SET_SIZE);
198    }
199}
200
201/// Resolve inference variables to concrete values after unification.  The
202/// method was renamed `resolve_vars_if_possible` → `deeply_resolve_ignoring_regions`
203/// in nightly 2026-09-11.
204#[cfg(rapx_has_deeply_resolve_ignoring_regions)]
205fn resolve_var<'tcx, T: rustc_middle::ty::TypeFoldable<TyCtxt<'tcx>>>(
206    infcx: &InferCtxt<'tcx>,
207    value: T,
208) -> T {
209    infcx.deeply_resolve_ignoring_regions(value)
210}
211
212#[cfg(not(rapx_has_deeply_resolve_ignoring_regions))]
213fn resolve_var<'tcx, T: rustc_middle::ty::TypeFoldable<TyCtxt<'tcx>>>(
214    infcx: &InferCtxt<'tcx>,
215    value: T,
216) -> T {
217    infcx.resolve_vars_if_possible(value)
218}
219
220/// try to unfiy lhs = rhs,
221/// e.g.,
222/// try_unify(Vec<T>, Vec<i32>, ...) = Some(i32)
223/// try_unify(Vec<T>, i32, ...) = None
224fn unify_ty<'tcx>(
225    lhs: Ty<'tcx>,
226    rhs: Ty<'tcx>,
227    identity: &[ty::GenericArg<'tcx>],
228    infcx: &InferCtxt<'tcx>,
229    cause: &ObligationCause<'tcx>,
230    param_env: ty::ParamEnv<'tcx>,
231) -> Option<Mono<'tcx>> {
232    // rap_info!("check {} = {}", lhs, rhs);
233    infcx.probe(|_| {
234        match infcx
235            .at(cause, param_env)
236            .eq(DefineOpaqueTypes::Yes, lhs, rhs)
237        {
238            Ok(_infer_ok) => {
239                // rap_trace!("[infer_ok] {} = {} : {:?}", lhs, rhs, infer_ok);
240                let mono = identity
241                    .iter()
242                    .map(|arg| match arg.kind() {
243                        ty::GenericArgKind::Lifetime(region) => resolve_var(infcx, region).into(),
244                        ty::GenericArgKind::Type(ty) => resolve_var(infcx, ty).into(),
245                        ty::GenericArgKind::Const(ct) => resolve_var(infcx, ct).into(),
246                    })
247                    .collect();
248                Some(mono)
249            }
250            Err(_e) => {
251                // rap_trace!("[infer_err] {} = {} : {:?}", lhs, rhs, e);
252                None
253            }
254        }
255    })
256}
257
258fn is_args_fit_trait_bound<'tcx>(
259    fn_did: DefId,
260    args: &[ty::GenericArg<'tcx>],
261    tcx: TyCtxt<'tcx>,
262) -> bool {
263    let args = tcx.mk_args(args);
264    rap_trace!(
265        "fn: {:?} args: {:?} identity: {:?}",
266        fn_did,
267        args,
268        ty::GenericArgs::identity_for_item(tcx, fn_did)
269    );
270    let infcx = tcx.infer_ctxt().build(ty::TypingMode::PostAnalysis);
271    let param_env = tcx.param_env(fn_did);
272    let pred = crate::compat::predicates_of(tcx, fn_did);
273    let inst_pred = pred.instantiate(tcx, args);
274    rap_trace!(
275        "[trait bound] check {}",
276        tcx.def_path_str_with_args(fn_did, args)
277    );
278
279    #[cfg(not(rapx_ge_100))]
280    let iter = inst_pred.predicates.iter();
281    #[cfg(rapx_ge_100)]
282    let iter = inst_pred.clauses.iter();
283    for pred in iter {
284        #[cfg(rapx_ge_99)]
285        let pred = pred.skip_norm_wip();
286        let obligation = Obligation::new(
287            tcx,
288            ObligationCause::dummy(),
289            param_env,
290            pred.as_predicate(),
291        );
292        rap_trace!("[trait bound] check pred: {:?}", pred);
293
294        let res = infcx.evaluate_obligation(&obligation);
295        match res {
296            Ok(eva) => {
297                if !eva.may_apply() {
298                    rap_trace!("[trait bound] check fail for {pred:?}");
299                    return false;
300                }
301            }
302            Err(_) => {
303                rap_trace!("[trait bound] check fail for {pred:?}");
304                return false;
305            }
306        }
307    }
308    rap_trace!("[trait bound] check succ");
309    true
310}
311
312fn is_fn_solvable<'tcx>(fn_did: DefId, tcx: TyCtxt<'tcx>) -> bool {
313    let predicates = crate::compat::predicates_of(tcx, fn_did);
314    #[cfg(not(rapx_ge_100))]
315    let iter = predicates.instantiate_identity(tcx).predicates;
316    #[cfg(rapx_ge_100)]
317    let iter = predicates.instantiate_identity(tcx).clauses;
318    for pred in iter {
319        #[cfg(rapx_ge_99)]
320        let pred = pred.skip_norm_wip();
321        if let Some(pred) = pred.as_trait_clause() {
322            let trait_did = pred.skip_binder().trait_ref.def_id;
323            if tcx.is_lang_item(trait_did, LangItem::Fn)
324                || tcx.is_lang_item(trait_did, LangItem::FnMut)
325                || tcx.is_lang_item(trait_did, LangItem::FnOnce)
326            {
327                return false;
328            }
329        }
330    }
331    true
332}
333
334fn get_mono_set<'tcx>(
335    fn_did: DefId,
336    available_ty: &HashSet<TyWrapper<'tcx>>,
337    tcx: TyCtxt<'tcx>,
338) -> MonoSet<'tcx> {
339    let mut rng = rand::rng();
340
341    // sample from reachable types
342    rap_debug!("[get_mono_set] solve {}", tcx.def_path_str(fn_did));
343    let identity = ty::GenericArgs::identity_for_item(tcx, fn_did);
344    let infcx = tcx
345        .infer_ctxt()
346        .ignoring_regions()
347        .build(ty::TypingMode::PostAnalysis);
348    let param_env = tcx.param_env(fn_did);
349    let dummy_cause = ObligationCause::dummy();
350    let fresh_args = infcx.fresh_args_for_item(DUMMY_SP, fn_did);
351    // this replace generic types in fn_sig to infer var, e.g. fn(Vec<T>, i32) => fn(Vec<?0>, i32)
352    let fn_sig = fn_sig_with_generic_args(fn_did, fresh_args, tcx);
353    let identity_fnsig = fn_sig_with_generic_args(fn_did, identity, tcx);
354    let generics = tcx.generics_of(fn_did);
355
356    // print fresh_args for debugging
357    for i in 0..fresh_args.len() {
358        rap_trace!(
359            "[get_mono_set] arg#{}: {:?} -> {:?}",
360            i,
361            generics.param_at(i, tcx).name,
362            fresh_args[i]
363        );
364    }
365
366    let mut s = MonoSet::all(&fresh_args);
367
368    rap_trace!("[get_mono_set] initialize s: {:?}", s);
369
370    for (no, input_ty) in fn_sig.inputs().iter().enumerate() {
371        if !input_ty.has_infer_types() {
372            continue;
373        }
374        rap_debug!(
375            "[get_mono_set] input_ty#{}: {}",
376            no,
377            identity_fnsig.inputs()[no]
378        );
379
380        let reachable_set = available_ty
381            .iter()
382            .fold(MonoSet::new(), |mut reachable_set, ty| {
383                if let Some(mono) = unify_ty(
384                    *input_ty,
385                    (*ty).into(),
386                    &fresh_args,
387                    &infcx,
388                    &dummy_cause,
389                    param_env,
390                ) {
391                    reachable_set.insert(mono);
392                }
393                reachable_set
394            });
395        // reachable_set.random_sample(&mut rng);
396        rap_debug!(
397            "[get_mono_set] size of s: {}, size of input: {}",
398            s.count(),
399            reachable_set.count()
400        );
401        rap_trace!("[get_mono_set] input = {:?}", reachable_set);
402        s = s.merge(&reachable_set, tcx);
403        s.random_sample(&mut rng);
404        rap_trace!("[get_mono_set] after merge s = {:?}", reachable_set);
405    }
406
407    rap_debug!(
408        "[get_mono_set] after input filter, size of s: {}",
409        s.count()
410    );
411
412    let mut res = MonoSet::new();
413
414    for mono in s.monos {
415        solve_unbound_type_generics(
416            fn_did,
417            mono,
418            &mut res,
419            // &fresh_args,
420            &infcx,
421            &dummy_cause,
422            param_env,
423            tcx,
424        );
425    }
426
427    // erase infer region var
428    res.erase_region_var(tcx);
429
430    // if there is still unbound generic type, we try to instantiate it with predefined candidates
431    res.instantiate_unbound(tcx)
432}
433
434fn solve_unbound_type_generics<'tcx>(
435    did: DefId,
436    mono: Mono<'tcx>,
437    res: &mut MonoSet<'tcx>,
438    infcx: &InferCtxt<'tcx>,
439    cause: &ObligationCause<'tcx>,
440    param_env: ty::ParamEnv<'tcx>,
441    tcx: TyCtxt<'tcx>,
442) {
443    if !mono.has_infer_types() {
444        res.insert(mono);
445        return;
446    }
447    let args = tcx.mk_args(&mono.value);
448    let preds = crate::compat::predicates_of(tcx, did);
449    let preds = preds.instantiate(tcx, args);
450    let mut mset = MonoSet::all(args);
451    rap_debug!("[solve_unbound] did = {did:?}, mset={mset:?}");
452    #[cfg(not(rapx_ge_100))]
453    let pred_iter = preds.predicates.iter();
454    #[cfg(rapx_ge_100)]
455    let pred_iter = preds.clauses.iter();
456    for pred in pred_iter {
457        rap_debug!("[solve_unbound] pred = {:?}", pred);
458        #[cfg(rapx_ge_99)]
459        let pred = pred.skip_norm_wip();
460        if let Some(trait_pred) = pred.as_trait_clause() {
461            let trait_pred = trait_pred.skip_binder();
462
463            rap_trace!("[solve_unbound] pred: {:?}", trait_pred);
464
465            let trait_def_id = trait_pred.trait_ref.def_id;
466            // ignore Sized trait
467            if tcx.is_lang_item(trait_def_id, LangItem::Sized)
468                || tcx.is_lang_item(trait_def_id, LangItem::Copy)
469            {
470                continue;
471            }
472
473            let mut p = MonoSet::new();
474
475            for impl_did in tcx.all_impls(trait_def_id)
476            // .chain(tcx.inherent_impls(trait_def_id).iter().map(|did| *did))
477            {
478                // format: <arg0 as Trait<arg1, arg2>>
479                let impl_trait_ref = tcx.impl_trait_ref(impl_did).skip_binder();
480
481                // filter irrelevant implementation. We only consider implementation that:
482                // 1. it is local
483                // 2. it is not local, but its' self_ty is a primitive
484                if !impl_did.is_local() && !impl_trait_ref.self_ty().is_primitive() {
485                    continue;
486                }
487
488                if let Some(mono) = unify_trait(
489                    trait_pred.trait_ref,
490                    impl_trait_ref,
491                    args,
492                    &infcx,
493                    &cause,
494                    param_env,
495                    tcx,
496                ) {
497                    p.insert(mono);
498                }
499            }
500            mset = mset.merge(&p, tcx);
501            rap_trace!("[solve_unbound] mset: {:?}", mset);
502        }
503    }
504
505    rap_trace!("[solve_unbound] (final) mset: {:?}", mset);
506    for mono in mset.monos {
507        res.insert(mono);
508    }
509}
510
511/// only handle the case that rhs does not have any infer types
512/// e.g., `<T as Into<U>> == <Foo as Into<Bar>> => Some(T=Foo, U=Bar))`
513fn unify_trait<'tcx>(
514    lhs: ty::TraitRef<'tcx>,
515    rhs: ty::TraitRef<'tcx>,
516    identity: &[ty::GenericArg<'tcx>],
517    infcx: &InferCtxt<'tcx>,
518    cause: &ObligationCause<'tcx>,
519    param_env: ty::ParamEnv<'tcx>,
520    tcx: TyCtxt<'tcx>,
521) -> Option<Mono<'tcx>> {
522    rap_trace!("[unify_trait] lhs: {:?}, rhs: {:?}", lhs, rhs);
523    if lhs.def_id != rhs.def_id {
524        return None;
525    }
526
527    assert!(lhs.args.len() == rhs.args.len());
528    let mut s = Mono::new(identity);
529    for (lhs_arg, rhs_arg) in lhs.args.iter().zip(rhs.args.iter()) {
530        if let (GenericArgKind::Type(lhs_ty), GenericArgKind::Type(rhs_ty)) =
531            (lhs_arg.kind(), rhs_arg.kind())
532        {
533            if rhs_ty.has_infer_types() || rhs_ty.has_param() {
534                // if rhs has infer types, we cannot unify it with lhs
535                return None;
536            }
537            let mono = unify_ty(lhs_ty, rhs_ty, identity, infcx, cause, param_env)?;
538            rap_trace!("[unify_trait] unified mono: {:?}", mono);
539            s = s.merge(&mono, tcx)?;
540        }
541    }
542    Some(s)
543}
544
545pub fn resolve_mono_apis<'tcx>(
546    fn_did: DefId,
547    available_ty: &HashSet<TyWrapper<'tcx>>,
548    tcx: TyCtxt<'tcx>,
549) -> MonoSet<'tcx> {
550    // 1. check solvable condition
551    if !is_fn_solvable(fn_did, tcx) {
552        return MonoSet::empty();
553    }
554
555    // 2. get mono set from available types
556    let ret = get_mono_set(fn_did, &available_ty, tcx);
557
558    // 3. check trait bound & ty is stable
559    let ret = ret.filter(|mono| {
560        is_args_fit_trait_bound(fn_did, &mono.value, tcx)
561            && mono.value.iter().all(|arg| {
562                arg.as_type()
563                    .map_or(true, |ty| !utils::is_ty_unstable(ty, tcx))
564            })
565    });
566
567    rap_debug!(
568        "[resolve_mono_apis] fn_did: {:?}, size of mono: {:?}",
569        fn_did,
570        ret.count()
571    );
572
573    ret
574}
575
576/// if type parameter is unbound, e.g., `T` in `fn foo<T>()`,
577/// we use some predefined types to substitute it
578pub fn get_unbound_generic_candidates<'tcx>(tcx: TyCtxt<'tcx>) -> Vec<ty::Ty<'tcx>> {
579    vec![
580        tcx.types.bool,
581        tcx.types.char,
582        tcx.types.u8,
583        tcx.types.i8,
584        tcx.types.i32,
585        tcx.types.u32,
586        // tcx.types.i64,
587        // tcx.types.u64,
588        tcx.types.f32,
589        // tcx.types.f64,
590        Ty::new_imm_ref(
591            tcx,
592            tcx.lifetimes.re_erased,
593            Ty::new_slice(tcx, tcx.types.u8),
594        ),
595        Ty::new_mut_ref(
596            tcx,
597            tcx.lifetimes.re_erased,
598            Ty::new_slice(tcx, tcx.types.u8),
599        ),
600    ]
601}
602
603// calculate the complexity of monomorphic solution,
604// complexity = sum of complexity of each type argument
605pub fn get_mono_complexity<'tcx>(args: &GenericArgsRef<'tcx>) -> usize {
606    args.iter().fold(0, |acc, arg| {
607        if let Some(ty) = arg.as_type() {
608            acc + utils::ty_complexity(ty)
609        } else {
610            acc
611        }
612    })
613}
614
615pub fn get_impls<'tcx>(
616    tcx: TyCtxt<'tcx>,
617    fn_did: DefId,
618    args: GenericArgsRef<'tcx>,
619) -> HashSet<DefId> {
620    rap_debug!(
621        "get impls for fn: {:?} args: {:?}",
622        tcx.def_path_str_with_args(fn_did, args),
623        args
624    );
625    let mut impls = HashSet::new();
626    let preds = crate::compat::predicates_of(tcx, fn_did);
627    let preds = preds.instantiate(tcx, args);
628    for (pred, _) in preds {
629        #[cfg(rapx_ge_99)]
630        let pred = pred.skip_norm_wip();
631        if let Some(trait_pred) = pred.as_trait_clause() {
632            let trait_ref: rustc_type_ir::TraitRef<TyCtxt<'tcx>> = tcx
633                .liberate_late_bound_regions(fn_did, trait_pred)
634                .trait_ref;
635
636            let res = tcx.codegen_select_candidate(
637                TypingEnv::fully_monomorphized().as_query_input(trait_ref),
638            );
639            if let Ok(source) = res {
640                match source {
641                    ImplSource::UserDefined(data) => {
642                        if data.impl_def_id.is_local() {
643                            impls.insert(data.impl_def_id);
644                        }
645                    }
646                    _ => {}
647                }
648            }
649            // rap_debug!("{:?} => {:?}", trait_ref, res);
650        }
651    }
652    rap_trace!("fn: {:?} args: {:?} impls: {:?}", fn_did, args, impls);
653    impls
654}