compiler: Retain and check conditional call argument loans

b02adb25ffc8a5a4477939fee9e17d43495373f6fe27470ed7683a63a07c67ac
Alexis Sellier committed ago 1 parent b2c3593c
lib/std/lang/resolver.rad +99 -20
841 841
    precise: bool,
842 842
}
843 843
844 844
/// A reference binding that protects its source for one lexical scope.
845 845
record LocalLoan: Copy {
846 -
    /// Local symbol through which the source can be accessed.
847 -
    binding: *unsafe mut Symbol,
846 +
    /// Local symbol that provides access, or nil for a pending call argument.
847 +
    binding: ?*unsafe mut Symbol,
848 848
    /// Storage retained by the reference.
849 849
    place: BorrowPlace,
850 850
    /// Whether other reads of the source are excluded.
851 851
    exclusive: bool,
852 852
}
7951 7951
{
7952 7952
    let place = borrowPlace(checker.resolver, node);
7953 7953
    let root = place.root else return;
7954 7954
    for i in 0..checker.localLen {
7955 7955
        let loan = checker.locals[i];
7956 +
        let mut throughBinding = false;
7957 +
        if let binding = loan.binding {
7958 +
            set throughBinding = usesLocalLoan(checker.resolver, node, binding);
7959 +
        }
7956 7960
        if (exclusive or loan.exclusive) and placesOverlap(&place, &loan.place)
7957 -
            and not usesLocalLoan(checker.resolver, node, loan.binding)
7961 +
            and not throughBinding
7958 7962
        {
7959 7963
            throw emitError(checker.resolver, node, ErrorKind::BorrowConflict(root.name));
7960 7964
        }
7961 7965
    }
7962 7966
}
8276 8280
    checker: &mut LinearChecker,
8277 8281
    env: &mut LinearEnv,
8278 8282
    node: *ast::Node,
8279 8283
    call: ast::Call,
8280 8284
) throws (ResolveError) {
8285 +
    let localStart = checker.localLen;
8281 8286
    match checker.resolver.nodeData.entries[node.id].extra {
8282 8287
        case NodeExtra::SliceAppend { .. }, NodeExtra::SliceDelete { .. } => {
8283 8288
            let case ast::NodeValue::FieldAccess(access) = call.callee.value
8284 8289
                else panic "slice mutation without receiver";
8285 8290
            try checkPatternLoan(checker, access.parent);
8305 8310
        for arg in call.args {
8306 8311
            try checkLinearNode(checker, env, arg, LinearUse::Consume);
8307 8312
        }
8308 8313
        return;
8309 8314
    };
8310 -
    let mut places: [BorrowPlace; MAX_FN_PARAMS + 1] = undefined;
8315 +
    let mut arguments: [*ast::Node; MAX_FN_PARAMS + 1] = undefined;
8311 8316
    let mut exclusive: [bool; MAX_FN_PARAMS + 1] = undefined;
8312 -
    let mut placesLen: u32 = 0;
8317 +
    let mut argumentsLen: u32 = 0;
8313 8318
8314 8319
    // Method function types exclude their implicit receiver. Account for it
8315 8320
    // explicitly so owning receivers are consumed and reference receivers
8316 8321
    // participate in call-scoped loan conflict checks.
8317 8322
    if let case ast::NodeValue::FieldAccess(access) = call.callee.value {
8337 8342
                receiverMutable or receiverClass == types::PointerClass::Owned);
8338 8343
            if receiverMutable or receiverClass == types::PointerClass::Owned {
8339 8344
                try checkPatternLoan(checker, access.parent);
8340 8345
            }
8341 8346
            if receiverClass <> types::PointerClass::Unsafe {
8342 -
                let place = borrowPlace(checker.resolver, access.parent);
8343 -
                if place.root <> nil {
8344 -
                    set places[placesLen] = place;
8345 -
                    set exclusive[placesLen] =
8346 -
                        receiverClass == types::PointerClass::Owned or receiverMutable;
8347 -
                    set placesLen += 1;
8348 -
                }
8347 +
                set arguments[argumentsLen] = access.parent;
8348 +
                set exclusive[argumentsLen] =
8349 +
                    receiverClass == types::PointerClass::Owned or receiverMutable;
8350 +
                set argumentsLen += 1;
8349 8351
            }
8350 8352
            if receiverClass == types::PointerClass::Ref {
8351 8353
                try checkLinearNode(checker, env, access.parent, LinearUse::Borrow);
8354 +
                if createsExplicitBorrow(access.parent) {
8355 +
                    try retainCallLoan(checker, access.parent, receiverMutable);
8356 +
                }
8352 8357
            } else if receiverClass == types::PointerClass::Owned {
8353 8358
                try checkLinearNode(checker, env, access.parent, LinearUse::Consume);
8354 8359
            }
8355 8360
        }
8356 8361
    }
