Skip to main content

rapx/analysis/api_dependency/graph/
ty_wrapper.rs

1use std::hash::Hash;
2use std::ops::Deref;
3
4use super::transform::TransformKind;
5use rustc_infer::infer::TyCtxtInferExt;
6use rustc_infer::traits::{Obligation, ObligationCause};
7use rustc_middle::traits;
8use rustc_middle::ty::{self, Ty, TyCtxt};
9use rustc_trait_selection::infer::InferCtxtExt;
10use rustc_trait_selection::traits::query::evaluate_obligation::InferCtxtExt as _;
11
12/// TyWrapper is a wrapper of rustc_middle::ty::Ty
13#[derive(Clone, Copy, Eq, Debug)]
14pub struct TyWrapper<'tcx> {
15    ty: Ty<'tcx>,
16}
17
18impl<'tcx> TyWrapper<'tcx> {
19    pub fn ty(&self) -> Ty<'tcx> {
20        self.ty
21    }
22
23    pub fn into_ref(&self, tcx: TyCtxt<'tcx>) -> TyWrapper<'tcx> {
24        Ty::new_ref(tcx, tcx.lifetimes.re_erased, self.ty, ty::Mutability::Not).into()
25    }
26
27    pub fn into_ref_mut(&self, tcx: TyCtxt<'tcx>) -> TyWrapper<'tcx> {
28        Ty::new_ref(tcx, tcx.lifetimes.re_erased, self.ty, ty::Mutability::Mut).into()
29    }
30
31    pub fn transform(&self, kind: TransformKind, tcx: TyCtxt<'tcx>) -> TyWrapper<'tcx> {
32        match kind {
33            TransformKind::Ref(mutability) => {
34                let ty = match mutability {
35                    ty::Mutability::Not => self.into_ref(tcx),
36                    ty::Mutability::Mut => self.into_ref_mut(tcx),
37                };
38                ty
39            }
40            _ => {
41                todo!();
42            }
43        }
44    }
45}
46
47impl<'tcx> From<Ty<'tcx>> for TyWrapper<'tcx> {
48    fn from(ty: ty::Ty<'tcx>) -> TyWrapper<'tcx> {
49        TyWrapper { ty }
50    }
51}
52
53impl<'tcx> Into<Ty<'tcx>> for TyWrapper<'tcx> {
54    fn into(self) -> Ty<'tcx> {
55        self.ty
56    }
57}
58
59impl PartialEq for TyWrapper<'_> {
60    fn eq(&self, other: &Self) -> bool {
61        eq_ty(self.ty, other.ty)
62    }
63}
64
65impl Hash for TyWrapper<'_> {
66    fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
67        hash_ty(self.ty, state, &mut 0);
68    }
69}
70
71fn eq_ty<'tcx>(lhs: Ty<'tcx>, rhs: Ty<'tcx>) -> bool {
72    match (lhs.kind(), rhs.kind()) {
73        (ty::TyKind::Adt(adt_def1, generic_arg1), ty::TyKind::Adt(adt_def2, generic_arg2)) => {
74            if adt_def1.did() != adt_def2.did() {
75                return false;
76            }
77            for (arg1, arg2) in generic_arg1.iter().zip(generic_arg2.iter()) {
78                match (arg1.kind(), arg2.kind()) {
79                    (ty::GenericArgKind::Lifetime(_), ty::GenericArgKind::Lifetime(_)) => continue,
80                    (ty::GenericArgKind::Type(ty1), ty::GenericArgKind::Type(ty2)) => {
81                        if !eq_ty(ty1, ty2) {
82                            return false;
83                        }
84                    }
85                    (ty::GenericArgKind::Const(ct1), ty::GenericArgKind::Const(ct2)) => {
86                        if ct1 != ct2 {
87                            return false;
88                        }
89                    }
90                    _ => return false,
91                }
92            }
93            true
94        }
95        (
96            ty::TyKind::RawPtr(inner_ty1, mutability1),
97            ty::TyKind::RawPtr(inner_ty2, mutability2),
98        )
99        | (
100            ty::TyKind::Ref(_, inner_ty1, mutability1),
101            ty::TyKind::Ref(_, inner_ty2, mutability2),
102        ) => mutability1 == mutability2 && eq_ty(*inner_ty1, *inner_ty2),
103        (ty::TyKind::Array(inner_ty1, len1), ty::TyKind::Array(inner_ty2, len2)) => {
104            if len1 != len2 {
105                return false;
106            }
107            eq_ty(*inner_ty1, *inner_ty2)
108        }
109        (ty::TyKind::Pat(inner_ty1, _), ty::TyKind::Pat(inner_ty2, _))
110        | (ty::TyKind::Slice(inner_ty1), ty::TyKind::Slice(inner_ty2)) => {
111            eq_ty(*inner_ty1, *inner_ty2)
112        }
113        (ty::TyKind::Tuple(tys1), ty::TyKind::Tuple(tys2)) => {
114            if tys1.len() != tys2.len() {
115                return false;
116            }
117            tys1.iter()
118                .zip(tys2.iter())
119                .all(|(ty1, ty2)| eq_ty(ty1, ty2))
120        }
121        _ => lhs == rhs,
122    }
123}
124
125fn traverse_ty_with_lifetime<'tcx, F: Fn(ty::Region, usize)>(ty: Ty<'tcx>, no: &mut usize, f: &F) {
126    match ty.kind() {
127        ty::TyKind::Adt(adt_def, generic_arg) => {
128            for arg in generic_arg.iter() {
129                match arg.kind() {
130                    ty::GenericArgKind::Lifetime(lt) => {
131                        *no = *no + 1;
132                        f(lt, *no);
133                    }
134                    ty::GenericArgKind::Type(ty) => {
135                        traverse_ty_with_lifetime(ty, no, f);
136                    }
137                    ty::GenericArgKind::Const(ct) => {}
138                }
139            }
140        }
141
142        ty::TyKind::RawPtr(inner_ty, mutability) => {
143            traverse_ty_with_lifetime(*inner_ty, no, f);
144        }
145
146        ty::TyKind::Ref(region, inner_ty, mutability) => {
147            *no = *no + 1;
148            f(*region, *no);
149            traverse_ty_with_lifetime(*inner_ty, no, f);
150        }
151        ty::TyKind::Array(inner_ty, _)
152        | ty::TyKind::Pat(inner_ty, _)
153        | ty::TyKind::Slice(inner_ty) => {
154            traverse_ty_with_lifetime(*inner_ty, no, f);
155        }
156        ty::TyKind::Tuple(tys) => {
157            for inner_ty in tys.iter() {
158                traverse_ty_with_lifetime(inner_ty, no, f);
159            }
160        }
161        _ => {
162            unreachable!("unexpected ty kind");
163        }
164    }
165}
166
167// hashing Ty<'tcx>, but ignore the difference of lifetimes
168fn hash_ty<'tcx, H: std::hash::Hasher>(ty: Ty<'tcx>, state: &mut H, no: &mut usize) {
169    std::mem::discriminant(ty.kind()).hash(state);
170
171    // hash the content
172    match ty.kind() {
173        ty::TyKind::Adt(adt_def, generic_arg) => {
174            adt_def.did().hash(state);
175            for arg in generic_arg.iter() {
176                match arg.kind() {
177                    ty::GenericArgKind::Lifetime(lt) => {
178                        *no = *no + 1;
179                        no.hash(state);
180                    }
181                    ty::GenericArgKind::Type(ty) => {
182                        hash_ty(ty, state, no);
183                    }
184                    ty::GenericArgKind::Const(ct) => {
185                        ct.hash(state);
186                    }
187                }
188            }
189        }
190
191        ty::TyKind::RawPtr(inner_ty, mutability) => {
192            mutability.hash(state);
193            hash_ty(*inner_ty, state, no);
194        }
195        ty::TyKind::Ref(_, inner_ty, mutability) => {
196            mutability.hash(state);
197            *no = *no + 1;
198            no.hash(state);
199            hash_ty(*inner_ty, state, no);
200        }
201        ty::TyKind::Array(inner_ty, _) | ty::TyKind::Slice(inner_ty) => {
202            hash_ty(*inner_ty, state, no);
203        }
204        ty::TyKind::Tuple(tys) => {
205            for inner_ty in tys.iter() {
206                hash_ty(inner_ty, state, no);
207            }
208        }
209        _ => {
210            ty.hash(state);
211        }
212    }
213}
214
215pub fn desc_ty_str<'tcx>(ty: Ty<'tcx>, no: &mut usize, tcx: TyCtxt<'tcx>) -> String {
216    match ty.kind() {
217        ty::TyKind::Adt(adt_def, generic_arg) => {
218            let mut ty_str = tcx.def_path_str(adt_def.did());
219            if !generic_arg.is_empty() {
220                ty_str += "<";
221                ty_str += &generic_arg
222                    .iter()
223                    .map(|arg| match arg.kind() {
224                        ty::GenericArgKind::Lifetime(lt) => {
225                            let current_no = *no;
226                            *no = *no + 1;
227                            format!("'#{:?}", current_no)
228                        }
229                        ty::GenericArgKind::Type(ty) => desc_ty_str(ty, no, tcx),
230                        ty::GenericArgKind::Const(ct) => format!("{:?}", ct),
231                    })
232                    .collect::<Vec<String>>()
233                    .join(", ");
234                ty_str += ">";
235            }
236            ty_str
237        }
238
239        ty::TyKind::RawPtr(inner_ty, mutability) => {
240            format!(
241                "*{} {}",
242                mutability.ptr_str(),
243                desc_ty_str(*inner_ty, no, tcx)
244            )
245        }
246        ty::TyKind::Ref(_, inner_ty, mutability) => {
247            let current_no = *no;
248            *no = *no + 1;
249            format!(
250                "&'#{} {}{}",
251                current_no,
252                mutability.prefix_str(),
253                desc_ty_str(*inner_ty, no, tcx)
254            )
255        }
256        ty::TyKind::Array(inner_ty, len) => {
257            format!("[{};{}]", desc_ty_str(*inner_ty, no, tcx), len)
258        }
259
260        ty::TyKind::Slice(inner_ty) => {
261            format!("[{}]", desc_ty_str(*inner_ty, no, tcx))
262        }
263        ty::TyKind::Tuple(tys) => format!(
264            "({})",
265            tys.iter()
266                .map(|ty| desc_ty_str(ty, no, tcx,))
267                .collect::<Vec<String>>()
268                .join(", "),
269        ),
270        ty::TyKind::Pat(inner_ty, _) => {
271            unreachable!();
272        }
273        _ => format!("{:?}", ty),
274    }
275}
276
277impl<'tcx> TyWrapper<'tcx> {
278    pub fn desc_str(&self, tcx: TyCtxt<'tcx>) -> String {
279        desc_ty_str(self.ty, &mut 0, tcx)
280    }
281}