Require unsafe context for raw pointer patterns

7397ed177e331057ebfbdf420fc183ec3bc2f0b2e317cbdd6c544e5e971d7d53
Alexis Sellier committed ago 1 parent 682efa21
lib/std/lang/resolver.rad +9 -0
4460 4460
{
4461 4461
    match pat.kind {
4462 4462
        case ast::PatternKind::Case => {
4463 4463
            // Analyze pattern against scrutinee type.
4464 4464
            let scrutineeTy = try infer(self, pat.scrutinee);
4465 +
            if isUnsafePointerType(scrutineeTy) {
4466 +
                try requireUnsafe(self, pat.scrutinee);
4467 +
            }
4465 4468
            let subject = unwrapMatchSubject(scrutineeTy);
4466 4469
            try resolveCasePattern(self, pat.pattern, subject.effectiveTy, IdentMode::Compare, subject.by);
4467 4470
        }
4468 4471
        case ast::PatternKind::Binding => {
4469 4472
            // Scrutinee must be optional, bind the payload.
4524 4527
    scrutineeTy: Type,
4525 4528
    mode: IdentMode,
4526 4529
    matchBy: MatchBy
4527 4530
) throws (ResolveError) {
4528 4531
    if let case Type::Pointer { target, .. } = scrutineeTy; isDestructuringPattern(pattern) {
4532 +
        if isUnsafePointerType(scrutineeTy) {
4533 +
            try requireUnsafe(self, pattern);
4534 +
        }
4529 4535
        try resolveCasePattern(self, pattern, *target, mode, matchBy);
4530 4536
        return;
4531 4537
    }
4532 4538
    // TODO: Collapse these nested matches.
4533 4539
    match scrutineeTy {
4845 4851
/// the subject type.
4846 4852
unsafe fn resolveMatch(self: &mut Resolver, node: *ast::Node, sw: ast::Match) -> Type
4847 4853
    throws (ResolveError)
4848 4854
{
4849 4855
    let subjectTy = try infer(self, sw.subject);
4856 +
    if isUnsafePointerType(subjectTy) {
4857 +
        try requireUnsafe(self, sw.subject);
4858 +
    }
4850 4859
    let subject = unwrapMatchSubject(subjectTy);
4851 4860
4852 4861
    if let case Type::Optional(inner) = subject.effectiveTy {
4853 4862
        try resolveMatchOptional(self, node, sw, inner, subject.by);
4854 4863
    } else if let case Type::Nominal(NominalType::Union(u)) = subject.effectiveTy {
lib/std/lang/resolver/tests.rad +21 -0
5952 5952
@test unsafe fn testMutablePointerTargetsAllowed() throws (testing::TestError) {
5953 5953
    try expectAnalyzeOk("record R: Copy { n: u8 } fn (r: &mut R) change() { set r.n = 2; } fn run(p: *mut R) { set p.n = 1; p.change(); let q = &mut p.n; set *q = 3; }");
5954 5954
    try expectAnalyzeOk("static DATA: [u8; 2] = [0, 0]; fn run() { let p = &mut DATA; set p[0] = 1; set p[..] = 2; let q = &mut p[0]; set *q = 3; }");
5955 5955
    try expectAnalyzeOk("fn run(first: *u8, second: *u8) { let mut p = first; set p = second; }");
5956 5956
}
5957 +
5958 +
/// Pattern access through raw pointers requires an unsafe context.
5959 +
@test unsafe fn testRawPointerPatternRejected() throws (testing::TestError) {
5960 +
    let programs = &[
5961 +
        "union U: Copy { A(u8), B } fn run(p: *unsafe U) { match p { case U::A(n) => { *n; } else => {} } }",
5962 +
        "union U: Copy { A(u8), B } fn run(p: *unsafe U) { if let case U::A(n) = p { *n; } }",
5963 +
        "union U: Copy { A(u8), B } fn run(p: *unsafe U) { while let case U::A(n) = p { *n; break; } }",
5964 +
        "record R: Copy { n: u8 } fn run(p: *unsafe R) { let case R { n } = p else return; }",
5965 +
    ];
5966 +
    for program in programs {
5967 +
        let mut a = testResolver();
5968 +
        let result = try resolveProgramStr(&mut a, program);
5969 +
        try expectErrorKind(&result, super::ErrorKind::UnsafeOperation);
5970 +
    }
5971 +
}
5972 +
5973 +
/// Unsafe contexts permit pattern access through raw pointers.
5974 +
@test unsafe fn testRawPointerPatternsAllowed() throws (testing::TestError) {
5975 +
    try expectAnalyzeOk("union U: Copy { A(u8), B } unsafe fn run(p: *unsafe U) { match p { case U::A(n) => { *n; } else => {} } if let case U::A(n) = p { *n; } while let case U::A(n) = p { *n; break; } }");
5976 +
    try expectAnalyzeOk("record R: Copy { n: u8 } unsafe fn run(p: *unsafe R) { let case R { n } = p else return; }");
5977 +
}