8359 8364
        let expected = *info.paramTypes[i];
8360 8365
        let argExclusive = isExclusiveArgument(expected);
8361 8366
        if argExclusive {
8362 8367
            try checkPatternLoan(checker, arg);
8363 8368
        }
8364 -
        let place = borrowPlace(checker.resolver, arg);
8365 8369
        if not isUnsafePointerType(expected) {
8366 -
            if let rootSym = place.root {
8367 -
                for j in 0..placesLen {
8368 -
                    if (exclusive[j] or argExclusive) and placesOverlap(&places[j], &place) {
8369 -
                        throw emitError(checker.resolver, arg, ErrorKind::BorrowConflict(rootSym.name));
8370 +
            for j in 0..argumentsLen {
8371 +
                if exclusive[j] or argExclusive {
8372 +
                    if let name = callArgumentConflict(checker.resolver, arguments[j], arg) {
8373 +
                        throw emitError(checker.resolver, arg, ErrorKind::BorrowConflict(name));
8370 8374
                    }
8371 8375
                }
8372 -
                set places[placesLen] = place;
8373 -
                set exclusive[placesLen] = argExclusive;
8374 -
                set placesLen += 1;
8375 8376
            }
8377 +
            set arguments[argumentsLen] = arg;
8378 +
            set exclusive[argumentsLen] = argExclusive;
8379 +
            set argumentsLen += 1;
8376 8380
        }
8377 8381
        try checkLocalLoans(checker, arg, argExclusive);
8378 8382
        if isRefType(expected) {
8379 8383
            try checkLinearNode(checker, env, arg, LinearUse::Borrow);
8380 8384
        } else {
8381 8385
            try checkLinearNode(checker, env, arg, LinearUse::Consume);
8382 8386
        }
8387 +
        if isRefType(expected) and createsExplicitBorrow(arg) {
8388 +
            try retainCallLoan(checker, arg, argExclusive);
8389 +
        }
8383 8390
    }
8391 +
    set checker.localLen = localStart;
8384 8392
    if *info.returnType == Type::Never and info.throwList.len == 0 {
8385 8393
        set env.terminated = true;
8386 8394
    }
8387 8395
}
8388 8396
8397 +
/// Return the storage name when two call arguments can address the same place.
8398 +
unsafe fn callArgumentConflict(
8399 +
    self: &mut Resolver, left: *ast::Node, right: *ast::Node
8400 +
) -> ?*[u8] {
8401 +
    match left.value {
8402 +
        case ast::NodeValue::CondExpr(cond) => {
8403 +
            if let name = callArgumentConflict(self, cond.thenExpr, right) {
8404 +
                return name;
8405 +
            }
8406 +
            return callArgumentConflict(self, cond.elseExpr, right);
8407 +
        }
8408 +
        case ast::NodeValue::As(cast) => return callArgumentConflict(self, cast.value, right),
8409 +
        else => {}
8410 +
    }
8411 +
    match right.value {
8412 +
        case ast::NodeValue::CondExpr(cond) => {
8413 +
            if let name = callArgumentConflict(self, left, cond.thenExpr) {
8414 +
                return name;
8415 +
            }
8416 +
            return callArgumentConflict(self, left, cond.elseExpr);
8417 +
        }
8418 +
        case ast::NodeValue::As(cast) => return callArgumentConflict(self, left, cast.value),
8419 +
        else => {}
8420 +
    }
8421 +
    let leftPlace = borrowPlace(self, left);
8422 +
    let rightPlace = borrowPlace(self, right);
8423 +
    if placesOverlap(&leftPlace, &rightPlace) {
8424 +
        let root = rightPlace.root else panic "callArgumentConflict: overlap without root";
8425 +
        return root.name;
8426 +
    }
8427 +
    return nil;
8428 +
}
8429 +
8430 +
/// Return whether evaluating an argument creates an explicit address borrow.
8431 +
fn createsExplicitBorrow(node: *ast::Node) -> bool {
8432 +
    match node.value {
8433 +
        case ast::NodeValue::AddressOf(_) => return true,
8434 +
        case ast::NodeValue::As(cast) => return createsExplicitBorrow(cast.value),
8435 +
        case ast::NodeValue::CondExpr(cond) =>
8436 +
            return createsExplicitBorrow(cond.thenExpr) or createsExplicitBorrow(cond.elseExpr),
8437 +
        else => return false,
8438 +
    }
8439 +
}
8440 +
8441 +
/// Protect explicit address arguments until their call begins.
8442 +
unsafe fn retainCallLoan(
8443 +
    checker: &mut LinearChecker, node: *ast::Node, exclusive: bool
8444 +
) throws (ResolveError) {
8445 +
    match node.value {
8446 +
        case ast::NodeValue::CondExpr(cond) => {
8447 +
            try retainCallLoan(checker, cond.thenExpr, exclusive);
8448 +
            try retainCallLoan(checker, cond.elseExpr, exclusive);
8449 +
            return;
8450 +
        }
8451 +
        case ast::NodeValue::As(cast) => {
8452 +
            try retainCallLoan(checker, cast.value, exclusive);
8453 +
            return;
8454 +
        }
8455 +
        else => {}
8456 +
    }
8457 +
    let place = borrowPlace(checker.resolver, node);
8458 +
    if place.root == nil {
8459 +
        return;
8460 +
    }
8461 +
    if checker.localLen >= MAX_LINEAR_BINDINGS {
8462 +
        throw emitError(checker.resolver, node, ErrorKind::Internal);
8463 +
    }
8464 +
    set checker.locals[checker.localLen] = LocalLoan { binding: nil, place, exclusive };
8465 +
    set checker.localLen += 1;
8466 +
}
8467 +
8389 8468
/// Check a pattern conditional. Linear scrutinees require an exhaustive match.
8390 8469
unsafe fn checkLinearIfLet(
8391 8470
    checker: &mut LinearChecker,
8392 8471
    env: &mut LinearEnv,
8393 8472
    node: *ast::Node,
lib/std/lang/resolver/tests.rad +87 -0
6421 6421
        let mut res = testResolver();
6422 6422
        let result = try resolveProgramStr(&mut res, program);
6423 6423
        try expectNoErrors(&result);
6424 6424
    }
6425 6425
}
6426 +
6427 +
/// Earlier reference arguments protect their storage during later arguments.
6428 +
@test unsafe fn testCallArgumentLoans() throws (testing::TestError) {
6429 +
    for program in [
6430 +
        "fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &mut u32, n: u32) {} fn f (p: &mut u32) { outer(&mut *p, inner(p)); }",
6431 +
        "fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &u32, n: u32) {} fn f (p: &mut u32) { outer(&*p, inner(p)); }",
6432 +
        "fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &mut u32, n: u32) {} unsafe fn f (p: &mut u32) { outer(&mut *p, inner(p)); }",
6433 +
        "record R { n: u32 } fn (r: &mut R) call(n: u32) {} fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn f (r: &mut R) { (&mut *r).call(inner(&mut r.n)); }",
6434 +
    ] {
6435 +
        let mut res = testResolver();
6436 +
        let result = try resolveProgramStr(&mut res, program);
6437 +
        let error = try expectError(&result);
6438 +
        let case super::ErrorKind::BorrowConflict(_) = error.kind
6439 +
            else throw testing::TestError::Failed;
6440 +
    }
6441 +
}
6442 +
6443 +
/// Conditional explicit arguments protect every possible borrowed place.
6444 +
@test unsafe fn testConditionalCallArgumentLoans() throws (testing::TestError) {
6445 +
    for program in [
6446 +
        "fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &u32, n: u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(&*p if flag else &*q, inner(p)); }",
6447 +
        "fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &u32, n: u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(&*p if flag else &*q, inner(q)); }",
6448 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(&mut *p if flag else &mut *q, &mut *p); }",
6449 +
        "fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &u32, n: u32) {} unsafe fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(&*p if flag else &*q, inner(q)); }",
6450 +
        "fn outer(p: &mut u32, q: &mut u32) {} unsafe fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(&mut *p if flag else &mut *q, &mut *q); }",
6451 +
        "fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &u32, n: u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer((&*p if flag else &*q) as &u32, inner(p)); }",
6452 +
        "record R { n: u32 } fn (r: &R) call(n: u32) {} fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn f (p: &mut R, q: &mut R, flag: bool) { (&*p if flag else &*q).call(inner(&mut q.n)); }",
6453 +
    ] {
6454 +
        let mut res = testResolver();
6455 +
        let result = try resolveProgramStr(&mut res, program);
6456 +
        let error = try expectError(&result);
6457 +
        let case super::ErrorKind::BorrowConflict(_) = error.kind
6458 +
            else throw testing::TestError::Failed;
6459 +
    }
