compiler: Check conditional call arguments for aliasing

ed2fbb6cf6251139613e98575cb41dd2e5616e0a0dd73d5562e87d909fc9b7f5
Alexis Sellier committed ago 1 parent f848c85f
lib/std/lang/resolver.rad +48 -17
10673 10673
        for arg in call.args {
10674 10674
            try checkLinearNode(checker, env, arg, LinearUse::Consume);
10675 10675
        }
10676 10676
        return;
10677 10677
    };
10678 -
    let mut places: [BorrowPlace; MAX_FN_PARAMS + 1] = undefined;
10678 +
    let mut arguments: [*ast::Node; MAX_FN_PARAMS + 1] = undefined;
10679 10679
    let mut exclusive: [bool; MAX_FN_PARAMS + 1] = undefined;
10680 -
    let mut placesLen: u32 = 0;
10680 +
    let mut argumentsLen: u32 = 0;
10681 10681
10682 10682
    // Method function types exclude their implicit receiver. Account for it
10683 10683
    // explicitly so owning receivers are consumed and reference receivers
10684 10684
    // participate in call-scoped loan conflict checks.
10685 10685
    if let case ast::NodeValue::FieldAccess(access) = call.callee.value {
10705 10705
                receiverMutable or receiverClass == types::PointerClass::Owned);
10706 10706
            if receiverMutable or receiverClass == types::PointerClass::Owned {
10707 10707
                try checkPatternLoan(checker, access.parent);
10708 10708
            }
10709 10709
            if receiverClass <> types::PointerClass::Unsafe {
10710 -
                let place = borrowPlace(checker.resolver, access.parent);
10711 -
                if place.root <> nil {
10712 -
                    set places[placesLen] = place;
10713 -
                    set exclusive[placesLen] =
10714 -
                        receiverClass == types::PointerClass::Owned or receiverMutable;
10715 -
                    set placesLen += 1;
10716 -
                }
10710 +
                set arguments[argumentsLen] = access.parent;
10711 +
                set exclusive[argumentsLen] =
10712 +
                    receiverClass == types::PointerClass::Owned or receiverMutable;
10713 +
                set argumentsLen += 1;
10717 10714
            }
10718 10715
            if types::isReference(receiverClass) {
10719 10716
                try checkLinearNode(checker, env, access.parent, LinearUse::Borrow);
10720 10717
                if createsExplicitBorrow(access.parent) {
10721 10718
                    try retainCallLoan(checker, access.parent, receiverMutable);
10730 10727
        let expected = *info.paramTypes[i];
10731 10728
        let argExclusive = isExclusiveArgument(expected) or createsCellBorrow(arg);
10732 10729
        if argExclusive {
10733 10730
            try checkPatternLoan(checker, arg);
10734 10731
        }
10735 -
        let place = borrowPlace(checker.resolver, arg);
10736 10732
        if not isUnsafePointerType(expected) {
10737 -
            if let rootSym = place.root {
10738 -
                for j in 0..placesLen {
10739 -
                    if (exclusive[j] or argExclusive) and placesOverlap(&places[j], &place) {
10740 -
                        throw emitError(checker.resolver, arg, ErrorKind::BorrowConflict(rootSym.name));
10733 +
            for j in 0..argumentsLen {
10734 +
                if exclusive[j] or argExclusive {
10735 +
                    if let name = callArgumentConflict(checker.resolver, arguments[j], arg) {
10736 +
                        throw emitError(checker.resolver, arg, ErrorKind::BorrowConflict(name));
10741 10737
                    }
10742 10738
                }
10743 -
                set places[placesLen] = place;
10744 -
                set exclusive[placesLen] = argExclusive;
10745 -
                set placesLen += 1;
10746 10739
            }
10740 +
            set arguments[argumentsLen] = arg;
10741 +
            set exclusive[argumentsLen] = argExclusive;
10742 +
            set argumentsLen += 1;
10747 10743
        }
10748 10744
        try checkLocalLoans(checker, env, arg, argExclusive);
10749 10745
        if isBorrowedReferenceParameter(expected) {
10750 10746
            try checkLinearNode(checker, env, arg, LinearUse::Borrow);
10751 10747
        } else {
10759 10755
    if *info.returnType == Type::Never and info.throwList.len == 0 {
10760 10756
        set env.terminated = true;
10761 10757
    }
10762 10758
}
10763 10759
10760 +
/// Return the storage name when two call arguments can address the same place.
10761 +
unsafe fn callArgumentConflict 'arena (
10762 +
    self: &mut Resolver 'arena, left: *ast::Node, right: *ast::Node
10763 +
) -> ?*[u8] {
10764 +
    match left.value {
10765 +
        case ast::NodeValue::CondExpr(cond) => {
10766 +
            if let name = callArgumentConflict(self, cond.thenExpr, right) {
10767 +
                return name;
10768 +
            }
10769 +
            return callArgumentConflict(self, cond.elseExpr, right);
10770 +
        }
10771 +
        case ast::NodeValue::As(cast) => return callArgumentConflict(self, cast.value, right),
10772 +
        case ast::NodeValue::RegionApply { value, .. } => return callArgumentConflict(self, value, right),
10773 +
        else => {}
10774 +
    }
10775 +
    match right.value {
10776 +
        case ast::NodeValue::CondExpr(cond) => {
10777 +
            if let name = callArgumentConflict(self, left, cond.thenExpr) {
10778 +
                return name;
10779 +
            }
10780 +
            return callArgumentConflict(self, left, cond.elseExpr);
10781 +
        }
10782 +
        case ast::NodeValue::As(cast) => return callArgumentConflict(self, left, cast.value),
10783 +
        case ast::NodeValue::RegionApply { value, .. } => return callArgumentConflict(self, left, value),
10784 +
        else => {}
10785 +
    }
10786 +
    let leftPlace = borrowPlace(self, left);
10787 +
    let rightPlace = borrowPlace(self, right);
10788 +
    if placesOverlap(&leftPlace, &rightPlace) {
10789 +
        let root = rightPlace.root else panic "callArgumentConflict: overlap without root";
10790 +
        return root.name;
10791 +
    }
10792 +
    return nil;
10793 +
}
10794 +
10764 10795
/// Return whether evaluating an argument creates an explicit address borrow.
10765 10796
fn createsExplicitBorrow(node: *ast::Node) -> bool {
10766 10797
    match node.value {
10767 10798
        case ast::NodeValue::AddressOf(_) => return true,
10768 10799
        case ast::NodeValue::As(cast) => return createsExplicitBorrow(cast.value),
lib/std/lang/resolver/tests/regions.rad +43 -0
78 78
                else throw testing::TestError::Failed;
79 79
        }
80 80
    }
81 81
}
82 82
83 +
/// Conditional argument alternatives cannot overlap another exclusive argument.
84 +
@test unsafe fn testConditionalCallArgumentOverlap() throws (testing::TestError) {
85 +
    for program in [
86 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer(p if flag else q, p); }",
87 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer(p if flag else q, q); }",
88 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer(p, p if flag else q); }",
89 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer(q, p if flag else q); }",
90 +
        "fn outer(p: &mut u32, q: &mut u32) {} unsafe fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer(p if flag else q, p); }",
91 +
        "fn outer(p: &u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer(p if flag else q, q); }",
92 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer(p if flag else q, q if flag else p); }",
93 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, r: &'r mut u32, a: bool, b: bool) { outer((p if a else q) if b else r, q); }",
94 +
        "fn outer(p: &u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer((p if flag else q) as &'r u32, p); }",
