plum
git clone https://git.pyrossh.dev/plum
A statically typed, imperative programming language inspired by rust, python
d1a4183
— Peter John
2026-07-19T21:40:36+05:30
feat(plum-checker): class/method/self/match type checking
- plum-checker/src/lib.rs +237 -48
- plum-checker/tests/checker_tests.rs +99 -0
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
|
|
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
|
-
//
|
|
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
|
|
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
|
-
|
|
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
|
-
|
|
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_ =>
|
|
367
|
+
ast::Expr::Self_ => lookup(env, "self"),
|
|
241
|
-
ast::Expr::TypeName(n) =>
|
|
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(<, &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, <).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(<, &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
|
-
|
|
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
|
-
|
|
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
|
+
}
|