6460 +
}
6461 +
6462 +
/// Conditional argument alternatives cannot overlap another exclusive argument.
6463 +
@test unsafe fn testConditionalCallArgumentOverlap() throws (testing::TestError) {
6464 +
    for program in [
6465 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(p if flag else q, p); }",
6466 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(p if flag else q, q); }",
6467 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(p, p if flag else q); }",
6468 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(q, p if flag else q); }",
6469 +
        "fn outer(p: &mut u32, q: &mut u32) {} unsafe fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(p if flag else q, p); }",
6470 +
        "fn outer(p: &u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(p if flag else q, q); }",
6471 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(p if flag else q, q if flag else p); }",
6472 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, r: &mut u32, a: bool, b: bool) { outer((p if a else q) if b else r, q); }",
6473 +
        "fn outer(p: &u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer((p if flag else q) as &u32, p); }",
6474 +
        "record R { n: u32 } fn (r: &mut R) call(p: &R) {} fn f (p: &mut R, q: &mut R, flag: bool) { (p if flag else q).call(p); }",
6475 +
        "trait R { fn (&mut R) call(p: &opaque R); } fn f (p: &mut opaque R, q: &mut opaque R, flag: bool) { (p if flag else q).call(q); }",
6476 +
    ] {
6477 +
        let mut res = testResolver();
6478 +
        let result = try resolveProgramStr(&mut res, program);
6479 +
        let error = try expectError(&result);
6480 +
        let case super::ErrorKind::BorrowConflict(_) = error.kind
6481 +
            else throw testing::TestError::Failed;
6482 +
    }
