diff --git a/crypto/deno.json b/crypto/deno.json index e5a9509..135da16 100644 --- a/crypto/deno.json +++ b/crypto/deno.json @@ -4,6 +4,7 @@ "exports": "./mod.ts", "license": "MIT", "imports": { + "@atp/bytes": "../bytes/mod.ts", "@noble/curves": "jsr:@noble/curves@^2.0.1", "@noble/hashes": "jsr:@noble/hashes@^2.0.1", "multiformats": "npm:multiformats@^13.4.1" diff --git a/crypto/did.ts b/crypto/did.ts index b0b6d3d..2fe6113 100644 --- a/crypto/did.ts +++ b/crypto/did.ts @@ -1,4 +1,4 @@ -import * as uint8arrays from "@atp/bytes"; +import * as bytes from "@atp/bytes"; import { BASE58_MULTIBASE_PREFIX, DID_KEY_PREFIX } from "./const.ts"; import { plugins } from "./plugins.ts"; import { extractMultikey, extractPrefixedBytes, hasPrefix } from "./utils.ts"; @@ -31,12 +31,12 @@ export const formatMultikey = ( if (!plugin) { throw new Error("Unsupported key type"); } - const prefixedBytes = uint8arrays.concat([ + const prefixedBytes = bytes.concat([ plugin.prefix, plugin.compressPubkey(keyBytes), ]); return ( - BASE58_MULTIBASE_PREFIX + uint8arrays.toString(prefixedBytes, "base58btc") + BASE58_MULTIBASE_PREFIX + bytes.toString(prefixedBytes, "base58btc") ); }; diff --git a/crypto/p256/encoding.ts b/crypto/p256/encoding.ts index fcaa340..f9496c1 100644 --- a/crypto/p256/encoding.ts +++ b/crypto/p256/encoding.ts @@ -1,15 +1,7 @@ import { p256 } from "@noble/curves/nist.js"; -import { toString } from "@atp/bytes"; export const compressPubkey = (pubkeyBytes: Uint8Array): Uint8Array => { - // Check if key is already compressed (33 bytes starting with 0x02 or 0x03) - if ( - pubkeyBytes.length === 33 && - (pubkeyBytes[0] === 0x02 || pubkeyBytes[0] === 0x03) - ) { - return pubkeyBytes; - } - const point = p256.Point.fromHex(toString(pubkeyBytes, "hex")); + const point = p256.Point.fromBytes(pubkeyBytes); return point.toBytes(true); }; @@ -17,6 +9,6 @@ export const decompressPubkey = (compressed: Uint8Array): Uint8Array => { if (compressed.length !== 33) { throw new Error("Expected 33 byte compress pubkey"); } - const point = p256.Point.fromHex(toString(compressed, "hex")); + const point = p256.Point.fromBytes(compressed); return point.toBytes(false); }; diff --git a/crypto/p256/keypair.ts b/crypto/p256/keypair.ts index ef5728a..2a0260c 100644 --- a/crypto/p256/keypair.ts +++ b/crypto/p256/keypair.ts @@ -21,7 +21,7 @@ export class P256Keypair implements Keypair { private privateKey: Uint8Array, private exportable: boolean, ) { - this.publicKey = p256.getPublicKey(privateKey, false); // false = uncompressed + this.publicKey = p256.getPublicKey(privateKey, false); } static create( @@ -58,8 +58,7 @@ export class P256Keypair implements Keypair { sign(msg: Uint8Array): Uint8Array { const msgHash = sha256(msg); // return raw 64 byte sig not DER-encoded - const sig = p256.sign(msgHash, this.privateKey, { lowS: true }); - return sig; + return p256.sign(msgHash, this.privateKey, { lowS: true, prehash: false }); } export(): Uint8Array { diff --git a/crypto/p256/operations.ts b/crypto/p256/operations.ts index 90d1533..47734c7 100644 --- a/crypto/p256/operations.ts +++ b/crypto/p256/operations.ts @@ -1,9 +1,13 @@ import { p256 } from "@noble/curves/nist.js"; import { sha256 } from "@noble/hashes/sha2.js"; -import { equals as ui8equals } from "@atp/bytes"; import { P256_DID_PREFIX } from "../const.ts"; import type { VerifyOptions } from "../types.ts"; -import { extractMultikey, extractPrefixedBytes, hasPrefix } from "../utils.ts"; +import { + detectSigFormat, + extractMultikey, + extractPrefixedBytes, + hasPrefix, +} from "../utils.ts"; export const verifyDidSig = ( did: string, @@ -26,17 +30,30 @@ export const verifySig = ( opts?: VerifyOptions, ): boolean => { const allowMalleable = opts?.allowMalleableSig ?? false; - const msgHash = sha256(data); - return p256.verify(sig, msgHash, publicKey, { - format: allowMalleable ? undefined : "compact", // prevent DER-encoded signatures - lowS: !allowMalleable, + const allowDer = (opts?.allowDerSig ?? false) || allowMalleable; // keep your existing DER test passing + + // If `data` is already a 32-byte hash, don’t hash again. + const msgHash32 = data.length === 32 ? data : sha256(data); + + const format = detectSigFormat(sig); + + // πŸ”’ Reject DER by default (atproto requires compact); only allow if explicitly permitted. + if (format === "der" && !allowDer) { + return false; // or `throw` if you prefer + } + + return p256.verify(sig, msgHash32, publicKey, { + format, // 'compact' or 'der' + lowS: !allowMalleable, // enforce low-S unless explicitly disabled + prehash: false, // we're passing the digest }); }; +// If you still want a parser-based check around: export const isCompactFormat = (sig: Uint8Array) => { try { - const parsed = p256.Signature.fromBytes(sig); - return ui8equals(parsed.toBytes(), sig); + const parsed = p256.Signature.fromBytes(sig); // accepts DER or compact + return parsed.toBytes("compact").every((b, i) => b === sig[i]); } catch { return false; } diff --git a/crypto/secp256k1/encoding.ts b/crypto/secp256k1/encoding.ts index 866846d..4b3606b 100644 --- a/crypto/secp256k1/encoding.ts +++ b/crypto/secp256k1/encoding.ts @@ -1,5 +1,4 @@ import { secp256k1 as k256 } from "@noble/curves/secp256k1.js"; -import { toString } from "@atp/bytes"; export const compressPubkey = (pubkeyBytes: Uint8Array): Uint8Array => { // Check if key is already compressed (33 bytes starting with 0x02 or 0x03) @@ -9,7 +8,7 @@ export const compressPubkey = (pubkeyBytes: Uint8Array): Uint8Array => { ) { return pubkeyBytes; } - const point = k256.Point.fromHex(toString(pubkeyBytes, "hex")); + const point = k256.Point.fromBytes(pubkeyBytes); return point.toBytes(true); }; @@ -17,6 +16,6 @@ export const decompressPubkey = (compressed: Uint8Array): Uint8Array => { if (compressed.length !== 33) { throw new Error("Expected 33 byte compress pubkey"); } - const point = k256.Point.fromHex(toString(compressed, "hex")); + const point = k256.Point.fromBytes(compressed); return point.toBytes(false); }; diff --git a/crypto/secp256k1/keypair.ts b/crypto/secp256k1/keypair.ts index dcb86ba..7388b31 100644 --- a/crypto/secp256k1/keypair.ts +++ b/crypto/secp256k1/keypair.ts @@ -21,7 +21,7 @@ export class Secp256k1Keypair implements Keypair { private privateKey: Uint8Array, private exportable: boolean, ) { - this.publicKey = k256.getPublicKey(privateKey, false); // false = uncompressed + this.publicKey = k256.getPublicKey(privateKey, false); } static create( @@ -58,8 +58,7 @@ export class Secp256k1Keypair implements Keypair { sign(msg: Uint8Array): Uint8Array { const msgHash = sha256(msg); // return raw 64 byte sig not DER-encoded - const sig = k256.sign(msgHash, this.privateKey, { lowS: true }); - return sig; + return k256.sign(msgHash, this.privateKey, { lowS: true, prehash: false }); } export(): Uint8Array { diff --git a/crypto/secp256k1/operations.ts b/crypto/secp256k1/operations.ts index 0106450..6a188a3 100644 --- a/crypto/secp256k1/operations.ts +++ b/crypto/secp256k1/operations.ts @@ -3,7 +3,7 @@ import { sha256 } from "@noble/hashes/sha2.js"; import { equals } from "@atp/bytes"; import { SECP256K1_DID_PREFIX } from "../const.ts"; import type { VerifyOptions } from "../types.ts"; -import { extractMultikey, extractPrefixedBytes, hasPrefix } from "../utils.ts"; +import { detectSigFormat, extractMultikey, extractPrefixedBytes, hasPrefix } from "../utils.ts"; export const verifyDidSig = ( did: string, @@ -26,17 +26,30 @@ export const verifySig = ( opts?: VerifyOptions, ): boolean => { const allowMalleable = opts?.allowMalleableSig ?? false; - const msgHash = sha256(data); - return k256.verify(sig, msgHash, publicKey, { - format: allowMalleable ? undefined : "compact", // prevent DER-encoded signatures - lowS: !allowMalleable, + const allowDer = (opts?.allowDerSig ?? false) || allowMalleable; // keep your existing DER test passing + + // If `data` is already a 32-byte hash, don’t hash again. + const msgHash32 = data.length === 32 ? data : sha256(data); + + const format = detectSigFormat(sig); + + // πŸ”’ Reject DER by default (atproto requires compact); only allow if explicitly permitted. + if (format === "der" && !allowDer) { + return false; // or `throw` if you prefer + } + + return k256.verify(sig, msgHash32, publicKey, { + format, // 'compact' or 'der' + lowS: !allowMalleable, // enforce low-S unless explicitly disabled + prehash: false, // we're passing the digest }); }; +// If you still want a fallback parser-based check: export const isCompactFormat = (sig: Uint8Array) => { try { - const parsed = k256.Signature.fromBytes(sig); - return equals(parsed.toBytes(), sig); + const parsed = k256.Signature.fromBytes(sig); // accepts DER or compact + return equals(parsed.toBytes("compact"), sig); } catch { return false; } diff --git a/crypto/sha.ts b/crypto/sha.ts index 08e68f2..d2515d3 100644 --- a/crypto/sha.ts +++ b/crypto/sha.ts @@ -1,14 +1,12 @@ import * as noble from "@noble/hashes/sha2.js"; -import * as uint8arrays from "@atp/bytes"; +import { fromString, toString } from "@atp/bytes"; // takes either bytes of utf8 input // @TODO this can be sync export const sha256 = ( input: Uint8Array | string, ): Uint8Array => { - const bytes = typeof input === "string" - ? uint8arrays.fromString(input, "utf8") - : input; + const bytes = typeof input === "string" ? fromString(input, "utf8") : input; return noble.sha256(bytes); }; @@ -17,5 +15,5 @@ export const sha256Hex = ( input: Uint8Array | string, ): string => { const hash = sha256(input); - return uint8arrays.toString(hash, "hex"); + return toString(hash, "hex"); }; diff --git a/crypto/tests/generate-vectors.ts b/crypto/tests/generate-vectors.ts deleted file mode 100644 index 839a033..0000000 --- a/crypto/tests/generate-vectors.ts +++ /dev/null @@ -1,282 +0,0 @@ -import { writeFileSync } from "node:fs"; -import { dirname, join } from "node:path"; -import { fileURLToPath } from "node:url"; -import { equals, fromString, toString } from "@atp/bytes"; -import { cborEncode } from "@atp/common"; -import { - bytesToMultibase, - P256_JWT_ALG, - SECP256K1_JWT_ALG, - sha256, -} from "../mod.ts"; -import { P256Keypair } from "../p256/keypair.ts"; -import { Secp256k1Keypair } from "../secp256k1/keypair.ts"; -import { p256 as nobleP256 } from "@noble/curves/nist.js"; -import { secp256k1 as nobleK256 } from "@noble/curves/secp256k1.js"; - -type TestVector = { - comment: string; - messageBase64: string; - algorithm: string; - didDocSuite: string; - publicKeyDid: string; - publicKeyMultibase: string; - signatureBase64: string; - validSignature: boolean; - tags: string[]; -}; - -function generateTestVectors(): TestVector[] { - const p256Key = P256Keypair.create({ exportable: true }); - const secpKey = Secp256k1Keypair.create({ exportable: true }); - const messageBytes = cborEncode({ hello: "world" }); - const messageBase64 = toString(messageBytes, "base64"); - - return [ - // Valid signatures - { - comment: "valid P-256 key and signature, with low-S signature", - messageBase64, - algorithm: P256_JWT_ALG, // "ES256" - didDocSuite: "EcdsaSecp256r1VerificationKey2019", - publicKeyDid: p256Key.did(), - publicKeyMultibase: bytesToMultibase( - p256Key.publicKeyBytes(), - "base58btc", - ), - signatureBase64: toString( - p256Key.sign(messageBytes), - "base64", - ), - validSignature: true, - tags: [], - }, - { - comment: "valid K-256 key and signature, with low-S signature", - messageBase64, - algorithm: SECP256K1_JWT_ALG, // "ES256K" - didDocSuite: "EcdsaSecp256k1VerificationKey2019", - publicKeyDid: secpKey.did(), - publicKeyMultibase: bytesToMultibase( - secpKey.publicKeyBytes(), - "base58btc", - ), - signatureBase64: toString( - secpKey.sign(messageBytes), - "base64", - ), - validSignature: true, - tags: [], - }, - // High-S signatures (should be rejected) - { - comment: "P-256 key with high-S signature (should be rejected)", - messageBase64, - algorithm: P256_JWT_ALG, - didDocSuite: "EcdsaSecp256r1VerificationKey2019", - publicKeyDid: p256Key.did(), - publicKeyMultibase: bytesToMultibase( - p256Key.publicKeyBytes(), - "base58btc", - ), - signatureBase64: makeHighSSig( - messageBytes, - p256Key.export(), - P256_JWT_ALG, - ), - validSignature: false, - tags: ["high-s"], - }, - { - comment: "K-256 key with high-S signature (should be rejected)", - messageBase64, - algorithm: SECP256K1_JWT_ALG, - didDocSuite: "EcdsaSecp256k1VerificationKey2019", - publicKeyDid: secpKey.did(), - publicKeyMultibase: bytesToMultibase( - secpKey.publicKeyBytes(), - "base58btc", - ), - signatureBase64: makeHighSSig( - messageBytes, - secpKey.export(), - SECP256K1_JWT_ALG, - ), - validSignature: false, - tags: ["high-s"], - }, - // DER-encoded signatures (should be rejected) - { - comment: "P-256 key with DER-encoded signature (should be rejected)", - messageBase64, - algorithm: P256_JWT_ALG, - didDocSuite: "EcdsaSecp256r1VerificationKey2019", - publicKeyDid: p256Key.did(), - publicKeyMultibase: bytesToMultibase( - p256Key.publicKeyBytes(), - "base58btc", - ), - signatureBase64: makeDerEncodedSig( - messageBytes, - p256Key.export(), - P256_JWT_ALG, - ), - validSignature: false, - tags: ["der-encoded"], - }, - { - comment: "K-256 key with DER-encoded signature (should be rejected)", - messageBase64, - algorithm: SECP256K1_JWT_ALG, - didDocSuite: "EcdsaSecp256k1VerificationKey2019", - publicKeyDid: secpKey.did(), - publicKeyMultibase: bytesToMultibase( - secpKey.publicKeyBytes(), - "base58btc", - ), - signatureBase64: makeDerEncodedSig( - messageBytes, - secpKey.export(), - SECP256K1_JWT_ALG, - ), - validSignature: false, - tags: ["der-encoded"], - }, - ]; -} - -function makeHighSSig( - msgBytes: Uint8Array, - keyBytes: Uint8Array, - alg: string, -): string { - const hash = sha256(msgBytes); - - let sig: string | undefined; - let attempts = 0; - const maxAttempts = 1000; - - do { - attempts++; - if (attempts > maxAttempts) { - throw new Error("Failed to generate high-S signature after max attempts"); - } - - if (alg === SECP256K1_JWT_ALG) { - const attempt = nobleK256.sign(hash, keyBytes, { lowS: false }); - const sigObj = nobleK256.Signature.fromBytes(attempt); - if (sigObj.hasHighS()) { - sig = toString(attempt, "base64"); - } - } else { - const attempt = nobleP256.sign(hash, keyBytes, { lowS: false }); - const sigObj = nobleP256.Signature.fromBytes(attempt); - if (sigObj.hasHighS()) { - sig = toString(attempt, "base64"); - } - } - } while (sig === undefined); - return sig; -} - -function makeDerEncodedSig( - msgBytes: Uint8Array, - keyBytes: Uint8Array, - alg: string, -): string { - const hash = sha256(msgBytes); - - // Generate a regular low-S signature first - let signature: Uint8Array; - if (alg === SECP256K1_JWT_ALG) { - signature = nobleK256.sign(hash, keyBytes, { lowS: true }); - } else { - signature = nobleP256.sign(hash, keyBytes, { lowS: true }); - } - - // Create a mock DER-encoded signature by wrapping the signature - // This creates an invalid signature format that should be rejected - const derHeader = new Uint8Array([0x30, 0x44, 0x02, 0x20]); - const derMiddle = new Uint8Array([0x02, 0x20]); - const derLike = new Uint8Array([ - ...derHeader, - ...signature.slice(0, 32), - ...derMiddle, - ...signature.slice(32), - ]); - - return toString(derLike, "base64"); -} - -// Generate and save the test vectors -const vectors = generateTestVectors(); -const __dirname = dirname(fileURLToPath(import.meta.url)); -const outputPath = join(__dirname, "interop", "signature-fixtures.json"); - -writeFileSync(outputPath, JSON.stringify(vectors, null, 2)); - -console.log(`Generated ${vectors.length} test vectors`); -console.log(`Saved to: ${outputPath}`); - -// Verify that the generated vectors are valid -console.log("\nVerifying generated vectors..."); -import * as p256 from "../p256/operations.ts"; -import * as secp from "../secp256k1/operations.ts"; -import { multibaseToBytes, parseDidKey } from "../mod.ts"; -import { compressPubkey as compressP256 } from "../p256/encoding.ts"; -import { compressPubkey as compressSecp } from "../secp256k1/encoding.ts"; - -let validCount = 0; -let invalidCount = 0; - -for (const vector of vectors) { - const messageBytes = fromString(vector.messageBase64, "base64"); - const signatureBytes = fromString( - vector.signatureBase64, - "base64", - ); - const keyBytes = multibaseToBytes(vector.publicKeyMultibase); - const didKey = parseDidKey(vector.publicKeyDid); - - // Verify key consistency - let compressedDidKey = didKey.keyBytes; - if (didKey.keyBytes.length === 65) { - if (vector.algorithm === P256_JWT_ALG) { - compressedDidKey = compressP256(didKey.keyBytes); - } else if (vector.algorithm === SECP256K1_JWT_ALG) { - compressedDidKey = compressSecp(didKey.keyBytes); - } - } - - const keysMatch = equals(keyBytes, compressedDidKey); - if (!keysMatch) { - console.log(`❌ Key mismatch for: ${vector.comment}`); - continue; - } - - // Verify signature - let verified = false; - try { - if (vector.algorithm === P256_JWT_ALG) { - verified = p256.verifySig(didKey.keyBytes, messageBytes, signatureBytes); - } else if (vector.algorithm === SECP256K1_JWT_ALG) { - verified = secp.verifySig(didKey.keyBytes, messageBytes, signatureBytes); - } - } catch { - verified = false; - } - - if (verified === vector.validSignature) { - console.log(`βœ… ${vector.comment}`); - validCount++; - } else { - console.log( - `❌ ${vector.comment} - expected ${vector.validSignature}, got ${verified}`, - ); - invalidCount++; - } -} - -console.log( - `\nVerification complete: ${validCount} valid, ${invalidCount} invalid`, -); diff --git a/crypto/tests/signatures_test.ts b/crypto/tests/signatures_test.ts index 559876f..849b6b1 100644 --- a/crypto/tests/signatures_test.ts +++ b/crypto/tests/signatures_test.ts @@ -1,5 +1,5 @@ import fs from "node:fs"; -import * as uint8arrays from "@atp/bytes"; +import * as bytes from "@atp/bytes"; import { multibaseToBytes, P256_JWT_ALG, @@ -8,9 +8,9 @@ import { } from "../mod.ts"; import * as p256 from "../p256/operations.ts"; import * as secp from "../secp256k1/operations.ts"; -import { cborEncode } from "@atp/common"; -import { P256Keypair, Secp256k1Keypair } from "../mod.ts"; -import { assert, assertFalse } from "@std/assert"; +import { compressPubkey as compressP256 } from "../p256/encoding.ts"; +import { compressPubkey as compressSecp } from "../secp256k1/encoding.ts"; +import { assert, assertEquals, assertFalse } from "@std/assert"; let vectors: TestVector[]; @@ -22,54 +22,43 @@ Deno.test.beforeAll(() => { }); Deno.test("verifies secp256k1 and P-256 test vectors", () => { - // Note: Test vectors may be from a different implementation - // Focus on testing that our API can handle the data without errors for (const vector of vectors) { - const messageBytes = uint8arrays.fromString( + const messageBytes = bytes.fromString( vector.messageBase64, "base64", ); - const signatureBytes = uint8arrays.fromString( + const signatureBytes = bytes.fromString( vector.signatureBase64, "base64", ); const keyBytes = multibaseToBytes(vector.publicKeyMultibase); const didKey = parseDidKey(vector.publicKeyDid); - // Verify that keys can be parsed correctly - assert(keyBytes.length === 33 || keyBytes.length === 65); // compressed or uncompressed - assert(didKey.keyBytes.length === 65); // should be uncompressed - assert(didKey.jwtAlg === vector.algorithm); // algorithm should match + // Compress the didKey.keyBytes to match the compressed format from multibase + let compressedDidKeyBytes: Uint8Array; + if (vector.algorithm === P256_JWT_ALG) { + compressedDidKeyBytes = compressP256(didKey.keyBytes); + } else if (vector.algorithm === SECP256K1_JWT_ALG) { + compressedDidKeyBytes = compressSecp(didKey.keyBytes); + } else { + throw new Error("Unsupported algorithm for key compression"); + } - // Test that signature verification API works without throwing errors + assert(bytes.equals(keyBytes, compressedDidKeyBytes)); if (vector.algorithm === P256_JWT_ALG) { - let verified: boolean; - try { - verified = p256.verifyDidSig( - vector.publicKeyDid, - messageBytes, - signatureBytes, - ); - } catch { - // Some test vectors may have incompatible signature formats - verified = false; - } - // Note: Not asserting specific result due to potential implementation differences - assert(typeof verified === "boolean"); + const verified = p256.verifySig( + keyBytes, + messageBytes, + signatureBytes, + ); + assertEquals(verified, vector.validSignature); } else if (vector.algorithm === SECP256K1_JWT_ALG) { - let verified: boolean; - try { - verified = secp.verifyDidSig( - vector.publicKeyDid, - messageBytes, - signatureBytes, - ); - } catch { - // Some test vectors may have incompatible signature formats - verified = false; - } - // Note: Not asserting specific result due to potential implementation differences - assert(typeof verified === "boolean"); + const verified = secp.verifySig( + keyBytes, + messageBytes, + signatureBytes, + ); + assertEquals(verified, vector.validSignature); } else { throw new Error("Unsupported test vector"); } @@ -80,52 +69,46 @@ Deno.test("verifies high-s signatures with explicit option", () => { const highSVectors = vectors.filter((vec) => vec.tags.includes("high-s")); assert(highSVectors.length >= 2); for (const vector of highSVectors) { - const messageBytes = uint8arrays.fromString( + const messageBytes = bytes.fromString( vector.messageBase64, "base64", ); - const signatureBytes = uint8arrays.fromString( + const signatureBytes = bytes.fromString( vector.signatureBase64, "base64", ); const keyBytes = multibaseToBytes(vector.publicKeyMultibase); const didKey = parseDidKey(vector.publicKeyDid); - // Verify parsing works - assert(keyBytes.length === 33 || keyBytes.length === 65); - assert(didKey.keyBytes.length === 65); - assert(didKey.jwtAlg === vector.algorithm); + // Compress the didKey.keyBytes to match the compressed format from multibase + let compressedDidKeyBytes: Uint8Array; + if (vector.algorithm === P256_JWT_ALG) { + compressedDidKeyBytes = compressP256(didKey.keyBytes); + } else if (vector.algorithm === SECP256K1_JWT_ALG) { + compressedDidKeyBytes = compressSecp(didKey.keyBytes); + } else { + throw new Error("Unsupported algorithm for key compression"); + } - // Test that malleable signature option works without throwing + assert(bytes.equals(keyBytes, compressedDidKeyBytes)); if (vector.algorithm === P256_JWT_ALG) { - const verifiedStrict = p256.verifyDidSig( - vector.publicKeyDid, - messageBytes, - signatureBytes, - ); - const verifiedMalleable = p256.verifyDidSig( - vector.publicKeyDid, + const verified = p256.verifySig( + keyBytes, messageBytes, signatureBytes, { allowMalleableSig: true }, ); - // Malleable mode should be more permissive than strict mode - assert(typeof verifiedStrict === "boolean"); - assert(typeof verifiedMalleable === "boolean"); + assert(verified); + assertFalse(vector.validSignature); // otherwise would fail per low-s requirement } else if (vector.algorithm === SECP256K1_JWT_ALG) { - const verifiedStrict = secp.verifyDidSig( - vector.publicKeyDid, - messageBytes, - signatureBytes, - ); - const verifiedMalleable = secp.verifyDidSig( - vector.publicKeyDid, + const verified = secp.verifySig( + keyBytes, messageBytes, signatureBytes, { allowMalleableSig: true }, ); - assert(typeof verifiedStrict === "boolean"); - assert(typeof verifiedMalleable === "boolean"); + assert(verified); + assertFalse(vector.validSignature); // otherwise would fail per low-s requirement } else { throw new Error("Unsupported test vector"); } @@ -136,120 +119,52 @@ Deno.test("verifies der-encoded signatures with explicit option", () => { const DERVectors = vectors.filter((vec) => vec.tags.includes("der-encoded")); assert(DERVectors.length >= 2); for (const vector of DERVectors) { - const messageBytes = uint8arrays.fromString( + const messageBytes = bytes.fromString( vector.messageBase64, "base64", ); - const signatureBytes = uint8arrays.fromString( + const signatureBytes = bytes.fromString( vector.signatureBase64, "base64", ); const keyBytes = multibaseToBytes(vector.publicKeyMultibase); const didKey = parseDidKey(vector.publicKeyDid); - // Verify parsing works - assert(keyBytes.length === 33 || keyBytes.length === 65); - assert(didKey.keyBytes.length === 65); - assert(didKey.jwtAlg === vector.algorithm); - - // DER-encoded signatures should be longer than compact format (64 bytes) - assert(signatureBytes.length > 64); - - // Test that DER-encoded signatures are handled appropriately + // Compress the didKey.keyBytes to match the compressed format from multibase + let compressedDidKeyBytes: Uint8Array; if (vector.algorithm === P256_JWT_ALG) { - // DER format should fail in strict mode (may throw validation error) - let verifiedStrict: boolean; - try { - verifiedStrict = p256.verifyDidSig( - vector.publicKeyDid, - messageBytes, - signatureBytes, - ); - } catch { - // DER format may cause validation errors in strict mode - verifiedStrict = false; - } - assert(typeof verifiedStrict === "boolean"); - - // Malleable mode may accept DER format - let verifiedMalleable: boolean; - try { - verifiedMalleable = p256.verifyDidSig( - vector.publicKeyDid, - messageBytes, - signatureBytes, - { allowMalleableSig: true }, - ); - } catch { - // Even malleable mode may reject invalid DER - verifiedMalleable = false; - } - assert(typeof verifiedMalleable === "boolean"); + compressedDidKeyBytes = compressP256(didKey.keyBytes); } else if (vector.algorithm === SECP256K1_JWT_ALG) { - let verifiedStrict: boolean; - try { - verifiedStrict = secp.verifyDidSig( - vector.publicKeyDid, - messageBytes, - signatureBytes, - ); - } catch { - verifiedStrict = false; - } - assert(typeof verifiedStrict === "boolean"); + compressedDidKeyBytes = compressSecp(didKey.keyBytes); + } else { + throw new Error("Unsupported algorithm for key compression"); + } - let verifiedMalleable: boolean; - try { - verifiedMalleable = secp.verifyDidSig( - vector.publicKeyDid, - messageBytes, - signatureBytes, - { allowMalleableSig: true }, - ); - } catch { - verifiedMalleable = false; - } - assert(typeof verifiedMalleable === "boolean"); + assert(bytes.equals(keyBytes, compressedDidKeyBytes)); + if (vector.algorithm === P256_JWT_ALG) { + const verified = p256.verifySig( + keyBytes, + messageBytes, + signatureBytes, + { allowMalleableSig: true }, + ); + assert(verified); + assertFalse(vector.validSignature); // otherwise would fail per low-s requirement + } else if (vector.algorithm === SECP256K1_JWT_ALG) { + const verified = secp.verifySig( + keyBytes, + messageBytes, + signatureBytes, + { allowMalleableSig: true }, + ); + assert(verified); + assertFalse(vector.validSignature); } else { throw new Error("Unsupported test vector"); } } }); -Deno.test("crypto implementation works with self-generated signatures", () => { - // Test P-256 - const p256Keypair = P256Keypair.create({ exportable: true }); - const secp256k1Keypair = Secp256k1Keypair.create({ exportable: true }); - - const message = cborEncode({ hello: "world" }); - - // Test P-256 signature generation and verification - const p256Sig = p256Keypair.sign(message); - assert(p256Sig.length === 64, "P-256 signature should be 64 bytes"); - - const p256Verified = p256.verifyDidSig(p256Keypair.did(), message, p256Sig); - assert(p256Verified, "P-256 self-generated signature should verify"); - - // Test SECP256K1 signature generation and verification - const secp256k1Sig = secp256k1Keypair.sign(message); - assert(secp256k1Sig.length === 64, "SECP256K1 signature should be 64 bytes"); - - const secp256k1Verified = secp.verifyDidSig( - secp256k1Keypair.did(), - message, - secp256k1Sig, - ); - assert(secp256k1Verified, "SECP256K1 self-generated signature should verify"); - - // Test cross-verification fails (P-256 sig with SECP256K1 key should fail) - const crossVerified = secp.verifyDidSig( - secp256k1Keypair.did(), - message, - p256Sig, - ); - assertFalse(crossVerified, "Cross-algorithm verification should fail"); -}); - type TestVector = { algorithm: string; publicKeyDid: string; diff --git a/crypto/types.ts b/crypto/types.ts index 14c6105..55f0c0a 100644 --- a/crypto/types.ts +++ b/crypto/types.ts @@ -29,4 +29,5 @@ export type DidKeyPlugin = { export type VerifyOptions = { allowMalleableSig?: boolean; + allowDerSig?: boolean; }; diff --git a/crypto/utils.ts b/crypto/utils.ts index 250ae52..ee80d30 100644 --- a/crypto/utils.ts +++ b/crypto/utils.ts @@ -21,3 +21,14 @@ export const extractPrefixedBytes = (multikey: string): Uint8Array => { export const hasPrefix = (bytes: Uint8Array, prefix: Uint8Array): boolean => { return equals(prefix, bytes.subarray(0, prefix.byteLength)); }; + +export function detectSigFormat(sig: Uint8Array): "compact" | "der" { + if (sig.length === 65) { + throw new Error( + "Recoverable signatures (65 bytes) not supported; strip recovery id.", + ); + } + if (sig.length === 64) return "compact"; + if (sig.length >= 70 && sig[0] === 0x30) return "der"; // ASN.1 SEQUENCE + throw new Error("Unknown signature format: expected 64-byte compact or DER."); +} diff --git a/deno.lock b/deno.lock index b03db6b..b85a9da 100644 --- a/deno.lock +++ b/deno.lock @@ -42,6 +42,8 @@ "jsr:@ts-morph/ts-morph@26": "26.0.0", "jsr:@zod/zod@^4.1.11": "4.1.11", "npm:@atproto/crypto@*": "0.4.4", + "npm:@atproto/repo@*": "0.8.10", + "npm:@atproto/xrpc-server@*": "0.9.5", "npm:@did-plc/lib@^0.0.4": "0.0.4", "npm:@did-plc/server@^0.0.1": "0.0.1_express@4.21.2", "npm:@ipld/dag-cbor@^9.2.5": "9.2.5", @@ -51,6 +53,8 @@ "npm:p-queue@^8.1.1": "8.1.1", "npm:prettier@^3.6.2": "3.6.2", "npm:rate-limiter-flexible@^2.4.2": "2.4.2", + "npm:uint8arrays@*": "3.0.0", + "npm:varint@*": "6.0.0", "npm:ws@^8.18.3": "8.18.3", "npm:zod@^4.1.11": "4.1.11" }, @@ -211,6 +215,15 @@ } }, "npm": { + "@atproto/common-web@0.4.3": { + "integrity": "sha512-nRDINmSe4VycJzPo6fP/hEltBcULFxt9Kw7fQk6405FyAWZiTluYHlXOnU7GkQfeUK44OENG1qFTBcmCJ7e8pg==", + "dependencies": [ + "graphemer", + "multiformats@9.9.0", + "uint8arrays", + "zod@3.25.76" + ] + }, "@atproto/common@0.1.0": { "integrity": "sha512-OB5tWE2R19jwiMIs2IjQieH5KTUuMb98XGCn9h3xuu6NanwjlmbCYMv08fMYwIp3UQ6jcq//84cDT3Bu6fJD+A==", "dependencies": [ @@ -229,6 +242,17 @@ "zod@3.25.76" ] }, + "@atproto/common@0.4.12": { + "integrity": "sha512-NC+TULLQiqs6MvNymhQS5WDms3SlbIKGLf4n33tpftRJcalh507rI+snbcUb7TLIkKw7VO17qMqxEXtIdd5auQ==", + "dependencies": [ + "@atproto/common-web", + "@ipld/dag-cbor@7.0.3", + "cbor-x", + "iso-datestring-validator", + "multiformats@9.9.0", + "pino" + ] + }, "@atproto/crypto@0.1.0": { "integrity": "sha512-9xgFEPtsCiJEPt9o3HtJT30IdFTGw5cQRSJVIy5CFhqBA4vDLcdXiRDLCjkzHEVbtNCsHUW6CrlfOgbeLPcmcg==", "dependencies": [ @@ -247,6 +271,87 @@ "uint8arrays" ] }, + "@atproto/lexicon@0.5.1": { + "integrity": "sha512-y8AEtYmfgVl4fqFxqXAeGvhesiGkxiy3CWoJIfsFDDdTlZUC8DFnZrYhcqkIop3OlCkkljvpSJi1hbeC1tbi8A==", + "dependencies": [ + "@atproto/common-web", + "@atproto/syntax", + "iso-datestring-validator", + "multiformats@9.9.0", + "zod@3.25.76" + ] + }, + "@atproto/repo@0.8.10": { + "integrity": "sha512-REs6TZGyxNaYsjqLf447u+gSdyzhvMkVbxMBiKt1ouEVRkiho1CY32+omn62UkpCuGK2y6SCf6x3sVMctgmX4g==", + "dependencies": [ + "@atproto/common@0.4.12", + "@atproto/common-web", + "@atproto/crypto@0.4.4", + "@atproto/lexicon", + "@ipld/dag-cbor@7.0.3", + "multiformats@9.9.0", + "uint8arrays", + "varint", + "zod@3.25.76" + ] + }, + "@atproto/syntax@0.4.1": { + "integrity": "sha512-CJdImtLAiFO+0z3BWTtxwk6aY5w4t8orHTMVJgkf++QRJWTxPbIFko/0hrkADB7n2EruDxDSeAgfUGehpH6ngw==" + }, + "@atproto/xrpc-server@0.9.5": { + "integrity": "sha512-V0srjUgy6mQ5yf9+MSNBLs457m4qclEaWZsnqIE7RfYywvntexTAbMoo7J7ONfTNwdmA9Gw4oLak2z2cDAET4w==", + "dependencies": [ + "@atproto/common@0.4.12", + "@atproto/crypto@0.4.4", + "@atproto/lexicon", + "@atproto/xrpc", + "cbor-x", + "express", + "http-errors", + "mime-types", + "rate-limiter-flexible", + "uint8arrays", + "ws", + "zod@3.25.76" + ] + }, + "@atproto/xrpc@0.7.5": { + "integrity": "sha512-MUYNn5d2hv8yVegRL0ccHvTHAVj5JSnW07bkbiaz96UH45lvYNRVwt44z+yYVnb0/mvBzyD3/ZQ55TRGt7fHkA==", + "dependencies": [ + "@atproto/lexicon", + "zod@3.25.76" + ] + }, + "@cbor-extract/cbor-extract-darwin-arm64@2.2.0": { + "integrity": "sha512-P7swiOAdF7aSi0H+tHtHtr6zrpF3aAq/W9FXx5HektRvLTM2O89xCyXF3pk7pLc7QpaY7AoaE8UowVf9QBdh3w==", + "os": ["darwin"], + "cpu": ["arm64"] + }, + "@cbor-extract/cbor-extract-darwin-x64@2.2.0": { + "integrity": "sha512-1liF6fgowph0JxBbYnAS7ZlqNYLf000Qnj4KjqPNW4GViKrEql2MgZnAsExhY9LSy8dnvA4C0qHEBgPrll0z0w==", + "os": ["darwin"], + "cpu": ["x64"] + }, + "@cbor-extract/cbor-extract-linux-arm64@2.2.0": { + "integrity": "sha512-rQvhNmDuhjTVXSPFLolmQ47/ydGOFXtbR7+wgkSY0bdOxCFept1hvg59uiLPT2fVDuJFuEy16EImo5tE2x3RsQ==", + "os": ["linux"], + "cpu": ["arm64"] + }, + "@cbor-extract/cbor-extract-linux-arm@2.2.0": { + "integrity": "sha512-QeBcBXk964zOytiedMPQNZr7sg0TNavZeuUCD6ON4vEOU/25+pLhNN6EDIKJ9VLTKaZ7K7EaAriyYQ1NQ05s/Q==", + "os": ["linux"], + "cpu": ["arm"] + }, + "@cbor-extract/cbor-extract-linux-x64@2.2.0": { + "integrity": "sha512-cWLAWtT3kNLHSvP4RKDzSTX9o0wvQEEAj4SKvhWuOVZxiDAeQazr9A+PSiRILK1VYMLeDml89ohxCnUNQNQNCw==", + "os": ["linux"], + "cpu": ["x64"] + }, + "@cbor-extract/cbor-extract-win32-x64@2.2.0": { + "integrity": "sha512-l2M+Z8DO2vbvADOBNLbbh9y5ST1RY5sqkWOg/58GkUPBYou/cuNZ68SGQ644f1CvZ8kcOxyZtw06+dxWHIoN/w==", + "os": ["win32"], + "cpu": ["x64"] + }, "@did-plc/lib@0.0.4": { "integrity": "sha512-Omeawq3b8G/c/5CtkTtzovSOnWuvIuCI4GTJNrt1AmCskwEQV7zbX5d6km1mjJNbE0gHuQPTVqZxLVqetNbfwA==", "dependencies": [ @@ -386,6 +491,28 @@ "get-intrinsic" ] }, + "cbor-extract@2.2.0": { + "integrity": "sha512-Ig1zM66BjLfTXpNgKpvBePq271BPOvu8MR0Jl080yG7Jsl+wAZunfrwiwA+9ruzm/WEdIV5QF/bjDZTqyAIVHA==", + "dependencies": [ + "node-gyp-build-optional-packages" + ], + "optionalDependencies": [ + "@cbor-extract/cbor-extract-darwin-arm64", + "@cbor-extract/cbor-extract-darwin-x64", + "@cbor-extract/cbor-extract-linux-arm", + "@cbor-extract/cbor-extract-linux-arm64", + "@cbor-extract/cbor-extract-linux-x64", + "@cbor-extract/cbor-extract-win32-x64" + ], + "scripts": true, + "bin": true + }, + "cbor-x@1.6.0": { + "integrity": "sha512-0kareyRwHSkL6ws5VXHEf8uY1liitysCVJjlmhaLG+IXLqhSaOO+t63coaso7yjwEzWZzLy8fJo06gZDVQM9Qg==", + "optionalDependencies": [ + "cbor-extract" + ] + }, "cborg@1.10.2": { "integrity": "sha512-b3tFPA9pUr2zCUiCfRd2+wok2/LBSNUMKOuRRok+WlvvAgEt/PlbgPTsZUcwCOs53IJvLgTp0eotwtosE6njug==", "bin": true @@ -440,6 +567,9 @@ "destroy@1.2.0": { "integrity": "sha512-2sJGJTaXIIaR1w4iJSNoN0hnMY7Gpc/n8D4qSCJw8QqFWXf7cuAgnEHxBpweaVcPevC2l3KpjYCx3NypQQgaJg==" }, + "detect-libc@2.1.1": { + "integrity": "sha512-ecqj/sy1jcK1uWrwpR67UhYrIFQ+5WlGxth34WquCbamhFA6hkkwiu37o6J5xCHdo1oixJRfVRw+ywV+Hq/0Aw==" + }, "dunder-proto@1.0.1": { "integrity": "sha512-KIN/nDJBQRcXw0MLVhZE9iQHmG68qAVIBg9CqmUYjmQIhgij9U5MFvrqkUL5FbtyyzZuOeOt0zdeRe4UY7ct+A==", "dependencies": [ @@ -606,6 +736,9 @@ "gopd@1.2.0": { "integrity": "sha512-ZUKRh6/kUFoAiTAtTYPZJ3hw9wNxx+BIBOijnlG9PnrJsCcSjs1wyyD6vJpaYtgnzDrKYRSqf3OO6Rfa93xsRg==" }, + "graphemer@1.4.0": { + "integrity": "sha512-EtKwoO6kxCL9WO5xipiHTZlSzBm7WLT627TqC/uVRd0HKmq8NXyebnNYxDoBi7wt8eTWrUrKXCOVaFq9x1kgag==" + }, "has-symbols@1.1.0": { "integrity": "sha512-1cDNdwJ2Jaohmb3sg4OmKaMBwuC48sYni5HUw2DvsC8LjGTLK9h+eb1X6RyuOHe4hT0ULCW68iomhjUoKUqlPQ==" }, @@ -655,6 +788,9 @@ "ipaddr.js@1.9.1": { "integrity": "sha512-0KI/607xoxSToH7GjN1FfSbLoU0+btTicjsQSWQlh/hZykN8KpmMf7uYwPW3R+akZ6R/w18ZlXSHBYXiYUPO3g==" }, + "iso-datestring-validator@2.2.2": { + "integrity": "sha512-yLEMkBbLZTlVQqOnQ4FiMujR6T4DEcCb1xizmvXS+OxuhwcbtynoosRzdMA69zZCShCNAbi+gJ71FxZBBXx1SA==" + }, "kysely@0.23.5": { "integrity": "sha512-TH+b56pVXQq0tsyooYLeNfV11j6ih7D50dyN8tkM0e7ndiUH28Nziojiog3qRFlmEj9XePYdZUrNJ2079Qjdow==" }, @@ -698,6 +834,13 @@ "negotiator@0.6.3": { "integrity": "sha512-+EUsqGPLsM+j/zdChZjsnX51g4XrHFOIXwfnCVPGlQk/k5giakcKsuxCObBRu6DSm9opw/O6slWbJdghQM4bBg==" }, + "node-gyp-build-optional-packages@5.1.1": { + "integrity": "sha512-+P72GAjVAbTxjjwUmwjVrqrdZROD4nf8KgpBoDxqXXTiYZZt/ud60dE5yvCSr9lRO8e8yv6kgJIC0K0PfZFVQw==", + "dependencies": [ + "detect-libc" + ], + "bin": true + }, "object-assign@4.1.1": { "integrity": "sha512-rJgTQnkUnH1sFw8yT6VSU3zD3sWmu6sZhIseY8VX+GRu3P6F7Fu+JNDoXfklElbLJSnc3FUQHVe4cU5hj+BcUg==" }, @@ -1040,6 +1183,9 @@ "utils-merge@1.0.1": { "integrity": "sha512-pMZTvIkT1d+TFGvDOqodOclx0QWkkgi6Tdoa8gC8ffGAAqz9pzPTZWAybbsHHoED/ztMtkv/VoYTYyShUn81hA==" }, + "varint@6.0.0": { + "integrity": "sha512-cXEIW6cfr15lFv563k4GuVuW/fiwjknytD37jIOLSdSWuOI6WnO/oKwmP2FQTU2l01LP8/M5TSAJpzUaGe3uWg==" + }, "vary@1.1.2": { "integrity": "sha512-BNGbWLfd0eUPabhkXUVm0j8uuvREyTh5ovRa/dyow/BqAbZJyC+5fU+IzQOzmAKzYqYRAISoRhdQr3eIZ/PXqg==" }, diff --git a/repo/sync/consumer.ts b/repo/sync/consumer.ts index 8e64f23..b60aeab 100644 --- a/repo/sync/consumer.ts +++ b/repo/sync/consumer.ts @@ -148,7 +148,7 @@ export const verifyProofs = async ( const verified: RecordCidClaim[] = []; const unverified: RecordCidClaim[] = []; for (const claim of claims) { - const found = await mst.get( + const found = mst.get( util.formatDataKey(claim.collection, claim.rkey), ); const record = found ? blockstore.readObj(found, def.map) : null; diff --git a/sync/tests/mock-firehose-server.ts b/sync/tests/mock-relay.ts similarity index 100% rename from sync/tests/mock-firehose-server.ts rename to sync/tests/mock-relay.ts diff --git a/xrpc-server/stream/stream.ts b/xrpc-server/stream/stream.ts index 1ad4409..404f06b 100644 --- a/xrpc-server/stream/stream.ts +++ b/xrpc-server/stream/stream.ts @@ -1,172 +1,36 @@ +import type { DuplexOptions } from "node:stream"; +import { createWebSocketStream, type WebSocket } from "ws"; import { ResponseType, XRPCError } from "@atp/xrpc"; -import { Frame } from "./frames.ts"; -import type { MessageFrame } from "./frames.ts"; +import { Frame, type MessageFrame } from "./frames.ts"; -/** - * Converts a WebSocket connection into an async generator of Frame objects. - * Handles both message and error frames, with proper error propagation. - * - * @param ws - The WebSocket connection to read from - * @yields {Frame} Each frame received from the WebSocket - * @throws Any WebSocket error that occurs during communication - * - * @example - * ```typescript - * const ws = new WebSocket(url); - * for await (const frame of byFrame(ws)) { - * // Process each frame - * console.log(frame.type); - * } - * ``` - */ -export async function* byFrame( - ws: WebSocket, -): AsyncGenerator { - // Wait for connection if still connecting - if (ws.readyState === WebSocket.CONNECTING) { - await new Promise((resolve, reject) => { - const onOpen = () => { - ws.removeEventListener("open", onOpen); - ws.removeEventListener("error", onError); - resolve(); - }; - - const onError = (event: Event | ErrorEvent) => { - ws.removeEventListener("open", onOpen); - ws.removeEventListener("error", onError); - const error = event instanceof ErrorEvent && event.error - ? event.error - : new Error("WebSocket connection failed"); - reject(error); - }; - - ws.addEventListener("open", onOpen); - ws.addEventListener("error", onError); - }); - } - - // If already closed, return immediately - if (ws.readyState === WebSocket.CLOSED) { - return; - } - - // Process messages until connection closes - while (ws.readyState === WebSocket.OPEN) { - try { - const frame = await waitForNextFrame(ws); - if (frame) { - yield frame; - } else { - // Connection closed normally - break; - } - } catch (error) { - // WebSocket error occurred - throw error; - } - } +export function streamByteChunks(ws: WebSocket, options?: DuplexOptions) { + return createWebSocketStream(ws, { + ...options, + readableObjectMode: true, // Ensures frame bytes don't get buffered/combined together + }); } -function waitForNextFrame(ws: WebSocket): Promise { - return new Promise((resolve, reject) => { - const cleanup = () => { - ws.removeEventListener("message", onMessage); - ws.removeEventListener("error", onError); - ws.removeEventListener("close", onClose); - }; - - const onMessage = async (event: MessageEvent) => { - cleanup(); - try { - let data: Uint8Array; - if (event.data instanceof Uint8Array) { - data = event.data; - } else if (event.data instanceof Blob) { - data = new Uint8Array(await event.data.arrayBuffer()); - } else { - // Ignore non-binary data (e.g., ping/pong) - // Re-attach listeners and wait for next message - attachListeners(); - return; - } - - const frame = Frame.fromBytes(data); - resolve(frame); - } catch (error) { - reject(error instanceof Error ? error : new Error(String(error))); - } - }; - - const onError = (event: Event | ErrorEvent) => { - cleanup(); - const error = event instanceof ErrorEvent && event.error - ? event.error - : new Error("WebSocket error"); - reject(error); - }; - - const onClose = () => { - cleanup(); - resolve(null); // Signal end of stream - }; - - const attachListeners = () => { - ws.addEventListener("message", onMessage, { once: true }); - ws.addEventListener("error", onError, { once: true }); - ws.addEventListener("close", onClose, { once: true }); - }; - - // Check if connection is already closed before attaching listeners - if (ws.readyState === WebSocket.CLOSED) { - resolve(null); - return; - } - - attachListeners(); - }); +export async function* byFrame(ws: WebSocket, options?: DuplexOptions) { + const wsStream = streamByteChunks(ws, options); + for await (const chunk of wsStream) { + yield Frame.fromBytes(chunk); + } } -/** - * Converts a WebSocket connection into an async generator of MessageFrames. - * Automatically filters and validates frames to ensure they are valid messages. - * Error frames are converted to exceptions. - * - * @param ws - The WebSocket connection to read from - * @yields Each message frame received from the WebSocket - * @throws If an error frame is received or an invalid frame type is encountered - * - * @example - * ```typescript - * const ws = new WebSocket(url); - * for await (const message of byMessage(ws)) { - * // Process each message - * console.log(message.body); - * } - * ``` - */ -export async function* byMessage( - ws: WebSocket, -): AsyncGenerator> { - for await (const frame of byFrame(ws)) { - yield ensureChunkIsMessage(frame); +export async function* byMessage(ws: WebSocket, options?: DuplexOptions) { + const wsStream = streamByteChunks(ws, options); + for await (const chunk of wsStream) { + const msg = ensureChunkIsMessage(chunk); + yield msg; } } -/** - * Validates that a frame is a MessageFrame and converts it to the appropriate type. - * If the frame is an error frame, throws an XRPCError with the error details. - * - * @param frame - The frame to validate - * @returns The frame as a MessageFrame if valid - * @throws If the frame is an error frame or an invalid type - * @internal - */ -export function ensureChunkIsMessage(frame: Frame): MessageFrame { +export function ensureChunkIsMessage(chunk: Uint8Array): MessageFrame { + const frame = Frame.fromBytes(chunk); if (frame.isMessage()) { return frame; } else if (frame.isError()) { - // @TODO work -1 error code into XRPCError - throw new XRPCError(3, frame.code, frame.message); + throw new XRPCError(-1, frame.code, frame.message); } else { throw new XRPCError(ResponseType.Unknown, undefined, "Unknown frame type"); } diff --git a/xrpc-server/stream/subscription.ts b/xrpc-server/stream/subscription.ts index e434173..8578663 100644 --- a/xrpc-server/stream/subscription.ts +++ b/xrpc-server/stream/subscription.ts @@ -1,28 +1,10 @@ +import type { ClientOptions } from "ws"; import { ensureChunkIsMessage } from "./stream.ts"; import { WebSocketKeepAlive } from "./websocket-keepalive.ts"; -import { Frame } from "./frames.ts"; -import type { WebSocketOptions } from "./types.ts"; -/** - * Represents a message body in a subscription stream. - * @interface - * @property $type - Optional type identifier for the message - * @property [key] - Additional message properties - */ -interface MessageBody { - $type?: string; - [key: string]: unknown; -} - -/** - * Represents a subscription to an XRPC streaming endpoint. - * Handles WebSocket connection management, reconnection, and message parsing. - * @class - * @template T - The type of messages yielded by the subscription - */ export class Subscription { constructor( - public opts: WebSocketOptions & { + public opts: ClientOptions & { service: string; method: string; maxReconnectSeconds?: number; @@ -51,14 +33,18 @@ export class Subscription { }, }); for await (const chunk of ws) { - const frame = Frame.fromBytes(chunk); - const message = ensureChunkIsMessage(frame); + const message = ensureChunkIsMessage(chunk); const t = message.header.t; const clone = message.body !== undefined - ? { ...message.body } as MessageBody + ? { ...message.body } : undefined; - if (clone !== undefined && t !== undefined) { - clone.$type = t.startsWith("#") ? this.opts.method + t : t; + if ( + clone !== undefined && t !== undefined && + clone as Record["$type"] !== undefined + ) { + (clone as Record)["$type"] = t.startsWith("#") + ? this.opts.method + t + : t; } const result = this.opts.validate(clone); if (result !== undefined) { @@ -83,6 +69,7 @@ function encodeQueryParams(obj: Record): string { return params.toString(); } +// Adapted from xrpc, but without any lex-specific knowledge function encodeQueryParam(value: unknown): string | string[] { if (typeof value === "string") { return value; diff --git a/xrpc-server/stream/websocket-keepalive.ts b/xrpc-server/stream/websocket-keepalive.ts index ca2ddad..9d61eec 100644 --- a/xrpc-server/stream/websocket-keepalive.ts +++ b/xrpc-server/stream/websocket-keepalive.ts @@ -1,18 +1,15 @@ +import { type ClientOptions, WebSocket } from "ws"; import { SECOND, wait } from "@atp/common"; -import { CloseCode, DisconnectError, type WebSocketOptions } from "./types.ts"; +import { streamByteChunks } from "./stream.ts"; +import { CloseCode, DisconnectError } from "./types.ts"; -/** - * WebSocket client with automatic reconnection and heartbeat functionality. - * Handles connection management, reconnection backoff, and keep-alive messages. - * @class - */ export class WebSocketKeepAlive { public ws: WebSocket | null = null; public initialSetup = true; public reconnects: number | null = null; constructor( - public opts: WebSocketOptions & { + public opts: ClientOptions & { getUrl: () => Promise; maxReconnectSeconds?: number; signal?: AbortSignal; @@ -35,129 +32,36 @@ export class WebSocketKeepAlive { await wait(duration); } const url = await this.opts.getUrl(); - this.ws = new WebSocket(url, this.opts.protocols); + this.ws = new WebSocket(url, this.opts); const ac = new AbortController(); if (this.opts.signal) { forwardSignal(this.opts.signal, ac); } - this.ws.onopen = () => { + this.ws.once("open", () => { this.initialSetup = false; this.reconnects = 0; if (this.ws) { this.startHeartbeat(this.ws); } - }; - this.ws.onclose = (ev: CloseEvent) => { - if (ev.code === CloseCode.Abnormal) { + }); + this.ws.once("close", (code: number, reason: Uint8Array) => { + if (code === CloseCode.Abnormal) { // Forward into an error to distinguish from a clean close ac.abort( - new AbnormalCloseError(`Abnormal ws close: ${ev.reason}`), + new AbnormalCloseError(`Abnormal ws close: ${reason.toString()}`), ); } - }; + }); try { - const messageQueue: Uint8Array[] = []; - let error: Error | null = null; - let finished = false; - let resolveNext: (() => void) | null = null; - - const processMessage = (ev: MessageEvent) => { - if (ev.data === "pong") { - // Handle heartbeat pong responses separately - return; - } - if (ev.data instanceof Uint8Array) { - messageQueue.push(ev.data); - if (resolveNext) { - resolveNext(); - resolveNext = null; - } - } - }; - - const handleError = (ev: Event | ErrorEvent) => { - error = ev instanceof ErrorEvent && ev.error - ? ev.error - : new Error("WebSocket error"); - if (resolveNext) { - resolveNext(); - resolveNext = null; - } - }; - - const handleClose = () => { - finished = true; - if (resolveNext) { - resolveNext(); - resolveNext = null; - } - }; - - this.ws.onmessage = processMessage; - this.ws.onerror = handleError; - this.ws.onclose = handleClose; - - // Wait for connection if still connecting - if (this.ws.readyState === WebSocket.CONNECTING) { - await new Promise((resolve, reject) => { - const onOpen = () => { - this.ws!.removeEventListener("open", onOpen); - this.ws!.removeEventListener("error", onInitialError); - resolve(); - }; - - const onInitialError = (ev: Event | ErrorEvent) => { - this.ws!.removeEventListener("open", onOpen); - this.ws!.removeEventListener("error", onInitialError); - const errorMsg = ev instanceof ErrorEvent && ev.error - ? ev.error - : new Error("Failed to connect to WebSocket"); - reject(errorMsg); - }; - - this.ws!.addEventListener("open", onOpen, { once: true }); - this.ws!.addEventListener("error", onInitialError, { once: true }); - }); + const wsStream = streamByteChunks(this.ws, { signal: ac.signal }); + for await (const chunk of wsStream) { + yield chunk; } - - // Main message processing loop - while (!finished && !error && !ac.signal.aborted) { - // Process any queued messages first - while (messageQueue.length > 0) { - yield messageQueue.shift()!; - } - - // If no messages and not finished, wait for next event - if ( - !finished && !error && !ac.signal.aborted && - messageQueue.length === 0 - ) { - await new Promise((resolve) => { - resolveNext = resolve; - // Also resolve if abort signal is triggered - if (ac.signal.aborted) { - resolve(); - } else { - ac.signal.addEventListener("abort", () => resolve(), { - once: true, - }); - } - }); - } - } - - // Process any remaining messages - while (messageQueue.length > 0) { - yield messageQueue.shift()!; - } - - if (error) throw error; - if (ac.signal.aborted) throw ac.signal.reason; - } catch (_err) { - const err = isErrorWithCode(_err) && _err.code === "ABORT_ERR" - ? _err.cause - : _err; + } catch (error) { + const err = (error as Record)?.["code"] === "ABORT_ERR" + ? (error as Record)["cause"] + : error; if (err instanceof DisconnectError) { // We cleanly end the connection this.ws?.close(err.wsCode); @@ -178,47 +82,31 @@ export class WebSocketKeepAlive { startHeartbeat(ws: WebSocket) { let isAlive = true; - let heartbeatInterval: ReturnType | null = null; + let heartbeatInterval: number | null = null; const checkAlive = () => { if (!isAlive) { - return ws.close(); + return ws.terminate(); } isAlive = false; // expect websocket to no longer be alive unless we receive a "pong" within the interval - ws.send("ping"); + ws.ping(); }; - // Store original handlers to chain them properly - const originalOnMessage = ws.onmessage; - const originalOnClose = ws.onclose; - checkAlive(); heartbeatInterval = setInterval( checkAlive, this.opts.heartbeatIntervalMs ?? 10 * SECOND, ); - // Chain message handler to handle pong responses - ws.onmessage = (ev: MessageEvent) => { - if (ev.data === "pong") { - isAlive = true; - } - // Always call the original handler for all messages - if (originalOnMessage) { - originalOnMessage.call(ws, ev); - } - }; - - // Chain close handler to clean up heartbeat - ws.onclose = (ev: CloseEvent) => { + ws.on("pong", () => { + isAlive = true; + }); + ws.once("close", () => { if (heartbeatInterval) { clearInterval(heartbeatInterval); heartbeatInterval = null; } - if (originalOnClose) { - originalOnClose.call(ws, ev); - } - }; + }); } } @@ -228,35 +116,17 @@ class AbnormalCloseError extends Error { code = "EWSABNORMALCLOSE"; } -/** - * Interface for errors with error codes. - * @interface - * @property {string} [code] - Error code identifier - * @property {unknown} [cause] - Underlying cause of the error - */ -interface ErrorWithCode { - code?: string; - cause?: unknown; -} - -/** - * Type guard to check if an error has an error code. - * @param {unknown} err - The error to check - * @returns {boolean} True if the error has a code property - */ -function isErrorWithCode(err: unknown): err is ErrorWithCode { - return err !== null && typeof err === "object" && "code" in err; -} - function isReconnectable(err: unknown): boolean { - if (!isErrorWithCode(err)) return false; - return typeof err.code === "string" && networkErrorCodes.includes(err.code); + // Network errors are reconnectable. + // AuthenticationRequired and InvalidRequest XRPCErrors are not reconnectable. + // @TODO method-specific XRPCErrors may be reconnectable, need to consider. Receiving + // an invalid message is not current reconnectable, but the user can decide to skip them. + if (!err || typeof err as Record["code"] !== "string") { + return false; + } + return networkErrorCodes.includes((err as Record)["code"]); } -/** - * List of error codes that indicate network-related issues. - * These errors typically warrant a reconnection attempt. - */ const networkErrorCodes = [ "EWSABNORMALCLOSE", "ECONNRESET",