plum

#treesitter#compiler#wasm

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

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


4a494a8Peter 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 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
- match infer_expr(&f_stmt.iter, env, ctx) {
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
- for var in &f_stmt.vars {
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
- inner_env.insert(var.clone(), TypeScheme::mono(PlumType::TInt));
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 = &param_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
- if call.args.len() != param_types.len() {
597
+ if call.args.len() != param_types.len() {
544
- return Err(format!("call '{}': expected {} args, got {}", call.name, param_types.len(), call.args.len()));
598
+ return Err(format!("call '{}': expected {} args, got {}", call.name, param_types.len(), call.args.len()));
545
- }
599
+ }
546
- for (i, (arg, expected)) in call.args.iter().zip(param_types.iter()).enumerate() {
600
+ for (i, (arg, expected)) in call.args.iter().zip(param_types.iter()).enumerate() {
547
- let arg_expr = match arg {
601
+ let arg_expr = match arg {
548
- ast::Arg::Positional(e) => e,
602
+ ast::Arg::Positional(e) => e,
549
- ast::Arg::Keyword { value, .. } => value,
603
+ ast::Arg::Keyword { value, .. } => value,
550
- ast::Arg::Pair { value, .. } => value,
604
+ ast::Arg::Pair { value, .. } => value,
551
- };
605
+ };
552
- let actual = infer_expr(arg_expr, env, ctx)?;
606
+ let actual = infer_expr(arg_expr, env, ctx)?;
553
- unify(expected, &actual).map_err(|e| format!("call '{}' arg {}: {}", call.name, i, e))?;
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
+ }