compiler: Preserve shared regions in pattern bindings

05ad7df3fff5a8d96795435e1f88f22a7486b2547ed05f1ddcadc6168b41f2be
Alexis Sellier committed ago 1 parent c8d0ba89
lib/std/lang/lower.rad +2 -2
2647 2647
            // Non-void unions are passed by reference (need to load tag).
2648 2648
            // When matching by reference, always load from the pointer.
2649 2649
            if unionInfo.isAllVoid {
2650 2650
                let mut tagVal = subject.val;
2651 2651
                match subject.by {
2652 -
                    case resolver::MatchBy::Ref, resolver::MatchBy::MutRef => {
2652 +
                    case resolver::MatchBy::Ref(_), resolver::MatchBy::MutRef => {
2653 2653
                        let base = emitValToReg(self, subject.val);
2654 2654
                        set tagVal = loadTag(self, base, 0, il::Type::W8);
2655 2655
                    }
2656 2656
                    case resolver::MatchBy::Value => {}
2657 2657
                }
3482 3482
    valOffset: i32
3483 3483
) -> il::Val where 'arena: 'phase, 'phase: 'function {
3484 3484
    match matchBy {
3485 3485
        case resolver::MatchBy::Value =>
3486 3486
            return tvalPayloadVal(self, base, bindType, valOffset),
3487 -
        case resolver::MatchBy::Ref, resolver::MatchBy::MutRef =>
3487 +
        case resolver::MatchBy::Ref(_), resolver::MatchBy::MutRef =>
3488 3488
            return il::Val::Reg(emitPtrOffset(self, base, valOffset)),
3489 3489
    }
3490 3490
}
3491 3491
3492 3492
/// Bind a variable to a tagged value's payload.
lib/std/lang/resolver.rad +9 -5
885 885
/// How pattern bindings are created during match.
886 886
export union MatchBy: Copy {
887 887
    /// Match by value.
888 888
    Value,
889 889
    /// Match by immutable reference.
890 -
    Ref,
890 +
    Ref(types::PointerClass),
891 891
    /// Match by mutable reference.
892 892
    MutRef,
893 893
}
894 894
895 895
/// State of a match statement being resolved.
1030 1030
    loopDepth: u32,
1031 1031
}
1032 1032
1033 1033
/// Unwrap a pointer type for pattern matching.
1034 1034
export fn unwrapMatchSubject(ty: Type) -> MatchSubject {
1035 -
    if let case Type::Pointer { target, mutable, .. } = ty {
1036 -
        let by = MatchBy::MutRef if mutable else MatchBy::Ref;
1035 +
    if let case Type::Pointer { class, target, mutable } = ty {
1036 +
        let mut bindingClass = types::PointerClass::Ref;
1037 +
        if let case types::PointerClass::Region(_) = class {
1038 +
            set bindingClass = class;
1039 +
        }
1040 +
        let by = MatchBy::MutRef if mutable else MatchBy::Ref(bindingClass);
1037 1041
        return MatchSubject { effectiveTy: *target, by };
1038 1042
    }
1039 1043
    return MatchSubject { effectiveTy: ty, by: MatchBy::Value };
1040 1044
}
1041 1045
6681 6685
    throws (ResolveError)
6682 6686
{
6683 6687
    let mut bindTy = ty;
6684 6688
    match matchBy {
6685 6689
        case MatchBy::Value => {}
6686 -
        case MatchBy::Ref => set bindTy = Type::Pointer {
6687 -
            class: types::PointerClass::Ref,
6690 +
        case MatchBy::Ref(class) => set bindTy = Type::Pointer {
6691 +
            class,
6688 6692
            target: allocType(self, ty),
6689 6693
            mutable: false,
6690 6694
        },
6691 6695
        case MatchBy::MutRef => set bindTy = Type::Pointer {
6692 6696
            class: types::PointerClass::Ref,
lib/std/lang/resolver/tests/regions.rad +100 -0
1788 1788
            try super::expectErrorKind(&result, resolver::ErrorKind::ImmutableBinding);
1789 1789
        }
1790 1790
    }
1791 1791
}
1792 1792
1793 +
/// Named shared pattern bindings retain the subject reference region.
1794 +
@test unsafe fn testNamedPatternBorrowRegions() throws (testing::TestError) {
1795 +
    for program in [
1796 +
        "record B: 'r { values: &'r mut [u32] } fn f 'r (b: &'r B 'r) -> &'r [u32] { return &b.values[..]; }",
1797 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn f 'r (u: &'r U 'r) -> &'r [u32] { match u { case U::Some(b) => return &b.values[..], else => panic, } }",
1798 +
        "record B: 'r { values: &'r mut [u32] } record S: 'r { items: [?B 'r; 1] } fn f 'r (s: &'r S 'r) -> &'r [u32] { match &s.items[0] { case nil => panic, b => return &b.values[..], } }",
1799 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn values 'r (b: &'r B 'r) -> &'r [u32] { return &b.values[..]; } fn f 'r (u: &'r U 'r) -> &'r [u32] { match u { case U::Some(b) => return values(b), else => panic, } }",
1800 +
        "record B: 'r { values: &'r mut [u32] } record S: 'r { items: [?B 'r; 1] } fn values 'r (b: &'r B 'r) -> &'r [u32] { return &b.values[..]; } fn f 'r (s: &'r S 'r) -> &'r [u32] { let mut found: ?&'r B 'r = nil; match &s.items[0] { case nil => panic, b => set found = b, } let b = found else panic; return values(b); }",
1801 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn f 'r (u: &'r U 'r) -> &'r [u32] { if let case U::Some(b) = u { return &b.values[..]; } panic; }",
1802 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn f 'r (u: &'r U 'r) -> &'r [u32] { let case U::Some(b) = u else panic; return &b.values[..]; }",
1803 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn f 'r (u: &'r U 'r) -> &'r [u32] { while let case U::Some(b) = u { return &b.values[..]; } panic; }",
1804 +
    ] {
1805 +
        let mut arena = super::testArena();
1806 +
        let storage: 'test = &mut arena in {
1807 +
            let mut res = super::testResolver(storage);
1808 +
            let result = try super::resolveProgramStr(&mut res, program);
1809 +
            try super::expectNoErrors(&result);
1810 +
        }
1811 +
    }
1812 +
}
1813 +
1814 +
/// Owned and raw pattern subjects produce ordinary shared bindings.
1815 +
@test unsafe fn testPointerPatternBindingClasses() throws (testing::TestError) {
1816 +
    for program in [
1817 +
        "record B { value: u32 } union U { Some(B), None } fn f(u: *U) { match u { case U::Some(b) => { assert b.value == 0; }, else => {}, } }",
1818 +
        "record B { value: u32 } union U { Some(B), None } unsafe fn f(u: *unsafe U) { match u { case U::Some(b) => { assert b.value == 0; }, else => {}, } }",
1819 +
    ] {
1820 +
        let mut arena = super::testArena();
1821 +
        let storage: 'test = &mut arena in {
1822 +
            let mut res = super::testResolver(storage);
1823 +
            let result = try super::resolveProgramStr(&mut res, program);
1824 +
            try super::expectNoErrors(&result);
1825 +
            let func = try super::getBlockStmt(result.root, 2);
1826 +
            let fnType = resolver::typeFor(&res, func) else throw testing::TestError::Failed;
1827 +
            let case resolver::Type::Fn(info) = fnType else throw testing::TestError::Failed;
1828 +
            let subject = resolver::unwrapMatchSubject(*info.paramTypes[0]);
1829 +
            let case resolver::MatchBy::Ref(types::PointerClass::Ref) = subject.by
1830 +
                else throw testing::TestError::Failed;
1831 +
            let case ast::NodeValue::FnDecl(decl) = func.value else throw testing::TestError::Failed;
1832 +
            let body = decl.body else throw testing::TestError::Failed;
1833 +
            let statement = try super::getBlockStmt(body, 0);
1834 +
            let case ast::NodeValue::Match(matchStmt) = statement.value else throw testing::TestError::Failed;
1835 +
            let case ast::NodeValue::MatchProng(prong) = matchStmt.prongs[0].value
1836 +
                else throw testing::TestError::Failed;
1837 +
            let case ast::ProngArm::Case(patterns) = prong.arm else throw testing::TestError::Failed;
1838 +
            let case ast::NodeValue::Call(call) = patterns[0].value else throw testing::TestError::Failed;
1839 +
            let bindingType = resolver::typeFor(&res, call.args[0]) else throw testing::TestError::Failed;
1840 +
            let case resolver::Type::Pointer { class: types::PointerClass::Ref, .. } = bindingType
1841 +
                else throw testing::TestError::Failed;
1842 +
        }
1843 +
    }
1844 +
}
1845 +
1846 +
/// Mutable pattern payload bindings cannot escape or overlap source mutation.
1847 +
@test unsafe fn testMutablePatternBindingEscape() throws (testing::TestError) {
1848 +
    let programs = [
1849 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn f 'r (u: &'r mut U 'r) -> &'r [u32] { match u { case U::Some(b) => return &b.values[..], else => panic, } }",
1850 +
    ];
1851 +
    for program in programs {
1852 +
        let mut arena = super::testArena();
1853 +
        let storage: 'test = &mut arena in {
1854 +
            let mut res = super::testResolver(storage);
1855 +
            let result = try super::resolveProgramStr(&mut res, program);
1856 +
            let error = try super::expectError(&result);
1857 +
            let case resolver::ErrorKind::TypeMismatch(_) = error.kind
1858 +
                else throw testing::TestError::Failed;
1859 +
        }
1860 +
    }
