Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 14 additions & 7 deletions src/checker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -722,7 +722,12 @@ impl Checker {
}
Expr::Arena(block) => self.check_expr(*block, arena, decls),
Expr::Return(expr) => {
let ty = self.check_expr(*expr, arena, decls);
// A bare `return` returns nothing, so it has to line up with
// a function that declares no return type.
let ty = match expr {
Some(e) => self.check_expr(*e, arena, decls),
None => mk_type(Type::Void),
};

// Constrain against the enclosing function's return type.
// Without this, only a return in tail position is checked
Expand Down Expand Up @@ -1526,7 +1531,8 @@ fn escape_walk(
errors: &mut Vec<TypeError>,
) {
match &arena[expr] {
Expr::Return(e) => {
Expr::Return(None) => {}
Expr::Return(Some(e)) => {
// Walk first so any declarations inside `e` update `tainted`.
escape_walk(*e, arena, scope, tainted, errors);
if expr_is_tainted(*e, arena, scope, tainted) {
Expand Down Expand Up @@ -1695,11 +1701,12 @@ fn lambda_has_captures(
.map(|e| lambda_has_captures(e, arena, lambda_params, outer_scope))
.unwrap_or(false)
}
Expr::Return(e)
| Expr::Assume(e)
| Expr::Field(e, _)
| Expr::AsTy(e, _)
| Expr::Arena(e) => lambda_has_captures(*e, arena, lambda_params, outer_scope),
Expr::Return(e) => e
.map(|e| lambda_has_captures(e, arena, lambda_params, outer_scope))
.unwrap_or(false),
Expr::Assume(e) | Expr::Field(e, _) | Expr::AsTy(e, _) | Expr::Arena(e) => {
lambda_has_captures(*e, arena, lambda_params, outer_scope)
}
Expr::While(cond, body) | Expr::ArrayIndex(cond, body) | Expr::Array(cond, body) => {
lambda_has_captures(*cond, arena, lambda_params, outer_scope)
|| lambda_has_captures(*body, arena, lambda_params, outer_scope)
Expand Down
20 changes: 15 additions & 5 deletions src/expr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -86,8 +86,9 @@ pub enum Expr {
/// Block expression containing a list of expressions.
Block(Vec<ExprID>),

/// Return expression.
Return(ExprID),
/// Return expression. `None` is a bare `return`, which is only valid in a
/// function with no declared return type.
Return(Option<ExprID>),

/// Break out of the innermost loop.
Break,
Expand Down Expand Up @@ -299,11 +300,13 @@ impl Expr {
format!("{{\n{}\n{}}}", exprs_str, indent_str)
}

Expr::Return(expr) => {
Expr::Return(Some(expr)) => {
let expr_str = arena.exprs[*expr].pretty_print(arena, indent);
format!("return {}", expr_str)
}

Expr::Return(None) => "return".to_string(),

Expr::Assume(expr) => {
let expr_str = arena.exprs[*expr].pretty_print(arena, indent);
format!("assume {}", expr_str)
Expand Down Expand Up @@ -473,7 +476,7 @@ pub fn copy_expr(
dst_arena.add(Expr::Block(new_exprs), loc)
}
Expr::Return(expr) => {
let new_expr = copy_expr(*expr, src_arena, dst_arena, subst);
let new_expr = expr.map(|e| copy_expr(e, src_arena, dst_arena, subst));
dst_arena.add(Expr::Return(new_expr), loc)
}
Expr::Assume(expr) => {
Expand Down Expand Up @@ -784,11 +787,18 @@ mod tests {
let mut arena = ExprArena::new();

let value_id = arena.add(Expr::Int(42, None), test_loc());
let return_expr = Expr::Return(value_id);
let return_expr = Expr::Return(Some(value_id));

assert_eq!(return_expr.pretty_print(&arena, 0), "return 42");
}

#[test]
fn test_pretty_print_bare_return() {
let arena = ExprArena::new();

assert_eq!(Expr::Return(None).pretty_print(&arena, 0), "return");
}

#[test]
fn test_pretty_print_tuple() {
let mut arena = ExprArena::new();
Expand Down
6 changes: 3 additions & 3 deletions src/hoist.rs
Original file line number Diff line number Diff line change
Expand Up @@ -199,7 +199,7 @@ fn collect_written_fields(expr_id: ExprID, fdecl: &FuncDecl, written: &mut HashS
collect_written_fields(*base, fdecl, written);
collect_written_fields(*idx, fdecl, written);
}
Expr::Return(e) | Expr::Assume(e) => {
Expr::Return(Some(e)) | Expr::Assume(e) => {
collect_written_fields(*e, fdecl, written);
}
Expr::Var(_, init, _) => {
Expand Down Expand Up @@ -298,7 +298,7 @@ fn collect_invariant_field_reads(
collect_invariant_field_reads(*base, fdecl, decls, written, reads);
collect_invariant_field_reads(*idx, fdecl, decls, written, reads);
}
Expr::Return(e) | Expr::Assume(e) | Expr::AsTy(e, _) => {
Expr::Return(Some(e)) | Expr::Assume(e) | Expr::AsTy(e, _) => {
collect_invariant_field_reads(*e, fdecl, decls, written, reads);
}
Expr::Var(_, init, _) => {
Expand Down Expand Up @@ -384,7 +384,7 @@ fn replace_field_reads(expr_id: ExprID, fdecl: &mut FuncDecl, subst: &HashMap<(N
replace_field_reads(base, fdecl, subst);
replace_field_reads(idx, fdecl, subst);
}
Expr::Return(e) | Expr::Assume(e) | Expr::AsTy(e, _) => {
Expr::Return(Some(e)) | Expr::Assume(e) | Expr::AsTy(e, _) => {
replace_field_reads(e, fdecl, subst);
}
Expr::Var(_, init, _) => {
Expand Down
56 changes: 41 additions & 15 deletions src/jit.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1625,24 +1625,45 @@ impl<'a> FunctionTranslator<'a> {
self.builder.ins().iconst(I32, 0)
}
Expr::Return(expr_id) => {
let result = self.translate_expr(*expr_id, decl, decls);
let ret_ty = decl.types[*expr_id];
// Evaluate the operand before releasing the call depth: it
// may itself contain calls.
let operand = expr_id.map(|expr_id| {
(
self.translate_expr(expr_id, decl, decls),
decl.types[expr_id],
)
});

if !self.no_recursion {
self.emit_call_depth_release();
}
if returns_via_pointer(ret_ty) {
// Copy result to output pointer and return void.
let output = self
.output_ptr
.expect("output_ptr not set for pointer return");
let size = ret_ty.size(decls) as i64;
let size_val = self.builder.ins().iconst(I64, size);
self.builder
.call_memcpy(self.module.target_config(), output, result, size_val);
self.builder.ins().return_(&[]);
} else {
self.builder.ins().return_(&[result]);
match operand {
Some((result, ret_ty)) if returns_via_pointer(ret_ty) => {
// Copy result to output pointer and return void.
let output = self
.output_ptr
.expect("output_ptr not set for pointer return");
let size = ret_ty.size(decls) as i64;
let size_val = self.builder.ins().iconst(I64, size);
self.builder.call_memcpy(
self.module.target_config(),
output,
result,
size_val,
);
self.builder.ins().return_(&[]);
}
// A bare return, or one whose operand is itself void, has
// no value to hand back — matching the void epilogue above.
None => {
self.builder.ins().return_(&[]);
}
Some((_, ret_ty)) if *ret_ty == crate::Type::Void => {
self.builder.ins().return_(&[]);
}
Some((result, _)) => {
self.builder.ins().return_(&[result]);
}
}

// Create an unreachable block for any code after return.
Expand Down Expand Up @@ -2669,7 +2690,12 @@ fn collect_free_vars_rec(
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
}
Expr::Return(e) | Expr::Assume(e) => {
Expr::Return(e) => {
if let Some(e) = e {
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
}
Expr::Assume(e) => {
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
Expr::Field(e, _) => {
Expand Down
42 changes: 29 additions & 13 deletions src/llvm_jit.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2275,22 +2275,33 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> {
}
Expr::Return(ret_id) => {
let ret_id = *ret_id;
let result = self.translate_expr(ret_id, decl);
let ret_ty = decl.types[ret_id];
// Evaluate the operand before releasing the call depth: it
// may itself contain calls.
let operand =
ret_id.map(|ret_id| (self.translate_expr(ret_id, decl), decl.types[ret_id]));

if !self.state.no_recursion {
self.emit_call_depth_release();
}
if returns_via_pointer(ret_ty) {
let out = self.output_ptr.unwrap();
let size = ret_ty.size(self.decls) as u64;
self.emit_memcpy(out, result.into_pointer_value(), size);
self.builder().build_return(None).unwrap();
} else if *ret_ty == crate::Type::Void {
self.builder().build_return(None).unwrap();
} else {
let coerced = self.coerce_return(result, decl.ret);
self.builder().build_return(Some(&coerced)).unwrap();
match operand {
Some((result, ret_ty)) if returns_via_pointer(ret_ty) => {
let out = self.output_ptr.unwrap();
let size = ret_ty.size(self.decls) as u64;
self.emit_memcpy(out, result.into_pointer_value(), size);
self.builder().build_return(None).unwrap();
}
// A bare return, or one whose operand is itself void,
// hands back no value.
None => {
self.builder().build_return(None).unwrap();
}
Some((_, ret_ty)) if *ret_ty == crate::Type::Void => {
self.builder().build_return(None).unwrap();
}
Some((result, _)) => {
let coerced = self.coerce_return(result, decl.ret);
self.builder().build_return(Some(&coerced)).unwrap();
}
}

// Unreachable block after return — add terminator so epilogue is skipped.
Expand Down Expand Up @@ -3899,7 +3910,12 @@ fn collect_free_vars_rec_llvm(
collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen);
}
}
Expr::Return(e) | Expr::Assume(e) => {
Expr::Return(e) => {
if let Some(e) = e {
collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen)
}
}
Expr::Assume(e) => {
collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen)
}
Expr::Field(e, _) => {
Expand Down
7 changes: 6 additions & 1 deletion src/monomorph_pass.rs
Original file line number Diff line number Diff line change
Expand Up @@ -396,7 +396,12 @@ impl MonomorphPass {
Expr::Lambda { body, .. } => {
self.process_expr(*body, fdecl, decls)?;
}
Expr::Return(inner) | Expr::Assume(inner) => {
Expr::Return(inner) => {
if let Some(inner) = inner {
self.process_expr(*inner, fdecl, decls)?;
}
}
Expr::Assume(inner) => {
self.process_expr(*inner, fdecl, decls)?;
}
Expr::For {
Expand Down
10 changes: 9 additions & 1 deletion src/parser.rs
Original file line number Diff line number Diff line change
Expand Up @@ -787,7 +787,15 @@ fn parse_stmt(arena: &mut ExprArena, typevars: &[Name], cx: &mut ParseContext) -
Token::Return => {
let loc = cx.lex.loc;
cx.next();
let e = parse_expr(arena, typevars, cx);

// A bare `return` is one with nothing left in the statement:
// end of line, end of block, or end of input.
let e = if matches!(cx.lex.tok, Token::Endl | Token::Rbrace | Token::End) {
None
} else {
Some(parse_expr(arena, typevars, cx))
};

arena.add(Expr::Return(e), loc)
}
Token::Break => {
Expand Down
7 changes: 5 additions & 2 deletions src/safety_checker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -913,7 +913,9 @@ impl SafetyChecker {
}
}
Expr::Return(expr) => {
self.check_expr(*expr, decl, decls);
if let Some(e) = expr {
self.check_expr(*e, decl, decls);
}
IndexInterval::default()
}
Expr::Assume(cond) => {
Expand Down Expand Up @@ -1577,7 +1579,8 @@ impl SafetyChecker {
start, end, body, ..
} => vec![*start, *end, *body],
Expr::Block(es) | Expr::Tuple(es) => es.clone(),
Expr::Return(e) | Expr::Arena(e) | Expr::Assume(e) => vec![*e],
Expr::Return(e) => e.iter().copied().collect(),
Expr::Arena(e) | Expr::Assume(e) => vec![*e],
Expr::StructLit(_, fields) => fields.iter().map(|(_, e)| *e).collect(),
}
}
Expand Down
19 changes: 17 additions & 2 deletions src/stack_codegen.rs
Original file line number Diff line number Diff line change
Expand Up @@ -661,6 +661,11 @@ impl<'a> FunctionTranslator<'a> {
Expr::While(..) | Expr::For { .. } => {
self.translate_expr_inner(expr, func, true);
}
// Return never falls through, so there is no value to drop —
// a drop here would be unreachable code after the terminator.
Expr::Return(..) => {
self.translate_expr(expr, func);
}
// Everything else: translate normally then drop.
_ => {
self.translate_expr(expr, func);
Expand Down Expand Up @@ -915,7 +920,12 @@ impl<'a> FunctionTranslator<'a> {
}
}

Expr::Return(expr_id) => {
Expr::Return(None) => {
func.emit(StackOp::ReturnVoid);
self.has_returned = true;
}

Expr::Return(Some(expr_id)) => {
let expr_id = *expr_id;
let ret_ty = self.expr_type(expr_id);
self.translate_expr(expr_id, func);
Expand Down Expand Up @@ -2972,7 +2982,12 @@ fn collect_free_vars_rec(
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
}
Expr::Return(e) | Expr::Assume(e) => {
Expr::Return(e) => {
if let Some(e) = e {
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
}
Expr::Assume(e) => {
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
Expr::Field(e, _) => {
Expand Down
26 changes: 23 additions & 3 deletions src/vm_codegen.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1138,8 +1138,23 @@ impl<'a> FunctionTranslator<'a> {
result
}
Expr::Return(expr_id) => {
let result = self.translate_expr(*expr_id, func);
let ret_ty = self.expr_type(*expr_id);
// A bare return has no value to hand back, so use the same
// zero placeholder void-like expressions produce. The void
// epilogue below ignores r0 anyway.
let (result, ret_ty) = match expr_id {
Some(expr_id) => (
self.translate_expr(*expr_id, func),
self.expr_type(*expr_id),
),
None => {
let result = self.alloc_reg();
func.emit(Opcode::LoadImm {
dst: result,
value: 0,
});
(result, mk_type(Type::Void))
}
};

if returns_via_pointer(ret_ty) {
// Reload output pointer from local slot (r0 may have been
Expand Down Expand Up @@ -3611,7 +3626,12 @@ fn collect_free_vars_rec(
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
}
Expr::Return(e) | Expr::Assume(e) => {
Expr::Return(e) => {
if let Some(e) = e {
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
}
Expr::Assume(e) => {
collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen);
}
Expr::Field(e, _) => {
Expand Down
Loading
Loading