Skip to main content

palladium/typeck/
exhaustiveness.rs

1// Pattern exhaustiveness checking for Palladium
2// "Ensuring all possibilities are covered"
3
4use crate::ast::{Pattern, PatternData};
5use crate::errors::{CompileError, Result, Span};
6use std::collections::{HashMap, HashSet};
7
8/// Represents a pattern in exhaustiveness checking
9#[derive(Debug, Clone, PartialEq, Eq, Hash)]
10pub enum PatternKind {
11    /// Wildcard pattern (_) - matches anything
12    Wildcard,
13    /// Variable binding - matches anything and binds it
14    Binding(String),
15    /// Enum constructor pattern
16    Constructor {
17        enum_name: String,
18        variant: String,
19        arity: usize,
20    },
21}
22
23/// Information about enum variants for exhaustiveness checking
24#[derive(Debug, Clone)]
25pub struct EnumInfo {
26    pub name: String,
27    pub variants: Vec<VariantInfo>,
28}
29
30#[derive(Debug, Clone)]
31pub struct VariantInfo {
32    pub name: String,
33    #[allow(dead_code)]
34    pub arity: usize, // Number of fields (0 for unit variants)
35}
36
37/// Pattern exhaustiveness checker
38pub struct ExhaustivenessChecker {
39    /// Information about all enums in the program
40    enums: HashMap<String, EnumInfo>,
41}
42
43impl ExhaustivenessChecker {
44    pub fn new(enums: HashMap<String, EnumInfo>) -> Self {
45        Self { enums }
46    }
47
48    /// Check if a match expression is exhaustive
49    pub fn check_match(&self, matched_type: &str, patterns: &[Pattern], span: Span) -> Result<()> {
50        // If the matched type is an enum, check exhaustiveness
51        if let Some(enum_info) = self.enums.get(matched_type) {
52            self.check_enum_exhaustiveness(enum_info, patterns, span)
53        } else {
54            // For non-enum types, we need at least one wildcard or binding pattern
55            let has_catchall = patterns
56                .iter()
57                .any(|p| matches!(p, Pattern::Wildcard | Pattern::Ident(_)));
58            if !has_catchall {
59                Err(CompileError::NonExhaustiveMatch {
60                    missing_patterns: vec!["_ (wildcard pattern)".to_string()],
61                    span: Some(span),
62                })
63            } else {
64                Ok(())
65            }
66        }
67    }
68
69    /// Check if patterns are exhaustive for an enum
70    fn check_enum_exhaustiveness(
71        &self,
72        enum_info: &EnumInfo,
73        patterns: &[Pattern],
74        span: Span,
75    ) -> Result<()> {
76        // Track which variants are covered
77        let mut covered_variants = HashSet::new();
78        let mut has_wildcard = false;
79        let mut unreachable_patterns = Vec::new();
80
81        for (i, pattern) in patterns.iter().enumerate() {
82            match pattern {
83                Pattern::Wildcard | Pattern::Ident(_) => {
84                    // Wildcard or binding matches all remaining variants
85                    if has_wildcard || covered_variants.len() == enum_info.variants.len() {
86                        unreachable_patterns.push((i, pattern.to_string()));
87                    }
88                    has_wildcard = true;
89                }
90                Pattern::EnumPattern {
91                    enum_name, variant, ..
92                } => {
93                    if enum_name != &enum_info.name {
94                        return Err(CompileError::TypeMismatch {
95                            expected: enum_info.name.clone(),
96                            found: enum_name.clone(),
97                            span: Some(span),
98                        });
99                    }
100
101                    // Check if this variant exists
102                    if !enum_info.variants.iter().any(|v| &v.name == variant) {
103                        return Err(CompileError::Generic(format!(
104                            "Unknown variant '{}::{}' in match pattern",
105                            enum_name, variant
106                        )));
107                    }
108
109                    // Check if already covered by wildcard
110                    if has_wildcard || covered_variants.contains(variant) {
111                        unreachable_patterns.push((i, pattern.to_string()));
112                    } else {
113                        covered_variants.insert(variant.clone());
114                    }
115                }
116            }
117        }
118
119        // Report unreachable patterns
120        if !unreachable_patterns.is_empty() {
121            return Err(CompileError::UnreachablePattern {
122                patterns: unreachable_patterns.into_iter().map(|(_, p)| p).collect(),
123                span: Some(span),
124            });
125        }
126
127        // Check if all variants are covered
128        if !has_wildcard && covered_variants.len() < enum_info.variants.len() {
129            let missing_variants: Vec<String> = enum_info
130                .variants
131                .iter()
132                .filter(|v| !covered_variants.contains(&v.name))
133                .map(|v| format!("{}::{}", enum_info.name, v.name))
134                .collect();
135
136            return Err(CompileError::NonExhaustiveMatch {
137                missing_patterns: missing_variants,
138                span: Some(span),
139            });
140        }
141
142        Ok(())
143    }
144
145    /// Check for redundant patterns (patterns that can never match)
146    #[allow(dead_code)]
147    pub fn check_redundancy(patterns: &[Pattern]) -> Vec<(usize, String)> {
148        let mut redundant = Vec::new();
149        let mut seen_wildcard = false;
150        let mut seen_variants = HashSet::new();
151
152        for (i, pattern) in patterns.iter().enumerate() {
153            match pattern {
154                Pattern::Wildcard | Pattern::Ident(_) => {
155                    if seen_wildcard {
156                        redundant.push((i, "This pattern is unreachable".to_string()));
157                    }
158                    seen_wildcard = true;
159                }
160                Pattern::EnumPattern {
161                    enum_name, variant, ..
162                } => {
163                    let variant_key = format!("{}::{}", enum_name, variant);
164                    if seen_wildcard {
165                        redundant.push((i, "This pattern is unreachable (previous wildcard pattern covers all cases)".to_string()));
166                    } else if seen_variants.contains(&variant_key) {
167                        redundant.push((i, format!("Variant '{}' already covered", variant_key)));
168                    } else {
169                        seen_variants.insert(variant_key);
170                    }
171                }
172            }
173        }
174
175        redundant
176    }
177}
178
179/// Helper to extract pattern information from AST patterns
180impl Pattern {
181    /// Convert AST pattern to exhaustiveness checker pattern kind
182    pub fn to_pattern_kind(&self) -> PatternKind {
183        match self {
184            Pattern::Wildcard => PatternKind::Wildcard,
185            Pattern::Ident(name) => PatternKind::Binding(name.clone()),
186            Pattern::EnumPattern {
187                enum_name,
188                variant,
189                data,
190            } => {
191                let arity = match data {
192                    None => 0,
193                    Some(PatternData::Tuple(patterns)) => patterns.len(),
194                    Some(PatternData::Struct(fields)) => fields.len(),
195                };
196                PatternKind::Constructor {
197                    enum_name: enum_name.clone(),
198                    variant: variant.clone(),
199                    arity,
200                }
201            }
202        }
203    }
204}
205
206#[cfg(test)]
207mod tests {
208    use super::*;
209
210    fn create_option_enum() -> EnumInfo {
211        EnumInfo {
212            name: "Option".to_string(),
213            variants: vec![
214                VariantInfo {
215                    name: "Some".to_string(),
216                    arity: 1,
217                },
218                VariantInfo {
219                    name: "None".to_string(),
220                    arity: 0,
221                },
222            ],
223        }
224    }
225
226    #[allow(dead_code)]
227    fn create_result_enum() -> EnumInfo {
228        EnumInfo {
229            name: "Result".to_string(),
230            variants: vec![
231                VariantInfo {
232                    name: "Ok".to_string(),
233                    arity: 1,
234                },
235                VariantInfo {
236                    name: "Err".to_string(),
237                    arity: 1,
238                },
239            ],
240        }
241    }
242
243    #[test]
244    fn test_exhaustive_enum_match() {
245        let mut enums = HashMap::new();
246        enums.insert("Option".to_string(), create_option_enum());
247
248        let checker = ExhaustivenessChecker::new(enums);
249
250        let patterns = vec![
251            Pattern::EnumPattern {
252                enum_name: "Option".to_string(),
253                variant: "Some".to_string(),
254                data: Some(PatternData::Tuple(vec![Pattern::Ident("x".to_string())])),
255            },
256            Pattern::EnumPattern {
257                enum_name: "Option".to_string(),
258                variant: "None".to_string(),
259                data: None,
260            },
261        ];
262
263        assert!(checker
264            .check_match("Option", &patterns, Span::dummy())
265            .is_ok());
266    }
267
268    #[test]
269    fn test_non_exhaustive_enum_match() {
270        let mut enums = HashMap::new();
271        enums.insert("Option".to_string(), create_option_enum());
272
273        let checker = ExhaustivenessChecker::new(enums);
274
275        let patterns = vec![Pattern::EnumPattern {
276            enum_name: "Option".to_string(),
277            variant: "Some".to_string(),
278            data: Some(PatternData::Tuple(vec![Pattern::Ident("x".to_string())])),
279        }];
280
281        let result = checker.check_match("Option", &patterns, Span::dummy());
282        assert!(result.is_err());
283
284        if let Err(CompileError::NonExhaustiveMatch {
285            missing_patterns, ..
286        }) = result
287        {
288            assert_eq!(missing_patterns, vec!["Option::None"]);
289        } else {
290            panic!("Expected NonExhaustiveMatch error");
291        }
292    }
293
294    #[test]
295    fn test_wildcard_makes_exhaustive() {
296        let mut enums = HashMap::new();
297        enums.insert("Option".to_string(), create_option_enum());
298
299        let checker = ExhaustivenessChecker::new(enums);
300
301        let patterns = vec![
302            Pattern::EnumPattern {
303                enum_name: "Option".to_string(),
304                variant: "Some".to_string(),
305                data: Some(PatternData::Tuple(vec![Pattern::Ident("x".to_string())])),
306            },
307            Pattern::Wildcard,
308        ];
309
310        assert!(checker
311            .check_match("Option", &patterns, Span::dummy())
312            .is_ok());
313    }
314
315    #[test]
316    fn test_unreachable_pattern_after_wildcard() {
317        let mut enums = HashMap::new();
318        enums.insert("Option".to_string(), create_option_enum());
319
320        let checker = ExhaustivenessChecker::new(enums);
321
322        let patterns = vec![
323            Pattern::Wildcard,
324            Pattern::EnumPattern {
325                enum_name: "Option".to_string(),
326                variant: "None".to_string(),
327                data: None,
328            },
329        ];
330
331        let result = checker.check_match("Option", &patterns, Span::dummy());
332        assert!(result.is_err());
333
334        if let Err(CompileError::UnreachablePattern { .. }) = result {
335            // Expected
336        } else {
337            panic!("Expected UnreachablePattern error");
338        }
339    }
340
341    #[test]
342    fn test_duplicate_variant_pattern() {
343        let mut enums = HashMap::new();
344        enums.insert("Option".to_string(), create_option_enum());
345
346        let checker = ExhaustivenessChecker::new(enums);
347
348        let patterns = vec![
349            Pattern::EnumPattern {
350                enum_name: "Option".to_string(),
351                variant: "None".to_string(),
352                data: None,
353            },
354            Pattern::EnumPattern {
355                enum_name: "Option".to_string(),
356                variant: "None".to_string(),
357                data: None,
358            },
359            Pattern::EnumPattern {
360                enum_name: "Option".to_string(),
361                variant: "Some".to_string(),
362                data: Some(PatternData::Tuple(vec![Pattern::Wildcard])),
363            },
364        ];
365
366        let result = checker.check_match("Option", &patterns, Span::dummy());
367        assert!(result.is_err());
368
369        if let Err(CompileError::UnreachablePattern { .. }) = result {
370            // Expected
371        } else {
372            panic!("Expected UnreachablePattern error");
373        }
374    }
375}