1mod pat_util;
9
10pub(crate) mod pat_analysis;
11
12use hir_def::{
13 AdtId, EnumVariantId, LocalFieldId, Lookup, VariantId,
14 expr_store::{Body, path::Path},
15 hir::PatId,
16 item_tree::FieldsShape,
17 signatures::{StructSignature, UnionSignature},
18};
19use hir_expand::name::Name;
20use rustc_type_ir::inherent::IntoKind;
21use span::Edition;
22use stdx::{always, never, variance::PhantomCovariantLifetime};
23
24use crate::{
25 ByRef, InferenceResult,
26 db::HirDatabase,
27 display::{HirDisplay, HirDisplayError, HirFormatter},
28 infer::BindingMode,
29 next_solver::{GenericArgs, Mutability, Ty, TyKind},
30};
31
32use self::pat_util::EnumerateAndAdjustIterator;
33
34#[derive(Clone, Debug)]
35pub(crate) enum PatternError {
36 Unimplemented,
37 UnexpectedType,
38 UnresolvedVariant,
39 MissingField,
40 ExtraFields,
41}
42
43#[derive(Clone, Debug, PartialEq)]
44pub(crate) struct FieldPat<'db> {
45 pub(crate) field: LocalFieldId,
46 pub(crate) pattern: Pat<'db>,
47}
48
49#[derive(Clone, Debug, PartialEq)]
50pub(crate) struct Pat<'db> {
51 pub(crate) ty: Ty<'db>,
52 pub(crate) kind: Box<PatKind<'db>>,
53}
54
55#[derive(Clone, Debug, PartialEq)]
57pub(crate) enum PatKind<'db> {
58 Wild,
59 Never,
60
61 Binding {
63 name: Name,
64 subpattern: Option<Pat<'db>>,
65 },
66
67 Variant {
70 substs: GenericArgs<'db>,
71 enum_variant: EnumVariantId,
72 subpatterns: Vec<FieldPat<'db>>,
73 },
74
75 Leaf {
78 subpatterns: Vec<FieldPat<'db>>,
79 },
80
81 Deref {
83 subpattern: Pat<'db>,
84 },
85
86 LiteralBool {
88 value: bool,
89 },
90
91 Or {
94 pats: Vec<Pat<'db>>,
95 },
96}
97
98pub(crate) struct PatCtxt<'a, 'db> {
99 db: &'db dyn HirDatabase,
100 infer: &'db InferenceResult<'db>,
101 body: &'a Body,
102 pub(crate) errors: Vec<PatternError>,
103}
104
105impl<'a, 'db> PatCtxt<'a, 'db> {
106 pub(crate) fn new(
107 db: &'db dyn HirDatabase,
108 infer: &'db InferenceResult<'db>,
109 body: &'a Body,
110 ) -> Self {
111 Self { db, infer, body, errors: Vec::new() }
112 }
113
114 pub(crate) fn lower_pattern(&mut self, pat: PatId) -> Pat<'db> {
115 let unadjusted_pat = self.lower_pattern_unadjusted(pat);
121 self.infer.pat_adjustments.get(&pat).map(|it| &**it).unwrap_or_default().iter().rev().fold(
122 unadjusted_pat,
123 |subpattern, ref_ty| Pat {
124 ty: ref_ty.source.as_ref(),
125 kind: Box::new(PatKind::Deref { subpattern }),
126 },
127 )
128 }
129
130 fn lower_pattern_unadjusted(&mut self, pat: PatId) -> Pat<'db> {
131 let mut ty = self.infer.pat_ty(pat);
132 let variant = self.infer.variant_resolution_for_pat(pat);
133
134 let kind = match self.body[pat] {
135 hir_def::hir::Pat::Wild => PatKind::Wild,
136
137 hir_def::hir::Pat::Lit(expr) => self.lower_lit(expr),
138
139 hir_def::hir::Pat::Path(ref path) => {
140 return self.lower_path(pat, path);
141 }
142
143 hir_def::hir::Pat::Tuple { ref args, ellipsis } => {
144 let arity = match ty.kind() {
145 TyKind::Tuple(tys) => tys.len(),
146 _ => {
147 never!("unexpected type for tuple pattern: {:?}", ty);
148 self.errors.push(PatternError::UnexpectedType);
149 return Pat { ty, kind: PatKind::Wild.into() };
150 }
151 };
152 let subpatterns = self.lower_tuple_subpats(args, arity, ellipsis);
153 PatKind::Leaf { subpatterns }
154 }
155
156 hir_def::hir::Pat::Bind { id, subpat, .. } => {
157 let bm = self.infer.binding_modes[pat];
158 ty = self.infer.binding_ty(id);
159 let name = &self.body[id].name;
160 match (bm, ty.kind()) {
161 (BindingMode(ByRef::Yes(_), _), TyKind::Ref(_, rty, _)) => ty = rty,
162 (BindingMode(ByRef::Yes(_), _), _) => {
163 never!(
164 "`ref {}` has wrong type {:?}",
165 name.display(self.db, Edition::LATEST),
166 ty
167 );
168 self.errors.push(PatternError::UnexpectedType);
169 return Pat { ty, kind: PatKind::Wild.into() };
170 }
171 _ => (),
172 }
173 PatKind::Binding { name: name.clone(), subpattern: self.lower_opt_pattern(subpat) }
174 }
175
176 hir_def::hir::Pat::TupleStruct { ref args, ellipsis, .. } if variant.is_some() => {
177 let expected_len = variant.unwrap().fields(self.db).fields().len();
178 let subpatterns = self.lower_tuple_subpats(args, expected_len, ellipsis);
179 self.lower_variant_or_leaf(pat, ty, subpatterns)
180 }
181
182 hir_def::hir::Pat::Record { ref args, .. } if variant.is_some() => {
183 let variant_data = variant.unwrap().fields(self.db);
184 let subpatterns = args
185 .iter()
186 .map(|field| {
187 variant_data.field(&field.name).map(|lfield_id| FieldPat {
189 field: lfield_id,
190 pattern: self.lower_pattern(field.pat),
191 })
192 })
193 .collect();
194 match subpatterns {
195 Some(subpatterns) => self.lower_variant_or_leaf(pat, ty, subpatterns),
196 None => {
197 self.errors.push(PatternError::MissingField);
198 PatKind::Wild
199 }
200 }
201 }
202 hir_def::hir::Pat::TupleStruct { .. } | hir_def::hir::Pat::Record { .. } => {
203 self.errors.push(PatternError::UnresolvedVariant);
204 PatKind::Wild
205 }
206
207 hir_def::hir::Pat::Or(ref pats) => PatKind::Or { pats: self.lower_patterns(pats) },
208
209 _ => {
210 self.errors.push(PatternError::Unimplemented);
211 PatKind::Wild
212 }
213 };
214
215 Pat { ty, kind: Box::new(kind) }
216 }
217
218 fn lower_tuple_subpats(
219 &mut self,
220 pats: &[PatId],
221 expected_len: usize,
222 ellipsis: Option<u32>,
223 ) -> Vec<FieldPat<'db>> {
224 if pats.len() > expected_len {
225 self.errors.push(PatternError::ExtraFields);
226 return Vec::new();
227 }
228
229 pats.iter()
230 .enumerate_and_adjust(expected_len, ellipsis.map(|it| it as usize))
231 .map(|(i, &subpattern)| FieldPat {
232 field: LocalFieldId::from_raw((i as u32).into()),
233 pattern: self.lower_pattern(subpattern),
234 })
235 .collect()
236 }
237
238 fn lower_patterns(&mut self, pats: &[PatId]) -> Vec<Pat<'db>> {
239 pats.iter().map(|&p| self.lower_pattern(p)).collect()
240 }
241
242 fn lower_opt_pattern(&mut self, pat: Option<PatId>) -> Option<Pat<'db>> {
243 pat.map(|p| self.lower_pattern(p))
244 }
245
246 fn lower_variant_or_leaf(
247 &mut self,
248 pat: PatId,
249 ty: Ty<'db>,
250 subpatterns: Vec<FieldPat<'db>>,
251 ) -> PatKind<'db> {
252 match self.infer.variant_resolution_for_pat(pat) {
253 Some(variant_id) => {
254 if let VariantId::EnumVariantId(enum_variant) = variant_id {
255 let substs = match ty.kind() {
256 TyKind::Adt(_, substs) => substs,
257 kind => {
258 always!(
259 matches!(kind, TyKind::FnDef(..) | TyKind::Error(_)),
260 "inappropriate type for def: {:?}",
261 ty
262 );
263 self.errors.push(PatternError::UnexpectedType);
264 return PatKind::Wild;
265 }
266 };
267 PatKind::Variant { substs, enum_variant, subpatterns }
268 } else {
269 PatKind::Leaf { subpatterns }
270 }
271 }
272 None => {
273 self.errors.push(PatternError::UnresolvedVariant);
274 PatKind::Wild
275 }
276 }
277 }
278
279 fn lower_path(&mut self, pat: PatId, _path: &Path) -> Pat<'db> {
280 let ty = self.infer.pat_ty(pat);
281
282 let pat_from_kind = |kind| Pat { ty, kind: Box::new(kind) };
283
284 match self.infer.variant_resolution_for_pat(pat) {
285 Some(_) => pat_from_kind(self.lower_variant_or_leaf(pat, ty, Vec::new())),
286 None => {
287 self.errors.push(PatternError::UnresolvedVariant);
288 pat_from_kind(PatKind::Wild)
289 }
290 }
291 }
292
293 fn lower_lit(&mut self, expr: hir_def::hir::ExprId) -> PatKind<'db> {
294 use hir_def::hir::{Expr, Literal::Bool};
295
296 match self.body[expr] {
297 Expr::Literal(Bool(value)) => PatKind::LiteralBool { value },
298 _ => {
299 self.errors.push(PatternError::Unimplemented);
300 PatKind::Wild
301 }
302 }
303 }
304}
305
306impl<'db> HirDisplay<'db> for Pat<'db> {
307 fn hir_fmt(&self, f: &mut HirFormatter<'_, 'db>) -> Result<(), HirDisplayError> {
308 match &*self.kind {
309 PatKind::Wild => write!(f, "_"),
310 PatKind::Never => write!(f, "!"),
311 PatKind::Binding { name, subpattern } => {
312 write!(f, "{}", name.display(f.db, f.edition()))?;
313 if let Some(subpattern) = subpattern {
314 write!(f, " @ ")?;
315 subpattern.hir_fmt(f)?;
316 }
317 Ok(())
318 }
319 PatKind::Variant { subpatterns, .. } | PatKind::Leaf { subpatterns } => {
320 let variant = match *self.kind {
321 PatKind::Variant { enum_variant, .. } => Some(VariantId::from(enum_variant)),
322 _ => self.ty.as_adt().and_then(|(adt, _)| match adt {
323 AdtId::StructId(s) => Some(s.into()),
324 AdtId::UnionId(u) => Some(u.into()),
325 AdtId::EnumId(_) => None,
326 }),
327 };
328
329 if let Some(variant) = variant {
330 match variant {
331 VariantId::EnumVariantId(v) => {
332 let loc = v.lookup(f.db);
333 write!(f, "{}", loc.name.display(f.db, f.edition()))?;
334 }
335 VariantId::StructId(s) => write!(
336 f,
337 "{}",
338 StructSignature::of(f.db, s).name.display(f.db, f.edition())
339 )?,
340 VariantId::UnionId(u) => write!(
341 f,
342 "{}",
343 UnionSignature::of(f.db, u).name.display(f.db, f.edition())
344 )?,
345 };
346
347 let variant_data = variant.fields(f.db);
348 if variant_data.shape == FieldsShape::Record {
349 write!(f, " {{ ")?;
350
351 let mut printed = 0;
352 let subpats = subpatterns
353 .iter()
354 .filter(|p| !matches!(*p.pattern.kind, PatKind::Wild))
355 .map(|p| {
356 printed += 1;
357 WriteWith::new(|f| {
358 write!(
359 f,
360 "{}: ",
361 variant_data.fields()[p.field]
362 .name
363 .display(f.db, f.edition())
364 )?;
365 p.pattern.hir_fmt(f)
366 })
367 });
368 f.write_joined(subpats, ", ")?;
369
370 if printed < variant_data.fields().len() {
371 write!(f, "{}..", if printed > 0 { ", " } else { "" })?;
372 }
373
374 return write!(f, " }}");
375 }
376 }
377
378 let num_fields =
379 variant.map_or(subpatterns.len(), |v| v.fields(f.db).fields().len());
380 if num_fields != 0 || variant.is_none() {
381 write!(f, "(")?;
382 let subpats = (0..num_fields).map(|i| {
383 WriteWith::new(move |f| {
384 let fid = LocalFieldId::from_raw((i as u32).into());
385 if let Some(p) = subpatterns.get(i)
386 && p.field == fid
387 {
388 return p.pattern.hir_fmt(f);
389 }
390 if let Some(p) = subpatterns.iter().find(|p| p.field == fid) {
391 p.pattern.hir_fmt(f)
392 } else {
393 write!(f, "_")
394 }
395 })
396 });
397 f.write_joined(subpats, ", ")?;
398 if let (TyKind::Tuple(..), 1) = (self.ty.kind(), num_fields) {
399 write!(f, ",")?;
400 }
401 write!(f, ")")?;
402 }
403
404 Ok(())
405 }
406 PatKind::Deref { subpattern } => {
407 match self.ty.kind() {
408 TyKind::Ref(.., mutbl) => {
409 write!(f, "&{}", if mutbl == Mutability::Mut { "mut " } else { "" })?
410 }
411 _ => never!("{:?} is a bad Deref pattern type", self.ty),
412 }
413 subpattern.hir_fmt(f)
414 }
415 PatKind::LiteralBool { value } => write!(f, "{value}"),
416 PatKind::Or { pats } => f.write_joined(pats.iter(), " | "),
417 }
418 }
419}
420
421struct WriteWith<'db, F>(F, PhantomCovariantLifetime<'db>)
422where
423 F: Fn(&mut HirFormatter<'_, 'db>) -> Result<(), HirDisplayError>;
424
425impl<'db, F> WriteWith<'db, F>
426where
427 F: Fn(&mut HirFormatter<'_, 'db>) -> Result<(), HirDisplayError>,
428{
429 fn new(f: F) -> Self {
430 Self(f, PhantomCovariantLifetime::new())
431 }
432}
433
434impl<'db, F> HirDisplay<'db> for WriteWith<'db, F>
435where
436 F: Fn(&mut HirFormatter<'_, 'db>) -> Result<(), HirDisplayError>,
437{
438 fn hir_fmt(&self, f: &mut HirFormatter<'_, 'db>) -> Result<(), HirDisplayError> {
439 (self.0)(f)
440 }
441}