Require unsafe functions to mutate slice headers

a1fd8da69aa0797ba91f80625a9636996b6e5e8c40b4b201c78fbed8f2e7b779
Alexis Sellier committed ago 1 parent d3cc6c5f
lib/std/lang/resolver.rad +4 -1
6238 6238
                return mutable;
6239 6239
            }
6240 6240
            return false;
6241 6241
        }
6242 6242
        case ast::NodeValue::FieldAccess(access) => {
6243 -
            let _ = try infer(self, access.parent);
6243 +
            let parentTy = try infer(self, access.parent);
6244 +
            if let case Type::Slice { .. } = autoDeref(parentTy) {
6245 +
                try requireUnsafe(self, node);
6246 +
            }
6244 6247
            return try canBorrowMutFrom(self, access.parent);
6245 6248
        }
6246 6249
        case ast::NodeValue::ScopeAccess(_) => {
6247 6250
            // Module-qualified access to a top-level symbol. A `static`
6248 6251
            // binds as a mutable value; a `constant` does not.
lib/std/lang/resolver/tests.rad +25 -0
5528 5528
    try expectAnalyzeOk("unsafe fn run(p: *mut u8) -> *mut [u8] { return @sliceOf(p, 0, 100); }");
5529 5529
    try expectAnalyzeOk("unsafe fn run(p: &u8) -> u8 { return @sliceOf(p, 1)[0]; }");
5530 5530
    try expectAnalyzeOk("fn run(s: *[u8]) -> *[u8] { return &s[..]; }");
5531 5531
}
5532 5532
5533 +
/// Slice header writes and mutable field borrows require an unsafe function.
5534 +
@test unsafe fn testSliceHeaderMutationRequiresUnsafe() throws (testing::TestError) {
5535 +
    let programs = &[
5536 +
        "fn run(s: *mut [u8]) { set s.len = 100; }",
5537 +
        "fn run(s: *mut [u8]) { set s.cap = 100; }",
5538 +
        "fn run(s: *mut [u8], p: *mut u8) { set s.ptr = p; }",
5539 +
        "fn run(s: *mut [u8]) { set s.len += 1; }",
5540 +
        "fn change(n: &mut u32) { set *n = 100; } fn run(s: *mut [u8]) { change(&mut s.len); }",
5541 +
        "fn change(n: &mut u32) { set *n = 100; } fn run(s: *mut [u8]) { change(&mut s.cap); }",
5542 +
        "fn change(p: &mut *u8, q: *u8) { set *p = q; } fn run(q: *u8) { let mut s: *[u8] = &[]; change(&mut s.ptr, q); }",
5543 +
        "fn run(s: &mut *[u8]) { set s.len = 100; }",
5544 +
        "static DATA: [u8; 1] = [42]; fn run() -> u8 { let mut s = &DATA[..]; set s.len = 100; return s[99]; }",
5545 +
    ];
5546 +
    for program in programs {
5547 +
        let mut a = testResolver();
5548 +
        let result = try resolveProgramStr(&mut a, program);
5549 +
        try expectErrorKind(&result, super::ErrorKind::UnsafeOperation);
5550 +
    }
5551 +
    try expectAnalyzeOk("unsafe fn run(s: *mut [u8]) { set s.len = 0; set s.cap = 0; }");
5552 +
    try expectAnalyzeOk("unsafe fn run(s: *mut [u8], p: *mut u8) { set s.ptr = p; }");
5553 +
    try expectAnalyzeOk("fn change(n: &mut u32) { set *n = 0; } unsafe fn run(s: *mut [u8]) { change(&mut s.len); }");
5554 +
    try expectAnalyzeOk("fn run(s: *mut [u8]) { set s[0] = 1; }");
5555 +
    try expectAnalyzeOk("record R { len: u32 } fn run(r: &mut R) { set r.len = 1; }");
5556 +
}
5557 +
5533 5558
/// Unsafe declarations may compose unsafe operations and calls.
5534 5559
@test unsafe fn testUnsafePointerOperationsAllowed() throws (testing::TestError) {
5535 5560
    let program = "record Marker: Once {} unsafe fn load(pointer: *unsafe u32) -> u32 { return *pointer; } unsafe fn run(pointer: *unsafe u32) -> u32 { let next = pointer + 1; let same = pointer == next; return load(pointer); }";
5536 5561
    try expectAnalyzeOk(program);
5537 5562
}
test/tests/slice.header.unsafe.rad added +23 -0
1 +
//! returns: 0
2 +
3 +
/// Backing storage for valid slice header changes.
4 +
static DATA: [u8; 3] = [11, 22, 33];
5 +
6 +
/// Set a borrowed header field to a valid length.
7 +
fn shorten(length: &mut u32) {
8 +
    set *length = 1;
9 +
}
10 +
11 +
/// Change slice metadata within its backing storage in an unsafe function.
12 +
@default unsafe fn main() -> i32 {
13 +
    let mut s = &mut DATA[..];
14 +
    set s.len = 2;
15 +
    set s.cap = 2;
16 +
    set s.ptr = &mut DATA[1];
17 +
    if s[0] <> 22 or s[1] <> 33 { return 1; }
18 +
    shorten(&mut s.len);
19 +
    if s.len <> 1 { return 2; }
20 +
    set s.len += 1;
21 +
    if s[1] <> 33 { return 3; }
22 +
    return 0;
23 +
}