From e1fcb90e19255f844154d5c23c37f85193da0e14 Mon Sep 17 00:00:00 2001 From: Alois Date: Sun, 5 Jul 2026 21:45:43 +0200 Subject: [PATCH] (feat): crypto migrations --- example/type-maps.yaml | 6 + src/sdk/encrypted-device-secret.ts | 74 +++ src/sdk/encrypted-message.ts | 351 ++++++++++++ src/sdk/index.ts | 831 ++++++++++++++++++++++++++--- src/sdk/ratchet.ts | 70 +++ src/sdk/session.ts | 249 +++++++++ test/e2ee.mjs | 503 +++++++++++++++++ tsconfig.json | 8 +- type-map/build.rs | 23 +- 9 files changed, 2032 insertions(+), 83 deletions(-) create mode 100644 src/sdk/encrypted-device-secret.ts create mode 100644 src/sdk/encrypted-message.ts create mode 100644 src/sdk/ratchet.ts create mode 100644 src/sdk/session.ts create mode 100644 test/e2ee.mjs diff --git a/example/type-maps.yaml b/example/type-maps.yaml index 0289011..6e411e4 100644 --- a/example/type-maps.yaml +++ b/example/type-maps.yaml @@ -6,6 +6,7 @@ type_maps: DataTypes: "1.0": CommunicationTypes: + CommunicationType: 32 DataTypes: Data: 32 Flags: 33 @@ -15,8 +16,11 @@ type_maps: EncryptedPayload: 37 SignedPayload: 38 SecurePayload: 39 + CommunicationType: 40 + DataType: 41 "2.0": CommunicationTypes: + CommunicationType: 32 DataTypes: Data: 34 Flags: 33 @@ -26,3 +30,5 @@ type_maps: EncryptedPayload: 38 SignedPayload: 39 SecurePayload: 40 + CommunicationType: 41 + DataType: 42 diff --git a/src/sdk/encrypted-device-secret.ts b/src/sdk/encrypted-device-secret.ts new file mode 100644 index 0000000..bc4587d --- /dev/null +++ b/src/sdk/encrypted-device-secret.ts @@ -0,0 +1,74 @@ +export interface EncryptedDeviceSecretRecord { + userId: string; + deviceId: string; + secretId: string; + version: number; + encryptedSecret: Uint8Array; + wrappingPublicKeyId?: string; + wrappingScheme: string; + createdAt: number; + updatedAt: number; +} + +export interface MTPEncryptedDeviceSecretProvider { + setEncryptedDeviceSecret(record: EncryptedDeviceSecretRecord): Promise; + getEncryptedDeviceSecret(query: { + userId: string; + deviceId?: string; + secretId?: string; + }): Promise; +} + +function keyFor(record: Pick): string { + return `${record.userId}\0${record.deviceId}\0${record.secretId}`; +} + +function cloneRecord(record: EncryptedDeviceSecretRecord): EncryptedDeviceSecretRecord { + return { + ...record, + encryptedSecret: new Uint8Array(record.encryptedSecret), + }; +} + +function validateEncryptedRecord(record: EncryptedDeviceSecretRecord): void { + if (!record.userId || !record.deviceId || !record.secretId) { + throw new Error("encrypted device secret requires userId, deviceId, and secretId"); + } + if (!(record.encryptedSecret instanceof Uint8Array) || record.encryptedSecret.length === 0) { + throw new Error("encrypted device secret requires non-empty encryptedSecret bytes"); + } + if (!record.wrappingScheme) { + throw new Error("encrypted device secret requires wrappingScheme"); + } +} + +export class InMemoryEncryptedDeviceSecretProvider implements MTPEncryptedDeviceSecretProvider { + private store = new Map(); + + async setEncryptedDeviceSecret(record: EncryptedDeviceSecretRecord): Promise { + validateEncryptedRecord(record); + const now = Date.now(); + this.store.set(keyFor(record), cloneRecord({ ...record, updatedAt: record.updatedAt || now })); + } + + async getEncryptedDeviceSecret(query: { + userId: string; + deviceId?: string; + secretId?: string; + }): Promise { + if (!query.userId) { + throw new Error("userId is required"); + } + if (query.deviceId && query.secretId) { + const found = this.store.get(`${query.userId}\0${query.deviceId}\0${query.secretId}`); + return found ? cloneRecord(found) : null; + } + for (const record of this.store.values()) { + if (record.userId !== query.userId) continue; + if (query.deviceId && record.deviceId !== query.deviceId) continue; + if (query.secretId && record.secretId !== query.secretId) continue; + return cloneRecord(record); + } + return null; + } +} diff --git a/src/sdk/encrypted-message.ts b/src/sdk/encrypted-message.ts new file mode 100644 index 0000000..3f816ea --- /dev/null +++ b/src/sdk/encrypted-message.ts @@ -0,0 +1,351 @@ +import * as bindings from "mtp/raw"; +import { MTPRatchet } from "./ratchet.js"; +import type { MTPSessionState } from "./session"; + +export const MTP_E2EE_VERSION = 1; +export const FLAG_INIT = 0x01; +export const FLAG_DEVICE_SECRET = 0x02; +export const FLAG_KEY_ROTATION = 0x04; +export const MAX_RATCHET_SKIP = 100; +const SUPPORTED_FLAGS = FLAG_INIT | FLAG_DEVICE_SECRET | FLAG_KEY_ROTATION; +const HEADER_FIXED_LEN = 1 + 1 + 8 + 8 + 4 + 2 + 4; + +export interface ParsedEncryptedMessage { + version: 1; + flags: number; + senderClientId: bigint; + recipientClientId: bigint; + messageNumber: number; + kemCiphertext?: Uint8Array; + ciphertext: Uint8Array; + /** Compatibility alias for older SDK tests/callers. */ + header?: EncryptedMessageHeader; + /** Compatibility alias for older SDK tests/callers. */ + aeadPayload?: Uint8Array; +} + +export interface EncryptedMessageHeader { + version: 1; + flags: number; + senderClientId: bigint; + recipientClientId: bigint; + messageNumber: number; + kemCiphertext?: Uint8Array; +} +export interface SerializedEncryptedMessage { + header: EncryptedMessageHeader; + aeadPayload: Uint8Array; +} + +function writeU64BE(value: bigint): Uint8Array { + if (value < 0n || value > 0xffff_ffff_ffff_ffffn) { + throw new Error("u64 value out of range"); + } + const buf = new Uint8Array(8); + for (let i = 7; i >= 0; i--) { + buf[i] = Number(value & 0xffn); + value >>= 8n; + } + return buf; +} + +function readU64BE(bytes: Uint8Array, offset: number): bigint { + let value = 0n; + for (let i = 0; i < 8; i++) { + value = (value << 8n) | BigInt(bytes[offset + i]); + } + return value; +} + +function writeU32BE(value: number): Uint8Array { + if (!Number.isSafeInteger(value) || value < 0 || value > 0xffff_ffff) { + throw new Error("u32 value out of range"); + } + return new Uint8Array([ + (value >>> 24) & 0xff, + (value >>> 16) & 0xff, + (value >>> 8) & 0xff, + value & 0xff, + ]); +} + +function concatBytes(parts: Uint8Array[]): Uint8Array { + const out = new Uint8Array(parts.reduce((sum, part) => sum + part.length, 0)); + let offset = 0; + for (const part of parts) { + out.set(part, offset); + offset += part.length; + } + return out; +} + +function assertSupported(message: ParsedEncryptedMessage): void { + if (message.version !== MTP_E2EE_VERSION) { + throw new Error( + `Unsupported encrypted message version: ${message.version}`, + ); + } + if ((message.flags & ~SUPPORTED_FLAGS) !== 0) { + throw new Error(`Unsupported encrypted message flags: ${message.flags}`); + } + const isInit = (message.flags & FLAG_INIT) !== 0; + if (isInit && !message.kemCiphertext?.length) { + throw new Error("Init message must include KEM ciphertext"); + } + if (!isInit && message.kemCiphertext?.length) { + throw new Error("Non-init message must not include KEM ciphertext"); + } + if (!message.ciphertext.length) { + throw new Error("Encrypted message ciphertext must be non-empty"); + } +} + +export function serializeEncryptedMessage( + message: ParsedEncryptedMessage, +): Uint8Array; +export function serializeEncryptedMessage( + message: SerializedEncryptedMessage, +): Uint8Array; +export function serializeEncryptedMessage( + message: ParsedEncryptedMessage | SerializedEncryptedMessage, +): Uint8Array { + const normalized: ParsedEncryptedMessage = + "header" in message + ? { + ...message.header, + ciphertext: message.aeadPayload, + } + : message; + + assertSupported(normalized); + const kemCiphertext = normalized.kemCiphertext ?? new Uint8Array(0); + if (kemCiphertext.length > 0xffff) { + throw new Error("KEM ciphertext too long"); + } + + return concatBytes([ + new Uint8Array([normalized.version]), + new Uint8Array([normalized.flags]), + writeU64BE(normalized.senderClientId), + writeU64BE(normalized.recipientClientId), + writeU32BE(normalized.messageNumber), + new Uint8Array([ + (kemCiphertext.length >>> 8) & 0xff, + kemCiphertext.length & 0xff, + ]), + kemCiphertext, + writeU32BE(normalized.ciphertext.length), + normalized.ciphertext, + ]); +} + +export function parseEncryptedMessage( + bytes: Uint8Array, +): ParsedEncryptedMessage { + let offset = 0; + if (!(bytes instanceof Uint8Array)) { + throw new Error("Encrypted message must be bytes"); + } + if (bytes.length < HEADER_FIXED_LEN) { + throw new Error("Encrypted message too short"); + } + + const version = bytes[offset++]; + const flags = bytes[offset++]; + const senderClientId = readU64BE(bytes, offset); + offset += 8; + const recipientClientId = readU64BE(bytes, offset); + offset += 8; + const messageNumber = + ((bytes[offset] << 24) | + (bytes[offset + 1] << 16) | + (bytes[offset + 2] << 8) | + bytes[offset + 3]) >>> + 0; + offset += 4; + const kemLen = (bytes[offset] << 8) | bytes[offset + 1]; + offset += 2; + + let kemCiphertext: Uint8Array | undefined; + if (kemLen > 0) { + if (bytes.length < offset + kemLen + 4) { + throw new Error("Encrypted message KEM ciphertext truncated"); + } + kemCiphertext = bytes.slice(offset, offset + kemLen); + offset += kemLen; + } + + if (bytes.length < offset + 4) { + throw new Error("Encrypted message missing ciphertext length"); + } + const ciphertextLen = + ((bytes[offset] << 24) | + (bytes[offset + 1] << 16) | + (bytes[offset + 2] << 8) | + bytes[offset + 3]) >>> + 0; + offset += 4; + if (bytes.length < offset + ciphertextLen) { + throw new Error("Encrypted message ciphertext truncated"); + } + const ciphertext = bytes.slice(offset, offset + ciphertextLen); + offset += ciphertextLen; + if (offset !== bytes.length) { + throw new Error("Encrypted message has trailing data"); + } + + const parsed: ParsedEncryptedMessage = { + version: version as 1, + flags, + senderClientId, + recipientClientId, + messageNumber, + kemCiphertext, + ciphertext, + }; + parsed.header = { + version: parsed.version, + flags: parsed.flags, + senderClientId: parsed.senderClientId, + recipientClientId: parsed.recipientClientId, + messageNumber: parsed.messageNumber, + kemCiphertext: parsed.kemCiphertext, + }; + parsed.aeadPayload = parsed.ciphertext; + assertSupported(parsed); + return parsed; +} + +function buildAAD(header: EncryptedMessageHeader): Uint8Array { + return concatBytes([ + new Uint8Array([header.version]), + new Uint8Array([header.flags]), + writeU64BE(header.senderClientId), + writeU64BE(header.recipientClientId), + writeU32BE(header.messageNumber), + ]); +} + +export function encryptedMessageAAD( + header: EncryptedMessageHeader, + extra?: Uint8Array, +): Uint8Array { + return extra?.length + ? concatBytes([buildAAD(header), extra]) + : buildAAD(header); +} + +export async function encryptPayload(args: { + plaintext: Uint8Array; + session: MTPSessionState; + kemCiphertext?: Uint8Array; + aad?: Uint8Array; +}): Promise<{ + payload: Uint8Array; + session: MTPSessionState; +}> { + const step = await MTPRatchet.stepSend(args.session.sendChainKey); + const header: EncryptedMessageHeader = { + version: 1, + flags: args.kemCiphertext ? FLAG_INIT : 0, + senderClientId: args.session.ownClientId, + recipientClientId: args.session.peerClientId, + messageNumber: args.session.sendCount, + kemCiphertext: args.kemCiphertext, + }; + const aad = args.aad ?? encryptedMessageAAD(header); + const cipher = new bindings.WasmChaCha20Poly1305(step.key); + let ciphertext: Uint8Array; + try { + ciphertext = cipher.encrypt(args.plaintext, aad); + } finally { + cipher.free(); + step.key.fill(0); + } + + const payload = serializeEncryptedMessage({ ...header, ciphertext }); + return { + payload, + session: { + ...args.session, + sendChainKey: step.chainKey, + sendCount: args.session.sendCount + 1, + updatedAt: Date.now(), + }, + }; +} + +export async function decryptPayload(args: { + payload: Uint8Array; + session: MTPSessionState; + expectedRecipientClientId?: bigint; + aad?: Uint8Array; +}): Promise<{ + plaintext: Uint8Array; + session: MTPSessionState; +}> { + const parsed = parseEncryptedMessage(args.payload); + const expectedRecipientClientId = + args.expectedRecipientClientId ?? args.session.ownClientId; + if (parsed.recipientClientId !== expectedRecipientClientId) { + throw new Error("Encrypted message recipient mismatch"); + } + if (parsed.senderClientId !== args.session.peerClientId) { + throw new Error("Encrypted message sender mismatch"); + } + if (parsed.messageNumber < args.session.recvCount) { + throw new Error("Encrypted message replay or out-of-order message number"); + } + + let chainKey = args.session.recvChainKey; + let messageKey: Uint8Array | undefined; + const gap = parsed.messageNumber - args.session.recvCount; + if (gap > MAX_RATCHET_SKIP) { + throw new Error( + `Encrypted message receive gap exceeds max skip (${MAX_RATCHET_SKIP})`, + ); + } + + const steps = gap + 1; + for (let i = 0; i < steps; i += 1) { + const step = await MTPRatchet.stepRecv(chainKey); + if (i === steps - 1) { + messageKey = step.key; + } else { + step.key.fill(0); + } + if (chainKey !== args.session.recvChainKey) chainKey.fill(0); + chainKey = step.chainKey; + } + if (!messageKey) { + throw new Error("Failed to derive receive message key"); + } + + const header: EncryptedMessageHeader = { + version: parsed.version, + flags: parsed.flags, + senderClientId: parsed.senderClientId, + recipientClientId: parsed.recipientClientId, + messageNumber: parsed.messageNumber, + kemCiphertext: parsed.kemCiphertext, + }; + const aad = args.aad ?? encryptedMessageAAD(header); + const cipher = new bindings.WasmChaCha20Poly1305(messageKey); + let plaintext: Uint8Array; + try { + plaintext = cipher.decrypt(parsed.ciphertext, aad); + } finally { + cipher.free(); + messageKey.fill(0); + } + + return { + plaintext, + session: { + ...args.session, + recvChainKey: chainKey, + recvCount: parsed.messageNumber + 1, + updatedAt: Date.now(), + }, + }; +} diff --git a/src/sdk/index.ts b/src/sdk/index.ts index 4c0e19c..b7a7331 100644 --- a/src/sdk/index.ts +++ b/src/sdk/index.ts @@ -7,6 +7,24 @@ import initWasm, { import * as bindings from "mtp/raw"; import type * as RawBindings from "../raw/index"; import type { MTPCommunicationType } from "../type-map/index"; +import type { MTPSessionStorage, MTPSessionState } from "./session"; +import { + MTPSessionManager, + getConversationId, + deriveSessionKeys, +} from "./session.js"; +import type { + EncryptedDeviceSecretRecord, + MTPEncryptedDeviceSecretProvider, +} from "./encrypted-device-secret"; +import { InMemoryEncryptedDeviceSecretProvider } from "./encrypted-device-secret.js"; +import { InMemorySessionStorage } from "./session.js"; +import { + parseEncryptedMessage, + encryptPayload, + decryptPayload, + FLAG_INIT, +} from "./encrypted-message.js"; export type StorageValue = string | null; @@ -19,12 +37,25 @@ export interface MTPCredentialStorage { export type MTPStorage = MTPCredentialStorage; export type MTPLogEvent = - | { hint: "info" | "warning"; type: string; data: unknown; direction?: "send" | "recv" } - | { hint: "error"; type: string | "error"; error: string; data?: unknown; direction?: "send" | "recv" }; + | { + hint: "info" | "warning"; + type: string; + data: unknown; + direction?: "send" | "recv"; + } + | { + hint: "error"; + type: string | "error"; + error: string; + data?: unknown; + direction?: "send" | "recv"; + }; export type ParsedFrame = RawBindings.ParsedFrame; -export type Ed25519GenerateResult = ReturnType; +export type Ed25519GenerateResult = ReturnType< + typeof bindings.ed25519_generate +>; export type WasmEncapsulated = RawBindings.WasmEncapsulated; @@ -32,12 +63,26 @@ export interface MTPCrypto { generateKeyring(): Uint8Array; generateEd25519(): Ed25519GenerateResult; keyringFromEd25519(secretKey: Uint8Array, publicKey: Uint8Array): Uint8Array; - verifyEd25519(publicKey: Uint8Array, message: Uint8Array, signature: Uint8Array): void; - deriveEncryptionKey(ikm: Uint8Array, salt: Uint8Array, context: Uint8Array): Uint8Array; - hkdfExpand(ikm: Uint8Array, salt: Uint8Array, info: Uint8Array, len: number): Uint8Array; + verifyEd25519( + publicKey: Uint8Array, + message: Uint8Array, + signature: Uint8Array, + ): void; + deriveEncryptionKey( + ikm: Uint8Array, + salt: Uint8Array, + context: Uint8Array, + ): Uint8Array; + hkdfExpand( + ikm: Uint8Array, + salt: Uint8Array, + info: Uint8Array, + len: number, + ): Uint8Array; sha256(data: Uint8Array): Uint8Array; sha256Double(data: Uint8Array): Uint8Array; keyringToKeys(keyring: string | MTPBytesInput): MTPKeyringKeys; + publicKeyBundleToKeys(publicKeyBundle: string | MTPBytesInput): MTPPublicKeyBundleKeys; encrypt(key: Uint8Array, input: Uint8Array): Promise; decrypt(key: Uint8Array, input: Uint8Array): Promise; encryptText(key: Uint8Array, plaintext: string): Promise; @@ -49,13 +94,19 @@ export interface MTPCrypto { export const crypto: MTPCrypto = { generateKeyring: () => bindings.keyring_generate(), generateEd25519: () => bindings.ed25519_generate(), - keyringFromEd25519: (secretKey, publicKey) => bindings.keyring_from_ed25519(secretKey, publicKey), - verifyEd25519: (publicKey, message, signature) => bindings.ed25519_verify(publicKey, message, signature), - deriveEncryptionKey: (ikm, salt, context) => bindings.wasm_derive_encryption_key(ikm, salt, context), - hkdfExpand: (ikm, salt, info, len) => bindings.wasm_hkdf_expand(ikm, salt, info, len), + keyringFromEd25519: (secretKey, publicKey) => + bindings.keyring_from_ed25519(secretKey, publicKey), + verifyEd25519: (publicKey, message, signature) => + bindings.ed25519_verify(publicKey, message, signature), + deriveEncryptionKey: (ikm, salt, context) => + bindings.wasm_derive_encryption_key(ikm, salt, context), + hkdfExpand: (ikm, salt, info, len) => + bindings.wasm_hkdf_expand(ikm, salt, info, len), sha256: (data) => bindings.wasm_sha256(data), sha256Double: (data) => bindings.wasm_sha256_double(data), keyringToKeys: (keyring) => keyringToKeys(keyring), + publicKeyBundleToKeys: (publicKeyBundle) => + publicKeyBundleToKeys(publicKeyBundle), encrypt: async (key, input) => { const cipher = new bindings.WasmChaCha20Poly1305(key); @@ -78,7 +129,10 @@ export const crypto: MTPCrypto = { encryptText: async (key, plaintext) => { const cipher = new bindings.WasmChaCha20Poly1305(key); try { - const ciphertext = cipher.encrypt(utf8Encode(plaintext), new Uint8Array(0)); + const ciphertext = cipher.encrypt( + utf8Encode(plaintext), + new Uint8Array(0), + ); return bytesToBase64(ciphertext); } finally { cipher.free(); @@ -96,7 +150,8 @@ export const crypto: MTPCrypto = { } }, - encapsulate: (otherPublicKey) => bindings.wasm_kem_encapsulate(otherPublicKey), + encapsulate: (otherPublicKey) => + bindings.wasm_kem_encapsulate(otherPublicKey), decapsulate: (ownPrivateKey, ciphertext) => bindings.wasm_kem_decapsulate(ownPrivateKey, ciphertext), @@ -136,12 +191,20 @@ export interface MTPCodecOptions { } export interface MTPCodec { - encode(type: MTPCommunicationType, data: Record, options?: MTPCodecOptions): Uint8Array; + encode( + type: MTPCommunicationType, + data: Record, + options?: MTPCodecOptions, + ): Uint8Array; decode(frame: MTPBytesInput): ParsedFrame; format(frame: MTPBytesInput): string; } -export function encode(type: MTPCommunicationType, data: Record, options?: MTPCodecOptions): Uint8Array { +export function encode( + type: MTPCommunicationType, + data: Record, + options?: MTPCodecOptions, +): Uint8Array { return bindings.build_frame(type, data, options ?? {}); } @@ -184,6 +247,12 @@ export interface MTPKeyringKeys { sigClSecretKey: Uint8Array; } +export interface MTPPublicKeyBundleKeys { + kemPublicKey: Uint8Array; + sigPqPublicKey: Uint8Array; + sigClPublicKey: Uint8Array; +} + export interface MTPClientOptions { url: string; descriptor?: string; @@ -195,8 +264,15 @@ export interface MTPClientOptions { maxMessageSize?: number; authTimeoutMs?: number; pings?: boolean | { intervalMs?: number }; - wasm?: RawBindings.InitInput | Promise | { module_or_path: RawBindings.InitInput | Promise }; + wasm?: + | RawBindings.InitInput + | Promise + | { + module_or_path: RawBindings.InitInput | Promise; + }; logger?: (event: MTPLogEvent) => void; + sessionStorage?: MTPSessionStorage; + encryptedDeviceSecretProvider?: MTPEncryptedDeviceSecretProvider; } export type Unsubscribe = () => void; @@ -211,7 +287,10 @@ export interface MTPRequestOptions extends MTPSendOptions { responseType?: MTPCommunicationType; } -type InternalCredentials = Omit & { +type InternalCredentials = Omit< + MTPCredentials, + "clientId" | "keyring" | "hostPublicKey" +> & { clientId: bigint | null; keyringBytes: Uint8Array; hostPublicKey?: Uint8Array; @@ -231,22 +310,31 @@ function emit(logger, event) { } function isErrorType(type) { - return type === "Error" || type.startsWith("Error") || [ - "BadRequest", - "Unauthorized", - "Forbidden", - "NotFound", - "TooManyRequests", - "InternalServerError", - "BadGateway", - "ServiceUnavailable", - "GatewayTimeout", - ].includes(type); + return ( + type === "Error" || + type.startsWith("Error") || + [ + "BadRequest", + "Unauthorized", + "Forbidden", + "NotFound", + "TooManyRequests", + "InternalServerError", + "BadGateway", + "ServiceUnavailable", + "GatewayTimeout", + ].includes(type) + ); } function errorMessage(frame) { const data = frame?.data ?? {}; - return String(data.ErrorMessage ?? data.Error ?? data.Description ?? `Received ${frame?.type ?? "error"} frame`); + return String( + data.ErrorMessage ?? + data.Error ?? + data.Description ?? + `Received ${frame?.type ?? "error"} frame`, + ); } async function storageGet(storage, key) { @@ -404,7 +492,9 @@ function utf8Decode(bytes) { i += 2; } else if (b < 0xf0) { out += String.fromCharCode( - ((b & 0x0f) << 12) | ((bytes[i + 1] & 0x3f) << 6) | (bytes[i + 2] & 0x3f), + ((b & 0x0f) << 12) | + ((bytes[i + 1] & 0x3f) << 6) | + (bytes[i + 2] & 0x3f), ); i += 3; } else { @@ -449,7 +539,11 @@ export function secretKeyFromString(secret) { } const ikm = utf8Encode(trimmed); - return bindings.wasm_derive_encryption_key(ikm, SYMMETRIC_KEY_SALT, SYMMETRIC_KEY_SALT); + return bindings.wasm_derive_encryption_key( + ikm, + SYMMETRIC_KEY_SALT, + SYMMETRIC_KEY_SALT, + ); } function normalizeBytes(value, name) { @@ -483,9 +577,10 @@ function generateKeyringBytes() { } export function keyringToKeys(keyring) { - const bytes = typeof keyring === "string" - ? bytesFromString(keyring, "keyring") - : bytesFrom(keyring, "keyring"); + const bytes = + typeof keyring === "string" + ? bytesFromString(keyring, "keyring") + : bytesFrom(keyring, "keyring"); if (bytes.length < 12) { throw new TypeError("keyring data is too short to contain 6 keys"); @@ -510,11 +605,51 @@ export function keyringToKeys(keyring) { }; } +export function publicKeyBundleToKeys(publicKeyBundle) { + const bytes = + typeof publicKeyBundle === "string" + ? bytesFromString(publicKeyBundle, "publicKeyBundle") + : bytesFrom(publicKeyBundle, "publicKeyBundle"); + + if (bytes.length < 6) { + throw new TypeError("public key bundle data is too short to contain 3 keys"); + } + + let offset = 0; + const readKey = () => { + if (offset + 2 > bytes.length) { + throw new TypeError("public key bundle is truncated"); + } + const len = (bytes[offset] << 8) | bytes[offset + 1]; + offset += 2; + if (offset + len > bytes.length) { + throw new TypeError("public key bundle is truncated"); + } + const key = bytes.slice(offset, offset + len); + offset += len; + return key; + }; + + const result = { + kemPublicKey: readKey(), + sigPqPublicKey: readKey(), + sigClPublicKey: readKey(), + }; + + if (offset !== bytes.length) { + throw new TypeError("public key bundle has trailing data"); + } + + return result; +} + function serializeCredentials(credentials) { return JSON.stringify({ clientId: credentials.clientId?.toString() ?? null, keyring: Array.from(credentials.keyringBytes ?? []), - hostPublicKey: credentials.hostPublicKey ? Array.from(credentials.hostPublicKey) : undefined, + hostPublicKey: credentials.hostPublicKey + ? Array.from(credentials.hostPublicKey) + : undefined, }); } @@ -532,9 +667,10 @@ function deserializeCredentials(credentials) { return { clientId: toBigInt(normalized.clientId), keyringBytes: bytesFrom(keyring, "credentials.keyring"), - hostPublicKey: normalized.hostPublicKey == null - ? undefined - : normalizeBytes(normalized.hostPublicKey, "credentials.hostPublicKey"), + hostPublicKey: + normalized.hostPublicKey == null + ? undefined + : normalizeBytes(normalized.hostPublicKey, "credentials.hostPublicKey"), }; } @@ -567,10 +703,17 @@ function validateOptions(options) { } } } - if (options.maxMessageSize != null && (!Number.isSafeInteger(options.maxMessageSize) || options.maxMessageSize <= 0)) { + if ( + options.maxMessageSize != null && + (!Number.isSafeInteger(options.maxMessageSize) || + options.maxMessageSize <= 0) + ) { throw new TypeError("maxMessageSize must be a positive safe integer"); } - if (options.authTimeoutMs != null && (!Number.isSafeInteger(options.authTimeoutMs) || options.authTimeoutMs <= 0)) { + if ( + options.authTimeoutMs != null && + (!Number.isSafeInteger(options.authTimeoutMs) || options.authTimeoutMs <= 0) + ) { throw new TypeError("authTimeoutMs must be a positive safe integer"); } } @@ -604,10 +747,22 @@ export class MTPClient { readonly crypto = MTPClient.crypto; readonly codec = MTPClient.codec; - private constructor(options: NormalizedMTPClientOptions, client: RawBindings.WasmClient) { + readonly sessionManager: MTPSessionManager; + readonly encryptedDeviceSecretProvider: MTPEncryptedDeviceSecretProvider; + + private constructor( + options: NormalizedMTPClientOptions, + client: RawBindings.WasmClient, + ) { this.#options = options; this.#credentials = deserializeCredentials(options.credentials); this.raw = { client, bindings }; + this.encryptedDeviceSecretProvider = + options.encryptedDeviceSecretProvider ?? + new InMemoryEncryptedDeviceSecretProvider(); + this.sessionManager = new MTPSessionManager( + options.sessionStorage ?? new InMemorySessionStorage(), + ); } static async create(options: MTPClientOptions): Promise { @@ -616,28 +771,31 @@ export class MTPClient { const normalizedOptions = { ...options, - hostPublicKey: options.hostPublicKey == null - ? undefined - : normalizeBytes(options.hostPublicKey, "hostPublicKey"), + hostPublicKey: + options.hostPublicKey == null + ? undefined + : normalizeBytes(options.hostPublicKey, "hostPublicKey"), }; let sdk: MTPClient | undefined; const client = new WasmClient( - (state) => emit(normalizedOptions.logger, { - hint: "info", - type: "state", - data: ConnectionState[state] ?? state, - }), + (state) => + emit(normalizedOptions.logger, { + hint: "info", + type: "state", + data: ConnectionState[state] ?? state, + }), (frame) => { if (sdk) { sdk.#handleFrame(frame); } }, - (error) => emit(normalizedOptions.logger, { - hint: "error", - type: "Error", - error: String(error), - }), + (error) => + emit(normalizedOptions.logger, { + hint: "error", + type: "Error", + error: String(error), + }), ); sdk = new MTPClient(normalizedOptions, client); @@ -648,10 +806,22 @@ export class MTPClient { keyringBytes: generateKeyringBytes(), hostPublicKey: normalizedOptions.hostPublicKey, }; - } else if (!sdk.#credentials.hostPublicKey && normalizedOptions.hostPublicKey) { - sdk.#credentials = { ...sdk.#credentials, hostPublicKey: normalizedOptions.hostPublicKey }; - } else if (!normalizedOptions.hostPublicKey && sdk.#credentials.hostPublicKey) { - sdk.#options = { ...sdk.#options, hostPublicKey: sdk.#credentials.hostPublicKey }; + } else if ( + !sdk.#credentials.hostPublicKey && + normalizedOptions.hostPublicKey + ) { + sdk.#credentials = { + ...sdk.#credentials, + hostPublicKey: normalizedOptions.hostPublicKey, + }; + } else if ( + !normalizedOptions.hostPublicKey && + sdk.#credentials.hostPublicKey + ) { + sdk.#options = { + ...sdk.#options, + hostPublicKey: sdk.#credentials.hostPublicKey, + }; } return sdk; } @@ -660,7 +830,9 @@ export class MTPClient { return WasmClient.is_supported(); } - static async init(wasm?: MTPClientOptions["wasm"]): Promise>> { + static async init( + wasm?: MTPClientOptions["wasm"], + ): Promise>> { wasmInitPromise ??= initWasm(wasm); return await wasmInitPromise; } @@ -729,10 +901,17 @@ export class MTPClient { async #connectAuthenticated() { if (!this.#options.hostPublicKey) { - throw new Error("MTPClient.connect requires hostPublicKey for authenticated connections"); + throw new Error( + "MTPClient.connect requires hostPublicKey for authenticated connections", + ); } - if (!this.#credentials?.keyringBytes?.length || this.#credentials.clientId == null) { - throw new Error("MTPClient.connect requires credentials with clientId and keyring"); + if ( + !this.#credentials?.keyringBytes?.length || + this.#credentials.clientId == null + ) { + throw new Error( + "MTPClient.connect requires credentials with clientId and keyring", + ); } const config = this.#connectionConfig(); @@ -810,7 +989,8 @@ export class MTPClient { this.raw.client.stop_protocol_pings(); return; } - const intervalMs = typeof pings === "object" ? pings.intervalMs ?? 30_000 : 30_000; + const intervalMs = + typeof pings === "object" ? (pings.intervalMs ?? 30_000) : 30_000; this.raw.client.start_protocol_pings(intervalMs, clientId); } @@ -819,7 +999,9 @@ export class MTPClient { return typeOrFrame; } if (typeof typeOrFrame !== "string" || !typeOrFrame) { - throw new TypeError("message type must be a non-empty string or Uint8Array frame"); + throw new TypeError( + "message type must be a non-empty string or Uint8Array frame", + ); } if (data == null || typeof data !== "object" || Array.isArray(data)) { throw new TypeError("message data must be an object"); @@ -828,15 +1010,37 @@ export class MTPClient { } async send(message: Uint8Array): Promise; - async send(type: MTPCommunicationType, data: Record, options?: MTPSendOptions): Promise; - async send(typeOrFrame: Uint8Array | MTPCommunicationType, data?: Record, options?: MTPSendOptions): Promise { + async send( + type: MTPCommunicationType, + data: Record, + options?: MTPSendOptions, + ): Promise; + async send( + typeOrFrame: Uint8Array | MTPCommunicationType, + data?: Record, + options?: MTPSendOptions, + ): Promise { const message = this.#buildFrame(typeOrFrame, data, options); try { const frame = this.raw.bindings.parse_frame(message); - emit(this.#options.logger, isErrorType(frame.type) - ? { hint: "error", type: frame.type, error: errorMessage(frame), data: frame.data, direction: "send" } - : { hint: "info", type: frame.type, data: frame.data, direction: "send" }); + emit( + this.#options.logger, + isErrorType(frame.type) + ? { + hint: "error", + type: frame.type, + error: errorMessage(frame), + data: frame.data, + direction: "send", + } + : { + hint: "info", + type: frame.type, + data: frame.data, + direction: "send", + }, + ); } catch (error) { emit(this.#options.logger, { hint: "error", @@ -849,15 +1053,41 @@ export class MTPClient { await this.raw.client.send(message); } - async request(message: Uint8Array, data?: never, options?: MTPRequestOptions): Promise; - async request(type: MTPCommunicationType, data: Record, options?: MTPRequestOptions): Promise; - async request(typeOrFrame: Uint8Array | MTPCommunicationType, data?: Record, options: MTPRequestOptions = {}): Promise { + async request( + message: Uint8Array, + data?: never, + options?: MTPRequestOptions, + ): Promise; + async request( + type: MTPCommunicationType, + data: Record, + options?: MTPRequestOptions, + ): Promise; + async request( + typeOrFrame: Uint8Array | MTPCommunicationType, + data?: Record, + options: MTPRequestOptions = {}, + ): Promise { const frame = this.#buildFrame(typeOrFrame, data, options); try { const parsed = this.raw.bindings.parse_frame(frame); - emit(this.#options.logger, isErrorType(parsed.type) - ? { hint: "error", type: parsed.type, error: errorMessage(parsed), data: parsed.data, direction: "send" } - : { hint: "info", type: parsed.type, data: parsed.data, direction: "send" }); + emit( + this.#options.logger, + isErrorType(parsed.type) + ? { + hint: "error", + type: parsed.type, + error: errorMessage(parsed), + data: parsed.data, + direction: "send", + } + : { + hint: "info", + type: parsed.type, + data: parsed.data, + direction: "send", + }, + ); } catch (error) { emit(this.#options.logger, { hint: "error", @@ -869,7 +1099,10 @@ export class MTPClient { return await this.raw.client.request(frame, options.responseType ?? null); } - subscribe(type: MTPCommunicationType, handler: (message: ParsedFrame) => void): Unsubscribe { + subscribe( + type: MTPCommunicationType, + handler: (message: ParsedFrame) => void, + ): Unsubscribe { if (typeof type !== "string" || !type) { throw new TypeError("subscription type must be a non-empty string"); } @@ -899,6 +1132,420 @@ export class MTPClient { } } + #getKemPublicKey(): Uint8Array { + if (!this.#credentials?.keyringBytes?.length) { + throw new Error("No keyring available"); + } + const keys = keyringToKeys(this.#credentials.keyringBytes); + return keys.kemPublicKey; + } + + #getKemSecretKey(): Uint8Array { + if (!this.#credentials?.keyringBytes?.length) { + throw new Error("No keyring available"); + } + const keys = keyringToKeys(this.#credentials.keyringBytes); + return keys.kemSecretKey; + } + + async sendEncrypted( + type: number | string, + data: Record, + options: MTPSendOptions & { + recipientClientId: bigint | number | string; + recipientPublicKey: string | MTPBytesInput; + senderUserId?: string; + recipientUserId?: string; + recipientDeviceId?: string; + }, + ): Promise { + const ownId = this.#credentials?.clientId; + if (ownId == null) { + throw new Error("Client not registered"); + } + if (options.recipientClientId == null) { + throw new Error("recipientClientId is required"); + } + const recipientClientId = BigInt(options.recipientClientId); + + const plaintext = this.raw.bindings.build_frame(type as string, data, { + sender: ownId, + receiver: recipientClientId, + ...options, + }); + + let session = await this.sessionManager.getSession( + ownId, + recipientClientId, + ); + let kemCiphertext: Uint8Array | undefined; + + if (!session) { + if (options.recipientPublicKey == null) { + throw new Error("recipientPublicKey is required for new encrypted sessions"); + } + const recipientPublicKey = publicKeyBundleToKeys( + options.recipientPublicKey, + ); + const enc = bindings.wasm_kem_encapsulate( + recipientPublicKey.kemPublicKey, + ); + kemCiphertext = enc.ciphertext; + const conversationId = getConversationId(ownId, recipientClientId); + + session = await this.sessionManager.createSession({ + ownClientId: ownId, + peerClientId: recipientClientId, + peerPublicKey: recipientPublicKey.kemPublicKey, + sharedSecret: enc.shared_secret, + role: "initiator", + transcriptContext: { + senderUserId: (options as { senderUserId?: string }).senderUserId, + senderClientId: ownId, + recipientUserId: options.recipientUserId, + recipientClientId, + recipientPublicKey: recipientPublicKey.kemPublicKey, + kemCiphertext, + conversationId, + }, + }); + enc.shared_secret.fill(0); + } + + const { payload, session: newSession } = await encryptPayload({ + plaintext, + session, + kemCiphertext, + }); + + await this.sessionManager.saveSession(newSession); + + const messageId = String(Date.now()); + const createdAt = Date.now(); + const senderUserId = (options as { senderUserId?: string }).senderUserId; + const frame = this.raw.bindings.build_frame( + "EncryptedMessage", + { + MessageId: messageId, + ConversationId: session.conversationId, + SenderClientId: ownId, + RecipientClientId: recipientClientId, + SenderUserId: senderUserId, + RecipientUserId: options.recipientUserId, + CreatedAt: createdAt, + EncryptionVersion: 1, + EncryptedPayload: payload, + }, + { + sender: ownId, + receiver: recipientClientId, + }, + ); + await this.raw.client.send(frame); + + if (senderUserId && options.recipientUserId) { + const ownKemPublicKey = this.#getKemPublicKey(); + const archiveEnc = bindings.wasm_kem_encapsulate(ownKemPublicKey); + const archiveSession = await this.sessionManager.createSession({ + ownClientId: ownId, + peerClientId: ownId, + peerPublicKey: ownKemPublicKey, + sharedSecret: archiveEnc.shared_secret, + role: "initiator", + transcriptContext: { + senderUserId, + senderClientId: ownId, + recipientUserId: options.recipientUserId, + recipientClientId: ownId, + recipientPublicKey: ownKemPublicKey, + kemCiphertext: archiveEnc.ciphertext, + conversationId: `archive:${session.conversationId}:${messageId}`, + }, + }); + archiveEnc.shared_secret.fill(0); + const { payload: archivePayload } = await encryptPayload({ + plaintext, + session: archiveSession, + kemCiphertext: archiveEnc.ciphertext, + }); + const archiveFrame = this.raw.bindings.build_frame( + "EncryptedMessage", + { + MessageId: `${messageId}:sender`, + ConversationId: session.conversationId, + SenderClientId: ownId, + RecipientClientId: ownId, + SenderUserId: senderUserId, + RecipientUserId: options.recipientUserId, + CreatedAt: createdAt, + EncryptionVersion: 1, + EncryptedPayload: archivePayload, + }, + { + sender: ownId, + receiver: ownId, + }, + ); + await this.raw.client.send(archiveFrame); + } + } + + subscribeEncrypted( + type: number | string, + handler: (data: unknown, meta: ParsedFrame) => void | Promise, + ): Unsubscribe; + subscribeEncrypted( + handler: (data: { + type: string; + data: Record; + sender?: bigint; + receiver?: bigint; + }) => void | Promise, + ): Unsubscribe; + subscribeEncrypted( + typeOrHandler: + | number + | string + | ((data: { + type: string; + data: Record; + sender?: bigint; + receiver?: bigint; + }) => void | Promise), + maybeHandler?: (data: unknown, meta: ParsedFrame) => void | Promise, + ): Unsubscribe { + const expectedInnerType = + typeof typeOrHandler === "function" ? null : String(typeOrHandler); + const legacyHandler = + typeof typeOrHandler === "function" ? typeOrHandler : null; + const sub = this.raw.client.subscribe( + "EncryptedMessage", + async (frame: ParsedFrame) => { + const raw = + frame.data?.["encryptedPayload"] ?? + frame.data?.["EncryptedPayload"] ?? + frame.data?.["encrypted_payload"]; + if (!raw) return; + + let payloadBytes: Uint8Array; + if (raw instanceof Uint8Array) { + payloadBytes = raw; + } else if (Array.isArray(raw)) { + payloadBytes = new Uint8Array(raw); + } else { + return; + } + + try { + const parsed = parseEncryptedMessage(payloadBytes); + + const ownId = this.#credentials?.clientId; + if (ownId == null) return; + if (parsed.recipientClientId !== ownId) return; + + const peerClientId = parsed.senderClientId; + let session = await this.sessionManager.getSession( + ownId, + peerClientId, + ); + + if (!session) { + if (!(parsed.flags & FLAG_INIT) || !parsed.kemCiphertext) { + return; + } + + const ownKemSecret = this.#getKemSecretKey(); + const sharedSecret = bindings.wasm_kem_decapsulate( + ownKemSecret, + parsed.kemCiphertext, + ); + + session = await this.sessionManager.createSession({ + ownClientId: ownId, + peerClientId, + peerPublicKey: new Uint8Array(0), + sharedSecret, + role: "receiver", + transcriptContext: { + senderUserId: String( + frame.data?.["SenderUserId"] ?? + frame.data?.["senderUserId"] ?? + "", + ), + senderClientId: peerClientId, + recipientUserId: String( + frame.data?.["RecipientUserId"] ?? + frame.data?.["recipientUserId"] ?? + "", + ), + recipientClientId: ownId, + recipientPublicKey: this.#getKemPublicKey(), + kemCiphertext: parsed.kemCiphertext, + conversationId: getConversationId(peerClientId, ownId), + }, + }); + sharedSecret.fill(0); + } + + const { plaintext, session: newSession } = await decryptPayload({ + payload: payloadBytes, + session, + expectedRecipientClientId: ownId, + }); + + await this.sessionManager.saveSession(newSession); + + let parsedFrame: ParsedFrame; + try { + parsedFrame = this.raw.bindings.parse_frame(plaintext); + } catch { + return; + } + + if (expectedInnerType && parsedFrame.type !== expectedInnerType) { + return; + } + if (legacyHandler) { + await legacyHandler({ + type: parsedFrame.type, + data: parsedFrame.data, + sender: parsedFrame.sender, + receiver: parsedFrame.receiver, + }); + } else if (maybeHandler) { + await maybeHandler(parsedFrame.data, parsedFrame); + } + } catch (e) { + emit(this.#options.logger, { + hint: "error", + type: "E2EE", + error: String(e), + direction: "recv", + }); + } + }, + ); + + return () => this.raw.client.unsubscribe(sub); + } + + async decryptEncryptedRecord( + frameData: Record, + ): Promise { + const raw = + frameData["encryptedPayload"] ?? + frameData["EncryptedPayload"] ?? + frameData["encrypted_payload"]; + if (!raw) throw new Error("EncryptedPayload is required"); + + const payloadBytes = + raw instanceof Uint8Array + ? raw + : Array.isArray(raw) + ? new Uint8Array(raw) + : bytesFrom(raw as MTPBytesInput, "EncryptedPayload"); + + const parsed = parseEncryptedMessage(payloadBytes); + const ownId = this.#credentials?.clientId; + if (ownId == null) throw new Error("Client not registered"); + if (parsed.recipientClientId !== ownId) { + throw new Error("Encrypted message recipient mismatch"); + } + + const peerClientId = parsed.senderClientId; + let session = await this.sessionManager.getSession(ownId, peerClientId); + + const isSenderArchive = + parsed.senderClientId === ownId && parsed.recipientClientId === ownId; + + if (isSenderArchive && (parsed.flags & FLAG_INIT) && parsed.kemCiphertext) { + const ownKemSecret = this.#getKemSecretKey(); + const sharedSecret = bindings.wasm_kem_decapsulate( + ownKemSecret, + parsed.kemCiphertext, + ); + const archiveMessageId = String( + frameData["MessageId"] ?? frameData["messageId"] ?? "", + ).replace(/:sender$/, ""); + session = await this.sessionManager.createSession({ + ownClientId: ownId, + peerClientId: ownId, + peerPublicKey: this.#getKemPublicKey(), + sharedSecret, + role: "receiver", + transcriptContext: { + senderUserId: String( + frameData["SenderUserId"] ?? frameData["senderUserId"] ?? "", + ), + senderClientId: ownId, + recipientUserId: String( + frameData["RecipientUserId"] ?? frameData["recipientUserId"] ?? "", + ), + recipientClientId: ownId, + recipientPublicKey: this.#getKemPublicKey(), + kemCiphertext: parsed.kemCiphertext, + conversationId: `archive:${String(frameData["ConversationId"] ?? frameData["conversationId"] ?? "")}:${archiveMessageId}`, + }, + }); + sharedSecret.fill(0); + } else if (!session) { + if (!(parsed.flags & FLAG_INIT) || !parsed.kemCiphertext) { + throw new Error("No session for non-init encrypted message"); + } + const ownKemSecret = this.#getKemSecretKey(); + const sharedSecret = bindings.wasm_kem_decapsulate( + ownKemSecret, + parsed.kemCiphertext, + ); + + session = await this.sessionManager.createSession({ + ownClientId: ownId, + peerClientId, + peerPublicKey: new Uint8Array(0), + sharedSecret, + role: "receiver", + transcriptContext: { + senderUserId: String( + frameData["SenderUserId"] ?? frameData["senderUserId"] ?? "", + ), + senderClientId: peerClientId, + recipientUserId: String( + frameData["RecipientUserId"] ?? frameData["recipientUserId"] ?? "", + ), + recipientClientId: ownId, + recipientPublicKey: this.#getKemPublicKey(), + kemCiphertext: parsed.kemCiphertext, + conversationId: getConversationId(peerClientId, ownId), + }, + }); + sharedSecret.fill(0); + } + + const { plaintext, session: newSession } = await decryptPayload({ + payload: payloadBytes, + session, + expectedRecipientClientId: ownId, + }); + if (!isSenderArchive) { + await this.sessionManager.saveSession(newSession); + } + return this.raw.bindings.parse_frame(plaintext); + } + + async setEncryptedDeviceSecret( + record: EncryptedDeviceSecretRecord, + ): Promise { + await this.encryptedDeviceSecretProvider.setEncryptedDeviceSecret(record); + } + + async getEncryptedDeviceSecret(query: { + userId: string; + deviceId?: string; + secretId?: string; + }): Promise { + return this.encryptedDeviceSecretProvider.getEncryptedDeviceSecret(query); + } + disconnect(): void { this.raw.client.stop_protocol_pings(); this.raw.client.disconnect(); @@ -906,3 +1553,37 @@ export class MTPClient { } export { ConnectionState, bindings as raw }; + +// E2EE exports +export type { + MTPSessionState, + MTPSessionStorage, + MTPSessionTranscriptContext, +} from "./session"; +export { + MTPSessionManager, + InMemorySessionStorage, + getConversationId, + deriveSessionKeys, + buildSessionTranscript, +} from "./session.js"; +export { MTPRatchet } from "./ratchet.js"; +export type { RatchetStep } from "./ratchet.js"; +export { + serializeEncryptedMessage, + parseEncryptedMessage, + encryptPayload, + decryptPayload, + MTP_E2EE_VERSION, + FLAG_INIT, + MAX_RATCHET_SKIP, +} from "./encrypted-message.js"; +export type { + EncryptedMessageHeader, + SerializedEncryptedMessage, +} from "./encrypted-message"; +export type { + EncryptedDeviceSecretRecord, + MTPEncryptedDeviceSecretProvider, +} from "./encrypted-device-secret"; +export { InMemoryEncryptedDeviceSecretProvider } from "./encrypted-device-secret.js"; diff --git a/src/sdk/ratchet.ts b/src/sdk/ratchet.ts new file mode 100644 index 0000000..9ae8242 --- /dev/null +++ b/src/sdk/ratchet.ts @@ -0,0 +1,70 @@ +import * as bindings from "mtp/raw"; + +function utf8Encode(text: string): Uint8Array { + if (typeof TextEncoder !== "undefined") { + return new TextEncoder().encode(text); + } + if (typeof Buffer !== "undefined") { + return new Uint8Array(Buffer.from(text, "utf-8")); + } + const bytes = new Uint8Array(text.length * 4); + let len = 0; + for (let i = 0; i < text.length; i += 1) { + const code = text.codePointAt(i) as number; + if (code < 0x80) { + bytes[len++] = code; + } else if (code < 0x800) { + bytes[len++] = 0xc0 | (code >> 6); + bytes[len++] = 0x80 | (code & 0x3f); + } else if (code < 0x10000) { + bytes[len++] = 0xe0 | (code >> 12); + bytes[len++] = 0x80 | ((code >> 6) & 0x3f); + bytes[len++] = 0x80 | (code & 0x3f); + } else { + bytes[len++] = 0xf0 | (code >> 18); + bytes[len++] = 0x80 | ((code >> 12) & 0x3f); + bytes[len++] = 0x80 | ((code >> 6) & 0x3f); + bytes[len++] = 0x80 | (code & 0x3f); + i += 1; + } + } + return bytes.subarray(0, len); +} + +const HKDF_MESSAGE_KEY = "mtp-e2ee-v1-message-key"; +const HKDF_NEXT_CHAIN = "mtp-e2ee-v1-next-chain"; + +export interface RatchetStep { + key: Uint8Array; + chainKey: Uint8Array; +} + +export class MTPRatchet { + static async stepSend(chainKey: Uint8Array): Promise { + return this.step(chainKey); + } + + static async stepRecv(chainKey: Uint8Array): Promise { + return this.step(chainKey); + } + + static async step(chainKey: Uint8Array): Promise { + const messageKey = bindings.wasm_hkdf_expand( + chainKey, + new Uint8Array(0), + utf8Encode(HKDF_MESSAGE_KEY), + 32, + ); + const nextChainKey = bindings.wasm_hkdf_expand( + chainKey, + new Uint8Array(0), + utf8Encode(HKDF_NEXT_CHAIN), + 32, + ); + + return { + key: messageKey, + chainKey: nextChainKey, + }; + } +} diff --git a/src/sdk/session.ts b/src/sdk/session.ts new file mode 100644 index 0000000..c1ad0a4 --- /dev/null +++ b/src/sdk/session.ts @@ -0,0 +1,249 @@ +import * as bindings from "mtp/raw"; + +function utf8Encode(text: string): Uint8Array { + if (typeof TextEncoder !== "undefined") { + return new TextEncoder().encode(text); + } + if (typeof Buffer !== "undefined") { + return new Uint8Array(Buffer.from(text, "utf-8")); + } + const bytes = new Uint8Array(text.length * 4); + let len = 0; + for (let i = 0; i < text.length; i += 1) { + const code = text.codePointAt(i) as number; + if (code < 0x80) { + bytes[len++] = code; + } else if (code < 0x800) { + bytes[len++] = 0xc0 | (code >> 6); + bytes[len++] = 0x80 | (code & 0x3f); + } else if (code < 0x10000) { + bytes[len++] = 0xe0 | (code >> 12); + bytes[len++] = 0x80 | ((code >> 6) & 0x3f); + bytes[len++] = 0x80 | (code & 0x3f); + } else { + bytes[len++] = 0xf0 | (code >> 18); + bytes[len++] = 0x80 | ((code >> 12) & 0x3f); + bytes[len++] = 0x80 | ((code >> 6) & 0x3f); + bytes[len++] = 0x80 | (code & 0x3f); + i += 1; + } + } + return bytes.subarray(0, len); +} + +export const HKDF_SALT_ROOT = "mtp-e2ee-v1-root"; +const HKDF_INITIATOR_SEND = "mtp-e2ee-v1-initiator-send"; +const HKDF_INITIATOR_RECV = "mtp-e2ee-v1-initiator-recv"; + +export interface MTPSessionTranscriptContext { + senderUserId?: string; + senderClientId: bigint; + recipientUserId?: string; + recipientClientId: bigint; + recipientPublicKey: Uint8Array; + kemCiphertext: Uint8Array; + conversationId: string; +} + +export interface MTPSessionState { + version: 1; + conversationId: string; + ownClientId: bigint; + peerClientId: bigint; + peerPublicKey: Uint8Array; + sendChainKey: Uint8Array; + recvChainKey: Uint8Array; + sendCount: number; + recvCount: number; + createdAt: number; + updatedAt: number; +} + +export interface MTPSessionStorage { + getSession(conversationId: string): Promise; + setSession(state: MTPSessionState): Promise; + deleteSession(conversationId: string): Promise; +} + +export class InMemorySessionStorage implements MTPSessionStorage { + private store = new Map(); + + async getSession(conversationId: string): Promise { + return this.store.get(conversationId) ?? null; + } + + async setSession(state: MTPSessionState): Promise { + this.store.set(state.conversationId, { ...state }); + } + + async deleteSession(conversationId: string): Promise { + this.store.delete(conversationId); + } +} + +function writeU32BE(value: number): Uint8Array { + return new Uint8Array([ + (value >>> 24) & 0xff, + (value >>> 16) & 0xff, + (value >>> 8) & 0xff, + value & 0xff, + ]); +} + +function writeU64BE(value: bigint): Uint8Array { + if (value < 0n || value > 0xffff_ffff_ffff_ffffn) + throw new Error("u64 out of range"); + const buf = new Uint8Array(8); + for (let i = 7; i >= 0; i -= 1) { + buf[i] = Number(value & 0xffn); + value >>= 8n; + } + return buf; +} + +function concatBytes(parts: Uint8Array[]): Uint8Array { + const out = new Uint8Array(parts.reduce((sum, part) => sum + part.length, 0)); + let offset = 0; + for (const part of parts) { + out.set(part, offset); + offset += part.length; + } + return out; +} + +function transcriptField(label: string, value: Uint8Array): Uint8Array { + const labelBytes = utf8Encode(label); + return concatBytes([ + writeU32BE(labelBytes.length), + labelBytes, + writeU32BE(value.length), + value, + ]); +} + +export function buildSessionTranscript( + args: MTPSessionTranscriptContext, +): Uint8Array { + const recipientPublicKeyHash = bindings.wasm_sha256(args.recipientPublicKey); + const kemHash = bindings.wasm_sha256(args.kemCiphertext); + return concatBytes([ + transcriptField("domain", utf8Encode("mtp-e2ee-session-transcript-v1")), + transcriptField("version", utf8Encode("1")), + transcriptField("senderUserId", utf8Encode(args.senderUserId ?? "")), + transcriptField("senderClientId", writeU64BE(args.senderClientId)), + transcriptField("recipientUserId", utf8Encode(args.recipientUserId ?? "")), + transcriptField("recipientClientId", writeU64BE(args.recipientClientId)), + transcriptField("recipientPublicKeyHash", recipientPublicKeyHash), + transcriptField("kemCiphertextHash", kemHash), + transcriptField("conversationId", utf8Encode(args.conversationId)), + ]); +} + +export async function deriveSessionKeys( + sharedSecret: Uint8Array, + transcript: Uint8Array = new Uint8Array(0), +): Promise<{ + root: Uint8Array; + initiatorSend: Uint8Array; + initiatorRecv: Uint8Array; +}> { + const rootInfo = concatBytes([utf8Encode(HKDF_SALT_ROOT), transcript]); + const root = bindings.wasm_hkdf_expand( + sharedSecret, + new Uint8Array(0), + rootInfo, + 32, + ); + const initiatorSend = bindings.wasm_hkdf_expand( + root, + new Uint8Array(0), + utf8Encode(HKDF_INITIATOR_SEND), + 32, + ); + const initiatorRecv = bindings.wasm_hkdf_expand( + root, + new Uint8Array(0), + utf8Encode(HKDF_INITIATOR_RECV), + 32, + ); + return { root, initiatorSend, initiatorRecv }; +} + +export function getConversationId( + ownClientId: bigint, + peerClientId: bigint, +): string { + const ids = [ownClientId, peerClientId].sort((a, b) => + a < b ? -1 : a > b ? 1 : 0, + ); + return `${ids[0].toString(16)}:${ids[1].toString(16)}`; +} + +export class MTPSessionManager { + constructor(private storage: MTPSessionStorage) {} + + getConversationId( + ownClientId: bigint, + peerClientId: bigint, + ): Promise { + return Promise.resolve(getConversationId(ownClientId, peerClientId)); + } + + async getSession( + ownClientId: bigint, + peerClientId: bigint, + ): Promise { + return this.storage.getSession( + getConversationId(ownClientId, peerClientId), + ); + } + + async saveSession(state: MTPSessionState): Promise { + await this.storage.setSession({ ...state, updatedAt: Date.now() }); + } + + async deleteSession( + ownClientId: bigint, + peerClientId: bigint, + ): Promise { + await this.storage.deleteSession( + getConversationId(ownClientId, peerClientId), + ); + } + + async createSession(args: { + ownClientId: bigint; + peerClientId: bigint; + peerPublicKey: Uint8Array; + sharedSecret: Uint8Array; + role: "initiator" | "receiver"; + transcript?: Uint8Array; + transcriptContext?: MTPSessionTranscriptContext; + }): Promise { + const transcript = + args.transcript ?? + (args.transcriptContext + ? buildSessionTranscript(args.transcriptContext) + : undefined); + const { root, initiatorSend, initiatorRecv } = await deriveSessionKeys( + args.sharedSecret, + transcript, + ); + const now = Date.now(); + const state: MTPSessionState = { + version: 1, + conversationId: getConversationId(args.ownClientId, args.peerClientId), + ownClientId: args.ownClientId, + peerClientId: args.peerClientId, + peerPublicKey: args.peerPublicKey, + sendChainKey: args.role === "initiator" ? initiatorSend : initiatorRecv, + recvChainKey: args.role === "initiator" ? initiatorRecv : initiatorSend, + sendCount: 0, + recvCount: 0, + createdAt: now, + updatedAt: now, + }; + root.fill(0); + return state; + } +} diff --git a/test/e2ee.mjs b/test/e2ee.mjs new file mode 100644 index 0000000..fe05930 --- /dev/null +++ b/test/e2ee.mjs @@ -0,0 +1,503 @@ +import { initSync } from "../dist/raw/index.js"; +import fs from "fs"; +import path from "path"; +import { fileURLToPath } from "url"; +import assert from "node:assert/strict"; +import { describe, it } from "node:test"; + +const __dirname = path.dirname(fileURLToPath(import.meta.url)); +const wasmPath = path.resolve(__dirname, "../wasm/pkg/mtp_wasm_bg.wasm"); +const wasmBytes = fs.readFileSync(wasmPath); +const wasmModule = new WebAssembly.Module(wasmBytes); +initSync(wasmModule); + +const sdk = await import("../dist/sdk/index.js"); +const { MTPRatchet } = await import("../dist/sdk/ratchet.js"); +const { + serializeEncryptedMessage, + parseEncryptedMessage, + encryptPayload, + decryptPayload, + FLAG_INIT, + MTP_E2EE_VERSION, +} = await import("../dist/sdk/encrypted-message.js"); +const { + MTPSessionManager, + InMemorySessionStorage, + deriveSessionKeys, + getConversationId, +} = await import("../dist/sdk/session.js"); + +const bindings = sdk.raw; + +function concat(...arrays) { + const totalLen = arrays.reduce((sum, a) => sum + a.length, 0); + const result = new Uint8Array(totalLen); + let offset = 0; + for (const a of arrays) { + result.set(a, offset); + offset += a.length; + } + return result; +} + +function setupSessions(sharedSecret, aliceId = 1n, bobId = 2n) { + const aliceStorage = new InMemorySessionStorage(); + const aliceManager = new MTPSessionManager(aliceStorage); + const bobStorage = new InMemorySessionStorage(); + const bobManager = new MTPSessionManager(bobStorage); + + return { + aliceManager, + aliceStorage, + bobManager, + bobStorage, + async initSessions() { + const { initiatorSend, initiatorRecv } = await deriveSessionKeys( + sharedSecret, + new Uint8Array(0), + ); + const aliceSession = await aliceManager.createSession({ + ownClientId: aliceId, + peerClientId: bobId, + peerPublicKey: new Uint8Array(32), + sharedSecret, + role: "initiator", + }); + const bobSession = await bobManager.createSession({ + ownClientId: bobId, + peerClientId: aliceId, + peerPublicKey: new Uint8Array(32), + sharedSecret, + role: "receiver", + }); + return { aliceSession, bobSession, initiatorSend, initiatorRecv }; + }, + }; +} + +await describe("E2EE Session Derivation", async () => { + await it("Both sides derive same shared secret", async () => { + const sharedSecret = sdk.crypto.sha256(new Uint8Array([1, 2, 3, 4, 5])); + const transcript = new Uint8Array(0); + + const aliceKeys = await deriveSessionKeys(sharedSecret, transcript); + const bobKeys = await deriveSessionKeys(sharedSecret, transcript); + + // Deterministic: same inputs → same outputs + assert.deepEqual(aliceKeys.initiatorSend, bobKeys.initiatorSend); + assert.deepEqual(aliceKeys.initiatorRecv, bobKeys.initiatorRecv); + + // Init and recv keys are different + assert.notDeepEqual(aliceKeys.initiatorSend, aliceKeys.initiatorRecv); + }); + + await it("Session manager assigns correct chain keys per role", async () => { + const ss = sdk.crypto.sha256(new Uint8Array([1])); + const { initSessions } = setupSessions(ss); + const { aliceSession, bobSession, initiatorSend, initiatorRecv } = + await initSessions(); + + // Alice (initiator): send = initiatorSend, recv = initiatorRecv + assert.deepEqual(aliceSession.sendChainKey, initiatorSend); + assert.deepEqual(aliceSession.recvChainKey, initiatorRecv); + + // Bob (receiver): send = initiatorRecv, recv = initiatorSend + assert.deepEqual(bobSession.sendChainKey, initiatorRecv); + assert.deepEqual(bobSession.recvChainKey, initiatorSend); + + // Alice's send chain = Bob's recv chain + assert.deepEqual(aliceSession.sendChainKey, bobSession.recvChainKey); + // Alice's recv chain = Bob's send chain + assert.deepEqual(aliceSession.recvChainKey, bobSession.sendChainKey); + }); + + await it("Different transcripts produce different keys", async () => { + const ss = sdk.crypto.sha256(new Uint8Array([99])); + const aliceKeys1 = await deriveSessionKeys(ss, new Uint8Array(0)); + const aliceKeys2 = await deriveSessionKeys( + ss, + sdk.crypto.sha256(new Uint8Array([42])), + ); + assert.notDeepEqual(aliceKeys1.initiatorSend, aliceKeys2.initiatorSend); + }); +}); + +await describe("E2EE Ratchet", async () => { + await it("Repeated sends produce different message keys", async () => { + const chainKey = sdk.crypto.sha256(new Uint8Array([42])); + const step1 = await MTPRatchet.step(chainKey); + const step2 = await MTPRatchet.step(step1.chainKey); + const step3 = await MTPRatchet.step(step2.chainKey); + + assert.notDeepEqual(step1.key, step2.key); + assert.notDeepEqual(step2.key, step3.key); + assert.notDeepEqual(step1.key, step3.key); + assert.notDeepEqual(chainKey, step1.chainKey); + }); + + await it("Receiver can decrypt messages sent by sender in order", async () => { + const chainKey = sdk.crypto.sha256(new Uint8Array([7])); + + const send1 = await MTPRatchet.step(chainKey); + const send2 = await MTPRatchet.step(send1.chainKey); + const send3 = await MTPRatchet.step(send2.chainKey); + + const recv1 = await MTPRatchet.step(chainKey); + const recv2 = await MTPRatchet.step(recv1.chainKey); + const recv3 = await MTPRatchet.step(recv2.chainKey); + + assert.deepEqual(send1.key, recv1.key); + assert.deepEqual(send2.key, recv2.key); + assert.deepEqual(send3.key, recv3.key); + }); +}); + +await describe("E2EE Serialization", async () => { + await it("Roundtrips a basic message", () => { + const msg = { + header: { + version: 1, + flags: 0, + senderClientId: 0x1234567890abcdefn, + recipientClientId: 0xfedcba0987654321n, + messageNumber: 42, + }, + aeadPayload: new Uint8Array([1, 2, 3, 4, 5]), + }; + const bytes = serializeEncryptedMessage(msg); + const parsed = parseEncryptedMessage(bytes); + assert.equal(parsed.header.version, 1); + assert.equal(parsed.header.flags, 0); + assert.equal(parsed.header.senderClientId, msg.header.senderClientId); + assert.equal(parsed.header.recipientClientId, msg.header.recipientClientId); + assert.equal(parsed.header.messageNumber, 42); + assert.equal(parsed.header.kemCiphertext, undefined); + assert.deepEqual(parsed.aeadPayload, msg.aeadPayload); + }); + + await it("Roundtrips an init message with KEM ciphertext", () => { + const msg = { + header: { + version: 1, + flags: FLAG_INIT, + senderClientId: 1n, + recipientClientId: 2n, + messageNumber: 0, + kemCiphertext: new Uint8Array([0xde, 0xad, 0xbe, 0xef]), + }, + aeadPayload: new Uint8Array([10, 20, 30]), + }; + const bytes = serializeEncryptedMessage(msg); + const parsed = parseEncryptedMessage(bytes); + assert.equal(parsed.header.flags & FLAG_INIT, FLAG_INIT); + assert.deepEqual(parsed.header.kemCiphertext, msg.header.kemCiphertext); + }); + + await it("Roundtrip: serialize(parse(x)) === x", () => { + const msg = { + header: { + version: 1, + flags: 0, + senderClientId: 0xaaaabbbbccccddddn, + recipientClientId: 0xffff000011112222n, + messageNumber: 65535, + }, + aeadPayload: new Uint8Array(100).fill(0x42), + }; + const bytes = serializeEncryptedMessage(msg); + const parsed = parseEncryptedMessage(bytes); + const bytes2 = serializeEncryptedMessage(parsed); + assert.deepEqual(bytes, bytes2); + }); + + await it("Rejects malformed payloads", () => { + assert.throws(() => parseEncryptedMessage(new Uint8Array(0))); + assert.throws(() => parseEncryptedMessage(new Uint8Array([0x01]))); + assert.throws(() => + parseEncryptedMessage( + new Uint8Array([0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00]), + ), + ); + }); + + await it("Rejects unsupported version", () => { + const msg = { + header: { + version: 1, + flags: 0, + senderClientId: 0n, + recipientClientId: 0n, + messageNumber: 0, + }, + aeadPayload: new Uint8Array([1]), + }; + const bytes = serializeEncryptedMessage(msg); + bytes[0] = 99; + assert.throws(() => parseEncryptedMessage(bytes)); + }); + + await it("Rejects trailing data", () => { + const msg = { + header: { + version: 1, + flags: 0, + senderClientId: 0n, + recipientClientId: 0n, + messageNumber: 0, + }, + aeadPayload: new Uint8Array([1]), + }; + const bytes = concat( + serializeEncryptedMessage(msg), + new Uint8Array([0xff]), + ); + assert.throws(() => parseEncryptedMessage(bytes)); + }); +}); + +await describe("E2EE Encrypt/Decrypt", async () => { + await it("Alice encrypts and Bob decrypts successfully", async () => { + const keyring = sdk.crypto.generateKeyring(); + const bobKeys = sdk.crypto.keyringToKeys(keyring); + + // Alice encapsulates to Bob's KEM public key + const enc = sdk.crypto.encapsulate(bobKeys.kemPublicKey); + // Bob decapsulates the ciphertext + const bobSS = sdk.crypto.decapsulate(bobKeys.kemSecretKey, enc.ciphertext); + assert.deepEqual(bobSS, enc.shared_secret); + + const ss = enc.shared_secret; + const { initSessions } = setupSessions(ss); + const { aliceSession, bobSession } = await initSessions(); + + // Alice encrypts a message to Bob + const plaintext = sdk.codec.encode( + "Ping", + { Version: "hello from Alice" }, + { sender: 1n, receiver: 2n }, + ); + + const { payload, session: aliceNewSession } = await encryptPayload({ + plaintext, + session: aliceSession, + kemCiphertext: enc.ciphertext, + }); + + // Bob decrypts the message + const { plaintext: decrypted, session: bobNewSession } = + await decryptPayload({ + payload, + session: bobSession, + }); + + const frame = sdk.codec.decode(decrypted); + assert.equal(frame.type, "Ping"); + assert.equal(frame.data["Version"], "hello from Alice"); + + // Chain keys advanced correctly + assert.deepEqual(aliceNewSession.sendChainKey, bobNewSession.recvChainKey); + assert.equal(aliceNewSession.sendCount, 1); + assert.equal(bobNewSession.recvCount, 1); + assert.notDeepEqual( + aliceNewSession.sendChainKey, + aliceSession.sendChainKey, + ); + }); + + await it("Encrypt-decrypt multiple messages with chain advance", async () => { + const ss = sdk.crypto.sha256(new Uint8Array([1, 2, 3])); + const { initSessions } = setupSessions(ss); + let { aliceSession, bobSession } = await initSessions(); + + // Message 1 + const { payload: p1, session: aliceAfter1 } = await encryptPayload({ + plaintext: sdk.codec.encode("Ping", { Version: "msg1" }), + session: aliceSession, + }); + const { plaintext: d1, session: bobAfter1 } = await decryptPayload({ + payload: p1, + session: bobSession, + }); + assert.equal(sdk.codec.decode(d1).data["Version"], "msg1"); + assert.equal(bobAfter1.recvCount, 1); + + // Message 2 + const { payload: p2, session: aliceAfter2 } = await encryptPayload({ + plaintext: sdk.codec.encode("Ping", { Version: "msg2" }), + session: aliceAfter1, + }); + const { plaintext: d2, session: bobAfter2 } = await decryptPayload({ + payload: p2, + session: bobAfter1, + }); + assert.equal(sdk.codec.decode(d2).data["Version"], "msg2"); + assert.equal(bobAfter2.recvCount, 2); + + // Chain keys match after two messages + assert.deepEqual(aliceAfter2.sendChainKey, bobAfter2.recvChainKey); + assert.equal(aliceAfter2.sendCount, 2); + }); +}); + +await describe("E2EE Public Key Bundle", async () => { + await it("parses public key bundles from GetUserData.PublicKey", () => { + const keyring = sdk.crypto.generateKeyring(); + const keys = sdk.crypto.keyringToKeys(keyring); + const bundle = concat( + new Uint8Array([keys.kemPublicKey.length >> 8, keys.kemPublicKey.length & 0xff]), + keys.kemPublicKey, + new Uint8Array([keys.sigPqPublicKey.length >> 8, keys.sigPqPublicKey.length & 0xff]), + keys.sigPqPublicKey, + new Uint8Array([keys.sigClPublicKey.length >> 8, keys.sigClPublicKey.length & 0xff]), + keys.sigClPublicKey, + ); + + const parsed = sdk.crypto.publicKeyBundleToKeys(bundle); + assert.deepEqual(parsed.kemPublicKey, keys.kemPublicKey); + assert.deepEqual(parsed.sigPqPublicKey, keys.sigPqPublicKey); + assert.deepEqual(parsed.sigClPublicKey, keys.sigClPublicKey); + }); +}); + +await describe("E2EE Session Manager", async () => { + await it("getConversationId is consistent regardless of order", () => { + const id1 = getConversationId(5n, 10n); + const id2 = getConversationId(10n, 5n); + assert.equal(id1, id2); + }); + + await it("MTPSessionManager creates, retrieves, and deletes sessions", async () => { + const ss = sdk.crypto.sha256(new Uint8Array([1, 2, 3])); + const storage = new InMemorySessionStorage(); + const manager = new MTPSessionManager(storage); + + assert.equal(await manager.getSession(1n, 2n), null); + + const session = await manager.createSession({ + ownClientId: 1n, + peerClientId: 2n, + peerPublicKey: new Uint8Array(32), + sharedSecret: ss, + role: "initiator", + }); + assert.equal(session.version, 1); + assert.equal(session.sendCount, 0); + assert.equal(session.recvCount, 0); + + await manager.saveSession(session); + const retrieved = await manager.getSession(1n, 2n); + assert.notEqual(retrieved, null); + assert.equal(retrieved.conversationId, session.conversationId); + + await manager.deleteSession(1n, 2n); + assert.equal(await manager.getSession(1n, 2n), null); + }); +}); + +await describe("E2EE Full Flow: KEM + Session + Ratchet + AEAD", async () => { + await it("Alice encapsulates to Bob, both derive matching sessions, encrypt-decrypt works", async () => { + // Bob generates keyring + const bobKeyring = sdk.crypto.generateKeyring(); + const bobKeys = sdk.crypto.keyringToKeys(bobKeyring); + + // Alice encapsulates to Bob's KEM public key + const enc = sdk.crypto.encapsulate(bobKeys.kemPublicKey); + + // Bob decapsulates + const bobSharedSecret = sdk.crypto.decapsulate( + bobKeys.kemSecretKey, + enc.ciphertext, + ); + assert.deepEqual(bobSharedSecret, enc.shared_secret); + + const ss = enc.shared_secret; + const { initSessions } = setupSessions(ss); + const { aliceSession, bobSession } = await initSessions(); + + // Alice sends encrypted init message with KEM ciphertext + const msg1 = new Uint8Array([0x48, 0x65, 0x6c, 0x6c, 0x6f]); // "Hello" + const { payload: p1, session: aliceAfter1 } = await encryptPayload({ + plaintext: msg1, + session: aliceSession, + kemCiphertext: enc.ciphertext, + }); + + // Bob receives and decrypts + const { plaintext: d1, session: bobAfter1 } = await decryptPayload({ + payload: p1, + session: bobSession, + }); + assert.deepEqual(d1, msg1); + assert.deepEqual(aliceAfter1.sendChainKey, bobAfter1.recvChainKey); + assert.equal(aliceAfter1.sendCount, 1); + assert.equal(bobAfter1.recvCount, 1); + + // Second message (no KEM ciphertext) + const msg2 = new Uint8Array([0x57, 0x6f, 0x72, 0x6c, 0x64]); // "World" + const { payload: p2, session: aliceAfter2 } = await encryptPayload({ + plaintext: msg2, + session: aliceAfter1, + }); + + const { plaintext: d2, session: bobAfter2 } = await decryptPayload({ + payload: p2, + session: bobAfter1, + }); + assert.deepEqual(d2, msg2); + assert.deepEqual(aliceAfter2.sendChainKey, bobAfter2.recvChainKey); + assert.equal(aliceAfter2.sendCount, 2); + assert.equal(bobAfter2.recvCount, 2); + }); +}); + +await describe("E2EE Tamper Detection", async () => { + await it("Rejects modified ciphertext", async () => { + const ss = sdk.crypto.sha256(new Uint8Array([42])); + const { initSessions } = setupSessions(ss); + const { aliceSession, bobSession } = await initSessions(); + + const plaintext = new Uint8Array([0x01, 0x02, 0x03]); + const { payload } = await encryptPayload({ + plaintext, + session: aliceSession, + }); + + // Tamper with AEAD payload + const tampered = new Uint8Array(payload); + tampered[tampered.length - 1] ^= 0xff; + + await assert.rejects( + () => decryptPayload({ payload: tampered, session: bobSession }), + /decrypt failed/, + ); + }); + + await it("Rejects out-of-order message numbers", async () => { + const ss = sdk.crypto.sha256(new Uint8Array([7])); + const { initSessions } = setupSessions(ss); + const { aliceSession, bobSession } = await initSessions(); + + // Send two messages + const { payload: p1, session: aliceAfter1 } = await encryptPayload({ + plaintext: sdk.codec.encode("Ping", { Version: "a" }), + session: aliceSession, + }); + await encryptPayload({ + plaintext: sdk.codec.encode("Ping", { Version: "b" }), + session: aliceAfter1, + }); + + // Bob decrypts p1 + const { session: bobAfter1 } = await decryptPayload({ + payload: p1, + session: bobSession, + }); + + // Now bob expects msgNumber 1, but we try to replay msgNumber 0 + await assert.rejects( + () => decryptPayload({ payload: p1, session: bobAfter1 }), + /replay|out of order/, + ); + }); +}); diff --git a/tsconfig.json b/tsconfig.json index 5a697e9..3d94fd4 100644 --- a/tsconfig.json +++ b/tsconfig.json @@ -12,17 +12,17 @@ "ignoreDeprecations": "6.0", "paths": { "mtp/raw": ["src/raw/index.ts"], - "mtp/type-map": ["src/type-map/index.ts"] + "mtp/type-map": ["src/type-map/index.ts"], }, "strict": false, "skipLibCheck": true, "isolatedModules": true, - "verbatimModuleSyntax": true + "verbatimModuleSyntax": true, }, "include": [ "src/raw/**/*.ts", "src/sdk/**/*.ts", "src/type-map/**/*.ts", - "src/vite/**/*.ts" - ] + "src/vite/**/*.ts", + ], } diff --git a/type-map/build.rs b/type-map/build.rs index 9a0fafb..41fd144 100755 --- a/type-map/build.rs +++ b/type-map/build.rs @@ -185,10 +185,23 @@ fn main() { serde_yaml::from_str(&content).expect("Failed to parse type-maps.yaml") } Err(_) => { - eprint!("warning: MTP_TYPE_MAPS not set; generating types with reserved entries only"); - Config { - protocol_version: String::new(), - type_maps: BTreeMap::new(), + let manifest_dir = + std::path::PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); + let default_path = manifest_dir.join("../example/type-maps.yaml"); + if default_path.exists() { + println!("cargo:rerun-if-changed={}", default_path.display()); + let content = std::fs::read_to_string(&default_path) + .expect("Failed to read default example/type-maps.yaml"); + serde_yaml::from_str(&content) + .expect("Failed to parse default example/type-maps.yaml") + } else { + eprint!( + "warning: MTP_TYPE_MAPS not set; generating types with reserved entries only" + ); + Config { + protocol_version: String::new(), + type_maps: BTreeMap::new(), + } } } }; @@ -317,6 +330,7 @@ fn generate_comm_type_enum(out: &mut String, user_names: &BTreeSet<&str>) { "#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]" ) .unwrap(); + writeln!(out, "#[allow(clippy::enum_variant_names)]").unwrap(); writeln!(out, "pub enum CommunicationType {{").unwrap(); for entry in RESERVED_COMM_TYPES { @@ -401,6 +415,7 @@ fn generate_data_type_enum(out: &mut String, user_names: &BTreeSet<&str>) { "#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]" ) .unwrap(); + writeln!(out, "#[allow(clippy::enum_variant_names)]").unwrap(); writeln!(out, "pub enum DataType {{").unwrap(); for entry in RESERVED_DATA_TYPES {