1861 +
    for program in [
1862 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn f 'r (u: &'r mut U 'r) { match u { case U::Some(b) => { let held = &b.values[..]; set *u = U::None; assert held.len == 0; }, else => {}, } }",
1863 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn f 'r (u: &'r mut U 'r) { let mut held: ?&'r B 'r = nil; match &*u { case U::Some(b) => set held = b, else => {}, } set *u = U::None; let b = held else return; assert b.values.len == 0; }",
1864 +
    ] {
1865 +
        let mut arena = super::testArena();
1866 +
        let storage: 'test = &mut arena in {
1867 +
            let mut res = super::testResolver(storage);
1868 +
            let result = try super::resolveProgramStr(&mut res, program);
1869 +
            let error = try super::expectError(&result);
1870 +
            let case resolver::ErrorKind::BorrowConflict(_) = error.kind
1871 +
                else throw testing::TestError::Failed;
1872 +
        }
1873 +
    }
1874 +
}
1875 +
1876 +
/// Named shared pattern bindings cannot outlive a shorter subject reference.
1877 +
@test unsafe fn testNamedPatternBorrowRejectsShortOwner() throws (testing::TestError) {
1878 +
    for program in [
1879 +
        "record B: 'r { values: &'r mut [u32] } union U: 'r { Some(B 'r), None } fn f 'r 's (u: &'s U 'r) -> &'r [u32] where 'r: 's { match u { case U::Some(b) => return &b.values[..], else => panic, } }",
1880 +
        "record B: 'r { values: &'r mut [u32] } record S: 'r { items: [?B 'r; 1] } fn f 'r 's (s: &'s S 'r) -> &'r [u32] where 'r: 's { match &s.items[0] { case nil => panic, b => return &b.values[..], } }",
1881 +
    ] {
1882 +
        let mut arena = super::testArena();
1883 +
        let storage: 'test = &mut arena in {
1884 +
            let mut res = super::testResolver(storage);
1885 +
            let result = try super::resolveProgramStr(&mut res, program);
1886 +
            let error = try super::expectError(&result);
1887 +
            let case resolver::ErrorKind::TypeMismatch(_) = error.kind
1888 +
                else throw testing::TestError::Failed;
1889 +
        }
1890 +
    }
1891 +
}
1892 +
1793 1893
/// A reference through an exclusive field cannot outlive the owner's borrow.
1794 1894
@test unsafe fn testExclusiveOwnerBorrowLifetimes() throws (testing::TestError) {
1795 1895
    for program in [
1796 1896
        "record H: 'r { item: &'r mut u32 } fn f 'r 's (h: &'s H 'r) -> &'r u32 where 'r: 's { let p: &'r u32 = h.item; return p; }",
1797 1897
        "record H: 'r { item: ?&'r mut u32 } fn f 'r 's (h: &'s H 'r) -> ?&'r u32 where 'r: 's { let p: ?&'r u32 = h.item; return p; }",