plum
git clone https://git.pyrossh.dev/plum
A statically typed, imperative programming language inspired by rust, python
4a494a8
— Peter John
2026-07-23T19:52:22+05:30
feat(plum-checker): type-check variadic call arity and for-in-variadic iteration
- plum-checker/src/lib.rs +74 -18
- plum-checker/tests/checker_tests.rs +97 -0
plum-checker/src/lib.rs
CHANGED
|
@@ -174,6 +174,18 @@ fn check_fn(f: &ast::Fn, global_env: &TypeEnv, ctx: &CheckCtx) -> Vec<CheckError
|
|
|
174
174
|
let mut errors = Vec::new();
|
|
175
175
|
let mut env = global_env.clone();
|
|
176
176
|
|
|
177
|
+
let variadic_positions: Vec<usize> = f.params.iter().enumerate()
|
|
178
|
+
.filter(|(_, p)| matches!(p.ty, ast::ParamType::Variadic(_)))
|
|
179
|
+
.map(|(i, _)| i)
|
|
180
|
+
.collect();
|
|
181
|
+
if variadic_positions.len() > 1 {
|
|
182
|
+
errors.push(CheckError { message: format!("fn '{}': at most one variadic parameter is allowed", f.name) });
|
|
183
|
+
} else if let Some(&pos) = variadic_positions.first() {
|
|
184
|
+
if pos != f.params.len() - 1 {
|
|
185
|
+
errors.push(CheckError { message: format!("fn '{}': a variadic parameter must be last", f.name) });
|
|
186
|
+
}
|
|
187
|
+
}
|
|
188
|
+
|
|
177
189
|
// Methods (`name<Receiver>(...)`) get an implicit `self: Receiver` binding.
|
|
178
190
|
if let Some(recv) = &f.type_param {
|
|
179
191
|
let recv_ty = ast::Type { name: recv.clone(), generics: vec![] };
|
|
@@ -342,13 +354,28 @@ fn check_stmt(stmt: &ast::Stmt, env: &mut TypeEnv, declared_ret: &PlumType, fn_n
|
|
|
342
354
|
errors.append(&mut check_block(&w.body, env, declared_ret, fn_name, ctx));
|
|
343
355
|
}
|
|
344
356
|
ast::Stmt::For(f_stmt) => {
|
|
345
|
-
|
|
357
|
+
let iter_ty = infer_expr(&f_stmt.iter, env, ctx);
|
|
346
|
-
Ok(_) => {}
|
|
347
|
-
Err(msg) => errors.push(CheckError { message: format!("fn '{}': for iter: {}", fn_name, msg) }),
|
|
348
|
-
}
|
|
349
358
|
let mut inner_env = env.clone();
|
|
359
|
+
match &iter_ty {
|
|
360
|
+
Ok(PlumType::TVariadic(elem)) => {
|
|
361
|
+
if f_stmt.vars.len() != 1 {
|
|
362
|
+
errors.push(CheckError { message: format!("fn '{}': for-loop over a variadic param must bind exactly one variable", fn_name) });
|
|
363
|
+
}
|
|
350
|
-
|
|
364
|
+
for var in &f_stmt.vars {
|
|
365
|
+
inner_env.insert(var.clone(), TypeScheme::mono((**elem).clone()));
|
|
366
|
+
}
|
|
367
|
+
}
|
|
368
|
+
Ok(_) => {
|
|
369
|
+
for var in &f_stmt.vars {
|
|
351
|
-
|
|
370
|
+
inner_env.insert(var.clone(), TypeScheme::mono(PlumType::TInt));
|
|
371
|
+
}
|
|
372
|
+
}
|
|
373
|
+
Err(msg) => {
|
|
374
|
+
errors.push(CheckError { message: format!("fn '{}': for iter: {}", fn_name, msg) });
|
|
375
|
+
for var in &f_stmt.vars {
|
|
376
|
+
inner_env.insert(var.clone(), TypeScheme::mono(PlumType::TInt));
|
|
377
|
+
}
|
|
378
|
+
}
|
|
352
379
|
}
|
|
353
380
|
errors.append(&mut check_block(&f_stmt.body, &mut inner_env, declared_ret, fn_name, ctx));
|
|
354
381
|
}
|
|
@@ -540,19 +567,48 @@ pub fn infer_expr(expr: &ast::Expr, env: &TypeEnv, ctx: &CheckCtx) -> Result<Plu
|
|
|
540
567
|
}
|
|
541
568
|
match lookup(env, &call.name) {
|
|
542
569
|
Ok(PlumType::TFun(param_types, ret)) => {
|
|
570
|
+
match param_types.last() {
|
|
571
|
+
Some(PlumType::TVariadic(elem)) => {
|
|
572
|
+
let fixed = ¶m_types[..param_types.len() - 1];
|
|
573
|
+
if call.args.len() < fixed.len() {
|
|
574
|
+
return Err(format!("call '{}': expected at least {} arg(s), got {}", call.name, fixed.len(), call.args.len()));
|
|
575
|
+
}
|
|
576
|
+
for (i, (arg, expected)) in call.args.iter().zip(fixed.iter()).enumerate() {
|
|
577
|
+
let arg_expr = match arg {
|
|
578
|
+
ast::Arg::Positional(e) => e,
|
|
579
|
+
ast::Arg::Keyword { value, .. } => value,
|
|
580
|
+
ast::Arg::Pair { value, .. } => value,
|
|
581
|
+
};
|
|
582
|
+
let actual = infer_expr(arg_expr, env, ctx)?;
|
|
583
|
+
unify(expected, &actual).map_err(|e| format!("call '{}' arg {}: {}", call.name, i, e))?;
|
|
584
|
+
}
|
|
585
|
+
for (i, arg) in call.args.iter().enumerate().skip(fixed.len()) {
|
|
586
|
+
let arg_expr = match arg {
|
|
587
|
+
ast::Arg::Positional(e) => e,
|
|
588
|
+
ast::Arg::Keyword { value, .. } => value,
|
|
589
|
+
ast::Arg::Pair { value, .. } => value,
|
|
590
|
+
};
|
|
591
|
+
let actual = infer_expr(arg_expr, env, ctx)?;
|
|
592
|
+
unify(elem, &actual).map_err(|e| format!("call '{}' variadic arg {}: {}", call.name, i, e))?;
|
|
593
|
+
}
|
|
594
|
+
Ok(*ret)
|
|
595
|
+
}
|
|
596
|
+
_ => {
|
|
543
|
-
|
|
597
|
+
if call.args.len() != param_types.len() {
|
|
544
|
-
|
|
598
|
+
return Err(format!("call '{}': expected {} args, got {}", call.name, param_types.len(), call.args.len()));
|
|
545
|
-
|
|
599
|
+
}
|
|
546
|
-
|
|
600
|
+
for (i, (arg, expected)) in call.args.iter().zip(param_types.iter()).enumerate() {
|
|
547
|
-
|
|
601
|
+
let arg_expr = match arg {
|
|
548
|
-
|
|
602
|
+
ast::Arg::Positional(e) => e,
|
|
549
|
-
|
|
603
|
+
ast::Arg::Keyword { value, .. } => value,
|
|
550
|
-
|
|
604
|
+
ast::Arg::Pair { value, .. } => value,
|
|
551
|
-
|
|
605
|
+
};
|
|
552
|
-
|
|
606
|
+
let actual = infer_expr(arg_expr, env, ctx)?;
|
|
553
|
-
|
|
607
|
+
unify(expected, &actual).map_err(|e| format!("call '{}' arg {}: {}", call.name, i, e))?;
|
|
608
|
+
}
|
|
609
|
+
Ok(*ret)
|
|
610
|
+
}
|
|
554
611
|
}
|
|
555
|
-
Ok(*ret)
|
|
556
612
|
}
|
|
557
613
|
Ok(_) => Err(format!("'{}' is not a function", call.name)),
|
|
558
614
|
Err(_) => Ok(PlumType::TVar("_".to_string())), // unknown fn: allow, codegen will catch
|
plum-checker/tests/checker_tests.rs
CHANGED
|
@@ -655,3 +655,100 @@ use() -> Bool =
|
|
|
655
655
|
let result = check_source(&source);
|
|
656
656
|
assert!(result.is_ok(), "expected Ok, got {:?}", result.err());
|
|
657
657
|
}
|
|
658
|
+
|
|
659
|
+
#[test]
|
|
660
|
+
fn variadic_call_with_zero_trailing_args_passes() {
|
|
661
|
+
let src = "\
|
|
662
|
+
sumAll(nums: ...Int) -> Int =
|
|
663
|
+
0
|
|
664
|
+
|
|
665
|
+
useSumAll() -> Int =
|
|
666
|
+
sumAll()
|
|
667
|
+
";
|
|
668
|
+
let source = parse(src);
|
|
669
|
+
assert!(check_source(&source).is_ok(), "expected Ok");
|
|
670
|
+
}
|
|
671
|
+
|
|
672
|
+
#[test]
|
|
673
|
+
fn variadic_call_with_several_trailing_args_passes() {
|
|
674
|
+
let src = "\
|
|
675
|
+
sumAll(nums: ...Int) -> Int =
|
|
676
|
+
0
|
|
677
|
+
|
|
678
|
+
useSumAll() -> Int =
|
|
679
|
+
sumAll(1, 2, 3)
|
|
680
|
+
";
|
|
681
|
+
let source = parse(src);
|
|
682
|
+
assert!(check_source(&source).is_ok(), "expected Ok");
|
|
683
|
+
}
|
|
684
|
+
|
|
685
|
+
#[test]
|
|
686
|
+
fn variadic_call_with_mismatched_trailing_arg_type_is_error() {
|
|
687
|
+
let src = "\
|
|
688
|
+
sumAll(nums: ...Int) -> Int =
|
|
689
|
+
0
|
|
690
|
+
|
|
691
|
+
useSumAll() -> Int =
|
|
692
|
+
sumAll(1, \"two\")
|
|
693
|
+
";
|
|
694
|
+
let source = parse(src);
|
|
695
|
+
assert!(check_source(&source).is_err());
|
|
696
|
+
}
|
|
697
|
+
|
|
698
|
+
#[test]
|
|
699
|
+
fn variadic_call_with_fixed_prefix_passes() {
|
|
700
|
+
let src = "\
|
|
701
|
+
combine(prefix: Int, rest: ...Int) -> Int =
|
|
702
|
+
prefix
|
|
703
|
+
|
|
704
|
+
useCombine() -> Int =
|
|
705
|
+
combine(1, 2, 3)
|
|
706
|
+
";
|
|
707
|
+
let source = parse(src);
|
|
708
|
+
assert!(check_source(&source).is_ok(), "expected Ok");
|
|
709
|
+
}
|
|
710
|
+
|
|
711
|
+
#[test]
|
|
712
|
+
fn two_variadic_params_is_error() {
|
|
713
|
+
let src = "\
|
|
714
|
+
bad(a: ...Int, b: ...Int) -> Int =
|
|
715
|
+
0
|
|
716
|
+
";
|
|
717
|
+
let source = parse(src);
|
|
718
|
+
assert!(check_source(&source).is_err());
|
|
719
|
+
}
|
|
720
|
+
|
|
721
|
+
#[test]
|
|
722
|
+
fn variadic_param_not_last_is_error() {
|
|
723
|
+
let src = "\
|
|
724
|
+
bad(a: ...Int, b: Int) -> Int =
|
|
725
|
+
0
|
|
726
|
+
";
|
|
727
|
+
let source = parse(src);
|
|
728
|
+
assert!(check_source(&source).is_err());
|
|
729
|
+
}
|
|
730
|
+
|
|
731
|
+
#[test]
|
|
732
|
+
fn for_loop_over_variadic_binds_element_type() {
|
|
733
|
+
let src = "\
|
|
734
|
+
sumAll(nums: ...Int) -> Int =
|
|
735
|
+
total = 0
|
|
736
|
+
for v in nums
|
|
737
|
+
total = total + v
|
|
738
|
+
total
|
|
739
|
+
";
|
|
740
|
+
let source = parse(src);
|
|
741
|
+
assert!(check_source(&source).is_ok(), "expected Ok");
|
|
742
|
+
}
|
|
743
|
+
|
|
744
|
+
#[test]
|
|
745
|
+
fn for_loop_over_variadic_with_two_vars_is_error() {
|
|
746
|
+
let src = "\
|
|
747
|
+
bad(nums: ...Int) -> Int =
|
|
748
|
+
for v, i in nums
|
|
749
|
+
v
|
|
750
|
+
0
|
|
751
|
+
";
|
|
752
|
+
let source = parse(src);
|
|
753
|
+
assert!(check_source(&source).is_err());
|
|
754
|
+
}
|