]> git.codecow.com Git - nano25519.git/commitdiff
Hoist buffer pointers to only reach into wasm once per buffer instead of every call...
authorChris Duncan <chris@codecow.com>
Mon, 24 Aug 2026 22:26:47 +0000 (15:26 -0700)
committerChris Duncan <chris@codecow.com>
Mon, 24 Aug 2026 22:26:47 +0000 (15:26 -0700)
src/assembly/index.ts
src/lib/nano25519.ts

index 443e0b63d5ac2d0f16ac11c84668883a52584647..343ad1722e727666488592a3b3e9bb7e0c7c2f9b 100644 (file)
@@ -10,9 +10,9 @@ const PUBLICKEY_BYTES: i32 = 32
 const SIGNATURE_BYTES: i32 = 64
 
 // Static I/O buffers
-const INPUT_BUFFER_BYTES: i32 = SIGNATURE_BYTES + PUBLICKEY_BYTES
-const MESSAGE_BUFFER_BYTES: i32 = 32768
-const OUTPUT_BUFFER_BYTES: i32 = SIGNATURE_BYTES
+export const INPUT_BUFFER_BYTES: i32 = SIGNATURE_BYTES + PUBLICKEY_BYTES
+export const MESSAGE_BUFFER_BYTES: i32 = 32768
+export const OUTPUT_BUFFER_BYTES: i32 = SIGNATURE_BYTES
 const INPUT_BUFFER = new StaticArray<u8>(INPUT_BUFFER_BYTES)
 const MESSAGE_BUFFER = new StaticArray<u8>(MESSAGE_BUFFER_BYTES)
 const OUTPUT_BUFFER = new StaticArray<u8>(OUTPUT_BUFFER_BYTES)
index cdd064294177ce8a328bbe70ce94d28004a99819..a5e1b21a727f1e7a9a1d5eda28cfcb9d75f4417c 100644 (file)
@@ -12,6 +12,9 @@ type Exports = {
                getInputPointer: () => number
                getMessagePointer: () => number
                getOutputPointer: () => number
+               INPUT_BUFFER_BYTES: WebAssembly.Global
+               MESSAGE_BUFFER_BYTES: WebAssembly.Global
+               OUTPUT_BUFFER_BYTES: WebAssembly.Global
                memory: WebAssembly.Memory
        }
 }
@@ -42,6 +45,12 @@ const { exports } = new WebAssembly.Instance(module, {
                }
        }
 }) as Exports
+const IN_LEN = exports.INPUT_BUFFER_BYTES.value
+const MSG_LEN = exports.MESSAGE_BUFFER_BYTES.value
+const OUT_LEN = exports.OUTPUT_BUFFER_BYTES.value
+const IN_PTR = exports.getInputPointer()
+const MSG_PTR = exports.getMessagePointer()
+const OUT_PTR = exports.getOutputPointer()
 
 export function derive (k: unknown, out?: unknown): string | Uint8Array<ArrayBuffer> | void {
        if (typeof out !== 'undefined' && !(isBytes(out) && out.byteLength === 32)) {
@@ -52,15 +61,13 @@ export function derive (k: unknown, out?: unknown): string | Uint8Array<ArrayBuf
        let buffer = new Uint8Array(exports.memory.buffer)
        try {
                privateKey.set(normalize('private key', 32, 32, k))
-               let inPtr = exports.getInputPointer()
                for (let i = 0; i < 32; i++) {
-                       buffer[inPtr + i] = privateKey[i]
+                       buffer[IN_PTR + i] = privateKey[i]
                }
                exports.derive()
-               const outPtr = exports.getOutputPointer()
                buffer = new Uint8Array(exports.memory.buffer)
                for (let i = 0; i < 32; i++) {
-                       publicKey[i] = buffer[outPtr + i]
+                       publicKey[i] = buffer[OUT_PTR + i]
                }
                if (typeof k === 'string') {
                        let hex = ''
@@ -90,20 +97,17 @@ export function sign (m: unknown, k: unknown, s?: unknown): string | Uint8Array<
        let buffer = new Uint8Array(exports.memory.buffer)
        try {
                secretKey.set(normalize('secret key', 64, 64, k))
-               let inPtr = exports.getInputPointer()
                for (let i = 0; i < 64; i++) {
-                       buffer[inPtr + i] = secretKey[i]
+                       buffer[IN_PTR + i] = secretKey[i]
                }
                const message = normalize('message', 0, mLen, m)
-               let mPtr = exports.getMessagePointer()
                for (let i = 0; i < message.byteLength; i++) {
-                       buffer[mPtr + i] = message[i]
+                       buffer[MSG_PTR + i] = message[i]
                }
                exports.sign(message.byteLength)
-               const outPtr = exports.getOutputPointer()
                buffer = new Uint8Array(exports.memory.buffer)
                for (let i = 0; i < 64; i++) {
-                       signature[i] = buffer[outPtr + i]
+                       signature[i] = buffer[OUT_PTR + i]
                }
                if (typeof k === 'string') {
                        let hex = ''
@@ -130,17 +134,14 @@ export function verify (s: unknown, m: unknown, k: unknown): boolean {
                const signature = normalize('signature', 64, 64, s)
                const message = normalize('message', 0, mLen, m)
                const publicKey = normalize('public key', 32, 32, k)
-               let mPtr = exports.getMessagePointer()
-               let inPtr = exports.getInputPointer()
                for (let i = 0; i < message.byteLength; i++) {
-                       buffer[mPtr + i] = message[i]
+                       buffer[MSG_PTR + i] = message[i]
                }
                for (let i = 0; i < 64; i++) {
-                       buffer[inPtr + i] = signature[i]
+                       buffer[IN_PTR + i] = signature[i]
                }
-               inPtr += 64
                for (let i = 0; i < 32; i++) {
-                       buffer[inPtr + i] = publicKey[i]
+                       buffer[IN_PTR + 64 + i] = publicKey[i]
                }
                const v = exports.verify(message.byteLength)
                return v === 0
@@ -150,12 +151,9 @@ export function verify (s: unknown, m: unknown, k: unknown): boolean {
 }
 
 function clear (memory: Uint8Array): void {
-       let inPtr = exports.getInputPointer()
-       let msgPtr = exports.getMessagePointer()
-       let outPtr = exports.getOutputPointer()
-       memory.fill(0, inPtr, inPtr + 96)
-       memory.fill(0, msgPtr, msgPtr + mLen)
-       memory.fill(0, outPtr, outPtr + 64)
+       memory.fill(0, IN_PTR, IN_LEN)
+       memory.fill(0, MSG_PTR, MSG_LEN)
+       memory.fill(0, OUT_PTR, OUT_LEN)
 }
 
 function isBytes (a: unknown): a is Uint8Array<ArrayBuffer> {