diff --git a/.gitattributes b/.gitattributes index e375d2543..25300b4ef 100644 --- a/.gitattributes +++ b/.gitattributes @@ -10,3 +10,26 @@ vendor/sdl3/SDL3.dll filter=lfs diff=lfs merge=lfs -text vendor/sdl3/SDL3.lib filter=lfs diff=lfs merge=lfs -text vendor/sdl3/mixer/*.dll filter=lfs diff=lfs merge=lfs -text vendor/sdl3/mixer/*.lib filter=lfs diff=lfs merge=lfs -text +vendor/wgpu/lib/wgpu-windows-x86_64-msvc-release/lib/wgpu_native.dll filter=lfs diff=lfs merge=lfs -text +vendor/wgpu/lib/wgpu-windows-x86_64-msvc-release/lib/wgpu_native.dll.lib filter=lfs diff=lfs merge=lfs -text +vendor/wgpu/lib/wgpu-windows-x86_64-msvc-release/lib/wgpu_native.lib filter=lfs diff=lfs merge=lfs -text +vendor/wgpu/lib/wgpu-windows-x86_64-msvc-release/lib/wgpu_native.pdb filter=lfs diff=lfs merge=lfs -text +vendor/raylib/linux/libraygui.so filter=lfs diff=lfs merge=lfs -text +vendor/raylib/linux/libraygui.a filter=lfs diff=lfs merge=lfs -text +vendor/raylib/linux/libraylib.a filter=lfs diff=lfs merge=lfs -text +vendor/raylib/linux/libraylib.so.600 filter=lfs diff=lfs merge=lfs -text +vendor/raylib/linux-arm64/libraylib.a filter=lfs diff=lfs merge=lfs -text +vendor/raylib/linux-arm64/libraylib.so.600 filter=lfs diff=lfs merge=lfs -text +vendor/raylib/macos/libraygui-arm64.dylib filter=lfs diff=lfs merge=lfs -text +vendor/raylib/macos/libraygui.dylib filter=lfs diff=lfs merge=lfs -text +vendor/raylib/macos/libraylib.a filter=lfs diff=lfs merge=lfs -text +vendor/raylib/macos/libraylib.600.dylib filter=lfs diff=lfs merge=lfs -text +vendor/raylib/wasm/libraylib.web.a filter=lfs diff=lfs merge=lfs -text +vendor/raylib/windows/raylib.dll filter=lfs diff=lfs merge=lfs -text +vendor/raylib/windows/raylib.lib filter=lfs diff=lfs merge=lfs -text +vendor/raylib/windows/raylibdll.lib filter=lfs diff=lfs merge=lfs -text +vendor/raylib/windows/raygui.dll filter=lfs diff=lfs merge=lfs -text +vendor/raylib/windows/raygui.lib filter=lfs diff=lfs merge=lfs -text +vendor/raylib/windows/rayguidll.lib filter=lfs diff=lfs merge=lfs -text +vendor/box3d/lib/linux-amd64/libbox3d.a filter=lfs diff=lfs merge=lfs -text +vendor/box3d/lib/darwin/libbox3d.a filter=lfs diff=lfs merge=lfs -text diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index dd49b8ff8..35bc14a6c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -27,9 +27,9 @@ jobs: gmake release ./odin version ./odin report - gmake -C vendor/stb/src - gmake -C vendor/cgltf/src - gmake -C vendor/miniaudio/src + (cd ./vendor/stb/src/; ./build_stb.sh) + (cd ./vendor/cgltf/src/; ./build_cgltf.sh) + (cd ./vendor/miniaudio/src/; ./build_miniaudio.sh) ./odin check examples/all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:netbsd_amd64 ./odin check examples/all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:netbsd_arm64 ./odin check examples/all/sdl3 -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:netbsd_amd64 -no-entry-point @@ -60,9 +60,9 @@ jobs: gmake release ./odin version ./odin report - gmake -C vendor/stb/src - gmake -C vendor/cgltf/src - gmake -C vendor/miniaudio/src + (cd ./vendor/stb/src/; ./build_stb.sh) + (cd ./vendor/cgltf/src/; ./build_cgltf.sh) + (cd ./vendor/miniaudio/src/; ./build_miniaudio.sh) ./odin check examples/all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:freebsd_amd64 ./odin check examples/all/sdl3 -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:freebsd_amd64 -no-entry-point ./odin test tests/core/normal.odin -file -all-packages -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -define:ODIN_TEST_FANCY=false -define:ODIN_TEST_FAIL_ON_BAD_MEMORY=true @@ -122,9 +122,9 @@ jobs: brew install curl mbedtls - name: Compile needed Vendor run: | - make -C vendor/stb/src - make -C vendor/cgltf/src - make -C vendor/miniaudio/src + (cd ./vendor/stb/src/; ./build_stb.sh) + (cd ./vendor/cgltf/src/; ./build_cgltf.sh) + (cd ./vendor/miniaudio/src/; ./build_miniaudio.sh) - name: Odin check run: ./odin check examples/demo -vet - name: Odin run @@ -138,10 +138,13 @@ jobs: - name: Normal Core library tests run: ./odin test tests/core/normal.odin -file -all-packages -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -define:ODIN_TEST_FANCY=false -define:ODIN_TEST_FAIL_ON_BAD_MEMORY=true -sanitize:address - name: Optimized Core library tests - run: ./odin test tests/core/speed.odin -o:speed -file -all-packages -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -define:ODIN_TEST_FANCY=false -define:ODIN_TEST_FAIL_ON_BAD_MEMORY=true -sanitize:address + run: ./odin test tests/core/speed.odin -o:speed -file -all-packages -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -define:ODIN_TEST_FANCY=false -define:ODIN_TEST_FAIL_ON_BAD_MEMORY=true - name: Wycheproof tests - run: ./odin test tests/core/crypto/wycheproof -vet -vet-tabs -strict-style -vet-style -vet-cast -warnings-as-errors -disallow-do -o:speed + - name: Noise Protocol Framework tests + run: ./odin test tests/core/crypto/noise -vet -vet-tabs -strict-style -vet-style -vet-cast -warnings-as-errors -disallow-do -o:speed + - name: X.509 limbo tests + run: ./odin test tests/core/crypto/x509_limbo -vet -vet-tabs -strict-style -vet-style -vet-cast -warnings-as-errors -disallow-do -o:speed - name: Vendor library tests run: ./odin test tests/vendor -all-packages -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -define:ODIN_TEST_FANCY=false -define:ODIN_TEST_FAIL_ON_BAD_MEMORY=true -sanitize:address if: matrix.os != 'macos-15-intel' && matrix.os != 'macos-latest' @@ -238,12 +241,22 @@ jobs: shell: cmd run: | call "C:\Program Files\Microsoft Visual Studio\2022\Enterprise\VC\Auxiliary\Build\vcvars64.bat" - odin test tests/core/speed.odin -o:speed -file -all-packages -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -define:ODIN_TEST_FANCY=false -define:ODIN_TEST_FAIL_ON_BAD_MEMORY=true -sanitize:address + odin test tests/core/speed.odin -o:speed -file -all-packages -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -define:ODIN_TEST_FANCY=false -define:ODIN_TEST_FAIL_ON_BAD_MEMORY=true - name: Wycheproof tests shell: cmd run: | call "C:\Program Files\Microsoft Visual Studio\2022\Enterprise\VC\Auxiliary\Build\vcvars64.bat" odin test tests/core/crypto/wycheproof -vet -vet-tabs -strict-style -vet-style -vet-cast -warnings-as-errors -disallow-do -o:speed -microarch:native + - name: Noise Protocol Framework tests + shell: cmd + run: | + call "C:\Program Files\Microsoft Visual Studio\2022\Enterprise\VC\Auxiliary\Build\vcvars64.bat" + odin test tests/core/crypto/noise -vet -vet-tabs -strict-style -vet-style -vet-cast -warnings-as-errors -disallow-do -o:speed -microarch:native + - name: X.509 limbo tests + shell: cmd + run: | + call "C:\Program Files\Microsoft Visual Studio\2022\Enterprise\VC\Auxiliary\Build\vcvars64.bat" + odin test tests/core/crypto/x509_limbo -vet -vet-tabs -strict-style -vet-style -vet-cast -warnings-as-errors -disallow-do -o:speed -microarch:native - name: Vendor library tests shell: cmd run: | @@ -307,9 +320,9 @@ jobs: - name: Compile needed Vendor run: | - make -C vendor/stb/src - make -C vendor/cgltf/src - make -C vendor/miniaudio/src + (cd ./vendor/stb/src/; ./build_stb.sh) + (cd ./vendor/cgltf/src/; ./build_cgltf.sh) + (cd ./vendor/miniaudio/src/; ./build_miniaudio.sh) - name: Odin check examples/all run: ./odin check examples/all -target:linux_riscv64 -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do diff --git a/.github/workflows/nightly.yml b/.github/workflows/nightly.yml index c7858c82c..5d8fe2e8e 100644 --- a/.github/workflows/nightly.yml +++ b/.github/workflows/nightly.yml @@ -44,27 +44,34 @@ jobs: name: windows_artifacts path: dist build_linux: - name: Linux Build + strategy: + matrix: + arch: [amd64, arm64] + env: + ARCH: ${{ matrix.arch }} + DOCKER_IMAGE: ${{ matrix.arch == 'amd64' && 'alpine:3.24' || 'arm64v8/alpine:3.24' }} + name: Linux Build ${{ matrix.arch }} if: github.repository == 'odin-lang/Odin' - runs-on: ubuntu-latest + runs-on: ${{ matrix.arch == 'amd64' && 'ubuntu-latest' || 'ubuntu-24.04-arm' }} steps: - uses: actions/checkout@v4 with: lfs: true - name: (Linux) Download LLVM and Build Odin run: | - docker run --rm -v "$PWD:/src" -w /src alpine sh -c ' + docker run --rm -v "$PWD:/src" -w /src $DOCKER_IMAGE sh -c ' apk add --no-cache \ - musl-dev llvm20-dev clang20 git mold lz4 \ + musl-dev llvm20-dev clang20 build-base git mold lz4 \ libxml2-static llvm20-static zlib-static zstd-static \ make && + git config --global --add safe.directory /src && ./ci/build_linux_static.sh ' - name: Odin run run: ./odin run examples/demo - name: Copy artifacts run: | - FILE="odin-linux-amd64-nightly+$(date -I)" + FILE="odin-linux-$ARCH-nightly+$(date -I)" mkdir $FILE cp odin $FILE cp LICENSE $FILE @@ -78,59 +85,22 @@ jobs: tar -czvf dist.tar.gz $FILE - name: Odin run run: | - FILE="odin-linux-amd64-nightly+$(date -I)" + FILE="odin-linux-$ARCH-nightly+$(date -I)" $FILE/odin run examples/demo - name: Upload artifact uses: actions/upload-artifact@v4 with: - name: linux_artifacts - path: dist.tar.gz - build_linux_arm: - name: Linux ARM Build - if: github.repository == 'odin-lang/Odin' - runs-on: ubuntu-24.04-arm - steps: - - uses: actions/checkout@v4 - with: - lfs: true - - name: (Linux ARM) Download LLVM and Build Odin - run: | - docker run --rm -v "$PWD:/src" -w /src arm64v8/alpine sh -c ' - apk add --no-cache \ - musl-dev llvm20-dev clang20 git mold lz4 \ - libxml2-static llvm20-static zlib-static zstd-static \ - make && - ./ci/build_linux_static.sh - ' - - name: Odin run - run: ./odin run examples/demo - - name: Copy artifacts - run: | - FILE="odin-linux-arm64-nightly+$(date -I)" - mkdir $FILE - cp odin $FILE - cp LICENSE $FILE - cp -r shared $FILE - cp -r base $FILE - cp -r core $FILE - cp -r vendor $FILE - cp -r examples $FILE - ./ci/remove_windows_binaries.sh $FILE - # Creating a tarball so executable permissions are retained, see https://github.com/actions/upload-artifact/issues/38 - tar -czvf dist.tar.gz $FILE - - name: Odin run - run: | - FILE="odin-linux-arm64-nightly+$(date -I)" - $FILE/odin run examples/demo - - name: Upload artifact - uses: actions/upload-artifact@v4 - with: - name: linux_arm_artifacts + name: ${{ matrix.arch == 'amd64' && 'linux_artifacts' || 'linux_arm_artifacts' }} path: dist.tar.gz build_macos: - name: MacOS Build + strategy: + matrix: + arch: [amd64, arm64] + env: + ARCH: ${{ matrix.arch }} + name: Macos Build ${{ matrix.arch }} if: github.repository == 'odin-lang/Odin' - runs-on: macos-15-intel + runs-on: ${{ matrix.arch == 'amd64' && 'macos-15-intel' || 'macos-latest' }} steps: - uses: actions/checkout@v4 with: @@ -139,6 +109,7 @@ jobs: run: | brew update brew install llvm@20 dylibbundler lld@20 + brew link llvm@20 - name: build odin # These -L makes the linker prioritize system libraries over LLVM libraries, this is mainly to @@ -146,7 +117,7 @@ jobs: run: CXXFLAGS="-L/usr/lib/system -L/usr/lib" make nightly - name: Bundle run: | - FILE="odin-macos-amd64-nightly+$(date -I)" + FILE="odin-macos-$ARCH-nightly+$(date -I)" mkdir $FILE cp odin $FILE cp LICENSE $FILE @@ -161,57 +132,16 @@ jobs: tar -czvf dist.tar.gz $FILE - name: Odin run run: | - FILE="odin-macos-amd64-nightly+$(date -I)" + FILE="odin-macos-$ARCH-nightly+$(date -I)" $FILE/odin run examples/demo - name: Upload artifact uses: actions/upload-artifact@v4 with: - name: macos_artifacts - path: dist.tar.gz - build_macos_arm: - name: MacOS ARM Build - if: github.repository == 'odin-lang/Odin' - runs-on: macos-latest # ARM machine - steps: - - uses: actions/checkout@v4 - with: - lfs: true - - name: Download LLVM and setup PATH - run: | - brew update - brew install llvm@20 dylibbundler lld@20 - - - name: build odin - # These -L makes the linker prioritize system libraries over LLVM libraries, this is mainly to - # not link with libunwind bundled with LLVM but link with libunwind on the system. - run: CXXFLAGS="-L/usr/lib/system -L/usr/lib" make nightly - - name: Bundle - run: | - FILE="odin-macos-arm64-nightly+$(date -I)" - mkdir $FILE - cp odin $FILE - cp LICENSE $FILE - cp -r shared $FILE - cp -r base $FILE - cp -r core $FILE - cp -r vendor $FILE - cp -r examples $FILE - ./ci/remove_windows_binaries.sh $FILE - dylibbundler -b -x $FILE/odin -d $FILE/libs -od -p @executable_path/libs - # Creating a tarball so executable permissions are retained, see https://github.com/actions/upload-artifact/issues/38 - tar -czvf dist.tar.gz $FILE - - name: Odin run - run: | - FILE="odin-macos-arm64-nightly+$(date -I)" - $FILE/odin run examples/demo - - name: Upload artifact - uses: actions/upload-artifact@v4 - with: - name: macos_arm_artifacts + name: ${{ matrix.arch == 'amd64' && 'macos_artifacts' || 'macos_arm_artifacts' }} path: dist.tar.gz upload_b2: runs-on: [ubuntu-latest] - needs: [build_windows, build_macos, build_macos_arm, build_linux, build_linux_arm] + needs: [build_windows, build_macos, build_linux] steps: - uses: actions/checkout@v4 - uses: actions/setup-python@v5 diff --git a/.gitignore b/.gitignore index 00dbcc2a5..bf0efe2a0 100644 --- a/.gitignore +++ b/.gitignore @@ -277,7 +277,15 @@ odin *.bin demo.bin libLLVM*.so* -*.a + +# core:rexcode commits its generated encode/decode tables: each arch's tablegen +# emits human-readable Odin, serialized to tables/*.bin, which the library +# #loads at compile time. ISA packages live under core/rexcode/isa/; the +# broad x86/ rule above otherwise shadows the isa/x86 source package, so +# un-ignore that directory tree and every arch's committed table blobs. Stray +# *.bin build output inside the tree stays ignored -- only tables/*.bin is kept. +!/core/rexcode/isa/x86/ +!core/rexcode/isa/*/tables/*.bin # WASM *.wasm diff --git a/base/intrinsics/intrinsics.odin b/base/intrinsics/intrinsics.odin index ee99da1b3..e639aadbd 100644 --- a/base/intrinsics/intrinsics.odin +++ b/base/intrinsics/intrinsics.odin @@ -4,6 +4,7 @@ package intrinsics import "base:runtime" + // Package-Related is_package_imported :: proc(package_name: string) -> bool --- @@ -184,6 +185,8 @@ type_is_simd_vector :: proc($T: typeid) -> bool --- type_is_matrix :: proc($T: typeid) -> bool --- type_is_fixed_capacity_dynamic_array :: proc($T: typeid) -> bool --- +type_is_internally_pointer_like :: proc($T: typeid) -> bool --- + type_has_nil :: proc($T: typeid) -> bool --- type_is_matrix_row_major :: proc($T: typeid) -> bool where type_is_matrix(T) --- @@ -205,12 +208,17 @@ type_bit_set_underlying_type :: proc($T: typeid) -> typeid where type_is_bit_set type_has_field :: proc($T: typeid, $name: string) -> bool --- type_field_type :: proc($T: typeid, $name: string) -> typeid --- +type_field_bit_size :: proc($T: typeid, $name: string) -> int where type_is_bit_field(T) --- +type_field_bit_offset :: proc($T: typeid, $name: string) -> int where type_is_bit_field(T) --- + type_proc_parameter_count :: proc($T: typeid) -> int where type_is_proc(T) --- type_proc_return_count :: proc($T: typeid) -> int where type_is_proc(T) --- type_proc_parameter_type :: proc($T: typeid, index: int) -> typeid where type_is_proc(T) --- type_proc_return_type :: proc($T: typeid, index: int) -> typeid where type_is_proc(T) --- +type_proc_calling_convention :: proc($T: typeid) -> Odin_Calling_Convention where type_is_proc(T) --- + type_struct_field_count :: proc($T: typeid) -> int where type_is_struct(T) --- type_struct_has_implicit_padding :: proc($T: typeid) -> bool where type_is_struct(T) --- @@ -245,6 +253,8 @@ type_integer_to_signed :: proc($T: typeid) -> type where type_is_integer(T), t type_has_shared_fields :: proc($U, $V: typeid) -> bool where type_is_struct(U), type_is_struct(V) --- + + // Returns the canonicalized name of the type, of which is used to produce the pseudo-unique 'typeid' type_canonical_name :: proc($T: typeid) -> string --- @@ -345,7 +355,11 @@ simd_trunc :: proc(a: #simd[N]any_float) -> #simd[N]any_float --- // rounding to the nearest integral value; if two values are equally near, rounds to the even one simd_nearest :: proc(a: #simd[N]any_float) -> #simd[N]any_float --- -simd_to_bits :: proc(v: #simd[N]T) -> #simd[N]Integer where size_of(T) == size_of(Integer), type_is_unsigned(Integer) --- +simd_approx_recip :: proc(x: #simd[N]T) -> #simd[N]T where type_is_float(T)) --- +simd_approx_recip_sqrt :: proc(x: #simd[N]T) -> #simd[N]T where type_is_float(T)) --- + +simd_to_bits :: proc(v: #simd[N]T) -> #simd[N]Integer where size_of(T) == size_of(Integer), type_is_unsigned(Integer) --- +simd_to_bits_signed :: proc(v: #simd[N]T) -> #simd[N]Integer where size_of(T) == size_of(Integer), !type_is_unsigned(Integer) --- // equivalent to a swizzle with descending indices, e.g. reserve(a, 3, 2, 1, 0) simd_lanes_reverse :: proc(a: #simd[N]T) -> #simd[N]T --- @@ -353,6 +367,19 @@ simd_lanes_reverse :: proc(a: #simd[N]T) -> #simd[N]T --- simd_lanes_rotate_left :: proc(a: #simd[N]T, $offset: int) -> #simd[N]T --- simd_lanes_rotate_right :: proc(a: #simd[N]T, $offset: int) -> #simd[N]T --- +// return {b[0], a[1], b[2], a[3], ...} +simd_odd_even :: proc(a, b: #simd[N]T) -> #simd[N]T --- + +// Returns the sums of N consecutive lanes +simd_sums_of_n :: proc(a: #simd[LANES]T, $N: uint) -> #simd[LANES/N]T where is_power_of_two(N) --- + +simd_pairwise_add :: proc(a, b: #simd[LANES]T) -> #simd[LANES]T where LANES % 2 == 0 --- +simd_pairwise_sub :: proc(a, b: #simd[LANES]T) -> #simd[LANES]T where LANES % 2 == 0 --- + +simd_interleave :: proc(a, ..#simd[LANES/N]T) -> #simd[LANES]T where N >= 1 --- +simd_deinterleave :: proc(a: #simd[LANES]T, $N: uint) -> (..#simd[LANES/N]T) where N >= 1, LANES % N == 0 --- // returns N multiple vectors + + // Checks if the current target supports the given target features. // // Takes a constant comma-seperated string (eg: "sha512,sse4.1"), or a procedure type which has either @@ -387,6 +414,14 @@ x86_cpuid :: proc(ax, cx: u32) -> (eax, ebx, ecx, edx: u32) --- x86_xgetbv :: proc(cx: u32) -> (eax, edx: u32) --- +// C specific things +c_va_list :: struct{/*platform specific implementation*/} + +c_va_start :: proc(list: ^c_va_list, /*#c_vararg parameter*/ args: ..$T) --- +c_va_end :: proc(list: ^c_va_list) --- +c_va_copy :: proc(dst, src: ^c_va_list) --- +c_va_arg :: proc(list: ^c_va_list, $T: typeid) -> T --- + // Darwin targets only objc_object :: struct{} diff --git a/base/runtime/core.odin b/base/runtime/core.odin index dbd8ba468..41d620bdd 100644 --- a/base/runtime/core.odin +++ b/base/runtime/core.odin @@ -23,8 +23,27 @@ package runtime import "base:intrinsics" +/* +Fast_Math_Flag :: enum u8 { + Allow_Reassoc = 0, + No_NaNs = 1, + No_Infs = 2, + No_Signed_Zeros = 3, + Allow_Reciprocal = 4, + Allow_Contract = 5, + Approx_Func = 6, +} +*/ +Fast_Math_Flag :: intrinsics.Fast_Math_Flag + +// Fast_Math_Flags :: distinct bit_set[Fast_Math_Flag; u32] +Fast_Math_Flags :: intrinsics.Fast_Math_Flags + + // NOTE(bill): This must match the compiler's -Calling_Convention :: enum u8 { + +/* +enum u8 { Invalid = 0, Odin = 1, Contextless = 2, @@ -44,6 +63,8 @@ Calling_Convention :: enum u8 { Preserve_Most = 12, Preserve_All = 13, } +*/ +Calling_Convention :: type_of(ODIN_DEFAULT_CALLING_CONVENTION) Type_Info_Enum_Value :: distinct i64 @@ -279,7 +300,79 @@ when ODIN_OS == .Windows { dll_instance: rawptr } -// IMPORTANT NOTE(bill): Must be in this order (as the compiler relies upon it) + +// This is safe to change. The log2 size of a cache-line. At minimum it has to +// be six though. Higher cache line sizes are permitted. +MAP_CACHE_LINE_LOG2 :: 6 + +// The size of a cache-line. +MAP_CACHE_LINE_SIZE :: 1 << MAP_CACHE_LINE_LOG2 + +// The minimum cache-line size allowed by this implementation is 64 bytes since +// we need 6 bits in the base pointer to store the integer log2 capacity, which +// at maximum is 63. Odin uses signed integers to represent length and capacity, +// so only 63 bits are needed in the maximum case. +#assert(MAP_CACHE_LINE_SIZE >= 64) + +// Map_Cell type that packs multiple T in such a way to ensure that each T stays +// aligned by align_of(T) and such that align_of(Map_Cell(T)) % MAP_CACHE_LINE_SIZE == 0 +// +// This means a value of type T will never straddle a cache-line. +// +// When multiple Ts can fit in a single cache-line the data array will have more +// than one element. When it cannot, the data array will have one element and +// an array of Map_Cell(T) will be padded to stay a multiple of MAP_CACHE_LINE_SIZE. +// +// We rely on the type system to do all the arithmetic and padding for us here. +// +// The usual array[index] indexing for []T backed by a []Map_Cell(T) becomes a bit +// more involved as there now may be internal padding. The indexing now becomes +// +// N :: len(Map_Cell(T){}.data) +// i := index / N +// j := index % N +// cell[i].data[j] +// +// However, since len(Map_Cell(T){}.data) is a compile-time constant, there are some +// optimizations we can do to eliminate the need for any divisions as N will +// be bounded by [1, 64). +// +// In the optimal case, len(Map_Cell(T){}.data) = 1 so the cell array can be treated +// as a regular array of T, which is the case for hashes. +Map_Cell :: struct($T: typeid) #align(MAP_CACHE_LINE_SIZE) { + data: [MAP_CACHE_LINE_SIZE / size_of(T) when 0 < size_of(T) && size_of(T) < MAP_CACHE_LINE_SIZE else 1]T, +} + +// So we can operate on a cell data structure at runtime without any type +// information, we have a simple table that stores some traits about the cell. +// +// 32-bytes on 64-bit +// 16-bytes on 32-bit +Map_Cell_Info :: struct { + size_of_type: uintptr, // 8-bytes on 64-bit, 4-bytes on 32-bits + align_of_type: uintptr, // 8-bytes on 64-bit, 4-bytes on 32-bits + size_of_cell: uintptr, // 8-bytes on 64-bit, 4-bytes on 32-bits + elements_per_cell: uintptr, // 8-bytes on 64-bit, 4-bytes on 32-bits +} + +Map_Hash :: uintptr + +// When working with the type-erased structure at runtime we need information +// about the map to make working with it possible. This info structure stores +// that. +// +// `Map_Info` and `Map_Cell_Info` are read only data structures and cannot be +// modified after creation +// +// 32-bytes on 64-bit +// 16-bytes on 32-bit +Map_Info :: struct { + ks: ^Map_Cell_Info, // 8-bytes on 64-bit, 4-bytes on 32-bit + vs: ^Map_Cell_Info, // 8-bytes on 64-bit, 4-bytes on 32-bit + key_hasher: proc "contextless" (key: rawptr, seed: Map_Hash) -> Map_Hash, // 8-bytes on 64-bit, 4-bytes on 32-bit + key_equal: proc "contextless" (lhs, rhs: rawptr) -> bool, // 8-bytes on 64-bit, 4-bytes on 32-bit +} + Source_Code_Location :: struct { @@ -509,11 +602,9 @@ Raw_Quaternion256_Vector_Scalar :: struct {vector: [3]f64, scalar: f64} Windows, Darwin, Linux, - Essence, FreeBSD, OpenBSD, NetBSD, - Haiku, WASI, JS, Orca, @@ -565,7 +656,6 @@ Odin_Build_Mode_Type :: type_of(ODIN_BUILD_MODE) /* // Defined internally by the compiler Odin_Endian_Type :: enum int { - Unknown, Little, Big, } @@ -578,11 +668,9 @@ ALL_ODIN_OS_TYPES :: Odin_OS_Types{ .Windows, .Darwin, .Linux, - .Essence, .FreeBSD, .OpenBSD, .NetBSD, - .Haiku, .WASI, .JS, .Orca, diff --git a/base/runtime/core_builtin.odin b/base/runtime/core_builtin.odin index fe76ee5a5..e65cf56f0 100644 --- a/base/runtime/core_builtin.odin +++ b/base/runtime/core_builtin.odin @@ -2,6 +2,8 @@ package runtime import "base:intrinsics" +MAP_ENABLED :: !ODIN_BEDROCK + @builtin Maybe :: union($T: typeid) {T} @@ -65,7 +67,7 @@ when !NO_DEFAULT_TEMP_ALLOCATOR { // Initializes the global temporary allocator used as the default `context.temp_allocator`. // This is ignored when `NO_DEFAULT_TEMP_ALLOCATOR` is true. @(builtin, disabled=NO_DEFAULT_TEMP_ALLOCATOR) -init_global_temporary_allocator :: proc(size: int, backup_allocator := context.allocator) { +init_global_temporary_allocator :: proc "odin" (size: int, backup_allocator := context.allocator) { when !NO_DEFAULT_TEMP_ALLOCATOR { default_temp_allocator_init(&global_default_temp_allocator_data, size, backup_allocator) } @@ -387,7 +389,7 @@ pop_front_safe :: proc { @builtin clear :: proc{ clear_dynamic_array, - clear_map, + clear_map where MAP_ENABLED, clear_fixed_capacity_dynamic_array, clear_soa_dynamic_array, @@ -397,7 +399,7 @@ clear :: proc{ @builtin reserve :: proc{ reserve_dynamic_array, - reserve_map, + reserve_map where MAP_ENABLED, reserve_soa, } @@ -430,7 +432,7 @@ non_zero_resize :: proc{ @builtin shrink :: proc{ shrink_dynamic_array, - shrink_map, + shrink_map where MAP_ENABLED, } // `free` will try to free the passed pointer, with the given `allocator` if the allocator supports this operation. @@ -471,14 +473,6 @@ delete_dynamic_array :: proc(array: $T/[dynamic]$E, loc := #caller_location) -> delete_slice :: proc(array: $T/[]$E, allocator := context.allocator, loc := #caller_location) -> Allocator_Error { return mem_free_with_size(raw_data(array), len(array)*size_of(E), allocator, loc) } -// `delete_map` will try to free the underlying data of the passed map, with the given `allocator` if the allocator supports this operation. -// -// Note: Prefer the procedure group `delete`. -@builtin -delete_map :: proc(m: $T/map[$K]$V, loc := #caller_location) -> Allocator_Error { - return map_free_dynamic(transmute(Raw_Map)m, map_info(T), loc) -} - @builtin delete_string16 :: proc(str: string16, allocator := context.allocator, loc := #caller_location) -> Allocator_Error { @@ -489,6 +483,16 @@ delete_cstring16 :: proc(str: cstring16, allocator := context.allocator, loc := return mem_free((^u16)(str), allocator, loc) } +when MAP_ENABLED { + // `delete_map` will try to free the underlying data of the passed map, with the given `allocator` if the allocator supports this operation. + // + // Note: Prefer the procedure group `delete`. + @builtin + delete_map :: proc(m: $T/map[$K]$V, loc := #caller_location) -> Allocator_Error { + return map_free_dynamic(transmute(Raw_Map)m, map_info(T), loc) + } +} + // `delete` will try to free the underlying data of the passed built-in data structure (string, cstring, dynamic array, slice, or map), with the given `allocator` if the allocator supports this operation. // // Note: Prefer `delete` over the specific `delete_*` procedures where possible. @@ -498,7 +502,7 @@ delete :: proc{ delete_cstring, delete_dynamic_array, delete_slice, - delete_map, + delete_map where MAP_ENABLED, delete_soa_slice, delete_soa_dynamic_array, delete_string16, @@ -513,7 +517,7 @@ new :: proc($T: typeid, allocator := context.allocator, loc := #caller_location) t = (^T)(raw_data(mem_alloc_bytes(size_of(T), align_of(T), allocator, loc) or_return)) return } -@(require_results) +@(builtin, require_results) new_aligned :: proc($T: typeid, alignment: int, allocator := context.allocator, loc := #caller_location) -> (t: ^T, err: Allocator_Error) { t = (^T)(raw_data(mem_alloc_bytes(size_of(T), alignment, allocator, loc) or_return)) return @@ -530,7 +534,7 @@ new_clone :: proc(data: $T, allocator := context.allocator, loc := #caller_locat DEFAULT_DYNAMIC_ARRAY_CAPACITY :: 8 -@(require_results) +@(builtin, require_results) make_aligned :: proc($T: typeid/[]$E, #any_int len: int, alignment: int, allocator := context.allocator, loc := #caller_location) -> (res: T, err: Allocator_Error) #optional_allocator_error { err = _make_aligned_type_erased(&res, size_of(E), len, alignment, allocator, loc) return @@ -597,29 +601,32 @@ _make_dynamic_array_len_cap :: proc(array: ^Raw_Dynamic_Array, size_of_elem, ali return } -// `make_map` initializes a map with an allocator. Like `new`, the first argument is a type, not a value. -// Unlike `new`, `make`'s return value is the same as the type of its argument, not a pointer to it. -// -// Note: Prefer using the procedure group `make`. -@(builtin, require_results) -make_map :: proc($T: typeid/map[$K]$E, allocator := context.allocator, loc := #caller_location) -> (m: T) { - m.allocator = allocator - return m +when MAP_ENABLED { + // `make_map` initializes a map with an allocator. Like `new`, the first argument is a type, not a value. + // Unlike `new`, `make`'s return value is the same as the type of its argument, not a pointer to it. + // + // Note: Prefer using the procedure group `make`. + @(builtin, require_results) + make_map :: proc($T: typeid/map[$K]$E, allocator := context.allocator, loc := #caller_location) -> (m: T) { + m.allocator = allocator + return m + } + + // `make_map_cap` initializes a map with an allocator and allocates space using `capacity`. + // Like `new`, the first argument is a type, not a value. + // Unlike `new`, `make`'s return value is the same as the type of its argument, not a pointer to it. + // + // Note: Prefer using the procedure group `make`. + @(builtin, require_results) + make_map_cap :: proc($T: typeid/map[$K]$E, #any_int capacity: int, allocator := context.allocator, loc := #caller_location) -> (m: T, err: Allocator_Error) #optional_allocator_error { + make_map_expr_error_loc(loc, capacity) + context.allocator = allocator + + err = reserve_map(&m, capacity, loc) + return + } } -// `make_map_cap` initializes a map with an allocator and allocates space using `capacity`. -// Like `new`, the first argument is a type, not a value. -// Unlike `new`, `make`'s return value is the same as the type of its argument, not a pointer to it. -// -// Note: Prefer using the procedure group `make`. -@(builtin, require_results) -make_map_cap :: proc($T: typeid/map[$K]$E, #any_int capacity: int, allocator := context.allocator, loc := #caller_location) -> (m: T, err: Allocator_Error) #optional_allocator_error { - make_map_expr_error_loc(loc, capacity) - context.allocator = allocator - - err = reserve_map(&m, capacity, loc) - return -} // `make_multi_pointer` allocates and initializes a multi-pointer. Like `new`, the first argument is a type, not a value. // Unlike `new`, `make`'s return value is the same as the type of its argument, not a pointer to it. // @@ -649,8 +656,8 @@ make :: proc{ make_dynamic_array, make_dynamic_array_len, make_dynamic_array_len_cap, - make_map, - make_map_cap, + make_map where MAP_ENABLED, + make_map_cap where MAP_ENABLED, make_multi_pointer, make_soa_slice, @@ -659,64 +666,67 @@ make :: proc{ make_soa_dynamic_array_len_cap, } +when MAP_ENABLED { + // `clear_map` will set the length of a passed map to `0` + // + // Note: Prefer the procedure group `clear` + @builtin + clear_map :: proc "contextless" (m: ^$T/map[$K]$V) { + if m == nil { + return + } + map_clear_dynamic((^Raw_Map)(m), map_info(T)) + } -// `clear_map` will set the length of a passed map to `0` -// -// Note: Prefer the procedure group `clear` -@builtin -clear_map :: proc "contextless" (m: ^$T/map[$K]$V) { - if m == nil { + // `reserve_map` will try to reserve memory of a passed map to the requested element count (setting the `cap`). + // + // Note: Prefer the procedure group `reserve` + @builtin + reserve_map :: proc(m: ^$T/map[$K]$V, #any_int capacity: int, loc := #caller_location) -> Allocator_Error { + return __dynamic_map_reserve((^Raw_Map)(m), map_info(T), uint(capacity), loc) + } + + // Shrinks the capacity of a map down to the current length. + // + // Note: Prefer the procedure group `shrink` + @builtin + shrink_map :: proc(m: ^$T/map[$K]$V, loc := #caller_location) -> (did_shrink: bool, err: Allocator_Error) { + if m != nil { + return map_shrink_dynamic((^Raw_Map)(m), map_info(T), loc) + } return } - map_clear_dynamic((^Raw_Map)(m), map_info(T)) -} -// `reserve_map` will try to reserve memory of a passed map to the requested element count (setting the `cap`). -// -// Note: Prefer the procedure group `reserve` -@builtin -reserve_map :: proc(m: ^$T/map[$K]$V, #any_int capacity: int, loc := #caller_location) -> Allocator_Error { - return __dynamic_map_reserve((^Raw_Map)(m), map_info(T), uint(capacity), loc) -} - -// Shrinks the capacity of a map down to the current length. -// -// Note: Prefer the procedure group `shrink` -@builtin -shrink_map :: proc(m: ^$T/map[$K]$V, loc := #caller_location) -> (did_shrink: bool, err: Allocator_Error) { - if m != nil { - return map_shrink_dynamic((^Raw_Map)(m), map_info(T), loc) - } - return -} - -// The delete_key built-in procedure deletes the element with the specified key (m[key]) from the map. -// If m is nil, or there is no such element, this procedure is a no-op -@builtin -delete_key :: proc(m: ^$T/map[$K]$V, key: K) -> (deleted_key: K, deleted_value: V) { - if m != nil { - key := key - old_k, old_v, ok := map_erase_dynamic((^Raw_Map)(m), map_info(T), uintptr(&key)) - if ok { - deleted_key = (^K)(old_k)^ - deleted_value = (^V)(old_v)^ + // The delete_key built-in procedure deletes the element with the specified key (m[key]) from the map. + // If m is nil, or there is no such element, this procedure is a no-op + // It is safe to use `delete_key` while iterating a map. + // But if you iterate across a map and insert a new key, it could resize which means you are not iterating across all of the elements. + @builtin + delete_key :: proc(m: ^$T/map[$K]$V, key: K) -> (deleted_key: K, deleted_value: V) { + if m != nil { + key := key + old_k, old_v, ok := map_erase_dynamic((^Raw_Map)(m), map_info(T), uintptr(&key)) + if ok { + deleted_key = (^K)(old_k)^ + deleted_value = (^V)(old_v)^ + } } + return } - return } -_append_elem :: #force_no_inline proc(array: ^Raw_Dynamic_Array, size_of_elem, align_of_elem: int, arg_ptr: rawptr, should_zero: bool, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +_append_elem :: #force_no_inline proc(array: ^Raw_Dynamic_Array, size_of_elem, align_of_elem: int, arg_ptr: rawptr, should_zero: bool, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { if array == nil { return } if array.cap < array.len+1 { // Same behavior as _append_elems but there's only one arg, so we always just add DEFAULT_DYNAMIC_ARRAY_CAPACITY. - cap := 2 * array.cap + DEFAULT_DYNAMIC_ARRAY_CAPACITY + cap := max(2 * array.cap, DEFAULT_DYNAMIC_ARRAY_CAPACITY) // do not 'or_return' here as it could be a partial success - err = _reserve_dynamic_array(array, size_of_elem, align_of_elem, cap, should_zero, loc) + err = _reserve_dynamic_array_unsafe(array, size_of_elem, align_of_elem, cap, should_zero, loc) } if array.cap-array.len > 0 { data := ([^]byte)(array.data) @@ -724,20 +734,71 @@ _append_elem :: #force_no_inline proc(array: ^Raw_Dynamic_Array, size_of_elem, a data = data[array.len*size_of_elem:] intrinsics.mem_copy_non_overlapping(data, arg_ptr, size_of_elem) array.len += 1 - n = 1 + num_appended = 1 } return } +_append_elem_ptr :: #force_no_inline proc(array: ^Raw_Dynamic_Array, arg: rawptr, should_zero: bool, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { + if array == nil { + return + } + + if array.cap < array.len+1 { + // Same behavior as _append_elems but there's only one arg, so we always just add DEFAULT_DYNAMIC_ARRAY_CAPACITY. + cap := max(2 * array.cap, DEFAULT_DYNAMIC_ARRAY_CAPACITY) + + // do not 'or_return' here as it could be a partial success + err = _reserve_dynamic_array_unsafe(array, size_of(rawptr), align_of(rawptr), cap, should_zero, loc) + } + if array.cap-array.len > 0 { + data := ([^]rawptr)(array.data) + assert(data != nil, loc=loc) + data[array.len] = arg + array.len += 1 + num_appended = 1 + } + return +} + + // `append_elem` appends an element to the end of a dynamic array. @builtin -append_elem :: proc(array: ^$T/[dynamic]$E, #no_broadcast arg: E, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +append_elem :: proc(array: ^$T/[dynamic]$E, #no_broadcast arg: E, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { when size_of(E) == 0 { + if array == nil { + return + } (^Raw_Dynamic_Array)(array).len += 1 return 1, nil - } else { + } else when intrinsics.type_is_internally_pointer_like(E) { + return _append_elem_ptr((^Raw_Dynamic_Array)(array), rawptr(arg), should_zero=true, loc=loc) + } else when ODIN_OPTIMIZATION_MODE <= .Size { arg := arg - return _append_elem((^Raw_Dynamic_Array)(array), size_of(E), align_of(E), &arg, true, loc=loc) + return _append_elem((^Raw_Dynamic_Array)(array), size_of(E), align_of(E), &arg, should_zero=true, loc=loc) + } else { + if array == nil { + return + } + arg := arg + arr := (^Raw_Dynamic_Array)(array) + if arr.cap < arr.len+1 { + // Same behavior as _append_elems but there's only one arg, so we always just add DEFAULT_DYNAMIC_ARRAY_CAPACITY. + cap := max(2 * arr.cap, DEFAULT_DYNAMIC_ARRAY_CAPACITY) + + // do not 'or_return' here as it could be a partial success + err = _reserve_dynamic_array_unsafe(arr, size_of(E), align_of(E), cap, should_zero=true, loc=loc) + } + if arr.cap-arr.len > 0 { + // NOTE(bill, 2026-06-19): When this is in the hot path with -o:speed or -o:aggressive enabled, + // this code path cannot rely on type erasure and `mem_copy_non_overlapping`. + // So directly inlining the call and storing the argument like this helps the optimize a lot + assert(arr.data != nil, loc=loc) + ([^]E)(arr.data)[arr.len] = arg + arr.len += 1 + num_appended = 1 + } + return } } @@ -745,17 +806,42 @@ append_elem :: proc(array: ^$T/[dynamic]$E, #no_broadcast arg: E, loc := #caller // // Note: Prefer using the procedure group `non_zero_append @builtin -non_zero_append_elem :: proc(array: ^$T/[dynamic]$E, #no_broadcast arg: E, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +non_zero_append_elem :: proc(array: ^$T/[dynamic]$E, #no_broadcast arg: E, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { when size_of(E) == 0 { (^Raw_Dynamic_Array)(array).len += 1 return 1, nil - } else { + } else when intrinsics.type_is_internally_pointer_like(E) { + return _append_elem_ptr((^Raw_Dynamic_Array)(array), rawptr(arg), should_zero=false, loc=loc) + } else when ODIN_OPTIMIZATION_MODE <= .Size { arg := arg - return _append_elem((^Raw_Dynamic_Array)(array), size_of(E), align_of(E), &arg, false, loc=loc) + return _append_elem((^Raw_Dynamic_Array)(array), size_of(E), align_of(E), &arg, should_zero=false, loc=loc) + } else { + if array == nil { + return + } + arg := arg + arr := (^Raw_Dynamic_Array)(array) + if arr.cap < arr.len+1 { + // Same behavior as _append_elems but there's only one arg, so we always just add DEFAULT_DYNAMIC_ARRAY_CAPACITY. + cap := max(2 * arr.cap, DEFAULT_DYNAMIC_ARRAY_CAPACITY) + + // do not 'or_return' here as it could be a partial success + err = _reserve_dynamic_array_unsafe(arr, size_of(E), align_of(E), cap, should_zero=false, loc=loc) + } + if arr.cap-arr.len > 0 { + // NOTE(bill, 2026-06-19): When this is in the hot path with -o:speed or -o:aggressive enabled, + // this code path cannot rely on type erasure and `mem_copy_non_overlapping`. + // So directly inlining the call and storing the argument like this helps the optimize a lot + assert(arr.data != nil, loc=loc) + ([^]E)(arr.data)[arr.len] = arg + arr.len += 1 + num_appended = 1 + } + return } } -_append_elems :: #force_no_inline proc(array: ^Raw_Dynamic_Array, size_of_elem, align_of_elem: int, should_zero: bool, loc := #caller_location, args: rawptr, arg_len: int) -> (n: int, err: Allocator_Error) #optional_allocator_error { +_append_elems :: #force_no_inline proc(array: ^Raw_Dynamic_Array, size_of_elem, align_of_elem: int, should_zero: bool, loc := #caller_location, args: rawptr, arg_len: int) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { if array == nil { return 0, nil } @@ -786,7 +872,7 @@ _append_elems :: #force_no_inline proc(array: ^Raw_Dynamic_Array, size_of_elem, // // Note: Prefer using the procedure group `append`. @builtin -append_elems :: proc(array: ^$T/[dynamic]$E, #no_broadcast args: ..E, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +append_elems :: proc(array: ^$T/[dynamic]$E, #no_broadcast args: ..E, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { when size_of(E) == 0 { a := (^Raw_Dynamic_Array)(array) a.len += len(args) @@ -800,7 +886,7 @@ append_elems :: proc(array: ^$T/[dynamic]$E, #no_broadcast args: ..E, loc := #ca // // Note: Prefer using the procedure group `non_zero_append @builtin -non_zero_append_elems :: proc(array: ^$T/[dynamic]$E, #no_broadcast args: ..E, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +non_zero_append_elems :: proc(array: ^$T/[dynamic]$E, #no_broadcast args: ..E, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { when size_of(E) == 0 { a := (^Raw_Dynamic_Array)(array) a.len += len(args) @@ -811,7 +897,7 @@ non_zero_append_elems :: proc(array: ^$T/[dynamic]$E, #no_broadcast args: ..E, l } // The append_string built-in procedure appends a string to the end of a [dynamic]u8 like type -_append_elem_string :: proc(array: ^$T/[dynamic]$E/u8, arg: $A/string, should_zero: bool, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +_append_elem_string :: proc(array: ^$T/[dynamic]$E/u8, arg: $A/string, should_zero: bool, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { return _append_elems((^Raw_Dynamic_Array)(array), 1, 1, should_zero, loc, raw_data(arg), len(arg)) } @@ -819,14 +905,14 @@ _append_elem_string :: proc(array: ^$T/[dynamic]$E/u8, arg: $A/string, should_ze // // Note: Prefer using the procedure group `append`. @builtin -append_elem_string :: proc(array: ^$T/[dynamic]$E/u8, arg: $A/string, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +append_elem_string :: proc(array: ^$T/[dynamic]$E/u8, arg: $A/string, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { return _append_elem_string(array, arg, true, loc) } // `non_zero_append_elem_string` appends a string to the end of a dynamic array of bytes, without zeroing any reserved memory // // Note: Prefer using the procedure group `non_zero_append`. @builtin -non_zero_append_elem_string :: proc(array: ^$T/[dynamic]$E/u8, arg: $A/string, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +non_zero_append_elem_string :: proc(array: ^$T/[dynamic]$E/u8, arg: $A/string, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { return _append_elem_string(array, arg, false, loc) } @@ -834,7 +920,7 @@ non_zero_append_elem_string :: proc(array: ^$T/[dynamic]$E/u8, arg: $A/string, l // // Note: Prefer using the procedure group `non_zero_append`. @builtin -non_zero_append_elem_fixed_capacity_string :: proc "contextless" (array: ^$T/[dynamic; $N]$E/u8, arg: $A/string) -> (n: int) { +non_zero_append_elem_fixed_capacity_string :: proc "contextless" (array: ^$T/[dynamic; $N]$E/u8, arg: $A/string) -> (num_appended: int) { return append_fixed_capacity_elem(array, transmute([]byte)arg) } @@ -844,11 +930,11 @@ non_zero_append_elem_fixed_capacity_string :: proc "contextless" (array: ^$T/[dy // // Note: Prefer using the procedure group `append`. @builtin -append_string :: proc(array: ^$T/[dynamic]$E/u8, args: ..string, loc := #caller_location) -> (n: int, err: Allocator_Error) #optional_allocator_error { +append_string :: proc(array: ^$T/[dynamic]$E/u8, args: ..string, loc := #caller_location) -> (num_appended: int, err: Allocator_Error) #optional_allocator_error { n_arg: int for arg in args { n_arg, err = append(array, ..transmute([]E)(arg), loc=loc) - n += n_arg + num_appended += n_arg if err != nil { return } @@ -859,7 +945,7 @@ append_string :: proc(array: ^$T/[dynamic]$E/u8, args: ..string, loc := #caller_ // `append_fixed_capacity_elem` appends an element to the end of a fixed capacity dynamic array. Returns 0 on failure @builtin -append_fixed_capacity_elem :: proc "contextless" (array: ^$T/[dynamic; $N]$E, #no_broadcast arg: E) -> (n: int) { +append_fixed_capacity_elem :: proc "contextless" (array: ^$T/[dynamic; $N]$E, #no_broadcast arg: E) -> (num_appended: int) { Raw :: Raw_Fixed_Capacity_Dynamic_Array(N, E) if (^Raw)(array).len >= N { @@ -876,29 +962,29 @@ append_fixed_capacity_elem :: proc "contextless" (array: ^$T/[dynamic; $N]$E, #n // `append_fixed_capacity_elem` appends an element to the end of a fixed capacity dynamic array. Returns 0 on failure @builtin -append_fixed_capacity_elems :: proc "contextless" (array: ^$T/[dynamic; $N]$E, #no_broadcast args: ..E) -> (n: int) { +append_fixed_capacity_elems :: proc "contextless" (array: ^$T/[dynamic; $N]$E, #no_broadcast args: ..E) -> (num_appended: int) { Raw :: Raw_Fixed_Capacity_Dynamic_Array(N, E) raw := (^Raw)(array) - n = min(N - len(array), len(args)) + num_appended = min(N - len(array), len(args)) #no_bounds_check when size_of(E) != 0 { - intrinsics.mem_copy(&raw.data[raw.len], raw_data(args), n*size_of(E)) + intrinsics.mem_copy(&raw.data[raw.len], raw_data(args), num_appended*size_of(E)) } - raw.len += n - return n + raw.len += num_appended + return num_appended } // The append_fixed_capacity_string built-in procedure appends multiple strings to the end of a [dynamic]u8 like type // // Note: Prefer using the procedure group `append`. @builtin -append_fixed_capacity_string :: proc "contextless" (array: ^$T/[dynamic; $N]$E/u8, args: ..string) -> (n: int) { +append_fixed_capacity_string :: proc "contextless" (array: ^$T/[dynamic; $N]$E/u8, args: ..string) -> (num_appended: int) { n_arg: int for arg in args { n_arg = append_fixed_capacity_elems(array, ..transmute([]E)(arg)) - n += n_arg + num_appended += n_arg if n_arg < len(arg) { return } @@ -1191,7 +1277,7 @@ assign_at_elem_fixed_capacity_dynamic_array :: proc "contextless" (array: ^$T/[d array[index] = arg ok = true } else { - resize(array, index+1, loc) or_return + resize(array, index+1) or_return array[index] = arg ok = true } @@ -1210,7 +1296,7 @@ assign_at_elems_fixed_capacity_dynamic_array :: proc "contextless" (array: ^$T/[ copy(array[index:], args) ok = true } else { - resize(array, new_size, loc) or_return + resize(array, new_size) or_return copy(array[index:], args) ok = true } @@ -1228,7 +1314,7 @@ assign_at_elem_string_fixed_capacity_dynamic_array :: proc "contextless" (array: copy(array[index:], arg) ok = true } else { - resize(array, new_size, loc) or_return + resize(array, new_size) or_return copy(array[index:], arg) ok = true } @@ -1309,6 +1395,35 @@ _reserve_dynamic_array :: #force_no_inline proc(a: ^Raw_Dynamic_Array, size_of_e return nil } +_reserve_dynamic_array_unsafe :: #force_no_inline proc(a: ^Raw_Dynamic_Array, size_of_elem, align_of_elem: int, capacity: int, should_zero: bool, loc := #caller_location) -> Allocator_Error { + if capacity <= a.cap { + return nil + } + + if a.allocator.procedure == nil { + a.allocator = context.allocator + assert(a.allocator.procedure != nil) + } + + old_size := a.cap * size_of_elem + new_size := capacity * size_of_elem + allocator := a.allocator + + new_data: []byte + if should_zero { + new_data = mem_resize(a.data, old_size, new_size, align_of_elem, allocator, loc) or_return + } else { + new_data = non_zero_mem_resize(a.data, old_size, new_size, align_of_elem, allocator, loc) or_return + } + if new_data == nil && new_size > 0 { + return .Out_Of_Memory + } + + a.data = raw_data(new_data) + a.cap = capacity + return nil +} + // `reserve_dynamic_array` will try to reserve memory of a passed dynamic array or map to the requested element count (setting the `cap`). // // When a memory resize allocation is required, the memory will be asked to be zeroed (i.e. it calls `mem_resize`). @@ -1468,52 +1583,54 @@ _shrink_dynamic_array :: proc(a: ^Raw_Dynamic_Array, size_of_elem, align_of_elem return true, nil } -@builtin -map_insert :: proc(m: ^$T/map[$K]$V, key: K, value: V, loc := #caller_location) -> (ptr: ^V) { - key, value := key, value - return (^V)(__dynamic_map_set_without_hash((^Raw_Map)(m), map_info(T), rawptr(&key), rawptr(&value), loc)) -} - -// Explicitly inserts a key and value into a map `m`, the same as `map_insert`, but the return values differ. -// - `prev_key` will return the previous pointer of a key if it exists, check `found_previous` if was previously found -// - `value_ptr` will return the pointer of the memory where the insertion happens, and `nil` if the map failed to resize -// - `found_previous` will be true a previous key was found -@(builtin, require_results) -map_upsert :: proc(m: ^$T/map[$K]$V, key: K, value: V, loc := #caller_location) -> (prev_key: K, value_ptr: ^V, found_previous: bool) { - key, value := key, value - kp, vp := __dynamic_map_set_extra_without_hash((^Raw_Map)(m), map_info(T), rawptr(&key), rawptr(&value), loc) - if kp != nil { - prev_key = (^K)(kp)^ - found_previous = true +when MAP_ENABLED { + @builtin + map_insert :: proc(m: ^$T/map[$K]$V, key: K, value: V, loc := #caller_location) -> (ptr: ^V) { + key, value := key, value + return (^V)(__dynamic_map_set_without_hash((^Raw_Map)(m), map_info(T), rawptr(&key), rawptr(&value), loc)) } - value_ptr = (^V)(vp) - return -} -/* -Retrieves a pointer to the key and value for a possibly just inserted entry into the map. + // Explicitly inserts a key and value into a map `m`, the same as `map_insert`, but the return values differ. + // - `prev_key` will return the previous pointer of a key if it exists, check `found_previous` if was previously found + // - `value_ptr` will return the pointer of the memory where the insertion happens, and `nil` if the map failed to resize + // - `found_previous` will be true a previous key was found + @(builtin, require_results) + map_upsert :: proc(m: ^$T/map[$K]$V, key: K, value: V, loc := #caller_location) -> (prev_key: K, value_ptr: ^V, found_previous: bool) { + key, value := key, value + kp, vp := __dynamic_map_set_extra_without_hash((^Raw_Map)(m), map_info(T), rawptr(&key), rawptr(&value), loc) + if kp != nil { + prev_key = (^K)(kp)^ + found_previous = true + } + value_ptr = (^V)(vp) + return + } -If the `key` was not in the map `m`, an entry is inserted with the zero value and `just_inserted` will be `true`. -Otherwise the existing entry is left untouched and pointers to its key and value are returned. + /* + Retrieves a pointer to the key and value for a possibly just inserted entry into the map. -If the map has to grow in order to insert the entry and the allocation fails, `err` is set and returned. + If the `key` was not in the map `m`, an entry is inserted with the zero value and `just_inserted` will be `true`. + Otherwise the existing entry is left untouched and pointers to its key and value are returned. -If `err` is `nil`, `key_ptr` and `value_ptr` are valid pointers and will not be `nil`. + If the map has to grow in order to insert the entry and the allocation fails, `err` is set and returned. -WARN: User modification of the key pointed at by `key_ptr` should only be done if the new key is equal to (in hash) the old key. -If that is not the case you will corrupt the map. -*/ -@(builtin, require_results) -map_entry :: proc(m: ^$T/map[$K]$V, key: K, loc := #caller_location) -> (key_ptr: ^K, value_ptr: ^V, just_inserted: bool, err: Allocator_Error) { - key := key - zero: V + If `err` is `nil`, `key_ptr` and `value_ptr` are valid pointers and will not be `nil`. - _key_ptr, _value_ptr: rawptr - _key_ptr, _value_ptr, just_inserted, err = __dynamic_map_entry((^Raw_Map)(m), map_info(T), &key, &zero, loc) + WARN: User modification of the key pointed at by `key_ptr` should only be done if the new key is equal to (in hash) the old key. + If that is not the case you will corrupt the map. + */ + @(builtin, require_results) + map_entry :: proc(m: ^$T/map[$K]$V, key: K, loc := #caller_location) -> (key_ptr: ^K, value_ptr: ^V, just_inserted: bool, err: Allocator_Error) { + key := key + zero: V - key_ptr = (^K)(_key_ptr) - value_ptr = (^V)(_value_ptr) - return + _key_ptr, _value_ptr: rawptr + _key_ptr, _value_ptr, just_inserted, err = __dynamic_map_entry((^Raw_Map)(m), map_info(T), &key, &zero, loc) + + key_ptr = (^K)(_key_ptr) + value_ptr = (^V)(_value_ptr) + return + } } @@ -1622,14 +1739,14 @@ ensure_contextless :: proc "contextless" (condition: bool, message := #caller_ex } } -// Panics the program with a message to indicate something has yet to be implemented. +// Panics the program with a message. // This uses the `default_assertion_contextless_failure_proc` to assert. @builtin panic_contextless :: proc "contextless" (message: string, loc := #caller_location) -> ! { default_assertion_contextless_failure_proc("panic", message, loc) } -// Panics the program with a message. +// Panics the program with a message to indicate something has yet to be implemented. // This uses the `default_assertion_contextless_failure_proc` to assert. @builtin unimplemented_contextless :: proc "contextless" (message := "", loc := #caller_location) -> ! { diff --git a/base/runtime/default_temp_allocator_arena.odin b/base/runtime/default_temp_allocator_arena.odin index eafaf9fbe..f1d09792a 100644 --- a/base/runtime/default_temp_allocator_arena.odin +++ b/base/runtime/default_temp_allocator_arena.odin @@ -97,15 +97,6 @@ alloc_from_memory_block :: proc(block: ^Memory_Block, min_size, alignment: uint) @(require_results) arena_alloc :: proc(arena: ^Arena, size, alignment: uint, loc := #caller_location) -> (data: []byte, err: Allocator_Error) { - align_forward_uint :: proc "contextless" (ptr, align: uint) -> uint { - p := ptr - modulo := p & (align-1) - if modulo != 0 { - p += align - modulo - } - return p - } - assert(alignment & (alignment-1) == 0, "non-power of two alignment", loc) size := size diff --git a/base/runtime/default_temporary_allocator.odin b/base/runtime/default_temporary_allocator.odin index 2017570bb..857463912 100644 --- a/base/runtime/default_temporary_allocator.odin +++ b/base/runtime/default_temporary_allocator.odin @@ -7,7 +7,7 @@ when NO_DEFAULT_TEMP_ALLOCATOR { // `Default_Temp_Allocator` is a `nil_allocator` when `NO_DEFAULT_TEMP_ALLOCATOR` is `true`. Default_Temp_Allocator :: struct {} - default_temp_allocator_init :: proc(s: ^Default_Temp_Allocator, size: int, backing_allocator := context.allocator) {} + default_temp_allocator_init :: proc(s: ^Default_Temp_Allocator, size: int, backing_allocator: Allocator) {} default_temp_allocator_destroy :: proc "contextless" (s: ^Default_Temp_Allocator) {} @@ -30,7 +30,7 @@ when NO_DEFAULT_TEMP_ALLOCATOR { arena: Arena, } - default_temp_allocator_init :: proc(s: ^Default_Temp_Allocator, size: int, backing_allocator := context.allocator) { + default_temp_allocator_init :: proc(s: ^Default_Temp_Allocator, size: int, backing_allocator: Allocator) { _ = arena_init(&s.arena, uint(size), backing_allocator) } diff --git a/base/runtime/dynamic_map_internal.odin b/base/runtime/dynamic_map_internal.odin index 73a152845..c9b3668c3 100644 --- a/base/runtime/dynamic_map_internal.odin +++ b/base/runtime/dynamic_map_internal.odin @@ -1,3 +1,4 @@ +#+build !bedrock package runtime import "base:intrinsics" @@ -47,60 +48,6 @@ MAP_MIN_LOG2_CAPACITY :: 3 // 8 elements // Has to be less than 100% though. #assert(MAP_LOAD_FACTOR < 100) -// This is safe to change. The log2 size of a cache-line. At minimum it has to -// be six though. Higher cache line sizes are permitted. -MAP_CACHE_LINE_LOG2 :: 6 - -// The size of a cache-line. -MAP_CACHE_LINE_SIZE :: 1 << MAP_CACHE_LINE_LOG2 - -// The minimum cache-line size allowed by this implementation is 64 bytes since -// we need 6 bits in the base pointer to store the integer log2 capacity, which -// at maximum is 63. Odin uses signed integers to represent length and capacity, -// so only 63 bits are needed in the maximum case. -#assert(MAP_CACHE_LINE_SIZE >= 64) - -// Map_Cell type that packs multiple T in such a way to ensure that each T stays -// aligned by align_of(T) and such that align_of(Map_Cell(T)) % MAP_CACHE_LINE_SIZE == 0 -// -// This means a value of type T will never straddle a cache-line. -// -// When multiple Ts can fit in a single cache-line the data array will have more -// than one element. When it cannot, the data array will have one element and -// an array of Map_Cell(T) will be padded to stay a multiple of MAP_CACHE_LINE_SIZE. -// -// We rely on the type system to do all the arithmetic and padding for us here. -// -// The usual array[index] indexing for []T backed by a []Map_Cell(T) becomes a bit -// more involved as there now may be internal padding. The indexing now becomes -// -// N :: len(Map_Cell(T){}.data) -// i := index / N -// j := index % N -// cell[i].data[j] -// -// However, since len(Map_Cell(T){}.data) is a compile-time constant, there are some -// optimizations we can do to eliminate the need for any divisions as N will -// be bounded by [1, 64). -// -// In the optimal case, len(Map_Cell(T){}.data) = 1 so the cell array can be treated -// as a regular array of T, which is the case for hashes. -Map_Cell :: struct($T: typeid) #align(MAP_CACHE_LINE_SIZE) { - data: [MAP_CACHE_LINE_SIZE / size_of(T) when 0 < size_of(T) && size_of(T) < MAP_CACHE_LINE_SIZE else 1]T, -} - -// So we can operate on a cell data structure at runtime without any type -// information, we have a simple table that stores some traits about the cell. -// -// 32-bytes on 64-bit -// 16-bytes on 32-bit -Map_Cell_Info :: struct { - size_of_type: uintptr, // 8-bytes on 64-bit, 4-bytes on 32-bits - align_of_type: uintptr, // 8-bytes on 64-bit, 4-bytes on 32-bits - size_of_cell: uintptr, // 8-bytes on 64-bit, 4-bytes on 32-bits - elements_per_cell: uintptr, // 8-bytes on 64-bit, 4-bytes on 32-bits -} - // map_cell_info :: proc "contextless" ($T: typeid) -> ^Map_Cell_Info {...} map_cell_info :: intrinsics.type_map_cell_info @@ -226,8 +173,6 @@ map_data :: #force_inline proc "contextless" (m: Raw_Map) -> uintptr { } -Map_Hash :: uintptr - TOMBSTONE_MASK :: 1<<(size_of(Map_Hash)*8 - 1) // Procedure to check if a slot is empty for a given hash. This is represented @@ -288,23 +233,6 @@ map_probe_distance :: #force_inline proc "contextless" (m: Raw_Map, hash: Map_Ha return (slot - uintptr(hash)) & (capacity - 1) // NOTE(bill): this is equivalent to the above, but less operations } -// When working with the type-erased structure at runtime we need information -// about the map to make working with it possible. This info structure stores -// that. -// -// `Map_Info` and `Map_Cell_Info` are read only data structures and cannot be -// modified after creation -// -// 32-bytes on 64-bit -// 16-bytes on 32-bit -Map_Info :: struct { - ks: ^Map_Cell_Info, // 8-bytes on 64-bit, 4-bytes on 32-bit - vs: ^Map_Cell_Info, // 8-bytes on 64-bit, 4-bytes on 32-bit - key_hasher: proc "contextless" (key: rawptr, seed: Map_Hash) -> Map_Hash, // 8-bytes on 64-bit, 4-bytes on 32-bit - key_equal: proc "contextless" (lhs, rhs: rawptr) -> bool, // 8-bytes on 64-bit, 4-bytes on 32-bit -} - - // The Map_Info structure is basically a pseudo-table of information for a given K and V pair. // map_info :: proc "contextless" ($T: typeid/map[$K]$V) -> ^Map_Info {...} map_info :: intrinsics.type_map_info @@ -905,22 +833,22 @@ __dynamic_map_check_grow :: proc "odin" (#no_alias m: ^Raw_Map, #no_alias info: return nil, false } -__dynamic_map_set_without_hash :: proc "odin" (#no_alias m: ^Raw_Map, #no_alias info: ^Map_Info, key, value: rawptr, loc := #caller_location) -> rawptr { +__dynamic_map_set_without_hash :: proc "odin" (#no_alias m: ^Raw_Map, #no_alias info: ^Map_Info, key, value: rawptr, loc := #caller_location) -> (value_ptr: rawptr, err: Allocator_Error) #optional_allocator_error { return __dynamic_map_set(m, info, info.key_hasher(key, map_seed(m^)), key, value, loc) } // IMPORTANT: USED WITHIN THE COMPILER -__dynamic_map_set :: proc "odin" (#no_alias m: ^Raw_Map, #no_alias info: ^Map_Info, hash: Map_Hash, key, value: rawptr, loc := #caller_location) -> rawptr { +__dynamic_map_set :: proc "odin" (#no_alias m: ^Raw_Map, #no_alias info: ^Map_Info, hash: Map_Hash, key, value: rawptr, loc := #caller_location) -> (value_ptr: rawptr, err: Allocator_Error) #optional_allocator_error { if found := __dynamic_map_get(m, info, hash, key); found != nil { intrinsics.mem_copy_non_overlapping(found, value, info.vs.size_of_type) - return found + return found, nil } hash := hash - err, has_grown := __dynamic_map_check_grow(m, info, loc) - if err != nil { - return nil + err_grow, has_grown := __dynamic_map_check_grow(m, info, loc) + if err_grow != nil { + return nil, err_grow } if has_grown { hash = info.key_hasher(key, map_seed(m^)) @@ -930,7 +858,7 @@ __dynamic_map_set :: proc "odin" (#no_alias m: ^Raw_Map, #no_alias info: ^Map_In if result != 0 { m.len += 1 } - return rawptr(result) + return rawptr(result), nil } __dynamic_map_set_extra_without_hash :: proc "odin" (#no_alias m: ^Raw_Map, #no_alias info: ^Map_Info, key, value: rawptr, loc := #caller_location) -> (prev_key_ptr, value_ptr: rawptr) { return __dynamic_map_set_extra(m, info, info.key_hasher(key, map_seed(m^)), key, value, loc) @@ -1060,4 +988,4 @@ default_hasher_quaternion256 :: proc "contextless" (x, y, z, w: f64, seed: uintp seed = default_hasher_f64(z, seed) seed = default_hasher_f64(w, seed) return seed -} \ No newline at end of file +} diff --git a/base/runtime/entry_unix.odin b/base/runtime/entry_unix.odin index e2223d5d6..f02ead242 100644 --- a/base/runtime/entry_unix.odin +++ b/base/runtime/entry_unix.odin @@ -1,5 +1,5 @@ #+private -#+build linux, darwin, freebsd, openbsd, netbsd, haiku +#+build linux, darwin, freebsd, openbsd, netbsd #+no-instrumentation package runtime @@ -9,13 +9,15 @@ when ODIN_BUILD_MODE == .Dynamic { @(link_name="_odin_entry_point", linkage="strong", require/*, link_section=".init"*/) _odin_entry_point :: proc "c" () { context = default_context() - #force_no_inline _startup_runtime() + when !ODIN_BEDROCK { #force_no_inline _startup_runtime() } intrinsics.__entry_point() } @(link_name="_odin_exit_point", linkage="strong", require/*, link_section=".fini"*/) _odin_exit_point :: proc "c" () { context = default_context() - #force_no_inline _cleanup_runtime() + when !ODIN_BEDROCK { + #force_no_inline _cleanup_runtime() + } } @(link_name="main", linkage="strong", require) main :: proc "c" (argc: i32, argv: [^]cstring) -> i32 { @@ -42,9 +44,9 @@ when ODIN_BUILD_MODE == .Dynamic { _start_odin :: proc "c" (argc: i32, argv: [^]cstring) -> ! { args__ = argv[:argc] context = default_context() - #force_no_inline _startup_runtime() + when !ODIN_BEDROCK { #force_no_inline _startup_runtime() } intrinsics.__entry_point() - #force_no_inline _cleanup_runtime() + when !ODIN_BEDROCK { #force_no_inline _cleanup_runtime() } intrinsics.syscall(SYS_exit, 0) unreachable() } @@ -53,9 +55,9 @@ when ODIN_BUILD_MODE == .Dynamic { main :: proc "c" (argc: i32, argv: [^]cstring) -> i32 { args__ = argv[:argc] context = default_context() - #force_no_inline _startup_runtime() + when !ODIN_BEDROCK { #force_no_inline _startup_runtime() } intrinsics.__entry_point() - #force_no_inline _cleanup_runtime() + when !ODIN_BEDROCK { #force_no_inline _cleanup_runtime() } return 0 } } diff --git a/base/runtime/entry_windows.odin b/base/runtime/entry_windows.odin index dc8e9b82c..d3b38bc9c 100644 --- a/base/runtime/entry_windows.odin +++ b/base/runtime/entry_windows.odin @@ -16,10 +16,10 @@ when ODIN_BUILD_MODE == .Dynamic { switch dll_forward_reason { case .Process_Attach: - #force_no_inline _startup_runtime() + when !ODIN_BEDROCK { #force_no_inline _startup_runtime() } intrinsics.__entry_point() case .Process_Detach: - #force_no_inline _cleanup_runtime() + when !ODIN_BEDROCK { #force_no_inline _cleanup_runtime() } case .Thread_Attach: break case .Thread_Detach: @@ -35,18 +35,18 @@ when ODIN_BUILD_MODE == .Dynamic { main :: proc "c" (argc: i32, argv: [^]cstring) -> i32 { args__ = argv[:argc] context = default_context() - #force_no_inline _startup_runtime() + when !ODIN_BEDROCK { #force_no_inline _startup_runtime() } intrinsics.__entry_point() - #force_no_inline _cleanup_runtime() + when !ODIN_BEDROCK { #force_no_inline _cleanup_runtime() } return 0 } } else when ODIN_NO_CRT { @(link_name="mainCRTStartup", linkage="strong", require) mainCRTStartup :: proc "system" () -> i32 { context = default_context() - #force_no_inline _startup_runtime() + when !ODIN_BEDROCK { #force_no_inline _startup_runtime() } intrinsics.__entry_point() - #force_no_inline _cleanup_runtime() + when !ODIN_BEDROCK { #force_no_inline _cleanup_runtime() } return 0 } } else { @@ -54,9 +54,9 @@ when ODIN_BUILD_MODE == .Dynamic { main :: proc "c" (argc: i32, argv: [^]cstring) -> i32 { args__ = argv[:argc] context = default_context() - #force_no_inline _startup_runtime() + when !ODIN_BEDROCK { #force_no_inline _startup_runtime() } intrinsics.__entry_point() - #force_no_inline _cleanup_runtime() + when !ODIN_BEDROCK { #force_no_inline _cleanup_runtime() } return 0 } } diff --git a/base/runtime/heap_allocator_other.odin b/base/runtime/heap_allocator_other.odin index 507dbf318..63c2ade84 100644 --- a/base/runtime/heap_allocator_other.odin +++ b/base/runtime/heap_allocator_other.odin @@ -1,4 +1,4 @@ -#+build js, wasi, freestanding, essence +#+build js, wasi, freestanding #+private package runtime diff --git a/base/runtime/heap_allocator_unix.odin b/base/runtime/heap_allocator_unix.odin index f6e7ce39e..bd5dc81b6 100644 --- a/base/runtime/heap_allocator_unix.odin +++ b/base/runtime/heap_allocator_unix.odin @@ -1,4 +1,4 @@ -#+build linux, darwin, freebsd, openbsd, netbsd, haiku +#+build linux, darwin, freebsd, openbsd, netbsd #+private package runtime diff --git a/base/runtime/internal.odin b/base/runtime/internal.odin index 12f0dceca..e08d0e01d 100644 --- a/base/runtime/internal.odin +++ b/base/runtime/internal.odin @@ -8,10 +8,10 @@ IS_WASM :: ODIN_ARCH == .wasm32 || ODIN_ARCH == .wasm64p32 @(private) RUNTIME_LINKAGE :: "strong" when ODIN_USE_SEPARATE_MODULES else - "internal" when ODIN_NO_ENTRY_POINT && (ODIN_BUILD_MODE == .Static || ODIN_BUILD_MODE == .Dynamic || ODIN_BUILD_MODE == .Object) else - "strong" when ODIN_BUILD_MODE == .Dynamic else - "strong" when !ODIN_NO_CRT else - "internal" + "internal" when ODIN_NO_ENTRY_POINT && (ODIN_BUILD_MODE == .Static || ODIN_BUILD_MODE == .Dynamic || ODIN_BUILD_MODE == .Object) else + "strong" when ODIN_BUILD_MODE == .Dynamic else + "strong" when !ODIN_NO_CRT else + "internal" RUNTIME_REQUIRE :: false // !ODIN_TILDE @(private) @@ -23,12 +23,42 @@ HAS_HARDWARE_SIMD :: false when (ODIN_ARCH == .amd64 || ODIN_ARCH == .i386) && ! false when (ODIN_ARCH == .riscv64) && !intrinsics.has_target_feature("v") else true +// Size of a native SIMD register for the current compilation target +NATIVE_SIMD_BIT_WIDTH :: + 512 when (ODIN_ARCH == .amd64) && intrinsics.has_target_feature("avx512f") else + 256 when (ODIN_ARCH == .amd64) && (intrinsics.has_target_feature("avx2") || intrinsics.has_target_feature("avx")) else + // Fallback for no hardware SIMD, but also SSE, NEON, SVE, RVV and WASM SIMD128. + 128 @(private) byte_slice :: #force_inline proc "contextless" (data: rawptr, len: int) -> []byte #no_bounds_check { return ([^]byte)(data)[:max(len, 0)] } +@(require_results) +align_forward_uint :: #force_inline proc "odin" (ptr, align: uint) -> uint { + assert(is_power_of_two_uint(align)) + return (ptr + align-1) & ~(align-1) +} + +@(require_results) +align_forward_int :: #force_inline proc "odin" (ptr, align: int) -> int { + assert(is_power_of_two_int(align)) + return int(align_forward_uint(uint(ptr), uint(align))) +} + +@(require_results) +align_forward_uintptr :: #force_inline proc "odin" (ptr, align: uintptr) -> uintptr { + return uintptr(align_forward_uint(uint(ptr), uint(align))) +} + +align_forward :: proc { + align_forward_int, + align_forward_uint, + align_forward_uintptr, +} + +@(require_results) is_power_of_two_int :: #force_inline proc "contextless" (x: int) -> bool { if x <= 0 { return false @@ -36,51 +66,17 @@ is_power_of_two_int :: #force_inline proc "contextless" (x: int) -> bool { return (x & (x-1)) == 0 } -align_forward_int :: #force_inline proc "odin" (ptr, align: int) -> int { - assert(is_power_of_two_int(align)) - - p := ptr - modulo := p & (align-1) - if modulo != 0 { - p += align - modulo - } - return p -} - +@(require_results) is_power_of_two_uint :: #force_inline proc "contextless" (x: uint) -> bool { - if x <= 0 { + if x == 0 { return false } return (x & (x-1)) == 0 } -align_forward_uint :: #force_inline proc "odin" (ptr, align: uint) -> uint { - assert(is_power_of_two_uint(align)) - - p := ptr - modulo := p & (align-1) - if modulo != 0 { - p += align - modulo - } - return p -} - +@(require_results) is_power_of_two_uintptr :: #force_inline proc "contextless" (x: uintptr) -> bool { - if x <= 0 { - return false - } - return (x & (x-1)) == 0 -} - -align_forward_uintptr :: #force_inline proc "odin" (ptr, align: uintptr) -> uintptr { - assert(is_power_of_two_uintptr(align)) - - p := ptr - modulo := p & (align-1) - if modulo != 0 { - p += align - modulo - } - return p + return is_power_of_two_uint(uint(x)) } is_power_of_two :: proc { @@ -89,12 +85,6 @@ is_power_of_two :: proc { is_power_of_two_uintptr, } -align_forward :: proc { - align_forward_int, - align_forward_uint, - align_forward_uintptr, -} - mem_zero :: proc "contextless" (data: rawptr, len: int) -> rawptr { if data == nil { return nil @@ -1174,236 +1164,6 @@ extendhfsf2 :: proc "c" (value: __float16) -> f32 { return gnu_h2f_ieee(value) } - - -@(link_name="__floattidf", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -floattidf :: proc "c" (a: i128) -> f64 { - DBL_MANT_DIG :: 53 - if a == 0 { - return 0.0 - } - a := a - N :: size_of(i128) * 8 - s := a >> (N-1) - a = (a ~ s) - s - sd: = N - intrinsics.count_leading_zeros(a) // number of significant digits - e := i32(sd - 1) // exponent - if sd > DBL_MANT_DIG { - switch sd { - case DBL_MANT_DIG + 1: - a <<= 1 - case DBL_MANT_DIG + 2: - // okay - case: - a = i128(u128(a) >> u128(sd - (DBL_MANT_DIG+2))) | - i128(u128(a) & (~u128(0) >> u128(N + DBL_MANT_DIG+2 - sd)) != 0) - } - - a |= i128((a & 4) != 0) - a += 1 - a >>= 2 - - if a & (i128(1) << DBL_MANT_DIG) != 0 { - a >>= 1 - e += 1 - } - } else { - a <<= u128(DBL_MANT_DIG - sd) & 127 - } - fb: [2]u32 - fb[1] = (u32(s) & 0x80000000) | // sign - (u32(e + 1023) << 20) | // exponent - u32((u64(a) >> 32) & 0x000FFFFF) // mantissa-high - fb[0] = u32(a) // mantissa-low - return transmute(f64)fb -} - - -@(link_name="__floattidf_unsigned", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -floattidf_unsigned :: proc "c" (a: u128) -> f64 { - DBL_MANT_DIG :: 53 - if a == 0 { - return 0.0 - } - a := a - N :: size_of(u128) * 8 - sd: = N - intrinsics.count_leading_zeros(a) // number of significant digits - e := i32(sd - 1) // exponent - if sd > DBL_MANT_DIG { - switch sd { - case DBL_MANT_DIG + 1: - a <<= 1 - case DBL_MANT_DIG + 2: - // okay - case: - a = u128(u128(a) >> u128(sd - (DBL_MANT_DIG+2))) | - u128(u128(a) & (~u128(0) >> u128(N + DBL_MANT_DIG+2 - sd)) != 0) - } - - a |= u128((a & 4) != 0) - a += 1 - a >>= 2 - - if a & (1 << DBL_MANT_DIG) != 0 { - a >>= 1 - e += 1 - } - } else { - a <<= u128(DBL_MANT_DIG - sd) - } - fb: [2]u32 - fb[1] = (0) | // sign - u32((e + 1023) << 20) | // exponent - u32((u64(a) >> 32) & 0x000FFFFF) // mantissa-high - fb[0] = u32(a) // mantissa-low - return transmute(f64)fb -} - - - -@(link_name="__fixunsdfti", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -fixunsdfti :: #force_no_inline proc "c" (a: f64) -> u128 { - // TODO(bill): implement `fixunsdfti` correctly - x := u64(a) - return u128(x) -} - -@(link_name="__fixunsdfdi", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -fixunsdfdi :: #force_no_inline proc "c" (a: f64) -> i128 { - // TODO(bill): implement `fixunsdfdi` correctly - x := i64(a) - return i128(x) -} - - - - -@(link_name="__umodti3", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -umodti3 :: proc "c" (a, b: u128) -> u128 { - r: u128 = --- - _ = udivmod128(a, b, &r) - return r -} - - -@(link_name="__udivmodti4", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -udivmodti4 :: proc "c" (a, b: u128, rem: ^u128) -> u128 { - return udivmod128(a, b, rem) -} - -when !IS_WASM { - @(link_name="__udivti3", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) - udivti3 :: proc "c" (a, b: u128) -> u128 { - return udivmodti4(a, b, nil) - } -} - - -@(link_name="__modti3", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -modti3 :: proc "c" (a, b: i128) -> i128 { - s_a := a >> (128 - 1) - s_b := b >> (128 - 1) - an := (a ~ s_a) - s_a - bn := (b ~ s_b) - s_b - - r: u128 = --- - _ = udivmod128(u128(an), u128(bn), &r) - return (i128(r) ~ s_a) - s_a -} - - -@(link_name="__divmodti4", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -divmodti4 :: proc "c" (a, b: i128, rem: ^i128) -> i128 { - s_a := a >> (128 - 1) // -1 if negative or 0 - s_b := b >> (128 - 1) - an := (a ~ s_a) - s_a // absolute - bn := (b ~ s_b) - s_b - - s_b ~= s_a // quotient sign - u_s_b := u128(s_b) - u_s_a := u128(s_a) - - r: u128 = --- - u := i128((udivmodti4(u128(an), u128(bn), &r) ~ u_s_b) - u_s_b) // negate if negative - rem^ = i128((r ~ u_s_a) - u_s_a) - return u -} - -@(link_name="__divti3", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -divti3 :: proc "c" (a, b: i128) -> i128 { - s_a := a >> (128 - 1) // -1 if negative or 0 - s_b := b >> (128 - 1) - an := (a ~ s_a) - s_a // absolute - bn := (b ~ s_b) - s_b - - s_a ~= s_b // quotient sign - u_s_a := u128(s_a) - - return i128((udivmodti4(u128(an), u128(bn), nil) ~ u_s_a) - u_s_a) // negate if negative -} - - -@(link_name="__fixdfti", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) -fixdfti :: proc "c" (a: u64) -> i128 { - significandBits :: 52 - typeWidth :: (size_of(u64)*8) - exponentBits :: (typeWidth - significandBits - 1) - maxExponent :: ((1 << exponentBits) - 1) - exponentBias :: (maxExponent >> 1) - - implicitBit :: (u64(1) << significandBits) - significandMask :: (implicitBit - 1) - signBit :: (u64(1) << (significandBits + exponentBits)) - absMask :: (signBit - 1) - exponentMask :: (absMask ~ significandMask) - - // Break a into sign, exponent, significand - aRep := a - aAbs := aRep & absMask - sign := i128(-1 if aRep & signBit != 0 else 1) - exponent := u64((aAbs >> significandBits) - exponentBias) - significand := u64((aAbs & significandMask) | implicitBit) - - // If exponent is negative, the result is zero. - if exponent < 0 { - return 0 - } - - // If the value is too large for the integer type, saturate. - if exponent >= size_of(i128) * 8 { - return max(i128) if sign == 1 else min(i128) - } - - // If 0 <= exponent < significandBits, right shift to get the result. - // Otherwise, shift left. - if exponent < significandBits { - return sign * i128(significand >> (significandBits - exponent)) - } else { - return sign * (i128(significand) << (exponent - significandBits)) - } - -} - - - -__write_bits :: proc "contextless" (dst, src: [^]byte, offset: uintptr, size: uintptr) { - for i in 0..>3]) & (1<<(i&7)) != 0) - dst[j>>3] &~= 1<<(j&7) - dst[j>>3] |= the_bit<<(j&7) - } -} - -__read_bits :: proc "contextless" (dst, src: [^]byte, offset: uintptr, size: uintptr) { - for j in 0..>3]) & (1<<(i&7)) != 0) - dst[j>>3] &~= 1<<(j&7) - dst[j>>3] |= the_bit<<(j&7) - } -} - when .Address in ODIN_SANITIZER_FLAGS { foreign { @(require) diff --git a/base/runtime/internal_i128.odin b/base/runtime/internal_i128.odin new file mode 100644 index 000000000..f4a0f90ed --- /dev/null +++ b/base/runtime/internal_i128.odin @@ -0,0 +1,217 @@ +#+vet !cast +#+build !bedrock +package runtime + +import "base:intrinsics" + +@(private="file") +IS_WASM :: ODIN_ARCH == .wasm32 || ODIN_ARCH == .wasm64p32 + +@(link_name="__floattidf", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +floattidf :: proc "c" (a: i128) -> f64 { + DBL_MANT_DIG :: 53 + if a == 0 { + return 0.0 + } + a := a + N :: size_of(i128) * 8 + s := a >> (N-1) + a = (a ~ s) - s + sd: = N - intrinsics.count_leading_zeros(a) // number of significant digits + e := i32(sd - 1) // exponent + if sd > DBL_MANT_DIG { + switch sd { + case DBL_MANT_DIG + 1: + a <<= 1 + case DBL_MANT_DIG + 2: + // okay + case: + a = i128(u128(a) >> u128(sd - (DBL_MANT_DIG+2))) | + i128(u128(a) & (~u128(0) >> u128(N + DBL_MANT_DIG+2 - sd)) != 0) + } + + a |= i128((a & 4) != 0) + a += 1 + a >>= 2 + + if a & (i128(1) << DBL_MANT_DIG) != 0 { + a >>= 1 + e += 1 + } + } else { + a <<= u128(DBL_MANT_DIG - sd) & 127 + } + fb: [2]u32 + fb[1] = (u32(s) & 0x80000000) | // sign + (u32(e + 1023) << 20) | // exponent + u32((u64(a) >> 32) & 0x000FFFFF) // mantissa-high + fb[0] = u32(a) // mantissa-low + return transmute(f64)fb +} + + +@(link_name="__floattidf_unsigned", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +floattidf_unsigned :: proc "c" (a: u128) -> f64 { + DBL_MANT_DIG :: 53 + if a == 0 { + return 0.0 + } + a := a + N :: size_of(u128) * 8 + sd: = N - intrinsics.count_leading_zeros(a) // number of significant digits + e := i32(sd - 1) // exponent + if sd > DBL_MANT_DIG { + switch sd { + case DBL_MANT_DIG + 1: + a <<= 1 + case DBL_MANT_DIG + 2: + // okay + case: + a = u128(u128(a) >> u128(sd - (DBL_MANT_DIG+2))) | + u128(u128(a) & (~u128(0) >> u128(N + DBL_MANT_DIG+2 - sd)) != 0) + } + + a |= u128((a & 4) != 0) + a += 1 + a >>= 2 + + if a & (1 << DBL_MANT_DIG) != 0 { + a >>= 1 + e += 1 + } + } else { + a <<= u128(DBL_MANT_DIG - sd) + } + fb: [2]u32 + fb[1] = (0) | // sign + u32((e + 1023) << 20) | // exponent + u32((u64(a) >> 32) & 0x000FFFFF) // mantissa-high + fb[0] = u32(a) // mantissa-low + return transmute(f64)fb +} + + + +@(link_name="__fixunsdfti", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +fixunsdfti :: #force_no_inline proc "c" (a: f64) -> u128 { + // TODO(bill): implement `fixunsdfti` correctly + x := u64(a) + return u128(x) +} + +@(link_name="__fixunsdfdi", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +fixunsdfdi :: #force_no_inline proc "c" (a: f64) -> i128 { + // TODO(bill): implement `fixunsdfdi` correctly + x := i64(a) + return i128(x) +} + + + + +@(link_name="__umodti3", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +umodti3 :: proc "c" (a, b: u128) -> u128 { + r: u128 = --- + _ = udivmod128(a, b, &r) + return r +} + + +@(link_name="__udivmodti4", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +udivmodti4 :: proc "c" (a, b: u128, rem: ^u128) -> u128 { + return udivmod128(a, b, rem) +} + +when !IS_WASM { + @(link_name="__udivti3", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) + udivti3 :: proc "c" (a, b: u128) -> u128 { + return udivmodti4(a, b, nil) + } +} + + +@(link_name="__modti3", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +modti3 :: proc "c" (a, b: i128) -> i128 { + s_a := a >> (128 - 1) + s_b := b >> (128 - 1) + an := (a ~ s_a) - s_a + bn := (b ~ s_b) - s_b + + r: u128 = --- + _ = udivmod128(u128(an), u128(bn), &r) + return (i128(r) ~ s_a) - s_a +} + + +@(link_name="__divmodti4", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +divmodti4 :: proc "c" (a, b: i128, rem: ^i128) -> i128 { + s_a := a >> (128 - 1) // -1 if negative or 0 + s_b := b >> (128 - 1) + an := (a ~ s_a) - s_a // absolute + bn := (b ~ s_b) - s_b + + s_b ~= s_a // quotient sign + u_s_b := u128(s_b) + u_s_a := u128(s_a) + + r: u128 = --- + u := i128((udivmodti4(u128(an), u128(bn), &r) ~ u_s_b) - u_s_b) // negate if negative + rem^ = i128((r ~ u_s_a) - u_s_a) + return u +} + +@(link_name="__divti3", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +divti3 :: proc "c" (a, b: i128) -> i128 { + s_a := a >> (128 - 1) // -1 if negative or 0 + s_b := b >> (128 - 1) + an := (a ~ s_a) - s_a // absolute + bn := (b ~ s_b) - s_b + + s_a ~= s_b // quotient sign + u_s_a := u128(s_a) + + return i128((udivmodti4(u128(an), u128(bn), nil) ~ u_s_a) - u_s_a) // negate if negative +} + + +@(link_name="__fixdfti", linkage=RUNTIME_LINKAGE, require=RUNTIME_REQUIRE) +fixdfti :: proc "c" (a: u64) -> i128 { + significandBits :: 52 + typeWidth :: (size_of(u64)*8) + exponentBits :: (typeWidth - significandBits - 1) + maxExponent :: ((1 << exponentBits) - 1) + exponentBias :: (maxExponent >> 1) + + implicitBit :: (u64(1) << significandBits) + significandMask :: (implicitBit - 1) + signBit :: (u64(1) << (significandBits + exponentBits)) + absMask :: (signBit - 1) + exponentMask :: (absMask ~ significandMask) + + // Break a into sign, exponent, significand + aRep := a + aAbs := aRep & absMask + sign := i128(-1 if aRep & signBit != 0 else 1) + exponent := u64((aAbs >> significandBits) - exponentBias) + significand := u64((aAbs & significandMask) | implicitBit) + + // If exponent is negative, the result is zero. + if exponent < 0 { + return 0 + } + + // If the value is too large for the integer type, saturate. + if exponent >= size_of(i128) * 8 { + return max(i128) if sign == 1 else min(i128) + } + + // If 0 <= exponent < significandBits, right shift to get the result. + // Otherwise, shift left. + if exponent < significandBits { + return sign * i128(significand >> (significandBits - exponent)) + } else { + return sign * (i128(significand) << (exponent - significandBits)) + } + +} + diff --git a/base/runtime/os_specific_essence.odin b/base/runtime/os_specific_essence.odin deleted file mode 100644 index 9d4b0c968..000000000 --- a/base/runtime/os_specific_essence.odin +++ /dev/null @@ -1,14 +0,0 @@ -#+build essence -#+private -package runtime - -_HAS_RAND_BYTES :: false - -// TODO(bill): reimplement `os.write` -_stderr_write :: proc "contextless" (data: []byte) -> (int, _OS_Errno) { - return 0, -1 -} - -_exit :: proc "contextless" (code: int) -> ! { - trap() -} \ No newline at end of file diff --git a/base/runtime/os_specific_haiku.odin b/base/runtime/os_specific_haiku.odin deleted file mode 100644 index 7e53539a1..000000000 --- a/base/runtime/os_specific_haiku.odin +++ /dev/null @@ -1,33 +0,0 @@ -#+build haiku -#+private -package runtime - -foreign import libc "system:c" - -_HAS_RAND_BYTES :: true - -foreign libc { - @(link_name="write") - _unix_write :: proc(fd: i32, buf: rawptr, size: int) -> int --- - - _errnop :: proc() -> ^i32 --- - - arc4random_buf :: proc(buf: [^]byte, nbytes: uint) --- -} - -_stderr_write :: proc "contextless" (data: []byte) -> (int, _OS_Errno) { - ret := _unix_write(2, raw_data(data), len(data)) - if ret < len(data) { - err := _errnop() - return int(ret), _OS_Errno(err^ if err != nil else 0) - } - return int(ret), 0 -} - -_rand_bytes :: proc "contextless" (dst: []byte) { - arc4random_buf(raw_data(dst), len(dst)) -} - -_exit :: proc "contextless" (code: int) -> ! { - trap() -} \ No newline at end of file diff --git a/base/runtime/udivmod128.odin b/base/runtime/udivmod128.odin index 8cc70df55..107dc605d 100644 --- a/base/runtime/udivmod128.odin +++ b/base/runtime/udivmod128.odin @@ -1,3 +1,4 @@ +#+build !bedrock package runtime import "base:intrinsics" diff --git a/build.bat b/build.bat index b6e8b890d..cc9d6af05 100644 --- a/build.bat +++ b/build.bat @@ -21,7 +21,7 @@ if "%VSCMD_ARG_TGT_ARCH%" neq "x64" ( where /Q git.exe || goto skip_git_hash if not exist .git\ goto skip_git_hash -for /f "tokens=1,2" %%i IN ('git show "--pretty=%%cd %%h" "--date=format:%%Y-%%m-%%d" --no-patch --no-notes HEAD') do ( +for /f "tokens=1,2" %%i IN ('git -c "log.showSignature=false" show "--pretty=%%cd %%h" "--date=format:%%Y-%%m-%%d" --no-patch --no-notes HEAD') do ( set CURR_DATE_TIME=%%i set GIT_SHA=%%j ) @@ -107,16 +107,6 @@ set libs= ^ set odin_res=misc\odin.res set odin_rc=misc\odin.rc -rem DO NOT TOUCH! -rem THIS TILDE STUFF IS FOR DEVELOPMENT ONLY! -set tilde_backend=0 -if %tilde_backend% EQU 1 ( - set libs=%libs% src\tilde\tb.lib - set compiler_defines=%compiler_defines% -DODIN_TILDE_BACKEND -) -rem DO NOT TOUCH! - - set linker_flags= -incremental:no -opt:ref -subsystem:console -MANIFEST:EMBED if %release_mode% EQU 0 ( rem Debug @@ -148,4 +138,4 @@ if %release_mode% EQU 0 echo: & echo Debug compiler built. Note: run "build.bat del *.obj > NUL 2> NUL -:end_of_build \ No newline at end of file +:end_of_build diff --git a/build_odin.sh b/build_odin.sh index ee891e41e..9eab985e9 100755 --- a/build_odin.sh +++ b/build_odin.sh @@ -1,6 +1,10 @@ #!/usr/bin/env sh set -eu +SUPPORTED_LLVM_VERSIONS="22 21 20 19 18 17" +SUGGESTED_LLVM_VERSION="22" +MINIMUM_LLVM_VERSION="17" + : ${CPPFLAGS=} : ${CXXFLAGS=} : ${LDFLAGS=} @@ -13,8 +17,10 @@ OS_ARCH="$(uname -m)" OS_NAME="$(uname -s)" if [ -d ".git" ] && [ -n "$(command -v git)" ]; then - GIT_SHA=$(git show --pretty='%h' --no-patch --no-notes HEAD) - GIT_DATE=$(git show "--pretty=%cd" "--date=format:%Y-%m" --no-patch --no-notes HEAD) + # Counter the user's git config to show the signature in logs. + gitnosig="-c log.showSignature=false" + GIT_SHA=$(git $gitnosig show --pretty='%h' --no-patch --no-notes HEAD) + GIT_DATE=$(git $gitnosig show "--pretty=%cd" "--date=format:%Y-%m" --no-patch --no-notes HEAD) CPPFLAGS="$CPPFLAGS -DGIT_SHA=\"$GIT_SHA\"" else GIT_DATE=$(date +"%Y-%m") @@ -26,8 +32,6 @@ error() { exit 1 } -SUPPORTED_LLVM_VERSIONS="22 21 20 19 18 17 14" - # Brew advises people not to add llvm to their $PATH, so try and use brew to find it. if [ -z "$LLVM_CONFIG" ] && [ -n "$(command -v brew)" ]; then for V in $SUPPORTED_LLVM_VERSIONS; do @@ -61,7 +65,22 @@ if [ -z "$LLVM_CONFIG" ]; then done if [ -z "$LLVM_CONFIG" ]; then - error "No supported llvm-config command found. Set LLVM_CONFIG to proceed." + if [ -f "/etc/os-release" ]; then + . /etc/os-release + case "$ID" in + ubuntu|debian|linuxmint|pop|zorin|kali|elementary|raspbian) + echo "ERROR: No supported llvm-config command found. Set LLVM_CONFIG to proceed." + echo "We suggest installing LLVM $SUGGESTED_LLVM_VERSION from https://apt.llvm.org" + exit 1 + ;; + fedora|centos|rocky|almalinux|rhel|amzn|ol|centos-stream) + echo "ERROR: No supported llvm-config command found. Set LLVM_CONFIG to proceed." + echo "We suggest installing LLVM $SUGGESTED_LLVM_VERSION via COPR, e.g. dnf copr enable fedora-llvm-team/llvm-snapshots" + exit 1 + ;; + esac + fi + error "No supported llvm-config command found. Set LLVM_CONFIG to proceed. We suggest LLVM $SUGGESTED_LLVM_VERSION." fi fi @@ -78,8 +97,8 @@ LLVM_VERSION_MAJOR="$(echo $LLVM_VERSION | awk -F. '{print $1}')" LLVM_VERSION_MINOR="$(echo $LLVM_VERSION | awk -F. '{print $2}')" LLVM_VERSION_PATCH="$(echo $LLVM_VERSION | awk -F. '{print $3}')" -if [ $LLVM_VERSION_MAJOR -lt 14 ] || ([ $LLVM_VERSION_MAJOR -gt 14 ] && [ $LLVM_VERSION_MAJOR -lt 17 ]) || [ $LLVM_VERSION_MAJOR -gt 22 ]; then - error "Invalid LLVM version $LLVM_VERSION: must be 14, 17, 18, 19, 20, 21 or 22" +if [ $LLVM_VERSION_MAJOR -lt $MINIMUM_LLVM_VERSION ]; then + error "Unsupported LLVM version $LLVM_VERSION: must be 17, 18, 19, 20, 21 or 22" fi case "$OS_NAME" in @@ -114,12 +133,7 @@ Linux) ;; OpenBSD) CXXFLAGS="$CXXFLAGS -I/usr/local/include $($LLVM_CONFIG --cxxflags --ldflags)" - LDFLAGS="$LDFLAGS -lstdc++ -L/usr/local/lib -liconv" - LDFLAGS="$LDFLAGS $($LLVM_CONFIG --libs core native --system-libs)" - ;; -Haiku) - CXXFLAGS="$CXXFLAGS -D_GNU_SOURCE $($LLVM_CONFIG --cxxflags --ldflags) -I/system/develop/headers/private/shared -I/system/develop/headers/private/kernel" - LDFLAGS="$LDFLAGS -lstdc++ -liconv" + LDFLAGS="$LDFLAGS -lstdc++ -L/usr/local/lib -Wl,-rpath,$($LLVM_CONFIG --libdir) -liconv" LDFLAGS="$LDFLAGS $($LLVM_CONFIG --libs core native --system-libs)" ;; *) diff --git a/check_all.bat b/check_all.bat index c5f7ee399..aaf612d65 100644 --- a/check_all.bat +++ b/check_all.bat @@ -49,12 +49,8 @@ if "%1" == "freestanding" ( ) if "%1" == "rare" ( - echo Checking essence_amd64 - odin check examples\all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:essence_amd64 echo Checking freebsd_i386 odin check examples\all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:freebsd_i386 - echo Checking haiku_amd64 - odin check examples\all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:haiku_amd64 ) if "%1" == "wasm" ( @@ -72,4 +68,4 @@ if "%1" == "wasm" ( odin check examples\all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:orca_wasm32 echo Checking js_wasm64p32 odin check examples\all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:js_wasm64p32 -) \ No newline at end of file +) diff --git a/check_all.sh b/check_all.sh index 568ac55ba..eee08790e 100755 --- a/check_all.sh +++ b/check_all.sh @@ -19,12 +19,8 @@ freestanding) ;; rare) - echo Checking essence_amd64 - odin check examples/all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:essence_amd64 echo Checking freebsd_i386 odin check examples/all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:freebsd_i386 - echo Checking haiku_amd64 - odin check examples/all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:haiku_amd64 ;; wasm) @@ -75,4 +71,4 @@ wasm) odin check examples/all -vet -vet-tabs -strict-style -vet-style -warnings-as-errors -disallow-do -target:openbsd_amd64 ;; -esac \ No newline at end of file +esac diff --git a/ci/build_linux_static.sh b/ci/build_linux_static.sh index 2eb99116e..d88081f1e 100755 --- a/ci/build_linux_static.sh +++ b/ci/build_linux_static.sh @@ -6,7 +6,14 @@ LLVM_CONFIG="llvm-config-20" DISABLED_WARNINGS="-Wno-switch -Wno-macro-redefined -Wno-unused-value" -CPPFLAGS="-DODIN_VERSION_RAW=\"dev-$(date +"%Y-%m")\"" +if [ -d ".git" ] && [ -n "$(command -v git)" ]; then + GIT_SHA=$(git show --pretty='%h' --no-patch --no-notes HEAD) + GIT_DATE=$(git show "--pretty=%cd" "--date=format:%Y-%m" --no-patch --no-notes HEAD) + CPPFLAGS="$CPPFLAGS -DGIT_SHA=\"$GIT_SHA\"" +else + GIT_DATE=$(date +"%Y-%m") +fi +CPPFLAGS="$CPPFLAGS -DODIN_VERSION_RAW=\"dev-$GIT_DATE\"" CXXFLAGS="-std=c++14 $($LLVM_CONFIG --cxxflags --ldflags)" LDFLAGS="-static -lm -lzstd -lz -lffi -pthread -ldl -fuse-ld=mold" diff --git a/core/bufio/reader.odin b/core/bufio/reader.odin index e361612d2..f949bbd85 100644 --- a/core/bufio/reader.odin +++ b/core/bufio/reader.odin @@ -355,16 +355,15 @@ _reader_proc :: proc(stream_data: rawptr, mode: io.Stream_Mode, p: []byte, offse // -// reader_read_slice reads until the first occurrence of delim from the reader -// It returns a slice pointing at the bytes in the buffer -// The bytes stop being valid at the next read -// If reader_read_slice encounters an error before finding a delimiter -// reader_read_slice fails with error .Buffer_Full if the buffer fills without a delim -// Because the data returned from reader_read_slice will be overwritten on the -// next IO operation, reader_read_bytes or reader_read_string is usually preferred -// -// reader_read_slice returns err != nil if and only if line does not end in delim +// reader_read_slice reads until the first occurrence of delim in the input, +// returning a slice pointing at the bytes in the internal buffer. +// The returned slice is only valid until the next read call. +// If the buffer fills without finding delim, it returns .Buffer_Full. +// If the underlying reader returns an error before finding delim, that error is returned. +// Because the returned data will be overwritten by the next I/O operation, +// reader_read_bytes or reader_read_string is usually preferred. // +// reader_read_slice returns err != nil if and only if line does not end in delim. reader_read_slice :: proc(b: ^Reader, delim: byte) -> (line: []byte, err: io.Error) { s := 0 for { diff --git a/core/c/c.odin b/core/c/c.odin index bc847d566..cbff10e66 100644 --- a/core/c/c.odin +++ b/core/c/c.odin @@ -1,7 +1,8 @@ // Defines the basic types used by `C` programs for foreign function and data structure interop. package c -import builtin "base:builtin" +import builtin "base:builtin" +import intrinsics "base:intrinsics" char :: builtin.u8 // assuming -funsigned-char @@ -115,14 +116,6 @@ NDEBUG :: !ODIN_DEBUG CHAR_BIT :: 8 -// Since there are no types in C with an alignment larger than that of -// max_align_t, which cannot be larger than sizeof(long double) as any other -// exposed type wouldn't be valid C, the maximum alignment possible in a -// strictly conformant C implementation is 16 on the platforms we care about. -// The choice of 4096 bytes for storage of this type is more than enough on all -// relevant platforms. -va_list :: struct #align(16) { - _: [4096]u8, -} +va_list :: intrinsics.c_va_list FILE :: struct {} diff --git a/core/c/libc/errno.odin b/core/c/libc/errno.odin index 138d70a80..1f6df52f7 100644 --- a/core/c/libc/errno.odin +++ b/core/c/libc/errno.odin @@ -80,25 +80,6 @@ when ODIN_OS == .Darwin { ERANGE :: 34 } -when ODIN_OS == .Haiku { - @(private="file") - @(default_calling_convention="c") - foreign libc { - @(link_name="_errnop") - _get_errno :: proc() -> ^int --- - } - - _HAIKU_USE_POSITIVE_POSIX_ERRORS :: #config(HAIKU_USE_POSITIVE_POSIX_ERRORS, false) - _POSIX_ERROR_FACTOR :: -1 when _HAIKU_USE_POSITIVE_POSIX_ERRORS else 1 - - @(private="file") _GENERAL_ERROR_BASE :: min(int) - @(private="file") _POSIX_ERROR_BASE :: _GENERAL_ERROR_BASE + 0x7000 - - EDOM :: _POSIX_ERROR_FACTOR * (_POSIX_ERROR_BASE + 16) - EILSEQ :: _POSIX_ERROR_FACTOR * (_POSIX_ERROR_BASE + 38) - ERANGE :: _POSIX_ERROR_FACTOR * (_POSIX_ERROR_BASE + 17) -} - when ODIN_OS == .JS { _ :: libc _get_errno :: proc "c" () -> ^int { diff --git a/core/c/libc/locale.odin b/core/c/libc/locale.odin index 3216e0f90..59cbc6c6d 100644 --- a/core/c/libc/locale.odin +++ b/core/c/libc/locale.odin @@ -110,7 +110,7 @@ when ODIN_OS == .Windows { } } -when ODIN_OS == .Darwin || ODIN_OS == .FreeBSD || ODIN_OS == .NetBSD || ODIN_OS == .OpenBSD || ODIN_OS == .Haiku || ODIN_OS == .Windows { +when ODIN_OS == .Darwin || ODIN_OS == .FreeBSD || ODIN_OS == .NetBSD || ODIN_OS == .OpenBSD || ODIN_OS == .Windows { LC_ALL :: 0 LC_COLLATE :: 1 diff --git a/core/c/libc/signal.odin b/core/c/libc/signal.odin index cddf06916..cfe9768c9 100644 --- a/core/c/libc/signal.odin +++ b/core/c/libc/signal.odin @@ -34,7 +34,7 @@ when ODIN_OS == .Windows { SIGTERM :: 15 } -when ODIN_OS == .Linux || ODIN_OS == .FreeBSD || ODIN_OS == .Haiku || ODIN_OS == .OpenBSD || ODIN_OS == .NetBSD || ODIN_OS == .Darwin { +when ODIN_OS == .Linux || ODIN_OS == .FreeBSD || ODIN_OS == .OpenBSD || ODIN_OS == .NetBSD || ODIN_OS == .Darwin { SIG_ERR :: rawptr(~uintptr(0)) SIG_DFL :: rawptr(uintptr(0)) SIG_IGN :: rawptr(uintptr(1)) diff --git a/core/c/libc/stdarg.odin b/core/c/libc/stdarg.odin index 232471713..339e4cb54 100644 --- a/core/c/libc/stdarg.odin +++ b/core/c/libc/stdarg.odin @@ -4,29 +4,12 @@ package libc import "base:intrinsics" -import "core:c" +va_list :: intrinsics.c_va_list -@(private="file") -@(default_calling_convention="none") -foreign _ { - @(link_name="llvm.va_start") _va_start :: proc(arglist: ^i8) --- - @(link_name="llvm.va_end") _va_end :: proc(arglist: ^i8) --- - @(link_name="llvm.va_copy") _va_copy :: proc(dst, src: ^i8) --- -} +va_start :: intrinsics.c_va_start +va_end :: intrinsics.c_va_end +va_copy :: intrinsics.c_va_copy -va_list :: c.va_list - -va_start :: #force_inline proc(ap: ^va_list, _: any) { - _va_start(cast(^i8)ap) -} - -va_end :: #force_inline proc(ap: ^va_list) { - _va_end(cast(^i8)ap) -} - -va_copy :: #force_inline proc(dst, src: ^va_list) { - _va_copy(cast(^i8)dst, cast(^i8)src) -} // We cannot provide va_arg as there is no way to create "C" style procedures // in Odin which take variable arguments the C way. The #c_vararg attribute only diff --git a/core/c/libc/stdio.odin b/core/c/libc/stdio.odin index db5b202e7..ffdb06f67 100644 --- a/core/c/libc/stdio.odin +++ b/core/c/libc/stdio.odin @@ -200,36 +200,6 @@ when ODIN_OS == .Darwin { } } -when ODIN_OS == .Haiku { - fpos_t :: distinct i64 - - _IOFBF :: 0 - _IOLBF :: 1 - _IONBF :: 2 - - BUFSIZ :: 8192 - - EOF :: int(-1) - - FOPEN_MAX :: 128 - - FILENAME_MAX :: 256 - - L_tmpnam :: 512 - - SEEK_SET :: 0 - SEEK_CUR :: 1 - SEEK_END :: 2 - - TMP_MAX :: 32768 - - foreign libc { - stderr: ^FILE - stdin: ^FILE - stdout: ^FILE - } -} - when ODIN_OS == .NetBSD { @(private) LRENAME :: "__posix_rename" @(private) LFGETPOS :: "__fgetpos50" diff --git a/core/c/libc/stdlib.odin b/core/c/libc/stdlib.odin index ca906a5f0..f63f40cd0 100644 --- a/core/c/libc/stdlib.odin +++ b/core/c/libc/stdlib.odin @@ -42,22 +42,6 @@ when ODIN_OS == .Linux { } } -when ODIN_OS == .Haiku { - RAND_MAX :: 0x7fffffff - - // GLIBC and MUSL only - @(private="file") - @(default_calling_convention="c") - foreign libc { - __ctype_get_mb_cur_max :: proc() -> ushort --- - } - - MB_CUR_MAX :: #force_inline proc() -> size_t { - return size_t(__ctype_get_mb_cur_max()) - } -} - - when ODIN_OS == .Darwin || ODIN_OS == .FreeBSD || ODIN_OS == .OpenBSD { RAND_MAX :: 0x7fffffff diff --git a/core/c/libc/time.odin b/core/c/libc/time.odin index 6106923f5..d431a84d7 100644 --- a/core/c/libc/time.odin +++ b/core/c/libc/time.odin @@ -45,7 +45,7 @@ when ODIN_OS == .Windows { } } -when ODIN_OS == .Linux || ODIN_OS == .FreeBSD || ODIN_OS == .Darwin || ODIN_OS == .OpenBSD || ODIN_OS == .NetBSD || ODIN_OS == .Haiku || ODIN_OS == .JS { +when ODIN_OS == .Linux || ODIN_OS == .FreeBSD || ODIN_OS == .Darwin || ODIN_OS == .OpenBSD || ODIN_OS == .NetBSD || ODIN_OS == .JS { @(default_calling_convention="c") foreign libc { // 7.27.2 Time manipulation functions @@ -95,7 +95,7 @@ when ODIN_OS == .Linux || ODIN_OS == .FreeBSD || ODIN_OS == .Darwin || ODIN_OS = time_t :: distinct i64 - when ODIN_OS == .FreeBSD || ODIN_OS == .NetBSD || ODIN_OS == .Haiku { + when ODIN_OS == .FreeBSD || ODIN_OS == .NetBSD { clock_t :: distinct int32_t } else { clock_t :: distinct long diff --git a/core/c/libc/wchar.odin b/core/c/libc/wchar.odin index 248611409..73db7296c 100644 --- a/core/c/libc/wchar.odin +++ b/core/c/libc/wchar.odin @@ -17,12 +17,12 @@ foreign libc { fwscanf :: proc(stream: ^FILE, format: [^]wchar_t, #c_vararg arg: ..any) -> int --- swprintf :: proc(stream: ^FILE, n: size_t, format: [^]wchar_t, #c_vararg arg: ..any) -> int --- swscanf :: proc(s, format: [^]wchar_t, #c_vararg arg: ..any) -> int --- - vfwprintf :: proc(stream: ^FILE, format: [^]wchar_t, arg: va_list) -> int --- - vfwscanf :: proc(stream: ^FILE, format: [^]wchar_t, arg: va_list) -> int --- - vswprintf :: proc(s: [^]wchar_t, n: size_t, format: [^]wchar_t, arg: va_list) -> int --- - vswscanf :: proc(s, format: [^]wchar_t, arg: va_list) -> int --- - vwprintf :: proc(format: [^]wchar_t, arg: va_list) -> int --- - vwscanf :: proc(format: [^]wchar_t, arg: va_list) -> int --- + vfwprintf :: proc(stream: ^FILE, format: [^]wchar_t, arg: ^va_list) -> int --- + vfwscanf :: proc(stream: ^FILE, format: [^]wchar_t, arg: ^va_list) -> int --- + vswprintf :: proc(s: [^]wchar_t, n: size_t, format: [^]wchar_t, arg: ^va_list) -> int --- + vswscanf :: proc(s, format: [^]wchar_t, arg: ^va_list) -> int --- + vwprintf :: proc(format: [^]wchar_t, arg: ^va_list) -> int --- + vwscanf :: proc(format: [^]wchar_t, arg: ^va_list) -> int --- wprintf :: proc(format: [^]wchar_t, #c_vararg arg: ..any) -> int --- wscanf :: proc(format: [^]wchar_t, #c_vararg arg: ..any) -> int --- diff --git a/core/c/libc/wctype.odin b/core/c/libc/wctype.odin index 6526a14e2..7d39e75c0 100644 --- a/core/c/libc/wctype.odin +++ b/core/c/libc/wctype.odin @@ -30,10 +30,6 @@ when ODIN_OS == .Windows { wctrans_t :: distinct int wctype_t :: distinct ulong -} else when ODIN_OS == .Haiku { - wctrans_t :: distinct i32 - wctype_t :: distinct i32 - } @(default_calling_convention="c") diff --git a/core/compress/common.odin b/core/compress/common.odin index 07ea7b71c..fcda3fc20 100644 --- a/core/compress/common.odin +++ b/core/compress/common.odin @@ -360,21 +360,26 @@ refill_lsb_from_memory :: #force_inline proc(z: ^Context_Memory_Input, width := refill := u64(width) b := u64(0) - if z.num_bits > refill { - return - } - for { - if len(z.input_data) != 0 { - b = u64(z.input_data[0]) - z.input_data = z.input_data[1:] - } - - z.code_buffer |= b << u8(z.num_bits) - z.num_bits += 8 if z.num_bits > refill { break } + if z.code_buffer == 0 && z.num_bits > 63 { + z.num_bits = 0 + } + if z.code_buffer >= 1 << uint(z.num_bits) { + // Code buffer is malformed. + z.num_bits = max(u64) + return + } + if len(z.input_data) != 0 { + b = u64(z.input_data[0]) + z.input_data = z.input_data[1:] + } else { + return + } + z.code_buffer |= b << u8(z.num_bits) + z.num_bits += 8 } } diff --git a/core/compress/zlib/zlib.odin b/core/compress/zlib/zlib.odin index 72664846e..efa9cb906 100644 --- a/core/compress/zlib/zlib.odin +++ b/core/compress/zlib/zlib.odin @@ -322,11 +322,8 @@ decode_huffman_slowpath :: proc(z: ^$C, t: ^Huffman_Table) -> (r: u16, err: Erro @(optimization_mode="favor_size") decode_huffman :: proc(z: ^$C, t: ^Huffman_Table) -> (r: u16, err: Error) #no_bounds_check { if z.num_bits < 16 { - if z.num_bits > 63 { - return 0, .Code_Buffer_Malformed - } compress.refill_lsb(z) - if z.code_buffer == 0 { + if z.num_bits > 63 { return 0, .Stream_Too_Short } } @@ -491,7 +488,7 @@ inflate_raw :: proc(z: ^$C, expected_output_size := -1, allocator := context.all */ expected_output_size = max(max(expected_output_size, compress.COMPRESS_OUTPUT_ALLOCATE_MIN), 512) - // fmt.printfln("ZLIB: Expected Payload Size: %v", expected_output_size) + // fmt.printf("\nZLIB: Expected Payload Size: %v\n\n", expected_output_size); if expected_output_size > 0 && expected_output_size <= compress.COMPRESS_OUTPUT_ALLOCATE_MAX { /* @@ -522,16 +519,11 @@ inflate_raw :: proc(z: ^$C, expected_output_size := -1, allocator := context.all final := u32(0) type := u32(0) - defer if int(z.bytes_written) != len(z.output.buf) { - resize(&z.output.buf, int(z.bytes_written)) - } - for { final = compress.read_bits_lsb(z, 1) type = compress.read_bits_lsb(z, 2) - // fmt.printfln("len(z): %v", len(z.input_data)) - // fmt.printfln("Final: %v | Type: %v", final, type) + // fmt.printf("Final: %v | Type: %v\n", final, type) switch type { case 0: @@ -566,6 +558,7 @@ inflate_raw :: proc(z: ^$C, expected_output_size := -1, allocator := context.all case 3: return .BType_3 case: + // fmt.printf("Err: %v | Final: %v | Type: %v\n", err, final, type) if type == 1 { // Use fixed code lengths. build_huffman(z_repeat, Z_FIXED_LENGTH[:]) or_return @@ -594,6 +587,7 @@ inflate_raw :: proc(z: ^$C, expected_output_size := -1, allocator := context.all for n < ntot { c = decode_huffman(z, codelength_ht) or_return + if c < 0 || c >= 19 { return .Huffman_Bad_Code_Lengths } @@ -638,12 +632,15 @@ inflate_raw :: proc(z: ^$C, expected_output_size := -1, allocator := context.all } parse_huffman_block(z, z_repeat, z_offset) or_return } - if final == 1 { break } } + if int(z.bytes_written) != len(z.output.buf) { + resize(&z.output.buf, int(z.bytes_written)) or_return + } + return nil } diff --git a/core/container/avl/avl.odin b/core/container/avl/avl.odin index 1208cd213..7c99879a7 100644 --- a/core/container/avl/avl.odin +++ b/core/container/avl/avl.odin @@ -60,13 +60,13 @@ Iterator :: struct($Value: typeid) { _called_next: bool, } -// init initializes a tree. +// `init` initializes a tree. init :: proc { init_ordered, init_cmp, } -// init_cmp initializes a tree. +// `init_cmp` initializes a tree. init_cmp :: proc( t: ^$T/Tree($Value), cmp_fn: proc(a, b: Value) -> Ordering, @@ -78,7 +78,7 @@ init_cmp :: proc( t._size = 0 } -// init_ordered initializes a tree containing ordered items, with +// `init_ordered` initializes a tree containing ordered items, with // a comparison function that results in an ascending order sort. init_ordered :: proc( t: ^$T/Tree($Value), @@ -87,7 +87,7 @@ init_ordered :: proc( init_cmp(t, slice.cmp_proc(Value), node_allocator) } -// destroy de-initializes a tree. +// `destroy` de-initializes a tree. destroy :: proc(t: ^$T/Tree($Value), call_on_remove: bool = true) { iter := iterator(t, Direction.Forward) for _ in iterator_next(&iter) { @@ -95,24 +95,24 @@ destroy :: proc(t: ^$T/Tree($Value), call_on_remove: bool = true) { } } -// len returns the number of elements in the tree. +// `len` returns the number of elements in the tree. len :: proc "contextless" (t: ^$T/Tree($Value)) -> int { return t._size } -// first returns the first node in the tree (in-order) or nil if and only if (⟺) +// `first` returns the first node in the tree (in-order) or nil if and only if (⟺) // the tree is empty. first :: proc "contextless" (t: ^$T/Tree($Value)) -> ^Node(Value) { return tree_first_or_last_in_order(t, Direction.Backward) } -// last returns the last element in the tree (in-order) or nil if and only if (⟺) +// `last` returns the last element in the tree (in-order) or nil if and only if (⟺) // the tree is empty. last :: proc "contextless" (t: ^$T/Tree($Value)) -> ^Node(Value) { return tree_first_or_last_in_order(t, Direction.Forward) } -// find finds the value in the tree, and returns the corresponding +// `find` finds the value in the tree, and returns the corresponding // node or nil if and only if (⟺) the value is not present. find :: proc(t: ^$T/Tree($Value), value: Value) -> ^Node(Value) { cur := t._root @@ -130,7 +130,7 @@ find :: proc(t: ^$T/Tree($Value), value: Value) -> ^Node(Value) { return cur } -// find_or_insert attempts to insert the value into the tree, and returns +// `find_or_insert` attempts to insert the value into the tree, and returns // the node, a boolean indicating if the value was inserted, and the // node allocator error if relevant. If the value is already // present, the existing node is returned un-altered. @@ -168,7 +168,7 @@ find_or_insert :: proc( return } -// remove removes a node or value from the tree, and returns true if and only if (⟺) the +// `remove` removes a node or value from the tree, and returns true if and only if (⟺) the // removal was successful. While the node's value will be left intact, // the node itself will be freed via the tree's node allocator. remove :: proc { @@ -176,7 +176,7 @@ remove :: proc { remove_node, } -// remove_value removes a value from the tree, and returns true if and only if (⟺) the +// `remove_value` removes a value from the tree, and returns true if and only if (⟺) the // removal was successful. While the node's value will be left intact, // the node itself will be freed via the tree's node allocator. remove_value :: proc(t: ^$T/Tree($Value), value: Value, call_on_remove: bool = true) -> bool { @@ -187,7 +187,7 @@ remove_value :: proc(t: ^$T/Tree($Value), value: Value, call_on_remove: bool = t return remove_node(t, n, call_on_remove) } -// remove_node removes a node from the tree, and returns true if and only if (⟺) the +// `remove_node` removes a node from the tree, and returns true if and only if (⟺) the // removal was successful. While the node's value will be left intact, // the node itself will be freed via the tree's node allocator. remove_node :: proc(t: ^$T/Tree($Value), node: ^Node(Value), call_on_remove: bool = true) -> bool { @@ -249,7 +249,7 @@ remove_node :: proc(t: ^$T/Tree($Value), node: ^Node(Value), call_on_remove: boo return true } -// iterator returns a tree iterator in the specified direction. +// `iterator` returns a tree iterator in the specified direction. iterator :: proc "contextless" (t: ^$T/Tree($Value), direction: Direction) -> Iterator(Value) { it: Iterator(Value) it._tree = transmute(^Tree(Value))t @@ -260,7 +260,7 @@ iterator :: proc "contextless" (t: ^$T/Tree($Value), direction: Direction) -> It return it } -// iterator_from_pos returns a tree iterator in the specified direction, +// `iterator_from_pos` returns a tree iterator in the specified direction, // spanning the range [pos, last] (inclusive). iterator_from_pos :: proc "contextless" ( t: ^$T/Tree($Value), @@ -280,14 +280,14 @@ iterator_from_pos :: proc "contextless" ( return it } -// iterator_get returns the node currently pointed to by the iterator, +// `iterator_get` returns the node currently pointed to by the iterator, // or nil if and only if (⟺) the node has been removed, the tree is empty, or the end // of the tree has been reached. iterator_get :: proc "contextless" (it: ^$I/Iterator($Value)) -> ^Node(Value) { return it._cur } -// iterator_remove removes the node currently pointed to by the iterator, +// `iterator_remove` removes the node currently pointed to by the iterator, // and returns true if and only if (⟺) the removal was successful. Semantics are the // same as the Tree remove. iterator_remove :: proc(it: ^$I/Iterator($Value), call_on_remove: bool = true) -> bool { @@ -303,7 +303,7 @@ iterator_remove :: proc(it: ^$I/Iterator($Value), call_on_remove: bool = true) - return ok } -// iterator_next advances the iterator and returns the (node, true) or +// `iterator_next` advances the iterator and returns the (node, true) or // or (nil, false) if and only if (⟺) the end of the tree has been reached. // // Note: The first call to iterator_next will return the first node instead diff --git a/core/container/handle_map/dynamic_handle_map.odin b/core/container/handle_map/dynamic_handle_map.odin index 09ea92f82..952d9daea 100644 --- a/core/container/handle_map/dynamic_handle_map.odin +++ b/core/container/handle_map/dynamic_handle_map.odin @@ -1,7 +1,6 @@ package container_handle_map import "base:runtime" -import "base:builtin" import "base:intrinsics" @(require) import "core:container/xar" @@ -139,4 +138,4 @@ dynamic_iterate :: proc "contextless" (it: ^$DHI/Dynamic_Handle_Map_Iterator($D/ } it.index = 0 return -} \ No newline at end of file +} diff --git a/core/container/pool/pool.odin b/core/container/pool/pool.odin index 34a9e0683..c686d9d6b 100644 --- a/core/container/pool/pool.odin +++ b/core/container/pool/pool.odin @@ -29,6 +29,13 @@ Pool :: struct($T: typeid) { num_outstanding: int, num_ready: int, link_off: uintptr, + // Guards free_list. An untagged Treiber stack is vulnerable to ABA: + // between get's head load and its CAS, another thread can pop the head, + // reuse it, and push it back — the CAS then succeeds and installs a + // stale next pointer, handing an in-flight element to two owners. + // A tagged head would keep this lock-free but requires a double-width + // CAS, which is not portably available. + mu: sync.Mutex, free_list: ^T, } @@ -61,20 +68,20 @@ destroy :: proc(p: ^Pool($T)) { get :: proc(p: ^Pool($T)) -> (elem: ^T, err: runtime.Allocator_Error) #optional_allocator_error { defer sync.atomic_add_explicit(&p.num_outstanding, 1, .Relaxed) - for { - elem = sync.atomic_load_explicit(&p.free_list, .Acquire) - if elem == nil { - // NOTE: pool arena has an internal lock. - return new(T, _pool_arena_allocator(&p.arena)) - } - - if _, ok := sync.atomic_compare_exchange_weak_explicit(&p.free_list, elem, _get_next(p, elem), .Acquire, .Relaxed); ok { - _set_next(p, elem, nil) - _unpoison_elem(p, elem) - sync.atomic_sub_explicit(&p.num_ready, 1, .Relaxed) - return - } + sync.mutex_lock(&p.mu) + elem = p.free_list + if elem == nil { + sync.mutex_unlock(&p.mu) + // NOTE: pool arena has an internal lock. + return new(T, _pool_arena_allocator(&p.arena)) } + p.free_list = _get_next(p, elem) + sync.mutex_unlock(&p.mu) + + _set_next(p, elem, nil) + _unpoison_elem(p, elem) + sync.atomic_sub_explicit(&p.num_ready, 1, .Relaxed) + return } put :: proc(p: ^Pool($T), elem: ^T) { @@ -84,13 +91,10 @@ put :: proc(p: ^Pool($T), elem: ^T) { defer sync.atomic_sub_explicit(&p.num_outstanding, 1, .Relaxed) defer sync.atomic_add_explicit(&p.num_ready, 1, .Relaxed) - for { - head := sync.atomic_load_explicit(&p.free_list, .Relaxed) - _set_next(p, elem, head) - if _, ok := sync.atomic_compare_exchange_weak_explicit(&p.free_list, head, elem, .Release, .Relaxed); ok { - return - } - } + sync.mutex_lock(&p.mu) + _set_next(p, elem, p.free_list) + p.free_list = elem + sync.mutex_unlock(&p.mu) } num_outstanding :: proc(p: ^Pool($T)) -> int { diff --git a/core/container/queue/queue.odin b/core/container/queue/queue.odin index 37f3a923f..4b371f378 100644 --- a/core/container/queue/queue.odin +++ b/core/container/queue/queue.odin @@ -236,16 +236,6 @@ back_ptr :: proc(q: ^$Q/Queue($T), loc := #caller_location) -> ^T { } -@(deprecated="Use `front_ptr` instead") -peek_front :: proc(q: ^$Q/Queue($T), loc := #caller_location) -> ^T { - return front_ptr(q, loc) -} - -@(deprecated="Use `back_ptr` instead") -peek_back :: proc(q: ^$Q/Queue($T), loc := #caller_location) -> ^T { - return back_ptr(q, loc) -} - /* Push an element to the back of the queue. @@ -416,7 +406,7 @@ push_back_elems :: proc(q: ^$Q/Queue($T), elems: ..T, loc := #caller_location) - } /* -Consume `n` elements from the back of the queue. +Consume `n` elements from the front of the queue. This will raise a bounds checking error if the queue does not have enough elements. */ diff --git a/core/container/rbtree/rbtree.odin b/core/container/rbtree/rbtree.odin index c138838df..d3b790dcc 100644 --- a/core/container/rbtree/rbtree.odin +++ b/core/container/rbtree/rbtree.odin @@ -63,13 +63,13 @@ Iterator :: struct($Key: typeid, $Value: typeid) { _called_next: bool, } -// init initializes a tree. +// `init` initializes a tree. init :: proc { init_ordered, init_cmp, } -// init_cmp initializes a tree. +// `init_cmp` initializes a tree. init_cmp :: proc(t: ^$T/Tree($Key, $Value), cmp_fn: proc(a, b: Key) -> Ordering, node_allocator := context.allocator) { t._root = nil t._node_allocator = node_allocator @@ -77,13 +77,13 @@ init_cmp :: proc(t: ^$T/Tree($Key, $Value), cmp_fn: proc(a, b: Key) -> Ordering, t._size = 0 } -// init_ordered initializes a tree containing ordered keys, with +// `init_ordered` initializes a tree containing ordered keys, with // a comparison function that results in an ascending order sort. init_ordered :: proc(t: ^$T/Tree($Key, $Value), node_allocator := context.allocator) where intrinsics.type_is_ordered(Key) { init_cmp(t, slice.cmp_proc(Key), node_allocator) } -// destroy de-initializes a tree. +// `destroy` de-initializes a tree. destroy :: proc(t: ^$T/Tree($Key, $Value), call_on_remove: bool = true) { iter := iterator(t, .Forward) for _ in iterator_next(&iter) { @@ -95,19 +95,19 @@ len :: proc "contextless" (t: $T/Tree($Key, $Value)) -> (node_count: int) { return t._size } -// first returns the first node in the tree (in-order) or nil if and only if (⟺) +// `first` returns the first node in the tree (in-order) or nil if and only if (⟺) // the tree is empty. first :: proc "contextless" (t: ^$T/Tree($Key, $Value)) -> ^Node(Key, Value) { return tree_first_or_last_in_order(t, Direction.Backward) } -// last returns the last element in the tree (in-order) or nil if and only if (⟺) +// `last` returns the last element in the tree (in-order) or nil if and only if (⟺) // the tree is empty. last :: proc "contextless" (t: ^$T/Tree($Key, $Value)) -> ^Node(Key, Value) { return tree_first_or_last_in_order(t, Direction.Forward) } -// find finds the key in the tree, and returns the corresponding node, or nil if and only if (⟺) the value is not present. +// `find` finds the key in the tree, and returns the corresponding node, or nil if and only if (⟺) the value is not present. find :: proc(t: $T/Tree($Key, $Value), key: Key) -> (node: ^Node(Key, Value)) { node = t._root for node != nil { @@ -120,7 +120,7 @@ find :: proc(t: $T/Tree($Key, $Value), key: Key) -> (node: ^Node(Key, Value)) { return node } -// find_value finds the key in the tree, and returns the corresponding value, or nil if and only if (⟺) the value is not present. +// `find_value` finds the key in the tree, and returns the corresponding value, or nil if and only if (⟺) the value is not present. find_value :: proc(t: $T/Tree($Key, $Value), key: Key) -> (value: Value, ok: bool) #optional_ok { if n := find(t, key); n != nil { return n.value, true @@ -128,10 +128,36 @@ find_value :: proc(t: $T/Tree($Key, $Value), key: Key) -> (value: Value, ok: boo return } -// find_or_insert attempts to insert the key-value pair into the tree, and returns +// `find_or_insert` attempts to insert the key-value pair into the tree, and returns // the node, a boolean indicating if a new node was inserted, and the -// node allocator error if relevant. If the key is already present, the existing node is updated and returned. +// node allocator error if relevant. If the key is already present, the existing node is returned un-altered. find_or_insert :: proc(t: ^$T/Tree($Key, $Value), key: Key, value: Value) -> (n: ^Node(Key, Value), inserted: bool, err: runtime.Allocator_Error) { + n_ptr := &t._root + for n_ptr^ != nil { + n = n_ptr^ + switch t._cmp_fn(key, n.key) { + case .Less: + n_ptr = &n._left + case .Greater: + n_ptr = &n._right + case .Equal: + return + } + } + _parent := n + + n = new_clone(Node(Key, Value){key=key, value=value, _parent=_parent, _color=.Red}, t._node_allocator) or_return + n_ptr^ = n + insert_case1(t, n) + t._size += 1 + return n, true, nil +} + + +// `upsert` attempts to insert the key-value pair into the tree, and returns +// the node, a boolean indicating if a new node was inserted, and the +// node allocator error if relevant. If the key is already present, the existing node's value is updated. +upsert :: proc(t: ^$T/Tree($Key, $Value), key: Key, value: Value) -> (n: ^Node(Key, Value), inserted: bool, err: runtime.Allocator_Error) { n_ptr := &t._root for n_ptr^ != nil { n = n_ptr^ @@ -154,7 +180,7 @@ find_or_insert :: proc(t: ^$T/Tree($Key, $Value), key: Key, value: Value) -> (n: return n, true, nil } -// remove removes a node or value from the tree, and returns true if and only if (⟺) the +// `remove` removes a node or value from the tree, and returns true if and only if (⟺) the // removal was successful. While the node's value will be left intact, // the node itself will be freed via the tree's node allocator. remove :: proc { @@ -162,7 +188,7 @@ remove :: proc { remove_node, } -// remove_value removes a value from the tree, and returns true if and only if (⟺) the +// `remove_value` removes a value from the tree, and returns true if and only if (⟺) the // removal was successful. While the node's key + value will be left intact, // the node itself will be freed via the tree's node allocator. remove_key :: proc(t: ^$T/Tree($Key, $Value), key: Key, call_on_remove := true) -> bool { @@ -173,7 +199,7 @@ remove_key :: proc(t: ^$T/Tree($Key, $Value), key: Key, call_on_remove := true) return remove_node(t, n, call_on_remove) } -// remove_node removes a node from the tree, and returns true if and only if (⟺) the +// `remove_node` removes a node from the tree, and returns true if and only if (⟺) the // removal was successful. While the node's key + value will be left intact, // the node itself will be freed via the tree's node allocator. remove_node :: proc(t: ^$T/Tree($Key, $Value), node: ^$N/Node(Key, Value), call_on_remove := true) -> (found: bool) { @@ -207,7 +233,7 @@ remove_node :: proc(t: ^$T/Tree($Key, $Value), node: ^$N/Node(Key, Value), call_ return true } -// iterator returns a tree iterator in the specified direction. +// `iterator` returns a tree iterator in the specified direction. iterator :: proc "contextless" (t: ^$T/Tree($Key, $Value), direction: Direction) -> Iterator(Key, Value) { it: Iterator(Key, Value) it._tree = cast(^Tree(Key, Value))t @@ -218,7 +244,7 @@ iterator :: proc "contextless" (t: ^$T/Tree($Key, $Value), direction: Direction) return it } -// iterator_from_pos returns a tree iterator in the specified direction, +// `iterator_from_pos` returns a tree iterator in the specified direction, // spanning the range [pos, last] (inclusive). iterator_from_pos :: proc "contextless" (t: ^$T/Tree($Key, $Value), pos: ^Node(Key, Value), direction: Direction) -> Iterator(Key, Value) { it: Iterator(Key, Value) @@ -234,14 +260,14 @@ iterator_from_pos :: proc "contextless" (t: ^$T/Tree($Key, $Value), pos: ^Node(K return it } -// iterator_get returns the node currently pointed to by the iterator, +// `iterator_get` returns the node currently pointed to by the iterator, // or nil if and only if (⟺) the node has been removed, the tree is empty, or the end // of the tree has been reached. iterator_get :: proc "contextless" (it: ^$I/Iterator($Key, $Value)) -> ^Node(Key, Value) { return it._cur } -// iterator_remove removes the node currently pointed to by the iterator, +// `iterator_remove` removes the node currently pointed to by the iterator, // and returns true if and only if (⟺) the removal was successful. Semantics are the // same as the Tree remove. iterator_remove :: proc(it: ^$I/Iterator($Key, $Value), call_on_remove: bool = true) -> bool { @@ -257,7 +283,7 @@ iterator_remove :: proc(it: ^$I/Iterator($Key, $Value), call_on_remove: bool = t return ok } -// iterator_next advances the iterator and returns the (node, true) or +// `iterator_next` advances the iterator and returns the (node, true) or // or (nil, false) if and only if (⟺) the end of the tree has been reached. // // Note: The first call to iterator_next will return the first node instead diff --git a/core/container/xar/xar.odin b/core/container/xar/xar.odin index 08b2b2d32..d2a517464 100644 --- a/core/container/xar/xar.odin +++ b/core/container/xar/xar.odin @@ -265,7 +265,7 @@ array_push_back_elem :: proc(x: ^$X/Array($T, $SHIFT), value: T, loc := #caller_ chunk_idx, elem_idx, chunk_cap := _meta_get(SHIFT, uint(x.len)) if x.chunks[chunk_idx] == nil { - x.chunks[chunk_idx] = make([^]T, chunk_cap, x.allocator) or_return + x.chunks[chunk_idx] = make([^]T, chunk_cap, x.allocator, loc) or_return } x.chunks[chunk_idx][elem_idx] = value x.len += 1 @@ -328,7 +328,7 @@ array_push_back_elem_and_get_ptr :: proc(x: ^$X/Array($T, $SHIFT), value: T, loc chunk_idx, elem_idx, chunk_cap := _meta_get(SHIFT, uint(x.len)) if x.chunks[chunk_idx] == nil { - x.chunks[chunk_idx] = make([^]T, chunk_cap, x.allocator) or_return + x.chunks[chunk_idx] = make([^]T, chunk_cap, x.allocator, loc) or_return } x.chunks[chunk_idx][elem_idx] = value x.len += 1 @@ -375,7 +375,7 @@ array_pop_safe :: proc(x: ^$X/Array($T, $SHIFT)) -> (val: T, ok: bool) { Example: - import "core:encoding/xar" + import "core:container/xar" unordered_remove_example :: proc() { x: xar.Array(int, 4) diff --git a/core/crypto/README.md b/core/crypto/README.md index 5065a6d02..0d1a8711b 100644 --- a/core/crypto/README.md +++ b/core/crypto/README.md @@ -13,10 +13,11 @@ constant-time byte comparison. - The crypto packages are not thread-safe. - Best-effort is make to mitigate timing side-channels on reasonable architectures. Architectures that are known to be unreasonable include - but are not limited to i386, i486, and WebAssembly. + but are not limited to i386, i486, VIA Nano 2000, ARM7T/ARM9T/Cortex-M3, + and WASM. - Implementations assume a 64-bit architecture (64-bit integer arithmetic - is fast, and includes add-with-carry, sub-with-borrow, and full-result - multiply). + is fast, and includes contant-time add-with-carry, sub-with-borrow, and + full-result multiply). - Hardware sidechannels are explicitly out of scope for this package. Notable examples include but are not limited to: - Power/RF side-channels etc. @@ -29,4 +30,4 @@ constant-time byte comparison. ## License -This library is made available under the zlib license. \ No newline at end of file +This library is made available under the zlib license. diff --git a/core/crypto/_bigint/i31.odin b/core/crypto/_bigint/i31.odin new file mode 100644 index 000000000..e71a95310 --- /dev/null +++ b/core/crypto/_bigint/i31.odin @@ -0,0 +1,904 @@ +// Constant time Big Integers +package _bigint + +// Copyright (c) 2016 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "base:intrinsics" +import "core:crypto" +import subtle "core:crypto/_subtle" +import "core:slice" + +// Integers 'i31' +// -------------- +// +// The 'i31' functions implement computations on big integers using +// an internal representation as an array of 32-bit integers. For +// an array `x`: +// -- x[0] encodes the array length and the "announced bit length" +// of the integer: namely, if the announced bit length is k, +// then x[0] = ((k / 31) << 5) + (k % 31). +// -- x[1], x[2]... contain the value in little-endian order, 31 +// bits per word (x[1] contains the least significant 31 bits). +// The upper bit of each word is 0. +// +// Multiplications rely on the elementary 32x32->64 multiplication. +// +// The announced bit length specifies the number of bits that are +// significant in the subsequent 32-bit words. Unused bits in the +// last (most significant) word are set to 0; subsequent words are +// uninitialized and need not exist at all. +// +// The execution time and memory access patterns of all computations +// depend on the announced bit length, but not on the actual word +// values. For modular integers, the announced bit length of any integer +// modulo `n` is equal to the actual bit length of `n`; thus, computations +// on modular integers are "constant-time" (only the modulus length may leak). + +I31_MASK :: 0x7fff_ffff + +// Compute the bit length of a 32-bit integer. +// Returned value is between 0 and 32 (inclusive). +@(require_results) +_u32_bit_length :: proc "contextless" (x: u32) -> (length: u32) { + x := x + k := subtle.neq(x, 0) + c := subtle.gt(x, 0xFFFF); x = subtle.csel(x, x >> 16, c); k += c << 4 + c = subtle.gt(x, 0x00FF); x = subtle.csel(x, x >> 8, c); k += c << 3 + c = subtle.gt(x, 0x000F); x = subtle.csel(x, x >> 4, c); k += c << 2 + c = subtle.gt(x, 0x0003); x = subtle.csel(x, x >> 2, c); k += c << 1 + k += subtle.gt(x, 0x0001) + return k +} + +// Multiply two 31-bit integers, with a 62-bit result. This default +// implementation assumes that the basic multiplication operator +// yields constant-time code. +// +// The mul31_lo() returns only the low 31 bits of the product. +// +// Note/Odin: +// The original BearSSL code provides alternative implemenetations +// of these routines gated behind `BR_CT_MUL31`, however that macro +// is only useful on Intel 80386/80486, VIA Nano 2000, and ARM7T/ARM9T. +@(require_results) +_mul31 :: #force_inline proc "contextless" (x, y: u32) -> (res: u64) { + return u64(x) * u64(y) +} + +@(private="file", require_results) +_mul31_lo :: #force_inline proc "contextless" (x, y: u32) -> (res: u32) { + return (x * y) & I31_MASK +} + +// Wrapper for `div_rem`; the remainder is returned, and the quotient is +// discarded. +@(private, require_results) +_rem_u32 :: #force_inline proc "contextless" (hi: u32, lo: u32, d: u32) -> (res: u32) { + _, rem := div_rem_u32(hi, lo, d) + return rem +} + +// Wrapper for `div_rem`; the quotient is returned, and the remainder is +// discarded. +@(private="file", require_results) +_div_u32 :: #force_inline proc "contextless" (hi: u32, lo: u32, d: u32) -> (quo: u32) { + q, _ := div_rem_u32(hi, lo, d) + return q +} + +// Constant-time division. The dividend `hi:lo` is divided by the divisor `d`; +// the quotient and remainder are returned. +// +// If `hi == d`, then the quotient does not fit on 32 bits; returned value is thus truncated. +// If `hi > d`, returned values are indeterminate. +@(require_results) +div_rem_u32 :: proc "contextless" (hi: u32, lo: u32, d: u32) -> (quo: u32, rem: u32) { + // TODO: optimize this + hi := hi + lo := lo + ch := subtle.eq(hi, d) + hi = subtle.csel(hi, 0, ch) + for k := uint(31); k > 0; k -= 1 { + j := 32 - k + w := (hi << j) | (lo >> k) + ctl := subtle.ge(w, d) | (hi >> k) + hi2 := (w - d) >> j + lo2 := lo - (d << k) + hi = subtle.csel(hi, hi2, ctl) + lo = subtle.csel(lo, lo2, ctl) + quo |= ctl << k + } + cf := subtle.ge(lo, d) | hi + quo |= cf + rem = subtle.csel(lo, lo - d, cf) + return +} + +// i31_rem computes x / y and returns the remainder. +@(require_results) +i31_rem :: proc "contextless" (x: []u32, y: u32) -> u32 { + words := uint(x[0] + 31) >> 5 + x_ := x[1:] + + r: u32 + for i := int(words-1); i >= 0; i -= 1 { + r = _rem_u32(r, x_[i], y) + } + + return r +} + +// Test whether an integer `x` is zero. +@(optimization_mode="none", require_results) +i31_is_zero :: proc "contextless" (x: []u32) -> (res: u32) { + z: u32 + + for u := (x[0] + 31) >> 5; u > 0; u -= 1 { + z |= x[u] + } + return ~(z | -z) >> 31 +} + +// Add `b` to `a` and return the `carry` (`0` or `1`). if `ctl` is `1`. +// If `ctl` is `0`, `a` is left alone but the `carry` will still be computed. +// +// The slices `a` and `b` MUST have the same announced bit length (in subscript `0`) +// +// `a` and `b` MAY be the same array, but partial overlap is not allowed. +@(require_results) +i31_add :: proc "contextless" (a: []u32, b: []u32, ctl: u32) -> (carry: u32) { + words := uint(a[0] + 63) >> 5 + for u in 1..> 31 + a[u] = subtle.csel(aw, naw & I31_MASK, ctl) + } + return +} + +// Subtract `b` from `a` and return the `carry` (`0` or `1`), if `ctl` is `1`. +// If `ctl` is `0`, then `a` is unmodified, but the carry is still computed +// and returned. +// +// The slices `a` and `b` MUST have the same announced bit length (in subscript `0`) +// +// `a` and `b` MAY be the same array, but partial overlap is not allowed. +@(require_results) +i31_sub :: proc "contextless" (a: []u32, b: []u32, ctl: u32) -> (carry: u32) { + words := uint(a[0] + 63) >> 5 + for u in 1..> 31 + a[u] = subtle.csel(aw, naw & I31_MASK, ctl) + } + return +} + +// Compute the ENCODED actual bit length of an integer `x`. +// The argument `x` should point to the first (least significant) +// value word of the integer. +// +// The upper bit of each value word MUST be `0`. +// +// Returned value is `((k / 31) << 5) + (k % 31)` if the bit length is `k`. +// +// CT: value or length of `x` does not leak. +@(require_results) +i31_bit_length :: proc "contextless" (x: []u32) -> (res: u32) { + tw, twk: u32 + + xlen := len(x) + for xlen > 0 { + xlen -= 1 + c := subtle.eq(tw, 0) + w := x[xlen] + + tw = subtle.csel(tw, w, c) + twk = subtle.csel(twk, u32(xlen), c) + } + return (twk << 5) + _u32_bit_length(tw) +} + +// Decode an integer from its big-endian unsigned representation. The +// "true" bit length of the integer is computed and set in the encoded +// announced bit length (`x[0]`), but all words of `x` corresponding to +// the full slice of source bytes. +// +// `x` needs to have a minimum length of: `1 + ((len(src) * 8) + 31) / 31` +// +// CT: value or length of `x` does not leak. +i31_decode :: proc "contextless" (x: []u32, src: []byte) { + u := len(src) - 1 + v := 1 + acc := u32(0) + acc_len := uint(0) + for u >= 0 { + b := u32(src[u]) + acc |= b << acc_len + acc_len += 8 + if acc_len >= 31 { + x[v] = acc & I31_MASK + acc_len -= 31 + acc = b >> (8 - acc_len) + v += 1 + } + u -= 1 + } + if acc_len != 0 { + x[v] = acc + v += 1 + } + x[0] = i31_bit_length(x[1:]) +} + +// Decode an integer from its big-endian unsigned representation. +// The integer MUST be lower than `m`; the (encoded) announced bit length +// written in `x` will be equal to that of `m`. All bytes from the +// `src` slice are read. +// +// Returned value is `1` if the decode value fits within the modulus, `0` +// otherwise. In the latter case, the `x` buffer will be set to `0` (but +// still with the announced bit length of `m`). +// +// CT: value or length of `x` does not leak. Memory access pattern depends +// only `src`'s length and the announced bit length of `m`. Whether `x` fits or +// not does not leak either. +@(require_results) +i31_decode_mod :: proc "contextless" (x: []u32, src: []byte, m: []u32) -> (res: u32) { + // Two-pass algorithm: in the first pass, we determine whether the + // value fits; in the second pass, we do the actual write. + // + // During the first pass, `res` contains the comparison result so far: + // 0x00000000 value is equal to the modulus + // 0x00000001 value is greater than the modulus + // 0xFFFFFFFF value is lower than the modulus + // + // Since we iterate starting with the least significant bytes (at + // the end of `src`), each new comparison overrides the previous + // except when the comparison yields 0 (equal). + // + // During the second pass, `res` is either 0xFFFFFFFF (value fits) 0x00000000 (value does not fit). + // We must iterate over all bytes of the source, _and_ possibly + // some extra virtual bytes (with value 0) so as to cover the + // complete modulus as well. We also add 4 such extra bytes beyond + // the modulus length because it then guarantees that no accumulated + // partial word remains to be processed. + _len := uint(len(src)) + mlen := uint((m[0] + 31) >> 5) + tlen := uint(mlen << 2) + if tlen < _len { + tlen = _len + } + tlen += 4 + + for pass in 0..<2 { + v := uint(1) + acc := u32(0) + acc_len := u32(0) + + for u in uint(0)..= 31 { + xw := acc & I31_MASK + acc_len -= 31 + + acc = b >> (8 - acc_len) + if v <= mlen { + if pass == 1 { + x[v] = res & xw + } else { + cc := u32(subtle.cmp(xw, m[v])) + res = subtle.csel(cc, res, subtle.eq(cc, 0)) + } + } else { + if pass == 0 { + res = subtle.csel(1, res, subtle.eq(xw, 0)) + } + } + v += 1 + } + } + + // When we reach this point at the end of the first pass: + // r is either 0, 1 or -1; we want to set r to 0 if it + // is equal to 0 or 1, and leave it to -1 otherwise. + // + // When we reach this point at the end of the second pass: + // r is either 0 or -1; we want to leave that value + // untouched. This is a subcase of the previous. + res >>= 1 + res |= (res << 1) + } + + x[0] = m[0] + + return res & 1 +} + +// Zeroize integer `x`. The announced bit length is set to the provided value, +// and the corresponding words are set to 0. The ENCODED bit length is expected +//here. +i31_zero :: proc "contextless" (x: []u32, bit_len: u32) { + x[0] = bit_len + intrinsics.mem_zero(raw_data(x[1:]), ((bit_len + 31) >> 5) * size_of(u32)) +} + +// Make a random integer of the provided size. The size is encoded. +// The header word is untouched. +i31_mkrand :: proc(x: []u32, esize: u32) { + _len := (esize + 31) >> 5 + x_ := slice.reinterpret([]byte, x) + crypto.rand_bytes(x_[4:4 + _len * size_of(u32)]) + for u in 1..<_len { + x[u] &= I31_MASK + } + m := _len & 31 + if m == 0 { + x[_len] &= I31_MASK + } else { + x[_len] &= I31_MASK >> (31 - m) + } +} + +// Right-shift an integer. The shift amount must be lower than 31 bits. +i31_rshift :: proc "contextless" (x: []u32, shift_amount: i32) { + _len := uint(x[0] + 31) >> 5 + if _len == 0 { + return + } + + count := uint(shift_amount) + + r := x[1] >> count + for u in 2..= _len { + w := u32(x[u]) + + x[u - 1] = ((w << (31 - count)) | r) & I31_MASK + r = w >> count + } + x[_len] = r +} + +// Reduce integer `a` modulo `m`. The result is written to `x`, +// and its announced bit length is set to be equal to that of `m`. +// +// `x` MUST be distinct from `a` and `m`. +// +// CT: only announced bit lengths leak, not values of `x`, `a` or `m`. +i31_reduce :: proc "contextless" (x: []u32, a: []u32, m: []u32) { + m_bitlen := m[0] + mlen := uint(m_bitlen + 31) >> 5 + + x[0] = m_bitlen + if m_bitlen == 0 { + return + } + + // If the source is shorter, then simply copy all words from a[] + // and zero out the upper words. + a_bitlen := a[0] + alen := uint(a_bitlen + 31) >> 5 + if a_bitlen < m_bitlen { + copy(x[1:], a[1:][:alen]) + for u in alen.. 0; u -= 1 { + i31_muladd_small(x, a[u], m) + } +} + +// Decode an integer from its big-endian unsigned representation, and +// reduce it modulo the provided modulus `m`. The announced bit length +// of the result is set to be equal to that of the modulus. +// +// `x` MUST be distinct from `m`. +i31_decode_reduce :: proc "contextless" (x: []u32, src: []byte, m: []u32) { + // Get the encoded bit length. + m_ebitlen := m[0] + + // Special case for an invalid (null) modulus. + if m_ebitlen == 0 { + x[0] = 0 + return + } + + // Clear the destination. + i31_zero(x, m_ebitlen) + + // First decode directly as many bytes as possible. + // This requires computing the actual bit length. + m_rbitlen := m_ebitlen >> 5 + m_rbitlen = (m_ebitlen & 31) + (m_rbitlen << 5) - m_rbitlen + + mblen := uint(m_rbitlen + 7) >> 3 + k := mblen - 1 + _len := uint(len(src)) + + if k >= _len { + i31_decode(x, src) + x[0] = m_ebitlen + return + } + + i31_decode(x, src[:k]) + x[0] = m_ebitlen + + // Input remaining bytes, using 31-bit words. + acc := u32(0) + acc_len := uint(0) + + for { + v := u32(src[k]) + + if acc_len >= 23 { + acc_len -= 23 + acc <<= (8 - acc_len) + acc |= v >> acc_len + i31_muladd_small(x, acc, m) + acc = v & (0xFF >> (8 - acc_len)) + } else { + acc = (acc << 8) | v + acc_len += 8 + } + + if k += 1; k >= _len { + break + } + } + + // We may have some bits accumulated. We then perform a shift to + // be able to inject these bits as a full 31-bit word. + if acc_len != 0 { + acc = (acc | (x[1] << acc_len)) & I31_MASK + i31_rshift(x, i32(31 - acc_len)) + i31_muladd_small(x, acc, m) + } +} + +// Multiply `x` by 2^31 and then add integer `z`, modulo `m`. +// This function assumes that `x` and `m` have the same announced bit +// length, the announced bit length of `m` matches its true bit length. +// +// `x` and `m` MUST be distinct arrays. +// `z` MUST fit in 31 bits (upper bit set to 0). +// +// CT: only the common announced bit length of `x` and `m` leaks, not +// the values of `x`, `z` or `m`. +i31_muladd_small :: proc "contextless" (x: []u32, z: u32, m: []u32) { + // We can test on the modulus bit length since we accept to leak + // that length. + m_bitlen := m[0] + if m_bitlen == 0 { + return + } + hi: u32 + if m_bitlen <= 31 { + hi = x[1] >> 1 + lo := (x[1] << 31) | z + x[1] = _rem_u32(hi, lo, m[1]) + return + } + mlen := uint(m_bitlen + 31) >> 5 + mblr := uint(m_bitlen) & 31 + + // Principle: we estimate the quotient (x*2^31+z)/m by + // doing a 64/32 division with the high words. + // + // Let: + // w = 2^31 + // a = (w*a0 + a1) * w^N + a2 + // b = b0 * w^N + b2 + // such that: + // 0 <= a0 < w + // 0 <= a1 < w + // 0 <= a2 < w^N + // w/2 <= b0 < w + // 0 <= b2 < w^N + // a < w*b + // I.e. the two top words of a are a0:a1, the top word of b is + // b0, we ensured that b0 is "full" (high bit set), and a is + // such that the quotient q = a/b fits on one word (0 <= q < w). + // + // If a = b*q + r (with 0 <= r < q), we can estimate q by + // doing an Euclidean division on the top words: + // a0*w+a1 = b0*u + v (with 0 <= v < b0) + // Then the following holds: + // 0 <= u <= w + // u-2 <= q <= u + hi = x[mlen] + a0, a1, b0: u32 + if mblr == 0 { + a0 = x[mlen] + intrinsics.mem_copy(raw_data(x[2:]), raw_data(x[1:]), (mlen - 1) * size_of(u32)) + x[1] = z + a1 = x[mlen] + b0 = m[mlen] + } else { + a0 = ((x[mlen] << (31 - mblr)) | (x[mlen - 1] >> mblr)) & I31_MASK + intrinsics.mem_copy(raw_data(x[2:]), raw_data(x[1:]), (mlen - 1) * size_of(u32)) + x[1] = z + a1 = ((x[mlen] << (31 - mblr)) | (x[mlen - 1] >> mblr)) & I31_MASK + b0 = ((m[mlen] << (31 - mblr)) | (m[mlen - 1] >> mblr)) & I31_MASK + } + + // We estimate a divisor q. If the quotient returned by div() + // is g: + // -- If a0 == b0 then g == 0; we want q = 0x7FFFFFFF. + // -- Otherwise: + // -- if g == 0 then we set q = 0; + // -- otherwise, we set q = g - 1. + // The properties described above then ensure that the true + // quotient is q-1, q or q+1. + // + // Take care that a0, a1 and b0 are 31-bit words, not 32-bit. We + // must adjust the parameters to br_div() accordingly. + g := _div_u32(a0 >> 1, a1 | (a0 << 31), b0) + q := subtle.csel(subtle.csel(g - 1, 0, subtle.eq(g, 0)), I31_MASK, subtle.eq(a0, b0)) + + // We subtract q*m from x (with the extra high word of value 'hi'). + // Since q may be off by 1 (in either direction), we may have to + // add or subtract m afterwards. + // + // The 'tb' flag will be true (1) at the end of the loop if the + // result is greater than or equal to the modulus (not counting + // 'hi' or the carry). + cc := u32(0) + tb := u32(1) + for u in 1..= mlen { + mw := m[u] + zl := _mul31(mw, q) + u64(cc) + cc = u32(zl >> 31) + zw := u32(zl) & I31_MASK + xw := x[u] + nxw := xw - zw + cc += nxw >> 31 + nxw &= I31_MASK + x[u] = nxw + tb = subtle.csel(subtle.gt(nxw, mw), tb, subtle.eq(nxw, mw)) + } + + // If we underestimated q, then either cc < hi (one extra bit + // beyond the top array word), or cc == hi and tb is true (no + // extra bit, but the result is not lower than the modulus). In + // these cases we must subtract m once. + // + // Otherwise, we may have overestimated, which will show as + // cc > hi (thus a negative result). Correction is adding m once. + over := subtle.gt(cc, hi) + under := ~over & (tb | subtle.lt(cc, hi)) + _ = i31_add(x, m, over) + _ = i31_sub(x, m, under) +} + +// Encode an integer into its big-endian unsigned representation. The +// output length in bytes is provided (parameter 'len'); if the length +// is too short then the integer is appropriately truncated; if it is +// too long then the extra bytes are set to 0. +i31_encode :: proc "contextless" (dst: []byte, x: []u32) { + xlen := uint(x[0] + 31) >> 5 + if xlen == 0 { + intrinsics.mem_zero(raw_data(dst[:]), len(dst) * size_of(u32)) + return + } + _len := uint(len(dst)) + k := uint(1) + acc := u32(0) + acc_len := uint(0) + for _len != 0 { + w := (k <= xlen) ? x[k] : 0 + k += 1 + if (acc_len == 0) { + acc = w + acc_len = 31 + } else { + z := acc | (w << acc_len) + acc_len -= 1 + acc = w >> (31 - acc_len) + if _len >= 4 { + _len -= 4 + ptr := (^u32be)(raw_data(dst[_len:])) + intrinsics.unaligned_store(ptr, u32be(z)) + } else { + switch _len { + case 3: + dst[_len - 3] = byte(z >> 16) + fallthrough + case 2: + dst[_len - 2] = byte(z >> 8) + fallthrough + case 1: + dst[_len - 1] = byte(z) + } + return + } + } + } +} + +// Compute `-(1/x) % 2^31`. If `x` is even, then this function returns `0`. +i31_ninv31 :: proc "contextless" (x: u32) -> (y: u32) { + y = 2 - x + y *= 2 - y * x + y *= 2 - y * x + y *= 2 - y * x + y *= 2 - y * x + return subtle.csel(0, -y, x & 1) & I31_MASK +} + +// Compute a modular Montgomery multiplication. `d` is filled with the +// value of `x*y/R % m` (where `R` is the Montgomery factor). +// +// The array `d` MUST be distinct from `x`, `y` and `m`[]. +// `x` and `y` MUST be numerically lower than `m`. +// +// `x` and `y` MAY be the same array. +// +// The `m0i` parameter is equal to `-(1/m0) mod 2^31`, where `m0` is the least +// significant value word of `m` (this works only if `m` is an odd integer). +i31_montymul :: proc "contextless" (d: []u32, x: []u32, y: []u32, m: []u32, m0i: u32) { + // Each outer loop iteration computes: + // `d <- (d + xu*y + f*m) / 2^31` + // We have `xu <= 2^31-1` and `f <= 2^31-1`. + // Thus, if `d <= 2*m-1` on input, then: + // `2*m-1 + 2*(2^31-1)*m <= (2^32)*m-1` + // and the new `d` value is less than `2*m`. + // + // We represent `d` over 31-bit words, with an extra word `dh`, + // which can thus be only 0 or 1. + _len := uint((m[0] + 31) >> 5) + len4 := _len & ~uint(3) + i31_zero(d, m[0]) + dh := u32(0) + for u in 0..<_len { + // The carry for each operation fits on 32 bits: + // `d[v+1] <= 2^31-1` + // `xu*y[v+1] <= (2^31-1)*(2^31-1)` + // `f*m[v+1] <= (2^31-1)*(2^31-1)` + // `r <= 2^32-1` + // `(2^31-1) + 2*(2^31-1)*(2^31-1) + (2^32-1) = 2^63 - 2^31` + // + // After division by `2^31`, the new `r` is then at most `2^32-1` + // + // Using a 32-bit carry has performance benefits on 32-bit + // systems; however, on 64-bit architectures, we prefer to + // keep the carry (r) in a 64-bit register, thus avoiding some + // "clear high bits" operations. + xu := x[u + 1] + f := _mul31_lo((d[1] + _mul31_lo(xu, y[1])), m0i) + + r := u64(0) + v := uint(0) + for ; v < len4; v += 4 { + z := u64(d[v + 1]) + _mul31(xu, y[v + 1]) + _mul31(f, m[v + 1]) + r + r = z >> 31 + d[v + 0] = u32(z) & I31_MASK + z = u64(d[v + 2]) + _mul31(xu, y[v + 2]) + _mul31(f, m[v + 2]) + r + r = z >> 31 + d[v + 1] = u32(z) & I31_MASK + z = u64(d[v + 3]) + _mul31(xu, y[v + 3]) + _mul31(f, m[v + 3]) + r + r = z >> 31 + d[v + 2] = u32(z) & I31_MASK + z = u64(d[v + 4]) + _mul31(xu, y[v + 4]) + _mul31(f, m[v + 4]) + r + r = z >> 31 + d[v + 3] = u32(z) & I31_MASK + } + for ; v < _len; v += 1 { + z := u64(d[v + 1]) + _mul31(xu, y[v + 1]) + _mul31(f, m[v + 1]) + r + r = z >> 31 + d[v] = u32(z) & I31_MASK + } + + // Since the new `dh` can only be `0` or `1`, the addition of + // the old dh with the carry MUST fit on 32 bits, and + // thus can be done into dh itself. + dh += u32(r) + d[_len] = dh & I31_MASK + dh >>= 31 + } + + // We must write back the bit length because it was overwritten in + // the loop (not overwriting it would require a test in the loop, + // which would yield bigger and slower code). + d[0] = m[0] + + // `d` may still be greater than `m` at that point; notably, the `dh` + // word may be non-zero. + _ = i31_sub(d, m, subtle.neq(dh, 0) | subtle.not(i31_sub(d, m, 0))) +} + +// Convert a modular integer to Montgomery representation. +// +// The integer `x` MUST be lower than `m`, but with the same announced bit length. +i31_to_monty :: proc "contextless" (x: []u32, m: []u32) { + // uint32_t k; + for k := (m[0] + 31) >> 5; k > 0; k -= 1 { + i31_muladd_small(x, 0, m) + } +} + +// Convert a modular integer back from Montgomery representation. +// +// The integer `x` MUST be lower than `m`[], but with the same announced bit +// length. +// +// The `m0i` parameter is equal to `-(1/m0) mod 2^32`, where `m0` is the least +// significant value word of `m` (this works only if `m` is an odd integer). +i31_from_monty :: proc "contextless" (x: []u32, m: []u32, m0i: u32) { + _len := uint(m[0] + 31) >> 5 + for _ in 0..<_len { + f := _mul31_lo(x[1], m0i) + cc := u64(0) + for v in 0..<_len { + z := u64(x[v + 1]) + _mul31(f, m[v + 1]) + cc + cc = z >> 31 + if v != 0 { + x[v] = u32(z & I31_MASK) + } + } + x[_len] = u32(cc) + } + + // We may have to do an extra subtraction, but only if the value in `x` + // is indeed greater than or equal to that of `m`, which is why we must + // do two calls: + // - First call computes the carry + // - Second call performs the subtraction only if the carry is 0). + _ = i31_sub(x, m, subtle.not(i31_sub(x, m, 0))) +} + +// Compute a modular exponentiation. +// +// `x` MUST be an integer modulo `m` (same announced bit length, lower value). +// `m` MUST be odd. +// +// The exponent `e` is in big-endian unsigned notation. +// +// The `m0i` parameter is equal to `-(1/m0) mod 2^31`, where `m0` is the least +// significant value word of `m` (this works only if `m` is an odd integer). +// +// The `t1` and `t2` parameters must be temporary arrays, each large enough to +// accommodate an integer with the same size as `m`. +i31_modpow :: proc "contextless" (x: []u32, e: []byte, m: []u32, m0i: u32, t1: []u32, t2: []u32) { + // `mlen` is the length of `m` expressed in `u32`'s (including the + // "bit length" first field). + mlen := uint((m[0] + 63) >> 5) + elen := u32(len(e)) + + // Throughout the algorithm: + // -- `t1` is in Montgomery representation; it contains x, x^2, x^4, x^8... + // -- The result is accumulated, in normal representation, in the `x` array. + // -- `t2` is used as destination buffer for each multiplication. + // + // Note that there is no need to call `i32_from_monty()`. + copy(t1[:mlen], x[:mlen]) + i31_to_monty(t1, m) + i31_zero(x, m[0]) + x[1] = 1 + for k := u32(0); k < (elen << 3); k += 1 { + ctl := (e[elen - 1 - (k >> 3)] >> (k & 7)) & 1 + + i31_montymul(t2, x, t1, m, m0i) + + for &d, i in x[:mlen] { + d = subtle.csel(d, t2[i], ctl) + } + + i31_montymul(t2, t1, t1, m, m0i) + copy(t1[:mlen], t2[:mlen]) + } +} + + +// Compute a modular exponentiation. +// +// `x` MUST be an integer modulo `m` (same announced bit length, lower value). +// `m` MUST be odd. +// +// The exponent `e` is in big-endian unsigned notation. +// +// The `m0i` parameter is equal to `-(1/m0) mod 2^31`, where `m0` is the least +// significant value word of `m`[] (this works only if m[] is an odd integer). +// +// The `tmp` array is used for temporaries; it must be large enough to accommodate +// at least two temporary values with the same size as `m` (including the leading +// "bit length" word). +// +// If there is room for more temporaries, then this function may use the extra +// room for window-based optimisation, resulting in faster computations. +// +// Returned value is `true` on success, `false` on error. An error is reported if +// the provided `tmp`array is too short. +i31_modpow_opt :: proc "contextless" (x: []u32, e: []byte, m: []u32, m0i: u32, tmp: []u32) -> u32 { + // NOTE/yawning: This is only used by the rsa_i31 code, with the key + // generation taking a function pointer to either this routine, + // or the i62 variant. + // + // If we ever need to support the i32 version, it is used extensively, + // but non e-waste architecutures will all do the right thing with + // the i62 version, albeit with a perforance hit on 32-bit CPUs. + + unimplemented_contextless() + + // i31_mod_pow(x, e, m, m0i, tmp[:len(m)], tmp[len(m):]) + // return 1 +} + +// Compute `d+a*b`, result in `d`. +// +// The initial announced bit length of `d` MUST match that of `a`[]. +// +// The `d` array MUST be large enough to accommodate the full result, +// plus (possibly) an extra word. The resulting announced bit length +// of `d` will be the sum of the announced bit lengths of `a` and `b` +// (therefore, it may be larger than the actual bit length of the numerical result). +// +// `a` and `b` may be the same array. `d` must be disjoint from both `a` and `b`. +i31_mulacc :: proc "contextless" (d: []u32, a: []u32, b: []u32) { + a_len := uint((a[0] + 31) >> 5) + b_len := uint((b[0] + 31) >> 5) + + // We want to add the two bit lengths, but these are encoded, + // which requires some extra care. + d_l := (a[0] & 31) + (b[0] & 31) + d_h := (a[0] >> 5) + (b[0] >> 5) + d[0] = (d_h << 5) + d_l + (~u32(d_l - 31) >> 31) + + for u in 0..> 31 + d[1 + u + v] = u32(z) & I31_MASK + } + d[1 + u + a_len] = u32(cc) + } +} diff --git a/core/crypto/_bigint/i62.odin b/core/crypto/_bigint/i62.odin new file mode 100644 index 000000000..eb6d99907 --- /dev/null +++ b/core/crypto/_bigint/i62.odin @@ -0,0 +1,361 @@ +package _bigint + +// Copyright (c) 2017 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "base:intrinsics" +import "core:math/bits" +import subtle "core:crypto/_subtle" +import "core:slice" + +@(private="file") +I62_MASK :: 0x3fff_ffff_ffff_ffff + +// Compute x*y+v1+v2. Operands are 64-bit, and result is 128-bit, with +// high word in "hi" and low word in "lo". +@(private="file", require_results) +_fma1 :: #force_inline proc "contextless" (x, y, v1, v2: u64) -> (hi, lo: u64) { + hi, lo = bits.mul_u64(x, y) + + carry: u64 + lo, carry = bits.add_u64(lo, v1, 0) + hi += carry + + lo, carry = bits.add_u64(lo, v2, 0) + hi += carry + + return +} + +// Compute x1*y1+x2*y2+v1+v2. Operands are 64-bit, and result is 128-bit, +// with high word in "hi" and low word in "lo". +// +// Callers should ensure that the two inner products, and the v1 and v2 +// operands, are multiple of 4 (this is not used by this specific definition +// but may help other implementations). +@(private="file", require_results) +_fma2 :: #force_inline proc "contextless" (x1, y1, x2, y2, v1, v2: u64) -> (hi, lo: u64) { + hi_1, lo_1 := bits.mul_u64(x1, y1) + hi_2, lo_2 := bits.mul_u64(x2, y2) + + carry: u64 + lo, carry = bits.add_u64(lo_1, lo_2, 0) + hi, _ = bits.add_u64(hi_1, hi_2, carry) + + lo, carry = bits.add_u64(lo, v1, 0) + hi += carry + + lo, carry = bits.add_u64(lo, v2, 0) + hi += carry + + return +} + +@(private="file", require_results) +_mul62_lo :: #force_inline proc "contextless" (x, y: u64) -> u64 { + return (x * y) & I62_MASK +} + +// Subtract b from a, and return the final carry. If 'ctl32' is 0, then +// a[] is kept unmodified, but the final carry is still computed and +// returned. +@(private="file", require_results) +_i62_sub :: proc "contextless" (a, b: []u64, num: int, ctl32: u32) -> u32 { + cc: u64 + + ctl := -ctl32 + mask := u64(ctl) | (u64(ctl) << 32) + for u in 0..> 63 + dw &= I62_MASK + a[u] = aw ~ (mask & (dw ~ aw)) + } + + return u32(cc) +} + +// Montgomery multiplication, over arrays of 62-bit values. The +// destination array (d) must be distinct from the other operands +// (x, y and m). All arrays are in little-endian format (least +// significant word comes first) over 'num' words. +@(private="file") +_i62_montymul :: proc "contextless" (d, x, y, m: []u64, num: int, m0i: u64) { + dh: u64 + + num4 := 1 + u64((num - 1) & ~int(3)) + intrinsics.mem_zero(raw_data(d), num * size_of(u64)) + for u in 0..> 62) + d[v - 1] = lo >> 2 + hi, lo = _fma2(xu, y[v + 1], f, m[v + 1], d[v + 1] << 2, r << 2) + r = hi + (r >> 62) + d[v + 0] = lo >> 2 + hi, lo = _fma2(xu, y[v + 2], f, m[v + 2], d[v + 2] << 2, r << 2) + r = hi + (r >> 62) + d[v + 1] = lo >> 2 + hi, lo = _fma2(xu, y[v + 3], f, m[v + 3], d[v + 3] << 2, r << 2) + r = hi + (r >> 62) + d[v + 2] = lo >> 2 + } + for ; v < num; v += 1 { + hi, lo = _fma2(xu, y[v], f, m[v], d[v] << 2, r << 2) + r = hi + (r >> 62) + d[v - 1] = lo >> 2 + } + + zh := dh + r + d[num - 1] = zh & I62_MASK + dh = zh >> 62 + } + _ = _i62_sub(d, m, num, u32(dh) | subtle.not(_i62_sub(d, m, num, 0))) +} + +// Conversion back from Montgomery representation. +@(private="file") +_i62_frommonty :: proc "contextless" (x, m: []u64, num: int, m0i: u64) { + for _ in 0..> 2 + } + } + x[num - 1] = cc >> 2 + } + _ = _i62_sub(x, m, num, subtle.not(_i62_sub(x, m, num, 0))) +} + +// Variant of i31_modpow_opt() that internally uses 64x64->128 +// multiplications. It expects the same parameters as i31_modpow_opt(), +// except that the temporaries should be 64-bit integers, not 32-bit +// integers. +i62_modpow_opt :: proc "contextless" (x31: []u32, e: []byte, m31: []u32, m0i31: u32, tmp: []u64) -> u32 { + twlen := len(tmp) + + // Get modulus size, in words. + mw31num := int((m31[0] + 31) >> 5) + mw62num := int((mw31num + 1) >> 1) + + // In order to apply this function, we must have enough room to + // copy the operand and modulus into the temporary array, along + // with at least two temporaries. If there is not enough room, + // switch to br_i31_modpow(). We also use br_i31_modpow() if the + // modulus length is not at least four words (94 bits or more). + if mw31num < 4 || mw62num << 2 > twlen { + // We assume here that we can split an aligned uint64_t + // into two properly aligned uint32_t. Since both types + // are supposed to have an exact width with no padding, + // then this property must hold. + + txlen := mw31num + 1 + if twlen < txlen { + return 0 + } + + tmp_as_u32s := slice.reinterpret([]u32, tmp) + t1, t2 := tmp_as_u32s[:txlen], tmp_as_u32s[txlen:] + + i31_modpow(x31, e, m31, m0i31, t1, t2) + + return 1 + } + + // Convert x to Montgomery representation: this means that + // we replace x with x*2^z mod m, where z is the smallest multiple + // of the word size such that 2^z >= m. We want to reuse the 31-bit + // functions here (for constant-time operation), but we need z + // for a 62-bit word size. + for _ in 0..> 1 + if u + 1 == mw31num { + m[v] = u64(m31[u + 1]) + x[v] = u64(x31[u + 1]) + } else { + m[v] = u64(m31[u + 1]) + (u64(m31[u + 2]) << 31) + x[v] = u64(x31[u + 1]) + (u64(x31[u + 2]) << 31) + } + } + + // Compute window size. We support windows up to 5 bits; for a + // window of size k bits, we need 2^k+1 temporaries (for k = 1, + // we use special code that uses only 2 temporaries). + win_len: int + for win_len = 5; win_len > 1; win_len -= 1 { + if (1 << uint(win_len) + 1) * mw62num <= twlen { + break + } + } + + t1 := tmp_[:mw62num] + t2 := tmp_[mw62num:] + + // Compute m0i, which is equal to -(1/m0) mod 2^62. We were + // provided with m0i31, which already fulfills this property + // modulo 2^31; the single expression below is then sufficient. + m0i := u64(m0i31) + m0i = _mul62_lo(m0i, 2 + _mul62_lo(m0i, m[0])) + + // Compute window contents. If the window has size one bit only, + // then t2 is set to x; otherwise, t2[0] is left untouched, and + // t2[k] is set to x^k (for k >= 1). + if win_len == 1 { + copy(t2, x) + } else { + copy(t2[mw62num:], x) + + base := t2[mw62num:] + for u := 2; u < 1 << uint(win_len); u += 1 { + _i62_montymul(base[mw62num:], base, x, m, mw62num, m0i) + base = base[mw62num:] + } + } + + // Set x to 1, in Montgomery representation. We again use the + // 31-bit code. + i31_zero(x31, m31[0]) + x31[(m31[0] + 31) >> 5] = 1 + i31_muladd_small(x31, 0, m31) + if mw31num & 1 != 0 { + i31_muladd_small(x31, 0, m31) + } + for u := 0; u < mw31num; u+= 2 { + v := u >> 1 + if u + 1 == mw31num { + x[v] = u64(x31[u + 1]) + } else { + x[v] = u64(x31[u + 1]) + (u64(x31[u + 2]) << 31) + } + } + + e_, e_len := e, len(e) + // We process bits from most to least significant. At each + // loop iteration, we have acc_len bits in acc. + acc: u32 + acc_len: uint + for acc_len > 0 || e_len > 0 { + // Get the next bits. + k := uint(win_len) + if acc_len < uint(win_len) { + if e_len > 0 { + acc = (acc << 8) | u32(e_[0]) + e_ = e_[1:] + e_len -= 1 + acc_len += 8 + } else { + k = acc_len + } + } + bits := (acc >> (acc_len - k)) & ((u32(1) << k) - 1) + acc_len -= k + + // We could get exactly k bits. Compute k squarings. + for _ in 0.. 1 { + intrinsics.mem_zero(raw_data(t2), mw62num * size_of(u64)) + + base := t2[mw62num:] + for u := u32(1); u < u32(1) << k; u += 1 { + mask := -u64(subtle.eq(u, bits)) + for v in 0..> 1]) + x31[u + 1] = u32(zw) & I31_MASK + if u + 1 < mw31num { + x31[u + 2] = u32(zw >> 31) + } + } + + return 1 +} + +// Wrapper for i62_modpow_opt() that uses the same type as +// i31_modpow_opt(); however, it requires its 'tmp' argument to the +// 64-bit aligned. +i62_modpow_opt_as_i31 :: proc "contextless" (x31: []u32, e: []byte, m31: []u32, m0i31: u32, tmp: []u32) -> u32 { + // As documented, this function expects the 'tmp' argument to be + // 64-bit aligned. This is OK since this function is internal (it + // is not part of BearSSL's public API). + ensure_contextless(uintptr(raw_data(tmp)) & 7 == 0) + ensure_contextless(len(tmp) & 1 == 0) // Length MUST be even. + + tmp_as_u64s := slice.reinterpret([]u64, tmp) + + return i62_modpow_opt(x31, e, m31, m0i31, tmp_as_u64s) +} diff --git a/core/crypto/_bigint/i62_primes.odin b/core/crypto/_bigint/i62_primes.odin new file mode 100644 index 000000000..a9606db05 --- /dev/null +++ b/core/crypto/_bigint/i62_primes.odin @@ -0,0 +1,265 @@ +package _bigint + +// Copyright (c) 2017 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import subtle "core:crypto/_subtle" +import "core:math/big" +import "core:slice" + +// Perform trial divisions on a candidate prime. We opt for the simple +// route and "just" compute a series of trial divisions. +// +// Returned value is 1 on success (none of the small primes +// divides x), 0 on error (a non-trivial GCD is obtained). +@(private="file", require_results) +trial_divisions :: proc "contextless" (x: []u32) -> u32 { + for factor in big._private_prime_table { + if factor <= 11 { + continue + } + if i31_rem(x, u32(factor)) == 0 { + return 0 + } + } + + return 1 +} + +// Perform n rounds of Miller-Rabin on the candidate prime x. This +// function assumes that x = 3 mod 4. +// +// WARNING: t MUST be 64-bit aligned, and be large enough such that +// it can hold 4 encoded integers that have the same number of limbs +// as x. +// +// Returned value is 1 on success (all rounds completed successfully), +// 0 otherwise. +@(private="file", require_results) +i62_miller_rabin :: proc(x: []u32, n: int, t: []u32) -> u32 { + // Since x = 3 mod 4, the Miller-Rabin test is simple: + // - get a random base a (such that 1 < a < x-1) + // - compute z = a^((x-1)/2) mod x + // - if z != 1 and z != x-1, the number x is composite + // + // We generate bases 'a' randomly with a size which is + // one bit less than x, which ensures that a < x-1. It + // is not useful to verify that a > 1 because the probability + // that we get a value a equal to 0 or 1 is much smaller + // than the probability of our Miller-Rabin tests not to + // detect a composite, which is already quite smaller than the + // probability of the hardware misbehaving and return a + // composite integer because of some glitch (e.g. bad RAM + // or ill-timed cosmic ray). + + // Compute (x-1)/2 (encoded). + xm1d2 := slice.reinterpret([]byte, t) + xm1d2_len := ((x[0] - (x[0] >> 5)) + 7) >> 3 + i31_encode(xm1d2[:xm1d2_len], x) + cc: u32 + for u in 0..> 1) | cc) + cc = w << 7 + } + + // We used some words of the provided buffer for (x-1)/2. + xm1d2_len_u32 := (xm1d2_len + 3) >> 2 + t_ := t[xm1d2_len_u32:] + tlen := len(t_) + + xlen := (x[0] + 31) >> 5 + asize := x[0] - 1 - subtle.eq0(x[0] & 31) + x0i := i31_ninv31(x[1]) + for _ in 0..> 5 + + for { + // Generate random bits. We force the two top bits and the + // two bottom bits to 1. + i31_mkrand(x, esize) + if (esize & 31) == 0 { + x[_len] |= 0x60000000 + } else if (esize & 31) == 1 { + x[_len] |= 0x00000001 + x[_len - 1] |= 0x40000000 + } else { + x[_len] |= 0x00000003 << ((esize & 31) - 2) + } + x[1] |= 0x00000003 + + // Trial division with low primes (3, 5, 7 and 11). We + // use the following properties: + // + // 2^2 = 1 mod 3 + // 2^4 = 1 mod 5 + // 2^3 = 1 mod 7 + // 2^10 = 1 mod 11 + m3, m5, m7, m11: u32 + s7, s11: uint + for u in 0..<_len { + w := x[1 + u] + w3 := (w & 0xFFFF) + (w >> 16) // max: 98302 + w5 := (w & 0xFFFF) + (w >> 16) // max: 98302 + w7 := (w & 0x7FFF) + (w >> 15) // max: 98302 + w11 := (w & 0xFFFFF) + (w >> 20) // max: 1050622 + + m3 += w3 << (u & 1) + m3 = (m3 & 0xFF) + (m3 >> 8) // max: 1025 + + m5 += w5 << ((4 - u) & 3) + m5 = (m5 & 0xFFF) + (m5 >> 12) // max: 4479 + + m7 += w7 << s7 + m7 = (m7 & 0x1FF) + (m7 >> 9) // max: 1280 + s7 += 1 + if s7 == 3 { + s7 = 0 + } + + m11 += w11 << s11 + s11 += 1 + if s11 == 10 { + s11 = 0 + } + m11 = (m11 & 0x3FF) + (m11 >> 10) // max: 526847 + } + + m3 = (m3 & 0x3F) + (m3 >> 6) // max: 78 + m3 = (m3 & 0x0F) + (m3 >> 4) // max: 18 + m3 = ((m3 * 43) >> 5) & 3 + + m5 = (m5 & 0xFF) + (m5 >> 8) // max: 271 + m5 = (m5 & 0x0F) + (m5 >> 4) // max: 31 + m5 -= 20 & -subtle.gt(m5, 19) + m5 -= 10 & -subtle.gt(m5, 9) + m5 -= 5 & -subtle.gt(m5, 4) + + m7 = (m7 & 0x3F) + (m7 >> 6) // max: 82 + m7 = (m7 & 0x07) + (m7 >> 3) // max: 16 + m7 = ((m7 * 147) >> 7) & 7 + + // 2^5 = 32 = -1 mod 11. + m11 = (m11 & 0x3FF) + (m11 >> 10) // max: 1536 + m11 = (m11 & 0x3FF) + (m11 >> 10) // max: 1023 + m11 = (m11 & 0x1F) + 33 - (m11 >> 5) // max: 64 + m11 -= 44 & -subtle.gt(m11, 43) + m11 -= 22 & -subtle.gt(m11, 21) + m11 -= 11 & -subtle.gt(m11, 10) + + // If any of these modulo is 0, then the candidate is + // not prime. Also, if pubexp is 3, 5, 7 or 11, and the + // corresponding modulus is 1, then the candidate must + // be rejected, because we need e to be invertible + // modulo p-1. We can use simple comparisons here + // because they won't leak information on a candidate + // that we keep, only on one that we reject (and is thus + // not secret). + if m3 == 0 || m5 == 0 || m7 == 0 || m11 == 0 { + continue + } + if (pubexp == 3 && m3 == 1) || (pubexp == 5 && m5 == 1) || (pubexp == 7 && m7 == 1) || (pubexp == 11 && m11 == 1) { + continue + } + + // More trial divisions. + if trial_divisions(x) == 0 { + continue + } + + // Miller-Rabin algorithm. Since we selected a random + // integer, not a maliciously crafted integer, we can use + // relatively few rounds to lower the risk of a false + // positive (i.e. declaring prime a non-prime) under + // 2^(-80). It is not useful to lower the probability much + // below that, since that would be substantially below + // the probability of the hardware misbehaving. Sufficient + // numbers of rounds are extracted from the Handbook of + // Applied Cryptography, note 4.49 (page 149). + // + // Since we work on the encoded size (esize), we need to + // compare with encoded thresholds. + rounds: int + switch { + case esize < 309: + rounds = 12 + case esize < 464: + rounds = 9 + case esize < 670: + rounds = 6 + case esize < 877: + rounds = 4 + case esize < 1341: + rounds = 3 + case: + rounds = 2 + } + + if i62_miller_rabin(x, rounds, t) == 1 { + return + } + } +} diff --git a/core/crypto/_fiat/field_p256r1/field.odin b/core/crypto/_fiat/field_p256r1/field.odin index f7dd978aa..fcb7357a7 100644 --- a/core/crypto/_fiat/field_p256r1/field.odin +++ b/core/crypto/_fiat/field_p256r1/field.odin @@ -71,7 +71,7 @@ fe_equal :: proc "contextless" (arg1, arg2: ^Montgomery_Domain_Field_Element) -> // This will only underflow if and only if (⟺) arg1 == arg2, and we return the borrow, // which will be 1. - is_eq := subtle.u64_is_zero(fe_non_zero(&tmp)) + is_eq := subtle.eq0(fe_non_zero(&tmp)) fe_clear(&tmp) diff --git a/core/crypto/_fiat/field_p384r1/field.odin b/core/crypto/_fiat/field_p384r1/field.odin index 2bddff18c..af371873e 100644 --- a/core/crypto/_fiat/field_p384r1/field.odin +++ b/core/crypto/_fiat/field_p384r1/field.odin @@ -77,7 +77,7 @@ fe_equal :: proc "contextless" (arg1, arg2: ^Montgomery_Domain_Field_Element) -> // This will only underflow if and only if (⟺) arg1 == arg2, and we return the borrow, // which will be 1. - is_eq := subtle.u64_is_zero(fe_non_zero(&tmp)) + is_eq := subtle.eq0(fe_non_zero(&tmp)) fe_clear(&tmp) diff --git a/core/crypto/_fiat/field_scalarp384r1/field.odin b/core/crypto/_fiat/field_scalarp384r1/field.odin index cfc27f322..d465bae44 100644 --- a/core/crypto/_fiat/field_scalarp384r1/field.odin +++ b/core/crypto/_fiat/field_scalarp384r1/field.odin @@ -60,7 +60,7 @@ fe_from_bytes :: proc "contextless" ( reduced[3], borrow = bits.sub_u64(tmp[3], ELL[3], borrow) reduced[4], borrow = bits.sub_u64(tmp[4], ELL[4], borrow) reduced[5], borrow = bits.sub_u64(tmp[5], ELL[5], borrow) - need_reduced := subtle.u64_is_zero(borrow) + need_reduced := subtle.eq0(borrow) fe_cond_select(&tmp, &tmp, &reduced, int(need_reduced)) fe_to_montgomery(out1, &tmp) diff --git a/core/crypto/_mldsa/constants.odin b/core/crypto/_mldsa/constants.odin new file mode 100644 index 000000000..6a9b6e096 --- /dev/null +++ b/core/crypto/_mldsa/constants.odin @@ -0,0 +1,72 @@ +#+private +package _mldsa + +CRHBYTES :: 64 +TRBYTES :: 64 + +N :: 256 +Q :: 8380417 +D :: 13 + +K_MAX :: 8 +L_MAX :: 7 + +POLYZ_PACKEDBYTES_MAX :: 640 + +POLYT1_PACKEDBYTES :: 320 +POLYT0_PACKEDBYTES :: 416 + +POLYVECT1_PACKEDBYTES_MAX :: K_MAX * POLYT1_PACKEDBYTES +POLYW1_PACKEDBYTES_MAX :: 192 + +CTILDBYTES_MAX :: 64 + +@(require_results) +polyeta_packedbytes :: #force_inline proc "contextless" (params: ^Params) -> int { + POLYETA_PACKEDBYTES_2 :: 96 + POLYETA_PACKEDBYTES_4 :: 128 + + switch params.eta { + case 2: + return POLYETA_PACKEDBYTES_2 + case 4: + return POLYETA_PACKEDBYTES_4 + case: + unreachable() + } +} + +@(require_results) +polyz_packedbytes :: #force_inline proc "contextless" (params: ^Params) -> int { + POLYZ_PACKEDBYTES_GAMMA1_17 :: 576 + POLYZ_PACKEDBYTES_GAMMA1_19 :: 640 + + switch params.gamma1 { + case 1 << 17: + return POLYZ_PACKEDBYTES_GAMMA1_17 + case 1 << 19: + return POLYZ_PACKEDBYTES_GAMMA1_19 + case: + unreachable() + } +} + +@(require_results) +polyw1_packedbytes :: #force_inline proc "contextless" (params: ^Params) -> int { + POLYW1_PACKEDBYTES_GAMMA2_95232 :: 192 + POLYW1_PACKEDBYTES_GAMMA2_261888 :: 128 + + switch params.gamma2 { + case (Q-1)/88: + return POLYW1_PACKEDBYTES_GAMMA2_95232 + case (Q-1)/32: + return POLYW1_PACKEDBYTES_GAMMA2_261888 + case: + unreachable() + } +} + +@(require_results) +polyvech_packedbytes :: #force_inline proc "contextless" (params: ^Params) -> int { + return params.omega + params.k +} diff --git a/core/crypto/_mldsa/dsa_internal.odin b/core/crypto/_mldsa/dsa_internal.odin new file mode 100644 index 000000000..bfb6b7508 --- /dev/null +++ b/core/crypto/_mldsa/dsa_internal.odin @@ -0,0 +1,394 @@ +package _mldsa + +import "core:crypto" +import "core:crypto/shake" + +// This implementation is derived from the PQ-CRYSTALS reference +// implementation [[ https://github.com/pq-crystals/dilithium ]], +// primarily for licensing reasons. Arguably mldsa-native is +// a more "up to date" codebase, but the changes to the +// ref code is minor and they slapped an attribution-required +// license on something that was originally CC-0/Apache 2.0. + +SEEDBYTES :: 32 +RNDBYTES :: 32 +CTXBYTES_MAX :: 255 + +Params :: struct { + k: int, + l: int, + eta: i32, + tau: int, + beta: i32, + gamma1: i32, + gamma2: i32, + omega: int, + ctild_bytes: int, +} + +@(rodata) +Params_44 := Params{ + k = 4, + l = 4, + eta = 2, + tau = 39, + beta = 78, + gamma1 = 1 << 17, + gamma2 = (Q-1)/88, + omega = 80, + ctild_bytes = 32, +} + +@(rodata) +Params_65 := Params{ + k = 6, + l = 5, + eta = 4, + tau = 49, + beta = 196, + gamma1 = 1 << 19, + gamma2 = (Q-1)/32, + omega = 55, + ctild_bytes = 48, +} + +@(rodata) +Params_87 := Params{ + k = 8, + l = 7, + eta = 2, + tau = 60, + beta = 120, + gamma1 = 1 << 19, + gamma2 = (Q-1)/32, + omega = 75, + ctild_bytes = 64, +} + +Private_Key :: struct { + params: ^Params, + + rho: [SEEDBYTES]byte, + tr: [TRBYTES]byte, + key: [SEEDBYTES]byte, + t0: Polyvec_K, + s1: Polyvec_L, + s2: Polyvec_K, + + pub_key: Public_Key, + seed: [SEEDBYTES]byte, +} + +Public_Key :: struct { + params: ^Params, + + t1: Polyvec_K, + rho: [SEEDBYTES]byte, + mu: [TRBYTES]byte, +} + +@(private) +Signature :: struct { + params: ^Params, + + c: [CTILDBYTES_MAX]byte, + z: Polyvec_L, + h: Polyvec_K, +} + +dsa_keygen_internal :: proc( + priv_key: ^Private_Key, + seed: []byte, + params: ^Params, +) { + ensure(len(seed) == SEEDBYTES, "crypto/mldsa: invalid seed") + + pub_key := &priv_key.pub_key + pub_key.params = params + priv_key.params = params + + copy(priv_key.seed[:], seed) + + seedbuf: [2*SEEDBYTES + CRHBYTES]byte = --- + mat_: [K_MAX]Polyvec_L = --- + defer crypto.zero_explicit(&seedbuf, size_of(seedbuf)) + defer crypto.zero_explicit(&mat_, size_of(mat_)) + + // Expand randomness for rho, rhoprime and key + copy(seedbuf[:], seed) + seedbuf[SEEDBYTES] = byte(params.k) + seedbuf[SEEDBYTES+1] = byte(params.l) + shake256(seedbuf[:], seedbuf[:SEEDBYTES+2]) + copy(priv_key.rho[:], seedbuf[:SEEDBYTES]) + rhoprime := seedbuf[SEEDBYTES:SEEDBYTES+CRHBYTES] + copy(priv_key.key[:], seedbuf[SEEDBYTES+CRHBYTES:]) + + // Expand matrix + mat := mat_[:params.k] + polyvec_matrix_expand(mat, priv_key.rho[:], params) + + // Sample short vectors s1 and s2 + polyvec_l_uniform_eta(&priv_key.s1, rhoprime, 0, params) + polyvec_k_uniform_eta(&priv_key.s2, rhoprime, u16(params.l), params) + + // Matrix-vector multiplication + s1hat: Polyvec_L = --- + defer crypto.zero_explicit(&s1hat, size_of(Polyvec_L)) + polyvec_copy(&s1hat, &priv_key.s1, params) + polyvec_l_ntt(&s1hat, params) + polyvec_matrix_pointwise_montgomery(&pub_key.t1, mat, &s1hat, params) + polyvec_k_reduce(&pub_key.t1, params) + polyvec_k_invntt_tomont(&pub_key.t1, params) + + // Add error vector s2 + polyvec_k_add(&pub_key.t1, &pub_key.t1, &priv_key.s2, params) + + // Extract t1 and write public key + pk_bytes_: [SEEDBYTES+POLYVECT1_PACKEDBYTES_MAX]byte = --- + pk_bytes := pk_bytes_[:public_key_size(params)] + polyvec_k_caddq(&pub_key.t1, params) + polyvec_k_power2round(&pub_key.t1, &priv_key.t0, &pub_key.t1, params) + copy(pub_key.rho[:], priv_key.rho[:]) + _ = pack_pk(pk_bytes, pub_key) + + // Compute H(rho, t1) and write secret key + shake256(pub_key.mu[:], pk_bytes) + copy(priv_key.tr[:], pub_key.mu[:]) +} + +dsa_sign_internal :: proc( + sig_bytes: []byte, + m: []byte, + ctx: []byte, + rnd: []byte, + priv_key: ^Private_Key, + external_mu: []byte = nil +) -> bool { + params := priv_key.params + switch params { + case &Params_44, &Params_65, &Params_87: + case: + return false + } + if len(sig_bytes) != signature_size(params) { + return false + } + ensure(len(ctx) <= CTXBYTES_MAX, "crypto/mlkem: invalid contxt size") + ensure(len(rnd) == RNDBYTES, "crypto/mlkem: invalid rnd size") + + mu, rhoprime: [CRHBYTES]byte = ---, --- + mat_: [K_MAX]Polyvec_L + w1_bytes_: [SEEDBYTES+POLYVECT1_PACKEDBYTES_MAX]byte = --- + s1, y: Polyvec_L = ---, --- + t0, s2, w1, w0: Polyvec_K = ---, ---, ---, --- + cp: Poly + + polyvec_copy(&s1, &priv_key.s1, params) + polyvec_copy(&s2, &priv_key.s2, params) + polyvec_copy(&t0, &priv_key.t0, params) + + defer crypto.zero_explicit(&mu, size_of(mu)) + defer crypto.zero_explicit(&rhoprime, size_of(rhoprime)) + defer crypto.zero_explicit(&mat_, size_of(mat_)) + defer crypto.zero_explicit(&w1_bytes_, size_of(w1_bytes_)) + defer polyvec_clear([]^Polyvec_L{&s1, &y}) + defer polyvec_clear([]^Polyvec_K{&t0, &s2, &w1, &w0}) + defer crypto.zero_explicit(&cp, size_of(cp)) + + sig: Signature = --- + sig.params = params + h := &sig.h + z := &sig.z + c := sig.c[:params.ctild_bytes] + + w1_bytes := w1_bytes_[:params.k*polyw1_packedbytes(params)] + + // Compute mu = CRH(tr, pre, msg) + if len(external_mu) == 0 { + // The FIPS publication handles the shake prefix + // in the public sign operation, but doing it + // here makes more sense. + ctx_buf: [2]byte + shake_ctx: shake.Context = --- + defer shake.reset(&shake_ctx) + + ctx_len := len(ctx) + + shake.init_256(&shake_ctx) + shake.write(&shake_ctx, priv_key.tr[:]) + if ctx_len > 0 { + ctx_buf[1] = byte(ctx_len) + } + shake.write(&shake_ctx, ctx_buf[:]) + if ctx_len > 0 { + shake.write(&shake_ctx, ctx) + } + shake.write(&shake_ctx, m) + shake.read(&shake_ctx, mu[:]) + } else { + ensure(len(external_mu) == CRHBYTES, "crypto/mlkem: invalid external mu") + copy(mu[:], external_mu) + } + + // Compute rhoprime = CRH(key, rnd, mu) + shake256(rhoprime[:], priv_key.key[:], rnd, mu[:]) + + // Expand matrix and transform vectors + mat := mat_[:params.k] + polyvec_matrix_expand(mat, priv_key.rho[:], params) + polyvec_l_ntt(&s1, params) + polyvec_k_ntt(&s2, params) + polyvec_k_ntt(&t0, params) + + // Rejection-sampling loop + iv: u32 // ref uses u16, but ML-DSA-87 will reuse the IV at p = ~2^{-23400} + for { + // Sample intermediate vector y + polyvec_l_uniform_gamma1(&y, rhoprime[:], iv, params) + iv += 1 + + // Matrix-vector multiplication + polyvec_copy(z, &y, params) + polyvec_l_ntt(z, params) + polyvec_matrix_pointwise_montgomery(&w1, mat, z, params) + polyvec_k_reduce(&w1, params) + polyvec_k_invntt_tomont(&w1, params) + + // Decompose w and call the random oracle + polyvec_k_caddq(&w1, params) + polyvec_k_decompose(&w1, &w0, &w1, params) + polyvec_k_pack_w1(w1_bytes, &w1, params) + + shake256(c, mu[:], w1_bytes) + poly_challenge(&cp, c, params) + poly_ntt(&cp) + + // Compute z, reject if it reveals secret + polyvec_l_pointwise_poly_montgomery(z, &cp, &s1, params) + polyvec_l_invntt_tomont(z, params) + polyvec_l_add(z, z, &y, params) + polyvec_l_reduce(z, params) + if polyvec_l_chknorm(z, params.gamma1 - params.beta, params) { + continue + } + + // Check that subtracting cs2 does not change high bits of w + // and low bits do not reveal secret information + polyvec_k_pointwise_poly_montgomery(h, &cp, &s2, params) + polyvec_k_invntt_tomont(h, params) + polyvec_k_sub(&w0, &w0, h, params) + polyvec_k_reduce(&w0, params) + if polyvec_k_chknorm(&w0, params.gamma2 - params.beta, params) { + continue + } + + // Compute hints for w1 + polyvec_k_pointwise_poly_montgomery(h, &cp, &t0, params) + polyvec_k_invntt_tomont(h, params) + polyvec_k_reduce(h, params) + if polyvec_k_chknorm(h, params.gamma2, params) { + continue + } + + polyvec_k_add(&w0, &w0, h, params) + n := polyvec_k_make_hint(h, &w0, &w1, params) + if n <= uint(params.omega) { + break + } + } + + // Write signature + return pack_sig(sig_bytes, &sig) +} + +dsa_verify_internal :: proc( + sig_bytes: []byte, + m: []byte, + ctx: []byte, + pub_key: ^Public_Key, +) -> bool { + ensure(len(ctx) <= CTXBYTES_MAX, "crypto/mlkem: invalid contxt size") + + params := pub_key.params + switch params { + case &Params_44, &Params_65, &Params_87: + case: + return false + } + + sig: Signature = --- + if !unpack_sig(&sig, sig_bytes, params) { + return false + } + if polyvec_l_chknorm(&sig.z, params.gamma1 - params.beta, params) { + return false + } + c := sig.c[:params.ctild_bytes] + z := &sig.z + h := &sig.h + + t1: Polyvec_K = --- + polyvec_copy(&t1, &pub_key.t1, params) + rho := pub_key.rho[:] + + // Compute CRH(H(rho, t1), pre, msg) + mu: [CRHBYTES]byte + { + // The FIPS publication handles the shake prefix + // in the public sign operation, but doing it + // here makes more sense. + ctx_buf: [2]byte + shake_ctx: shake.Context = --- + defer shake.reset(&shake_ctx) + + ctx_len := len(ctx) + + shake.init_256(&shake_ctx) + shake.write(&shake_ctx, pub_key.mu[:]) + if ctx_len > 0 { + ctx_buf[1] = byte(ctx_len) + } + shake.write(&shake_ctx, ctx_buf[:]) + if ctx_len > 0 { + shake.write(&shake_ctx, ctx) + } + shake.write(&shake_ctx, m) + shake.read(&shake_ctx, mu[:]) + } + + // Matrix-vector multiplication; compute Az - c2^dt1 + mat_: [K_MAX]Polyvec_L + w1: Polyvec_K = --- + cp: Poly = --- + mat := mat_[:params.l] + + poly_challenge(&cp, c, params) + polyvec_matrix_expand(mat, rho, params) + + polyvec_l_ntt(z, params) + polyvec_matrix_pointwise_montgomery(&w1, mat, z, params) + + poly_ntt(&cp) + polyvec_k_shiftl(&t1, params) + polyvec_k_ntt(&t1, params) + polyvec_k_pointwise_poly_montgomery(&t1, &cp, &t1, params) + + polyvec_k_sub(&w1, &w1, &t1, params) + polyvec_k_reduce(&w1, params) + polyvec_k_invntt_tomont(&w1, params) + + // Reconstruct w1 + buf_: [K_MAX*POLYW1_PACKEDBYTES_MAX]byte = --- + buf := buf_[:params.k*polyw1_packedbytes(params)] + polyvec_k_caddq(&w1, params) + polyvec_k_use_hint(&w1, &w1, h, params) + polyvec_k_pack_w1(buf, &w1, params) + + // Call random oracle and verify challenge + c2_: [CTILDBYTES_MAX]byte + c2 := c2_[:params.ctild_bytes] + shake256(c2, mu[:], buf) + + // Note/perf: Can be vartime + return crypto.compare_constant_time(c, c2) == 1 +} diff --git a/core/crypto/_mldsa/ntt.odin b/core/crypto/_mldsa/ntt.odin new file mode 100644 index 000000000..c36f88327 --- /dev/null +++ b/core/crypto/_mldsa/ntt.odin @@ -0,0 +1,75 @@ +#+private +package _mldsa + +@(rodata) +ZETAS := [N]i32 { + 0, 25847, -2608894, -518909, 237124, -777960, -876248, 466468, + 1826347, 2353451, -359251, -2091905, 3119733, -2884855, 3111497, 2680103, + 2725464, 1024112, -1079900, 3585928, -549488, -1119584, 2619752, -2108549, + -2118186, -3859737, -1399561, -3277672, 1757237, -19422, 4010497, 280005, + 2706023, 95776, 3077325, 3530437, -1661693, -3592148, -2537516, 3915439, + -3861115, -3043716, 3574422, -2867647, 3539968, -300467, 2348700, -539299, + -1699267, -1643818, 3505694, -3821735, 3507263, -2140649, -1600420, 3699596, + 811944, 531354, 954230, 3881043, 3900724, -2556880, 2071892, -2797779, + -3930395, -1528703, -3677745, -3041255, -1452451, 3475950, 2176455, -1585221, + -1257611, 1939314, -4083598, -1000202, -3190144, -3157330, -3632928, 126922, + 3412210, -983419, 2147896, 2715295, -2967645, -3693493, -411027, -2477047, + -671102, -1228525, -22981, -1308169, -381987, 1349076, 1852771, -1430430, + -3343383, 264944, 508951, 3097992, 44288, -1100098, 904516, 3958618, + -3724342, -8578, 1653064, -3249728, 2389356, -210977, 759969, -1316856, + 189548, -3553272, 3159746, -1851402, -2409325, -177440, 1315589, 1341330, + 1285669, -1584928, -812732, -1439742, -3019102, -3881060, -3628969, 3839961, + 2091667, 3407706, 2316500, 3817976, -3342478, 2244091, -2446433, -3562462, + 266997, 2434439, -1235728, 3513181, -3520352, -3759364, -1197226, -3193378, + 900702, 1859098, 909542, 819034, 495491, -1613174, -43260, -522500, + -655327, -3122442, 2031748, 3207046, -3556995, -525098, -768622, -3595838, + 342297, 286988, -2437823, 4108315, 3437287, -3342277, 1735879, 203044, + 2842341, 2691481, -2590150, 1265009, 4055324, 1247620, 2486353, 1595974, + -3767016, 1250494, 2635921, -3548272, -2994039, 1869119, 1903435, -1050970, + -1333058, 1237275, -3318210, -1430225, -451100, 1312455, 3306115, -1962642, + -1279661, 1917081, -2546312, -1374803, 1500165, 777191, 2235880, 3406031, + -542412, -2831860, -1671176, -1846953, -2584293, -3724270, 594136, -3776993, + -2013608, 2432395, 2454455, -164721, 1957272, 3369112, 185531, -1207385, + -3183426, 162844, 1616392, 3014001, 810149, 1652634, -3694233, -1799107, + -3038916, 3523897, 3866901, 269760, 2213111, -975884, 1717735, 472078, + -426683, 1723600, -1803090, 1910376, -1667432, -1104333, -260646, -3833893, + -2939036, -2235985, -420899, -2286327, 183443, -976891, 1612842, -3545687, + -554416, 3919660, -48306, -1362209, 3937738, 1400424, -846154, 1976782, +} + +ntt :: proc "contextless" (a: ^[N]i32) #no_bounds_check { + j, k := 0, 1 + for l := 128; l > 0; l >>= 1 { + for start := 0; start < N; start = j + l { + zeta := ZETAS[k] + k += 1 + for j = start; j < start + l; j += 1 { + t := montgomery_reduce(i64(zeta) * i64(a[j + l])) + a[j + l] = a[j] - t + a[j] = a[j] + t + } + } + } +} + +invntt_tomont :: proc "contextless" (a: ^[N]i32) #no_bounds_check { + F :: 41978 // mont^2/256 + + j, k := 0, 255 + for l := 1; l < N; l <<= 1 { + for start := 0; start < N; start = j + l { + zeta := -ZETAS[k] + k -= 1 + for j = start; j < start + l; j += 1 { + t := a[j] + a[j] = t + a[j + l] + a[j + l] = t - a[j + l] + a[j + l] = montgomery_reduce(i64(zeta) * i64(a[j + l])) + } + } + } + + for i in 0.. bool { + if len(pk_bytes) != public_key_size(pub_key.params) { + return false + } + + seed_bytes, t1_bytes := pk_bytes[:SEEDBYTES], pk_bytes[SEEDBYTES:] + + copy(seed_bytes, pub_key.rho[:]) + + for i in 0.. bool { + if len(pk_bytes) != public_key_size(params) { + return false + } + + seed_bytes, t1_bytes := pk_bytes[:SEEDBYTES], pk_bytes[SEEDBYTES:] + + pub_key.params = params + + copy(pub_key.rho[:], seed_bytes) + + for i in 0.. bool { + params := priv_key.params + if len(sk_bytes) != private_key_size(params) { + return false + } + + sk_bytes := sk_bytes + polyeta_len := polyeta_packedbytes(params) + + copy(sk_bytes, priv_key.rho[:]) + sk_bytes = sk_bytes[SEEDBYTES:] + + copy(sk_bytes, priv_key.key[:]) + sk_bytes = sk_bytes[SEEDBYTES:] + + copy(sk_bytes, priv_key.tr[:]) + sk_bytes = sk_bytes[TRBYTES:] + + for i in 0.. bool { + if len(sig_bytes) != signature_size(sig.params) { + return false + } + + sig_bytes := sig_bytes + polyz_len := polyz_packedbytes(sig.params) + + copy(sig_bytes, sig.c[:sig.params.ctild_bytes]) + sig_bytes = sig_bytes[sig.params.ctild_bytes:] + + for i in 0.. bool { + if len(sig_bytes) != signature_size(params) { + return false + } + + intrinsics.mem_zero(sig, size_of(Signature)) + + sig_bytes := sig_bytes + polyz_len := polyz_packedbytes(params) + omega := params.omega + + copy(sig.c[:], sig_bytes[:params.ctild_bytes]) + sig_bytes = sig_bytes[params.ctild_bytes:] + + for i in 0.. byte(omega) { + return false + } + + for j := k; j < int(sig_bytes[omega + i]); j += 1 { + // Coefficients are ordered for strong unforgeability + if j > k && sig_bytes[j] <= sig_bytes[j-1] { + return false + } + sig.h.vec[i].coeffs[sig_bytes[j]] = 1 + } + + k = int(sig_bytes[omega + i]) + } + + // Extra indices are zero for strong unforgeability + for j := k; j < omega; j += 1 { + if sig_bytes[j] != 0 { + return false + } + } + + sig.params = params + + return true +} + +@(private,require_results) +public_key_size :: #force_inline proc "contextless" (params: ^Params) -> int { + return SEEDBYTES + params.k * POLYT1_PACKEDBYTES +} + +@(private,require_results) +private_key_size :: #force_inline proc "contextless" (params: ^Params) -> int { + return 2*SEEDBYTES + TRBYTES + (params.l + params.k) * polyeta_packedbytes(params) + params.k * POLYT0_PACKEDBYTES +} + +@(private,require_results) +signature_size :: #force_inline proc "contextless" (params: ^Params) -> int { + return params.ctild_bytes + params.l * polyz_packedbytes(params) + polyvech_packedbytes(params) +} diff --git a/core/crypto/_mldsa/poly.odin b/core/crypto/_mldsa/poly.odin new file mode 100644 index 000000000..985a807d2 --- /dev/null +++ b/core/crypto/_mldsa/poly.odin @@ -0,0 +1,564 @@ +#+private +package _mldsa + +import "base:intrinsics" +import "core:crypto" +import "core:crypto/shake" + +Poly :: struct { + coeffs: [N]i32, +} + +poly_reduce :: proc "contextless" (a: ^Poly) { + for v, i in a.coeffs { + a.coeffs[i] = reduce32(v) + } +} + +poly_caddq :: proc "contextless" (a: ^Poly) { + for v, i in a.coeffs { + a.coeffs[i] = caddq(v) + } +} + +poly_add :: proc "contextless" (c, a, b: ^Poly) #no_bounds_check { + for i in 0.. uint #no_bounds_check { + s: uint + + for i in 0.. bool #no_bounds_check { + // It is ok to leak which coefficient violates the bound since + // the probability for each coefficient is independent of secret + // data but we must not leak the sign of the centralized + // representative. + for i in 0..> 31 + t = a.coeffs[i] - (t & 2 * a.coeffs[i]) + + if t >= bound { + return true + } + } + + return false +} + +unchecked_get_u24le :: #force_inline proc "contextless" (b: []byte) -> u32 #no_bounds_check { + r := u32(b[0]) + r |= u32(b[1]) << 8 + r |= u32(b[2]) << 16 + return r +} + +rej_uniform :: proc "contextless" (a: []i32, buf: []byte) -> int #no_bounds_check { + ctr, pos: int + + a_len, b_len := len(a), len(buf) + for ctr < a_len && pos + 3 <= b_len { + t := unchecked_get_u24le(buf[pos:]) + t &= 0x7FFFFF + pos += 3 + + if t < Q { + a[ctr] = i32(t) + ctr += 1 + } + } + + return ctr +} + +poly_uniform :: proc(a: ^Poly, seed: []byte, iv: u16) #no_bounds_check { + // Note/yawning: The dilithium reference code does something + // inexplicably more complicated, but this is identical in + // behavior, and simpler. + #assert(STREAM128_BLOCKBYTES % 3 == 0) + POLY_UNIFORM_NBLOCKS :: ((768 + STREAM128_BLOCKBYTES - 1)/STREAM128_BLOCKBYTES) + + buf: [POLY_UNIFORM_NBLOCKS*STREAM128_BLOCKBYTES]byte = --- + defer crypto.zero_explicit(&buf, size_of(buf)) + + ctx: shake.Context = --- + defer shake.reset(&ctx) + stream128_init(&ctx, seed, iv) + + shake.read(&ctx, buf[:]) + ctr := rej_uniform(a.coeffs[:], buf[:]) + + b := buf[:STREAM128_BLOCKBYTES] + for ctr < N { + shake.read(&ctx, b) + ctr += rej_uniform(a.coeffs[ctr:], b) + } +} + +rej_eta :: proc "contextless" (a: []i32, buf: []byte, params: ^Params) -> int { + ctr, pos: int + a_len, b_len := len(a), len(buf) + switch params.eta { + case 2: + for ctr < a_len && pos < b_len { + t0 := u32(buf[pos] & 0x0F) + t1 := u32(buf[pos] >> 4) + pos += 1 + + if t0 < 15 { + t0 = t0 - (205 * t0 >> 10) * 5 + a[ctr] = i32(2 - t0) + ctr += 1 + } + if t1 < 15 && ctr < a_len { + t1 = t1 - (205 * t1 >> 10) * 5 + a[ctr] = i32(2 - t1) + ctr += 1 + } + } + case 4: + for ctr < a_len && pos < b_len { + t0 := u32(buf[pos] & 0x0F) + t1 := u32(buf[pos] >> 4) + pos += 1 + + if t0 < 9 { + a[ctr] = i32(4 - t0) + ctr += 1 + } + if t1 < 9 && ctr < a_len { + a[ctr] = i32(4 - t1) + ctr += 1 + } + } + case: + unreachable() + } + + return ctr +} + +poly_uniform_eta :: proc(a: ^Poly, seed: []byte, iv: u16, params: ^Params) { + POLY_UNIFORM_ETA2_NBLOCKS :: ((136 + STREAM256_BLOCKBYTES - 1)/STREAM256_BLOCKBYTES) + POLY_UNIFORM_ETA4_NBLOCKS :: ((227 + STREAM256_BLOCKBYTES - 1)/STREAM256_BLOCKBYTES) + + buf_: [POLY_UNIFORM_ETA4_NBLOCKS*STREAM256_BLOCKBYTES]byte = --- + buf: []byte + switch params.eta { + case 2: + buf = buf_[:POLY_UNIFORM_ETA2_NBLOCKS*STREAM256_BLOCKBYTES] + case 4: + buf = buf_[:POLY_UNIFORM_ETA4_NBLOCKS*STREAM256_BLOCKBYTES] + case: + unreachable() + } + defer crypto.zero_explicit(&buf_, size_of(buf_)) + + ctx: shake.Context = --- + defer shake.reset(&ctx) + stream256_init(&ctx, seed, iv) + + shake.read(&ctx, buf) + ctr := rej_eta(a.coeffs[:], buf, params) + + b := buf[:STREAM256_BLOCKBYTES] + for ctr < N { + shake.read(&ctx, b) + ctr += rej_eta(a.coeffs[ctr:], b, params) + } +} + +poly_uniform_gamma1 :: proc(a: ^Poly, seed: []byte, iv: u16, params: ^Params) { + POLY_UNIFORM_GAMMA1_NBLOCKS_MAX :: ((POLYZ_PACKEDBYTES_MAX + STREAM256_BLOCKBYTES - 1)/STREAM256_BLOCKBYTES) + + n_blocks := (polyz_packedbytes(params) + STREAM256_BLOCKBYTES - 1)/STREAM256_BLOCKBYTES + + buf_: [POLY_UNIFORM_GAMMA1_NBLOCKS_MAX*STREAM256_BLOCKBYTES]byte = --- + buf := buf_[:n_blocks*STREAM256_BLOCKBYTES] + defer crypto.zero_explicit(&buf_, size_of(buf_)) + + ctx: shake.Context = --- + defer shake.reset(&ctx) + stream256_init(&ctx, seed, iv) + + shake.read(&ctx, buf) + polyz_unpack(a, buf, params) +} + +poly_challenge :: proc(c: ^Poly, seed: []byte, params: ^Params) #no_bounds_check { + buf: [STREAM256_BLOCKBYTES]byte = --- + defer crypto.zero_explicit(&buf, size_of(buf)) + + ctx: shake.Context = --- + defer shake.reset(&ctx) + + shake.init_256(&ctx) + shake.write(&ctx, seed) + shake.read(&ctx, buf[:]) + + signs: u64 + for i in uint(0)..<8 { + signs |= u64(buf[i]) << (8*i) + } + pos := 8 + + b: int + intrinsics.mem_zero(c, size_of(Poly)) + for i := N - params.tau; i < N; i+= 1 { + for { + if pos >= STREAM256_BLOCKBYTES { + shake.read(&ctx, buf[:]) + pos = 0 + } + + b = int(buf[pos]) + pos += 1 + if b <= i { + break + } + } + + c.coeffs[i] = c.coeffs[b] + c.coeffs[b] = i32(1 - 2 * (signs & 1)) + signs >>= 1 + } +} + +polyeta_pack :: proc "contextless" (r: []byte, a: ^Poly, params: ^Params) #no_bounds_check { + t: [8]byte = --- + defer crypto.zero_explicit(&t, size_of(t)) + + eta := params.eta + switch eta { + case 2: + for i in 0..> 0) | (t[1] << 3) | (t[2] << 6) + r[3*i+1] = (t[2] >> 2) | (t[3] << 1) | (t[4] << 4) | (t[5] << 7) + r[3*i+2] = (t[5] >> 1) | (t[6] << 2) | (t[7] << 5) + } + case 4: + for i in 0..> 0) & 7) + r.coeffs[8*i+1] = i32((a[3*i+0] >> 3) & 7) + r.coeffs[8*i+2] = i32(((a[3*i+0] >> 6) | (a[3*i+1] << 2)) & 7) + r.coeffs[8*i+3] = i32((a[3*i+1] >> 1) & 7) + r.coeffs[8*i+4] = i32((a[3*i+1] >> 4) & 7) + r.coeffs[8*i+5] = i32(((a[3*i+1] >> 7) | (a[3*i+2] << 1)) & 7) + r.coeffs[8*i+6] = i32((a[3*i+2] >> 2) & 7) + r.coeffs[8*i+7] = i32((a[3*i+2] >> 5) & 7) + + r.coeffs[8*i+0] = eta - r.coeffs[8*i+0] + r.coeffs[8*i+1] = eta - r.coeffs[8*i+1] + r.coeffs[8*i+2] = eta - r.coeffs[8*i+2] + r.coeffs[8*i+3] = eta - r.coeffs[8*i+3] + r.coeffs[8*i+4] = eta - r.coeffs[8*i+4] + r.coeffs[8*i+5] = eta - r.coeffs[8*i+5] + r.coeffs[8*i+6] = eta - r.coeffs[8*i+6] + r.coeffs[8*i+7] = eta - r.coeffs[8*i+7] + } + case 4: + for i in 0..> 4) + r.coeffs[2*i+0] = eta - r.coeffs[2*i+0] + r.coeffs[2*i+1] = eta - r.coeffs[2*i+1] + } + case: + unreachable() + } +} + +polyt1_pack :: proc "contextless" (r: []byte, a: ^Poly) #no_bounds_check { + for i in 0..> 0) + r[5*i+1] = byte((a.coeffs[4*i+0] >> 8) | (a.coeffs[4*i+1] << 2)) + r[5*i+2] = byte((a.coeffs[4*i+1] >> 6) | (a.coeffs[4*i+2] << 4)) + r[5*i+3] = byte((a.coeffs[4*i+2] >> 4) | (a.coeffs[4*i+3] << 6)) + r[5*i+4] = byte(a.coeffs[4*i+3] >> 2) + } +} + +polyt1_unpack :: proc "contextless" (r: ^Poly, a: []byte) #no_bounds_check { + for i in 0..> 0) | (u32(a[5*i+1]) << 8)) & 0x3FF) + r.coeffs[4*i+1] = i32((u32(a[5*i+1] >> 2) | (u32(a[5*i+2]) << 6)) & 0x3FF) + r.coeffs[4*i+2] = i32((u32(a[5*i+2] >> 4) | (u32(a[5*i+3]) << 4)) & 0x3FF) + r.coeffs[4*i+3] = i32((u32(a[5*i+3] >> 6) | (u32(a[5*i+4]) << 2)) & 0x3FF) + } +} + +polyt0_pack :: proc "contextless" (r: []byte, a: ^Poly) #no_bounds_check { + t: [8]byte = --- + defer crypto.zero_explicit(&t, size_of(t)) + + for i in 0..> 8 + r[13*i+ 1] |= t[1] << 5 + r[13*i+ 2] = t[1] >> 3 + r[13*i+ 3] = t[1] >> 11 + r[13*i+ 3] |= t[2] << 2 + r[13*i+ 4] = t[2] >> 6 + r[13*i+ 4] |= t[3] << 7 + r[13*i+ 5] = t[3] >> 1 + r[13*i+ 6] = t[3] >> 9 + r[13*i+ 6] |= t[4] << 4 + r[13*i+ 7] = t[4] >> 4 + r[13*i+ 8] = t[4] >> 12 + r[13*i+ 8] |= t[5] << 1 + r[13*i+ 9] = t[5] >> 7 + r[13*i+ 9] |= t[6] << 6 + r[13*i+10] = t[6] >> 2 + r[13*i+11] = t[6] >> 10 + r[13*i+11] |= t[7] << 3 + r[13*i+12] = t[7] >> 5 + } +} + +polyt0_unpack :: proc "contextless" (r: ^Poly, a: []byte) #no_bounds_check { + for i in 0..> 5) + r.coeffs[8*i+1] |= i32(u32(a[13*i+2]) << 3) + r.coeffs[8*i+1] |= i32(u32(a[13*i+3]) << 11) + r.coeffs[8*i+1] &= 0x1FFF + + r.coeffs[8*i+2] = i32(a[13*i+3] >> 2) + r.coeffs[8*i+2] |= i32(u32(a[13*i+4]) << 6) + r.coeffs[8*i+2] &= 0x1FFF + + r.coeffs[8*i+3] = i32(a[13*i+4] >> 7) + r.coeffs[8*i+3] |= i32(u32(a[13*i+5]) << 1) + r.coeffs[8*i+3] |= i32(u32(a[13*i+6]) << 9) + r.coeffs[8*i+3] &= 0x1FFF + + r.coeffs[8*i+4] = i32(a[13*i+6] >> 4) + r.coeffs[8*i+4] |= i32(u32(a[13*i+7]) << 4) + r.coeffs[8*i+4] |= i32(u32(a[13*i+8]) << 12) + r.coeffs[8*i+4] &= 0x1FFF + + r.coeffs[8*i+5] = i32(a[13*i+8] >> 1) + r.coeffs[8*i+5] |= i32(u32(a[13*i+9]) << 7) + r.coeffs[8*i+5] &= 0x1FFF + + r.coeffs[8*i+6] = i32(a[13*i+9] >> 6) + r.coeffs[8*i+6] |= i32(u32(a[13*i+10]) << 2) + r.coeffs[8*i+6] |= i32(u32(a[13*i+11]) << 10) + r.coeffs[8*i+6] &= 0x1FFF + + r.coeffs[8*i+7] = i32(a[13*i+11] >> 3) + r.coeffs[8*i+7] |= i32(u32(a[13*i+12]) << 5) + r.coeffs[8*i+7] &= 0x1FFF + + r.coeffs[8*i+0] = (1 << (D-1)) - r.coeffs[8*i+0] + r.coeffs[8*i+1] = (1 << (D-1)) - r.coeffs[8*i+1] + r.coeffs[8*i+2] = (1 << (D-1)) - r.coeffs[8*i+2] + r.coeffs[8*i+3] = (1 << (D-1)) - r.coeffs[8*i+3] + r.coeffs[8*i+4] = (1 << (D-1)) - r.coeffs[8*i+4] + r.coeffs[8*i+5] = (1 << (D-1)) - r.coeffs[8*i+5] + r.coeffs[8*i+6] = (1 << (D-1)) - r.coeffs[8*i+6] + r.coeffs[8*i+7] = (1 << (D-1)) - r.coeffs[8*i+7] + } +} + +polyz_pack :: proc "contextless" (r: []byte, a: ^Poly, params: ^Params) #no_bounds_check { + t: [4]u32 = --- + defer crypto.zero_explicit(&t, size_of(t)) + + gamma1 := params.gamma1 + switch gamma1 { + case 1 << 17: + for i in 0..> 8) + r[9*i+2] = byte(t[0] >> 16) + r[9*i+2] |= byte(t[1] << 2) + r[9*i+3] = byte(t[1] >> 6) + r[9*i+4] = byte(t[1] >> 14) + r[9*i+4] |= byte(t[2] << 4) + r[9*i+5] = byte(t[2] >> 4) + r[9*i+6] = byte(t[2] >> 12) + r[9*i+6] |= byte(t[3] << 6) + r[9*i+7] = byte(t[3] >> 2) + r[9*i+8] = byte(t[3] >> 10) + } + case 1 << 19: + for i in 0..> 8) + r[5*i+2] = byte(t[0] >> 16) + r[5*i+2] |= byte(t[1] << 4) + r[5*i+3] = byte(t[1] >> 4) + r[5*i+4] = byte(t[1] >> 12) + } + case: + unreachable() + } +} + +polyz_unpack :: proc "contextless" (r: ^Poly, a: []byte, params: ^Params) #no_bounds_check { + gamma1 := params.gamma1 + switch gamma1 { + case 1 << 17: + for i in 0..> 2) + r.coeffs[4*i+1] |= i32(u32(a[9*i+3]) << 6) + r.coeffs[4*i+1] |= i32(u32(a[9*i+4]) << 14) + r.coeffs[4*i+1] &= 0x3FFFF + + r.coeffs[4*i+2] = i32(a[9*i+4] >> 4) + r.coeffs[4*i+2] |= i32(u32(a[9*i+5]) << 4) + r.coeffs[4*i+2] |= i32(u32(a[9*i+6]) << 12) + r.coeffs[4*i+2] &= 0x3FFFF + + r.coeffs[4*i+3] = i32(a[9*i+6] >> 6) + r.coeffs[4*i+3] |= i32(u32(a[9*i+7]) << 2) + r.coeffs[4*i+3] |= i32(u32(a[9*i+8]) << 10) + r.coeffs[4*i+3] &= 0x3FFFF + + r.coeffs[4*i+0] = gamma1 - r.coeffs[4*i+0] + r.coeffs[4*i+1] = gamma1 - r.coeffs[4*i+1] + r.coeffs[4*i+2] = gamma1 - r.coeffs[4*i+2] + r.coeffs[4*i+3] = gamma1 - r.coeffs[4*i+3] + } + case 1 << 19: + for i in 0..> 4) + r.coeffs[2*i+1] |= i32(u32(a[5*i+3]) << 4) + r.coeffs[2*i+1] |= i32(u32(a[5*i+4]) << 12) + /* r.coeffs[2*i+1] &= 0xFFFFF */ /* No effect, since we're anyway at 20 bits */ + + r.coeffs[2*i+0] = gamma1 - r.coeffs[2*i+0] + r.coeffs[2*i+1] = gamma1 - r.coeffs[2*i+1] + } + case: + unreachable() + } +} + +polyw1_pack :: proc "contextless" (r: []byte, a: ^Poly, params: ^Params) #no_bounds_check { + switch params.gamma2 { + case (Q-1)/88: + for i in 0..> 2) + r[3*i+1] |= byte(a.coeffs[4*i+2] << 4) + r[3*i+2] = byte(a.coeffs[4*i+2] >> 4) + r[3*i+2] |= byte(a.coeffs[4*i+3] << 2) + } + case (Q-1)/32: + for i in 0.. bool #no_bounds_check { + for i in 0.. bool #no_bounds_check { + for i in 0.. uint #no_bounds_check { + s: uint + + for i in 0.. i32 { + QINV :: 58728449 // q^(-1) mod 2^32 + + t := i32(i64(i32(a)) * QINV) + t = i32((a - i64(t) * Q) >> 32) + return t +} + +@(require_results) +reduce32 :: #force_inline proc "contextless" (a: i32) -> i32 { + t := (a + (1 << 22)) >> 23 + t = a - t * Q + return t +} + +@(require_results) +caddq :: #force_inline proc "contextless" (a: i32) -> i32 { + a := a + a += (a >> 31) & Q + return a +} + +// @(require_results) +// freeze :: #force_inline proc "contextless" (a: i32) -> i32 { +// a := a +// a = reduce32(a) +// a = caddq(a) +// return a +// } diff --git a/core/crypto/_mldsa/rounding.odin b/core/crypto/_mldsa/rounding.odin new file mode 100644 index 000000000..a879d6159 --- /dev/null +++ b/core/crypto/_mldsa/rounding.odin @@ -0,0 +1,56 @@ +#+private +package _mldsa + +power2round :: proc "contextless" (a: i32) -> (i32, i32) { + a1 := (a + (1 << (D-1)) - 1) >> D + a0 := a - (a1 << D) + return a0, a1 +} + +decompose :: proc "contextless" (a: i32, gamma2: i32) -> (i32, i32) { + a1 := (a + 127) >> 7 + switch gamma2 { + case (Q - 1)/32: + a1 = (a1 * 1025 + (1 << 21)) >> 22 + a1 &= 15 + case (Q - 1)/88: + a1 = (a1 * 11275 + (1 << 23)) >> 24 + a1 ~= ((43 - a1) >> 31) & a1 + } + + a0 := a - a1 * 2 * gamma2 + a0 -= (((Q - 1)/2 - a0) >> 31) & Q + return a0, a1 +} + +make_hint :: proc "contextless" (a0, a1: i32, gamma2: i32) -> uint { + if (a0 > gamma2 || a0 < -gamma2 || (a0 == -gamma2 && a1 != 0)) { + return 1 + } + return 0 +} + +use_hint :: proc "contextless" (a: i32, hint: uint, gamma2: i32) -> i32 { + a0, a1 := decompose(a, gamma2) + if hint == 0 { + return a1 + } + + switch gamma2 { + case (Q - 1)/32: + if (a0 > 0) { + return (a1 + 1) & 15 + } else { + return (a1 -1) & 15 + } + case (Q - 1)/88: + if (a0 > 0) { + return (a1 == 43) ? 0 : a1 + 1 + } else { + return (a1 == 0) ? 43 : a1 - 1 + } + } + + unreachable() +} + diff --git a/core/crypto/_mldsa/symmetric_shake.odin b/core/crypto/_mldsa/symmetric_shake.odin new file mode 100644 index 000000000..5c0aa99d6 --- /dev/null +++ b/core/crypto/_mldsa/symmetric_shake.odin @@ -0,0 +1,39 @@ +#+private +package _mldsa + +import "core:crypto/_sha3" +import "core:crypto/shake" + +STREAM128_BLOCKBYTES :: _sha3.RATE_128 +STREAM256_BLOCKBYTES :: _sha3.RATE_256 + +stream128_init :: proc(ctx: ^shake.Context, seed: []byte, iv: u16) { + t: [2]byte = --- + t[0] = byte(iv) + t[1] = byte(iv >> 8) + + shake.init_128(ctx) + shake.write(ctx, seed) + shake.write(ctx, t[:]) +} + +stream256_init :: proc(ctx: ^shake.Context, seed: []byte, iv: u16) { + t: [2]byte = --- + t[0] = byte(iv) + t[1] = byte(iv >> 8) + + shake.init_256(ctx) + shake.write(ctx, seed) + shake.write(ctx, t[:]) +} + +shake256 :: proc(dst: []byte, srcs: ..[]byte) { + ctx: shake.Context = --- + defer shake.reset(&ctx) + + shake.init_256(&ctx) + for src in srcs { + shake.write(&ctx, src) + } + shake.read(&ctx, dst) +} diff --git a/core/crypto/_mlkem/cbd.odin b/core/crypto/_mlkem/cbd.odin new file mode 100644 index 000000000..301551ff5 --- /dev/null +++ b/core/crypto/_mlkem/cbd.odin @@ -0,0 +1,52 @@ +#+private +package _mlkem + +import "core:encoding/endian" + +unchecked_get_u24le :: #force_inline proc "contextless" (b: []byte) -> u32 #no_bounds_check { + r := u32(b[0]) + r |= u32(b[1]) << 8 + r |= u32(b[2]) << 16 + return r +} + +cbd3 :: proc "contextless" (r: ^Poly, buf: ^[3*N/4]byte) #no_bounds_check { + for i in 0..>1) & 0x00249249 + d += (t>>2) & 0x00249249 + + for j in uint(0)..<4 { + a := i16((d >> (6*j+0)) & 0x7) + b := i16((d >> (6*j+3)) & 0x7) + r.coeffs[4*i+int(j)] = a - b + } + } +} + +cbd2 :: proc "contextless" (r: ^Poly, buf: ^[2*N/4]byte) #no_bounds_check { + for i in 0..>1) & 0x55555555 + + for j in uint(0)..<8 { + a := i16((d >> (4*j+0)) & 0x3) + b := i16((d >> (4*j+2)) & 0x3) + r.coeffs[8*i+int(j)] = a - b + } + } +} + +poly_cbd_eta1_512 :: proc "contextless" (r: ^Poly, buf: ^[ETA1_512*N/4]byte) { + cbd3(r, buf) +} + +poly_cbd_eta1 :: proc "contextless" (r: ^Poly, buf: ^[ETA1*N/4]byte) { + cbd2(r, buf) +} + +poly_cbd_eta2 :: proc "contextless" (r: ^Poly, buf: ^[ETA2*N/4]byte) { + cbd2(r, buf) +} diff --git a/core/crypto/_mlkem/constants.odin b/core/crypto/_mlkem/constants.odin new file mode 100644 index 000000000..7195d8b61 --- /dev/null +++ b/core/crypto/_mlkem/constants.odin @@ -0,0 +1,53 @@ +package _mlkem + +K_512 :: 2 +K_768 :: 3 +K_1024 :: 4 +K_MAX :: K_1024 + +N :: 256 +Q :: 3329 + +ETA1_512 :: 3 +ETA1 :: 2 +ETA2 :: 2 + +POLYBYTES :: 384 +SYMBYTES :: 32 + +POLYCOMPRESSEDBYTES_512 :: 128 +POLYCOMPRESSEDBYTES_768 :: 128 +POLYCOMPRESSEDBYTES_1024 :: 160 + +POLYVECBYTES_512 :: K_512 * POLYBYTES +POLYVECBYTES_768 :: K_768 * POLYBYTES +POLYVECBYTES_1024 :: K_1024 * POLYBYTES + +POLYVECCOMPRESSEDBYTES_512 :: K_512 * 320 +POLYVECCOMPRESSEDBYTES_768 :: K_768 * 320 +POLYVECCOMPRESSEDBYTES_1024 :: K_1024 * 352 + +INDCPA_MSGBYTES :: SYMBYTES +INDCPA_PUBLICKEYBYTES_512 :: POLYVECBYTES_512 + SYMBYTES +INDCPA_SECRETKEYBYTES_512 :: POLYVECBYTES_512 +INDCPA_PUBLICKEYBYTES_768 :: POLYVECBYTES_768 + SYMBYTES +INDCPA_SECRETKEYBYTES_768 :: POLYVECBYTES_768 +INDCPA_PUBLICKEYBYTES_1024 :: POLYVECBYTES_1024 + SYMBYTES +INDCPA_SECRETKEYBYTES_1024 :: POLYVECBYTES_1024 +INDCPA_PUBLICKEYBYTES_MAX :: INDCPA_PUBLICKEYBYTES_1024 +INDCPA_BYTES_512 :: POLYVECCOMPRESSEDBYTES_512 + POLYCOMPRESSEDBYTES_512 +INDCPA_BYTES_768 :: POLYVECCOMPRESSEDBYTES_768 + POLYCOMPRESSEDBYTES_768 +INDCPA_BYTES_1024 :: POLYVECCOMPRESSEDBYTES_1024 + POLYCOMPRESSEDBYTES_1024 + +ENCAPSKEYBYTES_512 :: INDCPA_PUBLICKEYBYTES_512 +ENCAPSKEYBYTES_768 :: INDCPA_PUBLICKEYBYTES_768 +ENCAPSKEYBYTES_1024 :: INDCPA_PUBLICKEYBYTES_1024 + +DECAPSKEYBYTES_512 :: INDCPA_SECRETKEYBYTES_512 + INDCPA_PUBLICKEYBYTES_512 + 2 * SYMBYTES +DECAPSKEYBYTES_768 :: INDCPA_SECRETKEYBYTES_768 + INDCPA_PUBLICKEYBYTES_768 + 2 * SYMBYTES +DECAPSKEYBYTES_1024 :: INDCPA_SECRETKEYBYTES_1024 + INDCPA_PUBLICKEYBYTES_1024 + 2 * SYMBYTES + +CIPHERTEXTBYTES_512 :: INDCPA_BYTES_512 +CIPHERTEXTBYTES_768 :: INDCPA_BYTES_768 +CIPHERTEXTBYTES_1024 :: INDCPA_BYTES_1024 +CIPHERTEXTBYTES_MAX :: INDCPA_BYTES_1024 diff --git a/core/crypto/_mlkem/k_pke.odin b/core/crypto/_mlkem/k_pke.odin new file mode 100644 index 000000000..d8b87ad7a --- /dev/null +++ b/core/crypto/_mlkem/k_pke.odin @@ -0,0 +1,349 @@ +#+private +package _mlkem + +import "core:crypto" +import "core:crypto/shake" + +@(require_results) +pack_pk :: proc "contextless" (r: []byte, pk: ^Polyvec, seed: []byte, k: int) -> bool { + pk_len := polyvec_byte_size(k) + switch { + case pk_len == 0: + return false + case len(seed) != SYMBYTES || len(r) != pk_len + SYMBYTES: + return false + } + + polyvec_tobytes(r[:pk_len], pk, k) + copy(r[pk_len:], seed) + + return true +} + +@(require_results) +unpack_pk :: proc "contextless" (pk: ^Polyvec, seed, packedpk: []byte) -> bool { + pk_len := len(packedpk) - SYMBYTES + k: int + switch { + case pk_len == POLYVECBYTES_512: + k = K_512 + case pk_len == POLYVECBYTES_768: + k = K_768 + case pk_len == POLYVECBYTES_1024: + k = K_1024 + case len(packedpk) - pk_len != SYMBYTES: + return false + case len(seed) != SYMBYTES: + return false + } + if k == 0 { + return false + } + + ok := polyvec_frombytes(pk, packedpk[:pk_len], k) + copy(seed, packedpk[pk_len:]) + + return ok +} + +@(require_results) +pack_sk :: proc "contextless" (r: []byte, sk: ^Polyvec, k: int) -> bool { + r_len := len(r) + if r_len == 0 || r_len != polyvec_byte_size(k) { + return false + } + + polyvec_tobytes(r, sk, k) + + return true +} + +@(require_results) +unpack_sk :: proc "contextless" (sk: ^Polyvec, packedsk: []byte) -> bool { + k: int + switch len(packedsk) { + case POLYVECBYTES_512: + k = K_512 + case POLYVECBYTES_768: + k = K_768 + case POLYVECBYTES_1024: + k = K_1024 + case: + return false + } + if k == 0 { + return false + } + + return polyvec_frombytes(sk, packedsk, k) +} + +@(require_results) +pack_ciphertext :: proc "contextless" (r: []byte, b: ^Polyvec, v: ^Poly, k: int) -> bool { + b_len := polyvec_compressed_byte_size(k) + if len(r) != b_len + poly_compressed_bytes(k) { + return false + } + + polyvec_compress(r[:b_len], b, k) + poly_compress(r[b_len:], v) + + return true +} + +@(require_results) +unpack_ciphertext :: proc "contextless" (b: ^Polyvec, v: ^Poly, c: []byte) -> int { + b_len: int + k: int + switch len(c) { + case INDCPA_BYTES_512: + b_len = POLYVECCOMPRESSEDBYTES_512 + k = K_512 + case INDCPA_BYTES_768: + b_len = POLYVECCOMPRESSEDBYTES_768 + k = K_768 + case INDCPA_BYTES_1024: + b_len = POLYVECCOMPRESSEDBYTES_1024 + k = K_1024 + case: + return 0 + } + + polyvec_decompress(b, c[:b_len], k) + poly_decompress(v, c[b_len:]) + + return k +} + +@(require_results) +rej_uniform :: proc "contextless" (r: []i16, buf: []byte) -> int { + r_len, b_len := len(r), len(buf) + + ctr, pos: int + for ctr < r_len && pos + 3 <= b_len { + val0 := (u16(buf[pos+0] >> 0) | (u16(buf[pos+1]) << 8)) & 0xFFF + val1 := (u16(buf[pos+1] >> 4) | (u16(buf[pos+2]) << 4)) & 0xFFF + pos += 3 + + if val0 < Q { + r[ctr] = i16(val0) + ctr += 1 + } + if(ctr < r_len && val1 < Q) { + r[ctr] = i16(val1) + ctr += 1 + } + } + + return ctr +} + +gen_matrix :: proc(a: []Polyvec, seed: []byte, transposed: bool, k: int) { + GEN_MATRIX_NBLOCKS :: ((12*N/8*(1 << 12)/Q + XOF_BLOCKBYTES)/XOF_BLOCKBYTES) + + buf: [GEN_MATRIX_NBLOCKS*XOF_BLOCKBYTES]byte = --- + ctx: shake.Context = --- + ctr: int + + defer shake.reset(&ctx) + defer crypto.zero_explicit(&buf, size_of(buf)) + + for i in 0.. bool { + ensure(len(m) == INDCPA_MSGBYTES, "crypto/mlkem: invalid K-PKE m") + ensure(len(r) == SYMBYTES, "crypto/mlkem: invalid K-PKE r") + + k := ek.k + + at_: [K_MAX]Polyvec = --- + sp, ep, b: Polyvec = ---, ---, --- + kay, epp, v: Poly = ---, ---, --- + defer crypto.zero_explicit(&at_, size_of(Polyvec) * k) + defer polyvec_clear(&sp, &ep, &b) + defer poly_clear(&kay, &epp, &v) + + poly_frommsg(&kay, m) + + at := at_[:k] + + gen_matrix(at, ek.p[:], true, k) + + n := byte(0) + for i in 0.. bool { + if len(plaintext) != INDCPA_MSGBYTES { + return false + } + + k := dk.k + + b: Polyvec = --- + v, mp: Poly = ---, --- + defer poly_clear(&v, &mp) + + if unpack_ciphertext(&b, &v, c) != k { + return false + } + + polyvec_ntt(&b, k) + polyvec_basemul_acc_montgomery(&mp, &dk.pv, &b, k) + poly_invntt_tomont(&mp) + + poly_sub(&mp, &v, &mp) + poly_reduce(&mp) + + poly_tomsg(plaintext, &mp) + + return true +} diff --git a/core/crypto/_mlkem/kem_internal.odin b/core/crypto/_mlkem/kem_internal.odin new file mode 100644 index 000000000..b50aa4b58 --- /dev/null +++ b/core/crypto/_mlkem/kem_internal.odin @@ -0,0 +1,195 @@ +package _mlkem + +import "core:crypto" +import subtle "core:crypto/_subtle" + +// This implementation is derived from the PQ-CRYSTALS reference +// implementation [[ https://github.com/pq-crystals/kyber ]], +// primarily for licensing reasons. Arguably mlkem-native is +// a more "up to date" codebase, but the changes to the +// ref code is minor and they slapped an attribution-required +// license on something that was originally CC-0/Apache 2.0. + +// "Private Key" +Decapsulation_Key :: struct { + pke_dk: K_PKE_Decryption_Key, + ek: Encapsulation_Key, + seed: [SYMBYTES*2]byte, // (d, z) +} + +// "Public Key" +Encapsulation_Key :: struct { + pke_ek: K_PKE_Encryption_Key, + raw_bytes: [INDCPA_PUBLICKEYBYTES_MAX]byte, + h: [SYMBYTES]byte, +} + +decapsulation_key_expanded_bytes :: proc( + dk: ^Decapsulation_Key, + dst: []byte, +) { + sk := &dk.pke_dk + pv_len := polyvec_byte_size(sk.k) + ek_len := pv_len + SYMBYTES + + ek_bytes := dk.ek.raw_bytes[:ek_len] + + dst := dst + _ = pack_sk(dst[:pv_len], &sk.pv, sk.k) + + dst = dst[pv_len:] + copy(dst, ek_bytes) + dst = dst[ek_len:] + hash_h(dst[:SYMBYTES], ek_bytes) + dst = dst[SYMBYTES:] + copy(dst, dk.seed[SYMBYTES:]) +} + +@(require_results) +encapsulation_key_set_bytes :: proc( + ek: ^Encapsulation_Key, + k: int, + b: []byte, +) -> bool { + k_len: int + switch k { + case K_512: + k_len = ENCAPSKEYBYTES_512 + case K_768: + k_len = ENCAPSKEYBYTES_768 + case K_1024: + k_len = ENCAPSKEYBYTES_1024 + case: + return false + } + if len(b) != k_len { + return false + } + + pke_ek := &ek.pke_ek + ok := unpack_pk(&pke_ek.pv, pke_ek.p[:], b) + pke_ek.k = k + copy(ek.raw_bytes[:k_len], b) + hash_h(ek.h[:], b) + + // FIPS 203 unlike Kyber requires canonical encoding of + // encapsulation keys (Section 7,2), which is checked in + // unpack_pk. + + if !ok { + crypto.zero_explicit(ek, size_of(Encapsulation_Key)) + } + + return ok +} + +encapsulation_key_set_decaps :: proc(ek: ^Encapsulation_Key, dk: ^Decapsulation_Key) { + dk_ek := &dk.ek.pke_ek + ensure(dk_ek.k == K_512 || dk_ek.k == K_768 || dk_ek.k == K_1024, "crypto/mlkem: invalid decaps k") + + k_pke_encryption_key_set(&ek.pke_ek, dk_ek) + copy(ek.raw_bytes[:], dk.ek.raw_bytes[:]) + copy(ek.h[:], dk.ek.h[:]) +} + +// NIST's version of this also returns an encapsulation key, but our +// internal representation includes it as part of the decapsulation key +// in a more traditional "keypair" approach. +kem_keygen_internal :: proc( + dk: ^Decapsulation_Key, + seed: []byte, // (d, z) + k: int, +) { + ensure(len(seed) == 2 * SYMBYTES, "crypto/mlkem: invalid seed") + + dk_ek := &dk.ek + d, z := seed[:SYMBYTES], seed[SYMBYTES:] + + k_pke_keygen(&dk_ek.pke_ek, &dk.pke_dk, d, k) + + ek_len := polyvec_byte_size(k) + SYMBYTES + ek_bytes := dk_ek.raw_bytes[:ek_len] + ensure( + pack_pk(ek_bytes, &dk_ek.pke_ek.pv, dk_ek.pke_ek.p[:], k), + "crypto/mlkem: failed to pack K-PKE ek", + ) + hash_h(dk_ek.h[:], ek_bytes) + copy(dk.seed[:SYMBYTES], d) + copy(dk.seed[SYMBYTES:], z) +} + +// The `_internal` "de-randomized" versions of ML-KEM.Encaps and +// ML-KEM.Decaps are only ever to be called by the actual non-interal +// implementation or test cases. + +kem_encaps_internal :: proc( + shared_secret: []byte, + ciphertext: []byte, + ek: ^Encapsulation_Key, + randomness: []byte, +) { + ensure(len(shared_secret) == SYMBYTES, "crypto/mlkem: invalid K") + ensure(len(randomness) == SYMBYTES, "crypto/mlkem: invalid m") + ensure( + len(ciphertext) == ct_len_for_k(ek.pke_ek.k), + "crypto/mlkem: invalid ciphertext length", + ) + + buf: [2*SYMBYTES]byte = --- + defer crypto.zero_explicit(&buf, size_of(buf)) + + hash_g(buf[:], randomness, ek.h[:]) + + // Can't fail, ciphertext length is valid. + _ = k_pke_encrypt(ciphertext, &ek.pke_ek, randomness, buf[SYMBYTES:]) + + copy(shared_secret, buf[:SYMBYTES]) +} + +kem_decaps_internal :: proc( + shared_secret: []byte, + dk: ^Decapsulation_Key, + ciphertext: []byte, +) { + ct_len := ct_len_for_k(dk.pke_dk.k) + ensure( + len(ciphertext) == ct_len, + "crypto/mlkem: invalid ciphertext length", + ) + + m_: [SYMBYTES]byte + defer crypto.zero_explicit(&m_, size_of(m_)) + + // Can't fail, ciphertext length is valid. + _ = k_pke_decrypt(m_[:], &dk.pke_dk, ciphertext) + + buf: [2*SYMBYTES]byte = --- + defer crypto.zero_explicit(&buf, size_of(buf)) + + ek := &dk.ek + hash_g(buf[:], m_[:], ek.h[:]) + + rkprf(shared_secret, dk.seed[SYMBYTES:], ciphertext) + + ct_buf: [CIPHERTEXTBYTES_MAX]byte = --- + defer crypto.zero_explicit(&ct_buf, size_of(ct_buf)) + ct_ := ct_buf[:ct_len] + _ = k_pke_encrypt(ct_, &ek.pke_ek, m_[:], buf[SYMBYTES:]) + + ok := crypto.compare_constant_time(ciphertext, ct_) + subtle.cmov_bytes(shared_secret, buf[:SYMBYTES], ok) +} + +@(private="file") +ct_len_for_k :: proc(k: int) -> int { + switch k { + case K_512: + return CIPHERTEXTBYTES_512 + case K_768: + return CIPHERTEXTBYTES_768 + case K_1024: + return CIPHERTEXTBYTES_1024 + case: + panic("crypto/mlkem: invalid k for ciphertext length") + } +} diff --git a/core/crypto/_mlkem/ntt.odin b/core/crypto/_mlkem/ntt.odin new file mode 100644 index 000000000..c99c73487 --- /dev/null +++ b/core/crypto/_mlkem/ntt.odin @@ -0,0 +1,75 @@ +#+private +package _mlkem + +@(rodata) +ZETAS := [128]i16 { + -1044, -758, -359, -1517, 1493, 1422, 287, 202, + -171, 622, 1577, 182, 962, -1202, -1474, 1468, + 573, -1325, 264, 383, -829, 1458, -1602, -130, + -681, 1017, 732, 608, -1542, 411, -205, -1571, + 1223, 652, -552, 1015, -1293, 1491, -282, -1544, + 516, -8, -320, -666, -1618, -1162, 126, 1469, + -853, -90, -271, 830, 107, -1421, -247, -951, + -398, 961, -1508, -725, 448, -1065, 677, -1275, + -1103, 430, 555, 843, -1251, 871, 1550, 105, + 422, 587, 177, -235, -291, -460, 1574, 1653, + -246, 778, 1159, -147, -777, 1483, -602, 1119, + -1590, 644, -872, 349, 418, 329, -156, -75, + 817, 1097, 603, 610, 1322, -1285, -1465, 384, + -1215, -136, 1218, -1335, -874, 220, -1187, -1659, + -1185, -1530, -1278, 794, -1510, -854, -870, 478, + -108, -308, 996, 991, 958, -1460, 1522, 1628, +} + +@(require_results) +fqmul :: #force_inline proc "contextless" (a, b: i16) -> i16 { + return montgomery_reduce(i32(a) * i32(b)) +} + +ntt :: proc "contextless" (r: ^[N]i16) #no_bounds_check { + j, k := 0, 1 + for l := 128; l >= 2; l >>= 1 { + for start := 0; start < N; start = j + l { + zeta := ZETAS[k] + k += 1 + for j = start; j < start + l; j += 1 { + t := fqmul(zeta, r[j+l]) + r[j+l] = r[j] - t + r[j] = r[j] + t + } + } + } +} + +invntt :: proc "contextless" (r: ^[N]i16) #no_bounds_check { + F : i16 : 1441 // mont^2/128 + + j, k := 0, 127 + for l := 2; l <= 128; l <<= 1 { + for start := 0; start < 256; start = j+l { + zeta := ZETAS[k] + k -= 1 + for j = start; j < start + l; j += 1 { + t := r[j] + r[j] = barrett_reduce(t + r[j+l]) + r[j+l] = r[j+l] - t + r[j+l] = fqmul(zeta, r[j+l]) + } + } + } + + for v, i in r { + r[i] = fqmul(v, F) + } +} + +@(require_results) +base_case_multiply :: proc "contextless" (a_0, a_1, b_0, b_1, zeta: i16) -> (i16, i16) { + r_0 := fqmul(a_1, b_1) + r_0 = fqmul(r_0, zeta) + r_0 += fqmul(a_0, b_0) + r_1 := fqmul(a_0, b_1) + r_1 += fqmul(a_1, b_0) + + return r_0, r_1 +} diff --git a/core/crypto/_mlkem/poly.odin b/core/crypto/_mlkem/poly.odin new file mode 100644 index 000000000..13dd54930 --- /dev/null +++ b/core/crypto/_mlkem/poly.odin @@ -0,0 +1,241 @@ +#+private +package _mlkem + +import "core:crypto" +import subtle "core:crypto/_subtle" + +// Elements of R_q = Z_q[X]/(X^n + 1). Represents polynomial +// coeffs[0] + X*coeffs[1] + X^2*coeffs[2] + ... + X^{n-1}*coeffs[n-1] +Poly :: struct { + coeffs: [N]i16, +} + +poly_compress :: proc "contextless" (r: []byte, a: ^Poly) #no_bounds_check { + t: [8]byte = --- + defer crypto.zero_explicit(&t, size_of(t)) + + r := r + switch len(r) { + case POLYCOMPRESSEDBYTES_768: // Also covers _512 + for i in 0..> 15) & Q + // t[j] = ((((uint16_t)u << 4) + Q/2)/Q) & 15 + d0 := u32(u) << 4 + d0 += 1665 + d0 *= 80635 + d0 >>= 28 + t[j] = byte(d0) & 0xf + } + + r[0] = t[0] | (t[1] << 4) + r[1] = t[2] | (t[3] << 4) + r[2] = t[4] | (t[5] << 4) + r[3] = t[6] | (t[7] << 4) + r = r[4:] + } + case POLYCOMPRESSEDBYTES_1024: + for i in 0..> 15) & Q + // t[j] = ((((uint16_t)u << 5) + Q/2)/Q) & 31 + d0 := u32(u) << 5 + d0 += 1664 + d0 *= 40318 + d0 >>= 27 + t[j] = byte(d0) & 0x1f + } + + r[0] = (t[0] >> 0) | (t[1] << 5) + r[1] = (t[1] >> 3) | (t[2] << 2) | (t[3] << 7) + r[2] = (t[3] >> 1) | (t[4] << 4) + r[3] = (t[4] >> 4) | (t[5] << 1) | (t[6] << 6) + r[4] = (t[6] >> 2) | (t[7] << 3) + r = r[5:] + } + case: + unreachable() + } +} + +poly_decompress :: proc "contextless" (r: ^Poly, a: []byte) { + a := a + switch len(a) { + case POLYCOMPRESSEDBYTES_768: // Also covers _512 + for i in 0..> 4) + r.coeffs[2*i+1] = i16(((u16(a[0] >> 4) * Q) + 8) >> 4) + a = a[1:] + } + case POLYCOMPRESSEDBYTES_1024: + t: [8]byte = --- + defer crypto.zero_explicit(&t, size_of(t)) + + for i in 0..> 0) + t[1] = (a[0] >> 5) | (a[1] << 3) + t[2] = (a[1] >> 2) + t[3] = (a[1] >> 7) | (a[2] << 1) + t[4] = (a[2] >> 4) | (a[3] << 4) + t[5] = (a[3] >> 1) + t[6] = (a[3] >> 6) | (a[4] << 2) + t[7] = (a[4] >> 3) + a = a[5:] + + for j in 0..<8 { + r.coeffs[8*i+j] = i16((u32(t[j] & 31) * Q + 16) >> 5) + } + } + case: + unreachable() + } +} + +poly_tobytes :: proc "contextless" (r: []byte, a: ^Poly) #no_bounds_check { + ensure_contextless(len(r) >= POLYBYTES) + + for i in 0..> 15) & Q) + t1 := u16(a.coeffs[2*i+1]) + t1 += u16((i16(t1) >> 15) & Q) + r[3*i+0] = byte(t0 >> 0) + r[3*i+1] = byte(t0 >> 8) | byte(t1 << 4) + r[3*i+2] = byte(t1 >> 4) + } +} + +@(require_results) +poly_frombytes :: proc "contextless" (r: ^Poly, a: []byte) -> bool #no_bounds_check { + ensure_contextless(len(a) >= POLYBYTES) + + ok := true + for i in 0..> 0) | (u16(a[3*i+1]) << 8)) & 0xFFF) + r.coeffs[2*i+1] = i16(((u16(a[3*i+1]) >> 4) | (u16(a[3*i+2]) << 4)) & 0xFFF) + ok &= r.coeffs[2*i] < Q && r.coeffs[2*i+1] < Q + } + + return ok +} + +poly_frommsg :: proc "contextless" (r: ^Poly, msg: []byte) #no_bounds_check { + #assert(INDCPA_MSGBYTES == N/8) + ensure_contextless(len(msg) == INDCPA_MSGBYTES) + + for i in 0..> uint(j))&1) + } + } +} + +poly_tomsg :: proc "contextless" (msg: []byte, a: ^Poly) #no_bounds_check { + ensure_contextless(len(msg) == INDCPA_MSGBYTES) + + for i in 0..> 15) & Q + // t = (((t << 1) + Q/2)/Q) & 1 + t <<= 1 + t += 1665 + t *= 80635 + t >>= 28 + t &= 1 + msg[i] |= byte(t << j) + } + } +} + +poly_getnoise_eta1_512 :: proc(r: ^Poly, seed: []byte, iv: byte) { + buf: [ETA1_512*N/4]byte = --- + defer crypto.zero_explicit(&buf, size_of(buf)) + + prf(buf[:], seed, iv) + poly_cbd_eta1_512(r, &buf) +} + +poly_getnoise_eta1 :: proc(r: ^Poly, seed: []byte, iv: byte) { + buf: [ETA1*N/4]byte = --- + defer crypto.zero_explicit(&buf, size_of(buf)) + + prf(buf[:], seed, iv) + poly_cbd_eta1(r, &buf) +} + +poly_getnoise_eta2 :: proc(r: ^Poly, seed: []byte, iv: byte) { + buf: [ETA2*N/4]byte = --- + defer crypto.zero_explicit(&buf, size_of(buf)) + + prf(buf[:], seed, iv) + poly_cbd_eta2(r, &buf) +} + +poly_ntt :: proc "contextless" (r: ^Poly) { + ntt(&r.coeffs) + poly_reduce(r) +} + +poly_invntt_tomont :: proc "contextless" (r: ^Poly) { + invntt(&r.coeffs) +} + +poly_basemul_montgomery :: proc "contextless" (r, a, b: ^Poly) #no_bounds_check { + for i in 0.. int { + switch k { + case K_512: + return POLYCOMPRESSEDBYTES_512 + case K_768: + return POLYCOMPRESSEDBYTES_768 + case K_1024: + return POLYCOMPRESSEDBYTES_1024 + case: + unreachable() + } +} diff --git a/core/crypto/_mlkem/polyvec.odin b/core/crypto/_mlkem/polyvec.odin new file mode 100644 index 000000000..2b971d2df --- /dev/null +++ b/core/crypto/_mlkem/polyvec.odin @@ -0,0 +1,224 @@ +#+private +package _mlkem + +import "core:crypto" + +Polyvec :: struct { + vec: [K_MAX]Poly, +} + +polyvec_compress :: proc "contextless" (r: []byte, a: ^Polyvec, kay: int) #no_bounds_check { + d0: u64 + + r := r + switch len(r) { + case POLYVECCOMPRESSEDBYTES_512, POLYVECCOMPRESSEDBYTES_768: + ensure_contextless(kay == K_512 || kay == K_768) + + t: [4]u16 = --- + defer crypto.zero_explicit(&t, size_of(t)) + + for i in 0..> 15) & Q) + // t[k] = ((((uint32_t)t[k] << 10) + Q/2)/Q) & 0x3ff + d0 = u64(t[k]) + d0 <<= 10 + d0 += 1665 + d0 *= 1290167 + d0 >>= 32 + t[k] = u16(d0 & 0x3ff) + } + + r[0] = byte(t[0] >> 0) + r[1] = byte((t[0] >> 8) | (t[1] << 2)) + r[2] = byte((t[1] >> 6) | (t[2] << 4)) + r[3] = byte((t[2] >> 4) | (t[3] << 6)) + r[4] = byte(t[3] >> 2) + r = r[5:] + } + } + case POLYVECCOMPRESSEDBYTES_1024: + ensure_contextless(kay == K_1024) + + t: [8]u16 = --- + defer crypto.zero_explicit(&t, size_of(t)) + + for i in 0..> 15) & Q) + // t[k] = ((((uint32_t)t[k] << 11) + Q/2)/Q) & 0x7ff + d0 = u64(t[k]) + d0 <<= 11 + d0 += 1664 + d0 *= 645084 + d0 >>= 31 + t[k] = u16(d0 & 0x7ff) + } + + r[0] = byte(t[0] >> 0) + r[1] = byte((t[0] >> 8) | (t[1] << 3)) + r[2] = byte((t[1] >> 5) | (t[2] << 6)) + r[3] = byte(t[2] >> 2) + r[4] = byte((t[2] >> 10) | (t[3] << 1)) + r[5] = byte((t[3] >> 7) | (t[4] << 4)) + r[6] = byte((t[4] >> 4) | (t[5] << 7)) + r[7] = byte(t[5] >> 1) + r[8] = byte((t[5] >> 9) | (t[6] << 2)) + r[9] = byte((t[6] >> 6) | (t[7] << 5)) + r[10] = byte(t[7] >> 3) + r = r[11:] + } + } + case: + unreachable() + } +} + +polyvec_decompress :: proc "contextless" (r: ^Polyvec, a: []byte, kay: int) #no_bounds_check { + a := a + switch len(a) { + case POLYVECCOMPRESSEDBYTES_512, POLYVECCOMPRESSEDBYTES_768: + ensure_contextless(kay == K_512 || kay == K_768) + + t: [4]u16 = --- + defer crypto.zero_explicit(&t, size_of(t)) + + for i in 0..> 0) | (u16(a[1]) << 8) + t[1] = u16(a[1] >> 2) | (u16(a[2]) << 6) + t[2] = u16(a[2] >> 4) | (u16(a[3]) << 4) + t[3] = u16(a[3] >> 6) | (u16(a[4]) << 2) + a = a[5:] + + for k in 0..<4 { + r.vec[i].coeffs[4*j+k] = i16((u32(t[k] & 0x3FF) * Q + 512) >> 10) + } + } + } + case POLYVECCOMPRESSEDBYTES_1024: + t: [8]u16 = --- + defer crypto.zero_explicit(&t, size_of(t)) + + for i in 0..> 0) | (u16(a[1]) << 8) + t[1] = u16(a[1] >> 3) | (u16(a[2]) << 5) + t[2] = u16(a[2] >> 6) | (u16(a[3]) << 2) | (u16(a[4]) << 10) + t[3] = u16(a[4] >> 1) | (u16(a[5]) << 7) + t[4] = u16(a[5] >> 4) | (u16(a[6]) << 4) + t[5] = u16(a[6] >> 7) | (u16(a[7]) << 1) | (u16(a[8]) << 9) + t[6] = u16(a[8] >> 2) | (u16(a[9]) << 6) + t[7] = u16(a[9] >> 5) | (u16(a[10]) << 3) + a = a[11:] + + for k in 0..<8 { + r.vec[i].coeffs[8*j+k] = i16((u32(t[k] & 0x7FF) * Q + 1024) >> 11) + } + } + } + case: + unreachable() + } +} + +polyvec_tobytes :: proc "contextless" (r: []byte, a: ^Polyvec, k: int) #no_bounds_check { + ensure_contextless(len(r) == k * POLYBYTES, "crypto/mlkem: invalid buffer") + + r := r + for i in 0.. bool #no_bounds_check { + switch k { + case K_512, K_768, K_1024: + case: + panic_contextless("crypto/mlkem: invalid POLYVECBYTES") + } + ensure_contextless(len(a) == k * POLYBYTES, "crypto/mlkem: invalid buffer") + + a := a + ok := true + for i in 0.. int { + switch k { + case K_512, K_768, K_1024: + return k * POLYBYTES + case: + return 0 + } +} + +@(require_results) +polyvec_compressed_byte_size :: #force_inline proc "contextless" (k: int) -> int { + switch k { + case K_512: + return POLYVECCOMPRESSEDBYTES_512 + case K_768: + return POLYVECCOMPRESSEDBYTES_768 + case K_1024: + return POLYVECCOMPRESSEDBYTES_1024 + case: + return 0 + } +} + +polyvec_ntt :: proc "contextless" (r: ^Polyvec, k: int) { + for i in 0.. i16 { + QINV :: -3327 // q^-1 mod 2^16 + + t := i16(a) * QINV + return i16((a - i32(t) * Q) >> 16) +} + +@(require_results) +barrett_reduce :: #force_inline proc "contextless" (a: i16) -> i16 { + V : i16 : ((1<<26) + Q / 2) / Q + + t := i16((i32(V)*i32(a) + (1<<25)) >> 26) + t *= Q + return a - t +} diff --git a/core/crypto/_mlkem/symmetric_shake.odin b/core/crypto/_mlkem/symmetric_shake.odin new file mode 100644 index 000000000..bf851f6f6 --- /dev/null +++ b/core/crypto/_mlkem/symmetric_shake.odin @@ -0,0 +1,61 @@ +#+private +package _mlkem + +import "core:crypto" +import "core:crypto/_sha3" +import "core:crypto/sha3" +import "core:crypto/shake" + +XOF_BLOCKBYTES :: _sha3.RATE_128 +#assert(XOF_BLOCKBYTES % 3 == 0) + +prf :: proc(out, key: []byte, iv: byte) { + ctx: shake.Context = --- + defer shake.reset(&ctx) + + shake.init_256(&ctx) + shake.write(&ctx, key) + shake.write(&ctx, []byte{iv}) + shake.read(&ctx, out) +} + +rkprf :: proc(out, key, input: []byte) { + ctx: shake.Context = --- + defer shake.reset(&ctx) + + shake.init_256(&ctx) + shake.write(&ctx, key) + shake.write(&ctx, input) + shake.read(&ctx, out) +} + +xof_absorb :: proc(ctx: ^shake.Context, seed: []byte, x, y: byte) { + shake.init_128(ctx) + + extseed: [SYMBYTES+2]byte = --- + defer crypto.zero_explicit(&extseed, size_of(extseed)) + + copy(extseed[:], seed) + extseed[SYMBYTES+0] = x + extseed[SYMBYTES+1] = y + + shake.write(ctx, extseed[:]) +} + +hash_h :: proc(dst, src: []byte) { + ctx: sha3.Context = --- + + sha3.init_256(&ctx) + sha3.update(&ctx, src) + sha3.final(&ctx, dst) +} + +hash_g :: proc(dst: []byte, srcs: ..[]byte) { + ctx: sha3.Context = --- + + sha3.init_512(&ctx) + for src in srcs { + sha3.update(&ctx, src) + } + sha3.final(&ctx, dst) +} diff --git a/core/crypto/_sha3/sha3.odin b/core/crypto/_sha3/sha3.odin index b2d032b31..cfcafdc98 100644 --- a/core/crypto/_sha3/sha3.odin +++ b/core/crypto/_sha3/sha3.odin @@ -31,21 +31,24 @@ DS_SHA3 :: 0x06 DS_SHAKE :: 0x1f DS_CSHAKE :: 0x04 +DS_TURBOSHAKE_DEFAULT :: 0x1f + Context :: struct { - st: struct #raw_union { + st: struct #raw_union { b: [200]u8, q: [25]u64, }, - pt: int, - rsiz: int, - mdlen: int, - dsbyte: byte, - is_initialized: bool, - is_finalized: bool, // For SHAKE (unlimited squeeze is allowed) + pt: int, + rsiz: int, + mdlen: int, + keccak_round_start: i32, // Round at which keccak_permute starts + dsbyte: byte, + is_initialized: bool, + is_finalized: bool, // For SHAKE (unlimited squeeze is allowed) } @(private, rodata) -keccakf_rndc := [?]u64 { +keccak_rndc := [?]u64 { 0x0000000000000001, 0x0000000000008082, 0x800000000000808a, 0x8000000080008000, 0x000000000000808b, 0x0000000080000001, 0x8000000080008081, 0x8000000000008009, 0x000000000000008a, @@ -57,20 +60,23 @@ keccakf_rndc := [?]u64 { } @(private, rodata) -keccakf_rotc := [?]int { +keccak_rotc := [?]int { 1, 3, 6, 10, 15, 21, 28, 36, 45, 55, 2, 14, 27, 41, 56, 8, 25, 43, 62, 18, 39, 61, 20, 44, } @(private, rodata) -keccakf_piln := [?]i32 { +keccak_piln := [?]i32 { 10, 7, 11, 17, 18, 3, 5, 16, 8, 21, 24, 4, 15, 23, 19, 13, 12, 2, 20, 14, 22, 9, 6, 1, } +// Covers the Keccak permutation rounds, allows for a starting round parameter to support for example Keccak-p[1600,12] @(private) -keccakf :: proc "contextless" (st: ^[25]u64) { - i, j, r: i32 = ---, ---, --- +keccak_permute :: proc "contextless" (st: ^[25]u64, starting_round: i32 = 0) { + ensure_contextless((starting_round >= 0 && starting_round <= 23), "crypto/sha3: invalid Keccak permutation starting round") + + i, j, r: i32 = ---, ---, starting_round t: u64 = --- bc: [5]u64 = --- @@ -80,7 +86,7 @@ keccakf :: proc "contextless" (st: ^[25]u64) { } } - for r = 0; r < ROUNDS; r += 1 { + for ; r < ROUNDS; r += 1 { // theta for i = 0; i < 5; i += 1 { bc[i] = st[i] ~ st[i + 5] ~ st[i + 10] ~ st[i + 15] ~ st[i + 20] @@ -96,9 +102,9 @@ keccakf :: proc "contextless" (st: ^[25]u64) { // rho pi t = st[1] for i = 0; i < 24; i += 1 { - j = keccakf_piln[i] + j = keccak_piln[i] bc[0] = st[j] - st[j] = bits.rotate_left64(t, keccakf_rotc[i]) + st[j] = bits.rotate_left64(t, keccak_rotc[i]) t = bc[0] } @@ -112,7 +118,7 @@ keccakf :: proc "contextless" (st: ^[25]u64) { } } - st[0] ~= keccakf_rndc[r] + st[0] ~= keccak_rndc[r] } when ODIN_ENDIAN != .Little { @@ -142,7 +148,7 @@ update :: proc "contextless" (ctx: ^Context, data: []byte) { ctx.st.b[j] ~= data[i] j += 1 if j >= ctx.rsiz { - keccakf(&ctx.st.q) + keccak_permute(&ctx.st.q, ctx.keccak_round_start) j = 0 } } @@ -164,7 +170,7 @@ final :: proc "contextless" (ctx: ^Context, hash: []byte, finalize_clone: bool = ctx.st.b[ctx.pt] ~= ctx.dsbyte ctx.st.b[ctx.rsiz - 1] ~= 0x80 - keccakf(&ctx.st.q) + keccak_permute(&ctx.st.q, ctx.keccak_round_start) for i := 0; i < ctx.mdlen; i += 1 { hash[i] = ctx.st.b[i] } @@ -188,7 +194,7 @@ shake_xof :: proc "contextless" (ctx: ^Context) { ctx.st.b[ctx.pt] ~= ctx.dsbyte ctx.st.b[ctx.rsiz - 1] ~= 0x80 - keccakf(&ctx.st.q) + keccak_permute(&ctx.st.q, ctx.keccak_round_start) ctx.pt = 0 ctx.is_finalized = true // No more absorb, unlimited squeeze. @@ -201,11 +207,11 @@ shake_out :: proc "contextless" (ctx: ^Context, hash: []byte) { j := ctx.pt for i := 0; i < len(hash); i += 1 { if j >= ctx.rsiz { - keccakf(&ctx.st.q) + keccak_permute(&ctx.st.q, ctx.keccak_round_start) j = 0 } hash[i] = ctx.st.b[j] j += 1 } ctx.pt = j -} +} \ No newline at end of file diff --git a/core/crypto/_sha3/turboshake.odin b/core/crypto/_sha3/turboshake.odin new file mode 100644 index 000000000..e78498cbb --- /dev/null +++ b/core/crypto/_sha3/turboshake.odin @@ -0,0 +1,12 @@ +package _sha3 + +init_turboshake :: proc "contextless" (ctx: ^Context, d: byte, sec_strength: int) { + ensure_contextless((d >= 0x01 && d <= 0x7f), "crypto/sha3: invalid TurboSHAKE domain separation byte, allowed: >= 0x01 && <= 0x7f") + + ctx.mdlen = sec_strength / 8 + ctx.dsbyte = d + + init(ctx) + + ctx.keccak_round_start = 12 +} diff --git a/core/crypto/_subtle/subtle.odin b/core/crypto/_subtle/subtle.odin index 454066e4a..e57d4ca36 100644 --- a/core/crypto/_subtle/subtle.odin +++ b/core/crypto/_subtle/subtle.odin @@ -3,40 +3,248 @@ Various useful bit operations in constant time. */ package _subtle -import "core:math/bits" +import "base:intrinsics" -// byte_eq returns 1 if and only if (⟺) a == b, 0 otherwise. -@(optimization_mode="none") -byte_eq :: proc "contextless" (a, b: byte) -> int { +// Copyright (c) 2016 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +// Constant-time primitives. These functions manipulate integer values in +// order to provide constant-time comparisons and multiplexers. +// +// Boolean values (the "ctl" bits) MUST have value 0 or 1. +// +// Implementation notes: +// ===================== +// +// The uintN_t types are unsigned and with width exactly N bits; the C +// standard guarantees that computations are performed modulo 2^N, and +// there can be no overflow. Negation (unary '-') works on unsigned types +// as well. +// +// The intN_t types are guaranteed to have width exactly N bits, with no +// padding bit, and using two's complement representation. Casting +// intN_t to uintN_t really is conversion modulo 2^N. Beware that intN_t +// types, being signed, trigger implementation-defined behaviour on +// overflow (including raising some signal): with GCC, while modular +// arithmetics are usually applied, the optimizer may assume that +// overflows don't occur (unless the -fwrapv command-line option is +// added); Clang has the additional -ftrapv option to explicitly trap on +// integer overflow or underflow. + +// This code only works on a two's complement system. +#assert((-1 & 3) == 3) + +// not negates a boolean which MUST be `0` or `1` +@(optimization_mode="none", require_results) +not :: proc "contextless" (ctrl: $T) -> T where intrinsics.type_is_unsigned(T) { + return ctrl ~ 1 +} + +@(optimization_mode="none", require_results) +byte_eq :: proc "contextless" (a, b: byte) -> byte { v := a ~ b // v == 0 if and only if (⟺) a == b. The subtraction will underflow, setting the // sign bit, which will get returned. - return int((u32(v)-1) >> 31) + return byte((u32(v)-1) >> 31) } -// u64_eq returns 1 if and only if (⟺) a == b, 0 otherwise. -@(optimization_mode="none") +@(optimization_mode="none", require_results) +u32_eq :: proc "contextless" (a, b: u32) -> u32 { + q := a ~ b + return ((q | -q) >> 31) ~ 1 +} + +@(optimization_mode="none", require_results) u64_eq :: proc "contextless" (a, b: u64) -> u64 { - _, borrow := bits.sub_u64(0, a ~ b, 0) - return (~borrow) & 1 + q := a ~ b + return ((q | -q) >> 63) ~ 1 } +// eq returns 1 if and only if (⟺) a == b, 0 otherwise. eq :: proc { byte_eq, + u32_eq, u64_eq, } -// u64_is_zero returns 1 if and only if (⟺) a == 0, 0 otherwise. -@(optimization_mode="none") -u64_is_zero :: proc "contextless" (a: u64) -> u64 { - _, borrow := bits.sub_u64(a, 1, 0) - return borrow +@(require_results) +byte_neq :: proc "contextless" (a, b: byte) -> byte { + return #force_inline byte_eq(a, b) ~ 1 } -// u64_is_non_zero returns 1 if and only if (⟺) a != 0, 0 otherwise. -@(optimization_mode="none") -u64_is_non_zero :: proc "contextless" (a: u64) -> u64 { - is_zero := u64_is_zero(a) - return (~is_zero) & 1 +@(optimization_mode="none", require_results) +u32_neq :: proc "contextless" (a, b: u32) -> u32 { + q := a ~ b + return (q | -q) >> 31 +} + +@(optimization_mode="none", require_results) +u64_neq :: proc "contextless" (a, b: u64) -> u64 { + q := a ~ b + return (q | -q) >> 63 +} + +// neq returns 1 if and only if (⟺) a != b, 0 otherwise. +neq :: proc { + byte_neq, + u32_neq, + u64_neq, +} + +@(optimization_mode="none", require_results) +u32_gt :: proc "contextless" (x, y: u32) -> u32 { + /* + * If both x < 2^31 and y < 2^31, then y-x will have its high + * bit set if x > y, cleared otherwise. + * + * If either x >= 2^31 or y >= 2^31 (but not both), then the + * result is the high bit of x. + * + * If both x >= 2^31 and y >= 2^31, then we can virtually + * subtract 2^31 from both, and we are back to the first case. + * Since (y-2^31)-(x-2^31) = y-x, the subtraction is already + * fine. + */ + z := y - x + return (z ~ ((x ~ y) & (x ~ z))) >> 31 +} + +@(optimization_mode="none", require_results) +u64_gt :: proc "contextless" (x, y: u64) -> u64 { + z := y - x + return (z ~ ((x ~ y) & (x ~ z))) >> 63 +} + +// gt returns 1 if x > y, 0 otherwise. +gt :: proc { + u32_gt, + u64_gt, +} + +// gt returns 1 if x >= y, 0 otherwise. +@(require_results) +ge :: proc "contextless" (x, y: $T) -> T where T == u32 || T == u64 { + return #force_inline(gt(y, x)) ~ 1 +} + +// lt returns 1 if x < y, 0 otherwise. +@(require_results) +lt :: proc "contextless" (x, y: $T) -> T where T == u32 || T == u64 { + return #force_inline(gt(y, x)) +} + +// le returns 1 if x <= y, 0 otherwise. +@(require_results) +le :: proc "contextless" (x, y: $T) -> T where T == u32 || T == u64 { + return #force_inline(gt(x, y)) ~ 1 +} + +@(require_results) +u32_cmp :: proc "contextless" (x, y: u32) -> i32 { + return i32(#force_inline gt(x, y)) | -i32(#force_inline gt(y, x)) +} + +@(require_results) +u64_cmp :: proc "contextless" (x, y: u64) -> i64 { + return i64(#force_inline gt(x, y)) | -i64(#force_inline gt(y, x)) +} + +// cmp returns -1, 0, or 1, depending on wheter x is lower than, equal +// to, or greater than y. +cmp :: proc { + u32_cmp, + u64_cmp, +} + +// eq0 returns 1 if and only if (⟺) a == 0, 0 otherwise. +@(require_results) +eq0 :: proc "contextless" (a: $T) -> T where T == u32 || T == u64 { + return #force_inline eq(a, 0) +} + +// neq0 returns 1 if and only if (⟺) a != 0, 0 otherwise. +@(require_results) +neq0 :: proc "contextless" (a: $T) -> T where T == u32 || T == u64 { + return #force_inline eq(a, 0) ~ 1 +} + +cmov_bytes :: proc "contextless" (dst, src: []byte, #any_int ctrl: int) { + ensure_contextless(len(src) == len(dst), "crypto: cmov length mismatch") + + cmov_impl(dst, src, ctrl) +} + +cmov_u32s :: proc "contextless" (dst, src: []u32, #any_int ctrl: int) { + ensure_contextless(len(src) == len(dst), "crypto: cmov length mismatch") + + cmov_impl(dst, src, ctrl) +} + +@(private="file", optimization_mode="none") +cmov_impl :: proc "contextless"(dst, src: []$T, ctrl: int) { + s_len := len(src) + + c := -(T)(ctrl) + for i in 0.. i16 { + c := -ctrl + return a ~ i16(c & u16(a ~ b)) +} + +@(optimization_mode="none", require_results) +csel_u16 :: proc "contextless" (a, b: u16, #any_int ctrl: u16) -> u16 { + c := -ctrl + return a ~ (c & (a ~ b)) +} + +@(optimization_mode="none", require_results) +csel_u32 :: proc "contextless" (a, b: u32, #any_int ctrl: u32) -> u32 { + c := -ctrl + return a ~ (c & (a ~ b)) +} + +@(optimization_mode="none", require_results) +csel_u64 :: proc "contextless" (a, b: u64, #any_int ctrl: u64) -> u64 { + c := -ctrl + return a ~ (c & (a ~ b)) +} + +// csel returns `a` if ctl == `0`, `b` if ctl == `1`. +csel :: proc { + csel_i16, + csel_u16, + csel_u32, + csel_u64, } diff --git a/core/crypto/_weierstrass/fe.odin b/core/crypto/_weierstrass/fe.odin index 8ff6fe346..2b93d58a4 100644 --- a/core/crypto/_weierstrass/fe.odin +++ b/core/crypto/_weierstrass/fe.odin @@ -196,10 +196,10 @@ fe_gen_y_p384r1 :: proc "contextless" (fe: ^Field_Element_p384r1) { @(require_results) fe_is_zero_p256r1 :: proc "contextless" (fe: ^Field_Element_p256r1) -> int { - return int(subtle.u64_is_zero(p256r1.fe_non_zero(fe))) + return int(subtle.eq0(p256r1.fe_non_zero(fe))) } @(require_results) fe_is_zero_p384r1 :: proc "contextless" (fe: ^Field_Element_p384r1) -> int { - return int(subtle.u64_is_zero(p384r1.fe_non_zero(fe))) + return int(subtle.eq0(p384r1.fe_non_zero(fe))) } diff --git a/core/crypto/_weierstrass/sc.odin b/core/crypto/_weierstrass/sc.odin index 1ecfea6f9..6a3c7b68c 100644 --- a/core/crypto/_weierstrass/sc.odin +++ b/core/crypto/_weierstrass/sc.odin @@ -133,10 +133,10 @@ sc_is_zero :: proc { @(require_results) sc_is_zero_p256r1 :: proc "contextless" (fe: ^Scalar_p256r1) -> int { - return int(subtle.u64_is_zero(p256r1.fe_non_zero(fe))) + return int(subtle.eq0(p256r1.fe_non_zero(fe))) } @(require_results) sc_is_zero_p384r1 :: proc "contextless" (fe: ^Scalar_p384r1) -> int { - return int(subtle.u64_is_zero(p384r1.fe_non_zero(fe))) + return int(subtle.eq0(p384r1.fe_non_zero(fe))) } diff --git a/core/crypto/_weierstrass/scalar_mul.odin b/core/crypto/_weierstrass/scalar_mul.odin index eb3fa1459..1af3ba989 100644 --- a/core/crypto/_weierstrass/scalar_mul.odin +++ b/core/crypto/_weierstrass/scalar_mul.odin @@ -293,7 +293,7 @@ when crypto.COMPACT_IMPLS == false { // conditionally select the right result. pt_add_mixed(tmp, point, &tmp.x, &tmp.y) - ctrl := subtle.u64_is_non_zero(idx) + ctrl := subtle.neq0(idx) pt_cond_select(point, point, tmp, int(ctrl)) } } diff --git a/core/crypto/aes/aes_ecb.odin b/core/crypto/aes/aes_ecb.odin index cac62de5d..9ec1a9a37 100644 --- a/core/crypto/aes/aes_ecb.odin +++ b/core/crypto/aes/aes_ecb.odin @@ -21,7 +21,7 @@ init_ecb :: proc(ctx: ^Context_ECB, key: []byte, impl := DEFAULT_IMPLEMENTATION) encrypt_ecb :: proc(ctx: ^Context_ECB, dst, src: []byte) { ensure(ctx._is_initialized) ensure(len(dst) == BLOCK_SIZE, "crypto/aes: invalid dst size") - ensure(len(dst) == BLOCK_SIZE, "crypto/aes: invalid src size") + ensure(len(src) == BLOCK_SIZE, "crypto/aes: invalid src size") switch &impl in ctx._impl { case ct64.Context: @@ -35,7 +35,7 @@ encrypt_ecb :: proc(ctx: ^Context_ECB, dst, src: []byte) { decrypt_ecb :: proc(ctx: ^Context_ECB, dst, src: []byte) { ensure(ctx._is_initialized) ensure(len(dst) == BLOCK_SIZE, "crypto/aes: invalid dst size") - ensure(len(dst) == BLOCK_SIZE, "crypto/aes: invalid src size") + ensure(len(src) == BLOCK_SIZE, "crypto/aes: invalid src size") switch &impl in ctx._impl { case ct64.Context: diff --git a/core/crypto/argon2id/argon2id.odin b/core/crypto/argon2id/argon2id.odin index 3bff5a3a9..e2f5e487b 100644 --- a/core/crypto/argon2id/argon2id.odin +++ b/core/crypto/argon2id/argon2id.odin @@ -143,7 +143,7 @@ derive :: proc( m_ := 4 * u64(p) * (m / u64(4 * p)) b := mem.alloc_bytes_non_zeroed( int(m_) * BLOCK_SIZE_BYTES, - alignment = mem.DEFAULT_PAGE_SIZE, + alignment = mem.PAGE_SIZE, allocator = allocator, ) or_return defer delete(b, allocator) diff --git a/core/crypto/crypto.odin b/core/crypto/crypto.odin index f4ddbfbe7..fff66fdd0 100644 --- a/core/crypto/crypto.odin +++ b/core/crypto/crypto.odin @@ -16,7 +16,8 @@ HAS_RAND_BYTES :: runtime.HAS_RAND_BYTES // // The execution time of this routine is constant regardless of the contents // of the slices being compared, as long as the length of the slices is equal. -// If the length of the two slices is dif and only if (⟺)erent, it will early-return 0. +// If and only if (⟺) the length of the two slices is diferent, it will +// early-return 0. compare_constant_time :: proc "contextless" (a, b: []byte) -> int { // If the length of the slices is different, early return. // @@ -48,7 +49,7 @@ compare_byte_ptrs_constant_time :: proc "contextless" (a, b: ^byte, n: int) -> i // After the loop, v == 0 if and only if (⟺) a == b. The subtraction will underflow // if and only if (⟺) v == 0, setting the sign-bit, which gets returned. - return subtle.eq(0, v) + return int(subtle.eq(0, v)) } // is_zero_constant_time returns 1 if and only if (⟺) b is all 0s, 0 otherwise. @@ -58,7 +59,7 @@ is_zero_constant_time :: proc "contextless" (b: []byte) -> int { v |= b_ } - return subtle.byte_eq(0, v) + return int(subtle.byte_eq(0, v)) } /* @@ -84,18 +85,6 @@ zero_explicit :: proc "contextless" (data: rawptr, len: int) -> rawptr { return data } -/* -Set each byte of a memory range to a specific value. - -This procedure copies value specified by the `value` parameter into each of the -`len` bytes of a memory range, located at address `data`. - -This procedure returns the pointer to `data`. -*/ -set :: proc "contextless" (data: rawptr, value: byte, len: int) -> rawptr { - return runtime.memset(data, i32(value), len) -} - // rand_bytes fills the dst buffer with cryptographic entropy taken from // the system entropy source. This routine will block if the system entropy // source is not ready yet. All system entropy source failures are treated diff --git a/core/crypto/ecdh/ecdh.odin b/core/crypto/ecdh/ecdh.odin index 6a8f6e466..3da3e0e41 100644 --- a/core/crypto/ecdh/ecdh.odin +++ b/core/crypto/ecdh/ecdh.odin @@ -106,6 +106,7 @@ Public_Key :: struct { // private_key_generate uses the system entropy source to generate a new // Private_Key. This will only fail if and only if (⟺) the system entropy source is // missing or broken. +@(require_results) private_key_generate :: proc(priv_key: ^Private_Key, curve: Curve) -> bool { private_key_clear(priv_key) @@ -143,6 +144,7 @@ private_key_generate :: proc(priv_key: ^Private_Key, curve: Curve) -> bool { // private_key_set_bytes decodes a byte-encoded private key, and returns // true if and only if (⟺) the operation was successful. +@(require_results) private_key_set_bytes :: proc(priv_key: ^Private_Key, curve: Curve, b: []byte) -> bool { private_key_clear(priv_key) @@ -222,6 +224,40 @@ private_key_generate_public :: proc(priv_key: ^Private_Key) { priv_key._pub_key._curve = priv_key._curve } +// private_key_set sets priv_key to src. +private_key_set :: proc(priv_key, src: ^Private_Key) { + if src == nil || src._curve == .Invalid { + private_key_clear(priv_key) + return + } + + priv_key._curve = src._curve + + reflect.set_union_variant_typeid( + priv_key._impl, + _PRIV_IMPL_IDS[priv_key._curve], + ) + + #partial switch priv_key._curve { + case .SECP256R1: + secec.sc_set(&priv_key._impl.(secec.Scalar_p256r1), &src._impl.(secec.Scalar_p256r1)) + case .SECP384R1: + secec.sc_set(&priv_key._impl.(secec.Scalar_p384r1), &src._impl.(secec.Scalar_p384r1)) + case .X25519: + priv_buf := &(priv_key._impl.(X25519_Buf)) + src_buf := &(src._impl.(X25519_Buf)) + copy(priv_buf[:], src_buf[:]) + case .X448: + priv_buf := &(priv_key._impl.(X448_Buf)) + src_buf := &(src._impl.(X448_Buf)) + copy(priv_buf[:], src_buf[:]) + case: + panic("crypto/ecdh: invalid curve") + } + + public_key_set(&priv_key._pub_key, &src._pub_key) +} + // private_key_bytes sets dst to byte-encoding of priv_key. private_key_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { ensure(priv_key._curve != .Invalid, "crypto/ecdh: uninitialized private key") @@ -245,8 +281,15 @@ private_key_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { } } +// private_key_public_bytes sets dst to the byte-encoding of the public +// key corresponding to priv_key. +private_key_public_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { + public_key_bytes(&priv_key._pub_key, dst) +} + // private_key_equal returns true if and only if (⟺) the private keys are equal, // in constant time. +@(require_results) private_key_equal :: proc(p, q: ^Private_Key) -> bool { if p._curve != q._curve { return false @@ -277,6 +320,7 @@ private_key_clear :: proc "contextless" (priv_key: ^Private_Key) { // public_key_set_bytes decodes a byte-encoded public key, and returns // true if and only if (⟺) the operation was successful. +@(require_results) public_key_set_bytes :: proc(pub_key: ^Public_Key, curve: Curve, b: []byte) -> bool { public_key_clear(pub_key) @@ -325,6 +369,16 @@ public_key_set_bytes :: proc(pub_key: ^Public_Key, curve: Curve, b: []byte) -> b return true } +// public_key_set sets pub_key to src. +public_key_set :: proc(pub_key, src: ^Public_Key) { + if src == nil || src._curve == .Invalid { + public_key_clear(pub_key) + return + } + + pub_key^ = src^ +} + // public_key_set_priv sets pub_key to the public component of priv_key. public_key_set_priv :: proc(pub_key: ^Public_Key, priv_key: ^Private_Key) { ensure(priv_key._curve != .Invalid, "crypto/ecdh: uninitialized private key") @@ -367,6 +421,7 @@ public_key_bytes :: proc(pub_key: ^Public_Key, dst: []byte) { // public_key_equal returns true if and only if (⟺) the public keys are equal, // in constant time. +@(require_results) public_key_equal :: proc(p, q: ^Public_Key) -> bool { if p._curve != q._curve { return false @@ -435,12 +490,14 @@ ecdh :: proc(priv_key: ^Private_Key, pub_key: ^Public_Key, dst: []byte) -> bool } // curve returns the Curve used by a Private_Key or Public_Key instance. -curve :: proc(k: ^$T) -> Curve where(T == Private_Key || T == Public_Key) { +@(require_results) +curve :: proc(k: ^$T) -> Curve where (T == Private_Key || T == Public_Key) { return k._curve } // key_size returns the key size of a Private_Key or Public_Key in bytes. -key_size :: proc(k: ^$T) -> int where(T == Private_Key || T == Public_Key) { +@(require_results) +key_size :: proc(k: ^$T) -> int where (T == Private_Key || T == Public_Key) { when T == Private_Key { return PRIVATE_KEY_SIZES[k._curve] } else { @@ -450,6 +507,7 @@ key_size :: proc(k: ^$T) -> int where(T == Private_Key || T == Public_Key) { // shared_secret_size returns the shared secret size of a key exchange // in bytes. -shared_secret_size :: proc(k: ^$T) -> int where(T == Private_Key || T == Public_Key) { +@(require_results) +shared_secret_size :: proc(k: ^$T) -> int where (T == Private_Key || T == Public_Key) { return SHARED_SECRET_SIZES[k._curve] } diff --git a/core/crypto/ecdsa/ecdsa.odin b/core/crypto/ecdsa/ecdsa.odin index 6c71feef7..e63539fb1 100644 --- a/core/crypto/ecdsa/ecdsa.odin +++ b/core/crypto/ecdsa/ecdsa.odin @@ -2,7 +2,6 @@ package ecdsa import "core:crypto" import secec "core:crypto/_weierstrass" -import "core:mem" import "core:reflect" // Curve the curve identifier associated with a given Private_Key @@ -81,6 +80,7 @@ Public_Key :: struct { // private_key_generate uses the system entropy source to generate a new // Private_Key. This will only fail if and only if (⟺) the system entropy source is // missing or broken. +@(require_results) private_key_generate :: proc(priv_key: ^Private_Key, curve: Curve) -> bool { private_key_clear(priv_key) @@ -112,6 +112,7 @@ private_key_generate :: proc(priv_key: ^Private_Key, curve: Curve) -> bool { // private_key_set_bytes decodes a byte-encoded private key, and returns // true if and only if (⟺) the operation was successful. +@(require_results) private_key_set_bytes :: proc(priv_key: ^Private_Key, curve: Curve, b: []byte) -> bool { private_key_clear(priv_key) @@ -194,8 +195,41 @@ private_key_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { } } +// private_key_public_bytes sets dst to the byte-encoding of the public +// key corresponding to priv_key. +private_key_public_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { + public_key_bytes(&priv_key._pub_key, dst) +} + +// private_key_set sets priv_key to src. +private_key_set :: proc(priv_key, src: ^Private_Key) { + if src == nil || src._curve == .Invalid { + private_key_clear(priv_key) + return + } + + priv_key._curve = src._curve + + reflect.set_union_variant_typeid( + priv_key._impl, + _PRIV_IMPL_IDS[priv_key._curve], + ) + + #partial switch priv_key._curve { + case .SECP256R1: + secec.sc_set(&priv_key._impl.(secec.Scalar_p256r1), &src._impl.(secec.Scalar_p256r1)) + case .SECP384R1: + secec.sc_set(&priv_key._impl.(secec.Scalar_p384r1), &src._impl.(secec.Scalar_p384r1)) + case: + panic("crypto/ecdh: invalid curve") + } + + public_key_set(&priv_key._pub_key, &src._pub_key) +} + // private_key_equal returns true if and only if (⟺) the private keys are equal, // in constant time. +@(require_results) private_key_equal :: proc(p, q: ^Private_Key) -> bool { if p._curve != q._curve { return false @@ -215,11 +249,12 @@ private_key_equal :: proc(p, q: ^Private_Key) -> bool { // private_key_clear clears priv_key to the uninitialized state. private_key_clear :: proc "contextless" (priv_key: ^Private_Key) { - mem.zero_explicit(priv_key, size_of(Private_Key)) + crypto.zero_explicit(priv_key, size_of(Private_Key)) } // public_key_set_bytes decodes a byte-encoded public key, and returns // true if and only if (⟺) the operation was successful. +@(require_results) public_key_set_bytes :: proc(pub_key: ^Public_Key, curve: Curve, b: []byte) -> bool { public_key_clear(pub_key) @@ -262,6 +297,16 @@ public_key_set_bytes :: proc(pub_key: ^Public_Key, curve: Curve, b: []byte) -> b return true } +// public_key_set sets pub_key to src. +public_key_set :: proc(pub_key, src: ^Public_Key) { + if src == nil || src._curve == .Invalid { + public_key_clear(pub_key) + return + } + + pub_key^ = src^ +} + // public_key_set_priv sets pub_key to the public component of priv_key. public_key_set_priv :: proc(pub_key: ^Public_Key, priv_key: ^Private_Key) { ensure(priv_key._curve != .Invalid, "crypto/ecdsa: uninitialized private key") @@ -298,6 +343,7 @@ public_key_bytes :: proc(pub_key: ^Public_Key, dst: []byte) { // public_key_equal returns true if and only if (⟺) the public keys are equal, // in constant time. +@(require_results) public_key_equal :: proc(p, q: ^Public_Key) -> bool { if p._curve != q._curve { return false @@ -317,5 +363,21 @@ public_key_equal :: proc(p, q: ^Public_Key) -> bool { // public_key_clear clears pub_key to the uninitialized state. public_key_clear :: proc "contextless" (pub_key: ^Public_Key) { - mem.zero_explicit(pub_key, size_of(Public_Key)) + crypto.zero_explicit(pub_key, size_of(Public_Key)) +} + +// curve returns the Curve used by a Private_Key or Public_Key instance. +@(require_results) +curve :: proc(k: ^$T) -> Curve where (T == Private_Key || T == Public_Key) { + return k._curve +} + +// key_size returns the key size of a Private_Key or Public_Key in bytes. +@(require_results) +key_size :: proc(k: ^$T) -> int where (T == Private_Key || T == Public_Key) { + when T == Private_Key { + return PRIVATE_KEY_SIZES[k._curve] + } else { + return PUBLIC_KEY_SIZES[k._curve] + } } diff --git a/core/crypto/ecdsa/ecdsa_asn1.odin b/core/crypto/ecdsa/ecdsa_asn1.odin index 74c9d65e6..0bad6a263 100644 --- a/core/crypto/ecdsa/ecdsa_asn1.odin +++ b/core/crypto/ecdsa/ecdsa_asn1.odin @@ -1,17 +1,13 @@ package ecdsa +import "core:encoding/asn1" import secec "core:crypto/_weierstrass" -// ASN.1 format ECDSA signatures are`SEQUENCE { r INTEGER, s INTEGER }` -// this implements enough to generate/parse signatures. Eventually when -// we have a full ASN.1 DER library, these routines will be removed. +// ASN.1 ECDSA signatures are `SEQUENCE { r INTEGER, s INTEGER }`. These thin +// wrappers over core:encoding/asn1 generate/parse that structure the DER +// minimal-encoding rules. -@(private="file") -TAG_SEQUENCE :: 0x30 -@(private="file") -TAG_INTEGER :: 0x02 - -@(private,require_results) +@(private, require_results) generate_asn1_sig :: proc(r, s: ^$T, allocator := context.allocator) -> []byte { when T == secec.Scalar_p256r1 { SC_SZ :: secec.SC_SIZE_P256R1 @@ -21,163 +17,35 @@ generate_asn1_sig :: proc(r, s: ^$T, allocator := context.allocator) -> []byte { #panic("crypto/ecdsa: invalid curve") } - INT_TLP :: 3 // tag, tength, (optional) leading zero-byte - encode_uint :: proc(b: []byte) -> []byte { - b := b + r_buf, s_buf: [SC_SZ]byte = ---, --- + secec.sc_bytes(r_buf[:], r) + secec.sc_bytes(s_buf[:], s) - // DER requires minimal encoding. - off := INT_TLP - for v in b[off:] { - if v != 0 { - break - } - off += 1 - } - // If the sign big is set, add a leading zero. - if b[off] & 0x80 == 0x80 { - off -= 1 - b[off] = 0 - } - - // Encode the length (up to 127 octets, adequate for ECDSA). - l := len(b[off:]) - off -= 1 - b[off] = byte(l) - - // Encode the tag - off -= 1 - b[off] = TAG_INTEGER - - return b[off:] + sig, err := asn1.marshal( + asn1.sequence({asn1.integer_unsigned(r_buf[:]), asn1.integer_unsigned(s_buf[:])}), + allocator, + ) + if err != .None { + return nil } - - r_buf, s_buf: [INT_TLP+SC_SZ]byte = ---, --- - secec.sc_bytes(r_buf[INT_TLP:], r) - secec.sc_bytes(s_buf[INT_TLP:], s) - - r_bytes, s_bytes := encode_uint(r_buf[:]), encode_uint(s_buf[:]) - seq_len := len(r_bytes) + len(s_bytes) - - // WARNING: If secp521r1 support is added, this needs to support - // long-form length encoding. - ensure(seq_len <= 127, "BUG: crypto/ecdsa: signature length too large") - b := make([]byte, seq_len + 2, allocator) - b[0] = TAG_SEQUENCE - b[1] = byte(seq_len) - copy(b[2:], r_bytes) - copy(b[2+len(r_bytes):], s_bytes) - - return b + return sig } -@(private,require_results) +@(private, require_results) parse_asn1_sig :: proc(sig: []byte) -> (r, s: []byte, ok: bool) { - read_seq :: proc(b: []byte) -> (v: []byte, rest: []byte, ok: bool) { - b_len := len(b) - if b_len < 3 { - return nil, nil, false - } - if b[0] != TAG_SEQUENCE { - return nil, nil, false - } - seq_len, off: int - if b[1] & 0x80 == 0x80 { - if b[1] != 0x81 || b_len < 4 { // 2-length octets is sufficient for ecdsa. - return nil, nil, false - } - if b[2] & 0x80 == 0x80 || b[3] & 0x80 == 80 { - return nil, nil, false - } - seq_len = int(b[2]) * 127 + int(b[3]) - off = 4 - } else { - seq_len = int(b[1]) - off = 2 - } - if b_len - off < seq_len { - return nil, nil, false - } - return b[off:off+seq_len], b[off+seq_len:], true - } - - read_int :: proc(b: []byte) -> (v: []byte, rest: []byte, ok: bool) { - b_len := len(b) - if b_len < 3 { - return nil, nil, false - } - if b[0] != TAG_INTEGER { - return nil, nil, false - } - v_len := int(b[1]) - if v_len > 0x80 || b_len - 2 < v_len { // 127-bytes max. - return nil, nil, false - } - - return b[2:2+v_len], b[2+v_len:], true - } - - // SEQUENCE - seq_bytes, rest: []byte - seq_bytes, rest, ok = read_seq(sig) - if !ok { - return nil, nil, false - } - if len(rest) != 0 { + cur: asn1.Cursor + asn1.cursor_init(&cur, sig) + seq, e0 := asn1.read_sequence(&cur) + if e0 != .None || asn1.done(&cur) != .None { return nil, nil, false } - // INTEGER (r) - r, rest, ok = read_int(seq_bytes) - if !ok { + // r and s are unsigned; read_unsigned_integer_bytes validates the INTEGER + // and strips the DER sign octet, returning the magnitude as a view of sig. + rb, e1 := asn1.read_unsigned_integer_bytes(&seq) + sb, e2 := asn1.read_unsigned_integer_bytes(&seq) + if e1 != .None || e2 != .None || asn1.done(&seq) != .None { return nil, nil, false } - - // INTEGER (s) - s, rest, ok = read_int(rest) - if !ok { - return nil, nil, false - } - if len(rest) != 0 { - return nil, nil, false - } - - // DER requires a leading 0 if and only if (⟺) the sign bit of the leading byte - // is set to distinguish between positive and negative integers, - // and the minimal length representation. `r` and `s` are always - // going to be unsigned, so we validate malformed DER and strip - // the leading 0 as needed. - fixup_der_uint :: proc(b: []byte) -> ([]byte, bool) { - switch len(b) { - case 0: - // 0 length is invalid - return nil, false - case 1: - // Missing leading zero - if b[0] & 0x80 == 0x80 { - return nil, false - } - case: - if b[0] == 0 { - // Sign bit not set - if b[1] & 0x80 != 0x80 { - return nil, false - } - return b[1:], true - } else if b[0] & 0x80 == 0x80 { - // Missing leading zero - return nil, false - } - } - - return b, true - } - - if r, ok = fixup_der_uint(r); !ok { - return nil, nil, false - } - if s, ok = fixup_der_uint(s); !ok { - return nil, nil, false - } - - return r, s, true + return rb, sb, true } diff --git a/core/crypto/ecdsa/ecdsa_sign.odin b/core/crypto/ecdsa/ecdsa_sign.odin index c6fec56dc..a594bc601 100644 --- a/core/crypto/ecdsa/ecdsa_sign.odin +++ b/core/crypto/ecdsa/ecdsa_sign.odin @@ -13,8 +13,8 @@ import secec "core:crypto/_weierstrass" // The signature format is ASN1. `SEQUECE `{ r INTEGER, s INTEGER }`. @(require_results) sign_asn1 :: proc(priv_key: ^Private_Key, hash_algo: hash.Algorithm, msg: []byte, allocator: runtime.Allocator, deterministic := !crypto.HAS_RAND_BYTES) -> ([]byte, bool) { - ensure(hash_algo != .Invalid, "crypto/edsa: invalid hash algorithm") - ensure(priv_key._curve != .Invalid, "crypto/edsa: invalid curve") + ensure(hash_algo != .Invalid, "crypto/ecdsa: invalid hash algorithm") + ensure(priv_key._curve != .Invalid, "crypto/ecdsa: invalid curve") if !deterministic && !crypto.HAS_RAND_BYTES { return nil, false @@ -49,8 +49,8 @@ sign_asn1 :: proc(priv_key: ^Private_Key, hash_algo: hash.Algorithm, msg: []byte // The signature format is `r | s`. @(require_results) sign_raw :: proc(priv_key: ^Private_Key, hash_algo: hash.Algorithm, msg, sig: []byte, deterministic := !crypto.HAS_RAND_BYTES) -> bool { - ensure(hash_algo != .Invalid, "crypto/edsa: invalid hash algorithm") - ensure(priv_key._curve != .Invalid, "crypto/edsa: invalid curve") + ensure(hash_algo != .Invalid, "crypto/ecdsa: invalid hash algorithm") + ensure(priv_key._curve != .Invalid, "crypto/ecdsa: invalid curve") ensure(len(sig) == RAW_SIGNATURE_SIZES[priv_key._curve], "crypto/ecdsa: invalid destination size") if !deterministic && !crypto.HAS_RAND_BYTES { diff --git a/core/crypto/ecdsa/ecdsa_verify.odin b/core/crypto/ecdsa/ecdsa_verify.odin index bd973a8df..ddc1df9e6 100644 --- a/core/crypto/ecdsa/ecdsa_verify.odin +++ b/core/crypto/ecdsa/ecdsa_verify.odin @@ -10,8 +10,8 @@ import secec "core:crypto/_weierstrass" // The signature format is `r | s`. @(require_results) verify_raw :: proc(pub_key: ^Public_Key, hash_algo: hash.Algorithm, msg, sig: []byte) -> bool { - ensure(hash_algo != .Invalid, "crypto/edsa: invalid hash algorithm") - ensure(pub_key._curve != .Invalid, "crypto/edsa: invalid curve") + ensure(hash_algo != .Invalid, "crypto/ecdsa: invalid hash algorithm") + ensure(pub_key._curve != .Invalid, "crypto/ecdsa: invalid curve") if len(sig) != RAW_SIGNATURE_SIZES[pub_key._curve] { return false @@ -40,8 +40,8 @@ verify_raw :: proc(pub_key: ^Public_Key, hash_algo: hash.Algorithm, msg, sig: [] // The signature format is ASN.1 `SEQUENCE { r INTEGER, s INTEGER }`. @(require_results) verify_asn1 :: proc(pub_key: ^Public_Key, hash_algo: hash.Algorithm, msg, sig: []byte) -> bool { - ensure(hash_algo != .Invalid, "crypto/edsa: invalid hash algorithm") - ensure(pub_key._curve != .Invalid, "crypto/edsa: invalid curve") + ensure(hash_algo != .Invalid, "crypto/ecdsa: invalid hash algorithm") + ensure(pub_key._curve != .Invalid, "crypto/ecdsa: invalid curve") r_bytes, s_bytes, ok := parse_asn1_sig(sig) if !ok { diff --git a/core/crypto/ed25519/ed25519.odin b/core/crypto/ed25519/ed25519.odin index 2020c0633..cc7f93e07 100644 --- a/core/crypto/ed25519/ed25519.odin +++ b/core/crypto/ed25519/ed25519.odin @@ -50,6 +50,7 @@ Public_Key :: struct { // private_key_generate uses the system entropy source to generate a new // Private_Key. This will only fail if and only if (⟺) the system entropy source is // missing or broken. +@(require_results) private_key_generate :: proc(priv_key: ^Private_Key) -> bool { private_key_clear(priv_key) @@ -61,13 +62,12 @@ private_key_generate :: proc(priv_key: ^Private_Key) -> bool { defer crypto.zero_explicit(&b, size_of(b)) crypto.rand_bytes(b[:]) - private_key_set_bytes(priv_key, b[:]) - - return true + return private_key_set_bytes(priv_key, b[:]) } // private_key_set_bytes decodes a byte-encoded private key, and returns // true if and only if (⟺) the operation was successful. +@(require_results) private_key_set_bytes :: proc(priv_key: ^Private_Key, b: []byte) -> bool { if len(b) != PRIVATE_KEY_SIZE { return false @@ -97,6 +97,21 @@ private_key_set_bytes :: proc(priv_key: ^Private_Key, b: []byte) -> bool { return true } +// private_key_set sets priv_key to src. +private_key_set :: proc(priv_key, src: ^Private_Key) { + if src == nil || !src._is_initialized { + private_key_clear(priv_key) + return + } + + copy(priv_key._b[:], src._b[:]) + grp.sc_set(&priv_key._s, &src._s) + copy(priv_key._hdigest2[:], src._hdigest2[:]) + public_key_set(&priv_key._pub_key, &src._pub_key) + + priv_key._is_initialized = true +} + // private_key_bytes sets dst to byte-encoding of priv_key. private_key_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { ensure(priv_key._is_initialized, "crypto/ed25519: uninitialized private key") @@ -105,6 +120,12 @@ private_key_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { copy(dst, priv_key._b[:]) } +// private_key_public_bytes sets dst to the byte-encoding of the public +// key corresponding to priv_key. +private_key_public_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { + public_key_bytes(&priv_key._pub_key, dst) +} + // private_key_clear clears priv_key to the uninitialized state. private_key_clear :: proc "contextless" (priv_key: ^Private_Key) { crypto.zero_explicit(priv_key, size_of(Private_Key)) @@ -168,6 +189,7 @@ sign :: proc(priv_key: ^Private_Key, msg, sig: []byte) { // public_key_set_bytes decodes a byte-encoded public key, and returns // true if and only if (⟺) the operation was successful. +@(require_results) public_key_set_bytes :: proc "contextless" (pub_key: ^Public_Key, b: []byte) -> bool { if len(b) != PUBLIC_KEY_SIZE { return false @@ -186,6 +208,16 @@ public_key_set_bytes :: proc "contextless" (pub_key: ^Public_Key, b: []byte) -> return true } +// public_key_set sets pub_key to src. +public_key_set :: proc(pub_key, src: ^Public_Key) { + if src == nil || !src._is_initialized { + public_key_clear(pub_key) + return + } + + pub_key^ = src^ +} + // public_key_set_priv sets pub_key to the public component of priv_key. public_key_set_priv :: proc(pub_key: ^Public_Key, priv_key: ^Private_Key) { ensure(priv_key._is_initialized, "crypto/ed25519: uninitialized private key") @@ -206,18 +238,25 @@ public_key_bytes :: proc(pub_key: ^Public_Key, dst: []byte) { } // public_key_equal returns true if and only if (⟺) pub_key is equal to other. +@(require_results) public_key_equal :: proc(pub_key, other: ^Public_Key) -> bool { ensure(pub_key._is_initialized && other._is_initialized, "crypto/ed25519: uninitialized public key") return crypto.compare_constant_time(pub_key._b[:], other._b[:]) == 1 } +// public_key_clear clears pub_key to the uninitialized state. +public_key_clear :: proc "contextless" (pub_key: ^Public_Key) { + crypto.zero_explicit(pub_key, size_of(Public_Key)) +} + // verify returns true if and only if (⟺) sig is a valid signature by pub_key over msg. // // The optional `allow_small_order_A` parameter will make this // implementation strictly compatible with FIPS 186-5, at the expense of // SBS-security. Doing so is NOT recommended, and the disallowed // public keys all have a known discrete-log. +@(require_results) verify :: proc(pub_key: ^Public_Key, msg, sig: []byte, allow_small_order_A := false) -> bool { switch { case !pub_key._is_initialized: diff --git a/core/crypto/legacy/md5/md5.odin b/core/crypto/legacy/md5/md5.odin index 4bbc5d32a..ddc795b7d 100644 --- a/core/crypto/legacy/md5/md5.odin +++ b/core/crypto/legacy/md5/md5.odin @@ -18,6 +18,7 @@ package md5 zhibog, dotbmp: Initial implementation. */ +import "base:intrinsics" import "core:crypto" import "core:encoding/endian" import "core:math/bits" @@ -100,7 +101,7 @@ final :: proc(ctx: ^Context, hash: []byte, finalize_clone: bool = false) { i += 1 } transform(ctx, ctx.data[:]) - crypto.set(&ctx.data, 0, 56) + intrinsics.mem_zero(&ctx.data, 56) } ctx.bitlen += u64(ctx.datalen * 8) diff --git a/core/crypto/legacy/sha1/sha1.odin b/core/crypto/legacy/sha1/sha1.odin index 892f893a6..bf3ad9602 100644 --- a/core/crypto/legacy/sha1/sha1.odin +++ b/core/crypto/legacy/sha1/sha1.odin @@ -19,6 +19,7 @@ package sha1 zhibog, dotbmp: Initial implementation. */ +import "base:intrinsics" import "core:crypto" import "core:encoding/endian" import "core:math/bits" @@ -107,7 +108,7 @@ final :: proc(ctx: ^Context, hash: []byte, finalize_clone: bool = false) { i += 1 } transform(ctx, ctx.data[:]) - crypto.set(&ctx.data, 0, 56) + intrinsics.mem_zero(&ctx.data, 56) } ctx.bitlen += u64(ctx.datalen * 8) diff --git a/core/crypto/mldsa/api.odin b/core/crypto/mldsa/api.odin new file mode 100644 index 000000000..ce4f9cffd --- /dev/null +++ b/core/crypto/mldsa/api.odin @@ -0,0 +1,290 @@ +package mldsa + +import "core:crypto" +import "core:crypto/_mldsa" + +// Parameters are the supported ML-DSA parameter sets. +Parameters :: enum { + Invalid, + ML_DSA_44, + ML_DSA_65, + ML_DSA_87, +} + +// PRIVATE_KEY_SEED_SIZE is the size of a private key in bytes. +PRIVATE_KEY_SEED_SIZE :: _mldsa.SEEDBYTES // 32-bytes + +// MAX_CTX_SIZE is the maximum size of the signature context +// (domain separation tag) in bytes. +MAX_CTX_SIZE :: _mldsa.CTXBYTES_MAX // 255-bytes + +// PUBLIC_KEY_SIZES are the per-parameter sizes of a public +// key in bytes. +PUBLIC_KEY_SIZES := [Parameters]int { + .Invalid = 0, + .ML_DSA_44 = 1312, + .ML_DSA_65 = 1952, + .ML_DSA_87 = 2592, +} + +// SIGNATURE_SIZES are the per-parameter sizes of a signature +// in byte. +SIGNATURE_SIZES := [Parameters]int { + .Invalid = 0, + .ML_DSA_44 = 2420, + .ML_DSA_65 = 3309, + .ML_DSA_87 = 4627, +} + +@(private="file") +_PARAMS_TO_INTERNAL := [Parameters]^_mldsa.Params { + .Invalid = nil, + .ML_DSA_44 = &_mldsa.Params_44, + .ML_DSA_65 = &_mldsa.Params_65, + .ML_DSA_87 = &_mldsa.Params_87, +} + +// Private_Key is a ML-DSA private key. +Private_Key :: _mldsa.Private_Key + +// Public_Key is a ML-DSA public key. +Public_Key :: _mldsa.Public_Key + +// private_key_generate uses the system entropy source to generate a new +// Private_Key. This will only fail if and only if (⟺) the system entropy +// source is missing or broken. +@(require_results) +private_key_generate :: proc(priv_key: ^Private_Key, params: Parameters) -> bool { + private_key_clear(priv_key) + + if !crypto.HAS_RAND_BYTES { + return false + } + + params_ := _PARAMS_TO_INTERNAL[params] + if params_ == nil { + return false + } + + seed: [PRIVATE_KEY_SEED_SIZE]byte = --- + defer crypto.zero_explicit(&seed, size_of(seed)) + + crypto.rand_bytes(seed[:]) + + _mldsa.dsa_keygen_internal(priv_key, seed[:], params_) + + return true +} + +// private_key_set_bytes decodes a byte-encoded private key in "seed" format, +// and returns true if and only if (⟺) the operation was successful. +@(require_results) +private_key_set_bytes :: proc(priv_key: ^Private_Key, params: Parameters, b: []byte) -> bool { + private_key_clear(priv_key) + + params_ := _PARAMS_TO_INTERNAL[params] + if params_ == nil { + return false + } + if len(b) != PRIVATE_KEY_SEED_SIZE { + return false + } + + _mldsa.dsa_keygen_internal(priv_key, b, params_) + + return true +} + +// private_key_bytes sets dst to byte-encoding of priv_key in the "seed" +// format. +private_key_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { + ensure(priv_key.params != nil, "crypto/mldsa: uninitialized private key") + ensure(len(dst) == PRIVATE_KEY_SEED_SIZE, "crypto/mldsa: invalid destination size") + + copy(dst, priv_key.seed[:]) +} + +// private_key_public_bytes sets dst to the byte-encoding of the public +// key corresponding to priv_key. +private_key_public_bytes :: proc(priv_key: ^Private_Key, dst: []byte) { + public_key_bytes(&priv_key.pub_key, dst) +} + +// private_key_set sets priv_key to src. +private_key_set :: proc(priv_key, src: ^Private_Key) { + if src == nil || internal_to_params(src.params) == .Invalid { + private_key_clear(priv_key) + return + } + + _mldsa.set_sk(priv_key, src) +} + +// private_key_equal returns true if and only if (⟺) the private keys are +// equal, in constant time. +@(require_results) +private_key_equal :: proc(p, q: ^Private_Key) -> bool { + if p.params != q.params { + return false + } + if p.params == nil { + return true + } + + // Just compare the seed that was passed to dsa_keygen_internal, + // since the process is completely deterministic. + return crypto.compare_constant_time(p.seed[:], q.seed[:]) == 1 +} + +// private_key_clear clears priv_key to the uninitialized state. +private_key_clear :: proc "contextless" (priv_key: ^Private_Key) { + _mldsa.clear_sk(priv_key) +} + +// public_key_set_bytes decodes a byte-encoded public key, and returns +// true if and only if (⟺) the operation was successful. +@(require_results) +public_key_set_bytes :: proc(pub_key: ^Public_Key, params: Parameters, b: []byte) -> bool { + params_ := _PARAMS_TO_INTERNAL[params] + if params_ == nil { + return false + } + + return _mldsa.unpack_pk(pub_key, b, params_) +} + +// public_key_set sets pub_key to src. +public_key_set :: proc(pub_key, src: ^Public_Key) { + if src == nil || internal_to_params(src.params) == .Invalid { + public_key_clear(pub_key) + return + } + + _mldsa.set_pk(pub_key, src) +} + +// public_key_set_priv sets pub_key to the public component of priv_key. +public_key_set_priv :: proc(pub_key: ^Public_Key, priv_key: ^Private_Key) { + ensure(priv_key.params != nil, "crypto/mldsa: uninitialized private key") + public_key_set(pub_key, &priv_key.pub_key) +} + +// public_key_bytes sets dst to byte-encoding of pub_key. +public_key_bytes :: proc(pub_key: ^Public_Key, dst: []byte) { + ensure(pub_key.params != nil, "crypto/mldsa: uninitialized public key") + params := internal_to_params(pub_key.params) + ensure(len(dst) == PUBLIC_KEY_SIZES[params], "crypto/mldsa: invalid destination size") + + _ = _mldsa.pack_pk(dst, pub_key) +} + +// public_key_equal returns true if and only if (⟺) the public keys are equal, +// in constant time. +@(require_results) +public_key_equal :: proc(p, q: ^Public_Key) -> bool { + if p.params != q.params { + return false + } + if p.params == nil { + return true + } + + // Comparing the pre-computed hash should be enough, but pack + // both public keys and do the comparisons. + PUBLIC_KEY_SIZE_MAX :: 2592 + + l := PUBLIC_KEY_SIZES[internal_to_params(p.params)] + p_buf_, q_buf_: [PUBLIC_KEY_SIZE_MAX]byte = ---, --- + p_buf, q_buf := p_buf_[:l], q_buf_[:l] + + _ = _mldsa.pack_pk(p_buf, p) + _ = _mldsa.pack_pk(q_buf, q) + + return crypto.compare_constant_time(p_buf, q_buf) == 1 +} + +// public_key_clear clears pub_key to the uninitialized state. +public_key_clear :: proc "contextless" (pub_key: ^Public_Key) { + _mldsa.clear_pk(pub_key) +} + +// sign writes the signature by priv_key over (ctx, msg) to sig and +// returns true if and only if (⟺) the signing succeeded. +// +// ctx is an optional domain separation tag and may be omitted (nil). +@(require_results) +sign :: proc(priv_key: ^Private_Key, ctx, msg, sig: []byte, deterministic := !crypto.HAS_RAND_BYTES) -> bool { + params := internal_to_params(priv_key.params) + ensure(params != .Invalid, "crypto/mldsa: invalid private key") + ensure(len(sig) == SIGNATURE_SIZES[params], "crypto/mldsa: invalid destination size") + + if !deterministic && !crypto.HAS_RAND_BYTES { + return false + } + if len(ctx) > MAX_CTX_SIZE { + return false + } + + rnd: [_mldsa.RNDBYTES]byte + defer crypto.zero_explicit(&rnd, size_of(rnd)) + + if !deterministic { + crypto.rand_bytes(rnd[:]) + } + + return _mldsa.dsa_sign_internal(sig, msg, ctx, rnd[:], priv_key) +} + +// verify returns true if and only if (⟺) sig is a valid signature by pub_key +// over (ctx, msg). +@(require_results) +verify :: proc(pub_key: ^Public_Key, ctx, msg, sig: []byte) -> bool { + params := internal_to_params(pub_key.params) + ensure(params != .Invalid, "crypto/mldsa: invalid public key") + + if len(sig) != SIGNATURE_SIZES[params] { + return false + } + if len(ctx) > MAX_CTX_SIZE { + return false + } + + return _mldsa.dsa_verify_internal(sig, msg, ctx, pub_key) +} + +// params returns the Parameters used by a Private_Key or Public_Key +// instance. +@(require_results) +params :: proc(k: ^$T) -> Parameters where (T == Private_Key || T == Public_Key) { + return internal_to_params(k.params) +} + +// key_size returns the key size of a Private_Key or Public_Key in bytes. +@(require_results) +key_size :: proc(k: ^$T) -> int where (T == Private_Key || T == Public_Key) { + when T == Private_Key { + return PRIVATE_KEY_SEED_SIZE + } else { + return PUBLIC_KEY_SIZES[internal_to_params(k.params)] + } +} + +// signature_size returns the key size of a signature in bytes. +@(require_results) +signature_size :: proc(k: ^$T) -> int where (T == Private_Key || T == Public_Key) { + return SIGNATURE_SIZES[internal_to_params(k.params)] +} + +@(private="file",require_results) +internal_to_params :: proc "contextless" (params: ^_mldsa.Params) -> Parameters { + switch params { + case &_mldsa.Params_44: + return .ML_DSA_44 + case &_mldsa.Params_65: + return .ML_DSA_65 + case &_mldsa.Params_87: + return .ML_DSA_87 + case: + return .Invalid + } +} diff --git a/core/crypto/mldsa/doc.odin b/core/crypto/mldsa/doc.odin new file mode 100644 index 000000000..4801c070e --- /dev/null +++ b/core/crypto/mldsa/doc.odin @@ -0,0 +1,7 @@ +/* +Module-Lattice-Based Digital Signature Algorithm. + +See: +- [[ https://nvlpubs.nist.gov/nistpubs/FIPS/NIST.FIPS.204.pdf ]] +*/ +package mldsa diff --git a/core/crypto/mlkem/api.odin b/core/crypto/mlkem/api.odin new file mode 100644 index 000000000..3e3956ba8 --- /dev/null +++ b/core/crypto/mlkem/api.odin @@ -0,0 +1,304 @@ +package mlkem + +import "core:crypto" +import "core:crypto/_mlkem" + +// Parameters are the supported ML-KEM parameter sets. +Parameters :: enum { + Invalid, + ML_KEM_512, + ML_KEM_768, + ML_KEM_1024, +} + +// DECAPSULATION_KEY_SEED_SIZE is the size of a Decapsulation key in bytes. +DECAPSULATION_KEY_SEED_SIZE :: 64 // (d, z) in NIST terms. + +// DECAPSULATION_KEY_EXPANDED_SIZES are the per-parameter sizes of the +// decapsulation key in bytes. +DECAPSULATION_KEY_EXPANDED_SIZES := [Parameters]int { + .Invalid = 0, + .ML_KEM_512 = _mlkem.DECAPSKEYBYTES_512, // 1632-bytes + .ML_KEM_768 = _mlkem.DECAPSKEYBYTES_768, // 2400-bytes + .ML_KEM_1024 = _mlkem.DECAPSKEYBYTES_1024, // 3168-bytes +} + +// ENCAPSULATION_KEY_SIZES are the per-parameter sizes of the encapsulation +// key in bytes. +ENCAPSULATION_KEY_SIZES := [Parameters]int { + .Invalid = 0, + .ML_KEM_512 = _mlkem.ENCAPSKEYBYTES_512, // 800-bytes + .ML_KEM_768 = _mlkem.ENCAPSKEYBYTES_768, // 1184-bytes + .ML_KEM_1024 = _mlkem.ENCAPSKEYBYTES_1024, // 1568-bytes +} + +// CIPHERTEXT_SIZES are the per-parameter set sizes of the ciphertext +// in bytes. +CIPHERTEXT_SIZES := [Parameters]int { + .Invalid = 0, + .ML_KEM_512 = _mlkem.CIPHERTEXTBYTES_512, // 768-bytes + .ML_KEM_768 = _mlkem.CIPHERTEXTBYTES_768, // 1088-bytes + .ML_KEM_1024 = _mlkem.CIPHERTEXTBYTES_1024, // 1568-bytes +} + +// SHARED_SECRET_SIZE is the size of the final shared secret in bytes. +SHARED_SECRET_SIZE :: 32 + +// Decapsulation_Key is a ML-KEM decapsulation (aka "private") key. +// This implementation opts to include the encapsulation (aka "public") +// key as well for cases where the decapsulation key is reused (eg: HPKE +// with X-Wing). +Decapsulation_Key :: _mlkem.Decapsulation_Key + +// Encapsulation_Key is a ML-KEM encapsulation (aka "public") key. +Encapsulation_Key :: _mlkem.Encapsulation_Key + +// decapsulation_key_generate uses the system entropy source to generate +// a decapsulation key. This will only fail if and only if (⟺) the system +// entropy source is missing or broken. +@(require_results) +decapsulation_key_generate :: proc(dk: ^Decapsulation_Key, params: Parameters) -> bool { + decapsulation_key_clear(dk) + + if !crypto.HAS_RAND_BYTES { + return false + } + + k := params_to_k(params) + if k == 0 { + panic("crypto/mlkem: invalid parameter set") + } + + seed: [DECAPSULATION_KEY_SEED_SIZE]byte = --- + defer crypto.zero_explicit(&seed, size_of(seed)) + + crypto.rand_bytes(seed[:]) + _mlkem.kem_keygen_internal(dk, seed[:], k) + + return true +} + +// decapsulation_key_set_bytes decodes a byte-encoded decapsulation key +// in (d, z) "seed" format, and returns true if and only if (⟺) the +// operation was successful. +@(require_results) +decapsulation_key_set_bytes :: proc(dk: ^Decapsulation_Key, params: Parameters, seed: []byte) -> bool { + k := params_to_k(params) + if k == 0 { + return false + } + if len(seed) != DECAPSULATION_KEY_SEED_SIZE { + return false + } + + _mlkem.kem_keygen_internal(dk, seed, k) + + return true +} + +// decapsulation_key_bytes sets dst to byte-encoding of dk in the (d, z) +// "seed" format. +decapsulation_key_bytes :: proc(dk: ^Decapsulation_Key, dst: []byte) { + ensure(dk.pke_dk.k != 0, "crypto/mlkem: uninitialized Decapsulation_Key") + ensure(len(dst) == DECAPSULATION_KEY_SEED_SIZE, "crypto/mlkem: invalid destination size") + + copy(dst, dk.seed[:]) +} + +// decapsulation_key_expanded_bytes sets dst to the byte-encoding of dk. +// in the expanded FIPS 203 format. This primarily exists for export +// purposes. +decapsulation_key_expanded_bytes :: proc(dk: ^Decapsulation_Key, dst: []byte) { + dk_len: int + switch dk.pke_dk.k { + case _mlkem.K_512: + dk_len = DECAPSULATION_KEY_EXPANDED_SIZES[.ML_KEM_512] + case _mlkem.K_768: + dk_len = DECAPSULATION_KEY_EXPANDED_SIZES[.ML_KEM_768] + case _mlkem.K_1024: + dk_len = DECAPSULATION_KEY_EXPANDED_SIZES[.ML_KEM_1024] + case: + panic("crypto/mlkem: uninitialized Decapsulation_Key") + } + ensure(len(dst) == dk_len, "crypto/mlkem: invalid destination size") + + _mlkem.decapsulation_key_expanded_bytes(dk, dst) +} + +// decapsulation_key_encaps_bytes sets dst to the byte-encoding of the +// encasulation key corresponding to dk. +decapsulation_key_encaps_bytes :: proc(dk: ^Decapsulation_Key, dst: []byte) { + encapsulation_key_bytes(&dk.ek, dst) +} + +// decapsulation_key_clear clears dk to the uninitialized state. +decapsulation_key_clear :: proc(dk: ^Decapsulation_Key) { + crypto.zero_explicit(dk, size_of(Decapsulation_Key)) +} + +// encapsulation_key_set_bytes decodes a byte-encoded encapsulation key, +// and returns true if and only if (⟺) the operation was successful. +@(require_results) +encapsulation_key_set_bytes :: proc(ek: ^Encapsulation_Key, params: Parameters, b: []byte) -> bool { + k := params_to_k(params) + if k == 0 { + return false + } + if len(b) != ENCAPSULATION_KEY_SIZES[params] { + return false + } + + return _mlkem.encapsulation_key_set_bytes(ek, k, b) +} + +// encapsulation_key_set_decaps sets ek to the encapsulation key corresponding +// to dk. +encapsulation_key_set_decaps :: proc(ek: ^Encapsulation_Key, dk: ^Decapsulation_Key) { + ensure(dk.pke_dk.k != 0, "crypto/mlkem: uninitialized Decapsulation_Key") + _mlkem.encapsulation_key_set_decaps(ek, dk) +} + +// encapsulation_key_encaps_bytes sets dst to the byte-encoding of ek. +encapsulation_key_bytes :: proc(ek: ^Encapsulation_Key, dst: []byte) { + ensure(ek.pke_ek.k != 0, "crypto/mlkem: uninitialized Encapsulation_Key") + + k_len: int + switch ek.pke_ek.k { + case _mlkem.K_512: + k_len = ENCAPSULATION_KEY_SIZES[.ML_KEM_512] + case _mlkem.K_768: + k_len = ENCAPSULATION_KEY_SIZES[.ML_KEM_768] + case _mlkem.K_1024: + k_len = ENCAPSULATION_KEY_SIZES[.ML_KEM_1024] + case: + panic("crypto/mlkem: invalid destination size") + } + + copy(dst, ek.raw_bytes[:k_len]) +} + +// encapsulation_key_clear clears ek to the uninitialized state. +encapsulation_key_clear :: proc(ek: ^Encapsulation_Key) { + crypto.zero_explicit(ek, size_of(Encapsulation_Key)) +} + +// encaps_raw_ek_bytes uses the byte encoded encapsulation key to generate +// a shared secret and an associated ciphertext. This routine will fail +// if the system entropy source is unavailable, or of the encapsulation key +// is invalid. +@(require_results) +encaps_ek_raw_bytes :: proc(params: Parameters, raw_ek, shared_secret, ciphertext: []byte) -> bool { + ek: Encapsulation_Key = --- + if !encapsulation_key_set_bytes(&ek, params, raw_ek) { + return false + } + defer encapsulation_key_clear(&ek) + + return encaps_ek(&ek, shared_secret, ciphertext) +} + +// encaps_ek uses the encapsulation key to generate a shared secret and an +// associated ciphertext. This routine will fail if the system entropy source +// is unavailable. +@(require_results) +encaps_ek :: proc(ek: ^Encapsulation_Key, shared_secret, ciphertext: []byte) -> bool { + ensure(len(shared_secret) == SHARED_SECRET_SIZE, "crypto/mlkem: invalid shared_seret size") + + if !crypto.HAS_RAND_BYTES { + return false + } + + m: [_mlkem.SYMBYTES]byte = --- + defer crypto.zero_explicit(&m, size_of(m)) + + crypto.rand_bytes(m[:]) + _mlkem.kem_encaps_internal(shared_secret, ciphertext, ek, m[:]) + + return true +} + +encaps :: proc { + encaps_ek, + encaps_ek_raw_bytes, +} + +// decaps uses the decapsulation key to generate a shared secret from a +// ciphertext. Due to ML-KEM's implicit rejection mechanism, this function +// will only return false if and only if (⟺) the lengths of the inputs +// are invalid or the decapsulation key is uninitialized. +// +// This routine returning true does not guarantee that the shared secret +// matches that generated by the peer. +@(require_results) +decaps :: proc(dk: ^Decapsulation_Key, ciphertext, shared_secret: []byte) -> bool { + ensure(len(shared_secret) == SHARED_SECRET_SIZE, "crypto/mlkem: invalid shared_seret size") + + ct_len: int + switch dk.pke_dk.k { + case _mlkem.K_512: + ct_len = CIPHERTEXT_SIZES[.ML_KEM_512] + case _mlkem.K_768: + ct_len = CIPHERTEXT_SIZES[.ML_KEM_768] + case _mlkem.K_1024: + ct_len = CIPHERTEXT_SIZES[.ML_KEM_1024] + case: + return false + } + if len(ciphertext) != ct_len { + return false + } + + _mlkem.kem_decaps_internal(shared_secret, dk, ciphertext) + + return true +} + +// params returns the Parameters used by a Decapsulation_Key or +// Encapsulation_Key instance. +@(require_results) +params :: proc(k: ^$T) -> Parameters where (T == Encapsulation_Key || T == Decapsulation_Key) { + when T == Encapsulation_Key { + return k_to_params(k.pke_ek.k) + } else { + return k_to_params(k.pke_dk.k) + } +} + +// key_size returns the key size of a Decapsulation_Key or Encapsulation_Key +// in bytes. +@(require_results) +key_size :: proc(k: ^$T) -> int where (T == Encapsulation_Key || T == Decapsulation_Key) { + when T == Encapsulation_Key { + return ENCAPSULATION_KEY_SIZES[k.pke_ek.k] + } else { + return DECAPSULATION_KEY_SEED_SIZE + } +} + +@(private="file") +params_to_k :: #force_inline proc "contextless" (params: Parameters) -> int { + #partial switch params { + case .ML_KEM_512: + return _mlkem.K_512 + case .ML_KEM_768: + return _mlkem.K_768 + case .ML_KEM_1024: + return _mlkem.K_1024 + } + + return 0 +} + +@(private="file") +k_to_params :: #force_inline proc "contextless" (k: int) -> Parameters { + switch k { + case _mlkem.K_512: + return .ML_KEM_512 + case _mlkem.K_768: + return .ML_KEM_768 + case _mlkem.K_1024: + return .ML_KEM_1024 + } + + return .Invalid +} diff --git a/core/crypto/mlkem/doc.odin b/core/crypto/mlkem/doc.odin new file mode 100644 index 000000000..ec32def25 --- /dev/null +++ b/core/crypto/mlkem/doc.odin @@ -0,0 +1,7 @@ +/* +ML-KEM Module-Lattice-Based Key-Encapsulation Mechanism. + +See: +- [[ https://nvlpubs.nist.gov/nistpubs/FIPS/NIST.FIPS.203.pdf ]] +*/ +package mlkem diff --git a/core/crypto/noise/api.odin b/core/crypto/noise/api.odin new file mode 100644 index 000000000..bae58a399 --- /dev/null +++ b/core/crypto/noise/api.odin @@ -0,0 +1,590 @@ +package noise + +import "base:runtime" +import "core:crypto/aead" +import "core:crypto/ecdh" +import "core:crypto/hash" +import "core:strings" + +// MAX_PACKET_SIZE is the maximum Noise message size, including TAG_SIZE +// if relevant (`seal_message`, `open_message`). +MAX_PACKET_SIZE :: 65535 + +// PSK_SIZE is the size of an optional handshake pre-shared symmetric key. +PSK_SIZE :: 32 +// TAG_SIZE is the size of the AEAD authentication tag. +TAG_SIZE :: 16 +// MAX_STEP_MSG_SIZE is the maximum per-handshake step message size, +// excluding the optional payload. +// +// `e` is DH_LEN, `s` is either DH_LEN or DH_LEN + TAG_SIZE, and there +// is a maximum of one per each message, and a possible mandatory tag. +MAX_STEP_MSG_SIZE :: (MAX_DH_SIZE*2)+TAG_SIZE+TAG_SIZE + +// Status is the status of Noise protocol operation. +Status :: enum { + Ok, + + // States + Handshake_Pending, + Handshake_Complete, + Handshake_Split, + Handshake_Failed, + + // Errors + Invalid_Protocol_String, + Invalid_Pre_Shared_Key, + Invalid_DH_Key, + No_Self_Identity, + No_Peer_Identity, + Unexpected_Peer_Identity, + Unexpected_Pre_Shared_Key, + + DH_Failure, + Invalid_Handshake_Message, + + Decryption_Failure, + IV_Exhausted, + Invalid_Cipher_State, + Invalid_Destination_Buffer, + Invalid_Payload_Message, + Max_Packet_Size, + + Out_Of_Memory, +} + +// Handshake_State is the per-handshake state. +Handshake_State :: struct { + s: ecdh.Private_Key, + e: ecdh.Private_Key, + rs: ecdh.Public_Key, + re: ecdh.Public_Key, + psk: [PSK_SIZE]byte, + + symmetric_state: Symmetric_State, + message_pattern: ^Message_Pattern, + current_message: int, + + status: Status, + + initiator: bool, + pre_set_e: bool, +} + +// Cipher_States are the keyed AEAD instances and associated state, +// derived from a successful handshake. +Cipher_States :: struct { + c1_i_to_r: Cipher_State, + c2_r_to_i: Cipher_State, + + initiator: bool, +} + +// handshake_init initializes a Handshake_State with the provided parameters. +// The relevant values are copied into the Handshake_State instance, and +// can be discarded/sanitized right after handshake_init returns (eg: psk). +// +// Note: While this implementation supports setting `e`, this is primarily +// intended for testing, or cases where the runtime cryptographic entropy +// source is unavailable. Use of this functionality is STRONGLY +// discouraged. +@(require_results) +handshake_init :: proc( + self: ^Handshake_State, + initiator: bool, + prologue: []byte, + s: ^ecdh.Private_Key, // Our static key + rs: ^ecdh.Public_Key, // Peer static key + protocol_name: string, + psk: []byte = nil, + _e: ^ecdh.Private_Key = nil, // Our ephemeral key (for testing/RNG-less systems) +) -> Status { + return handshakestate_initialize( + self, + initiator, + prologue, + s, + _e, + rs, + nil, + protocol_name, + psk, + ) +} + +// handshake_initiator_step takes an input_message received from the responder +// if any and an optional payload to be sent to the responder, and performs +// one step of the Noise handshake process, returning the message to be sent +// to the responder if any, the payload received from the responder if any, +// and the status of the handshake. +// +// The output message MUST be sent to the responder even if the status code +// returned is .Handshake_Complete. +// +// If the dst parameter is provided, the message and payload will be written +// to dst, otherwise new buffers will be allocated. +@(require_results) +handshake_initiator_step :: proc( + self: ^Handshake_State, + input_message: []byte, + payload: []byte = nil, + dst: []byte = nil, + allocator := context.allocator, +) -> ([]byte, []byte, Status) { + output_message: []byte + payload_buffer: []byte + status: Status + + dst := dst + if input_message == nil { + output_message, status = handshakestate_write_message(self, payload, dst, allocator) + } else { + payload_buffer, status = handshakestate_read_message(self, input_message, dst, allocator) + if status == .Handshake_Pending { + if dst != nil { + dst = dst[len(payload_buffer):] + } + output_message, status = handshakestate_write_message(self, payload, dst, allocator) + } + } + + return output_message, payload_buffer, status +} + +// handshake_responder_step takes a input_message received from the initiator, +// and and an optional payload to be sent to the initiator, and performs +// one step of the Noise handshake process, returning the message to be sent +// to the initiator if any, the payload received from the initiator if any, +// and the status of the handshake. +// +// The output message MUST be sent to the initiator even if the status code +// returned is .Handshake_Complete. +// +// If the dst parameter is provided, the message and payload will be written +// to dst, otherwise new buffers will be allocated. +@(require_results) +handshake_responder_step :: proc( + self: ^Handshake_State, + input_message: []byte, + payload: []byte = nil, + dst: []byte = nil, + allocator := context.allocator, +) -> ([]byte, []byte, Status) { + output_message: []byte + + if input_message == nil { + return nil, nil, .Invalid_Handshake_Message + } + + dst := dst + payload_buffer, status := handshakestate_read_message(self, input_message, dst, allocator) + if status == .Handshake_Pending { + if dst != nil { + dst = dst[len(payload_buffer):] + } + output_message, status = handshakestate_write_message(self, payload, dst, allocator) + } + + return output_message, payload_buffer, status +} + +// handshake_write_message calls the Noise HandshakeState's WriteMessage +// function directly. In most cases you are better off using +// handshake_initiator_step or handshake_responder_step. +// +// If the dst parameter is provided, the message and payload will be written +// to dst, otherwise new buffers will be allocated. +@(require_results) +handshake_write_message :: proc( + self: ^Handshake_State, + payload: []byte, + dst: []byte = nil, + allocator := context.allocator, +) -> ([]byte, Status) { + return handshakestate_write_message(self, payload, dst, allocator) +} + +// handshake_read_message calls the Noise HandshakeState's ReadMessage +// function directly. In most cases you are better off using +// handshake_initiator_step or handshake_responder_step. +// +// If the dst parameter is provided, the message and payload will be written +// to dst, otherwise new buffers will be allocated. +@(require_results) +handshake_read_message :: proc( + self: ^Handshake_State, + message: []byte, + dst: []byte = nil, + allocator := context.allocator, +) -> ([]byte, Status) { + return handshakestate_read_message(self, message, dst, allocator) +} + +// handshake_split initializes a Cipher_States instance from a completed +// handshake. This can be called once and only once per Handshake_State +// instance. +@(require_results) +handshake_split :: proc(self: ^Handshake_State, cipher_states: ^Cipher_States) -> Status { + if self.status != .Handshake_Complete { + return self.status + } + + symmetricstate_split(&self.symmetric_state, cipher_states) + if self.message_pattern.is_one_way { + cipherstate_reset(&cipher_states.c2_r_to_i) + cipher_states.c2_r_to_i.is_invalid = true + } + cipher_states.initiator = self.initiator + self.status = .Handshake_Split + + return .Ok +} + +// handshake_peer_identity returns the peer's static DH key used by +// a completed handshake. +// +// This returns a pointer to the Handshake_State's copy of the peer's +// public key, that will get wiped by handshake_reset. If the key is +// needed after a call to handshake_reset, it must be copied. +@(require_results) +handshake_peer_identity :: proc(self: ^Handshake_State) -> (^ecdh.Public_Key, Status) { + #partial switch self.status { + case .Handshake_Complete, .Handshake_Split: + case: + return nil, self.status + } + + if ecdh.curve(&self.rs) == .Invalid { + return nil, .No_Peer_Identity + } + + return &self.rs, .Ok +} + +// handshake_hash returns the handshake transcript hash of a completed +// handshake, for the purposes of channel binding. See 11.2 of the +// specification for details on usage. +// +// This returns a slice to an internal buffer that will get wiped by +// handshake_reset. If the hash is needed after a call to handshake_reset, +// the slice must be copied. +@(require_results) +handshake_hash :: proc(self: ^Handshake_State) -> ([]byte, Status) { + #partial switch self.status { + case .Handshake_Complete, .Handshake_Split: + case: + return nil, self.status + } + + return symmetricstate_get_handshake_hash(&self.symmetric_state), .Ok +} + +// handshake_reset sanitizes the Handshake_State. It is both safe and +// recommended to call this as soon as practical (after any calls to +// handshake_peer_identity, handshake_hash, and handshake_split are +// complete). +handshake_reset :: proc(self: ^Handshake_State) { + handshakestate_reset(self) +} + +// seal_message encrypts the provided data, authenticates the aad and +// ciphertext, and returns the resulting ciphertext. The ciphertext +// will ALWAYS be `len(plaintext) + TAG_SIZE` bytes in length. +// +// If the dst parameter is provided, the ciphertext will be written +// to dst, otherwise a new buffer will be allocated. +@(require_results) +seal_message :: proc(self: ^Cipher_States, aad, plaintext: []byte, dst: []byte = nil, allocator := context.allocator) -> ([]byte, Status) { + data_len := len(plaintext) + + dst := dst + did_alloc: bool + switch { + case dst == nil: + err: runtime.Allocator_Error + dst, err = make([]byte, data_len + TAG_SIZE, allocator) + if err != nil { + return nil, .Out_Of_Memory + } + did_alloc = true + case: + if len(dst) != data_len + TAG_SIZE { + return nil, .Invalid_Destination_Buffer + } + } + + status: Status + switch self.initiator { + case true: + dst, status = cipherstate_encrypt_with_ad(&self.c1_i_to_r, aad, plaintext, dst) + case false: + dst, status = cipherstate_encrypt_with_ad(&self.c2_r_to_i, aad, plaintext, dst) + } + if status != .Ok && did_alloc { + delete(dst, allocator) + dst = nil + } + + return dst, status +} + +// open_message authenticates the aad and ciphertext, decrypts the +// ciphertext and returns the resulting plaintext. The plaintext will +// ALWAYS be `len(ciphertext) - TAG_SIZE` bytes in length. +// +// If the dst parameter is provided, the plaintext will be written to +// dst, otherwise a new buffer will be allocated. +@(require_results) +open_message :: proc(self: ^Cipher_States, aad, ciphertext: []byte, dst: []byte = nil, allocator := context.allocator) -> ([]byte, Status) { + if len(ciphertext) < TAG_SIZE { + return nil, .Invalid_Payload_Message + } + + data_len := len(ciphertext) - TAG_SIZE + + dst := dst + did_alloc: bool + switch { + case dst == nil: + if data_len > 0 { + err: runtime.Allocator_Error + dst, err = make([]byte, data_len, allocator) + if err != nil { + return nil, .Out_Of_Memory + } + did_alloc = true + } + case: + if len(dst) != data_len { + return nil, .Invalid_Destination_Buffer + } + } + + status: Status + switch self.initiator { + case true: + dst, status = cipherstate_decrypt_with_ad(&self.c2_r_to_i, aad, ciphertext, dst) + case false: + dst, status = cipherstate_decrypt_with_ad(&self.c1_i_to_r, aad, ciphertext, dst) + } + if status != .Ok && did_alloc { + delete(dst, allocator) + dst = nil + } + + return dst, status +} + +// cipherstates_rekey updates the selected AEAD key, using a one way function. +// See 11.3 of the specification for examples of usage. +// +// Note: If one side updates the seal_key, the other side must update +// the non-seal_key and vice versa. +@(require_results) +cipherstates_rekey :: proc(self: ^Cipher_States, seal_key: bool) -> Status { + cs := cipherstates_cs(self, seal_key) + if cs.is_invalid { + return .Invalid_Cipher_State + } + if !cipherstate_has_key(cs) { + return .Handshake_Pending + } + + cipherstate_rekey(cs) + + return .Ok +} + +// cipherstates_set_n sets the interal counter used to generate the AEAD +// IV to an explicit value. This can be used to deal with out-of-order +// transport messages. See 11.4 of the specification. +// +// WARNING: Reusing n across different aad/messages with the same Cipher_States +// will result in catastrophic loss of security. +@(require_results) +cipherstates_set_n :: proc(self: ^Cipher_States, seal_key: bool, n: u64) -> Status { + cs := cipherstates_cs(self, seal_key) + if cs.is_invalid { + return .Invalid_Cipher_State + } + if !cipherstate_has_key(cs) { + return .Handshake_Pending + } + + cs.n = n + + return .Ok +} + +// cipherstates_n returns the interal counter used to generate the AEAD +// IV. This can be used to deal with out-of-order transport messages. +// See 11.4 of the specification. +// +// WARNING: Reusing n across different aad/messages with the same Cipher_States +// will result in catastrophic loss of security. +@(require_results) +cipherstates_n :: proc(self: ^Cipher_States, seal_key: bool, n: u64) -> (u64, Status) { + cs := cipherstates_cs(self, seal_key) + if cs.is_invalid { + return 0, .Invalid_Cipher_State + } + if !cipherstate_has_key(cs) { + return 0, .Handshake_Pending + } + + return cs.n, .Ok +} + +// cipherstates_reset sanitizes the Cipher_States. +cipherstates_reset :: proc(self: ^Cipher_States) { + self.initiator = false + cipherstate_reset(&self.c1_i_to_r) + cipherstate_reset(&self.c2_r_to_i) +} + +@(private = "file") +cipherstates_cs :: proc(self: ^Cipher_States, seal_key: bool) -> ^Cipher_State { + switch self.initiator { + case true: + switch seal_key { + case true: + return &self.c1_i_to_r + case false: + return &self.c2_r_to_i + } + case false: + switch seal_key { + case true: + return &self.c2_r_to_i + case false: + return &self.c1_i_to_r + } + } + unreachable() +} + +// split_protocol_string splits a protocol string into individual components. +@(require_results) +split_protocol_string :: proc(protocol_name: string) -> (Handshake_Pattern, ecdh.Curve, aead.Algorithm, hash.Algorithm, Status) { + str := protocol_name + + if len(str) > 255 { + return .Invalid, .Invalid, .Invalid, .Invalid, .Invalid_Protocol_String + } + + s, ok := strings.split_by_byte_iterator(&str, '_') + if !ok || s != "Noise" { + return .Invalid, .Invalid, .Invalid, .Invalid, .Invalid_Protocol_String + } + + if s, ok = strings.split_by_byte_iterator(&str, '_'); !ok { + return .Invalid, .Invalid, .Invalid, .Invalid, .Invalid_Protocol_String + } + + pattern: Handshake_Pattern + switch s { + case "N" : pattern = .N + case "K" : pattern = .K + case "X" : pattern = .X + case "XX": pattern = .XX + case "NK": pattern = .NK + case "NN": pattern = .NN + case "KN": pattern = .KN + case "KK": pattern = .KK + case "NX": pattern = .NX + case "KX": pattern = .KX + case "XN": pattern = .XN + case "IN": pattern = .IN + case "XK": pattern = .XK + case "IK": pattern = .IK + case "IX": pattern = .IX + case "NK1": pattern = .NK1 + case "NX1": pattern = .NX1 + case "X1N": pattern = .X1N + case "X1K": pattern = .X1K + case "XK1": pattern = .XK1 + case "X1K1": pattern = .X1K1 + case "X1X": pattern = .X1X + case "XX1": pattern = .XX1 + case "X1X1": pattern = .X1X1 + case "K1N": pattern = .K1N + case "K1K": pattern = .K1K + case "KK1": pattern = .KK1 + case "K1K1": pattern = .K1K1 + case "K1X": pattern = .K1X + case "KX1": pattern = .KX1 + case "K1X1": pattern = .K1X1 + case "I1N": pattern = .I1N + case "I1K": pattern = .I1K + case "IK1": pattern = .IK1 + case "I1K1": pattern = .I1K1 + case "I1X": pattern = .I1X + case "IX1": pattern = .IX1 + case "I1X1": pattern = .I1X1 + case "Npsk0": pattern = .Npsk0 + case "Kpsk0": pattern = .Kpsk0 + case "Xpsk1": pattern = .Xpsk1 + case "NNpsk0": pattern = .NNpsk0 + case "NNpsk2": pattern = .NNpsk2 + case "NKpsk0": pattern = .NKpsk0 + case "NKpsk2": pattern = .NKpsk2 + case "NXpsk2": pattern = .NXpsk2 + case "XNpsk3": pattern = .XNpsk3 + case "XKpsk3": pattern = .XKpsk3 + case "XXpsk3": pattern = .XXpsk3 + case "KNpsk0": pattern = .KNpsk0 + case "KNpsk2": pattern = .KNpsk2 + case "KKpsk0": pattern = .KKpsk0 + case "KKpsk2": pattern = .KKpsk2 + case "KXpsk2": pattern = .KXpsk2 + case "INpsk1": pattern = .INpsk1 + case "INpsk2": pattern = .INpsk2 + case "IKpsk1": pattern = .IKpsk1 + case "IKpsk2": pattern = .IKpsk2 + case "IXpsk2": pattern = .IXpsk2 + case: pattern = .Invalid + } + + if s, ok = strings.split_by_byte_iterator(&str, '_'); !ok { + return .Invalid, .Invalid, .Invalid, .Invalid, .Invalid_Protocol_String + } + dh: ecdh.Curve + switch s { + case "25519": dh = .X25519 + case "448": dh = .X448 + case: dh = .Invalid + } + + if s, ok = strings.split_by_byte_iterator(&str, '_'); !ok { + return .Invalid, .Invalid, .Invalid, .Invalid, .Invalid_Protocol_String + } + cipher: aead.Algorithm + switch s { + case "AESGCM": cipher = .AES_GCM_256 + case "ChaChaPoly": cipher = .CHACHA20POLY1305 + case: cipher = .Invalid + } + + if s, ok = strings.split_by_byte_iterator(&str, '_'); !ok { + return .Invalid, .Invalid, .Invalid, .Invalid, .Invalid_Protocol_String + } + hash: hash.Algorithm + switch s { + case "SHA512": hash = .SHA512 + case "SHA256": hash = .SHA256 + case "BLAKE2s": hash = .BLAKE2S + case "BLAKE2b": hash = .BLAKE2B + case: hash = .Invalid + } + + status: Status + + if len(str) != 0 { + status = .Invalid_Protocol_String + } + if pattern == .Invalid || dh == .Invalid || cipher == .Invalid || hash == .Invalid { + status = .Invalid_Protocol_String + } + + return pattern, dh, cipher, hash, status +} diff --git a/core/crypto/noise/doc.odin b/core/crypto/noise/doc.odin new file mode 100644 index 000000000..734a1dc6e --- /dev/null +++ b/core/crypto/noise/doc.odin @@ -0,0 +1,35 @@ +/* +An implementation of the Noise Protocol Framework (Revision 34). + +The `fallback` modifier and multi-PSK patterns are not supported +for the sake of simplicity. + +See: +- [[ https://noiseprotocol.org/ ]] +*/ +package noise + +// In general, to complete a noise handshake you must: +// +// - If you are initiating the connection, call `handshake_initiator_step` +// passing `nil` as the `input_message` parameter. +// +// - Send the resulting `[]byte` to the responder (generally a server) via +// the method of your choice. This MUST be done even if the status code +// returned is `.Handshake_Complete`. +// +// - If the status code returned by `handshake_initiator_step` was +// `.Handshake_Complete`, the handshake completed successfully, +// and it is now possible to validate the peer identity, obtain the +// handshake transcript hash, and most usefully call `handshake_split` +// to populate the `Cipher_States` struct that will be used to +// encrypt/decrypt data. +// +// Otherwise, read the response from the responder and feed the response +// data as the `input_message` to the next `handshake_initiator_step` +// until it returns `.Handshake_Complete`. +// +// - If you are the responder, the method is much the same, except you +// must pass a valid `input_message` received from an initiator to the +// first call to `handshake_responder_step`. Repeat until the returned +// status is `.Handshake_Complete`. diff --git a/core/crypto/noise/patterns.odin b/core/crypto/noise/patterns.odin new file mode 100644 index 000000000..14e5b6498 --- /dev/null +++ b/core/crypto/noise/patterns.odin @@ -0,0 +1,1098 @@ +package noise + +import "core:slice" + +@(private) +Pre_Token :: enum { + res_s, + ini_s, +} + +@(private) +Token :: enum { + e, + s, + ee, + es, + se, + ss, + psk, +} + +@(private) +Message_Pattern :: struct { + pre_messages: []Pre_Token, + messages: [][]Token, + is_psk: bool, + is_one_way: bool, +} + +// Handshake_Pattern is the list of currently supported Noise Handshake +// Patterns. +Handshake_Pattern :: enum { + Invalid, + + // One way patterns + N, + K, + X, + + // Fundamental patterns + XX, + NK, + NN, + KN, + KK, + NX, + KX, + XN, + IN, + XK, + IK, + IX, + + // Deferred patterns + NK1, + NX1, + X1N, + X1K, + XK1, + X1K1, + X1X, + XX1, + X1X1, + K1N, + K1K, + KK1, + K1K1, + K1X, + KX1, + K1X1, + I1N, + I1K, + IK1, + I1K1, + I1X, + IX1, + I1X1, + + // Recommended PSK patterns + Npsk0, + Kpsk0, + Xpsk1, + NNpsk0, + NNpsk2, + NKpsk0, + NKpsk2, + NXpsk2, + XNpsk3, + XKpsk3, + XXpsk3, + KNpsk0, + KNpsk2, + KKpsk0, + KKpsk2, + KXpsk2, + INpsk1, + INpsk2, + IKpsk1, + IKpsk2, + IXpsk2, +} + +@(require_results) +pattern_requires_initiator_s :: proc(pattern: Handshake_Pattern) -> (pre: bool, hs: bool) { + p := HANDSHAKE_PATTERNS[pattern] + if slice.contains(p.pre_messages, Pre_Token.ini_s) { + pre = true + } + for msg, i in p.messages { + if i & 1 != 0 { + continue + } + if slice.contains(msg, Token.s) { + hs = true + break + } + } + return pre, hs +} + +@(require_results) +pattern_requires_responder_s :: proc(pattern: Handshake_Pattern) -> (pre: bool, hs: bool) { + p := HANDSHAKE_PATTERNS[pattern] + if slice.contains(p.pre_messages, Pre_Token.res_s) { + pre = true + } + for msg, i in p.messages { + if i & 1 == 0 { + continue + } + if slice.contains(msg, Token.s) { + hs = true + break + } + } + return pre, hs +} + +@(require_results) +pattern_is_psk :: proc(pattern: Handshake_Pattern) -> bool { + return HANDSHAKE_PATTERNS[pattern].is_psk +} + +@(require_results) +pattern_is_one_way :: proc(pattern: Handshake_Pattern) -> bool { + return HANDSHAKE_PATTERNS[pattern].is_one_way +} + +@(require_results) +pattern_num_messages :: proc(pattern: Handshake_Pattern) -> int { + return len(HANDSHAKE_PATTERNS[pattern].messages) +} + +@(private) +HANDSHAKE_PATTERNS := [Handshake_Pattern]^Message_Pattern { + .Invalid = nil, + .N = &PATTERN_N, + .K = &PATTERN_K, + .X = &PATTERN_X, + + .XX = &PATTERN_XX, + .NK = &PATTERN_NK, + .NN = &PATTERN_NN, + .KN = &PATTERN_KN, + .KK = &PATTERN_KK, + .NX = &PATTERN_NX, + .KX = &PATTERN_KX, + .XN = &PATTERN_XN, + .IN = &PATTERN_IN, + .XK = &PATTERN_XK, + .IK = &PATTERN_IK, + .IX = &PATTERN_IX, + + .NK1 = &PATTERN_NK1, + .NX1 = &PATTERN_NX1, + .X1N = &PATTERN_X1N, + .X1K = &PATTERN_X1K, + .XK1 = &PATTERN_XK1, + .X1K1 = &PATTERN_X1K1, + .X1X = &PATTERN_X1X, + .XX1 = &PATTERN_XX1, + .X1X1 = &PATTERN_X1X1, + .K1N = &PATTERN_K1N, + .K1K = & PATTERN_K1K, + .KK1 = &PATTERN_KK1, + .K1K1 = &PATTERN_K1K1, + .K1X = &PATTERN_K1X, + .KX1 = &PATTERN_KX1, + .K1X1 = &PATTERN_K1X1, + .I1N = &PATTERN_I1N, + .I1K = &PATTERN_I1K, + .IK1 = &PATTERN_IK1, + .I1K1 = &PATTERN_I1K1, + .I1X = &PATTERN_I1X, + .IX1 = &PATTERN_IX1, + .I1X1 = &PATTERN_I1X1, + + .Npsk0 = &PATTERN_Npsk0, + .Kpsk0 = &PATTERN_Kpsk0, + .Xpsk1 = &PATTERN_Xpsk1, + .NNpsk0 = &PATTERN_NNpsk0, + .NNpsk2 = &PATTERN_NNpsk2, + .NKpsk0 = &PATTERN_NKpsk0, + .NKpsk2 = &PATTERN_NKpsk2, + .NXpsk2 = &PATTERN_NXpsk2, + .XNpsk3 = &PATTERN_XNpsk3, + .XKpsk3 = &PATTERN_XKpsk3, + .XXpsk3 = &PATTERN_XXpsk3, + .KNpsk0 = &PATTERN_KNpsk0, + .KNpsk2 = &PATTERN_KNpsk2, + .KKpsk0 = &PATTERN_KKpsk0, + .KKpsk2 = &PATTERN_KKpsk2, + .KXpsk2 = &PATTERN_KXpsk2, + .INpsk1 = &PATTERN_INpsk1, + .INpsk2 = &PATTERN_INpsk2, + .IKpsk1 = &PATTERN_IKpsk1, + .IKpsk2 = &PATTERN_IKpsk2, + .IXpsk2 = &PATTERN_IXpsk2, +} + +// ------------- ONE WAY PATTERNS --------------------------------------------------------- + +// N: +// <- s +// ... +// -> e, es +@(private,rodata) +PATTERN_N : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es}, + }, + is_one_way = true, +} + +// K: +// -> s +// <- s +// ... +// -> e, es, ss +@(private,rodata) +PATTERN_K : Message_Pattern = { + pre_messages = {.ini_s, .res_s}, + messages = { + {.e, .es, .ss}, + }, + is_one_way = true, +} + +// X: +// <- s +// ... +// -> e, es, s, ss +@(private,rodata) +PATTERN_X : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es, .s, .ss}, + }, + is_one_way = true, +} + +// ---------------------------------------------------------------------------------------- + +// ------------- FUNDAMENTAL PATTERNS ----------------------------------------------------- + +// XX: +// -> e +// <- e, ee, s, es +// -> s, se +@(private,rodata) +PATTERN_XX : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .s, .es}, + {.s, .se}, + }, +} + +// NK: +// <- s +// ... +// -> e, es +// <- e, ee +@(private,rodata) +PATTERN_NK : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es}, + {.e, .ee}, + }, +} + +// NN: +// -> e +// <- e, ee +@(private,rodata) +PATTERN_NN : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee}, + }, +} + +// KN: +// -> s +// ... +// -> e +// <- e, ee, se +@(private,rodata) +PATTERN_KN : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.e,}, + {.e, .ee, .se}, + }, +} + +// KK: +// -> s +// <- s +// ... +// -> e, es, ss +// <- e, ee, se +@(private,rodata) +PATTERN_KK : Message_Pattern = { + pre_messages = {.ini_s, .res_s}, + messages = { + {.e, .es, .ss}, + {.e, .ee, .se}, + }, +} + +// NX: +// -> e +// <- e, ee, s, es +@(private,rodata) +PATTERN_NX : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .s, .es}, + }, +} + +// KX: +// -> s +// ... +// -> e +// <- e, ee, se, s, es +@(private,rodata) +PATTERN_KX : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.e}, + {.e, .ee, .se, .s, .es}, + }, +} + +// XN: +// -> e +// <- e, ee +// -> s, se +@(private,rodata) +PATTERN_XN : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee}, + {.s, .se}, + }, +} + +// IN: +// -> e, s +// <- e, ee, se +@(private,rodata) +PATTERN_IN : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s}, + {.e, .ee, .se}, + }, +} + +// XK: +// <- s +// ... +// -> e, es +// <- e, ee +// -> s, se +@(private,rodata) +PATTERN_XK : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es}, + {.e, .ee}, + {.s, .se}, + }, +} + +// IK: +// <- s +// ... +// -> e, es, s, ss +// <- e, ee, se +@(private,rodata) +PATTERN_IK : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es, .s, .ss}, + {.e, .ee, .se}, + }, +} + +// IX: +// -> e, s +// <- e, ee, se, s, es +@(private,rodata) +PATTERN_IX : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s}, + {.e, .ee, .se, .s, .es}, + }, +} + +// ---------------------------------------------------------------------------------------- + +// ------------- DEFERRED PATTERNS -------------------------------------------------------- + +// NK1: +// <- s +// ... +// -> e +// <- e, ee, es +@(private,rodata) +PATTERN_NK1 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e}, + {.e, .ee, .es}, + }, +} + +// NX1: +// -> e +// <- e, ee, s +// -> es +@(private,rodata) +PATTERN_NX1 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .s}, + {.es}, + }, +} + +// X1N: +// -> e +// <- e, ee +// -> s +// <- se +@(private,rodata) +PATTERN_X1N : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee}, + {.s}, + {.se}, + }, +} + +// X1K: +// <- s +// ... +// -> e, es +// <- e, ee +// -> s +// <- se +@(private,rodata) +PATTERN_X1K : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es}, + {.e, .ee}, + {.s}, + {.se}, + }, +} + +// XK1: +// <- s +// ... +// -> e +// <- e, ee, es +// -> s, se +@(private,rodata) +PATTERN_XK1 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e}, + {.e, .ee, .es}, + {.s, .se}, + }, +} + +// X1K1: +// <- s +// ... +// -> e +// <- e, ee, es +// -> s +// <- se +@(private,rodata) +PATTERN_X1K1 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e}, + {.e, .ee, .es}, + {.s}, + {.se}, + }, +} + +// X1X: +// -> e +// <- e, ee, s, es +// -> s +// <- se +@(private,rodata) +PATTERN_X1X : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .s, .es}, + {.s}, + {.se}, + }, +} + +// XX1: +// -> e +// <- e, ee, s +// -> es, s, se +@(private,rodata) +PATTERN_XX1 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .s}, + {.es, .s, .se}, + }, +} + +// X1X1: +// -> e +// <- e, ee, s +// -> es, s +// <- se +@(private,rodata) +PATTERN_X1X1 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .s}, + {.es, .s}, + {.se}, + }, +} + +// K1N: +// -> s +// ... +// -> e +// <- e, ee +// -> se +@(private,rodata) +PATTERN_K1N : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.e,}, + {.e, .ee}, + {.se}, + }, +} + +// K1K: +// -> s +// <- s +// ... +// -> e, es +// <- e, ee +// -> se +@(private,rodata) +PATTERN_K1K : Message_Pattern = { + pre_messages = {.ini_s, .res_s}, + messages = { + {.e, .es}, + {.e, .ee}, + {.se}, + }, +} + +// KK1: +// -> s +// <- s +// ... +// -> e +// <- e, ee, se, es +@(private,rodata) +PATTERN_KK1 : Message_Pattern = { + pre_messages = {.ini_s, .res_s}, + messages = { + {.e}, + {.e, .ee, .se, .es}, + }, +} + +// K1K1: +// -> s +// <- s +// ... +// -> e +// <- e, ee, es +// -> se +@(private,rodata) +PATTERN_K1K1 : Message_Pattern = { + pre_messages = {.ini_s, .res_s}, + messages = { + {.e}, + {.e, .ee, .es}, + {.se}, + }, +} + +// K1X: +// -> s +// ... +// -> e +// <- e, ee, s, es +// -> se +@(private,rodata) +PATTERN_K1X : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.e}, + {.e, .ee, .s, .es}, + {.se}, + }, +} + +// KX1: +// -> s +// ... +// -> e +// <- e, ee, se, s +// -> es +@(private,rodata) +PATTERN_KX1 : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.e}, + {.e, .ee, .se, .s}, + {.es}, + }, +} + +// K1X1: +// -> s +// ... +// -> e +// <- e, ee, s +// -> se, es +@(private,rodata) +PATTERN_K1X1 : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.e}, + {.e, .ee, .s}, + {.se, .es}, + }, +} + +// I1N: +// -> e, s +// <- e, ee +// -> se +@(private,rodata) +PATTERN_I1N : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s}, + {.e, .ee}, + {.se}, + }, +} + +// I1K: +// <- s +// ... +// -> e, es, s +// <- e, ee +// -> se +@(private,rodata) +PATTERN_I1K : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es, .s}, + {.e, .ee}, + {.se}, + }, +} + +// IK1: +// <- s +// ... +// -> e, s +// <- e, ee, se, es +@(private,rodata) +PATTERN_IK1 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .s}, + {.e, .ee, .se, .es}, + }, +} + +// I1K1: +// <- s +// ... +// -> e, s +// <- e, ee, es +// -> se +@(private,rodata) +PATTERN_I1K1 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .s}, + {.e, .ee, .es}, + {.se}, + }, +} + +// I1X: +// -> e, s +// <- e, ee, s, es +// -> se +@(private,rodata) +PATTERN_I1X : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s}, + {.e, .ee, .s, .es}, + {.se}, + }, +} + +// IX1: +// -> e, s +// <- e, ee, se, s +// -> es +@(private,rodata) +PATTERN_IX1 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s}, + {.e, .ee, .se, .s}, + {.es}, + }, +} + +// I1X1: +// -> e, s +// <- e, ee, s +// -> se, es +@(private,rodata) +PATTERN_I1X1 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s}, + {.e, .ee, .s}, + {.se, .es}, + }, +} + +// ---------------------------------------------------------------------------------------- + +// ------------- PSK PATTERNS ------------------------------------------------------------- + +// Npsk0: +// <- s +// ... +// -> psk, e, es +@(private,rodata) +PATTERN_Npsk0 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.psk, .e, .es}, + }, + is_psk = true, + is_one_way = true, +} + +// K: +// -> s +// <- s +// ... +// -> psk, e, es, ss +@(private,rodata) +PATTERN_Kpsk0 : Message_Pattern = { + pre_messages = {.ini_s, .res_s}, + messages = { + {.psk, .e, .es, .ss}, + }, + is_psk = true, + is_one_way = true, +} + +// X: +// <- s +// ... +// -> e, es, s, ss, psk +@(private,rodata) +PATTERN_Xpsk1 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es, .s, .ss, .psk}, + }, + is_psk = true, + is_one_way = true, +} + +// NNpsk0: +// -> psk, e +// <- e, ee +@(private,rodata) +PATTERN_NNpsk0 : Message_Pattern = { + pre_messages = nil, + messages = { + {.psk, .e}, + {.e, .ee}, + }, + is_psk = true, +} + +// NNpsk2: +// -> e +// <- e, ee, psk +@(private,rodata) +PATTERN_NNpsk2 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .psk}, + }, + is_psk = true, +} + +// NKpsk0: +// <- s +// ... +// -> psk, e, es +// <- e, ee +@(private,rodata) +PATTERN_NKpsk0 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.psk, .e, .es}, + {.e, .ee}, + }, + is_psk = true, +} + +// NKpsk2: +// <- s +// ... +// -> e, es +// <- e, ee, psk +@(private,rodata) +PATTERN_NKpsk2 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es}, + {.e, .ee, .psk}, + }, + is_psk = true, +} + +// NXpsk2: +// -> e +// <- e, ee, s, es, psk +@(private,rodata) +PATTERN_NXpsk2 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .s, .es, .psk}, + }, + is_psk = true, +} + +// XNpsk3: +// -> e +// <- e, ee +// -> s, se, psk +@(private,rodata) +PATTERN_XNpsk3 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee}, + {.s, .se, .psk}, + }, + is_psk = true, +} + +// XKpsk3: +// <- s +// ... +// -> e, es +// <- e, ee +// -> s, se, psk +@(private,rodata) +PATTERN_XKpsk3 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es}, + {.e, .ee}, + {.s, .se, .psk}, + }, + is_psk = true, +} + +// XXpsk3: +// -> e +// <- e, ee, s, es +// -> s, se, psk +@(private,rodata) +PATTERN_XXpsk3 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e}, + {.e, .ee, .s, .es}, + {.s, .se, .psk}, + }, + is_psk = true, +} + +// KNpsk0: +// -> s +// ... +// -> psk, e +// <- e, ee, se +@(private,rodata) +PATTERN_KNpsk0 : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.psk, .e}, + {.e, .ee, .se}, + }, + is_psk = true, +} + +// KNpsk2: +// -> s +// ... +// -> e +// <- e, ee, se, psk +@(private,rodata) +PATTERN_KNpsk2 : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.e}, + {.e, .ee, .se, .psk}, + }, + is_psk = true, +} + +// KKpsk0: +// -> s +// <- s +// ... +// -> psk, e, es, ss +// <- e, ee, se +@(private,rodata) +PATTERN_KKpsk0 : Message_Pattern = { + pre_messages = {.ini_s, .res_s}, + messages = { + {.psk, .e, .es, .ss}, + {.e, .ee, .se}, + }, + is_psk = true, +} + +// KKpsk2: +// -> s +// <- s +// ... +// -> e, es, ss +// <- e, ee, se, psk +@(private,rodata) +PATTERN_KKpsk2 : Message_Pattern = { + pre_messages = {.ini_s, .res_s}, + messages = { + {.e, .es, .ss}, + {.e, .ee, .se, .psk}, + }, + is_psk = true, +} + +// KXpsk2: +// -> s +// ... +// -> e +// <- e, ee, se, s, es, psk +@(private,rodata) +PATTERN_KXpsk2 : Message_Pattern = { + pre_messages = {.ini_s}, + messages = { + {.e}, + {.e, .ee, .se, .s, .es, .psk}, + }, + is_psk = true, +} + +// INpsk1: +// -> e, s, psk +// <- e, ee, se +@(private,rodata) +PATTERN_INpsk1 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s, .psk}, + {.e, .ee, .se}, + }, + is_psk = true, +} + +// INpsk2: +// -> e, s +// <- e, ee, se, psk +@(private,rodata) +PATTERN_INpsk2 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s}, + {.e, .ee, .se, .psk}, + }, + is_psk = true, +} + +// IKpsk1: +// <- s +// ... +// -> e, es, s, ss, psk +// <- e, ee, se +@(private,rodata) +PATTERN_IKpsk1 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es, .s, .ss, .psk}, + {.e, .ee, .se}, + }, + is_psk = true, +} + +// IKpsk2: +// <- s +// ... +// -> e, es, s, ss +// <- e, ee, se, psk +@(private,rodata) +PATTERN_IKpsk2 : Message_Pattern = { + pre_messages = {.res_s}, + messages = { + {.e, .es, .s, .ss}, + {.e, .ee, .se, .psk}, + }, + is_psk = true, +} + +// IXpsk2: +// -> e, s +// <- e, ee, se, s, es, psk +@(private,rodata) +PATTERN_IXpsk2 : Message_Pattern = { + pre_messages = nil, + messages = { + {.e, .s}, + {.e, .ee, .se, .s, .es, .psk}, + }, + is_psk = true, +} diff --git a/core/crypto/noise/protocol.odin b/core/crypto/noise/protocol.odin new file mode 100644 index 000000000..7327e638b --- /dev/null +++ b/core/crypto/noise/protocol.odin @@ -0,0 +1,989 @@ +#+private +package noise + +import "base:runtime" +import "core:crypto" +import "core:crypto/aead" +import "core:crypto/ecdh" +import "core:crypto/hash" +import "core:crypto/hkdf" +import "core:encoding/endian" +import "core:slice" + +AEAD_KEY_SIZE :: 32 + +MIN_DH_SIZE :: 32 +MAX_DH_SIZE :: 56 +MAX_HASH_SIZE :: 64 + +Protocol :: struct { + handshake_pattern: Handshake_Pattern, + dh: ecdh.Curve, + cipher: aead.Algorithm, + hash: hash.Algorithm, +} + +Symmetric_State :: struct { + protocol: Protocol, + cipher_state: Cipher_State, + + _ck: [MAX_HASH_SIZE]byte, + _h: [MAX_HASH_SIZE]byte, +} + +Cipher_State :: struct { + ctx: aead.Context, + n: u64, + n_exhausted: bool, + is_invalid: bool, +} + +@(require_results) +dh_len :: proc(protocol: ^Protocol) -> int { + return ecdh.PUBLIC_KEY_SIZES[protocol.dh] +} + +@(require_results) +hash_len :: proc(protocol: ^Protocol) -> int { + return hash.DIGEST_SIZES[protocol.hash] +} + +// Generates a new Diffie-Hellman key pair. A DH key pair consists of +// public_key and private_key elements. public_key represents an encoding +// of a DH public key into a byte sequence of length DHLEN. The public_key +// encoding details are specific to each set of DH functions. +generate_keypair :: proc(protocol: ^Protocol, private_key: ^ecdh.Private_Key) { + #partial switch protocol.dh { + case .X25519, .X448: + case: panic("crypto/noise: unsupported DH curve in protocol") + } + + if !ecdh.private_key_generate(private_key, protocol.dh) { + panic("crypto/noise: entropy source unavailable") + } +} + +// Performs a Diffie-Hellman calculation between the private key in key_pair +// and the public_key and returns an output sequence of bytes of length DHLEN. +// For security, the Gap-DH problem based on this function must be unsolvable +// by any practical cryptanalytic adversary [2]. +// +// The public_key either encodes some value which is a generator in a +// large prime-order group (which value may have multiple equivalent +// encodings), or is an invalid value. Implementations must handle invalid +// public keys either by returning some output which is purely a function +// of the public key and does not depend on the private key, or by signaling +// an error to the caller. +// +// The DH function may define more specific rules for handling invalid values. +@(require_results) +_dh :: proc(our_private_key: ^ecdh.Private_Key, their_public_key: ^ecdh.Public_Key, dst: []byte) -> Status { + if ok := ecdh.ecdh(our_private_key, their_public_key, dst); !ok { + return .DH_Failure + } + return .Ok +} + +// Encrypts plaintext using the cipher key k of 32 bytes and an 8-byte +// unsigned integer nonce n which must be unique for the key k. +// Returns the ciphertext. Encryption must be done with an "AEAD" +// encryption mode with the associated data(AD) (using the terminology +// from [1]) and returns a ciphertext that is the same size as the plaintext +// plus 16 bytes for authentication data. The entire ciphertext must be +// indistinguishable from random if the key is secret (note that this is +// an additional requirement that isn't necessarily met by all AEAD schemes). +_encrypt :: proc(ctx: ^aead.Context, n: u64, ad, plaintext, dst: []byte) { + pt_len := len(plaintext) + ensure(len(dst) == pt_len + TAG_SIZE, "crypto/noise: invalid AEAD encrypt destination") + + iv: [12]byte + #partial switch aead.algorithm(ctx) { + case .AES_GCM_256: endian.unchecked_put_u64be(iv[4:], n) + case .CHACHA20POLY1305: endian.unchecked_put_u64le(iv[4:], n) + } + + ciphertext, tag := dst[:pt_len], dst[pt_len:] + aead.seal_ctx(ctx, ciphertext, tag, iv[:], ad, plaintext) +} + +// Decrypts ciphertext using a cipher key k of 32 bytes, an 8-byte unsigned +// integer nonce n, and associated data ad. Returns the plaintext, unless +// authentication fails, in which case an error is signaled to the caller. +@(require_results) +_decrypt :: proc(ctx: ^aead.Context, n: u64, ad, ciphertext, dst: []byte) -> Status { + if len(ciphertext) < TAG_SIZE { + return .Decryption_Failure + } + + iv: [12]byte + #partial switch aead.algorithm(ctx) { + case .AES_GCM_256: endian.unchecked_put_u64be(iv[4:], n) + case .CHACHA20POLY1305: endian.unchecked_put_u64le(iv[4:], n) + } + + ct_len := len(ciphertext) - TAG_SIZE + ct, tag := ciphertext[:ct_len], ciphertext[ct_len:] + if ok := aead.open_ctx(ctx, dst, iv[:], ad, ct, tag); !ok { + return .Decryption_Failure + } + + return .Ok +} + +// Hashes some arbitrary-length data with a collision-resistant cryptographic +// hash function and returns an output of HASHLEN bytes. +_hash :: proc(dst: []byte, protocol: ^Protocol, data: ..[]byte) { + ctx: hash.Context + hash.init(&ctx, protocol.hash) + + for datum in data { + hash.update(&ctx, datum) + } + + hash.final(&ctx, dst) +} + +// Takes a chaining_key byte sequence of length HASHLEN, and an +// input_key_material byte sequence with length either zero bytes, +// 32 bytes, or DHLEN bytes. Returns a pair or triple of byte sequences +// each of length HASHLEN, depending on whether num_outputs is two or three: +// - Sets temp_key = HMAC-HASH(chaining_key, input_key_material). +// - Sets output1 = HMAC-HASH(temp_key, byte(0x01)). +// - Sets output2 = HMAC-HASH(temp_key, output1 || byte(0x02)). +// - If num_outputs == 2 then returns the pair (output1, output2). +// - Sets output3 = HMAC-HASH(temp_key, output2 || byte(0x03)). +// - Returns the triple (output1, output2, output3). +// +// Note that temp_key, output1, output2, and output3 are all HASHLEN +// bytes in length. Also note that the HKDF() function is simply HKDF +// from [4] with the chaining_key as HKDF salt, and zero-length HKDF info. +@(require_results) +_hkdf :: proc(dst, chaining_key, input_key_material: []byte, protocol: ^Protocol) -> ([]byte, []byte, []byte) { + assert(len(input_key_material) == 0 || len(input_key_material) == 32 || len(input_key_material) == dh_len(protocol)) + + hkdf.extract_and_expand(protocol.hash, chaining_key, input_key_material, nil, dst) + + h_len := hash_len(protocol) + assert(len(dst) == h_len * 2 || len(dst) == h_len * 3) + + r1, r2 := dst[:h_len], dst[h_len:h_len*2] + if len(dst) == h_len * 2 { + return r1, r2, nil + } + return r1, r2, dst[h_len*2:] +} + +// Sets k = key. Sets n = 0. +cipherstate_initialize_key :: proc(self: ^Cipher_State, key: []byte, protocol: ^Protocol) { + k_len := len(key) + switch { + case k_len == 0: + // k = empty + aead.reset(&self.ctx) + self.n = 0 + case k_len < AEAD_KEY_SIZE: + panic("crypto/noise: invalid AEAD key size") + case: + aead.init(&self.ctx, protocol.cipher, key[:AEAD_KEY_SIZE]) + self.n = 0 + } +} + +// Returns true if k is non-empty, false otherwise. +@(require_results) +cipherstate_has_key :: proc(self: ^Cipher_State) -> bool { + return aead.algorithm(&self.ctx) != .Invalid +} + +// If k is non-empty returns ENCRYPT(k, n++, ad, plaintext). Otherwise +// returns plaintext. +@(require_results) +cipherstate_encrypt_with_ad :: proc(self: ^Cipher_State, ad, plaintext, dst: []byte) -> ([]byte, Status) { + if self.is_invalid { + return nil, .Invalid_Cipher_State + } + if self.n_exhausted { + return nil, .IV_Exhausted + } + + pt_len := len(plaintext) + if pt_len > MAX_PACKET_SIZE - 16 { + return nil, .Max_Packet_Size + } + + if cipherstate_has_key(self) { + if len(dst) != pt_len + TAG_SIZE { + return nil, .Invalid_Destination_Buffer + } + _encrypt(&self.ctx, self.n, ad, plaintext, dst) + self.n += 1 + if self.n == 0 { + self.n_exhausted = true + } + } else { + if len(dst) != pt_len { + return nil, .Invalid_Destination_Buffer + } + if raw_data(dst) != raw_data(plaintext) { + copy(dst, plaintext) + } + } + + return dst, .Ok +} + +// If k is non-empty returns DECRYPT(k, n++, ad, ciphertext). Otherwise +// returns ciphertext. If an authentication failure occurs in DECRYPT() +// then n is not incremented and an error is signaled to the caller. +@(require_results) +cipherstate_decrypt_with_ad :: proc(self: ^Cipher_State, ad, ciphertext, dst: []byte) -> ([]byte, Status) { + if self.is_invalid { + return nil, .Invalid_Cipher_State + } + if self.n_exhausted { + return nil, .IV_Exhausted + } + + if cipherstate_has_key(self) { + if status := _decrypt(&self.ctx, self.n, ad, ciphertext, dst); status != .Ok { + return nil, status + } + self.n += 1 + if self.n == 0 { + self.n_exhausted = true + } + } else { + if len(dst) != len(ciphertext) { + return nil, .Invalid_Destination_Buffer + } + if raw_data(dst) != raw_data(ciphertext) { + copy(dst, ciphertext) + } + } + + return dst, .Ok +} + +// Sets k = REKEY(k). +cipherstate_rekey :: proc(self: ^Cipher_State) { + if cipherstate_has_key(self) { + algorithm := aead.algorithm(&self.ctx) + + // The "sensible" way to implement this is to inlike REKEY(k), + // so we do. + // + // Returns a new 32-byte cipher key as a pseudorandom function + // of k. If this function is not specifically defined for some + // set of cipher functions, then it defaults to returning the + // first 32 bytes from `ENCRYPT(k, maxnonce, zerolen, zeros)`, + // where maxnonce equals (2^64)-1, zerolen is a zero-length + // byte sequence, and zeros is a sequence of 32 bytes filled + // with zeros. + + zeroes: [AEAD_KEY_SIZE + TAG_SIZE]byte + defer crypto.zero_explicit(&zeroes, size_of(zeroes)) + + // 1 2 3 4 5 6 7 8 + n: u64 = 0xFF_FF_FF_FF_FF_FF_FF_FF + _encrypt(&self.ctx, n, nil, zeroes[:AEAD_KEY_SIZE], zeroes[:]) + aead.init(&self.ctx, algorithm, zeroes[:AEAD_KEY_SIZE]) + } +} + +cipherstate_reset :: proc(self: ^Cipher_State) { + aead.reset(&self.ctx) + crypto.zero_explicit(self, size_of(Cipher_State)) +} + +// Takes an arbitrary-length protocol_name byte sequence (see Section 8). +// Executes the following steps: +// - If protocol_name is less than or equal to HASHLEN bytes in length, +// sets h equal to protocol_name with zero bytes appended to make +// HASHLEN bytes. +// - Otherwise sets h = HASH(protocol_name). +// - Sets ck = h. +// - Calls InitializeKey(empty). +@(require_results) +symmetricstate_initialize :: proc(ss: ^Symmetric_State, protocol_name: string) -> Status { + if status := protocol_from_string(&ss.protocol, protocol_name); status != .Ok { + return status + } + + cipherstate_initialize_key(&ss.cipher_state, nil, &ss.protocol) + + h_len := hash_len(&ss.protocol) + h := ss._h[:h_len] + if len(protocol_name) <= h_len { + copy(h, protocol_name) + } else { + _hash(h, &ss.protocol, transmute([]byte)protocol_name) + } + + copy(ss._ck[:h_len], h) + + return .Ok +} + +// Sets h = HASH(h || data). +symmetricstate_mix_hash :: proc(self: ^Symmetric_State, data: ..[]byte) { + h := self._h[:hash_len(&self.protocol)] + if len(data) == 1 { + _hash(h, &self.protocol, h, data[0]) + } else if len(data) == 2 { + _hash(h, &self.protocol, h, data[0], data[1]) + } else if len(data) == 3 { + _hash(h, &self.protocol, h, data[0], data[1], data[2]) + } else { + panic("crypto/noise: invalid MixHash inputs") + } +} + +// Executes the following steps: +// - Sets ck, temp_k = HKDF(ck, input_key_material, 2). +// - If HASHLEN is 64, then truncates temp_k to 32 bytes. +// - Calls InitializeKey(temp_k). +symmetricstate_mix_key :: proc(self: ^Symmetric_State, input_key_material: []byte) { + h_len := hash_len(&self.protocol) + + dst_len := h_len * 2 + dst: [2*MAX_HASH_SIZE]byte = --- + defer crypto.zero_explicit(&dst, dst_len) + + ck, temp_k, _ := _hkdf(dst[:dst_len], self._ck[:h_len], input_key_material, &self.protocol) + copy(self._ck[:], ck) + cipherstate_initialize_key(&self.cipher_state, temp_k, &self.protocol) +} + +// This function is used for handling pre-shared symmetric keys, as described +// in Section 9. It executes the following steps: +// - Sets ck, temp_h, temp_k = HKDF(ck, input_key_material, 3). +// - Calls MixHash(temp_h). +// - If HASHLEN is 64, then truncates temp_k to 32 bytes. +// - Calls InitializeKey(temp_k). +symmetricstate_mix_key_and_hash :: proc(self: ^Symmetric_State, input_key_material: []byte) { + h_len := hash_len(&self.protocol) + + dst_len := h_len * 3 + dst: [3*MAX_HASH_SIZE]byte = --- + defer crypto.zero_explicit(&dst, dst_len) + + ck, temp_h, temp_k := _hkdf(dst[:dst_len], self._ck[:h_len], input_key_material, &self.protocol) + copy(self._ck[:], ck) + symmetricstate_mix_hash(self, temp_h) + cipherstate_initialize_key(&self.cipher_state, temp_k, &self.protocol) +} + +// Returns h. This function should only be called at the end of a handshake, +// i.e. after the Split() function has been called. +// +// This function is used for channel binding, as described in Section 11.2 +@(require_results) +symmetricstate_get_handshake_hash :: proc(self: ^Symmetric_State) -> []byte { + return self._h[:hash_len(&self.protocol)] +} + +// Sets ciphertext = EncryptWithAd(h, plaintext), calls MixHash(ciphertext), +// and returns ciphertext. +// +// Note that if k is empty, the EncryptWithAd() call will set ciphertext +// equal to plaintext. +@(require_results) +symmetricstate_encrypt_and_hash :: proc(self: ^Symmetric_State, plaintext, dst: []byte) -> ([]byte, Status) { + ciphertext, status := cipherstate_encrypt_with_ad(&self.cipher_state, self._h[:hash_len(&self.protocol)], plaintext, dst) + if status != .Ok { + return nil, status + } + symmetricstate_mix_hash(self, ciphertext) + return ciphertext, status +} + +// Sets plaintext = DecryptWithAd(h, ciphertext), calls MixHash(ciphertext), +// and returns plaintext. +// +// Note that if k is empty, the DecryptWithAd() call will set plaintext +// equal to ciphertext. +@(require_results) +symmetricstate_decrypt_and_hash :: proc(self: ^Symmetric_State, ciphertext, dst: []byte) -> ([]byte, Status) { + h_len := hash_len(&self.protocol) + + h: [MAX_HASH_SIZE]byte = --- + copy(h[:], self._h[:h_len]) + defer crypto.zero_explicit(&h, size_of(h)) + + // We reverse the order to save having to copy the ciphertext, in + // the case that ciphertext and dst alias. + symmetricstate_mix_hash(self, ciphertext) + return cipherstate_decrypt_with_ad(&self.cipher_state, h[:h_len], ciphertext, dst) +} + +// Returns a pair of CipherState objects for encrypting transport messages. +// Executes the following steps, where zerolen is a zero-length byte sequence: +// - Sets temp_k1, temp_k2 = HKDF(ck, zerolen, 2). +// - If HASHLEN is 64, then truncates temp_k1 and temp_k2 to 32 bytes. +// - Creates two new CipherState objects c1 and c2. +// - Calls c1.InitializeKey(temp_k1) and c2.InitializeKey(temp_k2). +// - Returns the pair (c1, c2). +symmetricstate_split :: proc(self: ^Symmetric_State, cipher_states: ^Cipher_States) { + h_len := hash_len(&self.protocol) + + dst_len := h_len * 2 + dst: [2*MAX_HASH_SIZE]byte = --- + defer crypto.zero_explicit(&dst, dst_len) + + temp_k1, temp_k2, _ := _hkdf(dst[:dst_len], self._ck[:h_len], nil, &self.protocol) + cipherstate_initialize_key(&cipher_states.c1_i_to_r, temp_k1, &self.protocol) + cipherstate_initialize_key(&cipher_states.c2_r_to_i, temp_k2, &self.protocol) +} + +symmetricstate_reset :: proc(self: ^Symmetric_State) { + cipherstate_reset(&self.cipher_state) + + crypto.zero_explicit(self, size_of(Symmetric_State)) +} + +// Takes a valid handshake_pattern (see Section 7) and an initiator boolean +// specifying this party's role as either initiator or responder. +// Takes a prologue byte sequence which may be zero-length, or which may +// contain context information that both parties want to confirm is identical +// (see Section 6). +// +// Takes a set of DH key pairs (s, e) and public keys (rs, re) for +// initializing local variables, any of which may be empty. Public keys +// are only passed in if the handshake_pattern uses pre-messages +// (see Section 7). The ephemeral values (e, re) are typically left empty, +// since they are created and exchanged during the handshake; but there +// are exceptions (see Section 10). +// +// Performs the following steps: +// - Derives a protocol_name byte sequence by combining the names for +// the handshake pattern and crypto functions, as specified in Section 8. +// - Calls InitializeSymmetric(protocol_name). +// - Calls MixHash(prologue). +// - Sets the initiator, s, e, rs, and re variables to the corresponding +// arguments. +// - Calls MixHash() once for each public key listed in the pre-messages +// from handshake_pattern, with the specified public key as input +// (see Section 7 for an explanation of pre-messages). +// - If both initiator and responder have pre-messages, the initiator's +// public keys are hashed first. +// - If multiple public keys are listed in either party's pre-message, +// the public keys are hashed in the order that they are listed. +// - Sets message_pattern to the message patterns from handshake_pattern. +@(require_results) +handshakestate_initialize :: proc( + handshake_state: ^Handshake_State, + initiator: bool, + prologue: []byte, + s: ^ecdh.Private_Key, + e: ^ecdh.Private_Key, // Only set for testing. + rs: ^ecdh.Public_Key, + re: ^ecdh.Public_Key, // Only set for testing. + protocol_name: string, + psk: []byte = nil, +) -> Status { + crypto.zero_explicit(handshake_state, size_of(Handshake_State)) + + symmetric_state := &handshake_state.symmetric_state + status: Status + do_init: { + if status = symmetricstate_initialize(symmetric_state, protocol_name); status != .Ok { + break do_init + } + + curve := symmetric_state.protocol.dh + if s != nil && ecdh.curve(s) != curve { + status = .Invalid_DH_Key + break do_init + } + if e != nil && ecdh.curve(e) != curve { + status = .Invalid_DH_Key + break do_init + } + if rs != nil && ecdh.curve(rs) != curve { + status = .Invalid_DH_Key + break do_init + } + if re != nil && ecdh.curve(re) != curve { + status = .Invalid_DH_Key + break do_init + } + + // Check if we will require s later down the line. + s_pre, s_hs: bool + if initiator { + s_pre, s_hs = pattern_requires_initiator_s(symmetric_state.protocol.handshake_pattern) + } else { + s_pre, s_hs = pattern_requires_responder_s(symmetric_state.protocol.handshake_pattern) + } + if (s_pre || s_hs) && s == nil { + status = .No_Self_Identity + break do_init + } + + message_pattern := HANDSHAKE_PATTERNS[symmetric_state.protocol.handshake_pattern] + if message_pattern.pre_messages != nil { + if initiator { + if slice.contains(message_pattern.pre_messages, Pre_Token.res_s) { + if rs == nil { + status = .No_Peer_Identity + break do_init + } + } + } else { + if slice.contains(message_pattern.pre_messages, Pre_Token.ini_s) { + if rs == nil { + status = .No_Peer_Identity + break do_init + } + } + } + } else { + if rs != nil { + status = .Unexpected_Peer_Identity + break do_init + } + } + + symmetricstate_mix_hash(symmetric_state, prologue) + + // In all supported patterns, `ini_s` will always precede `res_s`. + if message_pattern.pre_messages != nil { + tmp: [MAX_DH_SIZE]byte = --- + d_len := dh_len(&symmetric_state.protocol) + dst := tmp[:d_len] + + if initiator { + if slice.contains(message_pattern.pre_messages, Pre_Token.ini_s) { + ecdh.private_key_public_bytes(s, dst) + symmetricstate_mix_hash(symmetric_state, dst) + } + if slice.contains(message_pattern.pre_messages, Pre_Token.res_s) { + ecdh.public_key_bytes(rs, dst) + symmetricstate_mix_hash(symmetric_state, dst) + } + } else { + if slice.contains(message_pattern.pre_messages, Pre_Token.ini_s) { + ecdh.public_key_bytes(rs, dst) + symmetricstate_mix_hash(symmetric_state, dst) + } + if slice.contains(message_pattern.pre_messages, Pre_Token.res_s) { + ecdh.private_key_public_bytes(s, dst) + symmetricstate_mix_hash(symmetric_state, dst) + } + } + } + if message_pattern.is_psk { + if len(psk) != PSK_SIZE { + status = .Invalid_Pre_Shared_Key + break do_init + } + } else if len(psk) != 0 { + status = .Unexpected_Pre_Shared_Key + break do_init + } + } + if status != .Ok { + symmetricstate_reset(symmetric_state) + return status + } + + if s != nil { + ecdh.private_key_set(&handshake_state.s, s) + } + if e != nil { + ecdh.private_key_set(&handshake_state.e, e) + handshake_state.pre_set_e = true + } + if rs != nil { + ecdh.public_key_set(&handshake_state.rs, rs) + } + if re != nil { + ecdh.public_key_set(&handshake_state.re, re) + } + copy(handshake_state.psk[:], psk) + handshake_state.message_pattern = HANDSHAKE_PATTERNS[symmetric_state.protocol.handshake_pattern] + handshake_state.current_message = 0 + handshake_state.status = .Handshake_Pending + handshake_state.initiator = initiator + + return .Ok +} + +handshakestate_reset :: proc(self: ^Handshake_State) { + symmetricstate_reset(&self.symmetric_state) + ecdh.private_key_clear(&self.s) + ecdh.private_key_clear(&self.e) + + crypto.zero_explicit(self, size_of(Handshake_State)) +} + +// Takes a payload byte sequence which may be zero-length, and a +// message_buffer to write the output into. +// Performs the following steps, aborting if any EncryptAndHash() call +// returns an error: +// - Fetches and deletes the next message pattern from message_pattern, +// then sequentially processes each token from the message pattern: +// - For "e": Sets e (which must be empty) to GENERATE_KEYPAIR(). +// Appends e.public_key to the buffer. Calls MixHash(e.public_key). +// - For "s": Appends EncryptAndHash(s.public_key) to the buffer. +// - For "ee": Calls MixKey(DH(e, re)). +// - For "es": Calls MixKey(DH(e, rs)) if initiator, MixKey(DH(s, re)) +// if responder. +// - For "se": Calls MixKey(DH(s, re)) if initiator, MixKey(DH(e, rs)) +// if responder. +// - For "ss": Calls MixKey(DH(s, rs)). +// - Appends EncryptAndHash(payload) to the buffer. +// – (SKIPPED) If there are no more message patterns returns two new +// CipherState objects by calling Split(). +// +// Calling Split() is left to a separate function, although it is technically +// part of the specification. +@(require_results) +handshakestate_write_message :: proc(self: ^Handshake_State, payload, dst: []byte, allocator := context.allocator) -> ([]byte, Status) { + ensure(self.status == .Handshake_Pending, "crypto/noise: invalid state for WriteMessage") + + protocol := &self.symmetric_state.protocol + d_len := dh_len(protocol) + + pattern_buf: [dynamic; MAX_STEP_MSG_SIZE]byte + dh_buf: [MAX_DH_SIZE]byte = --- + defer crypto.zero_explicit(&dh_buf, size_of(dh_buf)) + + pattern := self.message_pattern.messages[self.current_message] + for token in pattern { + switch token { + case .e: + switch self.pre_set_e { + case true: + // Note: "which must be empty", but we allow pre-generated `e` + // for testing/rng-less systems. + self.pre_set_e = false + case false: + if ecdh.curve(&self.e) != .Invalid { + panic("crypto/noise: e was not empty when processing token 'e' during WriteMessage") + } + generate_keypair(protocol, &self.e) + } + e_public := dh_buf[:d_len] + ecdh.private_key_public_bytes(&self.e, e_public) + n := append(&pattern_buf, ..e_public) + ensure(n == d_len, "crypto/noise: truncated append `e`") + + symmetricstate_mix_hash(&self.symmetric_state, e_public) + if self.message_pattern.is_psk { + symmetricstate_mix_key(&self.symmetric_state, e_public) + } + + case .s: + s_public := dh_buf[:d_len] + ecdh.private_key_public_bytes(&self.s, s_public) + + tmp: [MAX_DH_SIZE+TAG_SIZE]byte = --- + dh_buf := tmp[:d_len+TAG_SIZE] + if !cipherstate_has_key(&self.symmetric_state.cipher_state) { + dh_buf = tmp[:d_len] + } + ct, status := symmetricstate_encrypt_and_hash(&self.symmetric_state, s_public, dh_buf) + if status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + n := append(&pattern_buf, ..ct) + ensure(n == len(ct), "crypto/noise: truncated append `s`") + + case .ee: + dh := dh_buf[:d_len] + if status := _dh(&self.e, &self.re, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + + case .es: + dh := dh_buf[:d_len] + if self.initiator { + if status := _dh(&self.e, &self.rs, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + } else { + if status := _dh(&self.s, &self.re, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + } + + case .se: + dh := dh_buf[:d_len] + if self.initiator { + if status := _dh(&self.s, &self.re, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + } else { + if status := _dh(&self.e, &self.rs, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + } + + case .ss: + dh := dh_buf[:d_len] + if status := _dh(&self.s, &self.rs, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + + case .psk: + symmetricstate_mix_key_and_hash(&self.symmetric_state, self.psk[:]) + } + } + self.current_message += 1 // Advance after the current message is successful. + + pattern_len := len(pattern_buf) + payload_len := len(payload) + msg_len := pattern_len + payload_len + if cipherstate_has_key(&self.symmetric_state.cipher_state) { + msg_len += TAG_SIZE + } + + msg: []byte + if msg_len != 0 { + did_alloc: bool + if dst != nil { + if len(dst) < msg_len { + self.status = .Handshake_Failed + return nil, .Out_Of_Memory + } + msg = dst[:msg_len] + } else { + err: runtime.Allocator_Error + msg, err = make([]byte, msg_len, allocator) + if err != nil { + self.status = .Handshake_Failed + return nil, .Out_Of_Memory + } + did_alloc = true + } + + copy(msg, pattern_buf[:]) + ciphertext := msg[pattern_len:] + if _, status := symmetricstate_encrypt_and_hash(&self.symmetric_state, payload, ciphertext); status != .Ok { + if did_alloc { + delete(msg) + } + self.status = .Handshake_Failed + return nil, status + } + } + + if self.current_message == len(self.message_pattern.messages) { + self.current_message = -1 + self.status = .Handshake_Complete + } + + return msg, self.status +} + +// Takes a byte sequence containing a Noise handshake message, and a +// payload_buffer to write the message's plaintext payload into. +// Performs the following steps, aborting if any DecryptAndHash() +// call returns an error: +// - Fetches and deletes the next message pattern from message_pattern, +// then sequentially processes each token from the message pattern: +// - For "e": Sets re (which must be empty) to the next DHLEN bytes +// from the message. Calls MixHash(re.public_key). +// - For "s": Sets temp to the next DHLEN + 16 bytes of the message +// if HasKey() == True, or to the next DHLEN bytes otherwise. +// Sets rs (which must be empty) to DecryptAndHash(temp). +// - For "ee": Calls MixKey(DH(e, re)). +// - For "es": Calls MixKey(DH(e, rs)) if initiator, MixKey(DH(s, re)) +// if responder. +// - For "se": Calls MixKey(DH(s, re)) if initiator, MixKey(DH(e, rs)) +// if responder. +// -For "ss": Calls MixKey(DH(s, rs)). +// - Calls DecryptAndHash() on the remaining bytes of the message and stores +// the output into payload_buffer. +// – (SKIPPED) If there are no more message patterns returns two new +// CipherState objects by calling Split(). +// +// Calling Split() is left to a separate function, although it is technically +// part of the specification. +@(require_results) +handshakestate_read_message :: proc(self: ^Handshake_State, message, dst: []byte, allocator := context.allocator) -> ([]byte, Status) { + ensure(self.status == .Handshake_Pending, "crypto/noise: invalid state for ReadMessage") + + protocol := &self.symmetric_state.protocol + d_len := dh_len(&self.symmetric_state.protocol) + + dh_buf: [MAX_DH_SIZE]byte = --- + defer crypto.zero_explicit(&dh_buf, size_of(dh_buf)) + + msg := message + + pattern := self.message_pattern.messages[self.current_message] + for token in pattern { + switch token { + case .e: + if len(msg) < d_len { + return nil, .Invalid_Handshake_Message + } + re := msg[:d_len] + + if ecdh.curve(&self.re) != .Invalid { + panic("crypto/noise: re was not empty when processing token 'e' during ReadMessage") + } + + if !ecdh.public_key_set_bytes(&self.re, protocol.dh, re) { + return nil, .Invalid_Handshake_Message + } + symmetricstate_mix_hash(&self.symmetric_state, re) + if self.message_pattern.is_psk { + symmetricstate_mix_key(&self.symmetric_state, re) + } + msg = msg[d_len:] + + case .s: + rs_len := d_len + if cipherstate_has_key(&self.symmetric_state.cipher_state) { + rs_len += TAG_SIZE + } + if len(msg) < rs_len { + self.status = .Handshake_Failed + return nil, .Invalid_Handshake_Message + } + + rs := dh_buf[:d_len] + if _, status := symmetricstate_decrypt_and_hash(&self.symmetric_state, msg[:rs_len], rs); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + + if ecdh.curve(&self.rs) != .Invalid { + panic("crypto/noise: rs was not empty when processing token 's' during ReadMessage") + } + + if !ecdh.public_key_set_bytes(&self.rs, protocol.dh, rs) { + self.status = .Handshake_Failed + return nil, .Invalid_Handshake_Message + } + msg = msg[rs_len:] + + case .ee: + dh := dh_buf[:d_len] + if status := _dh(&self.e, &self.re, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + + case .es: + dh := dh_buf[:d_len] + if self.initiator { + if status := _dh(&self.e, &self.rs, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + } else { + if status := _dh(&self.s, &self.re, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + } + + case .se: + dh := dh_buf[:d_len] + if self.initiator { + if status := _dh(&self.s, &self.re, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + } else { + if status := _dh(&self.e, &self.rs, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + } + + case .ss: + dh := dh_buf[:d_len] + if status := _dh(&self.s, &self.rs, dh); status != .Ok { + self.status = .Handshake_Failed + return nil, status + } + symmetricstate_mix_key(&self.symmetric_state, dh) + + case .psk: + symmetricstate_mix_key_and_hash(&self.symmetric_state, self.psk[:]) + } + } + self.current_message += 1 // Advance after the current message is successful. + + payload: []byte + payload_len := len(msg) + if cipherstate_has_key(&self.symmetric_state.cipher_state) { + if payload_len < TAG_SIZE { + self.status = .Handshake_Failed + return nil, self.status + } + payload_len -= TAG_SIZE + } + + did_alloc: bool + if dst != nil { + if len(dst) < payload_len { + self.status = .Handshake_Failed + return nil, .Out_Of_Memory + } + payload = dst[:payload_len] + } else if payload_len > 0 { + err: runtime.Allocator_Error + payload, err = make([]byte, payload_len, allocator) + if err != nil { + self.status = .Handshake_Failed + return nil, .Out_Of_Memory + } + did_alloc = true + } + + if _, status := symmetricstate_decrypt_and_hash(&self.symmetric_state, msg, payload); status != .Ok { + if did_alloc { + delete(payload) + } + self.status = .Handshake_Failed + return nil, self.status + } + + if self.current_message == len(self.message_pattern.messages) { + self.current_message = -1 + self.status = .Handshake_Complete + } + + return payload, self.status +} + +@(require_results) +protocol_from_string :: proc(self: ^Protocol, protocol_name: string) -> Status { + self^ = Protocol{} + + pattern, dh, cipher, hash, status := split_protocol_string(protocol_name) + if status != .Ok { + return status + } + + self.handshake_pattern = pattern + self.dh = dh + self.cipher = cipher + self.hash = hash + + return .Ok +} diff --git a/core/crypto/rsa/doc.odin b/core/crypto/rsa/doc.odin new file mode 100644 index 000000000..537333cfc --- /dev/null +++ b/core/crypto/rsa/doc.odin @@ -0,0 +1,7 @@ +/* +RSA (Rivest–Shamir–Adleman) cryptosystem. + +See: +- [[ https://www.rfc-editor.org/info/rfc8017/ ]] +*/ +package rsa diff --git a/core/crypto/rsa/rsa.odin b/core/crypto/rsa/rsa.odin new file mode 100644 index 000000000..29ef4dac8 --- /dev/null +++ b/core/crypto/rsa/rsa.odin @@ -0,0 +1,444 @@ +package rsa + +// Copyright (c) 2016 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "core:bytes" +import "core:crypto" +import subtle "core:crypto/_subtle" +import "core:encoding/endian" + +// Minimum size for a RSA modulus (in bits). +// +// Note: 1024-bits is arguably insufficient as of this writing, with +// 2048-bits being a more sensible value, however 1024-bits is likely +// still in frequent enough use. +// +// Note: CA signed TLS certificates have a strict requirement of a modulus +// size that is at least 2048-bits [[ https://cabforum.org/working-groups/server/baseline-requirements/documents/]]. +MODULUS_MIN_SIZE :: 1024 + +// Maximum size for a RSA modulus (in bits). +// +// This value MUST be a multiple of 64. This value MUST NOT exceed 47666 +// (some computations in RSA key generation rely on the factor size being +// no more than 23833 bits). RSA key sizes beyond 3072 bits don't make a +// lot of sense anyway. +MODULUS_MAX_SIZE :: 4096 + +// Maxmimum size for a RSA public exponent (in bits). +// +// Note: This implementation supports arbitrary size exponents, however +// limit it to something sensible (some implementations are known to +// choke on exponents >= 2^32), with the most common choice being +// `65537`. +EXPONENT_MAX_SIZE :: 32 + +// Maximum size for a RSA factor (in bits). This is for RSA private-key +// operations. Default is to support factors up to a bit more than half +// the maximum modulus size. +// +// This value MUST be a multiple of 32. +FACTOR_MAX_SIZE :: (MODULUS_MAX_SIZE + 64) >> 1 + +// Default size for a RSA key (in bits). +DEFAULT_MODULUS_SIZE :: 2048 + +// RSA public exponent used for key generation. This MUST be a prime +// number greater than 2. +@(private) +PUBLIC_EXPONENT :: 65537 + +#assert(EXPONENT_MAX_SIZE <= 32) + +// Private_Key is a RSA private key. +Private_Key :: struct { + _pub_key: Public_Key, + _d: Modulus, // Private exponent has the same size as n. + _p: Factor, + _q: Factor, + + // CRT coefficients. + _dp: Factor, // d % (p - 1) + _dq: Factor, // d % (q - 1) + _iq: Factor, // q^(-1) mod p + + _is_initialized: bool, +} + +// Public_Key is a RSA public key. +Public_Key :: struct { + _n: Modulus, + _e: u32, + _is_initialized: bool, +} + +// private_key_generate uses the system entropy source to generate a new +// Private_Key. The key size is specified in bits, and must be a multiple +// of 8. +@(require_results) +private_key_generate :: proc(priv_key: ^Private_Key, key_size := DEFAULT_MODULUS_SIZE) -> bool { + if !crypto.HAS_RAND_BYTES { + return false + } + if key_size < MODULUS_MIN_SIZE || key_size > MODULUS_MAX_SIZE { + return false + } + if key_size % 8 != 0 { + return false + } + + private_key_clear(priv_key) + defer if !priv_key._is_initialized { + private_key_clear(priv_key) + } + + for { + // The only way this can fail is if we get extremely unlucky + // and we fail to derive `iq` (1/d mod p). + if keygen_inner(priv_key, key_size) == 1 { + break + } + } + priv_key._is_initialized = true + priv_key._pub_key._is_initialized = true + + // Self-test the key. + priv_key._is_initialized = pkcs1_sig_selftest(priv_key) + + return priv_key._is_initialized +} + +// private_key_n copies the private key's public modulus to dst if dst is +// non-nil and of sufficient size, and returns the number of bytes +// copied/would be copied (ie: calling with `dst = nil` gets the required +// size). +@(require_results) +private_key_n :: proc(priv_key: ^Private_Key, dst: []byte) -> (n_len: int) { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return public_key_n(&priv_key._pub_key, dst) +} + +// private_key_e returns the private key's public exponent as a u32. +@(require_results) +private_key_e :: proc(priv_key: ^Private_Key) -> u32 { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return public_key_e(&priv_key._pub_key) +} + +// private_key_d copies the private key's private exponent `d` to dst if +// dst is non-nil and of sufficient size, and returns the number of bytes +// copied/would be copied (ie: calling with `dst = nil` gets the required +// size). +// +// Note: The data returned MUST be kept confidential. +@(require_results) +private_key_d :: proc(priv_key: ^Private_Key, dst: []byte) -> (n_len: int) { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return modulus_copyout(&priv_key._d, dst) +} + +// private_key_p copies the private key's first prime factor `p` to dst +// if dst is non-nil and of sufficient size, and returns the number of +// bytes copied/would be copied (ie: calling with `dst = nil` gets the +// required size). +// +// Note: The data returned MUST be kept confidential. +@(require_results) +private_key_p :: proc(priv_key: ^Private_Key, dst: []byte) -> (n_len: int) { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return factor_copyout(&priv_key._p, dst) +} + +// private_key_q copies the private key's second prime factor `q` to dst +// if dst is non-nil and of sufficient size, and returns the number of +// bytes copied/would be copied (ie: calling with `dst = nil` gets the +// required size). +// +// Note: The data returned MUST be kept confidential. +@(require_results) +private_key_q :: proc(priv_key: ^Private_Key, dst: []byte) -> (n_len: int) { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return factor_copyout(&priv_key._q, dst) +} + +// private_key_dp copies the private key's first reduced exponent +// `d % (p-1)` to dst if dst is non-nil and of sufficient size, and +// returns the number of bytes copied/would be copied (ie: calling with +//`dst = nil` gets the required size). +// +// Note: The data returned MUST be kept confidential. +@(require_results) +private_key_dp :: proc(priv_key: ^Private_Key, dst: []byte) -> (n_len: int) { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return factor_copyout(&priv_key._dp, dst) +} + +// private_key_dq copies the private key's second reduced exponent +// `d % (q-1)` to dst if dst is non-nil and of sufficient size, and +// returns the number of bytes copied/would be copied (ie: calling with +//`dst = nil` gets the required size). +// +// Note: The data returned MUST be kept confidential. +@(require_results) +private_key_dq :: proc(priv_key: ^Private_Key, dst: []byte) -> (n_len: int) { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return factor_copyout(&priv_key._dq, dst) +} + +// private_key_iq copies the private key's CRT coefficient `iq` to dst if +// dst is non-nil and of sufficient size, and returns the number of bytes +// copied/would be copied (ie: calling with`dst = nil` gets the required +// size). +// +// Note: The data returned MUST be kept confidential. +@(require_results) +private_key_iq :: proc(priv_key: ^Private_Key, dst: []byte) -> (n_len: int) { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return factor_copyout(&priv_key._iq, dst) +} + +// private_key_size returns the size of the private key's public modulus +// in bytes. All ciphertexts and signatures will also be this size. +@(require_results) +private_key_size :: proc(priv_key: ^Private_Key) -> int { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + + return priv_key._pub_key._n.v_len +} + +// private_key_set_bytes sets a private key from byte-encoded components, +// and returns true if and only if (⟺) the operation was successful. +// +// Note: All values are mandatory, and match the values included in the +// PKCS private key format. +// +// WARNING: This routine validates that it is possible to sign/verify with +// the deserialized values, however d is not checked at all, nor is the +// primality of p and q. +@(require_results) +private_key_set_bytes :: proc( + priv_key: ^Private_Key, + n: []byte, + e: []byte, + d: []byte, + p: []byte, + q: []byte, + dp: []byte, + dq: []byte, + iq: []byte, +) -> bool { + private_key_clear(priv_key) + defer if !priv_key._is_initialized { + private_key_clear(priv_key) + } + + if !public_key_set_bytes(&priv_key._pub_key, n, e) { + return false + } + + if !modulus_set_bytes(&priv_key._d, d) { + return false + } + if !factor_set_bytes(&priv_key._p, p) { + return false + } + if !factor_set_bytes(&priv_key._q, q) { + return false + } + if !factor_set_bytes(&priv_key._dp, dp) { + return false + } + if !factor_set_bytes(&priv_key._dq, dq) { + return false + } + if !factor_set_bytes(&priv_key._iq, iq) { + return false + } + + priv_key._is_initialized = true + + // Test the key. + // + // Note: This DOES NOT check that p/q are prime and if d is + // consistent (as it is not used by our implementation). + priv_key._is_initialized = pkcs1_sig_selftest(priv_key) + + return priv_key._is_initialized +} + +// private_key_set sets priv_key to src. +private_key_set :: proc(priv_key, src: ^Private_Key) { + if src == nil || !src._is_initialized { + private_key_clear(priv_key) + return + } + + public_key_set(&priv_key._pub_key, &src._pub_key) + modulus_set(&priv_key._d, &src._d) + factor_set(&priv_key._p, &src._p) + factor_set(&priv_key._q, &src._q) + factor_set(&priv_key._dp, &src._dp) + factor_set(&priv_key._dq, &src._dq) + factor_set(&priv_key._iq, &src._iq) + + priv_key._is_initialized = true +} + +// private_key_equal returns true if and only if (⟺) priv_key is equal to other. +@(require_results) +private_key_equal :: proc(priv_key, other: ^Private_Key) -> bool { + ensure(priv_key._is_initialized && other._is_initialized, "crypto/rsa: uninitialized private key") + + pk_eq := public_key_equal(&priv_key._pub_key, &other._pub_key) + + eq := crypto.compare_constant_time(modulus_bytes(&priv_key._d), modulus_bytes(&other._d)) + eq &= crypto.compare_constant_time(factor_bytes(&priv_key._p), factor_bytes(&other._p)) + eq &= crypto.compare_constant_time(factor_bytes(&priv_key._q), factor_bytes(&other._q)) + eq &= crypto.compare_constant_time(factor_bytes(&priv_key._dp), factor_bytes(&other._dp)) + eq &= crypto.compare_constant_time(factor_bytes(&priv_key._dq), factor_bytes(&other._dq)) + eq &= crypto.compare_constant_time(factor_bytes(&priv_key._iq), factor_bytes(&other._iq)) + + return pk_eq & (eq == 1) +} + +// private_key_clear clears priv_key to the uninitialized state. +private_key_clear :: proc "contextless" (priv_key: ^Private_Key) { + crypto.zero_explicit(priv_key, size_of(Private_Key)) +} + +// public_key_n copies the public key's modulus `n` to dst if dst is +// non-nil and of sufficient size, and returns the number of bytes +// copied/would be copied (ie: calling with `dst = nil` gets the +// required size). +@(require_results) +public_key_n :: proc(pub_key: ^Public_Key, dst: []byte) -> (n_len: int) { + ensure(pub_key._is_initialized, "crypto/rsa: uninitialized public key") + + return modulus_copyout(&pub_key._n, dst) +} + +// public_key_e returns the public key's exponent `e` as a u32. +@(require_results) +public_key_e :: proc(pub_key: ^Public_Key) -> u32 { + ensure(pub_key._is_initialized, "crypto/rsa: uninitialized public key") + + return pub_key._e +} + +// public_key_size returns the size of the public key's modulus in bytes. +// All ciphertexts and signatures will also be this size. +@(require_results) +public_key_size :: proc(pub_key: ^Public_Key) -> int { + ensure(pub_key._is_initialized, "crypto/rsa: uninitialized public key") + + return pub_key._n.v_len +} + +// public_key_set_bytes sets a public key from byte-encoded components, +// and returns true if and only if (⟺) the operation was successful. +@(require_results) +public_key_set_bytes :: proc(pub_key: ^Public_Key, n, e: []byte) -> bool { + public_key_clear(pub_key) + defer if !pub_key._is_initialized { + public_key_clear(pub_key) + } + + ok := modulus_set_bytes(&pub_key._n, n) + if !ok { + return false + } + if modulus_len(&pub_key._n) < MODULUS_MIN_SIZE >> 3 { + return false + } + if !modulus_is_odd(&pub_key._n) { + return false + } + + e_ := bytes.trim_left(e, []byte{0x00}) + e_len := len(e_) + if e_len > EXPONENT_MAX_SIZE >> 3 { + return false + } + e_buf: [4]byte + copy(e_buf[4 - e_len:], e) + e_u32 := endian.unchecked_get_u32be(e_buf[:]) + if e_u32 < 3 || e_u32 & 1 == 0 { + return false + } + pub_key._e = e_u32 + + pub_key._is_initialized = true + + return true +} + +// public_key_set sets pub_key to src. +public_key_set :: proc(pub_key, src: ^Public_Key) { + if src == nil || !src._is_initialized { + public_key_clear(pub_key) + return + } + + modulus_set(&pub_key._n, &src._n) + pub_key._e = src._e + pub_key._is_initialized = true +} + +// public_key_set_priv sets pub_key to the public component of priv_key. +public_key_set_priv :: proc(pub_key: ^Public_Key, priv_key: ^Private_Key) { + ensure(priv_key._is_initialized, "crypto/rsa: uninitialized private key") + pub_key^ = priv_key._pub_key +} + +// public_key_equal returns true if and only if (⟺) pub_key is equal to other. +public_key_equal :: proc(pub_key, other: ^Public_Key) -> bool { + ensure(pub_key._is_initialized && other._is_initialized, "crypto/rsa: uninitialized public key") + + eq := crypto.compare_constant_time(modulus_bytes(&pub_key._n), modulus_bytes(&other._n)) + eq &= int(subtle.eq(pub_key._e, other._e)) + + return eq == 1 +} + +// public_key_clear clears pub_key to the uninitialized state. +public_key_clear :: proc "contextless" (pub_key: ^Public_Key) { + crypto.zero_explicit(pub_key, size_of(Public_Key)) +} + +// size returns the size of the key's public modulus in bytes. +// All ciphertexts and signatures will also be this size. +size :: proc "contextless" (key: ^$T) -> int where T == Private_Key || T == Private_Key { + when T == Private_Key { + return private_key_size(key) + } else { + return public_key_size(key) + } +} diff --git a/core/crypto/rsa/rsa_dec_oaep.odin b/core/crypto/rsa/rsa_dec_oaep.odin new file mode 100644 index 000000000..37f48fea9 --- /dev/null +++ b/core/crypto/rsa/rsa_dec_oaep.odin @@ -0,0 +1,197 @@ +package rsa + +// Copyright (c) 2018 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "core:crypto" +import subtle "core:crypto/_subtle" +import "core:crypto/hash" + +// decrypt_oaep returns the plaintext and true if and only if (⟺) it +// successfully decrypts the ciphertext with OAEP parameterized by +// label, hash_algo, and mgf1_algo, and writes the plaintext into dst. +// If mgf1_algo is unspecified, hash_algo will be used. +// +// Note: dst MUST be large enough to contain the plaintext. +@(require_results) +decrypt_oaep :: proc( + priv_key: ^Private_Key, + hash_algo: hash.Algorithm, + ciphertext: []byte, + dst: []byte, + label: []byte = nil, + mgf1_algo := hash.Algorithm.Invalid, +) -> (plaintext: []byte, ok: bool) { + if !priv_key._is_initialized { + return + } + ct_len := len(ciphertext) + if ct_len != modulus_len(&priv_key._pub_key._n) { + return + } + if hash_algo == .Invalid { + return + } + mgf1_algo_ := mgf1_algo + if mgf1_algo == .Invalid { + mgf1_algo_ = hash_algo + } + + tmp: [MODULUS_MAX_SIZE >> 3]byte + pt_buf := tmp[:ct_len] + defer crypto.zero_explicit(raw_data(pt_buf), ct_len) + + copy(pt_buf, ciphertext) + r := private_modpow(pt_buf, priv_key) + r_, l := oaep_dec_unpad(hash_algo, mgf1_algo_, label, pt_buf) + + // Conditional branches are ok as we are past the padding + // verification. + if ok = r & r_ == 1; ok { + if l <= len(dst) { + copy(dst, pt_buf[:l]) + plaintext = dst[:l] + } else { + ok = false + } + } + + return +} + +// oaep_max_plaintext_size returns the maximum supported plaintext size +// for a given key, with OAEP parameterized by hash_algo and mgf1_algo. +// If mgf1_algo is unspecified, hash_algo will be used. +@(require_results) +oaep_max_plaintext_size :: proc( + k: ^$T, + hash_algo: hash.Algorithm, + mgf1_algo := hash.Algorithm.Invalid, +) -> int where T == Private_Key || T == Public_Key { + if !k._is_initialized { + return 0 + } + if hash_algo == .Invalid { + return 0 + } + mgf1_algo_ := mgf1_algo + if mgf1_algo == .Invalid { + mgf1_algo_ = hash_algo + } + + overhead := 2 + hash.DIGEST_SIZES[hash_algo] + hash.DIGEST_SIZES[mgf1_algo_] + + pub_key: ^Public_Key + when T == Private_Key { + pub_keyk = &k._pub_key + } else { + pub_key = k + } + return modulus_len(&k._n) - overhead +} + +@(private="file") +xor_hash_data :: proc(hash_algo: hash.Algorithm, dst: []byte, src: []byte) { + tmp: [hash.MAX_DIGEST_SIZE]byte = --- + hash_len := hash.DIGEST_SIZES[hash_algo] + digest := tmp[:hash_len] + defer crypto.zero_explicit(raw_data(digest), hash_len) + + hash.hash_bytes_to_buffer(hash_algo, src, digest) + for v, u in digest { + dst[u] ~= v + } +} + +@(private="file") +oaep_dec_unpad :: proc( + hash_algo: hash.Algorithm, + mgf1_algo: hash.Algorithm, + label: []byte, + data: []byte, +) -> (u32, int) { + hash_len := hash.DIGEST_SIZES[hash_algo] + k := len(data) + buf := data + + // There must be room for the padding. + if k < (hash_len << 1) + 2 { + return 0, 0 + } + + // Unmask the seed, then the DB value. + seed, db := buf[1:1+hash_len], buf[1+hash_len:] + mgf1_xor(seed, mgf1_algo, db) + mgf1_xor(db, mgf1_algo, seed) + + // Hash the label and XOR it with the value in the array; if + // they are equal then these should yield only zeros. + xor_hash_data(hash_algo, db, label) + + // At that point, if the padding was correct, when we should + // have: 0x00 || seed || 0x00 ... 0x00 0x01 || M + // Padding is valid as long as: + // - There is at least hlen+1 leading bytes of value 0x00. + // - There is at least one non-zero byte. + // - The first (leftmost) non-zero byte has value 0x01. + // + // Ultimately, we may leak the resulting message length, i.e. + // the position of the byte of value 0x01, but we must take care + // to do so only if the number of zero bytes has been verified + // to be at least hlen+1. + // + // The loop below counts the number of bytes of value 0x00, and + // checks that the next byte has value 0x01, in constant-time. + // + // - If the initial byte (before the seed) is not 0x00, then + // r and s are set to 0, and stay there. + // - Value r is 1 until the first non-zero byte is reached + // (after the seed); it switches to 0 at that point. + // - Value s is set to 1 if and only if the data encountered + // at the time of the transition of r from 1 to 0 has value + // exactly 0x01. + // - Value zlen counts the number of leading bytes of value zero + // (after the seed). + r := u32(subtle.eq(buf[0], 0)) + s, zlen: u32 + for u in hash_len + 1..> 8) + s |= nz & subtle.eq(w, 0x01) + r &= subtle.not(nz) + zlen += r + } + + // Padding is correct only if s == 1, _and_ zlen >= hlen. + s &= subtle.ge(zlen, u32(hash_len)) + + // At that point, padding was verified, and we are now allowed + // to make conditional jumps. + if s != 0 { + plen := 2 + hash_len + int(zlen) + k -= plen + copy(buf[:k], buf[plen:]) + } + return s, k +} diff --git a/core/crypto/rsa/rsa_dec_tls_pms.odin b/core/crypto/rsa/rsa_dec_tls_pms.odin new file mode 100644 index 000000000..cedd9032d --- /dev/null +++ b/core/crypto/rsa/rsa_dec_tls_pms.odin @@ -0,0 +1,55 @@ +package rsa + +// Copyright (c) 2016 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import subtle "core:crypto/_subtle" + +// unsafe_decrypt_tls_pms decrypts a TLS RSA-Encrypted Premaster Secret +// Message, unconditionally moves the decrypted plaintext to `data[:48]`, +// and returns 1 if and only if (⟺) the operation was successful. +// +// WARNING: This routine MUST only be used when implementing server-side +// support for TLS 1.2's Client Key Exchange message, and extreme care +// MUST be taken when handling failures. This key exchange scheme was +// removed in TLS 1.3, and not implementing support in the first place +// is strongly RECOMMENDED even for TLS 1.2 servers. +@(require_results) +unsafe_decrypt_tls_pms :: proc(priv_key: ^Private_Key, data: []byte) -> u32 { + // A first check on length. Since this test works only on the + // buffer length, it needs not (and cannot) be constant-time. + _len := len(data) + if _len < 59 || _len != priv_key._pub_key._n.v_len { + return 0 + } + x := private_modpow(data, priv_key) + + x &= u32(subtle.eq(data[0], 0x00)) + x &= u32(subtle.eq(data[1], 0x02)) + for u in 2..<(_len-49) { + x &= u32(subtle.neq(data[u], 0)) + } + x &= u32(subtle.eq(data[_len - 49], 0x00)) + copy(data[:48], data[_len - 48:]) + + return x +} diff --git a/core/crypto/rsa/rsa_enc_oaep.odin b/core/crypto/rsa/rsa_enc_oaep.odin new file mode 100644 index 000000000..7013b823b --- /dev/null +++ b/core/crypto/rsa/rsa_enc_oaep.odin @@ -0,0 +1,105 @@ +package rsa + +// Copyright (c) 2018 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "base:intrinsics" +import "core:crypto" +import "core:crypto/hash" + +// encrypt_oaep returns true if and only if (⟺) it successfully +// encrypts the plaintext with OAEP parameterized by label, hash_algo, +// and mgf1_algo, and writes the cipherttext into dst. If mgf1_algo is +// unspecified, hash_algo will be used. +// +// This routine will fail if the system entropy source is unavailable. +encrypt_oaep :: proc( + pub_key: ^Public_Key, + hash_algo: hash.Algorithm, + plaintext: []byte, + dst: []byte, + label: []byte = nil, + mgf1_algo := hash.Algorithm.Invalid, +) -> bool { + if !pub_key._is_initialized { + return false + } + if hash_algo == .Invalid { + return false + } + mgf1_algo_ := mgf1_algo + if mgf1_algo == .Invalid { + mgf1_algo_ = hash_algo + } + if len(dst) != modulus_len(&pub_key._n) { + return false + } + if len(plaintext) > oaep_max_plaintext_size(pub_key, hash_algo, mgf1_algo_) { + return false + } + + if oaep_enc_pad(hash_algo, mgf1_algo_, label, dst, plaintext) != 1 { + return false + } + + return public_modpow(dst, pub_key) == 1 +} + +@(private="file") +oaep_enc_pad :: proc( + hash_algo: hash.Algorithm, + mgf1_algo: hash.Algorithm, + label: []byte, + dst: []byte, + src: []byte, +) -> u32 { + hash_len := hash.DIGEST_SIZES[hash_algo] + src_len := len(src) + k := len(dst) + + // Note: Length checks are handled by the caller. + + // Apply padding. At this point, things cannot fail. + buf := dst + + // Assemble: DB = lHash || PS || 0x01 || M + // We first place the source message M with copy(), so that + // overlaps between source and destination buffers are supported. + copy(buf[k - src_len:], src) + hash.hash_bytes_to_buffer(hash_algo, label, buf[1+hash_len:1+hash_len << 1]) + intrinsics.mem_zero(raw_data(buf[1 + hash_len << 1:]), k - src_len - (hash_len << 1) - 2) + buf[k - src_len - 1] = 0x01 + + // Make the random seed. + seed, db := buf[1:1+hash_len], buf[1+hash_len:] + crypto.rand_bytes(seed) + + // Mask DB with the mask generated from the seed. + mgf1_xor(db, mgf1_algo, seed) + + // Mask the seed with the mask generated from the masked DB. + mgf1_xor(seed, mgf1_algo, db) + + // Padding result: EM = 0x00 || maskedSeed || maskedDB. + buf[0] = 0x00 + return 1 +} diff --git a/core/crypto/rsa/rsa_int.odin b/core/crypto/rsa/rsa_int.odin new file mode 100644 index 000000000..1e84c1fcb --- /dev/null +++ b/core/crypto/rsa/rsa_int.odin @@ -0,0 +1,110 @@ +#+private +package rsa + +import "core:bytes" + +Big_Int :: struct($N: int) { + v: [N]byte, + v_len: int, +} + +Modulus :: Big_Int(MODULUS_MAX_SIZE >> 3) +Factor :: Big_Int(FACTOR_MAX_SIZE >> 3) + +@(require_results) +modulus_set_bytes :: proc(n: ^Modulus, b: []byte) -> bool { + b_ := bytes.trim_left(b, []byte{0x00}) + b_len := len(b_) + + if b_len > size_of(n.v) || b_len == 0 { + return false + } + + copy(n.v[:], b_) + n.v_len = b_len + + return true +} + +modulus_set :: proc "contextless" (n, other: ^Modulus) { + // Copy the full thing. + copy(n.v[:], other.v[:]) + n.v_len = other.v_len +} + +@(require_results) +modulus_bytes :: #force_inline proc "contextless" (n: ^Modulus) -> []byte { + return n.v[:n.v_len] +} + +@(require_results) +modulus_len :: #force_inline proc "contextless" (n: ^Modulus) -> int { + return n.v_len +} + +@(require_results) +modulus_copyout :: proc(n: ^Modulus, dst: []byte) -> (n_len: int) { + if n_len = modulus_len(n); n_len == 0 { + return + } + + if len(dst) > 0 { + ensure(len(dst) >= n_len, "crypto/rsa: insufficent buffer size") + copy(dst, modulus_bytes(n)) + } + + return +} + +@(require_results) +modulus_is_odd :: proc "contextless" (n: ^Modulus) -> bool { + if n.v_len == 0 || n.v[n.v_len-1] & 1 == 0 { + return false + } + return true +} + +@(require_results) +factor_set_bytes :: proc(n: ^Factor, b: []byte) -> bool { + b_ := bytes.trim_left(b, []byte{0x00}) + b_len := len(b_) + + if b_len > size_of(n.v) || b_len == 0 { + return false + } + + copy(n.v[:], b_) + n.v_len = b_len + + return true +} + +factor_set :: proc "contextless" (n, other: ^Factor) { + // Copy the full thing. + copy(n.v[:], other.v[:]) + n.v_len = other.v_len +} + +@(require_results) +factor_bytes :: #force_inline proc "contextless" (n: ^Factor) -> []byte { + return n.v[:n.v_len] +} + +@(require_results) +factor_len :: #force_inline proc "contextless" (n: ^Factor) -> int { + return n.v_len +} + +@(require_results) +factor_copyout :: proc(n: ^Factor, dst: []byte) -> (n_len: int) { + if n_len = factor_len(n); n_len == 0 { + return + } + + if len(dst) > 0 { + ensure(len(dst) >= n_len, "crypto/rsa: insufficent buffer size") + copy(dst, factor_bytes(n)) + } + + return +} diff --git a/core/crypto/rsa/rsa_keygen.odin b/core/crypto/rsa/rsa_keygen.odin new file mode 100644 index 000000000..c0ca12043 --- /dev/null +++ b/core/crypto/rsa/rsa_keygen.odin @@ -0,0 +1,367 @@ +package rsa + +// Copyright (c) 2018 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "core:crypto" +import bigint "core:crypto/_bigint" +import subtle "core:crypto/_subtle" +import "core:slice" + +// Swap two buffers in RAM. They must be disjoint. +@(private="file") +bufswap_u32 :: proc "contextless" (b1, b2: []u32) { + l := len(b1) + + for u in 0.. u32 { + // We need temporary values for at least 7 integers of the same size + // as a factor (including header word); more space helps with performance + // (in modular exponentiations), but we much prefer to remain under + // 2 kilobytes in total, to save stack space. The macro TEMPS below + // exceeds 512 (which is a count in 32-bit words) when MODULUS_MAX_SIZE + // is greater than 4464 (default value is 4096, so the 2-kB limit is + // maintained unless MODULUS_MAX_SIZE was modified). + TEMPS :: max(512, ((((7 * ((((MODULUS_MAX_SIZE + 1) >> 1) + 61) / 31))) + 1) >> 1) << 1) + + assert(key_size >= MODULUS_MIN_SIZE && key_size <= MODULUS_MAX_SIZE) + + t64: [TEMPS >> 1]u64 + t32 := slice.reinterpret([]u32, t64[:]) + defer crypto.zero_explicit(&t64, size_of(t64)) + + esize_p := u32(key_size + 1) >> 1 + esize_q := u32(key_size) - esize_p + sk._p.v_len = int((esize_p + 7) >> 3) + sk._q.v_len = int((esize_q + 7) >> 3) + sk._dp.v_len = sk._p.v_len + sk._dq.v_len = sk._q.v_len + sk._iq.v_len = sk._p.v_len + + pk := &sk._pub_key + pk._n.v_len = (key_size + 7) >> 3 + pk._e = PUBLIC_EXPONENT + + sk._d.v_len = pk._n.v_len // Private exponent length is that of the modulus. + + // We now switch to encoded sizes. + // + // floor((x * 16913) / (2^19)) is equal to floor(x/31) for all + // integers x from 0 to 34966; the intermediate product fits on + // 30 bits, thus we can use MUL31(). + esize_p += u32(bigint._mul31(esize_p, 16913) >> 19) + esize_q += u32(bigint._mul31(esize_q, 16913) >> 19) + plen := (esize_p + 31) >> 5 + qlen := (esize_q + 31) >> 5 + p := t32 + q := p[1 + plen:] + t := q[1 + qlen:] + + // Since we use a prime exponent, when searching for candidate primes, + // checking if `GCD(e, prime - 1) = 1` is a simple matter of euclidian + // division. + for { + bigint.i62_mkprime(p, esize_p, PUBLIC_EXPONENT, t) + p[1] -= 1 + if bigint.i31_rem(p, PUBLIC_EXPONENT) != 0 { + p[1] += 1 + break + } + } + + for { + bigint.i62_mkprime(q, esize_q, PUBLIC_EXPONENT, t) + q[1] -= 1 + if bigint.i31_rem(q, PUBLIC_EXPONENT) != 0 { + q[1] += 1 + break + } + } + + // If p and q have the same size, then it is possible that q > p + // (when the target modulus size is odd, we generate p with a + // greater bit length than q). If q > p, we want to swap p and q + // for two reasons: + // - The final step below (inversion of q modulo p) is easier if + // p > q. + // - While BearSSL's RSA code is perfectly happy with RSA keys such + // that p < q, some other implementations have restrictions and + // require p > q. + // + // Note that we can do a simple non-constant-time swap here, + // because the only information we leak here is that we insist on + // returning p and q such that p > q, which is not a secret. + if esize_p == esize_q && bigint.i31_sub(p, q, 0) == 1 { + bufswap_u32(p[:1+plen], q) + } + + sk_p, sk_q := factor_bytes(&sk._p), factor_bytes(&sk._q) + bigint.i31_encode(sk_p, p) + bigint.i31_encode(sk_q, q) + // The odds of this happening are infinitesimally small, however + // checking for it is cheap. + if crypto.compare_constant_time(sk_p, sk_q) == 1 { + return 0 + } + + // Compute the public modulus too. + bigint.i31_zero(t, p[0]) + bigint.i31_mulacc(t, p, q) + bigint.i31_encode(modulus_bytes(&pk._n), t) + + // Compute the private exponent. + // + // Computing p - 1 and q - 1 this way is safe as p and q + // are guaranteed to be odd, thus the LSB will always be + // set. + p[1], q[1] = p[1] - 1, q[1] - 1 // p = p - 1, q = q - 1 + if compute_privexp(sk, p, q, pk._e, t) != 1 { + return 0 + } + + // Compute `d % (p - 1)`. + d_mod := t[:1+plen] + bigint.i31_decode_reduce(d_mod, modulus_bytes(&sk._d), p) + bigint.i31_encode(factor_bytes(&sk._dp), d_mod) + + // Compute `d % (q - 1)`. + bigint.i31_decode_reduce(d_mod, modulus_bytes(&sk._d), q) + bigint.i31_encode(factor_bytes(&sk._dq), d_mod) + + // Compute `q^(-1) mod p`. + p[1], q[1] = p[1] + 1, q[1] + 1 // Restore p, q. + return compute_qinv(sk, p, q, plen, t) +} + +@(private="file") +compute_qinv :: proc "contextless" (sk: ^Private_Key, p, q: []u32, plen: u32, t: []u32) -> u32 { + // Per Fermat's Little Theorem, `q^(-1) mod p = q^(p-2) mod p`. + // + // Note: p is guaranteed to be odd as it is a large prime. + + // Compute and encode `p-2`. + p_minus_two := t[:1+plen] + copy(p_minus_two, p[:1+plen]) + two := t[1+plen:] // Temporarily use this for 2. + bigint.i31_zero(two, p[0]) + bigint.i31_decode(two, []byte{2}) + _ = bigint.i31_sub(p_minus_two, two, 1) + iq := factor_bytes(&sk._iq) // Temporarily use this for p - 2. + bigint.i31_encode(iq, p_minus_two) + + // Enforce 64-bit alignment. + t_ := t + if len(t_) & 1 != 0 { + t_ = t_[1:] + } + + m0i := bigint.i31_ninv31(p[1]) + ret := bigint.i62_modpow_opt_as_i31(q, iq, p, m0i, t) + if ret != 0 { + bigint.i31_encode(iq, q) + } + + return ret +} + +@(private="file") +compute_privexp :: proc "contextless" (sk: ^Private_Key, p_minus_one, q_minus_one: []u32, e: u32, tmp: []u32) -> u32 { + // Compute phi = (p-1)*(q-1). The mulacc function sets the announced + // bit length of t to be the sum of the announced bit lengths of + // p-1 and q-1, which is usually exact but may overshoot by one 1 + // bit in some cases; we readjust it to its true length. + phi := tmp + bigint.i31_zero(phi, p_minus_one[0]) + bigint.i31_mulacc(phi, p_minus_one, q_minus_one) + _len := (phi[0] + 31) >> 5 + phi[0] = bigint.i31_bit_length(phi[1:1+_len]) + _len = (phi[0] + 31) >> 5 + + // Divide phi by public exponent e. The final remainder r must be + // non-zero (otherwise, the key is invalid). The quotient is k, + // which we write over phi, since we don't need phi after that. + r: u32 + for u := _len; u >= 1; u -= 1 { + // Upon entry, r < e, and phi[u] < 2^31; hence, + // hi:lo < e*2^31. Thus, the produced word k[u] + // must be lower than 2^31, and the new remainder r + // is lower than e. + hi := r >> 1 + lo := (r << 31) + phi[u] + phi[u], r = bigint.div_rem_u32(hi, lo, e) + } + if r == 0 { + return 0 + } + k := phi + + // Compute u and v such that u*e - v*r = GCD(e,r). We use + // a binary GCD algorithm, with 6 extra integers a, b, + // u0, u1, v0 and v1. Initial values are: + // a = e u0 = 1 v0 = 0 + // b = r u1 = r v1 = e-1 + // The following invariants are maintained: + // a = u0*e - v0*r + // b = u1*e - v1*r + // 0 < a <= e + // 0 < b <= r + // 0 <= u0 <= r + // 0 <= v0 <= e + // 0 <= u1 <= r + // 0 <= v1 <= e + // + // At each iteration, we reduce either a or b by one bit, and + // adjust u0, u1, v0 and v1 to maintain the invariants: + // - if a is even, then a <- a/2 + // - otherwise, if b is even, then b <- b/2 + // - otherwise, if a > b, then a <- (a-b)/2 + // - otherwise, if b > a, then b <- (b-a)/2 + // Algorithm stops when a = b. At that point, the common value + // is the GCD of e and r; it must be 1 (otherwise, the private + // key or public exponent is not valid). The (u0,v0) or (u1,v1) + // pairs are the solution we are looking for. + // + // Since either a or b is reduced by at least 1 bit at each + // iteration, 62 iterations are enough to reach the end + // condition. + // + // To maintain the invariants, we must compute the same operations + // on the u* and v* values that we do on a and b: + // - When a is divided by 2, u0 and v0 must be divided by 2. + // - When b is divided by 2, u1 and v1 must be divided by 2. + // - When b is subtracted from a, u1 and v1 are subtracted from + // u0 and v0, respectively. + // - When a is subtracted from b, u0 and v0 are subtracted from + // u1 and v1, respectively. + // + // However, we want to keep the u* and v* values in their proper + // ranges. The following remarks apply: + // + // - When a is divided by 2, then a is even. Therefore: + // + // * If r is odd, then u0 and v0 must have the same parity; + // if they are both odd, then adding r to u0 and e to v0 + // makes them both even, and the division by 2 brings them + // back to the proper range. + // + // * If r is even, then u0 must be even; if v0 is odd, then + // adding r to u0 and e to v0 makes them both even, and the + // division by 2 brings them back to the proper range. + // + // Thus, all we need to do is to look at the parity of v0, + // and add (r,e) to (u0,v0) when v0 is odd. In order to avoid + // a 32-bit overflow, we can add ((r+1)/2,(e/2)+1) after the + // division (r+1 does not overflow since r < e; and (e/2)+1 + // is equal to (e+1)/2 since e is odd). + // + // - When we subtract b from a, three cases may occur: + // + // * u1 <= u0 and v1 <= v0: just do the subtractions + // + // * u1 > u0 and v1 > v0: compute: + // (u0, v0) <- (u0 + r - u1, v0 + e - v1) + // + // * u1 <= u0 and v1 > v0: compute: + // (u0, v0) <- (u0 + r - u1, v0 + e - v1) + // + // The fourth case (u1 > u0 and v1 <= v0) is not possible + // because it would contradict "b < a" (which is the reason + // why we subtract b from a). + // + // The tricky case is the third one: from the equations, it + // seems that u0 may go out of range. However, the invariants + // and ranges of other values imply that, in that case, the + // new u0 does not actually exceed the range. + // + // We can thus handle the subtraction by adding (r,e) based + // solely on the comparison between v0 and v1. + a, b: u32 = e, r + u0, v0: u32 = 1, 0 + u1, v1: u32 = r, e - 1 + hr, he := (r + 1) >> 1, (e >> 1) + 1 + for _ in 0..<62 { + oa := a & 1 // 1 if a is odd + ob := b & 1 // 1 if b is odd + agtb := subtle.gt(a, b) // 1 if a > b + bgta := subtle.gt(b, a) // 1 if b > a + + sab := oa & ob & agtb // 1 if a <- a-b + sba := oa & ob & bgta // 1 if b <- b-a + + // a <- a-b, u0 <- u0-u1, v0 <- v0-v1 + ctl := subtle.gt(v1, v0) + a -= b & -sab + u0 -= (u1 - (r & -ctl)) & -sab + v0 -= (v1 - (e & -ctl)) & -sab + + // b <- b-a, u1 <- u1-u0 mod r, v1 <- v1-v0 mod e + ctl = subtle.gt(v0, v1) + b -= a & -sba + u1 -= (u0 - (r & -ctl)) & -sba + v1 -= (v0 - (e & -ctl)) & -sba + + da := subtle.not(oa) | sab // 1 if a <- a/2 + db := (oa & subtle.not(ob)) | sba // 1 if b <- b/2 + + // a <- a/2, u0 <- u0/2, v0 <- v0/2 + ctl = v0 & 1 + a ~= (a ~ (a >> 1)) & -da + u0 ~= (u0 ~ ((u0 >> 1) + (hr & -ctl))) & -da + v0 ~= (v0 ~ ((v0 >> 1) + (he & -ctl))) & -da + + // b <- b/2, u1 <- u1/2 mod r, v1 <- v1/2 mod e + ctl = v1 & 1 + b ~= (b ~ (b >> 1)) & -db + u1 ~= (u1 ~ ((u1 >> 1) + (hr & -ctl))) & -db + v1 ~= (v1 ~ ((v1 >> 1) + (he & -ctl))) & -db + } + + // Check that the GCD is indeed 1. If not, then the key is invalid + // (and there's no harm in leaking that piece of information). + if (a != 1) { + return 0 + } + + // Now we have u0*e - v0*r = 1. Let's compute the result as: + // d = u0 + v0*k + // We still have k in the tmp[] array, and its announced bit + // length is that of phi. + m := k[1+_len:] + m[0] = (1 << 5) + 1 // bit length is 32 bits, encoded + m[1] = v0 & bigint.I31_MASK + m[2] = v0 >> 31 + z := m[3:] + bigint.i31_zero(z, k[0]) + z[1] = u0 & bigint.I31_MASK + z[2] = u0 >> 31 + bigint.i31_mulacc(z, k, m) + + // Encode the result. + bigint.i31_encode(modulus_bytes(&sk._d), z) + + return 1 +} diff --git a/core/crypto/rsa/rsa_mgf1.odin b/core/crypto/rsa/rsa_mgf1.odin new file mode 100644 index 000000000..70500c65a --- /dev/null +++ b/core/crypto/rsa/rsa_mgf1.odin @@ -0,0 +1,49 @@ +package rsa + +// Copyright (c) 2018 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "core:crypto/hash" +import "core:encoding/endian" + +@(private) +mgf1_xor :: proc(data: []byte, hash_algo: hash.Algorithm, seed: []byte) { + tmp: [hash.MAX_DIGEST_SIZE]byte = --- + ctx: hash.Context = --- + + buf, blen := data, len(data) + hlen := hash.DIGEST_SIZES[hash_algo] + digest := tmp[:hlen] + for u, c := int(0), u32(0); u < blen; u, c = u + hlen, c + 1 { + hash.init(&ctx, hash_algo) + hash.update(&ctx, seed) + endian.unchecked_put_u32be(tmp[:], c) + hash.update(&ctx, tmp[:4]) + hash.final(&ctx, digest) + for v in 0..= blen { + break + } + buf[u + v] ~= digest[v] + } + } +} diff --git a/core/crypto/rsa/rsa_modpow_priv.odin b/core/crypto/rsa/rsa_modpow_priv.odin new file mode 100644 index 000000000..2de4ffb39 --- /dev/null +++ b/core/crypto/rsa/rsa_modpow_priv.odin @@ -0,0 +1,165 @@ +package rsa + +// Copyright (c) 2016 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "core:crypto" +import bigint "core:crypto/_bigint" +import "core:slice" + +@(private, require_results) +private_modpow :: proc(x: []byte, sk: ^Private_Key) -> u32 { + U :: (2 + ((FACTOR_MAX_SIZE + 30) / 31)) + TLEN :: (4 * U) // TLEN is counted in 64-bit words + + ensure(sk._is_initialized, "crypto/rsa: uninitialized private key") + + // Compute the actual lengths of p and q, in bytes. + // These lengths are not considered secret (we cannot really hide + // them anyway in constant-time code). + // + // Note/yawning: The factors should already be the correct size, + // with leading `0x00`s stripped. + p := factor_bytes(&sk._p) + plen := len(p) + for plen > 0 && p[0] == 0 { + p = p[1:] + plen -= 1 + } + q := factor_bytes(&sk._q) + qlen := len(q) + for qlen > 0 && q[0] == 0 { + q = q[1:] + qlen -= 1 + } + + // Compute the maximum factor length, in 31-bit words. + z := max(plen, qlen) << 3 + fwlen := 1 + for z > 0 { + z -= 31 + fwlen += 1 + } + + // Convert size to 62-bit words. + fwlen = (fwlen + 1) >> 1 + + // We need to fit at least 6 values in the stack buffer. + if 6 * fwlen > TLEN { + return 0 + } + + // Compute signature length (in bytes). + xlen := modulus_len(&sk._pub_key._n) + + tmp_: [TLEN]u64 // WARNING: This must be zeroed out. + defer crypto.zero_explicit(&tmp_, size_of(tmp_)) + tmp := tmp_[:] + + // Decode q. + mq := slice.reinterpret([]u32, tmp) + bigint.i31_decode(mq, q) + + // Decode p. + t1 := slice.reinterpret([]u32, tmp[fwlen:]) + bigint.i31_decode(t1, p) + + // Upstream recomputes the public modulus n, but we can just + // decode it as our key representation stores all PKCS#1 + // private key values, + t2 := slice.reinterpret([]u32, tmp[2*fwlen:]) + bigint.i31_decode(t2, modulus_bytes(&sk._pub_key._n)) + + // We encode the modulus into bytes, to perform the comparison + // with bytes. We know that the product length, in bytes, is + // exactly xlen. + // The comparison actually computes the carry when subtracting + // the modulus from the source value; that carry must be 1 for + // a value in the correct range. We keep it in r, which is our + // accumulator for the error code. + m_buf := slice.reinterpret([]byte, tmp[4*fwlen:]) + bigint.i31_encode(m_buf[:xlen], t2) + u := xlen + r: u32 + for u > 0 { + u -= 1 + wn := u32(m_buf[u]) + wx := u32(x[u]) + r = ((wx - (wn + r)) >> 8) & 1 + } + + // Move the decoded p to another temporary buffer. + mp := t2 + copy(mp, t1[:2*fwlen]) + + // Compute s2 = x^dq mod q. + q0i := bigint.i31_ninv31(mq[1]) + s2 := t1 + bigint.i31_decode_reduce(s2, x, mq) + r &= bigint.i62_modpow_opt(s2, factor_bytes(&sk._dq), mq, q0i, tmp[3*fwlen:]) + + // Compute s1 = x^dp mod p. + p0i := bigint.i31_ninv31(mp[1]) + s1 := slice.reinterpret([]u32, tmp[3*fwlen:]) + bigint.i31_decode_reduce(s1, x, mp) + r &= bigint.i62_modpow_opt(s1, factor_bytes(&sk._dp), mp, p0i, tmp[4*fwlen:]) + + // Compute: + // h = (s1 - s2)*(1/q) mod p + // s1 is an integer modulo p, but s2 is modulo q. PKCS#1 is + // unclear about whether p may be lower than q (some existing, + // widely deployed implementations of RSA don't tolerate p < q), + // but we want to support that occurrence, so we need to use the + // reduction function. + // + // Since we use br_i31_decode_reduce() for iq (purportedly, the + // inverse of q modulo p), we also tolerate improperly large + // values for this parameter. + t1 = slice.reinterpret([]u32, tmp[4*fwlen:]) + t2 = slice.reinterpret([]u32, tmp[5*fwlen:]) + bigint.i31_reduce(t2, s2, mp) + _ = bigint.i31_add(s1, mp, bigint.i31_sub(s1, t2, 1)) + bigint.i31_to_monty(s1, mp) + bigint.i31_decode_reduce(t1, factor_bytes(&sk._iq), mp) + bigint.i31_montymul(t2, s1, t1, mp, p0i) + + // h is now in t2. We compute the final result: + // s = s2 + q*h + // All these operations are non-modular. + // + // We need mq, s2 and t2. We use the t3 buffer as destination. + // The buffers mp, s1 and t1 are no longer needed, so we can + // reuse them for t3. Moreover, the first step of the computation + // is to copy s2 into t3, after which s2 is not needed. Right + // now, mq is in slot 0, s2 is in slot 1, and t2 is in slot 5. + // Therefore, we have ample room for t3 by simply using s2. + t3 := s2 + bigint.i31_mulacc(t3, mq, t2) + + // Encode the result. Since we already checked the value of xlen, + // we can just use it right away. + bigint.i31_encode(x, t3) + + // The only error conditions remaining at that point are invalid + // values for p and q (even integers). + return p0i & q0i & r +} diff --git a/core/crypto/rsa/rsa_modpow_pub.odin b/core/crypto/rsa/rsa_modpow_pub.odin new file mode 100644 index 000000000..f40bddf29 --- /dev/null +++ b/core/crypto/rsa/rsa_modpow_pub.odin @@ -0,0 +1,89 @@ +package rsa + +// Copyright (c) 2016 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import bigint "core:crypto/_bigint" +import "core:encoding/endian" +import "core:slice" + +@(private, require_results) +public_modpow :: proc(x: []byte, pk: ^Public_Key) -> u32 { + TLEN :: (2 * (2 + ((MODULUS_MAX_SIZE + 30) / 31))) + + ensure(pk._is_initialized, "crypto/rsa: uninitialized public key") + + // Get the actual length of the modulus, and see if it fits within + // our stack buffer. We also check that the length of x[] is valid. + // + // Note/yawning: The modulus should already be the correct size, + // with leading `0x00`s stripped. + n := modulus_bytes(&pk._n) + nlen := modulus_len(&pk._n) + for nlen > 0 && n[0] == 0 { + n = n[1:] + nlen -= 1 + } + if nlen == 0 || nlen > (MODULUS_MAX_SIZE >> 3) || len(x) != nlen { + return 0 + } + z := nlen << 3 + fwlen := 1 + for z > 0 { + z -= 31 + fwlen += 1 + } + // Convert fwlen to a count in 62-bit words. + fwlen = (fwlen + 1) >> 1 + + // The modulus gets decoded into m[]. + // The value to exponentiate goes into a[]. + tmp: [TLEN]u64 // WARNING: This must be zeroed out. + m := slice.reinterpret([]u32, tmp[:fwlen]) + a := slice.reinterpret([]u32, tmp[fwlen:2*fwlen]) + + // Decode the modulus. + bigint.i31_decode(m, n) + m0i := bigint.i31_ninv31(m[1]) + + // Note: if m[] is even, then m0i == 0. Otherwise, m0i must be + // an odd integer. + r := m0i & 1 + + // Decode x[] into a[]; we also check that its value is proper. + r &= bigint.i31_decode_mod(a, x, m) + + // Compute the modular exponentiation. + e_: [EXPONENT_MAX_SIZE >> 3]byte + e_off: int + endian.unchecked_put_u32be(e_[:], pk._e) + if e_[0] == 0 { + // `e = 65537` is the most common and sensible value, so this + // is the most sensible value. + e_off = 1 + } + bigint.i62_modpow_opt(a, e_[e_off:], m, m0i, tmp[2*fwlen:]) + + // Encode the result. + bigint.i31_encode(x, a) + return r +} diff --git a/core/crypto/rsa/rsa_sig_pkcs1.odin b/core/crypto/rsa/rsa_sig_pkcs1.odin new file mode 100644 index 000000000..dd01c654b --- /dev/null +++ b/core/crypto/rsa/rsa_sig_pkcs1.odin @@ -0,0 +1,233 @@ +package rsa + +// Copyright (c) 2017 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "core:bytes" +import "core:crypto" +import "core:crypto/hash" + +// PKCS1_HASH_OIDS maps common hash algorithms to the OIDs for +// use with PKCS#1 signatures. +@(rodata) +PKCS1_HASH_OIDS := #partial [hash.Algorithm][]byte { + // WARNING: Legacy verification ONLY. + .Insecure_SHA1 = []byte{ + 0x05, 0x2B, 0x0E, 0x03, 0x02, 0x1A, + }, + .SHA224 = []byte{ + 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02, 0x04, + }, + .SHA256 = []byte{ + 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02, 0x01, + }, + .SHA384 = []byte{ + 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02, 0x02, + }, + .SHA512 = []byte{ + 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02, 0x03, + }, + .SHA512_256 = []byte{ + 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02, 0x06, + }, +} + +@(private="file", rodata) +PKCS1_SELFTEST_DIGEST_SHA256 := []byte{ + 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, + 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f, 0x10, + 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, + 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f, 0x20, +} + +// verify_pkcs1 returns true if and only if (⟺) sig is a valid PKCS#1 +// signature by pub_key over msg, hased using hash_algo. If pre_hashed +// is set to true, it is assumed that msg is already hashed. +@(require_results) +verify_pkcs1 :: proc(pub_key: ^Public_Key, hash_algo: hash.Algorithm, msg, sig: []byte, is_prehashed := false) -> bool { + if !pub_key._is_initialized { + return false + } + if len(sig) != modulus_len(&pub_key._n) { + return false + } + + // Lookup the OID. + oid := PKCS1_HASH_OIDS[hash_algo] + if oid == nil { + return false + } + hash_len := hash.DIGEST_SIZES[hash_algo] + + // Compute the message hash. + msg_hash_buf: [hash.MAX_DIGEST_SIZE]byte = --- + msg_hash: []byte + switch is_prehashed { + case true: + if len(msg) != hash_len { + return false + } + msg_hash = msg + case false: + msg_hash = hash.hash_bytes_to_buffer(hash_algo, msg, msg_hash_buf[:]) + } + + // PKCS #1 V2.2 (RFC 8017) 8.2.2 specifies this as computing + // and comparing the padded hash, with unpadding and extracting + // the hash being an alternative. Upstream BearSSL implements + // the latter, which is not a problem if done correctly (which + // it does), however we will opt to go for implementing this + // as specified as it is more robust against implementation + // errors. + + // Compute the expected hash. + sig_buf, padded_hash_buf: [MODULUS_MAX_SIZE >> 3]byte = ---, --- + if len(sig) > len(sig_buf) { + return false + } + padded_hash_ := padded_hash_buf[:len(sig)] + if pkcs1_sig_pad(oid, msg_hash, padded_hash_) != 1 { + return false + } + + // Compute the signature's padded hash. + sig_ := sig_buf[:len(sig)] + copy(sig_, sig) + if public_modpow(sig_, pub_key) != 1 { + return false + } + + return bytes.equal(sig_, padded_hash_) +} + +// sign_pkcs1 returns true if and only if (⟺) it successfully writes +// the PKCS#1 signature by priv_key over msg, hashed using hash_algo. +// If pre_hashed is set to true, it is assumed that msg is already hashed. +@(require_results) +sign_pkcs1 :: proc(priv_key: ^Private_Key, hash_algo: hash.Algorithm, msg, sig: []byte, is_prehashed := false) -> bool { + if !priv_key._is_initialized { + return false + } + if len(sig) != modulus_len(&priv_key._pub_key._n) { + return false + } + + // Lookup the OID. + oid := PKCS1_HASH_OIDS[hash_algo] + if oid == nil { + return false + } + + // Compute the message hash. + msg_hash_buf: [hash.MAX_DIGEST_SIZE]byte = --- + msg_hash: []byte + switch is_prehashed { + case true: + if len(msg) != hash.DIGEST_SIZES[hash_algo] { + return false + } + msg_hash = msg + case false: + msg_hash = hash.hash_bytes_to_buffer(hash_algo, msg, msg_hash_buf[:]) + } + + if pkcs1_sig_pad(oid, msg_hash, sig) != 1 { + return false + } + + return private_modpow(sig, priv_key) == 1 +} + +@(private="file", require_results) +pkcs1_sig_pad :: proc "contextless" (hash_oid, hash, x: []byte) -> u32 { + // Padded hash value has format: + // 00 01 FF .. FF 00 30 x1 30 x2 06 x3 OID 05 00 04 x4 HASH + // + // with the following rules: + // + // -- Total length is equal to the modulus length (unsigned + // encoding). + // + // -- There must be at least eight bytes of value 0xFF. + // + // -- x4 is equal to the hash length (hash_len). + // + // -- x3 is equal to the encoded OID value length (hash_oid[0]). + // + // -- x2 = x3 + 4. + // + // -- x1 = x2 + x4 + 4 = x3 + x4 + 8. + // + // Note: the "05 00" is optional (signatures with and without + // that sequence exist in practice), but notes in PKCS#1 seem to + // indicate that the presence of that sequence (specifically, + // an ASN.1 NULL value for the hash parameters) may be slightly + // more "standard" than the opposite. + xlen, hash_len := len(x), len(hash) + + // Note/yawning: The hash OID is mandatory, as is the "05 00". + x3 := hash_oid[0] + + // Check that there is enough room for all the elements, + // including at least eight bytes of value 0xFF. + if xlen < int(x3) + hash_len + 21 { + return 0 + } + x[0] = 0x00 + x[1] = 0x01 + u := xlen - int(x3) - hash_len - 11 + for i in 2..< u { + x[i] = 0xff + } + x[u] = 0x00 + x[u + 1] = 0x30 + x[u + 2] = x3 + byte(hash_len) + 8 + x[u + 3] = 0x30 + x[u + 4] = x3 + 4 + x[u + 5] = 0x06 + copy(x[u+6:], hash_oid) + u += int(x3) + 7 + x[u] = 0x05 + u += 1 + x[u] = 0x00 + u += 1 + x[u] = 0x04 + u += 1 + x[u] = byte(hash_len) + u += 1 + copy(x[u:], hash) + + return 1 +} + +@(private) +pkcs1_sig_selftest :: proc(priv_key: ^Private_Key) -> bool { + sig_buf: [MODULUS_MAX_SIZE >> 3]byte = --- + defer crypto.zero_explicit(&sig_buf, size_of(sig_buf)) + + sig := sig_buf[:private_key_size(priv_key)] + if !sign_pkcs1(priv_key, .SHA256, PKCS1_SELFTEST_DIGEST_SHA256, sig, true) { + return false + } + + return verify_pkcs1(&priv_key._pub_key, .SHA256, PKCS1_SELFTEST_DIGEST_SHA256, sig, true) +} diff --git a/core/crypto/rsa/rsa_sig_pss.odin b/core/crypto/rsa/rsa_sig_pss.odin new file mode 100644 index 000000000..a2b19e509 --- /dev/null +++ b/core/crypto/rsa/rsa_sig_pss.odin @@ -0,0 +1,293 @@ +package rsa + +// Copyright (c) 2018 Thomas Pornin +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// THIS SOFTWARE IS PROVIDED BY THE AUTHORS “AS IS” AND ANY EXPRESS OR +// IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +// WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE AUTHORS BE LIABLE FOR ANY +// DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +// GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, +// WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING +// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF +// THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import "base:intrinsics" +import "core:crypto" +import bigint "core:crypto/_bigint" +import subtle "core:crypto/_subtle" +import "core:crypto/hash" + +// verify_pss returns true if and only if (⟺) sig is a valid PSS +// signature by pub_key over msg, hashed using hash_algo, and MGF1 +// parameterized by mgf1_algo and salt_len. If mgf1_algo is +// unspecified, hash_algo will be used. If pre_hashed is set +// to true, it is assumed that msg is already hashed. +@(require_results) +verify_pss :: proc( + pub_key: ^Public_Key, + hash_algo: hash.Algorithm, + salt_len: int, + msg: []byte, + sig: []byte, + is_prehashed := false, + mgf1_algo := hash.Algorithm.Invalid, +) -> bool { + if !pub_key._is_initialized { + return false + } + if hash_algo == .Invalid { + return false + } + mgf1_algo_ := mgf1_algo + if mgf1_algo == .Invalid { + mgf1_algo_ = hash_algo + } + if len(sig) != modulus_len(&pub_key._n) { + return false + } + + // Compute the message hash. + msg_hash_buf: [hash.MAX_DIGEST_SIZE]byte = --- + hash_len := hash.DIGEST_SIZES[hash_algo] + msg_hash: []byte + switch is_prehashed { + case true: + if len(msg) != hash_len { + return false + } + msg_hash = msg + case false: + msg_hash = hash.hash_bytes_to_buffer(hash_algo, msg, msg_hash_buf[:]) + } + + sig_buf: [MODULUS_MAX_SIZE >> 3]byte = --- + sig_ := sig_buf[:len(sig)] + copy(sig_, sig) + if public_modpow(sig_, pub_key) != 1 { + return false + } + + return pss_sig_unpad(hash_algo, mgf1_algo_, msg_hash, salt_len, pub_key, sig_) == 1 +} + +// sign_pss returns true if and only if (⟺) it successfully writes +// the PKCS#1 signature by priv_key over msg, hashed using hash_algo, and +// MGF1 parameterized by mgf1_algo and salt_len. If mgf1_algo is +// unspecified, hash_algo will be used. If pre_hashed is set to true, +// it is assumed that msg is already hashed. A reasonable choice for +// salt_len is the digest size of hash_algo, and FIPS 140-3 mandates +// that as the maximum permissible size. +// +// This routine will fail if the system entropy source is unavailable. +@(require_results) +sign_pss :: proc( + priv_key: ^Private_Key, + hash_algo: hash.Algorithm, + salt_len: int, + msg: []byte, + sig: []byte, + is_prehashed := false, + mgf1_algo := hash.Algorithm.Invalid, +) -> bool { + if !priv_key._is_initialized { + return false + } + if len(sig) != modulus_len(&priv_key._pub_key._n) { + return false + } + if hash_algo == .Invalid { + return false + } + mgf1_algo_ := mgf1_algo + if mgf1_algo == .Invalid { + mgf1_algo_ = hash_algo + } + if !crypto.HAS_RAND_BYTES && salt_len != 0 { + return false + } + + // Compute the message hash. + msg_hash_buf: [hash.MAX_DIGEST_SIZE]byte = --- + hash_len := hash.DIGEST_SIZES[hash_algo] + msg_hash: []byte + switch is_prehashed { + case true: + if len(msg) != hash_len { + return false + } + msg_hash = msg + case false: + msg_hash = hash.hash_bytes_to_buffer(hash_algo, msg, msg_hash_buf[:]) + } + + // Work out the exact length of n in bits. + n := modulus_bytes(&priv_key._pub_key._n) + assert(len(n) > 0 && n[0] != 0) + n_bitlen := int(bigint._u32_bit_length(u32(n[0]))) + (len(n) - 1) * 8 + + if pss_sig_pad(hash_algo, mgf1_algo_, msg_hash, salt_len, n_bitlen, sig) != 1 { + return false + } + + return private_modpow(sig, priv_key) == 1 +} + +@(private="file", require_results) +pss_sig_unpad :: proc( + data_algo: hash.Algorithm, + mgf1_algo: hash.Algorithm, + digest: []byte, + salt_len: int, + pk: ^Public_Key, + sig: []byte, +) -> u32 { + hash_len := hash.DIGEST_SIZES[data_algo] + x := sig + + // Value r will be set to a non-zero value is any test fails. + r: u32 + + // The value bit length (as an integer) must be strictly less than + // that of the modulus. + // + // Note/yawning: The modulus should already be the correct size, + // with leading `0x00`s stripped. + n := modulus_bytes(&pk._n) + nlen := modulus_len(&pk._n) + u: int + for u = 0; u < nlen; u += 1 { + if n[u] != 0 { + break + } + } + if u == nlen { + return 0 + } + n_bitlen := bigint._u32_bit_length(u32(n[u])) + (u32(nlen - u - 1) << 3) + n_bitlen -= 1 + if (n_bitlen & 7) == 0 { + r |= u32(x[0]) + x = x[1:] + } else { + r |= u32(x[0] & (0xFF << (n_bitlen & 7))) + } + xlen := int((n_bitlen + 7) >> 3) + + // Check that the modulus is large enough for the hash value + // length combined with the intended salt length. + if hash_len > xlen || salt_len > xlen || (hash_len + salt_len + 2) > xlen { + return 0 + } + + // Check value of rightmost byte. + r |= u32(x[xlen - 1] ~ 0xBC) + + // Generate the mask and XOR it into the first bytes to reveal PS; + // we must also mask out the leading bits. + seed := x[xlen - hash_len - 1:] + mgf1_xor(x[:xlen - hash_len - 1], mgf1_algo, seed[:hash_len]) + if (n_bitlen & 7) != 0 { + x[0] &= 0xFF >> (8 - (n_bitlen & 7)) + } + + // Check that all padding bytes have the expected value. + for u = 0; u < (xlen - hash_len - salt_len - 2); u += 1 { + r |= u32(x[u]) + } + r |= u32(x[xlen - hash_len - salt_len - 2] ~ 0x01) + + // Recompute H. + salt := x[xlen - hash_len - salt_len - 1:] + tmp: [hash.MAX_DIGEST_SIZE]byte + h := tmp[:hash_len] + ctx: hash.Context = --- + hash.init(&ctx, data_algo) + hash.update(&ctx, tmp[:8]) + hash.update(&ctx, digest) + hash.update(&ctx, salt[:salt_len]) + hash.final(&ctx, h) + + // Check that the recomputed H value matches the one appearing + // in the string. + x = x[xlen - hash_len - 1:] + r |= subtle.eq0(u32(crypto.compare_constant_time(h, x[:hash_len]))) + + return subtle.eq0(r) +} + +@(private="file", require_results) +pss_sig_pad :: proc( + data_algo: hash.Algorithm, + mgf1_algo: hash.Algorithm, + digest: []byte, + salt_len: int, + n_bitlen_: int, + sig: []byte, +) -> u32 { + x, n_bitlen := sig, n_bitlen_ + hash_len := hash.DIGEST_SIZES[data_algo] + + // The padded string is one bit smaller than the modulus; + // notably, if the modulus length is equal to 1 modulo 8, then + // the padded string will be one _byte_ smaller, and the first + // byte will be set to 0. We apply these transformations here. + n_bitlen -= 1 + if (n_bitlen & 7) == 0 { + x[0] = 0 + x = x[1:] + } + xlen := int((n_bitlen + 7) >> 3) + + // Check that the modulus is large enough for the hash value + // length combined with the intended salt length. + if hash_len > xlen || salt_len > xlen || (hash_len + salt_len + 2) > xlen { + return 0 + } + + // Produce a random salt. + salt := x[xlen - hash_len - salt_len - 1:] + salt = salt[:salt_len] + if salt_len != 0 { + crypto.rand_bytes(salt) + } + + // Compute the seed for MGF1. + seed := x[xlen - hash_len - 1:] + seed = seed[:hash_len] + ctx: hash.Context = --- + hash.init(&ctx, data_algo) + intrinsics.mem_zero(raw_data(seed), 8) + hash.update(&ctx, seed[:8]) + hash.update(&ctx, digest) + hash.update(&ctx, salt) + hash.final(&ctx, seed) + + // Prepare string PS (padded salt). The salt is already at the + // right place. + intrinsics.mem_zero(raw_data(x), xlen - salt_len - hash_len - 2) + x[xlen - salt_len - hash_len - 2] = 0x01 + + // Generate the mask and XOR it into PS. + mgf1_xor(x[:xlen - hash_len - 1], mgf1_algo, seed) + + // Clear the top bits to ensure the value is lower than the + // modulus. + x[0] &= 0xFF >> ((u32(xlen) << 3) - u32(n_bitlen)) + + // The seed (H) is already in the right place. We just set the + // last byte. + x[xlen - 1] = 0xBC + + return 1 +} diff --git a/core/crypto/rsa/rsa_test_key.odin b/core/crypto/rsa/rsa_test_key.odin new file mode 100644 index 000000000..28020d21f --- /dev/null +++ b/core/crypto/rsa/rsa_test_key.odin @@ -0,0 +1,180 @@ +package rsa + +// private_key_set_insecure_test sets the private key to the +// pregenerated INSECURE test key "testRSA2048" from RFC 9500 2.1. +// +// WARNING: This key MUST only be used for testing purposes. +@(require_results) +private_key_set_insecure_test :: proc(priv_key: ^Private_Key) -> bool { + // RFC 9500 2.1 "testRSA2048" + return private_key_set_bytes( + priv_key, + // n + []byte{ + 0xB0, 0xF9, 0xE8, 0x19, 0x43, 0xA7, 0xAE, 0x98, + 0x92, 0xAA, 0xDE, 0x17, 0xCA, 0x7C, 0x40, 0xF8, + 0x74, 0x4F, 0xED, 0x2F, 0x81, 0x48, 0xE6, 0xC8, + 0xEA, 0xA2, 0x7B, 0x7D, 0x00, 0x15, 0x48, 0xFB, + 0x51, 0x92, 0xAB, 0x28, 0xB5, 0x6C, 0x50, 0x60, + 0xB1, 0x18, 0xCC, 0xD1, 0x31, 0xE5, 0x94, 0x87, + 0x4C, 0x6C, 0xA9, 0x89, 0xB5, 0x6C, 0x27, 0x29, + 0x6F, 0x09, 0xFB, 0x93, 0xA0, 0x34, 0xDF, 0x32, + 0xE9, 0x7C, 0x6F, 0xF0, 0x99, 0x8C, 0xFD, 0x8E, + 0x6F, 0x42, 0xDD, 0xA5, 0x8A, 0xCD, 0x1F, 0xA9, + 0x79, 0x86, 0xF1, 0x44, 0xF3, 0xD1, 0x54, 0xD6, + 0x76, 0x50, 0x17, 0x5E, 0x68, 0x54, 0xB3, 0xA9, + 0x52, 0x00, 0x3B, 0xC0, 0x68, 0x87, 0xB8, 0x45, + 0x5A, 0xC2, 0xB1, 0x9F, 0x7B, 0x2F, 0x76, 0x50, + 0x4E, 0xBC, 0x98, 0xEC, 0x94, 0x55, 0x71, 0xB0, + 0x78, 0x92, 0x15, 0x0D, 0xDC, 0x6A, 0x74, 0xCA, + 0x0F, 0xBC, 0xD3, 0x54, 0x97, 0xCE, 0x81, 0x53, + 0x4D, 0xAF, 0x94, 0x18, 0x84, 0x4B, 0x13, 0xAE, + 0xA3, 0x1F, 0x9D, 0x5A, 0x6B, 0x95, 0x57, 0xBB, + 0xDF, 0x61, 0x9E, 0xFD, 0x4E, 0x88, 0x7F, 0x2D, + 0x42, 0xB8, 0xDD, 0x8B, 0xC9, 0x87, 0xEA, 0xE1, + 0xBF, 0x89, 0xCA, 0xB8, 0x5E, 0xE2, 0x1E, 0x35, + 0x63, 0x05, 0xDF, 0x6C, 0x07, 0xA8, 0x83, 0x8E, + 0x3E, 0xF4, 0x1C, 0x59, 0x5D, 0xCC, 0xE4, 0x3D, + 0xAF, 0xC4, 0x91, 0x23, 0xEF, 0x4D, 0x8A, 0xBB, + 0xA9, 0x3D, 0x39, 0x05, 0xE4, 0x02, 0x8D, 0x7B, + 0xA9, 0x14, 0x84, 0xA2, 0x75, 0x96, 0xE0, 0x7B, + 0x4B, 0x6E, 0xD9, 0x92, 0xF0, 0x77, 0xB5, 0x24, + 0xD3, 0xDC, 0xFE, 0x7D, 0xDD, 0x55, 0x49, 0xBE, + 0x7C, 0xCE, 0x8D, 0xA0, 0x35, 0xCF, 0xA0, 0xB3, + 0xFB, 0x8F, 0x9E, 0x46, 0xF7, 0x32, 0xB2, 0xA8, + 0x6B, 0x46, 0x01, 0x65, 0xC0, 0x8F, 0x53, 0x13, + }, + // e + []byte{0x01, 0x00, 0x01}, + // d + []byte{ + 0x41, 0x18, 0x8B, 0x20, 0xCF, 0xDB, 0xDB, 0xC2, + 0xCF, 0x1F, 0xFE, 0x75, 0x2D, 0xCB, 0xAA, 0x72, + 0x39, 0x06, 0x35, 0x2E, 0x26, 0x15, 0xD4, 0x9D, + 0xCE, 0x80, 0x59, 0x7F, 0xCF, 0x0A, 0x05, 0x40, + 0x3B, 0xEF, 0x00, 0xFA, 0x06, 0x51, 0x82, 0xF7, + 0x2D, 0xEC, 0xFB, 0x59, 0x6F, 0x4B, 0x0C, 0xE8, + 0xFF, 0x59, 0x70, 0xBA, 0xF0, 0x7A, 0x89, 0xA5, + 0x19, 0xEC, 0xC8, 0x16, 0xB2, 0xF4, 0xFF, 0xAC, + 0x50, 0x69, 0xAF, 0x1B, 0x06, 0xBF, 0xEF, 0x7B, + 0xF6, 0xBC, 0xD7, 0x9E, 0x4E, 0x81, 0xC8, 0xC5, + 0xA3, 0xA7, 0xD9, 0x13, 0x0D, 0xC3, 0xCF, 0xBA, + 0xDA, 0xE5, 0xF6, 0xD2, 0x88, 0xF9, 0xAE, 0xE3, + 0xF6, 0xFF, 0x92, 0xFA, 0xE0, 0xF8, 0x1A, 0xF5, + 0x97, 0xBE, 0xC9, 0x6A, 0xE9, 0xFA, 0xB9, 0x40, + 0x2C, 0xD5, 0xFE, 0x41, 0xF7, 0x05, 0xBE, 0xBD, + 0xB4, 0x7B, 0xB7, 0x36, 0xD3, 0xFE, 0x6C, 0x5A, + 0x51, 0xE0, 0xE2, 0x07, 0x32, 0xA9, 0x7B, 0x5E, + 0x46, 0xC1, 0xCB, 0xDB, 0x26, 0xD7, 0x48, 0x54, + 0xC6, 0xB6, 0x60, 0x4A, 0xED, 0x46, 0x37, 0x35, + 0xFF, 0x90, 0x76, 0x04, 0x65, 0x57, 0xCA, 0xF9, + 0x49, 0xBF, 0x44, 0x88, 0x95, 0xC2, 0x04, 0x32, + 0xC1, 0xE0, 0x9C, 0x01, 0x4E, 0xA7, 0x56, 0x60, + 0x43, 0x4F, 0x1A, 0x0F, 0x3B, 0xE2, 0x94, 0xBA, + 0xBC, 0x5D, 0x53, 0x0E, 0x6A, 0x10, 0x21, 0x3F, + 0x53, 0xB6, 0x03, 0x75, 0xFC, 0x84, 0xA7, 0x57, + 0x3F, 0x2A, 0xF1, 0x21, 0x55, 0x84, 0xF5, 0xB4, + 0xBD, 0xA6, 0xD4, 0xE8, 0xF9, 0xE1, 0x7A, 0x78, + 0xD9, 0x7E, 0x77, 0xB8, 0x6D, 0xA4, 0xA1, 0x84, + 0x64, 0x75, 0x31, 0x8A, 0x7A, 0x10, 0xA5, 0x61, + 0x01, 0x4E, 0xFF, 0xA2, 0x3A, 0x81, 0xEC, 0x56, + 0xE9, 0xE4, 0x10, 0x9D, 0xEF, 0x8C, 0xB3, 0xF7, + 0x97, 0x22, 0x3F, 0x7D, 0x8D, 0x0D, 0x43, 0x51, + }, + // p + []byte{ + 0xDD, 0x10, 0x57, 0x02, 0x38, 0x2F, 0x23, 0x2B, + 0x36, 0x81, 0xF5, 0x37, 0x91, 0xE2, 0x26, 0x17, + 0xC7, 0xBF, 0x4E, 0x9A, 0xCB, 0x81, 0xED, 0x48, + 0xDA, 0xF6, 0xD6, 0x99, 0x5D, 0xA3, 0xEA, 0xB6, + 0x42, 0x83, 0x9A, 0xFF, 0x01, 0x2D, 0x2E, 0xA6, + 0x28, 0xB9, 0x0A, 0xF2, 0x79, 0xFD, 0x3E, 0x6F, + 0x7C, 0x93, 0xCD, 0x80, 0xF0, 0x72, 0xF0, 0x1F, + 0xF2, 0x44, 0x3B, 0x3E, 0xE8, 0xF2, 0x4E, 0xD4, + 0x69, 0xA7, 0x96, 0x13, 0xA4, 0x1B, 0xD2, 0x40, + 0x20, 0xF9, 0x2F, 0xD1, 0x10, 0x59, 0xBD, 0x1D, + 0x0F, 0x30, 0x1B, 0x5B, 0xA7, 0xA9, 0xD3, 0x63, + 0x7C, 0xA8, 0xD6, 0x5C, 0x1A, 0x98, 0x15, 0x41, + 0x7D, 0x8E, 0xAB, 0x73, 0x4B, 0x0B, 0x4F, 0x3A, + 0x2C, 0x66, 0x1D, 0x9A, 0x1A, 0x82, 0xF3, 0xAC, + 0x73, 0x4C, 0x40, 0x53, 0x06, 0x69, 0xAB, 0x8E, + 0x47, 0x30, 0x45, 0xA5, 0x8E, 0x65, 0x53, 0x9D, + }, + // q + []byte{ + 0xCC, 0xF1, 0xE5, 0xBB, 0x90, 0xC8, 0xE9, 0x78, + 0x1E, 0xA7, 0x5B, 0xEB, 0xF1, 0x0B, 0xC2, 0x52, + 0xE1, 0x1E, 0xB0, 0x23, 0xA0, 0x26, 0x0F, 0x18, + 0x87, 0x55, 0x2A, 0x56, 0x86, 0x3F, 0x4A, 0x64, + 0x21, 0xE8, 0xC6, 0x00, 0xBF, 0x52, 0x3D, 0x6C, + 0xB1, 0xB0, 0xAD, 0xBD, 0xD6, 0x5B, 0xFE, 0xE4, + 0xA8, 0x8A, 0x03, 0x7E, 0x3D, 0x1A, 0x41, 0x5E, + 0x5B, 0xB9, 0x56, 0x48, 0xDA, 0x5A, 0x0C, 0xA2, + 0x6B, 0x54, 0xF4, 0xA6, 0x39, 0x48, 0x52, 0x2C, + 0x3D, 0x5F, 0x89, 0xB9, 0x4A, 0x72, 0xEF, 0xFF, + 0x95, 0x13, 0x4D, 0x59, 0x40, 0xCE, 0x45, 0x75, + 0x8F, 0x30, 0x89, 0x80, 0x90, 0x89, 0x56, 0x58, + 0x8E, 0xEF, 0x57, 0x5B, 0x3E, 0x4B, 0xC4, 0xC3, + 0x68, 0xCF, 0xE8, 0x13, 0xEE, 0x9C, 0x25, 0x2C, + 0x2B, 0x02, 0xE0, 0xDF, 0x91, 0xF1, 0xAA, 0x01, + 0x93, 0x8D, 0x38, 0x68, 0x5D, 0x60, 0xBA, 0x6F, + }, + // dp + []byte{ + 0x09, 0xED, 0x54, 0xEA, 0xED, 0x98, 0xF8, 0x4C, + 0x55, 0x7B, 0x4A, 0x86, 0xBF, 0x4F, 0x57, 0x84, + 0x93, 0xDC, 0xBC, 0x6B, 0xE9, 0x1D, 0xA1, 0x89, + 0x37, 0x04, 0x04, 0xA9, 0x08, 0x72, 0x76, 0xF4, + 0xCE, 0x51, 0xD8, 0xA1, 0x00, 0xED, 0x85, 0x7D, + 0xC2, 0xB0, 0x64, 0x94, 0x74, 0xF3, 0xF1, 0x5C, + 0xD2, 0x4C, 0x54, 0xDB, 0x28, 0x71, 0x10, 0xE5, + 0x6E, 0x5C, 0xB0, 0x08, 0x68, 0x2F, 0x91, 0x68, + 0xAA, 0x81, 0xF3, 0x14, 0x58, 0xB7, 0x43, 0x1E, + 0xCC, 0x1C, 0x44, 0x90, 0x6F, 0xDA, 0x87, 0xCA, + 0x89, 0x47, 0x10, 0xC3, 0x71, 0xE9, 0x07, 0x6C, + 0x1D, 0x49, 0xFB, 0xAE, 0x51, 0x27, 0x69, 0x34, + 0xF2, 0xAD, 0x78, 0x77, 0x89, 0xF4, 0x2D, 0x0F, + 0xA0, 0xB4, 0xC9, 0x39, 0x85, 0x5D, 0x42, 0x12, + 0x09, 0x6F, 0x70, 0x28, 0x0A, 0x4E, 0xAE, 0x7C, + 0x8A, 0x27, 0xD9, 0xC8, 0xD0, 0x77, 0x2E, 0x65, + }, + // dq + []byte{ + 0x8C, 0xB6, 0x85, 0x7A, 0x7B, 0xD5, 0x46, 0x5F, + 0x80, 0x04, 0x7E, 0x9B, 0x87, 0xBC, 0x00, 0x27, + 0x31, 0x84, 0x05, 0x81, 0xE0, 0x62, 0x61, 0x39, + 0x01, 0x2A, 0x5B, 0x50, 0x5F, 0x0A, 0x33, 0x84, + 0x7E, 0xB7, 0xB8, 0xC3, 0x28, 0x99, 0x49, 0xAD, + 0x48, 0x6F, 0x3B, 0x4B, 0x3D, 0x53, 0x9A, 0xB5, + 0xDA, 0x76, 0x30, 0x21, 0xCB, 0xC8, 0x2C, 0x1B, + 0xA2, 0x34, 0xA5, 0x66, 0x8D, 0xED, 0x08, 0x01, + 0xB8, 0x59, 0xF3, 0x43, 0xF1, 0xCE, 0x93, 0x04, + 0xE6, 0xFA, 0xA2, 0xB0, 0x02, 0xCA, 0xD9, 0xB7, + 0x8C, 0xDE, 0x5C, 0xDC, 0x2C, 0x1F, 0xB4, 0x17, + 0x1C, 0x42, 0x42, 0x16, 0x70, 0xA6, 0xAB, 0x0F, + 0x50, 0xCC, 0x4A, 0x19, 0x4E, 0xB3, 0x6D, 0x1C, + 0x91, 0xE9, 0x35, 0xBA, 0x01, 0xB9, 0x59, 0xD8, + 0x72, 0x8B, 0x9E, 0x64, 0x42, 0x6B, 0x3F, 0xC3, + 0xA7, 0x50, 0x6D, 0xEB, 0x52, 0x39, 0xA8, 0xA7, + }, + // iq (aka u) + []byte{ + 0x0A, 0x81, 0xD8, 0xA6, 0x18, 0x31, 0x4A, 0x80, + 0x3A, 0xF6, 0x1C, 0x06, 0x71, 0x1F, 0x2C, 0x39, + 0xB2, 0x66, 0xFF, 0x41, 0x4D, 0x53, 0x47, 0x6D, + 0x1D, 0xA5, 0x2A, 0x43, 0x18, 0xAA, 0xFE, 0x4B, + 0x96, 0xF0, 0xDA, 0x07, 0x15, 0x5F, 0x8A, 0x51, + 0x34, 0xDA, 0xB8, 0x8E, 0xE2, 0x9E, 0x81, 0x68, + 0x07, 0x6F, 0xCD, 0x78, 0xCA, 0x79, 0x1A, 0xC6, + 0x34, 0x42, 0xA8, 0x1C, 0xD0, 0x69, 0x39, 0x27, + 0xD8, 0x08, 0xE3, 0x35, 0xE8, 0xD8, 0xCB, 0xF2, + 0x12, 0x19, 0x07, 0x50, 0x9A, 0x57, 0x75, 0x9B, + 0x4F, 0x9A, 0x18, 0xFA, 0x3A, 0x7B, 0x33, 0x37, + 0x79, 0xED, 0xDE, 0x7A, 0x45, 0x93, 0x84, 0xF8, + 0x44, 0x4A, 0xDA, 0xEC, 0xFF, 0xEC, 0x95, 0xFD, + 0x55, 0x2B, 0x0C, 0xFC, 0xB6, 0xC7, 0xF6, 0x92, + 0x62, 0x6D, 0xDE, 0x1E, 0xF2, 0x68, 0xA4, 0x0D, + 0x2F, 0x67, 0xB5, 0xC8, 0xAA, 0x38, 0x7F, 0xF7, + }, + ) +} \ No newline at end of file diff --git a/core/crypto/turboshake/turboshake.odin b/core/crypto/turboshake/turboshake.odin new file mode 100644 index 000000000..9d2b88962 --- /dev/null +++ b/core/crypto/turboshake/turboshake.odin @@ -0,0 +1,57 @@ +/* +`TurboSHAKE` XOF algorithm family. + +The SHA3 hash algorithm can be found in the crypto/sha3. + +See: +- [[ https://www.rfc-editor.org/rfc/rfc9861 ]] +*/ +package turboshake + +import "../_sha3" + +// Context is a TurboSHAKE128 or TurboSHAKE256 instance. +Context :: distinct _sha3.Context + +// init_128 initializes a Context for TurboSHAKE128. +// Uses the default domain separation byte of 0x1f, unless one is supplied. +init_128 :: proc(ctx: ^Context, d: byte = _sha3.DS_TURBOSHAKE_DEFAULT) { + _sha3.init_turboshake((^_sha3.Context)(ctx), d, 128) +} + +// init_256 initializes a Context for TurboSHAKE256. +// Uses the default domain separation byte of 0x1f, unless one is supplied. +init_256 :: proc(ctx: ^Context, d: byte = _sha3.DS_TURBOSHAKE_DEFAULT) { + _sha3.init_turboshake((^_sha3.Context)(ctx), d, 256) +} + +// write writes more data into the TurboSHAKE instance. This MUST not be +// called after any reads have been done, and attempts to do so will panic. +write :: proc(ctx: ^Context, data: []byte) { + _sha3.update((^_sha3.Context)(ctx), data) +} + +// read reads output from the TurboSHAKE instance. There is no practical +// upper limit to the amount of data that can be read from TurboSHAKE. +// After read has been called one or more times, further calls to write +// will panic. +read :: proc(ctx: ^Context, dst: []byte) { + ctx_ := (^_sha3.Context)(ctx) + + if !ctx.is_finalized { + _sha3.shake_xof(ctx_) + } + + _sha3.shake_out(ctx_, dst) +} + +// clone clones the Context other into ctx. +clone :: proc(ctx, other: ^Context) { + _sha3.clone((^_sha3.Context)(ctx), (^_sha3.Context)(other)) +} + +// reset sanitizes the Context. The Context must be re-initialized to +// be used again. +reset :: proc(ctx: ^Context) { + _sha3.reset((^_sha3.Context)(ctx)) +} \ No newline at end of file diff --git a/core/crypto/x509/doc.odin b/core/crypto/x509/doc.odin new file mode 100644 index 000000000..3c0f9c95c --- /dev/null +++ b/core/crypto/x509/doc.odin @@ -0,0 +1,99 @@ +/* +X.509 v3 certificate parsing, signature verification, and chain (path) +validation. + +The parser is built on the strict DER reader in core:encoding/asn1 and +is zero-copy where possible: the returned Certificate's byte-slice +fields are views into the input DER, which must outlive it. The few +allocated fields (the extension/SAN tables) are released with +`destroy`. + +Input is DER. To parse a PEM certificate, decode it first with +core:encoding/pem (label "CERTIFICATE") and pass the resulting bytes. + +A successful `parse` means the bytes were well-formed, NOT that the +certificate is valid or trusted. The Certificate carries everything +the verifier needs: `raw_tbs` (the exact byte range a signature covers), +`raw_spki` (the range hashed for tls-server-end-point channel binding, +RFC 5929, and SPKI pinning), and `raw_issuer`/`raw_subject` (for the RFC +5280 binary-comparison rule). + +Hostname verification (`verify_hostname`) implements the RFC 6125 +subset modern clients use: subject alternative names only (no +CommonName fallback), with at most one wildcard as the entire +left-most label. + +Trust is established by `verify_chain`, which builds a path from a leaf +to a supplied trust anchor through supplied intermediates and checks, +for each certificate, validity, signature, name chaining, and the CA / +keyCertSign / pathLenConstraint rules; `verify_signature` exposes the +single-edge signature check on its own. + +LIMITATIONS: + + - Signature verification covers RSA PKCS#1 v1.5 and RSA-PSS + (SHA-256/384/512), ECDSA P-256/P-384, and Ed25519. These paths return + .Unsupported_Algorithm: SHA-1 (deprecated and rejected, RFC 9155), ECDSA + P-521 (effectively dead in web PKI), and RSA-PSS naming a digest or MGF + this package does not recognize. + - Name constraints (RFC 5280 4.2.1.10) are enforced for the dNSName and + iPAddress forms: a CA's permitted/excluded subtrees are checked against + every subordinate certificate's SANs, regardless of the extension's + criticality. A NameConstraints that uses any other base form + (directoryName, rfc822Name, URI, otherName), a minimum/maximum, or that + is malformed cannot be fully evaluated, so the whole chain is rejected + (fail closed) rather than accepted unchecked. NOT enforced: the RFC 5280 + rule that the extension be critical, and dNSName syntax validation (a + leading-period constraint is accepted, as OpenSSL does). + - REVOCATION IS NOT CHECKED. verify_chain performs NO CRL or OCSP + revocation checking. Callers that need revocation (e.g. TLS clients) + MUST supply it separately (OCSP stapling, CRLite, …). + - Certificate policies / policy constraints are not evaluated, and + there is no Public Suffix List: a (CABF-forbidden) wildcard such as + "*.com" would match "host.com". As a backstop, verify_chain still + fails closed on any uninterpreted CRITICAL extension + (.Unhandled_Critical_Extension). + - EKU is checked only when opts.required_eku is set, and then by RFC + 5280 semantics: a certificate with no EKU extension is unrestricted + (it is not required to assert the purpose); a certificate that DOES + assert EKU must include the purpose, enforced across the leaf and + every intermediate (EKU nesting). Leaf KeyUsage is not checked + against the intended protocol use. + +Parsing is deliberately lenient wherever strictness is a validation +concern rather than a structural one. Exception: Parser rejects +duplicate extension OIDs (Duplicate_Extension, RFC 5280 section 4.2). + + - Only dNSName and iPAddress subject alternative names are decoded + (into `dns_names` / `ip_addresses`). Other GeneralName forms (URI, + rfc822Name, directoryName, otherName) are skipped; the raw SAN + extension is still available via `extensions`. + - Only the extensions path validation needs are decoded + (BasicConstraints, KeyUsage, ExtKeyUsage, SubjectAltName, + Subject/Authority Key Identifier; NameConstraints is decoded at + verification time). All others (AIA, CRL distribution points, + certificate policies, …) are left raw in `extensions`. + - Subject and issuer are kept as raw DER (`raw_subject` / `raw_issuer`), + which is what name chaining compares (the RFC 5280 binary rule). The + attributes (CN, O, …) are decoded on demand by `parse_dn`, not at parse + time; `dn_get` / `dn_string` read them out and `serial_string` formats + the serial. + - Unsupported public-key curves yield Public_Key_Algorithm.Unknown. + - Non-conformant-but-extractable values are preserved: negative or + over-long serials, and validity dates far in the future. Validity + is stored as core:time.Time, which tops out near year 2262; dates + beyond that (the RFC 5280 "99991231235959Z" no-expiration sentinel) + saturate to that bound at parse time rather than failing, so they + read as "effectively never expires". + - Per-extension criticality rules (e.g. that subjectKeyIdentifier be + non-critical) are left to the caller via `Extension.critical`, and + a critical extension this package does not understand sets + `unhandled_critical` rather than failing the parse. + + +See: +- [[ https://www.rfc-editor.org/rfc/rfc5280 ]] +- [[ https://www.rfc-editor.org/rfc/rfc6125 ]] +- [[ https://www.rfc-editor.org/rfc/rfc5929 ]] +*/ +package x509 diff --git a/core/crypto/x509/ext.odin b/core/crypto/x509/ext.odin new file mode 100644 index 000000000..fbff6d88b --- /dev/null +++ b/core/crypto/x509/ext.odin @@ -0,0 +1,160 @@ +package x509 + +import "core:encoding/asn1" + +// Extension encoders. Each returns a complete DER Extension +// Extension ::= SEQUENCE { extnID OID, critical BOOLEAN DEFAULT FALSE, extnValue OCTET STRING } +// with extnValue carrying the encoded value (the inverse of the matching +// _parse_known_extension branch). The returned slice is allocated; +// a certificate builder splices these into the extensions [3] list via +// asn1.raw. critical=false is omitted per DEFAULT FALSE. + +// marshal_ext_basic_constraints encodes id-ce-basicConstraints: +// BasicConstraints ::= SEQUENCE { cA BOOLEAN DEFAULT FALSE, pathLenConstraint INTEGER OPTIONAL } +// cA is emitted only when `is_ca`; pathLenConstraint only when is_ca and +// `max_path_len` >= 0 (a negative value means "absent"). +@(require_results) +marshal_ext_basic_constraints :: proc(is_ca: bool, max_path_len: int, critical: bool, allocator := context.allocator) -> (der: []byte, err: Error) { + vals: [2]asn1.Value + path_buf: [8]byte + n := 0 + if is_ca { + vals[n] = asn1.boolean(true) + n += 1 + if max_path_len >= 0 { + vals[n] = asn1.integer_unsigned(_uint_be(max_path_len, path_buf[:])) + n += 1 + } + } + return _marshal_extension(_OID_EXT_BASIC_CONSTRAINTS, critical, asn1.sequence(vals[:n]), allocator) +} + +// marshal_ext_key_usage encodes id-ce-keyUsage as a KeyUsage BIT STRING, +// minimally (trailing zero bits dropped, per DER named bit strings). +@(require_results) +marshal_ext_key_usage :: proc(usage: Key_Usage, critical: bool, allocator := context.allocator) -> (der: []byte, err: Error) { + hi := -1 + for bit in Key_Usage_Bit { + if bit in usage { + hi = int(bit) // bits enumerate ascending, so the last present is the highest + } + } + buf: [3]byte // unused-bits octet + up to 2 payload octets (9 named bits) + content: []byte + if hi < 0 { + content = buf[:1] // no bits set: empty bit string, 0 unused + } else { + nbytes := hi / 8 + 1 + buf[0] = byte(nbytes * 8 - (hi + 1)) // unused-bits count + for bit in Key_Usage_Bit { + i := int(bit) + if i <= hi && bit in usage { + buf[1 + i / 8] |= 0x80 >> uint(i % 8) + } + } + content = buf[:1 + nbytes] + } + return _marshal_extension(_OID_EXT_KEY_USAGE, critical, asn1.primitive(asn1.universal(.Bit_String), content), allocator) +} + +// marshal_ext_ext_key_usage encodes id-ce-extKeyUsage: +// ExtKeyUsageSyntax ::= SEQUENCE SIZE (1..MAX) OF KeyPurposeId +// in EKU_Bit order. +@(require_results) +marshal_ext_ext_key_usage :: proc(eku: Ext_Key_Usage, critical: bool, allocator := context.allocator) -> (der: []byte, err: Error) { + purposes: [7]asn1.Value + n := 0 + for bit in EKU_Bit { + if bit in eku { + purposes[n] = asn1.object_identifier(_eku_oid(bit)) + n += 1 + } + } + return _marshal_extension(_OID_EXT_EXT_KEY_USAGE, critical, asn1.sequence(purposes[:n]), allocator) +} + +// marshal_ext_san encodes id-ce-subjectAltName: +// GeneralNames ::= SEQUENCE OF GeneralName +// emitting dNSName [2] IA5String entries (in order) followed by iPAddress [7] +// OCTET STRING entries. IP values are the raw 4- or 16-octet address. +@(require_results) +marshal_ext_san :: proc(dns_names: []string, ip_addresses: [][]byte, critical: bool, allocator := context.allocator) -> (der: []byte, err: Error) { + n := len(dns_names) + len(ip_addresses) + names, merr := make([]asn1.Value, n, allocator) // dynamic count: scaffolding, freed below + if merr != nil { + return nil, .Allocation_Failed + } + defer delete(names, allocator) + + i := 0 + for d in dns_names { + names[i] = asn1.context_primitive(2, transmute([]byte)d) // [2] IMPLICIT IA5String + i += 1 + } + for ip in ip_addresses { + names[i] = asn1.context_primitive(7, ip) // [7] IMPLICIT OCTET STRING + i += 1 + } + return _marshal_extension(_OID_EXT_SAN, critical, asn1.sequence(names[:]), allocator) +} + +// _marshal_extension wraps an extension value tree as a complete Extension. +// `value`'s borrowed backing lives in the caller's frame, which is active for +// the duration of this synchronous call, so the encode below sees it intact. +@(private, require_results) +_marshal_extension :: proc(oid: []byte, critical: bool, value: asn1.Value, allocator := context.allocator) -> (der: []byte, err: Error) { + out: []byte + merr: asn1.Error + if critical { + out, merr = asn1.marshal( + asn1.sequence({asn1.object_identifier(oid), asn1.boolean(true), asn1.octet_string_wrap({value})}), + allocator, + ) + } else { + out, merr = asn1.marshal(asn1.sequence({asn1.object_identifier(oid), asn1.octet_string_wrap({value})}), allocator) + } + if merr != .None { + return nil, .Allocation_Failed + } + return out, .None +} + +// _uint_be writes the minimal big-endian magnitude of a non-negative int into +// buf and returns the slice (empty for zero, which integer_unsigned encodes as 0). +@(private) +_uint_be :: proc(v: int, buf: []byte) -> []byte { + if v <= 0 { + return buf[:0] + } + n := 0 + x := v + for x > 0 { + n += 1 + x >>= 8 + } + for i in 0 ..< n { + buf[n - 1 - i] = byte(v >> uint(8 * i)) + } + return buf[:n] +} + +@(private) +_eku_oid :: proc(bit: EKU_Bit) -> []byte { + switch bit { + case .Server_Auth: + return _OID_EKU_SERVER_AUTH + case .Client_Auth: + return _OID_EKU_CLIENT_AUTH + case .Code_Signing: + return _OID_EKU_CODE_SIGNING + case .Email_Protection: + return _OID_EKU_EMAIL_PROTECTION + case .Time_Stamping: + return _OID_EKU_TIME_STAMPING + case .OCSP_Signing: + return _OID_EKU_OCSP_SIGNING + case .Any: + return _OID_EKU_ANY + } + return nil +} diff --git a/core/crypto/x509/marshal.odin b/core/crypto/x509/marshal.odin new file mode 100644 index 000000000..dd24da01d --- /dev/null +++ b/core/crypto/x509/marshal.odin @@ -0,0 +1,438 @@ +package x509 + +import "core:encoding/asn1" +import "core:time" + +// Distinguished-name attribute types the DN builder knows by name. `Other` +// carries any attribute outside this set via DN_Attribute.oid (used by +// parse_dn for attributes it does not recognize by name). +DN_Attribute_Type :: enum { + Common_Name, // CN, 2.5.4.3 + Country, // C, 2.5.4.6 + Locality, // L, 2.5.4.7 + State_Or_Province, // ST, 2.5.4.8 + Organization, // O, 2.5.4.10 + Organizational_Unit, // OU, 2.5.4.11 + Serial_Number, // 2.5.4.5 + Other, // any other attribute; its type OID is in `oid` +} + +// DN_Attribute is one relative distinguished name: a type and its value. When +// `type` is `Other`, `oid` holds the raw attribute-type OID content octets. +// Values are emitted as UTF8String, except Country and Serial_Number which +// are PrintableString (X.520), the policy RFC 5280 section 4.1.2.4 advises. +DN_Attribute :: struct { + type: DN_Attribute_Type, + value: string, + oid: []byte, // meaningful only when type == Other +} + +@(rodata, private) +_OID_AT_CN := []byte{0x55, 0x04, 0x03} +@(rodata, private) +_OID_AT_C := []byte{0x55, 0x04, 0x06} +@(rodata, private) +_OID_AT_L := []byte{0x55, 0x04, 0x07} +@(rodata, private) +_OID_AT_ST := []byte{0x55, 0x04, 0x08} +@(rodata, private) +_OID_AT_O := []byte{0x55, 0x04, 0x0A} +@(rodata, private) +_OID_AT_OU := []byte{0x55, 0x04, 0x0B} +@(rodata, private) +_OID_AT_SERIAL := []byte{0x55, 0x04, 0x05} + +// marshal_dn encodes `attrs` as a DER Name (RDNSequence): one single-valued +// RelativeDistinguishedName per attribute, in the given order. The returned +// slice is the caller's to free; the attribute value bytes are copied into +// it, so `attrs` need not outlive the call. +@(require_results) +marshal_dn :: proc(attrs: []DN_Attribute, allocator := context.allocator) -> (der: []byte, err: Error) { + n := len(attrs) + // Value-tree scaffolding, freed once the bytes are produced. Pre-sized so + // the sub-slices handed to set()/sequence() never move before the encode. + rdns, e1 := make([]asn1.Value, n, allocator) // one SET per attribute + if e1 != nil { + return nil, .Allocation_Failed + } + defer delete(rdns, allocator) + atvs, e2 := make([]asn1.Value, n, allocator) // the AttributeTypeAndValue SEQUENCE + if e2 != nil { + return nil, .Allocation_Failed + } + defer delete(atvs, allocator) + pairs, e3 := make([]asn1.Value, 2 * n, allocator) // {type OID, value} per attribute + if e3 != nil { + return nil, .Allocation_Failed + } + defer delete(pairs, allocator) + + for a, i in attrs { + pairs[2 * i] = asn1.object_identifier(_dn_oid(a)) + pairs[2 * i + 1] = asn1.primitive(_dn_string_tag(a.type), transmute([]byte)a.value) + atvs[i] = asn1.sequence(pairs[2 * i:2 * i + 2]) + rdns[i] = asn1.set(atvs[i:i + 1]) + } + + out, merr := asn1.marshal(asn1.sequence(rdns), allocator) + if merr != .None { + return nil, .Allocation_Failed + } + return out, .None +} + +@(private) +_dn_oid :: proc(a: DN_Attribute) -> []byte { + switch a.type { + case .Common_Name: + return _OID_AT_CN + case .Country: + return _OID_AT_C + case .Locality: + return _OID_AT_L + case .State_Or_Province: + return _OID_AT_ST + case .Organization: + return _OID_AT_O + case .Organizational_Unit: + return _OID_AT_OU + case .Serial_Number: + return _OID_AT_SERIAL + case .Other: + return a.oid + } + return nil +} + +@(private) +_dn_string_tag :: proc(type: DN_Attribute_Type) -> asn1.Tag { + #partial switch type { + case .Country, .Serial_Number: + return asn1.universal(.Printable_String) + } + return asn1.universal(.UTF8_String) +} + +@(rodata, private) +_DER_INT_ZERO := []byte{0x00} + +// marshal_csr_info encodes the CertificationRequestInfo, the to-be-signed +// portion of a PKCS#10 CSR (RFC 2986): version v1, the subject DN, the subject +// public key, and the attributes set. Sign the returned bytes and pass them +// with the signature to marshal_csr. The slice is the caller's to free; +// `subject`, `key`, and `extensions` need not outlive the call. +// +// `extensions` is a list of pre-encoded Extension DER (from the marshal_ext_* +// helpers); when non-empty it is requested via a single PKCS#9 +// extensionRequest attribute (the standard way a CSR asks the CA to place +// extensions, SANs, key usage in the issued certificate). Pass nil for the +// empty attributes set. +@(require_results) +marshal_csr_info :: proc(subject: []DN_Attribute, key: Public_Key, extensions: [][]byte = nil, allocator := context.allocator) -> (cri_der: []byte, err: Error) { + dn := marshal_dn(subject, allocator) or_return + defer delete(dn, allocator) + spki := marshal_spki(key, allocator) or_return + defer delete(spki, allocator) + attrs := _marshal_csr_attributes(extensions, allocator) or_return + defer delete(attrs, allocator) + + // CertificationRequestInfo ::= SEQUENCE { version, subject, subjectPKInfo, [0] attributes } + out, merr := asn1.marshal( + asn1.sequence( + { + asn1.integer_raw(_DER_INT_ZERO), // version v1 (0) + asn1.raw(dn), // subject Name + asn1.raw(spki), // subjectPublicKeyInfo + asn1.raw(attrs), // attributes [0] + }, + ), + allocator, + ) + if merr != .None { + return nil, .Allocation_Failed + } + return out, .None +} + +// _marshal_csr_attributes encodes the CertificationRequestInfo attributes +// field, [0] IMPLICIT SET OF Attribute. With no extensions it is the empty +// set (A0 00); otherwise it carries a single PKCS#9 extensionRequest attribute +// Attribute ::= SEQUENCE { type extensionRequest, values SET { Extensions } } +// whose value is the Extensions SEQUENCE OF Extension built from `extensions`. +@(private, require_results) +_marshal_csr_attributes :: proc(extensions: [][]byte, allocator := context.allocator) -> (der: []byte, err: Error) { + out: []byte + merr: asn1.Error + if len(extensions) == 0 { + out, merr = asn1.marshal(asn1.context_explicit(0, {}), allocator) + } else { + // Pre-sized scaffolding for the Extensions SEQUENCE OF, kept alive + // (never resized) through the marshal below. + ext_raws, e := make([]asn1.Value, len(extensions), allocator) + if e != nil { + return nil, .Allocation_Failed + } + defer delete(ext_raws, allocator) + for ext, i in extensions { + ext_raws[i] = asn1.raw(ext) + } + out, merr = asn1.marshal( + asn1.context_explicit( + 0, + {asn1.sequence({asn1.object_identifier(_OID_EXT_REQUEST), asn1.set({asn1.sequence(ext_raws[:])})})}, + ), + allocator, + ) + } + if merr != .None { + return nil, .Allocation_Failed + } + return out, .None +} + +// marshal_csr wraps a (separately signed) CertificationRequestInfo into a +// complete PKCS#10 CertificationRequest. `signature` is the raw signature +// value over `cri_der`, a DER ECDSA-Sig-Value for ECDSA, the 64-byte value +// for Ed25519, and `signature_algorithm` selects the matching +// AlgorithmIdentifier. The slice is the caller's to free. +@(require_results) +marshal_csr :: proc(cri_der: []byte, signature_algorithm: Signature_Algorithm, signature: []byte, allocator := context.allocator) -> (csr_der: []byte, err: Error) { + // CertificationRequest ::= SEQUENCE { CRI, signatureAlgorithm, signature BIT STRING } + return _marshal_signed(cri_der, signature_algorithm, signature, allocator) +} + +// _marshal_signed wraps an already-encoded body (a CertificationRequestInfo +// or a TBSCertificate) with its signature algorithm and signature BIT STRING: +// SEQUENCE { body, AlgorithmIdentifier, BIT STRING signature } +// the common shape of PKCS#10 CSRs and X.509 certificates. +@(private, require_results) +_marshal_signed :: proc(body_der: []byte, signature_algorithm: Signature_Algorithm, signature: []byte, allocator := context.allocator) -> (der: []byte, err: Error) { + oid, null_params, ok := _sig_alg_identifier(signature_algorithm) + if !ok { + return nil, .Unsupported_Algorithm + } + out: []byte + merr: asn1.Error + if null_params { + out, merr = asn1.marshal( + asn1.sequence( + { + asn1.raw(body_der), + asn1.sequence({asn1.object_identifier(oid), asn1.null()}), + asn1.bit_string_octets(signature), + }, + ), + allocator, + ) + } else { + out, merr = asn1.marshal( + asn1.sequence( + {asn1.raw(body_der), asn1.sequence({asn1.object_identifier(oid)}), asn1.bit_string_octets(signature)}, + ), + allocator, + ) + } + if merr != .None { + return nil, .Allocation_Failed + } + return out, .None +} + +// _marshal_alg_id encodes a standalone AlgorithmIdentifier (NULL params for +// RSA PKCS#1, absent for ECDSA/EdDSA) for embedding in a TBSCertificate. +@(private, require_results) +_marshal_alg_id :: proc(signature_algorithm: Signature_Algorithm, allocator := context.allocator) -> (der: []byte, err: Error) { + oid, null_params, ok := _sig_alg_identifier(signature_algorithm) + if !ok { + return nil, .Unsupported_Algorithm + } + out: []byte + merr: asn1.Error + if null_params { + out, merr = asn1.marshal(asn1.sequence({asn1.object_identifier(oid), asn1.null()}), allocator) + } else { + out, merr = asn1.marshal(asn1.sequence({asn1.object_identifier(oid)}), allocator) + } + if merr != .None { + return nil, .Allocation_Failed + } + return out, .None +} + +// Maps a signature algorithm to its AlgorithmIdentifier OID and whether the +// parameters field is an explicit NULL (RSA PKCS#1) or absent (ECDSA, EdDSA). +@(private) +_sig_alg_identifier :: proc(alg: Signature_Algorithm) -> (oid: []byte, null_params: bool, ok: bool) { + #partial switch alg { + case .RSA_SHA256: + return _OID_SIG_RSA_SHA256, true, true + case .RSA_SHA384: + return _OID_SIG_RSA_SHA384, true, true + case .RSA_SHA512: + return _OID_SIG_RSA_SHA512, true, true + case .ECDSA_SHA256: + return _OID_SIG_ECDSA_SHA256, false, true + case .ECDSA_SHA384: + return _OID_SIG_ECDSA_SHA384, false, true + case .ECDSA_SHA512: + return _OID_SIG_ECDSA_SHA512, false, true + case .Ed25519: + return _OID_ED25519, false, true + } + return nil, false, false +} + +@(rodata, private) +_DER_INT_V3 := []byte{0x02} // Version v3 (value 2) + +// _marshal_extensions_field encodes the TBSCertificate extensions field: +// [3] EXPLICIT Extensions, Extensions ::= SEQUENCE OF Extension +// from a list of pre-encoded Extension DER, or nil when there are none (the +// field is OPTIONAL, so an empty list omits it entirely). +@(private, require_results) +_marshal_extensions_field :: proc(extensions: [][]byte, allocator := context.allocator) -> (der: []byte, err: Error) { + if len(extensions) == 0 { + return nil, .None + } + raws, merr := make([]asn1.Value, len(extensions), allocator) + if merr != nil { + return nil, .Allocation_Failed + } + defer delete(raws, allocator) + for e, i in extensions { + raws[i] = asn1.raw(e) + } + out, ferr := asn1.marshal(asn1.context_explicit(3, {asn1.sequence(raws[:])}), allocator) + if ferr != .None { + return nil, .Allocation_Failed + } + return out, .None +} + +// TBS_Certificate gathers the fields of a TBSCertificate to encode. `issuer` +// and `subject` are RDNSequences; `serial` is the serialNumber's unsigned +// magnitude; `extensions` is a list of pre-encoded Extension DER (from the +// marshal_ext_* helpers), embedded in order. +TBS_Certificate :: struct { + serial: []byte, + signature_algorithm: Signature_Algorithm, + issuer: []DN_Attribute, + not_before: time.Time, + not_after: time.Time, + subject: []DN_Attribute, + public_key: Public_Key, + extensions: [][]byte, +} + +// marshal_tbs_certificate encodes a v3 TBSCertificate (RFC 5280 section 4.1), +// the to-be-signed portion of a certificate. Sign the returned bytes and pass +// them with the signature to marshal_certificate. The slice is the caller's to +// free; the inputs need not outlive the call. +@(require_results) +marshal_tbs_certificate :: proc(tbs: TBS_Certificate, allocator := context.allocator) -> (der: []byte, err: Error) { + issuer_dn := marshal_dn(tbs.issuer, allocator) or_return + defer delete(issuer_dn, allocator) + subject_dn := marshal_dn(tbs.subject, allocator) or_return + defer delete(subject_dn, allocator) + spki := marshal_spki(tbs.public_key, allocator) or_return + defer delete(spki, allocator) + sig_alg := _marshal_alg_id(tbs.signature_algorithm, allocator) or_return + defer delete(sig_alg, allocator) + ext_field := _marshal_extensions_field(tbs.extensions, allocator) or_return + defer delete(ext_field, allocator) + + // raw() splices the independently-marshalled pieces in place; raw(nil) for + // an absent extensions field contributes nothing. + out, merr := asn1.marshal( + asn1.sequence( + { + asn1.context_explicit(0, {asn1.integer_raw(_DER_INT_V3)}), // version [0] EXPLICIT v3 + asn1.integer_unsigned(tbs.serial), // serialNumber + asn1.raw(sig_alg), // signature AlgorithmIdentifier + asn1.raw(issuer_dn), // issuer + asn1.sequence({asn1.time(tbs.not_before), asn1.time(tbs.not_after)}), // validity + asn1.raw(subject_dn), // subject + asn1.raw(spki), // subjectPublicKeyInfo + asn1.raw(ext_field), // extensions [3] (absent when empty) + }, + ), + allocator, + ) + if merr != .None { + return nil, .Allocation_Failed + } + return out, .None +} + +// marshal_certificate wraps a (separately signed) TBSCertificate into a +// complete X.509 Certificate. `signature_algorithm` must match the one inside +// the TBS (RFC 5280 section 4.1.1.2). The slice is the caller's to free. +@(require_results) +marshal_certificate :: proc(tbs_der: []byte, signature_algorithm: Signature_Algorithm, signature: []byte, allocator := context.allocator) -> (der: []byte, err: Error) { + return _marshal_signed(tbs_der, signature_algorithm, signature, allocator) +} + +// Public_Key holds the subject public-key material to encode into a +// SubjectPublicKeyInfo, mirroring the fields parse() extracts onto a +// Certificate: rsa_n/rsa_e (unsigned magnitudes) for RSA; ec_point for +// ECDSA (the uncompressed point 0x04||X||Y) and Ed25519 (the 32-byte key). +Public_Key :: struct { + algorithm: Public_Key_Algorithm, + rsa_n: []byte, + rsa_e: []byte, + ec_point: []byte, +} + +// marshal_spki encodes `key` as a DER SubjectPublicKeyInfo, the inverse of +// the SPKI decoding in parse(); the returned slice is the caller's to free. +// Unknown / unsupported key algorithms yield .Unsupported_Algorithm. +@(require_results) +marshal_spki :: proc(key: Public_Key, allocator := context.allocator) -> (der: []byte, err: Error) { + out: []byte + merr: asn1.Error + switch key.algorithm { + case .RSA: + // SEQUENCE { SEQUENCE { OID rsaEncryption, NULL }, BIT STRING { RSAPublicKey } } + out, merr = asn1.marshal( + asn1.sequence( + { + asn1.sequence({asn1.object_identifier(_OID_KEY_RSA), asn1.null()}), + asn1.bit_string_wrap({asn1.sequence({asn1.integer_unsigned(key.rsa_n), asn1.integer_unsigned(key.rsa_e)})}), + }, + ), + allocator, + ) + case .ECDSA_P256, .ECDSA_P384, .ECDSA_P521: + curve_oid: []byte + #partial switch key.algorithm { + case .ECDSA_P256: + curve_oid = _OID_CURVE_P256 + case .ECDSA_P384: + curve_oid = _OID_CURVE_P384 + case .ECDSA_P521: + curve_oid = _OID_CURVE_P521 + } + // SEQUENCE { SEQUENCE { OID ecPublicKey, OID namedCurve }, BIT STRING point } + out, merr = asn1.marshal( + asn1.sequence( + { + asn1.sequence({asn1.object_identifier(_OID_KEY_EC), asn1.object_identifier(curve_oid)}), + asn1.bit_string_octets(key.ec_point), + }, + ), + allocator, + ) + case .Ed25519: + // SEQUENCE { SEQUENCE { OID Ed25519 }, BIT STRING key } (no params, RFC 8410) + out, merr = asn1.marshal( + asn1.sequence({asn1.sequence({asn1.object_identifier(_OID_ED25519)}), asn1.bit_string_octets(key.ec_point)}), + allocator, + ) + case .Unknown: + return nil, .Unsupported_Algorithm + } + if merr != .None { + return nil, .Allocation_Failed + } + return out, .None +} diff --git a/core/crypto/x509/name.odin b/core/crypto/x509/name.odin new file mode 100644 index 000000000..a8238e2ac --- /dev/null +++ b/core/crypto/x509/name.odin @@ -0,0 +1,198 @@ +package x509 + +import "core:bytes" +import "core:encoding/asn1" +import "core:strings" + +// Consumer-facing helpers for reading the primitives `parse` leaves as raw DER: +// distinguished-name (Name) decoding, plus a serial-number display formatter. +// These are conveniences over the raw fields, path validation does not need +// them, and `raw_subject` / `raw_issuer` remain available for the RFC 5280 +// binary-comparison rule (which is what issuer/subject matching uses). + +// parse_dn decodes a DER Name (an RDNSequence, e.g. cert.raw_subject or +// cert.raw_issuer) into its attributes, the inverse of marshal_dn. Attributes +// outside the recognized set (see DN_Attribute_Type) come back as `Other` with +// their type OID in `oid`. The returned slice is the caller's to free; every +// `value` (and `oid`) is a VIEW into `der`, which must outlive the result. +// +// Values are taken as the raw content octets of the attribute value: correct +// for the UTF8String / PrintableString / IA5String forms certificates use in +// practice, but Teletex/BMP/Universal strings are NOT transcoded (returned as +// their raw bytes). +@(require_results) +parse_dn :: proc(der: []byte, allocator := context.allocator) -> (attrs: []DN_Attribute, err: Error) { + cur: asn1.Cursor + asn1.cursor_init(&cur, der) + seq, e := asn1.read_sequence(&cur) + if e != .None || asn1.done(&cur) != .None { + return nil, .Malformed + } + + // Count AttributeTypeAndValue entries first (an RDN is a SET OF, usually one + // element), so the result is exact-sized. + count := 0 + { + tmp := seq + for !asn1.is_empty(&tmp) { + rdn, re := asn1.read_set(&tmp) + if re != .None { + return nil, .Malformed + } + for !asn1.is_empty(&rdn) { + if _, ae := asn1.read_sequence(&rdn); ae != .None { + return nil, .Malformed + } + count += 1 + } + } + } + if count == 0 { + return nil, .None // an empty Name is valid + } + + out, merr := make([]DN_Attribute, count, allocator) + if merr != nil { + return nil, .Allocation_Failed + } + i := 0 + for !asn1.is_empty(&seq) { + rdn, re := asn1.read_set(&seq) + if re != .None { + delete(out, allocator) + return nil, .Malformed + } + for !asn1.is_empty(&rdn) { + atv, ae := asn1.read_sequence(&rdn) + oid, oe := asn1.read_oid(&atv) + _, val, ve := asn1.read_any(&atv) + if ae != .None || oe != .None || ve != .None || asn1.done(&atv) != .None { + delete(out, allocator) + return nil, .Malformed + } + t, known := _dn_type_from_oid(oid) + out[i] = DN_Attribute { + type = t, + value = string(val), + } + if !known { + out[i].oid = oid + } + i += 1 + } + } + return out[:i], .None +} + +// dn_get returns the value of the first attribute of `type` (e.g. .Common_Name), +// and whether one was present. +dn_get :: proc(attrs: []DN_Attribute, type: DN_Attribute_Type) -> (value: string, ok: bool) { + for a in attrs { + if a.type == type { + return a.value, true + } + } + return "", false +} + +// dn_string renders `attrs` as an RFC 4514 string ("CN=leaf,O=Acme,C=US"): +// RDNs in reverse order, short names for the recognized attributes and the +// dotted OID for `Other`, with RFC 4514 section 2.4 special characters escaped. +// The returned string is the caller's to free. +@(require_results) +dn_string :: proc(attrs: []DN_Attribute, allocator := context.allocator) -> string { + b := strings.builder_make(allocator) + for i := len(attrs) - 1; i >= 0; i -= 1 { + a := attrs[i] + if i != len(attrs) - 1 { + strings.write_byte(&b, ',') + } + switch a.type { + case .Common_Name: + strings.write_string(&b, "CN") + case .Country: + strings.write_string(&b, "C") + case .Locality: + strings.write_string(&b, "L") + case .State_Or_Province: + strings.write_string(&b, "ST") + case .Organization: + strings.write_string(&b, "O") + case .Organizational_Unit: + strings.write_string(&b, "OU") + case .Serial_Number: + strings.write_string(&b, "serialNumber") + case .Other: + if s, e := asn1.oid_to_string(a.oid, context.temp_allocator); e == .None { + strings.write_string(&b, s) + } else { + strings.write_byte(&b, '?') + } + } + strings.write_byte(&b, '=') + _dn_escape(&b, a.value) + } + return strings.to_string(b) +} + +// serial_string formats the certificate serial as upper-case colon-separated +// hex ("07:44:76:…"). The serial is an opaque identifier (up to 20 octets) +// Allocated String +@(require_results) +serial_string :: proc(cert: ^Certificate, allocator := context.allocator) -> string { + HEX := "0123456789ABCDEF" + b := strings.builder_make(allocator) + for octet, i in cert.serial { + if i != 0 { + strings.write_byte(&b, ':') + } + strings.write_byte(&b, HEX[octet >> 4]) + strings.write_byte(&b, HEX[octet & 0x0F]) + } + return strings.to_string(b) +} + +// _dn_type_from_oid maps an attribute-type OID to a DN_Attribute_Type, the +// inverse of _dn_oid +@(private) +_dn_type_from_oid :: proc(oid: []byte) -> (type: DN_Attribute_Type, is_known_oid: bool) { + switch { + case bytes.equal(oid, _OID_AT_CN): + return .Common_Name, true + case bytes.equal(oid, _OID_AT_C): + return .Country, true + case bytes.equal(oid, _OID_AT_L): + return .Locality, true + case bytes.equal(oid, _OID_AT_ST): + return .State_Or_Province, true + case bytes.equal(oid, _OID_AT_O): + return .Organization, true + case bytes.equal(oid, _OID_AT_OU): + return .Organizational_Unit, true + case bytes.equal(oid, _OID_AT_SERIAL): + return .Serial_Number, true + } + return .Other, false +} + +// _dn_escape writes `s` into `b` with the RFC 4514 section 2.4 escapes: a +// leading '#' or space and a trailing space are escaped, as are the characters +// " + , ; < > \ and the NUL byte. +@(private) +_dn_escape :: proc(b: ^strings.Builder, s: string) { + for i in 0 ..< len(s) { + c := s[i] + lead := i == 0 && (c == ' ' || c == '#') + trail := i == len(s) - 1 && c == ' ' + switch c { + case '"', '+', ',', ';', '<', '>', '\\', 0x00: + strings.write_byte(b, '\\') + strings.write_byte(b, c) + case: + if lead || trail { + strings.write_byte(b, '\\') + } + strings.write_byte(b, c) + } + } +} diff --git a/core/crypto/x509/name_constraints.odin b/core/crypto/x509/name_constraints.odin new file mode 100644 index 000000000..bb1d2dcfa --- /dev/null +++ b/core/crypto/x509/name_constraints.odin @@ -0,0 +1,363 @@ +package x509 + +import "core:bytes" +import "core:encoding/asn1" +import "core:strings" + +// Name-constraint processing (RFC 5280 section 4.2.1.10 + the section 6.1.4 +// path-validation checks), scoped to the dNSName and iPAddress GeneralName +// forms. A NameConstraints extension that uses any other base form +// (directoryName, rfc822Name, uniformResourceIdentifier, otherName, …), a +// non-zero minimum, a maximum, or that fails to decode causes the whole chain +// to be rejected (fail closed): we never accept a constraint we cannot fully +// evaluate. + +// GeneralName context tags used in NameConstraints (X.509 GeneralName CHOICE). +@(private) +_GN_OTHER_NAME :: 0 // otherName [0] (constructed) +@(private) +_GN_RFC822 :: 1 // rfc822Name [1] +@(private) +_GN_DNS :: 2 // dNSName [2] +@(private) +_GN_DIR :: 4 // directoryName [4] (constructed) +@(private) +_GN_URI :: 6 // uniformResourceIdentifier [6] +@(private) +_GN_IP :: 7 // iPAddress [7] + +// _check_name_constraints enforces every NameConstraints extension in `chain` +// (leaf at index 0, trust anchor last). Each CA's constraints apply to every +// certificate below it; checking each CA independently against each +// subordinate yields the RFC 5280 permitted=intersection / excluded=union +// semantics without accumulator state. Returns true when the chain is +// acceptable, false when a name is forbidden (or a constraint cannot be +// evaluated, in which case we reject rather than guess). +@(private) +_check_name_constraints :: proc(chain: []^Certificate) -> bool { + // RFC 5280 4.2.1.10: NameConstraints MUST appear only in a CA certificate. + // A non-CA end entity that carries it is malformed — reject the chain. + if _, leaf_has := _find_name_constraints(chain[0]); leaf_has && !chain[0].is_ca { + return false + } + // ci walks the CAs (anchor down to the first intermediate above the leaf); + // the leaf (index 0) never constrains. + for ci := len(chain) - 1; ci >= 1; ci -= 1 { + nc, has := _find_name_constraints(chain[ci]) + if !has { + continue + } + if !_nc_decidable(nc) { + return false // a form/feature we cannot evaluate: fail closed + } + for sub := ci - 1; sub >= 0; sub -= 1 { + c := chain[sub] + // RFC 5280 section 6.1.4(b): constraints are not applied to a + // self-issued intermediate, only to the final (leaf) certificate. + if sub != 0 && bytes.equal(c.raw_subject, c.raw_issuer) { + continue + } + if !_names_permitted(nc, c) { + return false + } + } + } + return true +} + +// _find_name_constraints returns the raw extnValue DER of the cert's +// NameConstraints extension, if present. +@(private) +_find_name_constraints :: proc(cert: ^Certificate) -> (der: []byte, ok: bool) { + for ext in cert.extensions { + if bytes.equal(ext.oid, _OID_EXT_NAME_CONSTRAINTS) { + return ext.value, true + } + } + return nil, false +} + +// _names_permitted checks a subordinate certificate's dNSName and iPAddress +// SANs against one CA's NameConstraints: each name form, when the CA lists +// permitted subtrees for it, must match at least one; and no name may match an +// excluded subtree. +@(private) +_names_permitted :: proc(nc: []byte, sub: ^Certificate) -> bool { + for dns in sub.dns_names { + if _nc_section_has(nc, false, _GN_DNS) && !_nc_dns_match(nc, false, dns) { + return false + } + if _nc_dns_match(nc, true, dns) { + return false + } + } + for ip in sub.ip_addresses { + if _nc_section_has(nc, false, _GN_IP) && !_nc_ip_match(nc, false, ip) { + return false + } + if _nc_ip_match(nc, true, ip) { + return false + } + } + return true +} + +// _nc_decidable structurally validates a NameConstraints and reports whether we +// can fully evaluate it. It returns false, so the caller fails closed, when the +// extension is malformed (an element other than permittedSubtrees [0] / +// excludedSubtrees [1], neither section present, or an empty GeneralSubtrees, +// which SIZE (1..MAX) forbids), or when a subtree uses a base form other than +// dNSName / iPAddress or carries a minimum/maximum (barred by the RFC 5280 +// profile). Only a NameConstraints whose every subtree is a bare +// dNSName/iPAddress is decidable. +@(private) +_nc_decidable :: proc(nc: []byte) -> bool { + cur: asn1.Cursor + asn1.cursor_init(&cur, nc) + seq, e := asn1.read_sequence(&cur) + if e != .None || asn1.done(&cur) != .None { + return false + } + sections := 0 + for !asn1.is_empty(&seq) { + tag, content, re := asn1.read_any(&seq) + if re != .None { + return false + } + // The only permitted members are permittedSubtrees [0] / excludedSubtrees + // [1]; a NULL or any other element (CABF 7.1.2.5.2) is malformed. + if tag.class != .Context_Specific || (tag.number != 0 && tag.number != 1) { + return false + } + sections += 1 + if !_nc_section_wellformed(content) { + return false + } + } + return sections > 0 +} + +// _nc_section_wellformed checks one GeneralSubtrees body: at least one subtree +// (SIZE (1..MAX)), every base a bare dNSName/iPAddress, no minimum/maximum. +@(private) +_nc_section_wellformed :: proc(content: []byte) -> bool { + cur: asn1.Cursor + asn1.cursor_init(&cur, content) + count := 0 + for !asn1.is_empty(&cur) { + gs, e := asn1.read_sequence(&cur) + if e != .None { + return false + } + count += 1 + tag, _, be := asn1.read_any(&gs) + if be != .None || tag.class != .Context_Specific { + return false + } + if tag.number != _GN_DNS && tag.number != _GN_IP { + return false // directoryName / rfc822 / URI / otherName / … + } + // minimum [0] / maximum [1] must both be absent (RFC 5280 profile). + if asn1.done(&gs) != .None { + return false + } + } + return count > 0 +} + +// _nc_section returns the raw content octets of the permitted [0] (or excluded +// [1]) GeneralSubtrees, i.e. the concatenation of GeneralSubtree elements. +@(private) +_nc_section :: proc(nc: []byte, excluded: bool) -> (content: []byte, present: bool) { + want := u32(excluded ? 1 : 0) + cur: asn1.Cursor + asn1.cursor_init(&cur, nc) + seq, e := asn1.read_sequence(&cur) + if e != .None { + return nil, false + } + for !asn1.is_empty(&seq) { + tag, c, re := asn1.read_any(&seq) + if re != .None { + return nil, false + } + if tag.class == .Context_Specific && tag.number == want { + return c, true + } + } + return nil, false +} + +// _nc_section_has reports whether the given section contains at least one +// subtree whose base is `form`. +@(private) +_nc_section_has :: proc(nc: []byte, excluded: bool, form: u32) -> bool { + content, present := _nc_section(nc, excluded) + if !present { + return false + } + cur: asn1.Cursor + asn1.cursor_init(&cur, content) + for !asn1.is_empty(&cur) { + gs, e := asn1.read_sequence(&cur) + if e != .None { + return false + } + tag, _, be := asn1.read_any(&gs) + if be != .None { + return false + } + if tag.class == .Context_Specific && tag.number == form { + return true + } + } + return false +} + +// _nc_dns_match reports whether `name` matches any dNSName subtree in the +// given section. +@(private) +_nc_dns_match :: proc(nc: []byte, excluded: bool, name: string) -> bool { + content, present := _nc_section(nc, excluded) + if !present { + return false + } + cur: asn1.Cursor + asn1.cursor_init(&cur, content) + for !asn1.is_empty(&cur) { + gs, e := asn1.read_sequence(&cur) + if e != .None { + return false + } + tag, base, be := asn1.read_any(&gs) + if be != .None { + return false + } + if tag.class == .Context_Specific && tag.number == _GN_DNS { + c := string(base) + // A wildcard SAN "*.B" stands for every "