6483 +
}
6484 +
6485 +
/// Conditional argument alternatives permit shared access and disjoint places.
6486 +
@test unsafe fn testConditionalCallArgumentSeparation() throws (testing::TestError) {
6487 +
    for program in [
6488 +
        "fn outer(p: &u32, q: &u32) {} fn f (p: &u32, q: &u32, flag: bool) { outer(p if flag else q, p); }",
6489 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f (p: &mut u32, q: &mut u32, r: &mut u32, flag: bool) { outer(p if flag else q, r); }",
6490 +
        "record R { a: u32, b: u32, c: u32 } fn outer(p: &mut u32, q: &mut u32) {} fn f (r: &mut R, flag: bool) { outer(&mut r.a if flag else &mut r.b, &mut r.c); }",
6491 +
        "record R { n: u32 } fn (r: &mut R) call(p: &R) {} fn f (p: &mut R, q: &mut R, r: &R, flag: bool) { (p if flag else q).call(r); }",
6492 +
    ] {
6493 +
        let mut res = testResolver();
6494 +
        let result = try resolveProgramStr(&mut res, program);
6495 +
        try expectNoErrors(&result);
6496 +
    }
6497 +
}
6498 +
6499 +
/// Call loans allow shared reads, separate fields, and access after the call.
6500 +
@test unsafe fn testCallArgumentLoanScopes() throws (testing::TestError) {
6501 +
    for program in [
6502 +
        "fn read(p: &u32) -> u32 { return *p; } fn outer(p: &u32, n: u32) {} fn f (p: &mut u32, q: &mut u32, flag: bool) { outer(&*p if flag else &*q, read(p)); set *p = 3; set *q = 4; }",
6503 +
        "record R { a: u32, b: u32, c: u32 } fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &mut u32, n: u32) {} fn f (r: &mut R, flag: bool) { outer(&mut r.a if flag else &mut r.b, inner(&mut r.c)); set r.a = 3; }",
6504 +
        "fn outer(p: &mut u32, n: u32) {} fn f (p: &mut u32, flag: bool) { outer(&mut *p if flag else &mut *p, 0); set *p = 3; }",
6505 +
        "fn read(p: &u32) -> u32 { return *p; } fn outer(p: &u32, n: u32) {} fn f (p: &mut u32) { outer(&*p, read(p)); set *p = 3; }",
6506 +
        "record R { a: u32, b: u32 } fn inner(p: &mut u32) -> u32 { set *p = 2; return 0; } fn outer(p: &mut u32, n: u32) {} fn f (r: &mut R) { outer(&mut r.a, inner(&mut r.b)); set r.a = 3; }",
6507 +
    ] {
6508 +
        let mut res = testResolver();
6509 +
        let result = try resolveProgramStr(&mut res, program);
6510 +
        try expectNoErrors(&result);
6511 +
    }
6512 +
}
lib/std/vec.rad +6 -4
65 65
/// Returns false if the vector is at capacity.
66 66
export unsafe fn push(vec: &mut RawVec, elem: &opaque) -> bool {
67 67
    if vec.len >= capacity(vec) {
68 68
        return false;
69 69
    }
70 -
    let off: u32 = vec.len * vec.stride;
71 -
    copyBytes(&mut vec.data[off..off + vec.stride], @sliceOf(elem as &u8, vec.stride));
70 +
    let stride = vec.stride;
71 +
    let off: u32 = vec.len * stride;
72 +
    copyBytes(&mut vec.data[off..off + stride], @sliceOf(elem as &u8, stride));
72 73
    set vec.len += 1;
73 74
74 75
    return true;
75 76
}
76 77
95 96
/// Returns false if index is out of bounds.
96 97
export unsafe fn put(vec: &mut RawVec, index: u32, elem: &opaque) -> bool {
97 98
    if index >= vec.len {
98 99
        return false;
99 100
    }
100 -
    let off: u32 = index * vec.stride;
101 -
    copyBytes(&mut vec.data[off..off + vec.stride], @sliceOf(elem as &u8, vec.stride));
101 +
    let stride = vec.stride;
102 +
    let off: u32 = index * stride;
103 +
    copyBytes(&mut vec.data[off..off + stride], @sliceOf(elem as &u8, stride));
102 104
103 105
    return true;
104 106
}
105 107
106 108
/// Copy bytes from source to destination.
test/tests/mutref.loop.bug.rad +8 -4
19 19
/// merge tries to `load` through it as if it were a pointer.
20 20
fn testZeroIter(n: u32) -> u32 {
21 21
    let mut val: u32 = 42;
22 22
    let mut i: u32 = 0;
23 23
    while i < n {
24 -
        store(&mut val, val + 1);
24 +
        let next = val + 1;
25 +
        store(&mut val, next);
25 26
        set i += 1;
26 27
    }
27 28
    return val;
28 29
}
29 30
30 31
/// Multiple iterations: accumulate via &mut pointer in a loop.
31 32
fn testMultiIter() -> u32 {
32 33
    let mut acc: u32 = 0;
33 34
    let mut i: u32 = 0;
34 35
    while i < 5 {
35 -
        store(&mut acc, acc + i);
36 +
        let next = acc + i;
37 +
        store(&mut acc, next);
36 38
        set i += 1;
37 39
    }
38 40
    return acc;
39 41
}
40 42
42 44
fn testMultipleVars() -> u32 {
43 45
    let mut a: u32 = 0;
44 46
    let mut b: u32 = 100;
45 47
    let mut i: u32 = 0;
46 48
    while i < 3 {
47 -
        store(&mut a, a + 1);
48 -
        store(&mut b, b - 1);
49 +
        let nextA = a + 1;
50 +
        store(&mut a, nextA);
51 +
        let nextB = b - 1;
52 +
        store(&mut b, nextB);
49 53
        set i += 1;
50 54
    }
51 55
    return a + b;
52 56
}
53 57