plum

#treesitter#compiler#wasm

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

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


30f1008Peter John 2026-07-19T19:14:09+05:30
feat(plum-checker): implement type checker with unification
plum-checker/src/lib.rs CHANGED
@@ -1 +1,291 @@
1
1
  pub mod types;
2
+
3
+ use types::{PlumType, TypeEnv, TypeScheme, CheckError, CheckResult};
4
+ use plum_core::ast;
5
+
6
+ pub fn plum_type_from_ast(ty: &ast::Type) -> PlumType {
7
+ match ty.name.as_str() {
8
+ "Int" => PlumType::TInt,
9
+ "Float" => PlumType::TFloat,
10
+ "Bool" => PlumType::TBool,
11
+ "Str" => PlumType::TStr,
12
+ "Unit" => PlumType::TUnit,
13
+ other => PlumType::TNamed(other.to_string()),
14
+ }
15
+ }
16
+
17
+ pub fn unify(t1: &PlumType, t2: &PlumType) -> Result<(), String> {
18
+ match (t1, t2) {
19
+ (PlumType::TVar(_), _) | (_, PlumType::TVar(_)) => Ok(()),
20
+ (PlumType::TInt, PlumType::TInt) => Ok(()),
21
+ (PlumType::TFloat, PlumType::TFloat) => Ok(()),
22
+ (PlumType::TBool, PlumType::TBool) => Ok(()),
23
+ (PlumType::TStr, PlumType::TStr) => Ok(()),
24
+ (PlumType::TUnit, PlumType::TUnit) => Ok(()),
25
+ (PlumType::TNamed(a), PlumType::TNamed(b)) if a == b => Ok(()),
26
+ (PlumType::TFun(ps1, r1), PlumType::TFun(ps2, r2)) if ps1.len() == ps2.len() => {
27
+ for (p1, p2) in ps1.iter().zip(ps2.iter()) {
28
+ unify(p1, p2)?;
29
+ }
30
+ unify(r1, r2)
31
+ }
32
+ _ => Err(format!("type mismatch: expected {}, found {}", t1, t2)),
33
+ }
34
+ }
35
+
36
+ fn lookup(env: &TypeEnv, name: &str) -> Result<PlumType, String> {
37
+ env.get(name)
38
+ .map(|s| *s.body.clone())
39
+ .ok_or_else(|| format!("undefined name '{}'", name))
40
+ }
41
+
42
+ pub fn check_source(source: &ast::Source) -> CheckResult<()> {
43
+ let mut errors: Vec<CheckError> = Vec::new();
44
+ let mut global_env: TypeEnv = TypeEnv::new();
45
+
46
+ // First pass: register all top-level function signatures and consts
47
+ for item in &source.items {
48
+ match item {
49
+ ast::Item::Fn(f) => {
50
+ let param_types: Vec<PlumType> = f.params.iter().map(|p| {
51
+ match &p.ty {
52
+ ast::ParamType::Type(t) => plum_type_from_ast(t),
53
+ ast::ParamType::Variadic(t) => plum_type_from_ast(t),
54
+ }
55
+ }).collect();
56
+ let ret = f.returns.as_ref()
57
+ .map(|r| PlumType::TNamed(r.name.clone()))
58
+ .unwrap_or(PlumType::TUnit);
59
+ let scheme = TypeScheme::mono(PlumType::TFun(param_types, Box::new(ret)));
60
+ global_env.insert(f.name.clone(), scheme);
61
+ }
62
+ ast::Item::Const(c) => {
63
+ global_env.insert(c.name.clone(), TypeScheme::mono(PlumType::TVar("_".to_string())));
64
+ }
65
+ _ => {}
66
+ }
67
+ }
68
+
69
+ // Second pass: check each function body
70
+ for item in &source.items {
71
+ if let ast::Item::Fn(f) = item {
72
+ let mut local_errors = check_fn(f, &global_env);
73
+ errors.append(&mut local_errors);
74
+ }
75
+ }
76
+
77
+ if errors.is_empty() { Ok(()) } else { Err(errors) }
78
+ }
79
+
80
+ fn check_fn(f: &ast::Fn, global_env: &TypeEnv) -> Vec<CheckError> {
81
+ let mut errors = Vec::new();
82
+ let mut env = global_env.clone();
83
+
84
+ // Add params to env
85
+ for p in &f.params {
86
+ let ty = match &p.ty {
87
+ ast::ParamType::Type(t) => plum_type_from_ast(t),
88
+ ast::ParamType::Variadic(t) => plum_type_from_ast(t),
89
+ };
90
+ env.insert(p.name.clone(), TypeScheme::mono(ty));
91
+ }
92
+
93
+ let declared_ret = f.returns.as_ref()
94
+ .map(|r| {
95
+ let ast_ty = ast::Type { name: r.name.clone(), generics: vec![] };
96
+ plum_type_from_ast(&ast_ty)
97
+ })
98
+ .unwrap_or(PlumType::TUnit);
99
+
100
+ match &f.body {
101
+ ast::FnBody::Expr(e) => {
102
+ match infer_expr(e, &env) {
103
+ Ok(t) => {
104
+ if let Err(msg) = unify(&declared_ret, &t) {
105
+ errors.push(CheckError { message: format!("fn '{}': return type mismatch: {}", f.name, msg) });
106
+ }
107
+ }
108
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': {}", f.name, msg) }),
109
+ }
110
+ }
111
+ ast::FnBody::Block(block) => {
112
+ let mut block_errors = check_block(block, &mut env, &declared_ret, &f.name);
113
+ errors.append(&mut block_errors);
114
+ }
115
+ }
116
+ errors
117
+ }
118
+
119
+ fn check_block(block: &ast::Block, env: &mut TypeEnv, declared_ret: &PlumType, fn_name: &str) -> Vec<CheckError> {
120
+ let mut errors = Vec::new();
121
+ for stmt in &block.stmts {
122
+ let mut stmt_errors = check_stmt(stmt, env, declared_ret, fn_name);
123
+ errors.append(&mut stmt_errors);
124
+ }
125
+ errors
126
+ }
127
+
128
+ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_name: &str) -> Vec<CheckError> {
129
+ let mut errors = Vec::new();
130
+ match stmt {
131
+ ast::Stmt::Assign(a) => {
132
+ for (target, value) in a.targets.iter().zip(a.values.iter()) {
133
+ match infer_expr(value, env) {
134
+ Ok(t) => { env.insert(target.clone(), TypeScheme::mono(t)); }
135
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': assign '{}': {}", fn_name, target, msg) }),
136
+ }
137
+ }
138
+ }
139
+ ast::Stmt::Return(Some(e)) => {
140
+ match infer_expr(e, env) {
141
+ Ok(t) => {
142
+ if let Err(msg) = unify(declared_ret, &t) {
143
+ errors.push(CheckError { message: format!("fn '{}': return type mismatch: {}", fn_name, msg) });
144
+ }
145
+ }
146
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': return: {}", fn_name, msg) }),
147
+ }
148
+ }
149
+ ast::Stmt::Return(None) => {
150
+ if let Err(msg) = unify(declared_ret, &PlumType::TUnit) {
151
+ errors.push(CheckError { message: format!("fn '{}': bare return in non-Unit function: {}", fn_name, msg) });
152
+ }
153
+ }
154
+ ast::Stmt::If(if_) => {
155
+ match infer_expr(&if_.condition, env) {
156
+ Ok(t) => {
157
+ if let Err(msg) = unify(&PlumType::TBool, &t) {
158
+ errors.push(CheckError { message: format!("fn '{}': if condition must be Bool: {}", fn_name, msg) });
159
+ }
160
+ }
161
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': if condition: {}", fn_name, msg) }),
162
+ }
163
+ errors.append(&mut check_block(&if_.body, env, declared_ret, fn_name));
164
+ for ei in &if_.else_ifs {
165
+ match infer_expr(&ei.condition, env) {
166
+ Ok(t) => {
167
+ if let Err(msg) = unify(&PlumType::TBool, &t) {
168
+ errors.push(CheckError { message: format!("fn '{}': else if condition must be Bool: {}", fn_name, msg) });
169
+ }
170
+ }
171
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': else if condition: {}", fn_name, msg) }),
172
+ }
173
+ errors.append(&mut check_block(&ei.body, env, declared_ret, fn_name));
174
+ }
175
+ if let Some(else_block) = &if_.else_ {
176
+ errors.append(&mut check_block(else_block, env, declared_ret, fn_name));
177
+ }
178
+ }
179
+ ast::Stmt::While(w) => {
180
+ match infer_expr(&w.condition, env) {
181
+ Ok(t) => {
182
+ if let Err(msg) = unify(&PlumType::TBool, &t) {
183
+ errors.push(CheckError { message: format!("fn '{}': while condition must be Bool: {}", fn_name, msg) });
184
+ }
185
+ }
186
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': while condition: {}", fn_name, msg) }),
187
+ }
188
+ errors.append(&mut check_block(&w.body, env, declared_ret, fn_name));
189
+ }
190
+ ast::Stmt::For(f_stmt) => {
191
+ match infer_expr(&f_stmt.iter, env) {
192
+ Ok(_) => {}
193
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': for iter: {}", fn_name, msg) }),
194
+ }
195
+ let mut inner_env = env.clone();
196
+ for var in &f_stmt.vars {
197
+ inner_env.insert(var.clone(), TypeScheme::mono(PlumType::TInt));
198
+ }
199
+ errors.append(&mut check_block(&f_stmt.body, &mut inner_env, declared_ret, fn_name));
200
+ }
201
+ ast::Stmt::Expr(e) => {
202
+ if let Err(msg) = infer_expr(e, env) {
203
+ errors.push(CheckError { message: format!("fn '{}': {}", fn_name, msg) });
204
+ }
205
+ }
206
+ ast::Stmt::Assert(e) => {
207
+ match infer_expr(e, env) {
208
+ Ok(t) => {
209
+ if let Err(msg) = unify(&PlumType::TBool, &t) {
210
+ errors.push(CheckError { message: format!("fn '{}': assert must be Bool: {}", fn_name, msg) });
211
+ }
212
+ }
213
+ Err(msg) => errors.push(CheckError { message: format!("fn '{}': assert: {}", fn_name, msg) }),
214
+ }
215
+ }
216
+ ast::Stmt::Break | ast::Stmt::Continue | ast::Stmt::Todo => {}
217
+ ast::Stmt::Match(_) => {}
218
+ }
219
+ errors
220
+ }
221
+
222
+ fn infer_expr(expr: &ast::Expr, env: &TypeEnv) -> Result<PlumType, String> {
223
+ match expr {
224
+ ast::Expr::Int(_) => Ok(PlumType::TInt),
225
+ ast::Expr::Float(_) => Ok(PlumType::TFloat),
226
+ ast::Expr::String(_) => Ok(PlumType::TStr),
227
+ ast::Expr::Var(name) => lookup(env, name),
228
+ ast::Expr::Self_ => Ok(PlumType::TVar("Self".to_string())),
229
+ ast::Expr::TypeName(n) => Ok(PlumType::TNamed(n.clone())),
230
+ ast::Expr::Paren(inner) => infer_expr(inner, env),
231
+ ast::Expr::Not(inner) => {
232
+ let t = infer_expr(inner, env)?;
233
+ unify(&PlumType::TBool, &t)?;
234
+ Ok(PlumType::TBool)
235
+ }
236
+ ast::Expr::Unary(u) => infer_expr(&u.operand, env),
237
+ ast::Expr::Binary(b) => {
238
+ let lt = infer_expr(&b.left, env)?;
239
+ let rt = infer_expr(&b.right, env)?;
240
+ unify(&lt, &rt).map_err(|e| format!("binary op: {}", e))?;
241
+ match b.op {
242
+ ast::BinOp::Range => Ok(PlumType::TNamed("Range".to_string())),
243
+ _ => Ok(lt),
244
+ }
245
+ }
246
+ ast::Expr::Bool(b) => {
247
+ let lt = infer_expr(&b.left, env)?;
248
+ let rt = infer_expr(&b.right, env)?;
249
+ unify(&PlumType::TBool, &lt).map_err(|e| format!("bool op left: {}", e))?;
250
+ unify(&PlumType::TBool, &rt).map_err(|e| format!("bool op right: {}", e))?;
251
+ Ok(PlumType::TBool)
252
+ }
253
+ ast::Expr::Compare(c) => {
254
+ let lt = infer_expr(&c.left, env)?;
255
+ let rt = infer_expr(&c.right, env)?;
256
+ unify(&lt, &rt).map_err(|e| format!("compare op: {}", e))?;
257
+ Ok(PlumType::TBool)
258
+ }
259
+ ast::Expr::Ternary(t) => {
260
+ let ct = infer_expr(&t.condition, env)?;
261
+ unify(&PlumType::TBool, &ct).map_err(|e| format!("ternary condition: {}", e))?;
262
+ let tt = infer_expr(&t.then, env)?;
263
+ let et = infer_expr(&t.else_, env)?;
264
+ unify(&tt, &et).map_err(|e| format!("ternary branches: {}", e))?;
265
+ Ok(tt)
266
+ }
267
+ ast::Expr::FnCall(call) => {
268
+ match lookup(env, &call.name) {
269
+ Ok(PlumType::TFun(param_types, ret)) => {
270
+ if call.args.len() != param_types.len() {
271
+ return Err(format!("call '{}': expected {} args, got {}", call.name, param_types.len(), call.args.len()));
272
+ }
273
+ for (i, (arg, expected)) in call.args.iter().zip(param_types.iter()).enumerate() {
274
+ let arg_expr = match arg {
275
+ ast::Arg::Positional(e) => e,
276
+ ast::Arg::Keyword { value, .. } => value,
277
+ ast::Arg::Pair { value, .. } => value,
278
+ };
279
+ let actual = infer_expr(arg_expr, env)?;
280
+ unify(expected, &actual).map_err(|e| format!("call '{}' arg {}: {}", call.name, i, e))?;
281
+ }
282
+ Ok(*ret)
283
+ }
284
+ Ok(_) => Err(format!("'{}' is not a function", call.name)),
285
+ Err(_) => Ok(PlumType::TVar("_".to_string())), // unknown fn: allow, codegen will catch
286
+ }
287
+ }
288
+ ast::Expr::ClassCall(_) => Ok(PlumType::TVar("_".to_string())),
289
+ ast::Expr::Attribute(_) => Ok(PlumType::TVar("_".to_string())),
290
+ }
291
+ }
plum-checker/tests/checker_tests.rs CHANGED
@@ -1,4 +1,29 @@
1
1
  use plum_checker::types::*;
