Skip to main content

hir_ty/diagnostics/
match_check.rs

1//! Validation of matches.
2//!
3//! This module provides lowering from [hir_def::hir::Pat] to [self::Pat] and match
4//! checking algorithm.
5//!
6//! It is modeled on the rustc module `rustc_mir_build::thir::pattern`.
7
8mod 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/// Close relative to `rustc_mir_build::thir::pattern::PatKind`
56#[derive(Clone, Debug, PartialEq)]
57pub(crate) enum PatKind<'db> {
58    Wild,
59    Never,
60
61    /// `x`, `ref x`, `x @ P`, etc.
62    Binding {
63        name: Name,
64        subpattern: Option<Pat<'db>>,
65    },
66
67    /// `Foo(...)` or `Foo{...}` or `Foo`, where `Foo` is a variant name from an ADT with
68    /// multiple variants.
69    Variant {
70        substs: GenericArgs<'db>,
71        enum_variant: EnumVariantId,
72        subpatterns: Vec<FieldPat<'db>>,
73    },
74
75    /// `(...)`, `Foo(...)`, `Foo{...}`, or `Foo`, where `Foo` is a variant name from an ADT with
76    /// a single variant.
77    Leaf {
78        subpatterns: Vec<FieldPat<'db>>,
79    },
80
81    /// `&P`, `&mut P`, etc.
82    Deref {
83        subpattern: Pat<'db>,
84    },
85
86    // FIXME: for now, only bool literals are implemented
87    LiteralBool {
88        value: bool,
89    },
90
91    /// An or-pattern, e.g. `p | q`.
92    /// Invariant: `pats.len() >= 2`.
93    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        // XXX(iDawer): Collecting pattern adjustments feels imprecise to me.
116        // When lowering of & and box patterns are implemented this should be tested
117        // in a manner of `match_ergonomics_issue_9095` test.
118        // Pattern adjustment is part of RFC 2005-match-ergonomics.
119        // More info https://github.com/rust-lang/rust/issues/42640#issuecomment-313535089
120        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                        // XXX(iDawer): field lookup is inefficient
188                        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}