Require unsafe context for slice pointer access

7674c012af3fdc1ab5efbd04cc3a7666a6164f9f8814284db2b48ced057aca0c
Alexis Sellier committed ago 1 parent d9d9db3e
lib/std/io.rad +9 -3
2 2
use std::fmt;
3 3
use std::intrinsics;
4 4
5 5
/// Write the bytes to standard output.
6 6
export fn print(str: &[u8]) {
7 -
    intrinsics::ecall(64, 1, str.ptr as i64, str.len as i64, 0);
7 +
    unsafe {
8 +
        intrinsics::ecall(64, 1, str.ptr as i64, str.len as i64, 0);
9 +
    }
8 10
}
9 11
10 12
/// Write the bytes to standard error.
11 13
export fn printError(str: &[u8]) {
12 -
    intrinsics::ecall(64, 2, str.ptr as i64, str.len as i64, 0);
14 +
    unsafe {
15 +
        intrinsics::ecall(64, 2, str.ptr as i64, str.len as i64, 0);
16 +
    }
13 17
}
14 18
15 19
/// Write the bytes and a newline to standard output.
16 20
export fn printLn(str: &[u8]) {
17 21
    print(str);
39 43
    print(&buffer[start..]);
40 44
}
41 45
42 46
/// Read standard input into the buffer and return the system call result.
43 47
export fn read(buf: &mut [u8]) -> u32 {
44 -
    return intrinsics::ecall(63, 0, buf.ptr as i64, buf.len as i64, 0) as u32;
48 +
    unsafe {
49 +
        return intrinsics::ecall(63, 0, buf.ptr as i64, buf.len as i64, 0) as u32;
50 +
    }
45 51
}
46 52
47 53
export fn readToEnd(buf: *mut [u8]) -> *[u8] {
48 54
    let mut total: u32 = 0;
49 55
lib/std/lang/resolver.rad +1 -0
6152 6152
6153 6153
    if let case Type::Slice { class, item, mutable } = subjectTy {
6154 6154
        let fieldNode = access.child;
6155 6155
        let fieldName = try nodeName(self, fieldNode);
6156 6156
        if mem::eq(fieldName, PTR_FIELD) {
6157 +
            try requireUnsafe(self, node);
6157 6158
            setRecordFieldIndex(self, fieldNode, 0);
6158 6159
            return setNodeType(
6159 6160
                self,
6160 6161
                node,
6161 6162
                Type::Pointer { class, target: item, mutable },
lib/std/lang/resolver/tests.rad +28 -3
866 866
    try expectExprStmtType(&a, indexStmt, super::Type::I32);
867 867
}
868 868
869 869
@test unsafe fn testResolveSliceFields() throws (testing::TestError) {
870 870
    let mut a = testResolver();
871 -
    let program = "static xs: [i32; 3] = [1, 2, 3]; let slice: *[i32] = &xs[1..]; slice.len; slice.ptr;";
871 +
    let program = "static xs: [i32; 3] = [1, 2, 3]; let slice: *[i32] = &xs[1..]; slice.len; unsafe { slice.ptr; }";
872 872
    let result = try resolveProgramStr(&mut a, program);
873 873
    try expectNoErrors(&result);
874 874
875 875
    let lenStmt = try getBlockStmt(result.root, 2);
876 876
    let case ast::NodeValue::ExprStmt(lenExpr) = lenStmt.value
877 877
        else throw testing::TestError::Failed;
878 878
    let lenTy = try typeOf(&a, lenExpr);
879 879
    try testing::expect(lenTy == super::Type::U32);
880 880
881 881
    let ptrStmt = try getBlockStmt(result.root, 3);
882 -
    let case ast::NodeValue::ExprStmt(ptrExpr) = ptrStmt.value
882 +
    let case ast::NodeValue::Block(ptrBlock) = ptrStmt.value
883 +
        else throw testing::TestError::Failed;
884 +
    let case ast::NodeValue::ExprStmt(ptrExpr) = ptrBlock.statements[0].value
883 885
        else throw testing::TestError::Failed;
884 886
    let ptrTy = try typeOf(&a, ptrExpr);
885 887
    let targetTy = try expectPointerType(ptrTy, false);
886 888
    try testing::expect(targetTy == super::Type::I32);
887 889
}
5357 5359
    try expectAnalyzeOk("fn run() -> *[u32] { return &[1, 2]; }");
5358 5360
    try expectAnalyzeOk("fn run(value: *u32) -> *u32 { return &*value; }");
5359 5361
    try expectAnalyzeOk("fn run(values: *[u32]) -> *[u32] { return &values[..]; }");
5360 5362
    try expectAnalyzeOk("record Cell { value: u32 } fn run(value: *Cell) -> *u32 { return &value.value; }");
5361 5363
    try expectAnalyzeOk("fn run(values: *[u32]) -> *u32 { return &values[0]; }");
5362 -
    try expectAnalyzeOk("fn run(values: *[u32]) -> *u32 { return values.ptr; }");
5364 +
    try expectAnalyzeOk("unsafe fn run(values: *[u32]) -> *u32 { return values.ptr; }");
5363 5365
    try expectAnalyzeOk("unsafe fn run(value: *u32) -> *[u32] { return @sliceOf(value, 1); }");
5364 5366
}
5365 5367
5366 5368
/// References are rejected from every nested or storable type position.
5367 5369
@test unsafe fn testNestedRefPositionsRejected() throws (testing::TestError) {
5599 5601
    let mut a = testResolver();
5600 5602
    let result = try resolveProgramStr(&mut a, "fn run() { let n: u32 = 1; unsafe { let p = &n; } }");
5601 5603
    try expectErrorKind(&result, super::ErrorKind::RefBinding);
5602 5604
}
5603 5605
5606 +
/// Slice pointer access requires an unsafe context for every slice class.
5607 +
@test unsafe fn testSlicePointerRequiresUnsafe() throws (testing::TestError) {
5608 +
    let programs = &[
5609 +
        "fn run(s: *[u8]) -> *u8 { return s.ptr; }",
5610 +
        "fn run(s: *mut [u8]) -> *mut u8 { return s.ptr; }",
5611 +
        "fn run(s: &[u8]) -> u8 { return *s.ptr; }",
5612 +
        "fn run(s: &mut [u8]) { set *s.ptr = 1; }",
5613 +
        "fn run(s: &*[u8]) -> *u8 { return s.ptr; }",
5614 +
        "static DATA: [u8; 1] = [42]; fn run() -> u8 { let s = &DATA[1..]; return *s.ptr; }",
5615 +
        "fn run(s: *[u8]) -> u64 { return s.ptr as u64; }",
5616 +
    ];
5617 +
    for program in programs {
5618 +
        let mut a = testResolver();
5619 +
        let result = try resolveProgramStr(&mut a, program);
5620 +
        try expectErrorKind(&result, super::ErrorKind::UnsafeOperation);
5621 +
    }
5622 +
    try expectAnalyzeOk("unsafe fn run(s: *[u8]) -> *u8 { return s.ptr; }");
5623 +
    try expectAnalyzeOk("fn run(s: *[u8]) -> u64 { unsafe { return s.ptr as u64; } }");
5624 +
    try expectAnalyzeOk("fn run(s: *[u8]) -> *u8 { return &s[0]; }");
5625 +
    try expectAnalyzeOk("fn run(s: *[u8]) -> u32 { return s.len + s.cap; }");
5626 +
    try expectAnalyzeOk("record R { ptr: u32 } fn run(r: R) -> u32 { return r.ptr; }");
5627 +
}
5628 +
5604 5629
/// Unsafe declarations may compose unsafe operations and calls.
5605 5630
@test unsafe fn testUnsafePointerOperationsAllowed() throws (testing::TestError) {
5606 5631
    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); }";
5607 5632
    try expectAnalyzeOk(program);
5608 5633
}
lib/std/mem.rad +4 -2
56 56
/// Check whether two byte slices have the same length and contents.
57 57
export fn eq(a: &[u8], b: &[u8]) -> bool {
58 58
    if a.len <> b.len {
59 59
        return false;
60 60
    }
61 -
    if a.ptr == b.ptr {
62 -
        return true;
61 +
    unsafe {
62 +
        if a.ptr == b.ptr {
63 +
            return true;
64 +
        }
63 65
    }
64 66
    for i in 0..a.len {
65 67
        if a[i] <> b[i] {
66 68
            return false;
67 69
        }
lib/std/sys/unix.rad +12 -4
32 32
constant AT_FDCWD: i64 = -100;
33 33
34 34
/// Opens a file at the given path and returns a file descriptor.
35 35
/// Returns a negative value on error.
36 36
export fn open(path: &[u8], flags: OpenFlags) -> i64 {
37 -
    return intrinsics::ecall(56, AT_FDCWD, path.ptr as i64, *flags, 0);
37 +
    unsafe {
38 +
        return intrinsics::ecall(56, AT_FDCWD, path.ptr as i64, *flags, 0);
39 +
    }
38 40
}
39 41
40 42
/// Opens a file at the given path with mode, returns a file descriptor.
41 43
export fn openOpts(path: &[u8], flags: OpenFlags, mode: i64) -> i64 {
42 -
    return intrinsics::ecall(56, AT_FDCWD, path.ptr as i64, *flags, mode);
44 +
    unsafe {
45 +
        return intrinsics::ecall(56, AT_FDCWD, path.ptr as i64, *flags, mode);
46 +
    }
43 47
}
44 48
45 49
/// Reads from a file descriptor into the provided buffer.
46 50
/// Returns the number of bytes read, or a negative value on error.
47 51
export fn read(fd: i64, buf: &mut [u8]) -> i64 {
48 -
    return intrinsics::ecall(63, fd, buf.ptr as i64, buf.len as i64, 0);
52 +
    unsafe {
53 +
        return intrinsics::ecall(63, fd, buf.ptr as i64, buf.len as i64, 0);
54 +
    }
49 55
}
50 56
51 57
/// Reads from a file descriptor until EOF or buffer is full.
52 58
/// Returns the total number of bytes read, or a negative value on error.
53 59
export fn readToEnd(fd: i64, buf: &mut [u8]) -> i64 {
67 73
}
68 74
69 75
/// Writes to a file descriptor from the provided buffer.
70 76
/// Returns the number of bytes written, or a negative value on error.
71 77
export fn write(fd: i64, buf: &[u8]) -> i64 {
72 -
    return intrinsics::ecall(64, fd, buf.ptr as i64, buf.len as i64, 0);
78 +
    unsafe {
79 +
        return intrinsics::ecall(64, fd, buf.ptr as i64, buf.len as i64, 0);
80 +
    }
73 81
}
74 82
75 83
/// Writes the entire contents of a buffer to a file descriptor.
76 84
/// Returns `false` when the descriptor cannot accept the full buffer.
77 85
export fn writeAll(fd: i64, data: &[u8]) -> bool {
lib/std/tests.rad +9 -0
3 3
use std::fmt;
4 4
use std::mem;
5 5
use std::vec;
6 6
use std::testing;
7 7
use std::sys::unix;
8 +
use std::io;
9 +
10 +
/// Safe I/O and equality accept empty buffers.
11 +
@test fn testSafeEmptyBuffers() throws (testing::TestError) {
12 +
    io::print("");
13 +
    io::printError("");
14 +
    try testing::expect(unix::write(unix::STDOUT, "") == 0);
15 +
    try testing::expect(mem::eq("", ""));
16 +
}
8 17
9 18
// Data types //////////////////////////////////////////////////////////////////
10 19
11 20
record Point: Copy {
12 21
    x: i32,
test/tests/ecall.i64.rad +6 -4
6 6
@default fn main() -> i32 {
7 7
    // ecall(64, fd, buf, len, 0) = write(fd, buf, len).
8 8
    let msg: *[u8] = "ok\n";
9 9
10 10
    // Write to stdout. The pointer is passed as i64.
11 -
    let n = ecall(64, 1, msg.ptr as i64, msg.len as i64, 0);
11 +
    unsafe {
12 +
        let n = ecall(64, 1, msg.ptr as i64, msg.len as i64, 0);
12 13
13 -
    // Return value should be 3 (bytes written).
14 -
    if n <> 3 {
15 -
        return 1;
14 +
        // Return value should be 3 (bytes written).
15 +
        if n <> 3 {
16 +
            return 1;
17 +
        }
16 18
    }
17 19
    return 0;
18 20
}
test/tests/edge.cases.6.rad +5 -3
91 91
    set ANALYZER.entries[idx] = entry;
92 92
    set ANALYZER.len = idx + 1;
93 93
}
94 94
95 95
fn checkHeader(expected: *[Entry]) -> i32 {
96 -
    if ANALYZER.entries.ptr <> expected.ptr or ANALYZER.entries.len <> expected.len {
97 -
        // Slice header got clobbered instead of the backing storage.
98 -
        return 1;
96 +
    unsafe {
97 +
        if ANALYZER.entries.ptr <> expected.ptr or ANALYZER.entries.len <> expected.len {
98 +
            // Slice header got clobbered instead of the backing storage.
99 +
            return 1;
100 +
        }
99 101
    }
100 102
    if ANALYZER.len <> 1 {
101 103
        // Bookkeeping field was overwritten by the bad store.
102 104
        return 2;
103 105
    }
test/tests/pointer.slice.store.rad +2 -1
20 20
21 21
static STORAGE: [Entry; 2] = undefined;
22 22
static TABLE: Table = undefined;
23 23
static HOLDER: PtrBox = undefined;
24 24
25 -
@default fn main() -> i32 {
25 +
/// Exercise slice storage and inspect its pointer fields.
26 +
@default unsafe fn main() -> i32 {
26 27
    set TABLE.entries = &mut STORAGE[..];
27 28
    set TABLE.len     = 0;
28 29
29 30
    set HOLDER.ptr = &mut TABLE.entries;
30 31
test/tests/slice.basic.rad +1 -1
7 7
fn sliceCap(s: *[i32]) -> u32 {
8 8
    return s.cap;
9 9
}
10 10
11 11
/// Returns the pointer field from a slice header.
12 -
fn slicePtr(s: *[i32]) -> *i32 {
12 +
unsafe fn slicePtr(s: *[i32]) -> *i32 {
13 13
    return s.ptr;
14 14
}
15 15
16 16
/// Copies a slice header into a new binding and returns it.
17 17
fn sliceCopy(s: *[i32]) -> *[i32] {
test/tests/slice.ptr.checked.bounds.rad added +11 -0
1 +
//! returns: 133
2 +
3 +
/// Storage for an empty slice.
4 +
static DATA: [u8; 1] = [42];
5 +
6 +
/// A checked address requires an element within the slice.
7 +
@default fn main() -> u8 {
8 +
    let s = &DATA[1..];
9 +
    let p = &s[0];
10 +
    return *p;
11 +
}
test/tests/slice.ptr.unsafe.rad added +26 -0
1 +
//! returns: 0
2 +
3 +
/// Backing storage for pointer access.
4 +
static DATA: [u8; 2] = [11, 22];
5 +
6 +
/// Read the first byte through a checked address.
7 +
fn first(s: *[u8]) -> u8 {
8 +
    let p = &s[0];
9 +
    return *p;
10 +
}
11 +
12 +
/// Expose a slice pointer under the caller's unsafe contract.
13 +
unsafe fn pointer(s: *[u8]) -> *u8 {
14 +
    return s.ptr;
15 +
}
16 +
17 +
/// Exercise checked access and explicit unsafe pointer access.
18 +
@default fn main() -> i32 {
19 +
    if first(&DATA[..]) <> 11 { return 1; }
20 +
    let empty = &DATA[2..];
21 +
    unsafe {
22 +
        if empty.ptr <> &DATA[0] + 2 { return 2; }
23 +
        if *pointer(&DATA[..]) <> 11 { return 3; }
24 +
    }
25 +
    return 0;
26 +
}