95 +
        "record R { n: u32 } fn (r: &mut R) call(p: &R) {} fn f 'r (p: &'r mut R, q: &'r mut R, flag: bool) { (p if flag else q).call(p); }",
96 +
        "trait R { fn (&mut R) call(p: &opaque R); } fn f 'r (p: &'r mut opaque R, q: &'r mut opaque R, flag: bool) { (p if flag else q).call(q); }",
97 +
    ] {
98 +
        let mut arena = super::testArena();
99 +
        let storage: 'test = &mut arena in {
100 +
            let mut res = super::testResolver(storage);
101 +
            let result = try super::resolveProgramStr(&mut res, program);
102 +
            let error = try super::expectError(&result);
103 +
            let case resolver::ErrorKind::BorrowConflict(_) = error.kind
104 +
                else throw testing::TestError::Failed;
105 +
        }
106 +
    }
107 +
}
108 +
109 +
/// Conditional argument alternatives permit shared access and disjoint places.
110 +
@test unsafe fn testConditionalCallArgumentSeparation() throws (testing::TestError) {
111 +
    for program in [
112 +
        "fn outer(p: &u32, q: &u32) {} fn f 'r (p: &'r u32, q: &'r u32, flag: bool) { outer(p if flag else q, p); }",
113 +
        "fn outer(p: &mut u32, q: &mut u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, r: &'r mut u32, flag: bool) { outer(p if flag else q, r); }",
114 +
        "record R: 'r { a: &'r mut u32, b: &'r mut u32, c: &'r mut u32 } fn outer(p: &mut u32, q: &mut u32) {} fn f 'r (r: &mut R 'r, flag: bool) { outer(r.a if flag else r.b, r.c); }",
115 +
        "record R { n: u32 } fn (r: &mut R) call(p: &R) {} fn f 'r (p: &'r mut R, q: &'r mut R, r: &'r R, flag: bool) { (p if flag else q).call(r); }",
116 +
    ] {
117 +
        let mut arena = super::testArena();
118 +
        let storage: 'test = &mut arena in {
119 +
            let mut res = super::testResolver(storage);
120 +
            let result = try super::resolveProgramStr(&mut res, program);
121 +
            try super::expectNoErrors(&result);
122 +
        }
123 +
    }
124 +
}
125 +
83 126
/// Call loans allow shared reads, separate fields, and access after the call.
84 127
@test unsafe fn testRegionalCallArgumentLoanScopes() throws (testing::TestError) {
85 128
    for program in [
86 129
        "fn read(p: &u32) -> u32 { return *p; } fn outer(p: &u32, n: u32) {} fn f 'r (p: &'r mut u32, q: &'r mut u32, flag: bool) { outer(&*p if flag else &*q, read(p)); set *p = 3; set *q = 4; }",
87 130
        "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 (r: &'r mut R, flag: bool) { outer(&mut r.a if flag else &mut r.b, inner(&mut r.c)); set r.a = 3; }",