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 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#[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
220fn 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 infcx.probe(|_| {
234 match infcx
235 .at(cause, param_env)
236 .eq(DefineOpaqueTypes::Yes, lhs, rhs)
237 {
238 Ok(_infer_ok) => {
239 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 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 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 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 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 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 &infcx,
421 &dummy_cause,
422 param_env,
423 tcx,
424 );
425 }
426
427 res.erase_region_var(tcx);
429
430 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 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 {
478 let impl_trait_ref = tcx.impl_trait_ref(impl_did).skip_binder();
480
481 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
511fn 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 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 if !is_fn_solvable(fn_did, tcx) {
552 return MonoSet::empty();
553 }
554
555 let ret = get_mono_set(fn_did, &available_ty, tcx);
557
558 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
576pub 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.f32,
589 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
603pub 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 }
651 }
652 rap_trace!("fn: {:?} args: {:?} impls: {:?}", fn_did, args, impls);
653 impls
654}