plum

#treesitter#compiler#wasm

git clone https://git.pyrossh.dev/plum

A statically typed, imperative programming language inspired by rust, python


d1a4183Peter John 2026-07-19T21:40:36+05:30
feat(plum-checker): class/method/self/match type checking
plum-checker/src/lib.rs CHANGED
@@ -1,5 +1,6 @@
1
1
  pub mod types;
2
2
 
3
+ use std::collections::BTreeMap;
3
4
  use types::{PlumType, TypeEnv, TypeScheme, CheckError, CheckResult};
4
5
  use plum_core::ast;
5
6
 
@@ -33,17 +34,63 @@ pub fn unify(t1: &PlumType, t2: &PlumType) -> Result<(), String> {
33
34
  }
34
35
  }
35
36
 
36
- fn lookup(env: &TypeEnv, name: &str) -> Result<PlumType, String> {
37
+ pub fn lookup(env: &TypeEnv, name: &str) -> Result<PlumType, String> {
37
38
  env.get(name)
38
39
  .map(|s| *s.body.clone())
39
40
  .ok_or_else(|| format!("undefined name '{}'", name))
40
41
  }
41
42
 
43
+ /// Field names and types for every `type ClassName = ...` declaration in the source.
44
+ pub type ClassEnv = BTreeMap<String, Vec<(String, PlumType)>>;
45
+ /// `(receiver type, method name) -> TFun` for every `name<Receiver>(...)` method.
46
+ pub type MethodEnv = BTreeMap<(String, String), PlumType>;
47
+ /// Enum variant name -> owning enum name, e.g. `"True" -> "Bool"`.
48
+ pub type EnumVariants = BTreeMap<String, String>;
49
+
50
+ /// Shared, read-only lookup tables built once from the whole source, threaded through
51
+ /// every check/infer call alongside the (mutable, scope-local) `TypeEnv`.
52
+ pub struct CheckCtx<'a> {
53
+ pub classes: &'a ClassEnv,
54
+ pub methods: &'a MethodEnv,
55
+ pub enum_variants: &'a EnumVariants,
56
+ }
57
+
58
+ /// Builds the global lookup tables (function/const signatures, class fields, method
59
+ /// signatures, enum variants) from a whole source. Shared by `check_source` and by
60
+ /// `plum-wasm-codegen`, which needs the same tables to resolve `self`, field access,
61
+ /// and method dispatch during code generation.
42
- pub fn check_source(source: &ast::Source) -> CheckResult<()> {
62
+ pub fn build_global_tables(source: &ast::Source) -> (TypeEnv, ClassEnv, MethodEnv, EnumVariants) {
43
- let mut errors: Vec<CheckError> = Vec::new();
44
63
  let mut global_env: TypeEnv = TypeEnv::new();
64
+ let mut classes: ClassEnv = BTreeMap::new();
65
+ let mut methods: MethodEnv = BTreeMap::new();
66
+ let mut enum_variants: EnumVariants = BTreeMap::new();
67
+ // `Bool`'s variants are built in (see `infer_expr`'s TypeName handling) rather
68
+ // than requiring every source file to redeclare `enum Bool = | True | False`.
69
+ enum_variants.insert("True".to_string(), "Bool".to_string());
70
+ enum_variants.insert("False".to_string(), "Bool".to_string());
71
+
72
+ // First pass: register class fields and enum variants so later passes can
73
+ // resolve `self.field`, `ClassName(...)`, and bare enum-tag patterns.
74
+ for item in &source.items {
75
+ match item {
76
+ ast::Item::Class(c) => {
77
+ let fields = c.fields.iter()
78
+ .map(|f| (f.name.clone(), plum_type_from_ast(&f.ty)))
79
+ .collect();
80
+ classes.insert(c.name.clone(), fields);
81
+ }
82
+ ast::Item::Enum(e) => {
83
+ for v in &e.variants {
84
+ enum_variants.insert(v.name.clone(), e.name.clone());
85
+ }
86
+ }
87
+ _ => {}
88
+ }
89
+ }
45
90
 
46
- // First pass: register all top-level function signatures and consts
91
+ // Second pass: register top-level function/method signatures and consts.
92
+ // Methods (`name<Receiver>(...)`) live in `methods`, keyed by receiver type,
93
+ // so a bare call can't accidentally resolve to some other type's method.
47
94
  for item in &source.items {
48
95
  match item {
49
96
  ast::Item::Fn(f) => {
@@ -56,8 +103,12 @@ pub fn check_source(source: &ast::Source) -> CheckResult<()> {
56
103
  let ret = f.returns.as_ref()
57
104
  .map(|r| plum_type_from_ast(&ast::Type { name: r.name.clone(), generics: vec![] }))
58
105
  .unwrap_or(PlumType::TUnit);
59
- let scheme = TypeScheme::mono(PlumType::TFun(param_types, Box::new(ret)));
106
+ let fn_ty = PlumType::TFun(param_types, Box::new(ret));
107
+ if let Some(recv) = &f.type_param {
108
+ methods.insert((recv.clone(), f.name.clone()), fn_ty);
109
+ } else {
60
- global_env.insert(f.name.clone(), scheme);
110
+ global_env.insert(f.name.clone(), TypeScheme::mono(fn_ty));
111
+ }
61
112
  }
62
113
  ast::Item::Const(c) => {
63
114
  global_env.insert(c.name.clone(), TypeScheme::mono(PlumType::TVar("_".to_string())));
@@ -66,10 +117,17 @@ pub fn check_source(source: &ast::Source) -> CheckResult<()> {
66
117
  }
67
118
  }
68
119
 
69
- // Second pass: check each function body
120
+ (global_env, classes, methods, enum_variants)
121
+ }
122
+
123
+ pub fn check_source(source: &ast::Source) -> CheckResult<()> {
124
+ let mut errors: Vec<CheckError> = Vec::new();
125
+ let (global_env, classes, methods, enum_variants) = build_global_tables(source);
126
+ let ctx = CheckCtx { classes: &classes, methods: &methods, enum_variants: &enum_variants };
127
+
70
128
  for item in &source.items {
71
129
  if let ast::Item::Fn(f) = item {
72
- let mut local_errors = check_fn(f, &global_env);
130
+ let mut local_errors = check_fn(f, &global_env, &ctx);
73
131
  errors.append(&mut local_errors);
74
132
  }
75
133
  }
@@ -77,10 +135,15 @@ pub fn check_source(source: &ast::Source) -> CheckResult<()> {
77
135
  if errors.is_empty() { Ok(()) } else { Err(errors) }
78
136
  }
79
137
 
80
- fn check_fn(f: &ast::Fn, global_env: &TypeEnv) -> Vec<CheckError> {
138
+ fn check_fn(f: &ast::Fn, global_env: &TypeEnv, ctx: &CheckCtx) -> Vec<CheckError> {
81
139
  let mut errors = Vec::new();
82
140
  let mut env = global_env.clone();
83
141
 
142
+ // Methods (`name<Receiver>(...)`) get an implicit `self: Receiver` binding.
143
+ if let Some(recv) = &f.type_param {
144
+ env.insert("self".to_string(), TypeScheme::mono(PlumType::TNamed(recv.clone())));
145
+ }
146
+
84
147
  // Add params to env
85
148
  for p in &f.params {
86
149
  let ty = match &p.ty {
@@ -99,7 +162,7 @@ fn check_fn(f: &ast::Fn, global_env: &TypeEnv) -> Vec<CheckError> {
99
162
 
100
163
  match &f.body {
101
164
  ast::FnBody::Expr(e) => {
102
- match infer_expr(e, &env) {
165
+ match infer_expr(e, &env, ctx) {
103
166
  Ok(t) => {
104
167
  if let Err(msg) = unify(&declared_ret, &t) {
105
168
  errors.push(CheckError { message: format!("fn '{}': return type mismatch: {}", f.name, msg) });
@@ -109,11 +172,11 @@ fn check_fn(f: &ast::Fn, global_env: &TypeEnv) -> Vec<CheckError> {
109
172
  }
110
173
  }
111
174
  ast::FnBody::Block(block) => {
112
- let mut block_errors = check_block(block, &mut env, &declared_ret, &f.name);
175
+ let mut block_errors = check_block(block, &mut env, &declared_ret, &f.name, ctx);
113
176
  errors.append(&mut block_errors);
114
177
  // Check the type of the last expression statement against the declared return type
115
178
  if let Some(ast::Stmt::Expr(last_expr)) = block.stmts.last() {
116
- match infer_expr(last_expr, &env) {
179
+ match infer_expr(last_expr, &env, ctx) {
117
180
  Ok(t) => {
118
181
  if let Err(msg) = unify(&declared_ret, &t) {
119
182
  errors.push(CheckError { message: format!("fn '{}': return type mismatch: {}", f.name, msg) });
@@ -127,28 +190,28 @@ fn check_fn(f: &ast::Fn, global_env: &TypeEnv) -> Vec<CheckError> {
127
190
  errors
128
191
  }
129
192
 
130
- fn check_block(block: &ast::Block, env: &mut TypeEnv, declared_ret: &PlumType, fn_name: &str) -> Vec<CheckError> {
193
+ fn check_block(block: &ast::Block, env: &mut TypeEnv, declared_ret: &PlumType, fn_name: &str, ctx: &CheckCtx) -> Vec<CheckError> {
131
194
  let mut errors = Vec::new();
132
195
  for stmt in &block.stmts {
133
- let mut stmt_errors = check_stmt(stmt, env, declared_ret, fn_name);
196
+ let mut stmt_errors = check_stmt(stmt, env, declared_ret, fn_name, ctx);
134
197
  errors.append(&mut stmt_errors);
135
198
  }
136
199
  errors
137
200
  }
138
201
 
139
- fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_name: &str) -> Vec<CheckError> {
202
+ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_name: &str, ctx: &CheckCtx) -> Vec<CheckError> {
140
203
  let mut errors = Vec::new();
141
204
  match stmt {
142
205
  ast::Stmt::Assign(a) => {
143
206
  for (target, value) in a.targets.iter().zip(a.values.iter()) {
144
- match infer_expr(value, env) {
207
+ match infer_expr(value, env, ctx) {
145
208
  Ok(t) => { env.insert(target.clone(), TypeScheme::mono(t)); }
146
209
  Err(msg) => errors.push(CheckError { message: format!("fn '{}': assign '{}': {}", fn_name, target, msg) }),
147
210
  }
148
211
  }
149
212
  }
150
213
  ast::Stmt::Return(Some(e)) => {
151
- match infer_expr(e, env) {
214
+ match infer_expr(e, env, ctx) {
152
215
  Ok(t) => {
153
216
  if let Err(msg) = unify(declared_ret, &t) {
154
217
  errors.push(CheckError { message: format!("fn '{}': return type mismatch: {}", fn_name, msg) });
@@ -163,7 +226,7 @@ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_n
163
226
  }
164
227
  }
165
228
  ast::Stmt::If(if_) => {
166
- match infer_expr(&if_.condition, env) {
229
+ match infer_expr(&if_.condition, env, ctx) {
167
230
  Ok(t) => {
168
231
  if let Err(msg) = unify(&PlumType::TBool, &t) {
169
232
  errors.push(CheckError { message: format!("fn '{}': if condition must be Bool: {}", fn_name, msg) });
@@ -171,9 +234,9 @@ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_n
171
234
  }
172
235
  Err(msg) => errors.push(CheckError { message: format!("fn '{}': if condition: {}", fn_name, msg) }),
173
236
  }
174
- errors.append(&mut check_block(&if_.body, env, declared_ret, fn_name));
237
+ errors.append(&mut check_block(&if_.body, env, declared_ret, fn_name, ctx));
175
238
  for ei in &if_.else_ifs {
176
- match infer_expr(&ei.condition, env) {
239
+ match infer_expr(&ei.condition, env, ctx) {
177
240
  Ok(t) => {
178
241
  if let Err(msg) = unify(&PlumType::TBool, &t) {
179
242
  errors.push(CheckError { message: format!("fn '{}': else if condition must be Bool: {}", fn_name, msg) });
@@ -181,14 +244,14 @@ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_n
181
244
  }
182
245
  Err(msg) => errors.push(CheckError { message: format!("fn '{}': else if condition: {}", fn_name, msg) }),
183
246
  }
184
- errors.append(&mut check_block(&ei.body, env, declared_ret, fn_name));
247
+ errors.append(&mut check_block(&ei.body, env, declared_ret, fn_name, ctx));
185
248
  }
186
249
  if let Some(else_block) = &if_.else_ {
187
- errors.append(&mut check_block(else_block, env, declared_ret, fn_name));
250
+ errors.append(&mut check_block(else_block, env, declared_ret, fn_name, ctx));
188
251
  }
189
252
  }
190
253
  ast::Stmt::While(w) => {
191
- match infer_expr(&w.condition, env) {
254
+ match infer_expr(&w.condition, env, ctx) {
192
255
  Ok(t) => {
193
256
  if let Err(msg) = unify(&PlumType::TBool, &t) {
194
257
  errors.push(CheckError { message: format!("fn '{}': while condition must be Bool: {}", fn_name, msg) });
@@ -196,10 +259,10 @@ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_n
196
259
  }
197
260
  Err(msg) => errors.push(CheckError { message: format!("fn '{}': while condition: {}", fn_name, msg) }),
198
261
  }
199
- errors.append(&mut check_block(&w.body, env, declared_ret, fn_name));
262
+ errors.append(&mut check_block(&w.body, env, declared_ret, fn_name, ctx));
200
263
  }
201
264
  ast::Stmt::For(f_stmt) => {
202
- match infer_expr(&f_stmt.iter, env) {
265
+ match infer_expr(&f_stmt.iter, env, ctx) {
203
266
  Ok(_) => {}
204
267
  Err(msg) => errors.push(CheckError { message: format!("fn '{}': for iter: {}", fn_name, msg) }),
205
268
  }
@@ -207,15 +270,15 @@ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_n
207
270
  for var in &f_stmt.vars {
208
271
  inner_env.insert(var.clone(), TypeScheme::mono(PlumType::TInt));
209
272
  }
210
- errors.append(&mut check_block(&f_stmt.body, &mut inner_env, declared_ret, fn_name));
273
+ errors.append(&mut check_block(&f_stmt.body, &mut inner_env, declared_ret, fn_name, ctx));
211
274
  }
212
275
  ast::Stmt::Expr(e) => {
213
- if let Err(msg) = infer_expr(e, env) {
276
+ if let Err(msg) = infer_expr(e, env, ctx) {
214
277
  errors.push(CheckError { message: format!("fn '{}': {}", fn_name, msg) });
215
278
  }
216
279
  }
217
280
  ast::Stmt::Assert(e) => {
218
- match infer_expr(e, env) {
281
+ match infer_expr(e, env, ctx) {
219
282
  Ok(t) => {
220
283
  if let Err(msg) = unify(&PlumType::TBool, &t) {
221
284
  errors.push(CheckError { message: format!("fn '{}': assert must be Bool: {}", fn_name, msg) });
@@ -224,31 +287,98 @@ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_n
224
287
  Err(msg) => errors.push(CheckError { message: format!("fn '{}': assert: {}", fn_name, msg) }),
225
288
  }
226
289
  }
290
+ ast::Stmt::Match(m) => {
291
+ errors.append(&mut check_match(m, env, declared_ret, fn_name, ctx));
292
+ }
227
293
  ast::Stmt::Break | ast::Stmt::Continue | ast::Stmt::Todo => {}
228
- ast::Stmt::Match(_) => {}
229
294
  }
230
295
  errors
231
296
  }
232
297
 
298
+ fn check_match(m: &ast::Match, env: &TypeEnv, declared_ret: &PlumType, fn_name: &str, ctx: &CheckCtx) -> Vec<CheckError> {
299
+ let mut errors = Vec::new();
300
+
301
+ let mut subject_types: Vec<PlumType> = Vec::new();
302
+ for s in &m.subjects {
303
+ match infer_expr(s, env, ctx) {
304
+ Ok(t) => subject_types.push(t),
305
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': match subject: {}", fn_name, msg) }),
306
+ }
307
+ }
308
+ if subject_types.len() != m.subjects.len() {
309
+ return errors; // a subject failed to type — cases can't be checked meaningfully
310
+ }
311
+
312
+ for case in &m.cases {
313
+ let mut case_env = env.clone();
314
+ if case.patterns.len() == subject_types.len() {
315
+ for (pat, sty) in case.patterns.iter().zip(subject_types.iter()) {
316
+ if let Err(msg) = check_pattern(pat, sty, &mut case_env, ctx) {
317
+ errors.push(CheckError { message: format!("fn '{}': match case: {}", fn_name, msg) });
318
+ }
319
+ }
320
+ } else {
321
+ errors.push(CheckError {
322
+ message: format!(
323
+ "fn '{}': match case has {} pattern(s), expected {}",
324
+ fn_name, case.patterns.len(), subject_types.len()
325
+ ),
326
+ });
327
+ }
328
+ errors.append(&mut check_block(&case.body, &mut case_env, declared_ret, fn_name, ctx));
329
+ }
330
+ errors
331
+ }
332
+
333
+ /// Checks a single case pattern against the type of the subject it matches, binding any
334
+ /// new names it introduces into `env`. Constructor-payload sub-patterns (`Some(x)`) bind
335
+ /// against an unconstrained type since enum variants don't carry per-field type info (v1.5).
336
+ fn check_pattern(pat: &ast::CasePattern, subject_ty: &PlumType, env: &mut TypeEnv, ctx: &CheckCtx) -> Result<(), String> {
337
+ match pat {
338
+ ast::CasePattern::Wildcard => Ok(()),
339
+ ast::CasePattern::Int(_) => unify(subject_ty, &PlumType::TInt),
340
+ ast::CasePattern::Float(_) => unify(subject_ty, &PlumType::TFloat),
341
+ ast::CasePattern::String(_) => unify(subject_ty, &PlumType::TStr),
342
+ ast::CasePattern::Name(n) => {
343
+ let is_known_variant = n.chars().next().map(|c| c.is_uppercase()).unwrap_or(false)
344
+ && ctx.enum_variants.contains_key(n);
345
+ if is_known_variant {
346
+ Ok(()) // equality check against a known enum tag, e.g. `True`
347
+ } else {
348
+ env.insert(n.clone(), TypeScheme::mono(subject_ty.clone()));
349
+ Ok(())
350
+ }
351
+ }
352
+ ast::CasePattern::Class { name: _, fields } => {
353
+ for f in fields {
354
+ check_pattern(f, &PlumType::TVar("_".to_string()), env, ctx)?;
355
+ }
356
+ Ok(())
357
+ }
358
+ }
359
+ }
360
+
233
- fn infer_expr(expr: &ast::Expr, env: &TypeEnv) -> Result<PlumType, String> {
361
+ pub fn infer_expr(expr: &ast::Expr, env: &TypeEnv, ctx: &CheckCtx) -> Result<PlumType, String> {
234
362
  match expr {
235
363
  ast::Expr::Int(_) => Ok(PlumType::TInt),
236
364
  ast::Expr::Float(_) => Ok(PlumType::TFloat),
237
365
  ast::Expr::String(_) => Ok(PlumType::TStr),
238
- ast::Expr::Var(name) if name == "true" || name == "false" => Ok(PlumType::TBool),
239
366
  ast::Expr::Var(name) => lookup(env, name),
240
- ast::Expr::Self_ => Ok(PlumType::TVar("Self".to_string())),
367
+ ast::Expr::Self_ => lookup(env, "self"),
241
- ast::Expr::TypeName(n) => Ok(PlumType::TNamed(n.clone())),
368
+ ast::Expr::TypeName(n) => match n.as_str() {
369
+ "True" | "False" => Ok(PlumType::TBool),
370
+ other => Ok(PlumType::TNamed(other.to_string())),
371
+ },
242
- ast::Expr::Paren(inner) => infer_expr(inner, env),
372
+ ast::Expr::Paren(inner) => infer_expr(inner, env, ctx),
243
373
  ast::Expr::Not(inner) => {
244
- let t = infer_expr(inner, env)?;
374
+ let t = infer_expr(inner, env, ctx)?;
245
375
  unify(&PlumType::TBool, &t)?;
246
376
  Ok(PlumType::TBool)
247
377
  }
248
- ast::Expr::Unary(u) => infer_expr(&u.operand, env),
378
+ ast::Expr::Unary(u) => infer_expr(&u.operand, env, ctx),
249
379
  ast::Expr::Binary(b) => {
250
- let lt = infer_expr(&b.left, env)?;
380
+ let lt = infer_expr(&b.left, env, ctx)?;
251
- let rt = infer_expr(&b.right, env)?;
381
+ let rt = infer_expr(&b.right, env, ctx)?;
252
382
  unify(&lt, &rt).map_err(|e| format!("binary op: {}", e))?;
253
383
  match b.op {
254
384
  ast::BinOp::Range => Ok(PlumType::TNamed("Range".to_string())),
@@ -256,23 +386,23 @@ fn infer_expr(expr: &ast::Expr, env: &TypeEnv) -> Result<PlumType, String> {
256
386
  }
257
387
  }
258
388
  ast::Expr::Bool(b) => {
259
- let lt = infer_expr(&b.left, env)?;
389
+ let lt = infer_expr(&b.left, env, ctx)?;
260
- let rt = infer_expr(&b.right, env)?;
390
+ let rt = infer_expr(&b.right, env, ctx)?;
261
391
  unify(&PlumType::TBool, &lt).map_err(|e| format!("bool op left: {}", e))?;
262
392
  unify(&PlumType::TBool, &rt).map_err(|e| format!("bool op right: {}", e))?;
263
393
  Ok(PlumType::TBool)
264
394
  }
265
395
  ast::Expr::Compare(c) => {
266
- let lt = infer_expr(&c.left, env)?;
396
+ let lt = infer_expr(&c.left, env, ctx)?;
267
- let rt = infer_expr(&c.right, env)?;
397
+ let rt = infer_expr(&c.right, env, ctx)?;
268
398
  unify(&lt, &rt).map_err(|e| format!("compare op: {}", e))?;
269
399
  Ok(PlumType::TBool)
270
400
  }
271
401
  ast::Expr::Ternary(t) => {
272
- let ct = infer_expr(&t.condition, env)?;
402
+ let ct = infer_expr(&t.condition, env, ctx)?;
273
403
  unify(&PlumType::TBool, &ct).map_err(|e| format!("ternary condition: {}", e))?;
274
- let tt = infer_expr(&t.then, env)?;
404
+ let tt = infer_expr(&t.then, env, ctx)?;
275
- let et = infer_expr(&t.else_, env)?;
405
+ let et = infer_expr(&t.else_, env, ctx)?;
276
406
  unify(&tt, &et).map_err(|e| format!("ternary branches: {}", e))?;
277
407
  Ok(tt)
278
408
  }
@@ -288,7 +418,7 @@ fn infer_expr(expr: &ast::Expr, env: &TypeEnv) -> Result<PlumType, String> {
288
418
  ast::Arg::Keyword { value, .. } => value,
289
419
  ast::Arg::Pair { value, .. } => value,
290
420
  };
291
- let actual = infer_expr(arg_expr, env)?;
421
+ let actual = infer_expr(arg_expr, env, ctx)?;
292
422
  unify(expected, &actual).map_err(|e| format!("call '{}' arg {}: {}", call.name, i, e))?;
293
423
  }
294
424
  Ok(*ret)
@@ -297,7 +427,66 @@ fn infer_expr(expr: &ast::Expr, env: &TypeEnv) -> Result<PlumType, String> {
297
427
  Err(_) => Ok(PlumType::TVar("_".to_string())), // unknown fn: allow, codegen will catch
298
428
  }
299
429
  }
430
+ ast::Expr::ClassCall(call) => {
431
+ match ctx.classes.get(&call.type_name) {
432
+ Some(fields) => {
433
+ for fa in &call.fields {
434
+ match fields.iter().find(|(n, _)| n == &fa.name) {
435
+ Some((_, expected)) => {
436
+ let actual = infer_expr(&fa.value, env, ctx)?;
437
+ unify(expected, &actual)
438
+ .map_err(|e| format!("class '{}' field '{}': {}", call.type_name, fa.name, e))?;
439
+ }
440
+ None => return Err(format!("unknown field '{}' on class '{}'", fa.name, call.type_name)),
441
+ }
442
+ }
443
+ Ok(PlumType::TNamed(call.type_name.clone()))
444
+ }
445
+ // Unmodeled (e.g. builtin/std) type: allow, codegen will catch.
300
- ast::Expr::ClassCall(_) => Ok(PlumType::TVar("_".to_string())),
446
+ None => Ok(PlumType::TNamed(call.type_name.clone())),
447
+ }
448
+ }
449
+ ast::Expr::Attribute(attr) => {
450
+ let obj_ty = infer_expr(&attr.object, env, ctx)?;
451
+ match &attr.attr {
452
+ ast::AttrKind::Field(field_name) => match &obj_ty {
453
+ PlumType::TNamed(class_name) => match ctx.classes.get(class_name) {
454
+ Some(fields) => fields.iter()
455
+ .find(|(n, _)| n == field_name)
456
+ .map(|(_, t)| t.clone())
457
+ .ok_or_else(|| format!("no field '{}' on type '{}'", field_name, class_name)),
458
+ // Unmodeled type: allow, codegen will catch.
301
- ast::Expr::Attribute(_) => Ok(PlumType::TVar("_".to_string())),
459
+ None => Ok(PlumType::TVar("_".to_string())),
460
+ },
461
+ _ => Err(format!("cannot access field '{}' on non-class type {}", field_name, obj_ty)),
462
+ },
463
+ ast::AttrKind::Method(call) => match &obj_ty {
464
+ PlumType::TNamed(class_name) => match ctx.methods.get(&(class_name.clone(), call.name.clone())) {
465
+ Some(PlumType::TFun(param_types, ret)) => {
466
+ if call.args.len() != param_types.len() {
467
+ return Err(format!(
468
+ "method '{}.{}': expected {} args, got {}",
469
+ class_name, call.name, param_types.len(), call.args.len()
470
+ ));
471
+ }
472
+ for (i, (arg, expected)) in call.args.iter().zip(param_types.iter()).enumerate() {
473
+ let arg_expr = match arg {
474
+ ast::Arg::Positional(e) => e,
475
+ ast::Arg::Keyword { value, .. } => value,
476
+ ast::Arg::Pair { value, .. } => value,
477
+ };
478
+ let actual = infer_expr(arg_expr, env, ctx)?;
479
+ unify(expected, &actual)
480
+ .map_err(|e| format!("method '{}.{}' arg {}: {}", class_name, call.name, i, e))?;
481
+ }
482
+ Ok(*ret.clone())
483
+ }
484
+ // Unmodeled method (e.g. builtin/std): allow, codegen will catch.
485
+ _ => Ok(PlumType::TVar("_".to_string())),
486
+ },
487
+ _ => Ok(PlumType::TVar("_".to_string())),
488
+ },
489
+ }
490
+ }
302
491
  }
303
492
  }
plum-checker/tests/checker_tests.rs CHANGED
@@ -83,3 +83,102 @@ fn type_mismatch_in_binary_op_is_error() {
83
83
  let result = check_source(&source);
84
84
  assert!(result.is_err());
85
85
  }
86
+
87
+ #[test]
88
+ fn bool_literal_true_false_are_bool() {
89
+ let src = "isTrue() -> Bool =\n True\n";
90
+ let source = parse(src);
91
+ assert!(check_source(&source).is_ok(), "expected Ok, got {:?}", check_source(&source).err());
92
+ }
93
+
94
+ #[test]
95
+ fn bool_literal_wrong_return_type_is_error() {
96
+ let src = "bad() -> Int =\n False\n";
97
+ let source = parse(src);
98
+ let result = check_source(&source);
99
+ assert!(result.is_err());
100
+ }
101
+
102
+ #[test]
103
+ fn method_self_field_access_passes() {
104
+ let src = "type Cat =\n name: Str\n age: Int\n\ngetName<Cat>() -> Str =\n self.name\n";
105
+ let source = parse(src);
106
+ let result = check_source(&source);
107
+ assert!(result.is_ok(), "expected Ok, got {:?}", result.err());
108
+ }
109
+
110
+ #[test]
111
+ fn method_self_unknown_field_is_error() {
112
+ let src = "type Cat =\n name: Str\n\ngetAge<Cat>() -> Int =\n self.age\n";
113
+ let source = parse(src);
114
+ let result = check_source(&source);
115
+ assert!(result.is_err());
116
+ }
117
+
118
+ #[test]
119
+ fn self_outside_method_is_error() {
120
+ let src = "bad() -> Int =\n self\n";
121
+ let source = parse(src);
122
+ let result = check_source(&source);
123
+ assert!(result.is_err());
124
+ }
125
+
126
+ #[test]
127
+ fn class_call_checks_field_types() {
128
+ let src = "type Cat =\n name: Str\n age: Int\n\nmakeCat() -> Cat =\n Cat(name: \"x\", age: 1)\n";
129
+ let source = parse(src);
130
+ let result = check_source(&source);
131
+ assert!(result.is_ok(), "expected Ok, got {:?}", result.err());
132
+ }
133
+
134
+ #[test]
135
+ fn class_call_wrong_field_type_is_error() {
136
+ let src = "type Cat =\n name: Str\n age: Int\n\nmakeCat() -> Cat =\n Cat(name: \"x\", age: \"y\")\n";
137
+ let source = parse(src);
138
+ let result = check_source(&source);
139
+ assert!(result.is_err());
140
+ }
141
+
142
+ #[test]
143
+ fn class_call_unknown_field_is_error() {
144
+ let src = "type Cat =\n name: Str\n\nmakeCat() -> Cat =\n Cat(name: \"x\", age: 1)\n";
145
+ let source = parse(src);
146
+ let result = check_source(&source);
147
+ assert!(result.is_err());
148
+ }
149
+
150
+ #[test]
151
+ fn method_call_via_attribute_type_checks_args() {
152
+ let src = "type Cat =\n name: Str\n\nrename<Cat>(n: Str) -> Str =\n n\n\nuse(c: Cat) -> Str =\n c.rename(\"x\")\n";
153
+ let source = parse(src);
154
+ let result = check_source(&source);
155
+ assert!(result.is_ok(), "expected Ok, got {:?}", result.err());
156
+ }
157
+
158
+ #[test]
159
+ fn match_binds_name_pattern_to_subject_type() {
160
+ let src = "main(a: Int) -> Int =\n match a\n x =>\n x\n _ =>\n 0\n";
161
+ let source = parse(src);
162
+ let result = check_source(&source);
163
+ assert!(result.is_ok(), "expected Ok, got {:?}", result.err());
164
+ }
165
+
166
+ #[test]
167
+ fn match_true_false_are_variant_patterns_not_bindings_without_enum_decl() {
168
+ // True/False are built-in Bool variants — they must be recognized as tag
169
+ // comparisons even when the source doesn't redeclare `enum Bool`, so a
170
+ // later `_` wildcard arm remains reachable (each pattern binds/compares,
171
+ // it doesn't just re-bind the subject under the name "True").
172
+ let src = "pick(a: Bool) -> Int =\n match a\n True =>\n 1\n False =>\n 0\n";
173
+ let source = parse(src);
174
+ let result = check_source(&source);
175
+ assert!(result.is_ok(), "expected Ok, got {:?}", result.err());
176
+ }
177
+
178
+ #[test]
179
+ fn match_int_pattern_against_str_subject_is_error() {
180
+ let src = "main(a: Str) -> Int =\n match a\n 1 =>\n 1\n _ =>\n 0\n";
181
+ let source = parse(src);
182
+ let result = check_source(&source);
183
+ assert!(result.is_err());
184
+ }