2
+ use plum_checker::{plum_type_from_ast, unify};
3
+ use plum_core::ast::Type as AstType;
4
+
5
+ #[test]
6
+ fn ast_type_int_maps_to_tint() {
7
+ let ast_ty = AstType { name: "Int".to_string(), generics: vec![] };
8
+ assert_eq!(plum_type_from_ast(&ast_ty), PlumType::TInt);
9
+ }
10
+
11
+ #[test]
12
+ fn ast_type_unknown_maps_to_named() {
13
+ let ast_ty = AstType { name: "MyClass".to_string(), generics: vec![] };
14
+ assert_eq!(plum_type_from_ast(&ast_ty), PlumType::TNamed("MyClass".to_string()));
15
+ }
16
+
17
+ #[test]
18
+ fn unify_same_types_ok() {
19
+ assert!(unify(&PlumType::TInt, &PlumType::TInt).is_ok());
20
+ assert!(unify(&PlumType::TFloat, &PlumType::TFloat).is_ok());
21
+ }
22
+
23
+ #[test]
24
+ fn unify_different_types_err() {
25
+ assert!(unify(&PlumType::TInt, &PlumType::TFloat).is_err());
26
+ }
2
27
 
3
28
  #[test]
4
29
  fn fresh_vars_are_unique() {