Files

182 lines
6.9 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package io
import (
"encoding/binary"
"math"
"os"
"syscall"
"unsafe"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Memory-mapped arrays. A file of native-endian numbers can back an
// Array directly: the operating system maps the file's pages into the
// address space and the array reads them in place, so a data cube far
// larger than RAM opens instantly and only the touched pages ever
// reach memory. The mapping is read-only, which matches the arrays'
// immutability contract exactly, and the caller releases it with the
// returned function once the numbers are no longer needed.
// MapFloats maps n float64 values of path, starting at byte offset,
// into a read-only one-dimensional array. The values must have been
// written in the machine's native byte order (binary.NativeEndian).
// The array is a live view of the mapping: release unmaps it, and any
// use of the array afterwards is a use-after-free, so release comes
// strictly last. A negative offset, a non-positive count, a missing
// file or a file that does not hold all n values is an error.
func MapFloats(path string, offset int64, n int) (a *core.Array, release func() error, err error) {
const name = "MapFloats"
length, err := countBytes(name, n, 8)
if err != nil {
return nil, nil, err
}
raw, release, err := mapRegion(path, offset, length, 8, name)
if err != nil {
return nil, nil, err
}
// SAFETY: countBytes checked that n is positive and that n*8 does not
// overflow, and mapRegion checked that the file holds that many bytes
// from an offset that is a multiple of 8 and returned a mapping whose
// base is page-aligned with exactly that intra-page offset, so the
// pointer is 8-aligned and n float64 values lie inside the mapping.
// The mapping stays alive until the caller runs release.
values := unsafe.Slice((*float64)(unsafe.Pointer(&raw[0])), n)
a, err = core.FromFloatSlice(values, n)
if err != nil {
_ = release()
return nil, nil, base.Errf("%s: %w", name, err)
}
return a, release, nil
}
// MapFloat32s maps n float32 values of path into a read-only array,
// with the same contract as MapFloats.
func MapFloat32s(path string, offset int64, n int) (a *core.Array, release func() error, err error) {
const name = "MapFloat32s"
length, err := countBytes(name, n, 4)
if err != nil {
return nil, nil, err
}
raw, release, err := mapRegion(path, offset, length, 4, name)
if err != nil {
return nil, nil, err
}
// SAFETY: as in MapFloats, with the float32 alignment of 4.
values := unsafe.Slice((*float32)(unsafe.Pointer(&raw[0])), n)
a, err = core.FromFloat32Slice(values, n)
if err != nil {
_ = release()
return nil, nil, base.Errf("%s: %w", name, err)
}
return a, release, nil
}
// MapInts maps n int64 values of path into a read-only array, with
// the same contract as MapFloats.
func MapInts(path string, offset int64, n int) (a *core.Array, release func() error, err error) {
const name = "MapInts"
length, err := countBytes(name, n, 8)
if err != nil {
return nil, nil, err
}
raw, release, err := mapRegion(path, offset, length, 8, name)
if err != nil {
return nil, nil, err
}
// SAFETY: as in MapFloats; an int64 has the same size and alignment.
values := unsafe.Slice((*int64)(unsafe.Pointer(&raw[0])), n)
a, err = core.IntsFromArray(values, n)
if err != nil {
_ = release()
return nil, nil, base.Errf("%s: %w", name, err)
}
return a, release, nil
}
// countBytes converts an element count into a byte length, refusing the
// counts whose byte size does not fit an int64. Multiplying first and
// checking the product afterwards is the failure the review found: for
// n = 2^61+3 float64 values the product wraps to 24, which passes every
// file-size check, and the typed view is then built with the original n,
// which unsafe.Slice rejects with a panic instead of an error.
func countBytes(name string, n int, width int64) (int64, error) {
if n <= 0 {
return 0, base.Errf("%s: the element count must be positive, got %d", name, n)
}
if int64(n) > math.MaxInt64/width {
return 0, base.Errf("%s: %d elements of %d bytes take more bytes than a length can address",
name, n, width)
}
return int64(n) * width, nil
}
// mapRegion maps length bytes of path at offset read-only and returns
// the byte slice viewing them plus a release that unmaps. The count,
// the offset and the element alignment are checked against the file's
// size before the mapping, so a short file or a misaligned request
// fails with an error instead of a fault or a misaligned array.
func mapRegion(path string, offset, length, align int64, name string) ([]byte, func() error, error) {
if offset < 0 {
return nil, nil, base.Errf("%s: offset must not be negative, got %d", name, offset)
}
// The caller turns the bytes into a typed slice with unsafe, which
// requires the address to satisfy the element's alignment: an
// offset that is not a multiple of the element size would hand back
// a misaligned array.
if offset%align != 0 {
return nil, nil, base.Errf("%s: offset %d is not a multiple of %d, the element size", name, offset, align)
}
if length <= 0 {
return nil, nil, base.Errf("%s: length must be positive, got %d", name, length)
}
file, err := os.Open(path)
if err != nil {
return nil, nil, base.Errf("%s: %w", name, err)
}
defer file.Close()
info, err := file.Stat()
if err != nil {
return nil, nil, base.Errf("%s: %w", name, err)
}
if info.IsDir() {
return nil, nil, base.Errf("%s: %s is a directory", name, path)
}
if size := info.Size(); offset > size || length > size-offset {
return nil, nil, base.Errf("%s: %s holds %d bytes, %d needed from offset %d",
name, path, size, length, offset)
}
// The mapping call demands a page-aligned offset; misaligned
// requests map from the page boundary below and the returned slice
// skips the intra-page part.
page := int64(os.Getpagesize())
pageBase := offset / page * page
raw, err := syscall.Mmap(int(file.Fd()), pageBase, int(length+offset-pageBase), syscall.PROT_READ, syscall.MAP_SHARED)
if err != nil {
return nil, nil, base.Errf("%s: mapping %s failed: %w", name, path, err)
}
view := raw[offset-pageBase:]
return view, func() error {
if err := syscall.Munmap(raw); err != nil {
return base.Errf("%s: unmapping %s failed: %w", name, path, err)
}
return nil
}, nil
}
// SaveNativeFloats writes floats to path in the machine's native byte
// order, the format MapFloats reads back. It exists so a mapping test
// or tool can produce its own data without reaching for encoding
// details; a zero offset aligns it with MapFloats' contract.
func SaveNativeFloats(path string, values []float64) error {
buf := make([]byte, 8*len(values))
for i, v := range values {
binary.NativeEndian.PutUint64(buf[i*8:], math.Float64bits(v))
}
return os.WriteFile(path, buf, 0o644)
}